1use std::collections::{HashMap, HashSet};
7use std::future::Future;
8use std::pin::Pin;
9use std::sync::{Arc, RwLock};
10use std::task::{Context, Poll};
11
12use tower_service::Service;
13
14use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
15
16use crate::async_task::{MemoryTaskStore, TaskStore, TaskStoreError};
17use crate::context::{
18 CancellationToken, ClientRequesterHandle, NotificationSender, RequestContext,
19 ServerNotification,
20};
21use crate::error::{Error, JsonRpcError, Result};
22use crate::filter::{PromptFilter, ResourceFilter, ToolFilter};
23use crate::prompt::Prompt;
24use crate::protocol::*;
25#[cfg(feature = "dynamic-tools")]
26use crate::registry::{
27 DynamicPromptRegistry, DynamicPromptsInner, DynamicResourceRegistry,
28 DynamicResourceTemplateRegistry, DynamicResourceTemplatesInner, DynamicResourcesInner,
29 DynamicToolRegistry, DynamicToolsInner,
30};
31use crate::resource::{Resource, ResourceTemplate};
32use crate::session::SessionState;
33use crate::tool::Tool;
34
35pub(crate) type CompletionHandler = Arc<
37 dyn Fn(CompleteParams) -> Pin<Box<dyn Future<Output = Result<CompleteResult>> + Send>>
38 + Send
39 + Sync,
40>;
41
42fn decode_cursor(cursor: &str) -> Result<usize> {
46 let bytes = BASE64
47 .decode(cursor)
48 .map_err(|_| Error::JsonRpc(JsonRpcError::invalid_params("Invalid pagination cursor")))?;
49 let s = String::from_utf8(bytes)
50 .map_err(|_| Error::JsonRpc(JsonRpcError::invalid_params("Invalid pagination cursor")))?;
51 s.parse::<usize>()
52 .map_err(|_| Error::JsonRpc(JsonRpcError::invalid_params("Invalid pagination cursor")))
53}
54
55fn encode_cursor(offset: usize) -> String {
57 BASE64.encode(offset.to_string())
58}
59
60fn task_store_error(e: TaskStoreError) -> Error {
62 Error::JsonRpc(JsonRpcError::internal_error(format!(
63 "Task store error: {}",
64 e
65 )))
66}
67
68#[cfg(feature = "stateless")]
73fn is_final_protocol_request(extensions: &crate::context::Extensions) -> bool {
74 extensions
75 .get::<crate::stateless::StatelessRequestMeta>()
76 .and_then(|meta| meta.protocol_version.as_deref())
77 == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
78}
79
80#[cfg(not(feature = "stateless"))]
81fn is_final_protocol_request(_extensions: &crate::context::Extensions) -> bool {
82 false
83}
84
85#[cfg(feature = "stateless")]
90fn client_declares_tasks(extensions: &crate::context::Extensions) -> bool {
91 final_client_capabilities(extensions).is_some_and(|capabilities| {
92 capabilities.extensions.as_ref().is_some_and(|declared| {
93 declared.contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
94 })
95 })
96}
97
98#[cfg(not(feature = "stateless"))]
99fn client_declares_tasks(_extensions: &crate::context::Extensions) -> bool {
100 false
101}
102
103fn decode_input_responses(
110 responses: &std::collections::HashMap<String, serde_json::Value>,
111) -> crate::protocol::InputResponses {
112 responses
113 .iter()
114 .filter_map(|(key, value)| {
115 serde_json::from_value(value.clone())
116 .ok()
117 .map(|response| (key.clone(), response))
118 })
119 .collect()
120}
121
122#[cfg(feature = "oauth")]
129fn request_principal(extensions: &crate::context::Extensions) -> Option<String> {
130 extensions
131 .get::<crate::oauth::token::TokenClaims>()
132 .and_then(|claims| claims.sub.clone())
133}
134
135#[cfg(not(feature = "oauth"))]
136fn request_principal(_extensions: &crate::context::Extensions) -> Option<String> {
137 None
138}
139
140fn unknown_task_error(task_id: &str) -> JsonRpcError {
145 JsonRpcError::invalid_params(format!("Task not found: {task_id}"))
146}
147
148pub(crate) fn tasks_client_capabilities() -> crate::protocol::ClientCapabilities {
151 crate::protocol::ClientCapabilities {
152 extensions: Some(
153 [(
154 tower_mcp_types::protocol::TASKS_EXTENSION_ID.to_string(),
155 serde_json::json!({}),
156 )]
157 .into_iter()
158 .collect(),
159 ),
160 ..Default::default()
161 }
162}
163
164#[cfg(feature = "stateless")]
165fn final_client_capabilities(
166 extensions: &crate::context::Extensions,
167) -> Option<&ClientCapabilities> {
168 extensions
169 .get::<crate::stateless::StatelessRequestMeta>()
170 .and_then(|meta| meta.client_capabilities.as_ref())
171}
172
173#[cfg(not(feature = "stateless"))]
174fn final_client_capabilities(
175 _extensions: &crate::context::Extensions,
176) -> Option<&ClientCapabilities> {
177 None
178}
179
180#[cfg(feature = "stateless")]
185fn json_value_contains(actual: &serde_json::Value, required: &serde_json::Value) -> bool {
186 match (actual, required) {
187 (serde_json::Value::Object(actual), serde_json::Value::Object(required)) => {
188 required.iter().all(|(key, value)| {
189 actual
190 .get(key)
191 .is_some_and(|a| json_value_contains(a, value))
192 })
193 }
194 _ => actual == required,
195 }
196}
197
198#[cfg(feature = "stateless")]
199fn client_capabilities_satisfy(actual: &ClientCapabilities, required: &ClientCapabilities) -> bool {
200 let actual = serde_json::to_value(actual).expect("ClientCapabilities is always serializable");
201 let mut required =
202 serde_json::to_value(required).expect("ClientCapabilities is always serializable");
203 if required.pointer("/roots/listChanged") == Some(&serde_json::Value::Bool(false))
208 && let Some(roots) = required
209 .get_mut("roots")
210 .and_then(serde_json::Value::as_object_mut)
211 {
212 roots.remove("listChanged");
213 }
214 json_value_contains(&actual, &required)
215}
216
217#[cfg(feature = "stateless")]
218fn validate_input_required_result(
219 extensions: &crate::context::Extensions,
220 result: &InputRequiredResult,
221) -> Result<()> {
222 result.validate().map_err(|message| {
223 Error::invalid_params(format!("invalid InputRequiredResult: {message}"))
224 })?;
225
226 let meta = extensions
227 .get::<crate::stateless::StatelessRequestMeta>()
228 .filter(|meta| {
229 meta.protocol_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
230 })
231 .ok_or_else(|| {
232 Error::invalid_params(
233 "InputRequiredResult is only supported by the 2026-07-28 request lifecycle",
234 )
235 })?;
236 let actual = meta.client_capabilities.as_ref().ok_or_else(|| {
237 Error::invalid_params("clientCapabilities is required for InputRequiredResult")
238 })?;
239
240 if let Some(requests) = &result.input_requests {
241 for request in requests.values() {
242 let (supported, required) = match request {
243 InputRequest::CreateMessage(params) => {
244 let requires_tools = params.tools.is_some();
245 let requires_context = params
246 .include_context
247 .is_some_and(|mode| mode != IncludeContext::None);
248 let required_sampling = SamplingCapability {
249 tools: requires_tools.then(SamplingToolsCapability::default),
250 context: requires_context.then(SamplingContextCapability::default),
251 ..SamplingCapability::default()
252 };
253 let supported = actual.sampling.as_ref().is_some_and(|sampling| {
254 (!requires_tools || sampling.tools.is_some())
255 && (!requires_context || sampling.context.is_some())
256 });
257 (
258 supported,
259 ClientCapabilities {
260 sampling: Some(required_sampling),
261 ..ClientCapabilities::default()
262 },
263 )
264 }
265 InputRequest::ListRoots(_) => (
266 actual.roots.is_some(),
267 ClientCapabilities {
268 roots: Some(RootsCapability::default()),
269 ..ClientCapabilities::default()
270 },
271 ),
272 InputRequest::Elicit(ElicitRequestParams::Form(_)) => {
273 let supported = actual.elicitation.as_ref().is_some_and(|elicitation| {
274 elicitation.form.is_some()
275 || (elicitation.form.is_none() && elicitation.url.is_none())
276 });
277 (
278 supported,
279 ClientCapabilities {
280 elicitation: Some(ElicitationCapability {
281 form: Some(ElicitationFormCapability::default()),
282 ..ElicitationCapability::default()
283 }),
284 ..ClientCapabilities::default()
285 },
286 )
287 }
288 InputRequest::Elicit(ElicitRequestParams::Url(_)) => (
289 actual
290 .elicitation
291 .as_ref()
292 .is_some_and(|elicitation| elicitation.url.is_some()),
293 ClientCapabilities {
294 elicitation: Some(ElicitationCapability {
295 url: Some(ElicitationUrlCapability::default()),
296 ..ElicitationCapability::default()
297 }),
298 ..ClientCapabilities::default()
299 },
300 ),
301 _ => {
302 return Err(Error::invalid_params(
303 "unsupported input request method in InputRequiredResult",
304 ));
305 }
306 };
307 if !supported {
308 return Err(Error::JsonRpc(
309 JsonRpcError::missing_required_client_capability(required),
310 ));
311 }
312 }
313 }
314 Ok(())
315}
316
317fn paginate<T>(
321 items: Vec<T>,
322 cursor: Option<&str>,
323 page_size: Option<usize>,
324) -> Result<(Vec<T>, Option<String>)> {
325 let Some(page_size) = page_size else {
326 return Ok((items, None));
327 };
328
329 let offset = match cursor {
330 Some(c) => decode_cursor(c)?,
331 None => 0,
332 };
333
334 if offset >= items.len() {
335 return Ok((Vec::new(), None));
336 }
337
338 let end = (offset + page_size).min(items.len());
339 let next_cursor = if end < items.len() {
340 Some(encode_cursor(end))
341 } else {
342 None
343 };
344
345 let mut items = items;
346 let page = items.drain(offset..end).collect();
347 Ok((page, next_cursor))
348}
349
350#[derive(Clone)]
374pub struct McpRouter {
375 inner: Arc<McpRouterInner>,
376 session: SessionState,
377}
378
379impl std::fmt::Debug for McpRouter {
380 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
381 f.debug_struct("McpRouter")
382 .field("server_name", &self.inner.server_name)
383 .field("server_version", &self.inner.server_version)
384 .field("tools_count", &self.inner.tools.len())
385 .field("resources_count", &self.inner.resources.len())
386 .field("prompts_count", &self.inner.prompts.len())
387 .field("session_phase", &self.session.phase())
388 .finish()
389 }
390}
391
392#[derive(Clone, Debug)]
394struct AutoInstructionsConfig {
395 prefix: Option<String>,
396 suffix: Option<String>,
397}
398
399#[cfg(all(feature = "http", feature = "stateless"))]
400type ModernNotificationSink = Arc<dyn Fn(&ServerNotification) -> bool + Send + Sync + 'static>;
401
402#[derive(Clone)]
404struct McpRouterInner {
405 server_name: String,
406 server_version: String,
407 server_title: Option<String>,
409 server_description: Option<String>,
411 server_icons: Option<Vec<ToolIcon>>,
413 server_website_url: Option<String>,
415 instructions: Option<String>,
416 auto_instructions: Option<AutoInstructionsConfig>,
417 tools: HashMap<String, Arc<Tool>>,
418 resources: HashMap<String, Arc<Resource>>,
419 resource_templates: Vec<Arc<ResourceTemplate>>,
421 prompts: HashMap<String, Arc<Prompt>>,
422 in_flight: Arc<RwLock<HashMap<RequestId, CancellationToken>>>,
424 notification_tx: Option<NotificationSender>,
426 #[cfg(all(feature = "http", feature = "stateless"))]
431 modern_notification_sink: Arc<RwLock<Option<ModernNotificationSink>>>,
432 client_requester: Option<ClientRequesterHandle>,
434 task_store: Arc<dyn TaskStore>,
436 subscriptions: Arc<RwLock<HashSet<String>>>,
438 completion_handler: Option<CompletionHandler>,
440 tool_filter: Option<ToolFilter>,
442 resource_filter: Option<ResourceFilter>,
444 prompt_filter: Option<PromptFilter>,
446 extensions: Arc<crate::context::Extensions>,
448 protocol_extensions: HashMap<String, serde_json::Value>,
450 min_log_level: Arc<RwLock<LogLevel>>,
452 page_size: Option<usize>,
454 list_ttl_ms: Option<u64>,
458 read_ttl_ms: Option<u64>,
462 cache_scope: Option<CacheScope>,
467 logging_deprecated: Option<tower_mcp_types::protocol::DeprecationInfo>,
470 disabled_tools: Arc<RwLock<HashSet<String>>>,
472 disabled_resources: Arc<RwLock<HashSet<String>>>,
474 disabled_prompts: Arc<RwLock<HashSet<String>>>,
476 #[cfg(feature = "dynamic-tools")]
478 dynamic_tools: Option<Arc<DynamicToolsInner>>,
479 #[cfg(feature = "dynamic-tools")]
481 dynamic_prompts: Option<Arc<DynamicPromptsInner>>,
482 #[cfg(feature = "dynamic-tools")]
484 dynamic_resources: Option<Arc<DynamicResourcesInner>>,
485 #[cfg(feature = "dynamic-tools")]
487 dynamic_resource_templates: Option<Arc<DynamicResourceTemplatesInner>>,
488}
489
490impl McpRouterInner {
491 fn generate_instructions(&self, config: &AutoInstructionsConfig) -> String {
493 let mut parts = Vec::new();
494
495 if let Some(prefix) = &config.prefix {
496 parts.push(prefix.clone());
497 }
498
499 if !self.tools.is_empty() {
501 let mut lines = vec!["## Tools".to_string(), String::new()];
502 let mut tools: Vec<_> = self.tools.values().collect();
503 tools.sort_by(|a, b| a.name.cmp(&b.name));
504 for tool in tools {
505 let desc = tool.description.as_deref().unwrap_or("No description");
506 let tags = annotation_tags(tool.annotations.as_ref());
507 if tags.is_empty() {
508 lines.push(format!("- **{}**: {}", tool.name, desc));
509 } else {
510 lines.push(format!("- **{}**: {} [{}]", tool.name, desc, tags));
511 }
512 }
513 parts.push(lines.join("\n"));
514 }
515
516 if !self.resources.is_empty() || !self.resource_templates.is_empty() {
518 let mut lines = vec!["## Resources".to_string(), String::new()];
519 let mut resources: Vec<_> = self.resources.values().collect();
520 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
521 for resource in resources {
522 let desc = resource.description.as_deref().unwrap_or("No description");
523 lines.push(format!("- **{}**: {}", resource.uri, desc));
524 }
525 let mut templates: Vec<_> = self.resource_templates.iter().collect();
526 templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
527 for template in templates {
528 let desc = template.description.as_deref().unwrap_or("No description");
529 lines.push(format!("- **{}**: {}", template.uri_template, desc));
530 }
531 parts.push(lines.join("\n"));
532 }
533
534 if !self.prompts.is_empty() {
536 let mut lines = vec!["## Prompts".to_string(), String::new()];
537 let mut prompts: Vec<_> = self.prompts.values().collect();
538 prompts.sort_by(|a, b| a.name.cmp(&b.name));
539 for prompt in prompts {
540 let desc = prompt.description.as_deref().unwrap_or("No description");
541 lines.push(format!("- **{}**: {}", prompt.name, desc));
542 }
543 parts.push(lines.join("\n"));
544 }
545
546 if let Some(suffix) = &config.suffix {
547 parts.push(suffix.clone());
548 }
549
550 parts.join("\n\n")
551 }
552}
553
554fn annotation_tags(annotations: Option<&crate::protocol::ToolAnnotations>) -> String {
560 let Some(ann) = annotations else {
561 return String::new();
562 };
563 let mut tags = Vec::new();
564 if ann.is_read_only() {
565 tags.push("read-only");
566 }
567 if ann.is_idempotent() {
568 tags.push("idempotent");
569 }
570 tags.join(", ")
571}
572
573impl McpRouter {
574 pub fn new() -> Self {
576 Self {
577 inner: Arc::new(McpRouterInner {
578 server_name: "tower-mcp".to_string(),
579 server_version: env!("CARGO_PKG_VERSION").to_string(),
580 server_title: None,
581 server_description: None,
582 server_icons: None,
583 server_website_url: None,
584 instructions: None,
585 auto_instructions: None,
586 tools: HashMap::new(),
587 resources: HashMap::new(),
588 resource_templates: Vec::new(),
589 prompts: HashMap::new(),
590 in_flight: Arc::new(RwLock::new(HashMap::new())),
591 notification_tx: None,
592 #[cfg(all(feature = "http", feature = "stateless"))]
593 modern_notification_sink: Arc::new(RwLock::new(None)),
594 client_requester: None,
595 task_store: Arc::new(MemoryTaskStore::new()),
596 subscriptions: Arc::new(RwLock::new(HashSet::new())),
597 extensions: Arc::new(crate::context::Extensions::new()),
598 protocol_extensions: HashMap::new(),
599 completion_handler: None,
600 tool_filter: None,
601 resource_filter: None,
602 prompt_filter: None,
603 min_log_level: Arc::new(RwLock::new(LogLevel::Debug)),
604 page_size: None,
605 list_ttl_ms: None,
606 read_ttl_ms: None,
607 cache_scope: None,
608 logging_deprecated: None,
609 disabled_tools: Arc::new(RwLock::new(HashSet::new())),
610 disabled_resources: Arc::new(RwLock::new(HashSet::new())),
611 disabled_prompts: Arc::new(RwLock::new(HashSet::new())),
612 #[cfg(feature = "dynamic-tools")]
613 dynamic_tools: None,
614 #[cfg(feature = "dynamic-tools")]
615 dynamic_prompts: None,
616 #[cfg(feature = "dynamic-tools")]
617 dynamic_resources: None,
618 #[cfg(feature = "dynamic-tools")]
619 dynamic_resource_templates: None,
620 }),
621 session: SessionState::new(),
622 }
623 }
624
625 pub fn with_fresh_session(&self) -> Self {
633 Self {
634 inner: self.inner.clone(),
635 session: SessionState::new(),
636 }
637 }
638
639 pub fn tool_annotations_map(&self) -> ToolAnnotationsMap {
649 let disabled = self.inner.disabled_tools.read().unwrap();
650 let mut map = HashMap::new();
651 for (name, tool) in &self.inner.tools {
652 if disabled.contains(name) {
653 continue;
654 }
655 if let Some(annotations) = &tool.annotations {
656 map.insert(name.clone(), annotations.clone());
657 }
658 }
659 #[cfg(feature = "dynamic-tools")]
660 if let Some(dynamic) = &self.inner.dynamic_tools {
661 for tool in dynamic.list() {
662 if disabled.contains(&tool.name) {
663 continue;
664 }
665 if !map.contains_key(&tool.name)
667 && let Some(ref annotations) = tool.annotations
668 {
669 map.insert(tool.name.clone(), annotations.clone());
670 }
671 }
672 }
673 ToolAnnotationsMap { map: Arc::new(map) }
674 }
675
676 pub fn task_store(mut self, store: Arc<dyn TaskStore>) -> Self {
694 Arc::make_mut(&mut self.inner).task_store = store;
695 self
696 }
697
698 #[cfg(feature = "dynamic-tools")]
728 pub fn with_dynamic_tools(mut self) -> (Self, DynamicToolRegistry) {
729 let inner_dyn = Arc::new(DynamicToolsInner::new());
730 Arc::make_mut(&mut self.inner).dynamic_tools = Some(inner_dyn.clone());
731 (self, DynamicToolRegistry::new(inner_dyn))
732 }
733
734 #[cfg(feature = "dynamic-tools")]
757 pub fn with_dynamic_prompts(mut self) -> (Self, DynamicPromptRegistry) {
758 let inner_dyn = Arc::new(DynamicPromptsInner::new());
759 Arc::make_mut(&mut self.inner).dynamic_prompts = Some(inner_dyn.clone());
760 (self, DynamicPromptRegistry::new(inner_dyn))
761 }
762
763 #[cfg(feature = "dynamic-tools")]
786 pub fn with_dynamic_resources(mut self) -> (Self, DynamicResourceRegistry) {
787 let inner_dyn = Arc::new(DynamicResourcesInner::new());
788 Arc::make_mut(&mut self.inner).dynamic_resources = Some(inner_dyn.clone());
789 (self, DynamicResourceRegistry::new(inner_dyn))
790 }
791
792 #[cfg(feature = "dynamic-tools")]
814 pub fn with_dynamic_resource_templates(mut self) -> (Self, DynamicResourceTemplateRegistry) {
815 let inner_dyn = Arc::new(DynamicResourceTemplatesInner::new());
816 Arc::make_mut(&mut self.inner).dynamic_resource_templates = Some(inner_dyn.clone());
817 (self, DynamicResourceTemplateRegistry::new(inner_dyn))
818 }
819
820 #[cfg(feature = "stateless")]
828 #[cfg(feature = "http")]
829 pub(crate) fn with_request_notification_sender(mut self, tx: NotificationSender) -> Self {
830 Arc::make_mut(&mut self.inner).notification_tx = Some(tx);
831 self
832 }
833
834 pub fn with_notification_sender(mut self, tx: NotificationSender) -> Self {
838 let inner = Arc::make_mut(&mut self.inner);
839 #[cfg(feature = "dynamic-tools")]
842 if let Some(ref dynamic_tools) = inner.dynamic_tools {
843 dynamic_tools.add_notification_sender(tx.clone());
844 }
845 #[cfg(feature = "dynamic-tools")]
846 if let Some(ref dynamic_prompts) = inner.dynamic_prompts {
847 dynamic_prompts.add_notification_sender(tx.clone());
848 }
849 #[cfg(feature = "dynamic-tools")]
850 if let Some(ref dynamic_resources) = inner.dynamic_resources {
851 dynamic_resources.add_notification_sender(tx.clone());
852 }
853 #[cfg(feature = "dynamic-tools")]
854 if let Some(ref dynamic_resource_templates) = inner.dynamic_resource_templates {
855 dynamic_resource_templates.add_notification_sender(tx.clone());
856 }
857 inner.notification_tx = Some(tx);
858 self
859 }
860
861 #[cfg(all(feature = "http", feature = "stateless"))]
863 pub(crate) fn attach_modern_notification_sink(&self, sink: ModernNotificationSink) {
864 if let Ok(mut active) = self.inner.modern_notification_sink.write() {
865 *active = Some(sink);
866 }
867 }
868
869 pub fn notification_sender(&self) -> Option<&NotificationSender> {
871 self.inner.notification_tx.as_ref()
872 }
873
874 pub fn with_client_requester(mut self, requester: ClientRequesterHandle) -> Self {
879 Arc::make_mut(&mut self.inner).client_requester = Some(requester);
880 self
881 }
882
883 pub fn client_requester(&self) -> Option<&ClientRequesterHandle> {
885 self.inner.client_requester.as_ref()
886 }
887
888 pub fn with_state<T: Clone + Send + Sync + 'static>(mut self, state: T) -> Self {
931 let inner = Arc::make_mut(&mut self.inner);
932 Arc::make_mut(&mut inner.extensions).insert(state);
933 self
934 }
935
936 pub fn with_extension<T: Clone + Send + Sync + 'static>(self, value: T) -> Self {
941 self.with_state(value)
942 }
943
944 pub fn with_protocol_extension(mut self, extension: crate::ExtensionDeclaration) -> Self {
951 let (identifier, settings) = extension.into_parts();
952 Arc::make_mut(&mut self.inner)
953 .protocol_extensions
954 .insert(identifier, settings);
955 self
956 }
957
958 pub fn extensions(&self) -> &crate::context::Extensions {
960 &self.inner.extensions
961 }
962
963 pub fn create_context(
968 &self,
969 request_id: RequestId,
970 progress_token: Option<ProgressToken>,
971 ) -> RequestContext {
972 self.create_context_with_extensions(request_id, progress_token, &Extensions::new())
973 }
974
975 pub(crate) fn create_context_with_extensions(
980 &self,
981 request_id: RequestId,
982 progress_token: Option<ProgressToken>,
983 per_request: &Extensions,
984 ) -> RequestContext {
985 let ctx = RequestContext::new(request_id.clone());
986
987 let ctx = if let Some(token) = progress_token {
989 ctx.with_progress_token(token)
990 } else {
991 ctx
992 };
993
994 let ctx = if let Some(tx) = &self.inner.notification_tx {
996 ctx.with_notification_sender(tx.clone())
997 } else {
998 ctx
999 };
1000
1001 let mut merged = (*self.inner.extensions).clone();
1005 merged.merge(per_request);
1006 let negotiated_extensions = if is_final_protocol_request(per_request) {
1007 let server_capabilities =
1008 self.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
1009 final_client_capabilities(per_request)
1010 .map(|client_capabilities| {
1011 crate::NegotiatedExtensions::from_capabilities(
1012 client_capabilities,
1013 &server_capabilities,
1014 )
1015 })
1016 .unwrap_or_default()
1017 } else {
1018 self.session
1019 .get::<crate::NegotiatedExtensions>()
1020 .unwrap_or_default()
1021 };
1022 merged.insert(negotiated_extensions);
1023
1024 let ctx = if !is_final_protocol_request(per_request)
1029 && let Some(requester) = merged
1030 .get::<ClientRequesterHandle>()
1031 .cloned()
1032 .or_else(|| self.inner.client_requester.clone())
1033 {
1034 ctx.with_client_requester(requester)
1035 } else {
1036 ctx
1037 };
1038
1039 let ctx = if let Some(token) = merged.get::<CancellationToken>() {
1043 ctx.with_cancellation_token(token.clone())
1044 } else {
1045 ctx
1046 };
1047
1048 let ctx = ctx.with_extensions(Arc::new(merged));
1049
1050 let ctx = ctx.with_min_log_level(self.inner.min_log_level.clone());
1052
1053 let token = ctx.cancellation_token();
1055 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1056 in_flight.insert(request_id, token);
1057 }
1058
1059 ctx
1060 }
1061
1062 pub fn complete_request(&self, request_id: &RequestId) {
1064 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1065 in_flight.remove(request_id);
1066 }
1067 }
1068
1069 fn cancel_request(&self, request_id: &RequestId) -> bool {
1071 let Ok(in_flight) = self.inner.in_flight.read() else {
1072 return false;
1073 };
1074 let Some(token) = in_flight.get(request_id) else {
1075 return false;
1076 };
1077 token.cancel();
1078 true
1079 }
1080
1081 pub fn server_info(mut self, name: impl Into<String>, version: impl Into<String>) -> Self {
1083 let inner = Arc::make_mut(&mut self.inner);
1084 inner.server_name = name.into();
1085 inner.server_version = version.into();
1086 self
1087 }
1088
1089 pub fn page_size(mut self, size: usize) -> Self {
1096 Arc::make_mut(&mut self.inner).page_size = Some(size);
1097 self
1098 }
1099
1100 pub fn list_ttl(mut self, ms: u64) -> Self {
1106 Arc::make_mut(&mut self.inner).list_ttl_ms = Some(ms);
1107 self
1108 }
1109
1110 pub fn read_ttl(mut self, ms: u64) -> Self {
1117 Arc::make_mut(&mut self.inner).read_ttl_ms = Some(ms);
1118 self
1119 }
1120
1121 pub fn cache_scope(mut self, scope: CacheScope) -> Self {
1130 Arc::make_mut(&mut self.inner).cache_scope = Some(scope);
1131 self
1132 }
1133
1134 pub fn logging_deprecated(mut self, info: tower_mcp_types::protocol::DeprecationInfo) -> Self {
1140 Arc::make_mut(&mut self.inner).logging_deprecated = Some(info);
1141 self
1142 }
1143
1144 pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
1146 Arc::make_mut(&mut self.inner).instructions = Some(instructions.into());
1147 self
1148 }
1149
1150 pub fn auto_instructions(mut self) -> Self {
1182 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1183 prefix: None,
1184 suffix: None,
1185 });
1186 self
1187 }
1188
1189 pub fn auto_instructions_with(
1206 mut self,
1207 prefix: Option<impl Into<String>>,
1208 suffix: Option<impl Into<String>>,
1209 ) -> Self {
1210 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1211 prefix: prefix.map(Into::into),
1212 suffix: suffix.map(Into::into),
1213 });
1214 self
1215 }
1216
1217 pub fn server_title(mut self, title: impl Into<String>) -> Self {
1219 Arc::make_mut(&mut self.inner).server_title = Some(title.into());
1220 self
1221 }
1222
1223 pub fn server_description(mut self, description: impl Into<String>) -> Self {
1225 Arc::make_mut(&mut self.inner).server_description = Some(description.into());
1226 self
1227 }
1228
1229 pub fn server_icons(mut self, icons: Vec<ToolIcon>) -> Self {
1231 Arc::make_mut(&mut self.inner).server_icons = Some(icons);
1232 self
1233 }
1234
1235 pub fn server_website_url(mut self, url: impl Into<String>) -> Self {
1237 Arc::make_mut(&mut self.inner).server_website_url = Some(url.into());
1238 self
1239 }
1240
1241 pub fn tool(mut self, tool: Tool) -> Self {
1243 Arc::make_mut(&mut self.inner)
1244 .tools
1245 .insert(tool.name.clone(), Arc::new(tool));
1246 self
1247 }
1248
1249 pub fn tool_if(self, condition: bool, tool: Tool) -> Self {
1275 if condition { self.tool(tool) } else { self }
1276 }
1277
1278 pub fn resource(mut self, resource: Resource) -> Self {
1280 Arc::make_mut(&mut self.inner)
1281 .resources
1282 .insert(resource.uri.clone(), Arc::new(resource));
1283 self
1284 }
1285
1286 pub fn resource_if(self, condition: bool, resource: Resource) -> Self {
1305 if condition {
1306 self.resource(resource)
1307 } else {
1308 self
1309 }
1310 }
1311
1312 pub fn resource_template(mut self, template: ResourceTemplate) -> Self {
1346 Arc::make_mut(&mut self.inner)
1347 .resource_templates
1348 .push(Arc::new(template));
1349 self
1350 }
1351
1352 pub fn prompt(mut self, prompt: Prompt) -> Self {
1354 Arc::make_mut(&mut self.inner)
1355 .prompts
1356 .insert(prompt.name.clone(), Arc::new(prompt));
1357 self
1358 }
1359
1360 pub fn prompt_if(self, condition: bool, prompt: Prompt) -> Self {
1379 if condition { self.prompt(prompt) } else { self }
1380 }
1381
1382 pub fn tools(self, tools: impl IntoIterator<Item = Tool>) -> Self {
1408 tools
1409 .into_iter()
1410 .fold(self, |router, tool| router.tool(tool))
1411 }
1412
1413 pub fn tools_if(self, condition: bool, tools: impl IntoIterator<Item = Tool>) -> Self {
1417 if condition { self.tools(tools) } else { self }
1418 }
1419
1420 pub fn resources(self, resources: impl IntoIterator<Item = Resource>) -> Self {
1439 resources
1440 .into_iter()
1441 .fold(self, |router, resource| router.resource(resource))
1442 }
1443
1444 pub fn resources_if(
1448 self,
1449 condition: bool,
1450 resources: impl IntoIterator<Item = Resource>,
1451 ) -> Self {
1452 if condition {
1453 self.resources(resources)
1454 } else {
1455 self
1456 }
1457 }
1458
1459 pub fn prompts(self, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1478 prompts
1479 .into_iter()
1480 .fold(self, |router, prompt| router.prompt(prompt))
1481 }
1482
1483 pub fn prompts_if(self, condition: bool, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1487 if condition {
1488 self.prompts(prompts)
1489 } else {
1490 self
1491 }
1492 }
1493
1494 pub fn merge(mut self, other: McpRouter) -> Self {
1539 let inner = Arc::make_mut(&mut self.inner);
1540 let other_inner = other.inner;
1541
1542 for (name, tool) in &other_inner.tools {
1544 inner.tools.insert(name.clone(), tool.clone());
1545 }
1546
1547 for (uri, resource) in &other_inner.resources {
1549 inner.resources.insert(uri.clone(), resource.clone());
1550 }
1551
1552 for template in &other_inner.resource_templates {
1555 inner.resource_templates.push(template.clone());
1556 }
1557
1558 for (name, prompt) in &other_inner.prompts {
1560 inner.prompts.insert(name.clone(), prompt.clone());
1561 }
1562
1563 for (identifier, settings) in &other_inner.protocol_extensions {
1565 inner
1566 .protocol_extensions
1567 .insert(identifier.clone(), settings.clone());
1568 }
1569
1570 self
1571 }
1572
1573 pub fn nest(mut self, prefix: impl Into<String>, other: McpRouter) -> Self {
1613 let prefix = prefix.into();
1614 let inner = Arc::make_mut(&mut self.inner);
1615 let other_inner = other.inner;
1616
1617 for tool in other_inner.tools.values() {
1619 let prefixed_tool = tool.with_name_prefix(&prefix);
1620 inner
1621 .tools
1622 .insert(prefixed_tool.name.clone(), Arc::new(prefixed_tool));
1623 }
1624
1625 for (uri, resource) in &other_inner.resources {
1627 inner.resources.insert(uri.clone(), resource.clone());
1628 }
1629
1630 for template in &other_inner.resource_templates {
1632 inner.resource_templates.push(template.clone());
1633 }
1634
1635 for (name, prompt) in &other_inner.prompts {
1637 inner.prompts.insert(name.clone(), prompt.clone());
1638 }
1639
1640 for (identifier, settings) in &other_inner.protocol_extensions {
1643 inner
1644 .protocol_extensions
1645 .insert(identifier.clone(), settings.clone());
1646 }
1647
1648 self
1649 }
1650
1651 pub fn completion_handler<F, Fut>(mut self, handler: F) -> Self
1679 where
1680 F: Fn(CompleteParams) -> Fut + Send + Sync + 'static,
1681 Fut: Future<Output = Result<CompleteResult>> + Send + 'static,
1682 {
1683 Arc::make_mut(&mut self.inner).completion_handler =
1684 Some(Arc::new(move |params| Box::pin(handler(params))));
1685 self
1686 }
1687
1688 pub fn tool_filter(mut self, filter: ToolFilter) -> Self {
1723 Arc::make_mut(&mut self.inner).tool_filter = Some(filter);
1724 self
1725 }
1726
1727 pub fn resource_filter(mut self, filter: ResourceFilter) -> Self {
1758 Arc::make_mut(&mut self.inner).resource_filter = Some(filter);
1759 self
1760 }
1761
1762 pub fn prompt_filter(mut self, filter: PromptFilter) -> Self {
1791 Arc::make_mut(&mut self.inner).prompt_filter = Some(filter);
1792 self
1793 }
1794
1795 pub fn session(&self) -> &SessionState {
1797 &self.session
1798 }
1799
1800 pub fn log(&self, params: LoggingMessageParams) -> bool {
1822 let Some(tx) = &self.inner.notification_tx else {
1823 return false;
1824 };
1825 tx.try_send(ServerNotification::LogMessage(params)).is_ok()
1826 }
1827
1828 pub fn log_info(&self, message: &str) -> bool {
1832 self.log(LoggingMessageParams::new(
1833 LogLevel::Info,
1834 serde_json::json!({ "message": message }),
1835 ))
1836 }
1837
1838 pub fn log_warning(&self, message: &str) -> bool {
1840 self.log(LoggingMessageParams::new(
1841 LogLevel::Warning,
1842 serde_json::json!({ "message": message }),
1843 ))
1844 }
1845
1846 pub fn log_error(&self, message: &str) -> bool {
1848 self.log(LoggingMessageParams::new(
1849 LogLevel::Error,
1850 serde_json::json!({ "message": message }),
1851 ))
1852 }
1853
1854 pub fn log_debug(&self, message: &str) -> bool {
1856 self.log(LoggingMessageParams::new(
1857 LogLevel::Debug,
1858 serde_json::json!({ "message": message }),
1859 ))
1860 }
1861
1862 pub fn is_subscribed(&self, uri: &str) -> bool {
1864 if let Ok(subs) = self.inner.subscriptions.read() {
1865 return subs.contains(uri);
1866 }
1867 false
1868 }
1869
1870 pub fn subscribed_uris(&self) -> Vec<String> {
1872 if let Ok(subs) = self.inner.subscriptions.read() {
1873 return subs.iter().cloned().collect();
1874 }
1875 Vec::new()
1876 }
1877
1878 fn subscribe(&self, uri: &str) -> bool {
1880 if let Ok(mut subs) = self.inner.subscriptions.write() {
1881 return subs.insert(uri.to_string());
1882 }
1883 false
1884 }
1885
1886 fn unsubscribe(&self, uri: &str) -> bool {
1888 if let Ok(mut subs) = self.inner.subscriptions.write() {
1889 return subs.remove(uri);
1890 }
1891 false
1892 }
1893
1894 pub fn notify_resource_updated(&self, uri: &str) -> bool {
1901 let notification = ServerNotification::ResourceUpdated {
1902 uri: uri.to_string(),
1903 };
1904 let mut sent = false;
1905
1906 if self.is_subscribed(uri)
1907 && let Some(tx) = &self.inner.notification_tx
1908 {
1909 sent |= tx.try_send(notification.clone()).is_ok();
1910 }
1911
1912 #[cfg(all(feature = "http", feature = "stateless"))]
1913 if let Ok(active) = self.inner.modern_notification_sink.read()
1914 && let Some(sink) = active.as_ref()
1915 {
1916 sent |= sink(¬ification);
1917 }
1918
1919 sent
1920 }
1921
1922 pub async fn notify_task_status_changed(&self, task_id: &str) {
1937 self.notify_task_state(task_id).await;
1938 }
1939
1940 pub fn notify_resources_list_changed(&self) -> bool {
1944 let Some(tx) = &self.inner.notification_tx else {
1945 return false;
1946 };
1947 tx.try_send(ServerNotification::ResourcesListChanged)
1948 .is_ok()
1949 }
1950
1951 pub fn notify_tools_list_changed(&self) -> bool {
1955 let Some(tx) = &self.inner.notification_tx else {
1956 return false;
1957 };
1958 tx.try_send(ServerNotification::ToolsListChanged).is_ok()
1959 }
1960
1961 pub fn notify_prompts_list_changed(&self) -> bool {
1965 let Some(tx) = &self.inner.notification_tx else {
1966 return false;
1967 };
1968 tx.try_send(ServerNotification::PromptsListChanged).is_ok()
1969 }
1970
1971 pub fn disable_tool(&self, name: impl Into<String>) {
1982 let mut set = self.inner.disabled_tools.write().unwrap();
1983 set.insert(name.into());
1984 }
1985
1986 pub fn enable_tool(&self, name: &str) {
1989 let mut set = self.inner.disabled_tools.write().unwrap();
1990 set.remove(name);
1991 }
1992
1993 pub fn is_tool_enabled(&self, name: &str) -> bool {
1997 !self.inner.disabled_tools.read().unwrap().contains(name)
1998 }
1999
2000 pub fn disable_resource(&self, uri: impl Into<String>) {
2003 let mut set = self.inner.disabled_resources.write().unwrap();
2004 set.insert(uri.into());
2005 }
2006
2007 pub fn enable_resource(&self, uri: &str) {
2009 let mut set = self.inner.disabled_resources.write().unwrap();
2010 set.remove(uri);
2011 }
2012
2013 pub fn is_resource_enabled(&self, uri: &str) -> bool {
2015 !self.inner.disabled_resources.read().unwrap().contains(uri)
2016 }
2017
2018 pub fn disable_prompt(&self, name: impl Into<String>) {
2021 let mut set = self.inner.disabled_prompts.write().unwrap();
2022 set.insert(name.into());
2023 }
2024
2025 pub fn enable_prompt(&self, name: &str) {
2027 let mut set = self.inner.disabled_prompts.write().unwrap();
2028 set.remove(name);
2029 }
2030
2031 pub fn is_prompt_enabled(&self, name: &str) -> bool {
2033 !self.inner.disabled_prompts.read().unwrap().contains(name)
2034 }
2035
2036 pub(crate) fn implementation(&self) -> Implementation {
2046 Implementation {
2047 name: self.inner.server_name.clone(),
2048 version: self.inner.server_version.clone(),
2049 title: self.inner.server_title.clone(),
2050 description: self.inner.server_description.clone(),
2051 icons: self.inner.server_icons.clone(),
2052 website_url: self.inner.server_website_url.clone(),
2053 meta: None,
2054 }
2055 }
2056
2057 #[cfg(feature = "http")]
2063 pub(crate) fn tool_input_schema(&self, name: &str) -> Option<serde_json::Value> {
2064 if let Some(tool) = self.inner.tools.get(name) {
2065 return Some(tool.input_schema.clone());
2066 }
2067 #[cfg(feature = "dynamic-tools")]
2068 if let Some(tool) = self
2069 .inner
2070 .dynamic_tools
2071 .as_ref()
2072 .and_then(|tools| tools.get(name))
2073 {
2074 return Some(tool.input_schema.clone());
2075 }
2076 None
2077 }
2078
2079 fn capabilities(&self) -> ServerCapabilities {
2080 let has_resources =
2081 !self.inner.resources.is_empty() || !self.inner.resource_templates.is_empty();
2082 let has_notifications = self.inner.notification_tx.is_some();
2083
2084 #[cfg(feature = "dynamic-tools")]
2085 let has_dynamic_tools = self.inner.dynamic_tools.is_some();
2086 #[cfg(not(feature = "dynamic-tools"))]
2087 let has_dynamic_tools = false;
2088
2089 #[cfg(feature = "dynamic-tools")]
2090 let has_dynamic_prompts = self.inner.dynamic_prompts.is_some();
2091 #[cfg(not(feature = "dynamic-tools"))]
2092 let has_dynamic_prompts = false;
2093
2094 #[cfg(feature = "dynamic-tools")]
2095 let has_dynamic_resources = self.inner.dynamic_resources.is_some()
2096 || self.inner.dynamic_resource_templates.is_some();
2097 #[cfg(not(feature = "dynamic-tools"))]
2098 let has_dynamic_resources = false;
2099
2100 ServerCapabilities {
2101 tools: if self.inner.tools.is_empty() && !has_dynamic_tools {
2102 None
2103 } else {
2104 Some(ToolsCapability {
2105 list_changed: has_notifications,
2106 })
2107 },
2108 resources: if has_resources || has_dynamic_resources {
2109 Some(ResourcesCapability {
2110 subscribe: true,
2111 list_changed: has_notifications,
2112 })
2113 } else {
2114 None
2115 },
2116 prompts: if self.inner.prompts.is_empty() && !has_dynamic_prompts {
2117 None
2118 } else {
2119 Some(PromptsCapability {
2120 list_changed: has_notifications,
2121 })
2122 },
2123 logging: if self.inner.notification_tx.is_some() {
2125 Some(LoggingCapability {
2126 deprecated: self.inner.logging_deprecated.clone(),
2127 })
2128 } else {
2129 None
2130 },
2131 tasks: {
2137 let has_task_support = self
2138 .inner
2139 .tools
2140 .values()
2141 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2142 if has_task_support {
2143 Some(TasksCapability {
2144 list: None,
2148 cancel: Some(TasksCancelCapability {}),
2149 requests: Some(TasksRequestsCapability {
2150 tools: Some(TasksToolsRequestsCapability {
2151 call: Some(TasksToolsCallCapability {}),
2152 }),
2153 }),
2154 })
2155 } else {
2156 None
2157 }
2158 },
2159 completions: if self.inner.completion_handler.is_some() {
2161 Some(CompletionsCapability::default())
2162 } else {
2163 None
2164 },
2165 experimental: None,
2166 extensions: {
2167 let mut map = self.inner.protocol_extensions.clone();
2168 let has_task_support = self
2169 .inner
2170 .tools
2171 .values()
2172 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2173 if has_task_support {
2174 map.insert(
2175 tower_mcp_types::protocol::TASKS_EXTENSION_ID.to_string(),
2176 serde_json::json!({}),
2177 );
2178 }
2179 (!map.is_empty()).then_some(map)
2180 },
2181 }
2182 }
2183
2184 fn capabilities_for_protocol(&self, protocol_version: Option<&str>) -> ServerCapabilities {
2192 let mut capabilities = self.capabilities();
2193 if protocol_version == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
2194 capabilities.tasks = None;
2195 if !self.final_tasks_enabled()
2196 && let Some(extensions) = capabilities.extensions.as_mut()
2197 {
2198 extensions.remove(tower_mcp_types::protocol::TASKS_EXTENSION_ID);
2199 if extensions.is_empty() {
2200 capabilities.extensions = None;
2201 }
2202 }
2203 }
2204 capabilities
2205 }
2206
2207 pub(crate) fn final_tasks_enabled(&self) -> bool {
2212 self.inner
2213 .protocol_extensions
2214 .contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
2215 }
2216
2217 fn require_negotiated_tasks(
2222 &self,
2223 extensions: &crate::context::Extensions,
2224 method: &str,
2225 ) -> Result<()> {
2226 if !self.final_tasks_enabled() {
2227 return Err(Error::JsonRpc(JsonRpcError::method_not_found(method)));
2228 }
2229 if client_declares_tasks(extensions) {
2230 return Ok(());
2231 }
2232 Err(Error::JsonRpc(
2233 JsonRpcError::missing_required_client_capability(tasks_client_capabilities()),
2234 ))
2235 }
2236
2237 async fn authorize_task(
2243 &self,
2244 task_id: &str,
2245 extensions: &crate::context::Extensions,
2246 ) -> Result<()> {
2247 let owner = self
2248 .inner
2249 .task_store
2250 .task_owner(task_id)
2251 .await
2252 .map_err(task_store_error)?
2253 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2254
2255 if crate::async_task::owner_matches(&owner, request_principal(extensions).as_deref()) {
2256 Ok(())
2257 } else {
2258 tracing::debug!(
2259 target: "mcp::tasks",
2260 task_id = %task_id,
2261 "task operation refused: principal does not own the task"
2262 );
2263 Err(Error::JsonRpc(unknown_task_error(task_id)))
2264 }
2265 }
2266
2267 async fn final_get_task(&self, task_id: &str) -> Result<McpResponse> {
2269 let detailed = self.detailed_task(task_id).await?;
2270 Ok(McpResponse::FinalGetTask(crate::tasks::GetTaskResult::new(
2271 detailed,
2272 )))
2273 }
2274
2275 async fn detailed_task(&self, task_id: &str) -> Result<crate::tasks::DetailedTask> {
2281 let (task, result, error) = self
2282 .inner
2283 .task_store
2284 .get_task_result(task_id)
2285 .await
2286 .map_err(task_store_error)?
2287 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2288
2289 let mut metadata = crate::tasks::TaskMetadata::new(
2290 task.task_id.clone(),
2291 task.created_at.clone(),
2292 task.last_updated_at.clone(),
2293 task.ttl,
2294 );
2295 metadata.status_message = task.status_message.clone();
2296 metadata.poll_interval_ms = task.poll_interval;
2297
2298 Ok(match task.status {
2299 TaskStatus::Working => crate::tasks::DetailedTask::working(metadata),
2300 TaskStatus::InputRequired => {
2301 let outstanding = self
2304 .inner
2305 .task_store
2306 .outstanding_input_requests(task_id)
2307 .await
2308 .map_err(task_store_error)?
2309 .unwrap_or_default();
2310 crate::tasks::DetailedTask::input_required(metadata, outstanding)
2311 }
2312 TaskStatus::Completed => {
2313 let mut object = result
2316 .map(serde_json::to_value)
2317 .transpose()
2318 .map_err(|e| {
2319 Error::JsonRpc(JsonRpcError::internal_error(format!(
2320 "failed to encode task result: {e}"
2321 )))
2322 })?
2323 .and_then(|value| value.as_object().cloned())
2324 .unwrap_or_default();
2325 object.insert(
2329 "resultType".to_string(),
2330 serde_json::Value::String("complete".to_string()),
2331 );
2332 crate::tasks::DetailedTask::completed(metadata, object)
2333 }
2334 TaskStatus::Failed => crate::tasks::DetailedTask::failed(
2335 metadata,
2336 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
2337 ),
2338 TaskStatus::Cancelled => crate::tasks::DetailedTask::cancelled(metadata),
2339 _ => crate::tasks::DetailedTask::working(metadata),
2342 })
2343 }
2344
2345 async fn notify_task_state(&self, task_id: &str) {
2354 if !self.final_tasks_enabled() {
2355 return;
2356 }
2357
2358 let detailed = match self.detailed_task(task_id).await {
2359 Ok(detailed) => detailed,
2360 Err(error) => {
2361 tracing::debug!(
2362 target: "mcp::tasks",
2363 task_id = %task_id,
2364 %error,
2365 "skipping task notification: task state unavailable"
2366 );
2367 return;
2368 }
2369 };
2370
2371 let notification = ServerNotification::FinalTaskStatusChanged(
2372 crate::tasks::TaskStatusNotificationParams {
2373 task: detailed,
2374 meta: None,
2375 },
2376 );
2377
2378 #[cfg(all(feature = "http", feature = "stateless"))]
2383 if let Ok(active) = self.inner.modern_notification_sink.read()
2384 && let Some(sink) = active.as_ref()
2385 {
2386 sink(¬ification);
2387 return;
2388 }
2389
2390 if let Some(tx) = &self.inner.notification_tx {
2391 let _ = tx.try_send(notification);
2392 }
2393 }
2394
2395 fn effective_cache_scope(&self, ttl_ms: Option<u64>) -> Option<CacheScope> {
2402 self.inner
2403 .cache_scope
2404 .or_else(|| ttl_ms.map(|_| CacheScope::Private))
2405 }
2406
2407 fn apply_read_cache_hints(&self, mut result: ReadResourceResult) -> ReadResourceResult {
2412 if result.ttl_ms.is_none() {
2413 result.ttl_ms = self.inner.read_ttl_ms;
2414 }
2415 if result.cache_scope.is_none() {
2416 result.cache_scope = self.effective_cache_scope(result.ttl_ms);
2417 }
2418 result
2419 }
2420
2421 async fn handle(
2423 &self,
2424 request_id: RequestId,
2425 request: McpRequest,
2426 extensions: Extensions,
2427 ) -> Result<McpResponse> {
2428 let method = request.method_name();
2430 if !is_final_protocol_request(&extensions) && !self.session.is_request_allowed(method) {
2431 tracing::warn!(
2432 method = %method,
2433 phase = ?self.session.phase(),
2434 "Request rejected: session not initialized"
2435 );
2436 return Err(Error::JsonRpc(JsonRpcError::invalid_request(format!(
2437 "Session not initialized. Only 'initialize' and 'ping' are allowed before initialization. Got: {}",
2438 method
2439 ))));
2440 }
2441
2442 match request {
2443 McpRequest::Initialize(params) => {
2444 tracing::info!(
2445 client = %params.client_info.name,
2446 version = %params.client_info.version,
2447 "Client initializing"
2448 );
2449
2450 let protocol_support = extensions.get::<crate::ProtocolSupport>();
2454 let requested_is_legacy = crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2455 .contains(¶ms.protocol_version.as_str());
2456 let requested_is_supported = requested_is_legacy
2457 && protocol_support
2458 .is_none_or(|support| support.contains(¶ms.protocol_version));
2459 let protocol_version = if requested_is_supported {
2460 params.protocol_version
2461 } else {
2462 match protocol_support {
2463 None => crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
2464 Some(support) => support
2465 .versions()
2466 .iter()
2467 .find(|version| {
2468 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2469 .contains(&version.as_str())
2470 })
2471 .cloned()
2472 .ok_or_else(|| {
2473 Error::JsonRpc(JsonRpcError::unsupported_protocol_version(
2474 params.protocol_version,
2475 support.versions().iter().map(String::as_str),
2476 ))
2477 })?,
2478 }
2479 };
2480
2481 self.session.mark_initializing();
2483 let capabilities = self.capabilities_for_protocol(Some(&protocol_version));
2484 self.session.insert(params.capabilities.clone());
2485 self.session
2486 .insert(crate::NegotiatedExtensions::from_capabilities(
2487 ¶ms.capabilities,
2488 &capabilities,
2489 ));
2490
2491 Ok(McpResponse::Initialize(InitializeResult {
2492 protocol_version,
2493 capabilities,
2494 server_info: self.implementation(),
2495 instructions: if let Some(config) = &self.inner.auto_instructions {
2496 Some(self.inner.generate_instructions(config))
2497 } else {
2498 self.inner.instructions.clone()
2499 },
2500 meta: None,
2501 }))
2502 }
2503
2504 McpRequest::Discover(_) => {
2505 tracing::debug!("Stateless server/discover request");
2512 let server_info = self.implementation();
2513 let supported_versions = extensions.get::<crate::ProtocolSupport>().map_or_else(
2514 || {
2515 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2516 .iter()
2517 .map(|version| (*version).to_string())
2518 .collect()
2519 },
2520 |support| support.versions().to_vec(),
2521 );
2522 let capabilities = self
2527 .capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
2528 Ok(McpResponse::Discover(DiscoverResult {
2529 supported_versions,
2530 capabilities,
2531 ttl_ms: None,
2532 cache_scope: None,
2533 instructions: if let Some(config) = &self.inner.auto_instructions {
2534 Some(self.inner.generate_instructions(config))
2535 } else {
2536 self.inner.instructions.clone()
2537 },
2538 meta: Some(crate::protocol::ResultMeta {
2539 server_info: Some(server_info),
2540 }),
2541 }))
2542 }
2543
2544 McpRequest::ListTools(params) => {
2545 let final_protocol = is_final_protocol_request(&extensions);
2546 let final_tasks_negotiated = final_protocol
2547 && self.final_tasks_enabled()
2548 && client_declares_tasks(&extensions);
2549 let filter = self.inner.tool_filter.as_ref();
2550 let disabled = self.inner.disabled_tools.read().unwrap().clone();
2551 let is_visible = |t: &Tool| {
2552 !disabled.contains(&t.name)
2553 && !(final_protocol
2554 && matches!(t.task_support, TaskSupportMode::Required)
2555 && !final_tasks_negotiated)
2556 && filter
2557 .map(|f| f.is_visible(&self.session, t))
2558 .unwrap_or(true)
2559 };
2560 let definition = |t: &Tool| {
2561 let mut definition = t.definition();
2562 if final_protocol {
2563 definition.execution = None;
2564 }
2565 definition
2566 };
2567
2568 let mut tools: Vec<ToolDefinition> = self
2570 .inner
2571 .tools
2572 .values()
2573 .filter(|t| is_visible(t))
2574 .map(|t| definition(t))
2575 .collect();
2576
2577 #[cfg(feature = "dynamic-tools")]
2579 if let Some(ref dynamic) = self.inner.dynamic_tools {
2580 let static_names: HashSet<String> =
2581 tools.iter().map(|t| t.name.clone()).collect();
2582 for t in dynamic.list() {
2583 if !static_names.contains(&t.name) && is_visible(&t) {
2584 tools.push(definition(&t));
2585 }
2586 }
2587 }
2588
2589 tools.sort_by(|a, b| a.name.cmp(&b.name));
2590
2591 let (tools, next_cursor) =
2592 paginate(tools, params.cursor.as_deref(), self.inner.page_size)?;
2593
2594 Ok(McpResponse::ListTools(ListToolsResult {
2595 tools,
2596 next_cursor,
2597 ttl_ms: self.inner.list_ttl_ms,
2598 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
2599 meta: None,
2600 }))
2601 }
2602
2603 McpRequest::CallTool(params) => {
2604 if self
2606 .inner
2607 .disabled_tools
2608 .read()
2609 .unwrap()
2610 .contains(¶ms.name)
2611 {
2612 tracing::info!(
2613 target: "mcp::tools",
2614 tool = %params.name,
2615 status = "disabled",
2616 "tool call completed"
2617 );
2618 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2619 }
2620
2621 let tool = self.inner.tools.get(¶ms.name).cloned();
2623 #[cfg(feature = "dynamic-tools")]
2624 let tool = tool.or_else(|| {
2625 self.inner
2626 .dynamic_tools
2627 .as_ref()
2628 .and_then(|d| d.get(¶ms.name))
2629 });
2630
2631 let tool = match tool {
2632 Some(t) => t,
2633 None => {
2634 tracing::info!(
2635 target: "mcp::tools",
2636 tool = %params.name,
2637 status = "not_found",
2638 "tool call completed"
2639 );
2640 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2641 }
2642 };
2643
2644 if let Some(filter) = &self.inner.tool_filter
2646 && !filter.is_visible(&self.session, &tool)
2647 {
2648 tracing::info!(
2649 target: "mcp::tools",
2650 tool = %params.name,
2651 status = "denied",
2652 "tool call completed"
2653 );
2654 return Err(filter.denial_error(¶ms.name));
2655 }
2656
2657 let final_protocol = is_final_protocol_request(&extensions);
2661 let task_ttl = if final_protocol {
2662 if params.task.is_some() {
2663 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
2664 "The final Tasks extension does not allow a 'task' request parameter",
2665 )));
2666 }
2667
2668 let server_enabled = self.final_tasks_enabled();
2669 let tasks_negotiated = server_enabled && client_declares_tasks(&extensions);
2670 match tool.task_support {
2671 TaskSupportMode::Required if !server_enabled => {
2672 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
2675 ¶ms.name,
2676 )));
2677 }
2678 TaskSupportMode::Required if !tasks_negotiated => {
2679 return Err(Error::JsonRpc(
2680 JsonRpcError::missing_required_client_capability(
2681 tasks_client_capabilities(),
2682 ),
2683 ));
2684 }
2685 TaskSupportMode::Required | TaskSupportMode::Optional
2686 if tasks_negotiated =>
2687 {
2688 Some(None)
2689 }
2690 _ => None,
2691 }
2692 } else {
2693 match (¶ms.task, tool.task_support) {
2694 (Some(_), TaskSupportMode::Forbidden) => {
2695 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2696 "Tool '{}' does not support async tasks",
2697 params.name
2698 ))));
2699 }
2700 (None, TaskSupportMode::Required) => {
2701 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2702 "Tool '{}' requires async task execution (include 'task' in params)",
2703 params.name
2704 ))));
2705 }
2706 (Some(task), _) => Some(task.ttl),
2707 (None, _) => None,
2708 }
2709 };
2710
2711 #[cfg(feature = "stateless")]
2715 if let Some(required) = tool.required_client_capabilities()
2716 && let Some(meta) = extensions.get::<crate::stateless::StatelessRequestMeta>()
2717 && meta.protocol_version.as_deref()
2718 == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
2719 && !meta
2720 .client_capabilities
2721 .as_ref()
2722 .is_some_and(|actual| client_capabilities_satisfy(actual, required))
2723 {
2724 return Err(Error::JsonRpc(
2725 JsonRpcError::missing_required_client_capability(required.clone()),
2726 ));
2727 }
2728
2729 if let Some(task_ttl) = task_ttl {
2730 let (task_id, cancellation_token) = self
2732 .inner
2733 .task_store
2734 .create_task(
2735 ¶ms.name,
2736 params.arguments.clone(),
2737 task_ttl,
2738 request_principal(&extensions),
2739 )
2740 .await
2741 .map_err(task_store_error)?;
2742
2743 tracing::info!(task_id = %task_id, tool = %params.name, "Created async task");
2744
2745 let progress_token = params.meta.and_then(|m| m.progress_token);
2747 let ctx = self.create_context_with_extensions(
2748 request_id,
2749 progress_token,
2750 &extensions,
2751 );
2752
2753 let task_store = self.inner.task_store.clone();
2755 let tool = tool.clone();
2756 let arguments = params.arguments;
2757 let task_id_clone = task_id.clone();
2758
2759 let tool_name = params.name.clone();
2760 let notifier = self.clone();
2761 tokio::spawn(async move {
2762 if cancellation_token.is_cancelled() {
2764 tracing::debug!(task_id = %task_id_clone, "Task cancelled before execution");
2765 notifier.notify_task_state(&task_id_clone).await;
2766 return;
2767 }
2768
2769 let start = std::time::Instant::now();
2771 let result = tool.call_with_context(ctx, arguments).await;
2772 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
2773
2774 if cancellation_token.is_cancelled() {
2775 tracing::debug!(task_id = %task_id_clone, "Task cancelled during execution");
2776 notifier.notify_task_state(&task_id_clone).await;
2777 } else {
2778 let status = if result.is_error { "error" } else { "success" };
2783 let error_msg = result
2784 .is_error
2785 .then(|| result.first_text().unwrap_or("Tool execution failed"))
2786 .map(str::to_string);
2787 if let Err(e) = task_store.complete_task(&task_id_clone, result).await {
2788 tracing::warn!(task_id = %task_id_clone, error = %e, "failed to record task completion");
2789 }
2790 tracing::info!(
2791 target: "mcp::tools",
2792 tool = %tool_name,
2793 task_id = %task_id_clone,
2794 duration_ms,
2795 status,
2796 error = error_msg.as_deref().unwrap_or_default(),
2797 "tool call completed"
2798 );
2799 notifier.notify_task_state(&task_id_clone).await;
2800 }
2801 });
2802
2803 let task = self
2804 .inner
2805 .task_store
2806 .get_task(&task_id)
2807 .await
2808 .map_err(task_store_error)?
2809 .ok_or_else(|| {
2810 Error::JsonRpc(JsonRpcError::internal_error(
2811 "Failed to retrieve created task",
2812 ))
2813 })?;
2814
2815 if is_final_protocol_request(&extensions) {
2819 let mut metadata = crate::tasks::TaskMetadata::new(
2820 task.task_id.clone(),
2821 task.created_at.clone(),
2822 task.last_updated_at.clone(),
2823 task.ttl,
2824 );
2825 metadata.status_message = task.status_message.clone();
2826 metadata.poll_interval_ms = task.poll_interval;
2827 return Ok(McpResponse::FinalCreateTask(
2828 crate::tasks::CreateTaskResult::new(crate::tasks::Task::new(
2829 metadata,
2830 task.status,
2831 )),
2832 ));
2833 }
2834 Ok(McpResponse::CreateTask(CreateTaskResult::new(task)))
2835 } else {
2836 let progress_token = params.meta.and_then(|m| m.progress_token);
2838 let ctx = self.create_context_with_extensions(
2839 request_id,
2840 progress_token,
2841 &extensions,
2842 );
2843 #[cfg(feature = "stateless")]
2844 let ctx = {
2845 let mut ctx = ctx;
2846 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
2847 params.input_responses,
2848 params.request_state,
2849 ));
2850 ctx
2851 };
2852
2853 let start = std::time::Instant::now();
2854 let outcome = tool
2855 .call_outcome_with_context(ctx, params.arguments)
2856 .await?;
2857 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
2858
2859 match outcome {
2860 RequestOutcome::Complete(result) => {
2861 let status = if result.is_error { "error" } else { "success" };
2862 tracing::info!(
2863 target: "mcp::tools",
2864 tool = %params.name,
2865 duration_ms,
2866 status,
2867 "tool call completed"
2868 );
2869 Ok(McpResponse::CallTool(result))
2870 }
2871 RequestOutcome::InputRequired(result) => {
2872 #[cfg(feature = "stateless")]
2873 {
2874 validate_input_required_result(&extensions, &result)?;
2875 tracing::info!(
2876 target: "mcp::tools",
2877 tool = %params.name,
2878 duration_ms,
2879 status = "input_required",
2880 "tool call requires client input"
2881 );
2882 Ok(McpResponse::InputRequired(result))
2883 }
2884 #[cfg(not(feature = "stateless"))]
2885 {
2886 let _ = result;
2887 Err(Error::invalid_params(
2888 "InputRequiredResult support was not compiled",
2889 ))
2890 }
2891 }
2892 }
2893 }
2894 }
2895
2896 McpRequest::ListResources(params) => {
2897 let disabled = self.inner.disabled_resources.read().unwrap().clone();
2898 let is_visible = |r: &Resource| -> bool {
2899 !disabled.contains(&r.uri)
2900 && self
2901 .inner
2902 .resource_filter
2903 .as_ref()
2904 .map(|f| f.is_visible(&self.session, r))
2905 .unwrap_or(true)
2906 };
2907
2908 let mut resources: Vec<ResourceDefinition> = self
2909 .inner
2910 .resources
2911 .values()
2912 .filter(|r| is_visible(r))
2913 .map(|r| r.definition())
2914 .collect();
2915
2916 #[cfg(feature = "dynamic-tools")]
2918 if let Some(ref dynamic) = self.inner.dynamic_resources {
2919 let static_uris: HashSet<String> =
2920 resources.iter().map(|r| r.uri.clone()).collect();
2921 for r in dynamic.list() {
2922 if !static_uris.contains(&r.uri) && is_visible(&r) {
2923 resources.push(r.definition());
2924 }
2925 }
2926 }
2927
2928 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
2929
2930 let (resources, next_cursor) =
2931 paginate(resources, params.cursor.as_deref(), self.inner.page_size)?;
2932
2933 Ok(McpResponse::ListResources(ListResourcesResult {
2934 resources,
2935 next_cursor,
2936 ttl_ms: self.inner.list_ttl_ms,
2937 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
2938 meta: None,
2939 }))
2940 }
2941
2942 McpRequest::ListResourceTemplates(params) => {
2943 let mut resource_templates: Vec<ResourceTemplateDefinition> = self
2944 .inner
2945 .resource_templates
2946 .iter()
2947 .map(|t| t.definition())
2948 .collect();
2949
2950 #[cfg(feature = "dynamic-tools")]
2952 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
2953 let static_patterns: HashSet<String> = resource_templates
2954 .iter()
2955 .map(|t| t.uri_template.clone())
2956 .collect();
2957 for t in dynamic.list() {
2958 if !static_patterns.contains(&t.uri_template) {
2959 resource_templates.push(t.definition());
2960 }
2961 }
2962 }
2963
2964 resource_templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
2965
2966 let (resource_templates, next_cursor) = paginate(
2967 resource_templates,
2968 params.cursor.as_deref(),
2969 self.inner.page_size,
2970 )?;
2971
2972 Ok(McpResponse::ListResourceTemplates(
2973 ListResourceTemplatesResult {
2974 resource_templates,
2975 next_cursor,
2976 ttl_ms: self.inner.list_ttl_ms,
2977 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
2978 meta: None,
2979 },
2980 ))
2981 }
2982
2983 McpRequest::ReadResource(params) => {
2984 if self
2986 .inner
2987 .disabled_resources
2988 .read()
2989 .unwrap()
2990 .contains(¶ms.uri)
2991 {
2992 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
2993 ¶ms.uri,
2994 )));
2995 }
2996
2997 if let Some(resource) = self.inner.resources.get(¶ms.uri) {
2999 if let Some(filter) = &self.inner.resource_filter
3001 && !filter.is_visible(&self.session, resource)
3002 {
3003 return Err(filter.denial_error(¶ms.uri));
3004 }
3005
3006 tracing::debug!(uri = %params.uri, "Reading static resource");
3007 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3008 #[cfg(feature = "stateless")]
3009 let ctx = {
3010 let mut ctx = ctx;
3011 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3012 params.input_responses.clone(),
3013 params.request_state.clone(),
3014 ));
3015 ctx
3016 };
3017 return match resource.read_outcome_with_context(ctx).await? {
3018 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3019 self.apply_read_cache_hints(result),
3020 )),
3021 RequestOutcome::InputRequired(result) => {
3022 #[cfg(feature = "stateless")]
3023 {
3024 validate_input_required_result(&extensions, &result)?;
3025 Ok(McpResponse::InputRequired(result))
3026 }
3027 #[cfg(not(feature = "stateless"))]
3028 {
3029 let _ = result;
3030 Err(Error::invalid_params(
3031 "InputRequiredResult support was not compiled",
3032 ))
3033 }
3034 }
3035 };
3036 }
3037
3038 #[cfg(feature = "dynamic-tools")]
3040 #[allow(clippy::collapsible_if)]
3041 if let Some(ref dynamic) = self.inner.dynamic_resources {
3042 if let Some(resource) = dynamic.get(¶ms.uri) {
3043 if let Some(filter) = &self.inner.resource_filter
3044 && !filter.is_visible(&self.session, &resource)
3045 {
3046 return Err(filter.denial_error(¶ms.uri));
3047 }
3048 tracing::debug!(uri = %params.uri, "Reading dynamic resource");
3049 let ctx =
3050 self.create_context_with_extensions(request_id, None, &extensions);
3051 #[cfg(feature = "stateless")]
3052 let ctx = {
3053 let mut ctx = ctx;
3054 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3055 params.input_responses.clone(),
3056 params.request_state.clone(),
3057 ));
3058 ctx
3059 };
3060 return match resource.read_outcome_with_context(ctx).await? {
3061 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3062 self.apply_read_cache_hints(result),
3063 )),
3064 RequestOutcome::InputRequired(result) => {
3065 #[cfg(feature = "stateless")]
3066 {
3067 validate_input_required_result(&extensions, &result)?;
3068 Ok(McpResponse::InputRequired(result))
3069 }
3070 #[cfg(not(feature = "stateless"))]
3071 {
3072 let _ = result;
3073 Err(Error::invalid_params(
3074 "InputRequiredResult support was not compiled",
3075 ))
3076 }
3077 }
3078 };
3079 }
3080 }
3081
3082 for template in &self.inner.resource_templates {
3084 if let Some(variables) = template.match_uri(¶ms.uri) {
3085 tracing::debug!(
3086 uri = %params.uri,
3087 template = %template.uri_template,
3088 "Reading resource via template"
3089 );
3090 let ctx =
3091 self.create_context_with_extensions(request_id, None, &extensions);
3092 #[cfg(feature = "stateless")]
3093 let ctx = {
3094 let mut ctx = ctx;
3095 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3096 params.input_responses.clone(),
3097 params.request_state.clone(),
3098 ));
3099 ctx
3100 };
3101 return match template
3102 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3103 .await?
3104 {
3105 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3106 self.apply_read_cache_hints(result),
3107 )),
3108 RequestOutcome::InputRequired(result) => {
3109 #[cfg(feature = "stateless")]
3110 {
3111 validate_input_required_result(&extensions, &result)?;
3112 Ok(McpResponse::InputRequired(result))
3113 }
3114 #[cfg(not(feature = "stateless"))]
3115 {
3116 let _ = result;
3117 Err(Error::invalid_params(
3118 "InputRequiredResult support was not compiled",
3119 ))
3120 }
3121 }
3122 };
3123 }
3124 }
3125
3126 #[cfg(feature = "dynamic-tools")]
3128 #[allow(clippy::collapsible_if)]
3129 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3130 if let Some((template, variables)) = dynamic.match_uri(¶ms.uri) {
3131 tracing::debug!(
3132 uri = %params.uri,
3133 template = %template.uri_template,
3134 "Reading resource via dynamic template"
3135 );
3136 let ctx =
3137 self.create_context_with_extensions(request_id, None, &extensions);
3138 #[cfg(feature = "stateless")]
3139 let ctx = {
3140 let mut ctx = ctx;
3141 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3142 params.input_responses.clone(),
3143 params.request_state.clone(),
3144 ));
3145 ctx
3146 };
3147 return match template
3148 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3149 .await?
3150 {
3151 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3152 self.apply_read_cache_hints(result),
3153 )),
3154 RequestOutcome::InputRequired(result) => {
3155 #[cfg(feature = "stateless")]
3156 {
3157 validate_input_required_result(&extensions, &result)?;
3158 Ok(McpResponse::InputRequired(result))
3159 }
3160 #[cfg(not(feature = "stateless"))]
3161 {
3162 let _ = result;
3163 Err(Error::invalid_params(
3164 "InputRequiredResult support was not compiled",
3165 ))
3166 }
3167 }
3168 };
3169 }
3170 }
3171
3172 Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3174 ¶ms.uri,
3175 )))
3176 }
3177
3178 McpRequest::SubscribeResource(params) => {
3179 if !self.inner.resources.contains_key(¶ms.uri) {
3181 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3182 ¶ms.uri,
3183 )));
3184 }
3185
3186 tracing::debug!(uri = %params.uri, "Subscribing to resource");
3187 self.subscribe(¶ms.uri);
3188
3189 Ok(McpResponse::SubscribeResource(EmptyResult {}))
3190 }
3191
3192 McpRequest::UnsubscribeResource(params) => {
3193 if !self.inner.resources.contains_key(¶ms.uri) {
3195 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3196 ¶ms.uri,
3197 )));
3198 }
3199
3200 tracing::debug!(uri = %params.uri, "Unsubscribing from resource");
3201 self.unsubscribe(¶ms.uri);
3202
3203 Ok(McpResponse::UnsubscribeResource(EmptyResult {}))
3204 }
3205
3206 McpRequest::ListPrompts(params) => {
3207 let disabled = self.inner.disabled_prompts.read().unwrap().clone();
3208 let is_visible = |p: &Prompt| -> bool {
3209 !disabled.contains(&p.name)
3210 && self
3211 .inner
3212 .prompt_filter
3213 .as_ref()
3214 .map(|f| f.is_visible(&self.session, p))
3215 .unwrap_or(true)
3216 };
3217
3218 let mut prompts: Vec<PromptDefinition> = self
3219 .inner
3220 .prompts
3221 .values()
3222 .filter(|p| is_visible(p))
3223 .map(|p| p.definition())
3224 .collect();
3225
3226 #[cfg(feature = "dynamic-tools")]
3228 if let Some(ref dynamic) = self.inner.dynamic_prompts {
3229 let static_names: HashSet<String> =
3230 prompts.iter().map(|p| p.name.clone()).collect();
3231 for p in dynamic.list() {
3232 if !static_names.contains(&p.name) && is_visible(&p) {
3233 prompts.push(p.definition());
3234 }
3235 }
3236 }
3237
3238 prompts.sort_by(|a, b| a.name.cmp(&b.name));
3239
3240 let (prompts, next_cursor) =
3241 paginate(prompts, params.cursor.as_deref(), self.inner.page_size)?;
3242
3243 Ok(McpResponse::ListPrompts(ListPromptsResult {
3244 prompts,
3245 next_cursor,
3246 ttl_ms: self.inner.list_ttl_ms,
3247 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3248 meta: None,
3249 }))
3250 }
3251
3252 McpRequest::GetPrompt(params) => {
3253 if self
3255 .inner
3256 .disabled_prompts
3257 .read()
3258 .unwrap()
3259 .contains(¶ms.name)
3260 {
3261 return Err(Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3262 "Prompt not found: {}",
3263 params.name
3264 ))));
3265 }
3266
3267 let prompt = self.inner.prompts.get(¶ms.name).cloned();
3269 #[cfg(feature = "dynamic-tools")]
3270 let prompt = prompt.or_else(|| {
3271 self.inner
3272 .dynamic_prompts
3273 .as_ref()
3274 .and_then(|d| d.get(¶ms.name))
3275 });
3276 let prompt = prompt.ok_or_else(|| {
3277 Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3278 "Prompt not found: {}",
3279 params.name
3280 )))
3281 })?;
3282
3283 if let Some(filter) = &self.inner.prompt_filter
3285 && !filter.is_visible(&self.session, &prompt)
3286 {
3287 return Err(filter.denial_error(¶ms.name));
3288 }
3289
3290 tracing::debug!(name = %params.name, "Getting prompt");
3291 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3292 #[cfg(feature = "stateless")]
3293 let ctx = {
3294 let mut ctx = ctx;
3295 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3296 params.input_responses,
3297 params.request_state,
3298 ));
3299 ctx
3300 };
3301 let outcome = prompt
3302 .get_outcome_with_context(ctx, params.arguments)
3303 .await?;
3304
3305 match outcome {
3306 RequestOutcome::Complete(result) => Ok(McpResponse::GetPrompt(result)),
3307 RequestOutcome::InputRequired(result) => {
3308 #[cfg(feature = "stateless")]
3309 {
3310 validate_input_required_result(&extensions, &result)?;
3311 Ok(McpResponse::InputRequired(result))
3312 }
3313 #[cfg(not(feature = "stateless"))]
3314 {
3315 let _ = result;
3316 Err(Error::invalid_params(
3317 "InputRequiredResult support was not compiled",
3318 ))
3319 }
3320 }
3321 }
3322 }
3323
3324 McpRequest::Ping => Ok(McpResponse::Pong(EmptyResult {})),
3325
3326 McpRequest::GetTaskInfo(params) => {
3327 if is_final_protocol_request(&extensions) {
3328 self.require_negotiated_tasks(&extensions, "tasks/get")?;
3329 self.authorize_task(¶ms.task_id, &extensions).await?;
3330 return self.final_get_task(¶ms.task_id).await;
3331 }
3332 self.authorize_task(¶ms.task_id, &extensions).await?;
3333
3334 let (mut task, result, error) = self
3341 .inner
3342 .task_store
3343 .get_task_result(¶ms.task_id)
3344 .await
3345 .map_err(task_store_error)?
3346 .ok_or_else(|| {
3347 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3348 "Task not found: {}",
3349 params.task_id
3350 )))
3351 })?;
3352
3353 match task.status {
3354 TaskStatus::Completed => task.result = result,
3355 TaskStatus::Failed => {
3356 task.error = Some(
3360 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
3361 );
3362 }
3363 _ => {}
3364 }
3365
3366 Ok(McpResponse::GetTaskInfo(task))
3367 }
3368
3369 McpRequest::UpdateTask(params) => {
3370 if is_final_protocol_request(&extensions) {
3371 self.require_negotiated_tasks(&extensions, "tasks/update")?;
3372 self.authorize_task(¶ms.task_id, &extensions).await?;
3373 self.inner
3377 .task_store
3378 .apply_input_responses(
3379 ¶ms.task_id,
3380 decode_input_responses(¶ms.input_responses),
3381 )
3382 .await
3383 .map_err(task_store_error)?
3384 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3385 self.notify_task_state(¶ms.task_id).await;
3389 return Ok(McpResponse::FinalTaskAck(
3390 crate::tasks::TaskAcknowledgement::new(),
3391 ));
3392 }
3393
3394 self.authorize_task(¶ms.task_id, &extensions).await?;
3395
3396 let _ = self
3404 .inner
3405 .task_store
3406 .get_task(¶ms.task_id)
3407 .await
3408 .map_err(task_store_error)?
3409 .ok_or_else(|| {
3410 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3411 "Task not found: {}",
3412 params.task_id
3413 )))
3414 })?;
3415 Ok(McpResponse::UpdateTask(EmptyResult {}))
3416 }
3417
3418 McpRequest::CancelTask(params) => {
3419 if is_final_protocol_request(&extensions) {
3420 self.require_negotiated_tasks(&extensions, "tasks/cancel")?;
3421 self.authorize_task(¶ms.task_id, &extensions).await?;
3422 self.inner
3426 .task_store
3427 .cancel_task(¶ms.task_id, params.reason.as_deref())
3428 .await
3429 .map_err(task_store_error)?
3430 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3431 self.notify_task_state(¶ms.task_id).await;
3432 return Ok(McpResponse::FinalTaskAck(
3433 crate::tasks::TaskAcknowledgement::new(),
3434 ));
3435 }
3436
3437 self.authorize_task(¶ms.task_id, &extensions).await?;
3438
3439 let current = self
3441 .inner
3442 .task_store
3443 .get_task(¶ms.task_id)
3444 .await
3445 .map_err(task_store_error)?
3446 .ok_or_else(|| {
3447 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3448 "Task not found: {}",
3449 params.task_id
3450 )))
3451 })?;
3452
3453 if current.status.is_terminal() {
3454 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
3455 "Task {} is already in terminal state: {}",
3456 params.task_id, current.status
3457 ))));
3458 }
3459
3460 self.inner
3461 .task_store
3462 .cancel_task(¶ms.task_id, params.reason.as_deref())
3463 .await
3464 .map_err(task_store_error)?
3465 .ok_or_else(|| {
3466 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3467 "Task not found: {}",
3468 params.task_id
3469 )))
3470 })?;
3471
3472 Ok(McpResponse::CancelTask(EmptyResult {}))
3476 }
3477
3478 McpRequest::SetLoggingLevel(params) => {
3479 tracing::debug!(level = ?params.level, "Client set logging level");
3480 if let Ok(mut level) = self.inner.min_log_level.write() {
3481 *level = params.level;
3482 }
3483 Ok(McpResponse::SetLoggingLevel(EmptyResult {}))
3484 }
3485
3486 McpRequest::Complete(params) => {
3487 tracing::debug!(
3488 reference = ?params.reference,
3489 argument = %params.argument.name,
3490 "Completion request"
3491 );
3492
3493 if let Some(ref handler) = self.inner.completion_handler {
3495 let result = handler(params).await?;
3496 Ok(McpResponse::Complete(result))
3497 } else {
3498 Ok(McpResponse::Complete(CompleteResult::new(vec![])))
3500 }
3501 }
3502
3503 McpRequest::Unknown { method, .. } => {
3504 Err(Error::JsonRpc(JsonRpcError::method_not_found(&method)))
3505 }
3506 _ => Err(Error::JsonRpc(JsonRpcError::method_not_found(
3507 "unknown method",
3508 ))),
3509 }
3510 }
3511
3512 pub fn handle_notification(&self, notification: McpNotification) {
3514 match notification {
3515 McpNotification::Initialized => {
3516 let phase_before = self.session.phase();
3517 if self.session.mark_initialized() {
3518 if phase_before == crate::session::SessionPhase::Uninitialized {
3519 tracing::info!(
3520 "Session initialized from uninitialized state (race resolved)"
3521 );
3522 } else {
3523 tracing::info!("Session initialized, entering operation phase");
3524 }
3525 } else {
3526 tracing::warn!(
3527 phase = ?self.session.phase(),
3528 "Received initialized notification in unexpected state"
3529 );
3530 }
3531 }
3532 McpNotification::Cancelled(params) => {
3533 if let Some(ref request_id) = params.request_id {
3534 if self.cancel_request(request_id) {
3535 tracing::info!(
3536 request_id = ?request_id,
3537 reason = ?params.reason,
3538 "Request cancelled"
3539 );
3540 } else {
3541 tracing::debug!(
3542 request_id = ?request_id,
3543 reason = ?params.reason,
3544 "Cancellation requested for unknown request"
3545 );
3546 }
3547 } else {
3548 tracing::debug!(
3549 reason = ?params.reason,
3550 "Cancellation notification received without request_id"
3551 );
3552 }
3553 }
3554 McpNotification::Progress(params) => {
3555 tracing::trace!(
3556 token = ?params.progress_token,
3557 progress = params.progress,
3558 total = ?params.total,
3559 "Progress notification"
3560 );
3561 }
3569 McpNotification::RootsListChanged => {
3570 tracing::info!("Client roots list changed");
3571 }
3574 McpNotification::Unknown { method, .. } => {
3575 tracing::debug!(method = %method, "Unknown notification received");
3576 }
3577 _ => {
3578 tracing::debug!("Unrecognized notification variant received");
3579 }
3580 }
3581 }
3582}
3583
3584impl Default for McpRouter {
3585 fn default() -> Self {
3586 Self::new()
3587 }
3588}
3589
3590pub use crate::context::Extensions;
3596
3597#[derive(Debug, Clone)]
3622pub struct ToolAnnotationsMap {
3623 map: Arc<HashMap<String, ToolAnnotations>>,
3624}
3625
3626impl ToolAnnotationsMap {
3627 pub fn get(&self, tool_name: &str) -> Option<&ToolAnnotations> {
3631 self.map.get(tool_name)
3632 }
3633
3634 pub fn is_read_only(&self, tool_name: &str) -> bool {
3639 self.map.get(tool_name).is_some_and(|a| a.read_only_hint)
3640 }
3641
3642 pub fn is_destructive(&self, tool_name: &str) -> bool {
3647 self.map.get(tool_name).is_none_or(|a| a.destructive_hint)
3648 }
3649
3650 pub fn is_idempotent(&self, tool_name: &str) -> bool {
3655 self.map.get(tool_name).is_some_and(|a| a.idempotent_hint)
3656 }
3657}
3658
3659#[derive(Debug, Clone)]
3681pub struct RouterRequest {
3682 pub id: RequestId,
3684 pub inner: McpRequest,
3686 pub extensions: Extensions,
3688}
3689
3690impl RouterRequest {
3691 pub fn new(id: RequestId, inner: McpRequest) -> Self {
3693 Self {
3694 id,
3695 inner,
3696 extensions: Extensions::new(),
3697 }
3698 }
3699
3700 pub fn with_inner(self, inner: McpRequest) -> Self {
3706 Self {
3707 id: self.id,
3708 inner,
3709 extensions: self.extensions,
3710 }
3711 }
3712
3713 pub fn with_id_and_inner(self, id: RequestId, inner: McpRequest) -> Self {
3719 Self {
3720 id,
3721 inner,
3722 extensions: self.extensions,
3723 }
3724 }
3725
3726 pub fn clone_with_inner(&self, inner: McpRequest) -> Self {
3734 Self {
3735 id: self.id.clone(),
3736 inner,
3737 extensions: self.extensions.clone(),
3738 }
3739 }
3740}
3741
3742#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
3744pub struct RouterResponse {
3745 pub id: RequestId,
3747 pub inner: std::result::Result<McpResponse, JsonRpcError>,
3749}
3750
3751impl RouterResponse {
3752 pub fn is_error(&self) -> bool {
3768 self.inner.is_err()
3769 }
3770
3771 pub fn into_jsonrpc(self) -> JsonRpcResponse {
3773 match self.inner {
3774 Ok(response) => match serde_json::to_value(response) {
3775 Ok(result) => JsonRpcResponse::result(self.id, result),
3776 Err(e) => {
3777 tracing::error!(error = %e, "Failed to serialize response");
3778 JsonRpcResponse::error(
3779 Some(self.id),
3780 JsonRpcError::internal_error(format!("Serialization error: {}", e)),
3781 )
3782 }
3783 },
3784 Err(error) => JsonRpcResponse::error(Some(self.id), error),
3785 }
3786 }
3787}
3788
3789impl Service<RouterRequest> for McpRouter {
3790 type Response = RouterResponse;
3791 type Error = std::convert::Infallible; type Future =
3793 Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
3794
3795 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
3796 Poll::Ready(Ok(()))
3797 }
3798
3799 fn call(&mut self, req: RouterRequest) -> Self::Future {
3800 let router = self.clone();
3801 let request_id = req.id.clone();
3802 Box::pin(async move {
3803 let result = router.handle(req.id, req.inner, req.extensions).await;
3804 router.complete_request(&request_id);
3806 Ok(RouterResponse {
3807 id: request_id,
3808 inner: result.map_err(|e| match e {
3813 Error::JsonRpc(err) => err,
3814 Error::Tool(err) => JsonRpcError::internal_error(err.to_string()),
3815 e => JsonRpcError::internal_error(e.to_string()),
3816 }),
3817 })
3818 })
3819 }
3820}
3821
3822#[cfg(test)]
3823mod tests {
3824 use super::*;
3825 use crate::extract::{Context, Json};
3826 use crate::jsonrpc::JsonRpcService;
3827 use crate::tool::ToolBuilder;
3828 use schemars::JsonSchema;
3829 use serde::Deserialize;
3830 use tower::ServiceExt;
3831
3832 #[derive(Debug, Deserialize, JsonSchema)]
3833 struct AddInput {
3834 a: i64,
3835 b: i64,
3836 }
3837
3838 #[cfg(feature = "stateless")]
3839 fn final_extensions(client_capabilities: ClientCapabilities) -> Extensions {
3840 let mut extensions = Extensions::new();
3841 extensions.insert(crate::stateless::StatelessRequestMeta {
3842 protocol_version: Some(PROTOCOL_VERSION_2026_07_28.to_string()),
3843 client_capabilities: Some(client_capabilities),
3844 ..Default::default()
3845 });
3846 extensions
3847 }
3848
3849 #[cfg(feature = "stateless")]
3850 fn tasks_client_extensions() -> Extensions {
3851 final_extensions(ClientCapabilities {
3852 extensions: Some(
3853 [(TASKS_EXTENSION_ID.to_string(), serde_json::json!({}))]
3854 .into_iter()
3855 .collect(),
3856 ),
3857 ..Default::default()
3858 })
3859 }
3860
3861 #[cfg(feature = "stateless")]
3862 #[tokio::test]
3863 async fn final_tasks_require_server_opt_in_and_client_declaration() {
3864 let tool = || {
3865 ToolBuilder::new("optional_task")
3866 .task_support(TaskSupportMode::Optional)
3867 .handler(|input: AddInput| async move {
3868 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
3869 })
3870 .build()
3871 };
3872 let task_params = |task| CallToolParams {
3873 name: "optional_task".to_string(),
3874 arguments: serde_json::json!({"a": 1, "b": 2}),
3875 input_responses: None,
3876 request_state: None,
3877 meta: None,
3878 task,
3879 };
3880
3881 let implicit = McpRouter::new().tool(tool());
3885 let McpResponse::Discover(result) = implicit
3886 .handle(
3887 RequestId::Number(1),
3888 McpRequest::Discover(DiscoverParams::default()),
3889 Extensions::new(),
3890 )
3891 .await
3892 .unwrap()
3893 else {
3894 panic!("Expected Discover response");
3895 };
3896 assert!(
3897 result
3898 .capabilities
3899 .extensions
3900 .as_ref()
3901 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
3902 );
3903 let error = implicit
3904 .handle(
3905 RequestId::Number(2),
3906 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
3907 tasks_client_extensions(),
3908 )
3909 .await
3910 .unwrap_err();
3911 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
3912
3913 let router = McpRouter::new().tool(tool()).with_tasks();
3915 let McpResponse::Discover(result) = router
3916 .handle(
3917 RequestId::Number(3),
3918 McpRequest::Discover(DiscoverParams::default()),
3919 Extensions::new(),
3920 )
3921 .await
3922 .unwrap()
3923 else {
3924 panic!("Expected Discover response");
3925 };
3926 assert!(
3927 result
3928 .capabilities
3929 .extensions
3930 .as_ref()
3931 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID)),
3932 "with_tasks() must advertise the extension on the final path"
3933 );
3934 assert!(
3935 result.capabilities.tasks.is_none(),
3936 "the legacy capability shape is never advertised on the final path"
3937 );
3938
3939 let response = router
3942 .handle(
3943 RequestId::Number(4),
3944 McpRequest::CallTool(task_params(None)),
3945 final_extensions(ClientCapabilities::default()),
3946 )
3947 .await
3948 .unwrap();
3949 assert!(matches!(response, McpResponse::CallTool(_)));
3950
3951 let response = router
3954 .handle(
3955 RequestId::Number(5),
3956 McpRequest::CallTool(task_params(None)),
3957 tasks_client_extensions(),
3958 )
3959 .await
3960 .unwrap();
3961 assert!(
3962 matches!(response, McpResponse::FinalCreateTask(_)),
3963 "a negotiated request must receive a task, got {response:?}"
3964 );
3965
3966 let error = router
3969 .handle(
3970 RequestId::Number(6),
3971 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
3972 tasks_client_extensions(),
3973 )
3974 .await
3975 .unwrap_err();
3976 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
3977 }
3978
3979 #[cfg(feature = "stateless")]
3980 #[tokio::test]
3981 async fn final_task_methods_serve_the_negotiated_wire_shapes() {
3982 let router = McpRouter::new()
3983 .tool(
3984 ToolBuilder::new("optional_task")
3985 .task_support(TaskSupportMode::Optional)
3986 .handler(|input: AddInput| async move {
3987 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
3988 })
3989 .build(),
3990 )
3991 .with_tasks();
3992
3993 let McpResponse::FinalCreateTask(created) = router
3994 .handle(
3995 RequestId::Number(1),
3996 McpRequest::CallTool(CallToolParams {
3997 name: "optional_task".to_string(),
3998 arguments: serde_json::json!({"a": 1, "b": 2}),
3999 input_responses: None,
4000 request_state: None,
4001 meta: None,
4002 task: None,
4003 }),
4004 tasks_client_extensions(),
4005 )
4006 .await
4007 .unwrap()
4008 else {
4009 panic!("Expected a final create-task response");
4010 };
4011
4012 let wire = serde_json::to_value(&created).unwrap();
4014 assert_eq!(wire["resultType"], "task");
4015 assert!(wire.get("task").is_none(), "final results are flat: {wire}");
4016 assert!(wire["ttlMs"].is_number() || wire["ttlMs"].is_null());
4017 assert!(wire.get("ttl").is_none(), "legacy field name leaked");
4018 let task_id = created.task.metadata.task_id.clone();
4019
4020 let McpResponse::FinalGetTask(fetched) = router
4022 .handle(
4023 RequestId::Number(2),
4024 McpRequest::GetTaskInfo(GetTaskInfoParams {
4025 task_id: task_id.clone(),
4026 meta: None,
4027 }),
4028 tasks_client_extensions(),
4029 )
4030 .await
4031 .unwrap()
4032 else {
4033 panic!("Expected a final get-task response");
4034 };
4035 let wire = serde_json::to_value(&fetched).unwrap();
4036 assert_eq!(wire["resultType"], "complete");
4037 assert_eq!(wire["taskId"], serde_json::json!(task_id));
4038 assert!(wire["status"].is_string());
4039
4040 for (id, request) in [
4042 (
4043 3,
4044 McpRequest::UpdateTask(UpdateTaskParams {
4045 task_id: task_id.clone(),
4046 input_responses: HashMap::new(),
4047 meta: None,
4048 }),
4049 ),
4050 (
4051 4,
4052 McpRequest::CancelTask(CancelTaskParams {
4053 task_id: task_id.clone(),
4054 reason: None,
4055 meta: None,
4056 }),
4057 ),
4058 ] {
4059 let response = router
4060 .handle(RequestId::Number(id), request, tasks_client_extensions())
4061 .await
4062 .unwrap();
4063 let McpResponse::FinalTaskAck(ack) = response else {
4064 panic!("Expected a final ack for request {id}");
4065 };
4066 assert_eq!(
4067 serde_json::to_value(&ack).unwrap(),
4068 serde_json::json!({"resultType": "complete"})
4069 );
4070 }
4071
4072 let error = router
4074 .handle(
4075 RequestId::Number(5),
4076 McpRequest::GetTaskInfo(GetTaskInfoParams {
4077 task_id: "does-not-exist".to_string(),
4078 meta: None,
4079 }),
4080 tasks_client_extensions(),
4081 )
4082 .await
4083 .unwrap_err();
4084 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4085
4086 let error = router
4088 .handle(
4089 RequestId::Number(6),
4090 McpRequest::GetTaskInfo(GetTaskInfoParams {
4091 task_id: task_id.clone(),
4092 meta: None,
4093 }),
4094 final_extensions(ClientCapabilities::default()),
4095 )
4096 .await
4097 .unwrap_err();
4098 let Error::JsonRpc(error) = error else {
4099 panic!("expected a JSON-RPC error");
4100 };
4101 assert_eq!(error.code, -32021);
4102 assert_eq!(
4103 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4104 serde_json::json!({})
4105 );
4106 }
4107
4108 #[cfg(feature = "stateless")]
4109 #[tokio::test]
4110 async fn final_required_task_tools_follow_per_request_capabilities() {
4111 let router = McpRouter::new()
4112 .tool(
4113 ToolBuilder::new("required_task")
4114 .task_support(TaskSupportMode::Required)
4115 .handler(|input: AddInput| async move {
4116 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4117 })
4118 .build(),
4119 )
4120 .with_tasks();
4121 let params = || CallToolParams {
4122 name: "required_task".to_string(),
4123 arguments: serde_json::json!({"a": 1, "b": 2}),
4124 input_responses: None,
4125 request_state: None,
4126 meta: None,
4127 task: None,
4128 };
4129
4130 let McpResponse::ListTools(without_tasks) = router
4131 .handle(
4132 RequestId::Number(1),
4133 McpRequest::ListTools(ListToolsParams::default()),
4134 final_extensions(ClientCapabilities::default()),
4135 )
4136 .await
4137 .unwrap()
4138 else {
4139 panic!("expected tools/list")
4140 };
4141 assert!(without_tasks.tools.is_empty());
4142
4143 let McpResponse::ListTools(with_tasks) = router
4144 .handle(
4145 RequestId::Number(2),
4146 McpRequest::ListTools(ListToolsParams::default()),
4147 tasks_client_extensions(),
4148 )
4149 .await
4150 .unwrap()
4151 else {
4152 panic!("expected tools/list")
4153 };
4154 assert_eq!(with_tasks.tools.len(), 1);
4155 assert!(with_tasks.tools[0].execution.is_none());
4156
4157 let error = router
4158 .handle(
4159 RequestId::Number(3),
4160 McpRequest::CallTool(params()),
4161 final_extensions(ClientCapabilities::default()),
4162 )
4163 .await
4164 .unwrap_err();
4165 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32021));
4166
4167 let response = router
4168 .handle(
4169 RequestId::Number(4),
4170 McpRequest::CallTool(params()),
4171 tasks_client_extensions(),
4172 )
4173 .await
4174 .unwrap();
4175 assert!(matches!(response, McpResponse::FinalCreateTask(_)));
4176 }
4177
4178 #[cfg(all(feature = "oauth", feature = "stateless"))]
4179 #[tokio::test]
4180 async fn task_operations_are_bound_to_the_creating_principal() {
4181 fn as_principal(subject: &str) -> Extensions {
4182 let mut extensions = tasks_client_extensions();
4183 extensions.insert(crate::oauth::token::TokenClaims {
4184 sub: Some(subject.to_string()),
4185 iss: None,
4186 aud: None,
4187 exp: None,
4188 scope: None,
4189 client_id: None,
4190 extra: HashMap::new(),
4191 });
4192 extensions
4193 }
4194
4195 let router = McpRouter::new()
4196 .tool(
4197 ToolBuilder::new("optional_task")
4198 .task_support(TaskSupportMode::Optional)
4199 .handler(|input: AddInput| async move {
4200 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4201 })
4202 .build(),
4203 )
4204 .with_tasks();
4205
4206 let McpResponse::FinalCreateTask(created) = router
4207 .handle(
4208 RequestId::Number(1),
4209 McpRequest::CallTool(CallToolParams {
4210 name: "optional_task".to_string(),
4211 arguments: serde_json::json!({"a": 1, "b": 2}),
4212 input_responses: None,
4213 request_state: None,
4214 meta: None,
4215 task: None,
4216 }),
4217 as_principal("alice"),
4218 )
4219 .await
4220 .unwrap()
4221 else {
4222 panic!("Expected a final create-task response");
4223 };
4224 let task_id = created.task.metadata.task_id.clone();
4225
4226 assert!(
4228 router
4229 .handle(
4230 RequestId::Number(2),
4231 McpRequest::GetTaskInfo(GetTaskInfoParams {
4232 task_id: task_id.clone(),
4233 meta: None,
4234 }),
4235 as_principal("alice"),
4236 )
4237 .await
4238 .is_ok()
4239 );
4240
4241 for (id, label, context) in [
4244 (3, "another principal", as_principal("bob")),
4245 (4, "no principal", tasks_client_extensions()),
4246 ] {
4247 for (offset, request) in [
4248 McpRequest::GetTaskInfo(GetTaskInfoParams {
4249 task_id: task_id.clone(),
4250 meta: None,
4251 }),
4252 McpRequest::UpdateTask(UpdateTaskParams {
4253 task_id: task_id.clone(),
4254 input_responses: HashMap::new(),
4255 meta: None,
4256 }),
4257 McpRequest::CancelTask(CancelTaskParams {
4258 task_id: task_id.clone(),
4259 reason: None,
4260 meta: None,
4261 }),
4262 ]
4263 .into_iter()
4264 .enumerate()
4265 {
4266 let error = router
4267 .handle(
4268 RequestId::Number(id * 10 + offset as i64),
4269 request,
4270 context.clone(),
4271 )
4272 .await
4273 .unwrap_err();
4274 assert!(
4275 matches!(error, Error::JsonRpc(ref e) if e.code == -32602),
4276 "{label} was served: {error:?}"
4277 );
4278 let Error::JsonRpc(error) = error else {
4281 unreachable!()
4282 };
4283 assert!(
4284 error.message.contains("not found"),
4285 "refusal leaked that the task exists: {}",
4286 error.message
4287 );
4288 }
4289 }
4290
4291 assert!(
4293 router
4294 .handle(
4295 RequestId::Number(9),
4296 McpRequest::GetTaskInfo(GetTaskInfoParams {
4297 task_id: task_id.clone(),
4298 meta: None,
4299 }),
4300 as_principal("alice"),
4301 )
4302 .await
4303 .is_ok(),
4304 "a refused cancel must not have cancelled the task"
4305 );
4306 }
4307
4308 #[cfg(all(feature = "oauth", feature = "stateless"))]
4309 #[tokio::test]
4310 async fn final_tasks_work_across_independent_routers_with_a_shared_store() {
4311 fn as_principal(subject: &str) -> Extensions {
4312 let mut extensions = tasks_client_extensions();
4313 extensions.insert(crate::oauth::token::TokenClaims {
4314 sub: Some(subject.to_string()),
4315 iss: None,
4316 aud: None,
4317 exp: None,
4318 scope: None,
4319 client_id: None,
4320 extra: HashMap::new(),
4321 });
4322 extensions
4323 }
4324
4325 fn router_with_store(store: Arc<dyn TaskStore>) -> McpRouter {
4326 McpRouter::new()
4327 .tool(
4328 ToolBuilder::new("shared_task")
4329 .task_support(TaskSupportMode::Optional)
4330 .handler(|_input: serde_json::Value| async move {
4331 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
4332 Ok(CallToolResult::text("done"))
4333 })
4334 .build(),
4335 )
4336 .task_store(store)
4337 .with_tasks()
4338 }
4339
4340 let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
4341 let router_a = router_with_store(store.clone());
4342 let router_b = router_with_store(store);
4343
4344 let McpResponse::FinalCreateTask(created) = router_a
4345 .handle(
4346 RequestId::Number(1),
4347 McpRequest::CallTool(CallToolParams {
4348 name: "shared_task".to_string(),
4349 arguments: serde_json::json!({}),
4350 input_responses: None,
4351 request_state: None,
4352 meta: None,
4353 task: None,
4354 }),
4355 as_principal("alice"),
4356 )
4357 .await
4358 .unwrap()
4359 else {
4360 panic!("router A did not create a final task")
4361 };
4362 let task_id = created.task.metadata.task_id;
4363
4364 assert!(
4366 router_b
4367 .handle(
4368 RequestId::Number(2),
4369 McpRequest::GetTaskInfo(GetTaskInfoParams {
4370 task_id: task_id.clone(),
4371 meta: None,
4372 }),
4373 as_principal("alice"),
4374 )
4375 .await
4376 .is_ok()
4377 );
4378
4379 let denied = router_b
4381 .handle(
4382 RequestId::Number(3),
4383 McpRequest::GetTaskInfo(GetTaskInfoParams {
4384 task_id: task_id.clone(),
4385 meta: None,
4386 }),
4387 as_principal("bob"),
4388 )
4389 .await
4390 .unwrap_err();
4391 let unknown = router_b
4392 .handle(
4393 RequestId::Number(4),
4394 McpRequest::GetTaskInfo(GetTaskInfoParams {
4395 task_id: "unknown-task".to_string(),
4396 meta: None,
4397 }),
4398 as_principal("bob"),
4399 )
4400 .await
4401 .unwrap_err();
4402 let (Error::JsonRpc(denied), Error::JsonRpc(unknown)) = (denied, unknown) else {
4403 panic!("expected JSON-RPC task denials")
4404 };
4405 assert_eq!(denied.code, unknown.code);
4406 assert_eq!(
4407 denied.message.replace(&task_id, "<task-id>"),
4408 unknown.message.replace("unknown-task", "<task-id>")
4409 );
4410 assert_eq!(denied.data, unknown.data);
4411
4412 assert!(matches!(
4415 router_b
4416 .handle(
4417 RequestId::Number(5),
4418 McpRequest::CancelTask(CancelTaskParams {
4419 task_id: task_id.clone(),
4420 reason: None,
4421 meta: None,
4422 }),
4423 as_principal("alice"),
4424 )
4425 .await
4426 .unwrap(),
4427 McpResponse::FinalTaskAck(_)
4428 ));
4429 let McpResponse::FinalGetTask(fetched) = router_a
4430 .handle(
4431 RequestId::Number(6),
4432 McpRequest::GetTaskInfo(GetTaskInfoParams {
4433 task_id,
4434 meta: None,
4435 }),
4436 as_principal("alice"),
4437 )
4438 .await
4439 .unwrap()
4440 else {
4441 panic!("router A did not read the shared task")
4442 };
4443 assert_eq!(fetched.task.status(), TaskStatus::Cancelled);
4444 }
4445
4446 #[test]
4447 fn router_advertises_only_locally_declared_protocol_extensions() {
4448 let router = McpRouter::new().with_protocol_extension(
4449 crate::ExtensionDeclaration::new(
4450 "com.example/rendering",
4451 serde_json::json!({"formats": ["html"]}),
4452 )
4453 .unwrap(),
4454 );
4455
4456 let stable = router.capabilities();
4457 let final_capabilities =
4458 router.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
4459 for capabilities in [stable, final_capabilities] {
4460 let extensions = capabilities.extensions.unwrap();
4461 assert_eq!(extensions.len(), 1);
4462 assert_eq!(extensions["com.example/rendering"]["formats"][0], "html");
4463 assert!(!extensions.contains_key("com.example/client-only"));
4464 }
4465 }
4466
4467 #[tokio::test]
4468 async fn initialize_persists_negotiated_extensions_for_legacy_contexts() {
4469 let router = McpRouter::new().with_protocol_extension(
4470 crate::ExtensionDeclaration::new(
4471 "com.example/shared",
4472 serde_json::json!({"server": true}),
4473 )
4474 .unwrap(),
4475 );
4476 let client_capabilities = ClientCapabilities {
4477 extensions: Some(HashMap::from([
4478 (
4479 "com.example/shared".to_string(),
4480 serde_json::json!({"client": true}),
4481 ),
4482 ("com.example/client-only".to_string(), serde_json::json!({})),
4483 ])),
4484 ..ClientCapabilities::default()
4485 };
4486
4487 router
4488 .handle(
4489 RequestId::Number(1),
4490 McpRequest::Initialize(InitializeParams {
4491 protocol_version: crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
4492 capabilities: client_capabilities,
4493 client_info: Implementation {
4494 name: "extension-test".to_string(),
4495 version: "1.0.0".to_string(),
4496 title: None,
4497 description: None,
4498 icons: None,
4499 website_url: None,
4500 meta: None,
4501 },
4502 meta: None,
4503 }),
4504 Extensions::new(),
4505 )
4506 .await
4507 .unwrap();
4508
4509 let context = router.create_context(RequestId::Number(2), None);
4510 let negotiated = context.negotiated_extensions().unwrap();
4511 assert!(negotiated.contains("com.example/shared"));
4512 assert!(!negotiated.contains("com.example/client-only"));
4513 }
4514
4515 #[cfg(feature = "stateless")]
4516 #[test]
4517 fn final_request_context_exposes_only_negotiated_extensions() {
4518 let router = McpRouter::new().with_protocol_extension(
4519 crate::ExtensionDeclaration::new(
4520 "com.example/shared",
4521 serde_json::json!({"server": true}),
4522 )
4523 .unwrap(),
4524 );
4525 let per_request = final_extensions(ClientCapabilities {
4526 extensions: Some(HashMap::from([
4527 (
4528 "com.example/shared".to_string(),
4529 serde_json::json!({"client": true}),
4530 ),
4531 ("com.example/client-only".to_string(), serde_json::json!({})),
4532 ])),
4533 ..ClientCapabilities::default()
4534 });
4535
4536 let context =
4537 router.create_context_with_extensions(RequestId::Number(1), None, &per_request);
4538 let negotiated = context.negotiated_extensions().unwrap();
4539
4540 assert_eq!(negotiated.len(), 1);
4541 assert_eq!(
4542 negotiated
4543 .get("com.example/shared")
4544 .unwrap()
4545 .client_settings()["client"],
4546 true
4547 );
4548 assert!(!negotiated.contains("com.example/client-only"));
4549 }
4550
4551 #[cfg(feature = "stateless")]
4552 #[tokio::test]
4553 async fn final_protocol_withholds_incomplete_tasks_advertisement() {
4554 let optional = ToolBuilder::new("optional_task")
4555 .task_support(TaskSupportMode::Optional)
4556 .handler(|input: AddInput| async move {
4557 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4558 })
4559 .build();
4560 let required = ToolBuilder::new("required_task")
4561 .task_support(TaskSupportMode::Required)
4562 .handler(|input: AddInput| async move {
4563 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4564 })
4565 .build();
4566 let mut router = McpRouter::new().tool(optional).tool(required);
4567
4568 let stable_capabilities = router.capabilities();
4570 assert!(stable_capabilities.tasks.is_some());
4571 assert!(
4572 stable_capabilities
4573 .extensions
4574 .as_ref()
4575 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID))
4576 );
4577
4578 let response = router
4580 .handle(
4581 RequestId::Number(1),
4582 McpRequest::Discover(DiscoverParams::default()),
4583 Extensions::new(),
4584 )
4585 .await
4586 .unwrap();
4587 let McpResponse::Discover(result) = response else {
4588 panic!("Expected Discover response");
4589 };
4590 assert!(result.capabilities.tasks.is_none());
4591 assert!(
4592 result
4593 .capabilities
4594 .extensions
4595 .as_ref()
4596 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4597 );
4598
4599 init_router(&mut router).await;
4600
4601 let response = router
4603 .handle(
4604 RequestId::Number(2),
4605 McpRequest::ListTools(ListToolsParams::default()),
4606 Extensions::new(),
4607 )
4608 .await
4609 .unwrap();
4610 let McpResponse::ListTools(result) = response else {
4611 panic!("Expected ListTools response");
4612 };
4613 assert_eq!(result.tools.len(), 2);
4614 assert!(result.tools.iter().all(|tool| tool.execution.is_some()));
4615
4616 let response = router
4619 .handle(
4620 RequestId::Number(3),
4621 McpRequest::ListTools(ListToolsParams::default()),
4622 final_extensions(ClientCapabilities::default()),
4623 )
4624 .await
4625 .unwrap();
4626 let McpResponse::ListTools(result) = response else {
4627 panic!("Expected ListTools response");
4628 };
4629 assert_eq!(result.tools.len(), 1);
4630 assert_eq!(result.tools[0].name, "optional_task");
4631 assert!(result.tools[0].execution.is_none());
4632 }
4633
4634 #[cfg(feature = "stateless")]
4635 #[tokio::test]
4636 async fn final_protocol_enforces_tasks_negotiation() {
4637 let optional = ToolBuilder::new("optional_task")
4638 .task_support(TaskSupportMode::Optional)
4639 .handler(|input: AddInput| async move {
4640 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4641 })
4642 .build();
4643 let required = ToolBuilder::new("required_task")
4644 .task_support(TaskSupportMode::Required)
4645 .handler(|input: AddInput| async move {
4646 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4647 })
4648 .build();
4649 let mut router = McpRouter::new().tool(optional).tool(required).with_tasks();
4650 init_router(&mut router).await;
4651
4652 let response = router
4654 .handle(
4655 RequestId::Number(1),
4656 McpRequest::CallTool(CallToolParams {
4657 name: "optional_task".to_string(),
4658 arguments: serde_json::json!({"a": 1, "b": 2}),
4659 input_responses: None,
4660 request_state: None,
4661 meta: None,
4662 task: None,
4663 }),
4664 final_extensions(ClientCapabilities::default()),
4665 )
4666 .await
4667 .unwrap();
4668 assert!(matches!(response, McpResponse::CallTool(_)));
4669
4670 let error = router
4672 .handle(
4673 RequestId::Number(2),
4674 McpRequest::CallTool(CallToolParams {
4675 name: "optional_task".to_string(),
4676 arguments: serde_json::json!({"a": 1, "b": 2}),
4677 input_responses: None,
4678 request_state: None,
4679 meta: None,
4680 task: Some(TaskRequestParams { ttl: None }),
4681 }),
4682 final_extensions(ClientCapabilities::default()),
4683 )
4684 .await
4685 .unwrap_err();
4686 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32602));
4687
4688 let error = router
4692 .handle(
4693 RequestId::Number(3),
4694 McpRequest::CallTool(CallToolParams {
4695 name: "required_task".to_string(),
4696 arguments: serde_json::json!({"a": 1, "b": 2}),
4697 input_responses: None,
4698 request_state: None,
4699 meta: None,
4700 task: None,
4701 }),
4702 final_extensions(ClientCapabilities::default()),
4703 )
4704 .await
4705 .unwrap_err();
4706 let Error::JsonRpc(error) = error else {
4707 panic!("expected a JSON-RPC error");
4708 };
4709 assert_eq!(error.code, -32021);
4710 assert_eq!(
4711 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4712 serde_json::json!({}),
4713 "the error must name the extension the client needs to declare"
4714 );
4715
4716 let task_requests = [
4717 McpRequest::GetTaskInfo(GetTaskInfoParams {
4718 task_id: "task-unknown".to_string(),
4719 meta: None,
4720 }),
4721 McpRequest::UpdateTask(UpdateTaskParams {
4722 task_id: "task-unknown".to_string(),
4723 input_responses: HashMap::new(),
4724 meta: None,
4725 }),
4726 McpRequest::CancelTask(CancelTaskParams {
4727 task_id: "task-unknown".to_string(),
4728 reason: None,
4729 meta: None,
4730 }),
4731 ];
4732 for (index, request) in task_requests.into_iter().enumerate() {
4733 let error = router
4734 .handle(
4735 RequestId::Number(4 + index as i64),
4736 request,
4737 final_extensions(ClientCapabilities::default()),
4738 )
4739 .await
4740 .unwrap_err();
4741 let Error::JsonRpc(error) = error else {
4742 panic!("expected a JSON-RPC error");
4743 };
4744 assert_eq!(error.code, -32021);
4745 assert_eq!(
4746 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4747 serde_json::json!({})
4748 );
4749 }
4750
4751 let router_without_tasks = McpRouter::new();
4754 let error = router_without_tasks
4755 .handle(
4756 RequestId::Number(7),
4757 McpRequest::GetTaskInfo(GetTaskInfoParams {
4758 task_id: "task-unknown".to_string(),
4759 meta: None,
4760 }),
4761 final_extensions(tasks_client_capabilities()),
4762 )
4763 .await
4764 .unwrap_err();
4765 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32601));
4766 }
4767
4768 #[cfg(feature = "stateless")]
4769 #[test]
4770 fn input_required_capability_validation_uses_capability_semantics() {
4771 let roots = InputRequiredResult::with_requests(
4772 [(
4773 "roots".to_string(),
4774 InputRequest::ListRoots(ListRootsParams::default()),
4775 )]
4776 .into_iter()
4777 .collect(),
4778 );
4779 let extensions = final_extensions(ClientCapabilities {
4780 roots: Some(RootsCapability {
4781 list_changed: true,
4782 deprecated: None,
4783 }),
4784 ..Default::default()
4785 });
4786 validate_input_required_result(&extensions, &roots).unwrap();
4787 assert!(client_capabilities_satisfy(
4788 extensions
4789 .get::<crate::stateless::StatelessRequestMeta>()
4790 .and_then(|meta| meta.client_capabilities.as_ref())
4791 .unwrap(),
4792 &ClientCapabilities {
4793 roots: Some(RootsCapability::default()),
4794 ..Default::default()
4795 }
4796 ));
4797
4798 let sampling_with_tools = InputRequiredResult::with_requests(
4799 [(
4800 "sample".to_string(),
4801 InputRequest::CreateMessage(CreateMessageParams {
4802 tools: Some(Vec::new()),
4803 ..CreateMessageParams::new(vec![SamplingMessage::user("hello")], 10)
4804 }),
4805 )]
4806 .into_iter()
4807 .collect(),
4808 );
4809 let extensions = final_extensions(ClientCapabilities {
4810 sampling: Some(SamplingCapability::default()),
4811 ..Default::default()
4812 });
4813 assert!(validate_input_required_result(&extensions, &sampling_with_tools).is_err());
4814
4815 let form = InputRequiredResult::with_requests(
4816 [(
4817 "form".to_string(),
4818 InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
4819 mode: Some(ElicitMode::Form),
4820 message: "name".into(),
4821 requested_schema: ElicitFormSchema::new(),
4822 meta: None,
4823 })),
4824 )]
4825 .into_iter()
4826 .collect(),
4827 );
4828 let extensions = final_extensions(ClientCapabilities {
4829 elicitation: Some(ElicitationCapability::default()),
4830 ..Default::default()
4831 });
4832 validate_input_required_result(&extensions, &form).unwrap();
4833 }
4834
4835 async fn init_router(router: &mut McpRouter) {
4837 let init_req = RouterRequest {
4839 id: RequestId::Number(0),
4840 inner: McpRequest::Initialize(InitializeParams {
4841 protocol_version: "2025-11-25".to_string(),
4842 capabilities: ClientCapabilities {
4843 roots: None,
4844 sampling: None,
4845 elicitation: None,
4846 tasks: None,
4847 experimental: None,
4848 extensions: None,
4849 },
4850 client_info: Implementation {
4851 name: "test".to_string(),
4852 version: "1.0".to_string(),
4853 ..Default::default()
4854 },
4855 meta: None,
4856 }),
4857 extensions: Extensions::new(),
4858 };
4859 let _ = router.ready().await.unwrap().call(init_req).await.unwrap();
4860 router.handle_notification(McpNotification::Initialized);
4862 }
4863
4864 #[tokio::test]
4865 async fn test_router_list_tools() {
4866 let add_tool = ToolBuilder::new("add")
4867 .description("Add two numbers")
4868 .handler(|input: AddInput| async move {
4869 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4870 })
4871 .build();
4872
4873 let mut router = McpRouter::new().tool(add_tool);
4874
4875 init_router(&mut router).await;
4877
4878 let req = RouterRequest {
4879 id: RequestId::Number(1),
4880 inner: McpRequest::ListTools(ListToolsParams::default()),
4881 extensions: Extensions::new(),
4882 };
4883
4884 let resp = router.ready().await.unwrap().call(req).await.unwrap();
4885
4886 match resp.inner {
4887 Ok(McpResponse::ListTools(result)) => {
4888 assert_eq!(result.tools.len(), 1);
4889 assert_eq!(result.tools[0].name, "add");
4890 }
4891 _ => panic!("Expected ListTools response"),
4892 }
4893 }
4894
4895 #[tokio::test]
4896 async fn test_router_call_tool() {
4897 let add_tool = ToolBuilder::new("add")
4898 .description("Add two numbers")
4899 .handler(|input: AddInput| async move {
4900 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4901 })
4902 .build();
4903
4904 let mut router = McpRouter::new().tool(add_tool);
4905
4906 init_router(&mut router).await;
4908
4909 let req = RouterRequest {
4910 id: RequestId::Number(1),
4911 inner: McpRequest::CallTool(CallToolParams {
4912 input_responses: None,
4913 request_state: None,
4914 name: "add".to_string(),
4915 arguments: serde_json::json!({"a": 2, "b": 3}),
4916 meta: None,
4917 task: None,
4918 }),
4919 extensions: Extensions::new(),
4920 };
4921
4922 let resp = router.ready().await.unwrap().call(req).await.unwrap();
4923
4924 match resp.inner {
4925 Ok(McpResponse::CallTool(result)) => {
4926 assert!(!result.is_error);
4927 match &result.content[0] {
4929 Content::Text { text, .. } => assert_eq!(text, "5"),
4930 _ => panic!("Expected text content"),
4931 }
4932 }
4933 _ => panic!("Expected CallTool response"),
4934 }
4935 }
4936
4937 async fn init_jsonrpc_service(service: &mut JsonRpcService<McpRouter>, router: &McpRouter) {
4939 let init_req = JsonRpcRequest::new(0, "initialize").with_params(serde_json::json!({
4940 "protocolVersion": "2025-11-25",
4941 "capabilities": {},
4942 "clientInfo": { "name": "test", "version": "1.0" }
4943 }));
4944 let _ = service.call_single(init_req).await.unwrap();
4945 router.handle_notification(McpNotification::Initialized);
4946 }
4947
4948 #[tokio::test]
4949 async fn test_jsonrpc_service() {
4950 let add_tool = ToolBuilder::new("add")
4951 .description("Add two numbers")
4952 .handler(|input: AddInput| async move {
4953 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4954 })
4955 .build();
4956
4957 let router = McpRouter::new().tool(add_tool);
4958 let mut service = JsonRpcService::new(router.clone());
4959
4960 init_jsonrpc_service(&mut service, &router).await;
4962
4963 let req = JsonRpcRequest::new(1, "tools/list");
4964
4965 let resp = service.call_single(req).await.unwrap();
4966
4967 match resp {
4968 JsonRpcResponse::Result(r) => {
4969 assert_eq!(r.id, RequestId::Number(1));
4970 let tools = r.result.get("tools").unwrap().as_array().unwrap();
4971 assert_eq!(tools.len(), 1);
4972 }
4973 JsonRpcResponse::Error(_) => panic!("Expected success response"),
4974 _ => panic!("unexpected response variant"),
4975 }
4976 }
4977
4978 #[tokio::test]
4979 async fn test_batch_request() {
4980 let add_tool = ToolBuilder::new("add")
4981 .description("Add two numbers")
4982 .handler(|input: AddInput| async move {
4983 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4984 })
4985 .build();
4986
4987 let router = McpRouter::new().tool(add_tool);
4988 let mut service = JsonRpcService::new(router.clone())
4989 .protocol_versions(["2025-03-26"])
4990 .unwrap();
4991
4992 init_jsonrpc_service(&mut service, &router).await;
4994
4995 let requests = vec![
4997 JsonRpcRequest::new(1, "tools/list"),
4998 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
4999 "name": "add",
5000 "arguments": {"a": 10, "b": 20}
5001 })),
5002 JsonRpcRequest::new(3, "ping"),
5003 ];
5004
5005 let responses = service.call_batch(requests).await.unwrap();
5006
5007 assert_eq!(responses.len(), 3);
5008
5009 match &responses[0] {
5011 JsonRpcResponse::Result(r) => {
5012 assert_eq!(r.id, RequestId::Number(1));
5013 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5014 assert_eq!(tools.len(), 1);
5015 }
5016 JsonRpcResponse::Error(_) => panic!("Expected success for tools/list"),
5017 _ => panic!("unexpected response variant"),
5018 }
5019
5020 match &responses[1] {
5022 JsonRpcResponse::Result(r) => {
5023 assert_eq!(r.id, RequestId::Number(2));
5024 let content = r.result.get("content").unwrap().as_array().unwrap();
5025 let text = content[0].get("text").unwrap().as_str().unwrap();
5026 assert_eq!(text, "30");
5027 }
5028 JsonRpcResponse::Error(_) => panic!("Expected success for tools/call"),
5029 _ => panic!("unexpected response variant"),
5030 }
5031
5032 match &responses[2] {
5034 JsonRpcResponse::Result(r) => {
5035 assert_eq!(r.id, RequestId::Number(3));
5036 }
5037 JsonRpcResponse::Error(_) => panic!("Expected success for ping"),
5038 _ => panic!("unexpected response variant"),
5039 }
5040 }
5041
5042 #[tokio::test]
5043 async fn test_empty_batch_error() {
5044 let router = McpRouter::new();
5045 let mut service = JsonRpcService::new(router);
5046
5047 let result = service.call_batch(vec![]).await;
5048 assert!(result.is_err());
5049 }
5050
5051 #[tokio::test]
5056 async fn test_progress_token_extraction() {
5057 use crate::context::{ServerNotification, notification_channel};
5058 use crate::protocol::ProgressToken;
5059 use std::sync::Arc;
5060 use std::sync::atomic::{AtomicBool, Ordering};
5061
5062 let progress_reported = Arc::new(AtomicBool::new(false));
5064 let progress_ref = progress_reported.clone();
5065
5066 let tool = ToolBuilder::new("progress_tool")
5068 .description("Tool that reports progress")
5069 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5070 let reported = progress_ref.clone();
5071 async move {
5072 ctx.report_progress(50.0, Some(100.0), Some("Halfway"))
5074 .await;
5075 reported.store(true, Ordering::SeqCst);
5076 Ok(CallToolResult::text("done"))
5077 }
5078 })
5079 .build();
5080
5081 let (tx, mut rx) = notification_channel(10);
5083 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5084 let mut service = JsonRpcService::new(router.clone());
5085
5086 init_jsonrpc_service(&mut service, &router).await;
5088
5089 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5091 "name": "progress_tool",
5092 "arguments": {"a": 1, "b": 2},
5093 "_meta": {
5094 "progressToken": "test-token-123"
5095 }
5096 }));
5097
5098 let resp = service.call_single(req).await.unwrap();
5099
5100 match resp {
5102 JsonRpcResponse::Result(_) => {}
5103 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
5104 _ => panic!("unexpected response variant"),
5105 }
5106
5107 assert!(progress_reported.load(Ordering::SeqCst));
5109
5110 let notification = rx.try_recv().expect("Expected progress notification");
5112 match notification {
5113 ServerNotification::Progress(params) => {
5114 assert_eq!(
5115 params.progress_token,
5116 ProgressToken::String("test-token-123".to_string())
5117 );
5118 assert_eq!(params.progress, 50.0);
5119 assert_eq!(params.total, Some(100.0));
5120 assert_eq!(params.message.as_deref(), Some("Halfway"));
5121 }
5122 _ => panic!("Expected Progress notification"),
5123 }
5124 }
5125
5126 #[tokio::test]
5127 async fn test_tool_call_without_progress_token() {
5128 use crate::context::notification_channel;
5129 use std::sync::Arc;
5130 use std::sync::atomic::{AtomicBool, Ordering};
5131
5132 let progress_attempted = Arc::new(AtomicBool::new(false));
5133 let progress_ref = progress_attempted.clone();
5134
5135 let tool = ToolBuilder::new("no_token_tool")
5136 .description("Tool that tries to report progress without token")
5137 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5138 let attempted = progress_ref.clone();
5139 async move {
5140 ctx.report_progress(50.0, Some(100.0), None).await;
5142 attempted.store(true, Ordering::SeqCst);
5143 Ok(CallToolResult::text("done"))
5144 }
5145 })
5146 .build();
5147
5148 let (tx, mut rx) = notification_channel(10);
5149 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5150 let mut service = JsonRpcService::new(router.clone());
5151
5152 init_jsonrpc_service(&mut service, &router).await;
5153
5154 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5156 "name": "no_token_tool",
5157 "arguments": {"a": 1, "b": 2}
5158 }));
5159
5160 let resp = service.call_single(req).await.unwrap();
5161 assert!(matches!(resp, JsonRpcResponse::Result(_)));
5162
5163 assert!(progress_attempted.load(Ordering::SeqCst));
5165
5166 assert!(rx.try_recv().is_err());
5168 }
5169
5170 #[tokio::test]
5171 async fn test_batch_errors_returned_not_dropped() {
5172 let add_tool = ToolBuilder::new("add")
5173 .description("Add two numbers")
5174 .handler(|input: AddInput| async move {
5175 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5176 })
5177 .build();
5178
5179 let router = McpRouter::new().tool(add_tool);
5180 let mut service = JsonRpcService::new(router.clone())
5181 .protocol_versions(["2025-03-26"])
5182 .unwrap();
5183
5184 init_jsonrpc_service(&mut service, &router).await;
5185
5186 let requests = vec![
5188 JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5190 "name": "add",
5191 "arguments": {"a": 10, "b": 20}
5192 })),
5193 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5195 "name": "nonexistent_tool",
5196 "arguments": {}
5197 })),
5198 JsonRpcRequest::new(3, "ping"),
5200 ];
5201
5202 let responses = service.call_batch(requests).await.unwrap();
5203
5204 assert_eq!(responses.len(), 3);
5206
5207 match &responses[0] {
5209 JsonRpcResponse::Result(r) => {
5210 assert_eq!(r.id, RequestId::Number(1));
5211 }
5212 JsonRpcResponse::Error(_) => panic!("Expected success for first request"),
5213 _ => panic!("unexpected response variant"),
5214 }
5215
5216 match &responses[1] {
5218 JsonRpcResponse::Error(e) => {
5219 assert_eq!(e.id, Some(RequestId::Number(2)));
5220 assert!(e.error.message.contains("not found") || e.error.code == -32601);
5222 }
5223 JsonRpcResponse::Result(_) => panic!("Expected error for second request"),
5224 _ => panic!("unexpected response variant"),
5225 }
5226
5227 match &responses[2] {
5229 JsonRpcResponse::Result(r) => {
5230 assert_eq!(r.id, RequestId::Number(3));
5231 }
5232 JsonRpcResponse::Error(_) => panic!("Expected success for third request"),
5233 _ => panic!("unexpected response variant"),
5234 }
5235 }
5236
5237 #[tokio::test]
5242 async fn test_list_resource_templates() {
5243 use crate::resource::ResourceTemplateBuilder;
5244 use std::collections::HashMap;
5245
5246 let template = ResourceTemplateBuilder::new("file:///{path}")
5247 .name("Project Files")
5248 .description("Access project files")
5249 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5250 Ok(ReadResourceResult {
5251 contents: vec![ResourceContent {
5252 uri,
5253 mime_type: None,
5254 text: None,
5255 blob: None,
5256 meta: None,
5257 }],
5258 meta: None,
5259 ..Default::default()
5260 })
5261 });
5262
5263 let mut router = McpRouter::new().resource_template(template);
5264
5265 init_router(&mut router).await;
5267
5268 let req = RouterRequest {
5269 id: RequestId::Number(1),
5270 inner: McpRequest::ListResourceTemplates(ListResourceTemplatesParams::default()),
5271 extensions: Extensions::new(),
5272 };
5273
5274 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5275
5276 match resp.inner {
5277 Ok(McpResponse::ListResourceTemplates(result)) => {
5278 assert_eq!(result.resource_templates.len(), 1);
5279 assert_eq!(result.resource_templates[0].uri_template, "file:///{path}");
5280 assert_eq!(result.resource_templates[0].name, "Project Files");
5281 }
5282 _ => panic!("Expected ListResourceTemplates response"),
5283 }
5284 }
5285
5286 #[tokio::test]
5287 async fn test_read_resource_via_template() {
5288 use crate::resource::ResourceTemplateBuilder;
5289 use std::collections::HashMap;
5290
5291 let template = ResourceTemplateBuilder::new("db://users/{id}")
5292 .name("User Records")
5293 .handler(|uri: String, vars: HashMap<String, String>| async move {
5294 let id = vars.get("id").unwrap().clone();
5295 Ok(ReadResourceResult {
5296 contents: vec![ResourceContent {
5297 uri,
5298 mime_type: Some("application/json".to_string()),
5299 text: Some(format!(r#"{{"id": "{}"}}"#, id)),
5300 blob: None,
5301 meta: None,
5302 }],
5303 meta: None,
5304 ..Default::default()
5305 })
5306 });
5307
5308 let mut router = McpRouter::new().resource_template(template);
5309
5310 init_router(&mut router).await;
5312
5313 let req = RouterRequest {
5315 id: RequestId::Number(1),
5316 inner: McpRequest::ReadResource(ReadResourceParams {
5317 input_responses: None,
5318 request_state: None,
5319 uri: "db://users/123".to_string(),
5320 meta: None,
5321 }),
5322 extensions: Extensions::new(),
5323 };
5324
5325 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5326
5327 match resp.inner {
5328 Ok(McpResponse::ReadResource(result)) => {
5329 assert_eq!(result.contents.len(), 1);
5330 assert_eq!(result.contents[0].uri, "db://users/123");
5331 assert!(result.contents[0].text.as_ref().unwrap().contains("123"));
5332 }
5333 _ => panic!("Expected ReadResource response"),
5334 }
5335 }
5336
5337 #[tokio::test]
5338 async fn test_static_resource_takes_precedence_over_template() {
5339 use crate::resource::{ResourceBuilder, ResourceTemplateBuilder};
5340 use std::collections::HashMap;
5341
5342 let template = ResourceTemplateBuilder::new("file:///{path}")
5344 .name("Files Template")
5345 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5346 Ok(ReadResourceResult {
5347 contents: vec![ResourceContent {
5348 uri,
5349 mime_type: None,
5350 text: Some("from template".to_string()),
5351 blob: None,
5352 meta: None,
5353 }],
5354 meta: None,
5355 ..Default::default()
5356 })
5357 });
5358
5359 let static_resource = ResourceBuilder::new("file:///README.md")
5361 .name("README")
5362 .text("from static resource");
5363
5364 let mut router = McpRouter::new()
5365 .resource_template(template)
5366 .resource(static_resource);
5367
5368 init_router(&mut router).await;
5370
5371 let req = RouterRequest {
5373 id: RequestId::Number(1),
5374 inner: McpRequest::ReadResource(ReadResourceParams {
5375 input_responses: None,
5376 request_state: None,
5377 uri: "file:///README.md".to_string(),
5378 meta: None,
5379 }),
5380 extensions: Extensions::new(),
5381 };
5382
5383 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5384
5385 match resp.inner {
5386 Ok(McpResponse::ReadResource(result)) => {
5387 assert_eq!(
5389 result.contents[0].text.as_deref(),
5390 Some("from static resource")
5391 );
5392 }
5393 _ => panic!("Expected ReadResource response"),
5394 }
5395 }
5396
5397 #[tokio::test]
5398 async fn test_resource_not_found_when_no_match() {
5399 use crate::resource::ResourceTemplateBuilder;
5400 use std::collections::HashMap;
5401
5402 let template = ResourceTemplateBuilder::new("db://users/{id}")
5403 .name("Users")
5404 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5405 Ok(ReadResourceResult {
5406 contents: vec![ResourceContent {
5407 uri,
5408 mime_type: None,
5409 text: None,
5410 blob: None,
5411 meta: None,
5412 }],
5413 meta: None,
5414 ..Default::default()
5415 })
5416 });
5417
5418 let mut router = McpRouter::new().resource_template(template);
5419
5420 init_router(&mut router).await;
5422
5423 let req = RouterRequest {
5425 id: RequestId::Number(1),
5426 inner: McpRequest::ReadResource(ReadResourceParams {
5427 input_responses: None,
5428 request_state: None,
5429 uri: "db://posts/123".to_string(),
5430 meta: None,
5431 }),
5432 extensions: Extensions::new(),
5433 };
5434
5435 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5436
5437 match resp.inner {
5438 Err(err) => {
5439 assert!(err.message.contains("not found"));
5440 }
5441 Ok(_) => panic!("Expected error for non-matching URI"),
5442 }
5443 }
5444
5445 #[tokio::test]
5446 async fn test_capabilities_include_resources_with_only_templates() {
5447 use crate::resource::ResourceTemplateBuilder;
5448 use std::collections::HashMap;
5449
5450 let template = ResourceTemplateBuilder::new("file:///{path}")
5451 .name("Files")
5452 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5453 Ok(ReadResourceResult {
5454 contents: vec![ResourceContent {
5455 uri,
5456 mime_type: None,
5457 text: None,
5458 blob: None,
5459 meta: None,
5460 }],
5461 meta: None,
5462 ..Default::default()
5463 })
5464 });
5465
5466 let mut router = McpRouter::new().resource_template(template);
5467
5468 let init_req = RouterRequest {
5470 id: RequestId::Number(0),
5471 inner: McpRequest::Initialize(InitializeParams {
5472 protocol_version: "2025-11-25".to_string(),
5473 capabilities: ClientCapabilities {
5474 roots: None,
5475 sampling: None,
5476 elicitation: None,
5477 tasks: None,
5478 experimental: None,
5479 extensions: None,
5480 },
5481 client_info: Implementation {
5482 name: "test".to_string(),
5483 version: "1.0".to_string(),
5484 ..Default::default()
5485 },
5486 meta: None,
5487 }),
5488 extensions: Extensions::new(),
5489 };
5490 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5491
5492 match resp.inner {
5493 Ok(McpResponse::Initialize(result)) => {
5494 assert!(result.capabilities.resources.is_some());
5496 }
5497 _ => panic!("Expected Initialize response"),
5498 }
5499 }
5500
5501 #[tokio::test]
5506 async fn test_log_sends_notification() {
5507 use crate::context::notification_channel;
5508
5509 let (tx, mut rx) = notification_channel(10);
5510 let router = McpRouter::new().with_notification_sender(tx);
5511
5512 let sent = router.log_info("Test message");
5514 assert!(sent);
5515
5516 let notification = rx.try_recv().unwrap();
5518 match notification {
5519 ServerNotification::LogMessage(params) => {
5520 assert_eq!(params.level, LogLevel::Info);
5521 let data = params.data;
5522 assert_eq!(
5523 data.get("message").unwrap().as_str().unwrap(),
5524 "Test message"
5525 );
5526 }
5527 _ => panic!("Expected LogMessage notification"),
5528 }
5529 }
5530
5531 #[tokio::test]
5532 async fn test_log_with_custom_params() {
5533 use crate::context::notification_channel;
5534
5535 let (tx, mut rx) = notification_channel(10);
5536 let router = McpRouter::new().with_notification_sender(tx);
5537
5538 let params = LoggingMessageParams::new(
5540 LogLevel::Error,
5541 serde_json::json!({
5542 "error": "Connection failed",
5543 "host": "localhost"
5544 }),
5545 )
5546 .with_logger("database");
5547
5548 let sent = router.log(params);
5549 assert!(sent);
5550
5551 let notification = rx.try_recv().unwrap();
5552 match notification {
5553 ServerNotification::LogMessage(params) => {
5554 assert_eq!(params.level, LogLevel::Error);
5555 assert_eq!(params.logger.as_deref(), Some("database"));
5556 let data = params.data;
5557 assert_eq!(
5558 data.get("error").unwrap().as_str().unwrap(),
5559 "Connection failed"
5560 );
5561 }
5562 _ => panic!("Expected LogMessage notification"),
5563 }
5564 }
5565
5566 #[tokio::test]
5567 async fn test_log_without_channel_returns_false() {
5568 let router = McpRouter::new();
5570
5571 assert!(!router.log_info("Test"));
5573 assert!(!router.log_warning("Test"));
5574 assert!(!router.log_error("Test"));
5575 assert!(!router.log_debug("Test"));
5576 }
5577
5578 #[tokio::test]
5579 async fn test_logging_capability_with_channel() {
5580 use crate::context::notification_channel;
5581
5582 let (tx, _rx) = notification_channel(10);
5583 let mut router = McpRouter::new().with_notification_sender(tx);
5584
5585 let init_req = RouterRequest {
5587 id: RequestId::Number(0),
5588 inner: McpRequest::Initialize(InitializeParams {
5589 protocol_version: "2025-11-25".to_string(),
5590 capabilities: ClientCapabilities {
5591 roots: None,
5592 sampling: None,
5593 elicitation: None,
5594 tasks: None,
5595 experimental: None,
5596 extensions: None,
5597 },
5598 client_info: Implementation {
5599 name: "test".to_string(),
5600 version: "1.0".to_string(),
5601 ..Default::default()
5602 },
5603 meta: None,
5604 }),
5605 extensions: Extensions::new(),
5606 };
5607 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5608
5609 match resp.inner {
5610 Ok(McpResponse::Initialize(result)) => {
5611 assert!(result.capabilities.logging.is_some());
5613 }
5614 _ => panic!("Expected Initialize response"),
5615 }
5616 }
5617
5618 #[tokio::test]
5619 async fn test_no_logging_capability_without_channel() {
5620 let mut router = McpRouter::new();
5621
5622 let init_req = RouterRequest {
5624 id: RequestId::Number(0),
5625 inner: McpRequest::Initialize(InitializeParams {
5626 protocol_version: "2025-11-25".to_string(),
5627 capabilities: ClientCapabilities {
5628 roots: None,
5629 sampling: None,
5630 elicitation: None,
5631 tasks: None,
5632 experimental: None,
5633 extensions: None,
5634 },
5635 client_info: Implementation {
5636 name: "test".to_string(),
5637 version: "1.0".to_string(),
5638 ..Default::default()
5639 },
5640 meta: None,
5641 }),
5642 extensions: Extensions::new(),
5643 };
5644 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5645
5646 match resp.inner {
5647 Ok(McpResponse::Initialize(result)) => {
5648 assert!(result.capabilities.logging.is_none());
5650 }
5651 _ => panic!("Expected Initialize response"),
5652 }
5653 }
5654
5655 #[tokio::test]
5660 async fn test_create_task_via_call_tool() {
5661 let add_tool = ToolBuilder::new("add")
5662 .description("Add two numbers")
5663 .task_support(TaskSupportMode::Optional)
5664 .handler(|input: AddInput| async move {
5665 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5666 })
5667 .build();
5668
5669 let mut router = McpRouter::new().tool(add_tool);
5670 init_router(&mut router).await;
5671
5672 let req = RouterRequest {
5673 id: RequestId::Number(1),
5674 inner: McpRequest::CallTool(CallToolParams {
5675 input_responses: None,
5676 request_state: None,
5677 name: "add".to_string(),
5678 arguments: serde_json::json!({"a": 5, "b": 10}),
5679 meta: None,
5680 task: Some(TaskRequestParams { ttl: None }),
5681 }),
5682 extensions: Extensions::new(),
5683 };
5684
5685 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5686
5687 match resp.inner {
5688 Ok(McpResponse::CreateTask(result)) => {
5689 assert!(!result.task.task_id.is_empty());
5690 assert_eq!(result.task.status, TaskStatus::Working);
5691 }
5692 _ => panic!("Expected CreateTask response"),
5693 }
5694 }
5695
5696 struct CountingTaskStore {
5699 inner: MemoryTaskStore,
5700 creates: std::sync::atomic::AtomicUsize,
5701 gets: std::sync::atomic::AtomicUsize,
5702 completes: std::sync::atomic::AtomicUsize,
5703 }
5704
5705 impl CountingTaskStore {
5706 fn new() -> Self {
5707 Self {
5708 inner: MemoryTaskStore::new(),
5709 creates: std::sync::atomic::AtomicUsize::new(0),
5710 gets: std::sync::atomic::AtomicUsize::new(0),
5711 completes: std::sync::atomic::AtomicUsize::new(0),
5712 }
5713 }
5714 }
5715
5716 #[async_trait::async_trait]
5717 impl TaskStore for CountingTaskStore {
5718 async fn create_task(
5719 &self,
5720 tool_name: &str,
5721 arguments: serde_json::Value,
5722 ttl: Option<u64>,
5723 owner: crate::async_task::TaskOwner,
5724 ) -> crate::async_task::Result<(String, crate::async_task::CancellationToken)> {
5725 self.creates
5726 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5727 self.inner
5728 .create_task(tool_name, arguments, ttl, owner)
5729 .await
5730 }
5731
5732 async fn get_task(&self, task_id: &str) -> crate::async_task::Result<Option<TaskObject>> {
5733 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5734 self.inner.get_task(task_id).await
5735 }
5736
5737 async fn task_owner(
5738 &self,
5739 task_id: &str,
5740 ) -> crate::async_task::Result<Option<crate::async_task::TaskOwner>> {
5741 self.inner.task_owner(task_id).await
5742 }
5743
5744 async fn get_task_result(
5745 &self,
5746 task_id: &str,
5747 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
5748 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5751 self.inner.get_task_result(task_id).await
5752 }
5753
5754 async fn wait_for_completion(
5755 &self,
5756 task_id: &str,
5757 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
5758 self.inner.wait_for_completion(task_id).await
5759 }
5760
5761 async fn list_tasks(
5762 &self,
5763 status_filter: Option<TaskStatus>,
5764 ) -> crate::async_task::Result<Vec<TaskObject>> {
5765 self.inner.list_tasks(status_filter).await
5766 }
5767
5768 async fn require_input(
5769 &self,
5770 task_id: &str,
5771 requests: crate::protocol::InputRequests,
5772 message: Option<&str>,
5773 ) -> crate::async_task::Result<bool> {
5774 self.inner.require_input(task_id, requests, message).await
5775 }
5776
5777 async fn outstanding_input_requests(
5778 &self,
5779 task_id: &str,
5780 ) -> crate::async_task::Result<Option<crate::protocol::InputRequests>> {
5781 self.inner.outstanding_input_requests(task_id).await
5782 }
5783
5784 async fn apply_input_responses(
5785 &self,
5786 task_id: &str,
5787 responses: crate::protocol::InputResponses,
5788 ) -> crate::async_task::Result<Option<crate::async_task::AppliedInputResponses>> {
5789 self.inner.apply_input_responses(task_id, responses).await
5790 }
5791
5792 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> crate::async_task::Result<bool> {
5793 self.inner.set_ttl(task_id, ttl_ms).await
5794 }
5795
5796 async fn complete_task(
5797 &self,
5798 task_id: &str,
5799 result: CallToolResult,
5800 ) -> crate::async_task::Result<bool> {
5801 self.completes
5802 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5803 self.inner.complete_task(task_id, result).await
5804 }
5805
5806 async fn fail_task(
5807 &self,
5808 task_id: &str,
5809 error: JsonRpcError,
5810 ) -> crate::async_task::Result<bool> {
5811 self.inner.fail_task(task_id, error).await
5812 }
5813
5814 async fn cancel_task(
5815 &self,
5816 task_id: &str,
5817 reason: Option<&str>,
5818 ) -> crate::async_task::Result<Option<TaskObject>> {
5819 self.inner.cancel_task(task_id, reason).await
5820 }
5821 }
5822
5823 #[tokio::test]
5824 async fn test_injected_task_store_used_by_dispatch() {
5825 let store = Arc::new(CountingTaskStore::new());
5826
5827 let add_tool = ToolBuilder::new("add")
5828 .description("Add two numbers")
5829 .task_support(TaskSupportMode::Optional)
5830 .handler(|input: AddInput| async move {
5831 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5832 })
5833 .build();
5834
5835 let mut router = McpRouter::new()
5836 .tool(add_tool)
5837 .task_store(store.clone() as Arc<dyn TaskStore>);
5838 init_router(&mut router).await;
5839
5840 let req = RouterRequest {
5842 id: RequestId::Number(1),
5843 inner: McpRequest::CallTool(CallToolParams {
5844 input_responses: None,
5845 request_state: None,
5846 name: "add".to_string(),
5847 arguments: serde_json::json!({"a": 2, "b": 3}),
5848 meta: None,
5849 task: Some(TaskRequestParams { ttl: None }),
5850 }),
5851 extensions: Extensions::new(),
5852 };
5853 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5854 let task_id = match resp.inner {
5855 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
5856 other => panic!("Expected CreateTask response, got {other:?}"),
5857 };
5858
5859 assert_eq!(
5860 store.creates.load(std::sync::atomic::Ordering::Relaxed),
5861 1,
5862 "create_task must go through the injected store"
5863 );
5864
5865 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
5867 assert_eq!(
5868 store.completes.load(std::sync::atomic::Ordering::Relaxed),
5869 1,
5870 "complete_task must go through the injected store"
5871 );
5872
5873 let gets_before = store.gets.load(std::sync::atomic::Ordering::Relaxed);
5875 let req = RouterRequest {
5876 id: RequestId::Number(2),
5877 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
5878 task_id: task_id.clone(),
5879 meta: None,
5880 }),
5881 extensions: Extensions::new(),
5882 };
5883 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5884 match resp.inner {
5885 Ok(McpResponse::GetTaskInfo(info)) => {
5886 assert_eq!(info.task_id, task_id);
5887 assert_eq!(info.status, TaskStatus::Completed);
5888 }
5889 other => panic!("Expected GetTaskInfo response, got {other:?}"),
5890 }
5891 assert!(
5892 store.gets.load(std::sync::atomic::Ordering::Relaxed) > gets_before,
5893 "tasks/get must go through the injected store"
5894 );
5895 }
5896
5897 #[tokio::test]
5898 async fn test_removed_tasks_methods_get_method_not_found() {
5899 let mut router = McpRouter::new();
5903 init_router(&mut router).await;
5904
5905 for method in ["tasks/list", "tasks/result"] {
5906 let req = RouterRequest {
5907 id: RequestId::Number(1),
5908 inner: McpRequest::Unknown {
5909 method: method.to_string(),
5910 params: None,
5911 },
5912 extensions: Extensions::new(),
5913 };
5914
5915 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5916
5917 match resp.inner {
5918 Err(err) => {
5919 assert_eq!(err.code, -32601, "{method} must be MethodNotFound");
5920 }
5921 other => panic!("Expected MethodNotFound error for {method}, got {other:?}"),
5922 }
5923 }
5924 }
5925
5926 #[tokio::test]
5927 async fn test_task_lifecycle_complete() {
5928 let add_tool = ToolBuilder::new("add")
5929 .description("Add two numbers")
5930 .task_support(TaskSupportMode::Optional)
5931 .handler(|input: AddInput| async move {
5932 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5933 })
5934 .build();
5935
5936 let mut router = McpRouter::new().tool(add_tool);
5937 init_router(&mut router).await;
5938
5939 let req = RouterRequest {
5941 id: RequestId::Number(1),
5942 inner: McpRequest::CallTool(CallToolParams {
5943 input_responses: None,
5944 request_state: None,
5945 name: "add".to_string(),
5946 arguments: serde_json::json!({"a": 7, "b": 8}),
5947 meta: None,
5948 task: Some(TaskRequestParams { ttl: None }),
5949 }),
5950 extensions: Extensions::new(),
5951 };
5952
5953 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5954 let task_id = match resp.inner {
5955 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
5956 _ => panic!("Expected CreateTask response"),
5957 };
5958
5959 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
5961
5962 let req = RouterRequest {
5966 id: RequestId::Number(2),
5967 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
5968 task_id: task_id.clone(),
5969 meta: None,
5970 }),
5971 extensions: Extensions::new(),
5972 };
5973
5974 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5975
5976 match resp.inner {
5977 Ok(McpResponse::GetTaskInfo(info)) => {
5978 assert_eq!(info.task_id, task_id);
5979 assert_eq!(info.status, TaskStatus::Completed);
5980 }
5981 _ => panic!("Expected GetTaskInfo response"),
5982 }
5983 }
5984
5985 #[tokio::test]
5986 async fn test_task_cancellation() {
5987 let slow_tool = ToolBuilder::new("slow")
5989 .description("Slow tool")
5990 .task_support(TaskSupportMode::Optional)
5991 .handler(|_input: serde_json::Value| async move {
5992 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
5993 Ok(CallToolResult::text("done"))
5994 })
5995 .build();
5996
5997 let mut router = McpRouter::new().tool(slow_tool);
5998 init_router(&mut router).await;
5999
6000 let req = RouterRequest {
6002 id: RequestId::Number(1),
6003 inner: McpRequest::CallTool(CallToolParams {
6004 input_responses: None,
6005 request_state: None,
6006 name: "slow".to_string(),
6007 arguments: serde_json::json!({}),
6008 meta: None,
6009 task: Some(TaskRequestParams { ttl: None }),
6010 }),
6011 extensions: Extensions::new(),
6012 };
6013
6014 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6015 let task_id = match resp.inner {
6016 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6017 _ => panic!("Expected CreateTask response"),
6018 };
6019
6020 let req = RouterRequest {
6022 id: RequestId::Number(2),
6023 inner: McpRequest::CancelTask(CancelTaskParams {
6024 task_id: task_id.clone(),
6025 reason: Some("Test cancellation".to_string()),
6026 meta: None,
6027 }),
6028 extensions: Extensions::new(),
6029 };
6030
6031 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6032
6033 match resp.inner {
6035 Ok(McpResponse::CancelTask(EmptyResult {})) => {}
6036 other => panic!("Expected empty CancelTask ack, got {other:?}"),
6037 }
6038
6039 let req = RouterRequest {
6041 id: RequestId::Number(3),
6042 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6043 task_id: task_id.clone(),
6044 meta: None,
6045 }),
6046 extensions: Extensions::new(),
6047 };
6048 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6049 match resp.inner {
6050 Ok(McpResponse::GetTaskInfo(info)) => {
6051 assert_eq!(info.status, TaskStatus::Cancelled);
6052 }
6053 _ => panic!("Expected GetTaskInfo response"),
6054 }
6055 }
6056
6057 #[tokio::test]
6058 async fn test_get_task_info() {
6059 let add_tool = ToolBuilder::new("add")
6060 .description("Add two numbers")
6061 .task_support(TaskSupportMode::Optional)
6062 .handler(|input: AddInput| async move {
6063 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6064 })
6065 .build();
6066
6067 let mut router = McpRouter::new().tool(add_tool);
6068 init_router(&mut router).await;
6069
6070 let req = RouterRequest {
6072 id: RequestId::Number(1),
6073 inner: McpRequest::CallTool(CallToolParams {
6074 input_responses: None,
6075 request_state: None,
6076 name: "add".to_string(),
6077 arguments: serde_json::json!({"a": 1, "b": 2}),
6078 meta: None,
6079 task: Some(TaskRequestParams { ttl: Some(600_000) }),
6080 }),
6081 extensions: Extensions::new(),
6082 };
6083
6084 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6085 let task_id = match resp.inner {
6086 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6087 _ => panic!("Expected CreateTask response"),
6088 };
6089
6090 let req = RouterRequest {
6092 id: RequestId::Number(2),
6093 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6094 task_id: task_id.clone(),
6095 meta: None,
6096 }),
6097 extensions: Extensions::new(),
6098 };
6099
6100 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6101
6102 match resp.inner {
6103 Ok(McpResponse::GetTaskInfo(info)) => {
6104 assert_eq!(info.task_id, task_id);
6105 assert!(info.created_at.contains('T')); assert_eq!(info.ttl, Some(600_000));
6107 }
6108 _ => panic!("Expected GetTaskInfo response"),
6109 }
6110 }
6111
6112 #[tokio::test]
6113 async fn test_task_forbidden_tool_rejects_task_params() {
6114 let tool = ToolBuilder::new("sync_only")
6115 .description("Sync only tool")
6116 .handler(|_input: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
6117 .build();
6118
6119 let mut router = McpRouter::new().tool(tool);
6120 init_router(&mut router).await;
6121
6122 let req = RouterRequest {
6124 id: RequestId::Number(1),
6125 inner: McpRequest::CallTool(CallToolParams {
6126 input_responses: None,
6127 request_state: None,
6128 name: "sync_only".to_string(),
6129 arguments: serde_json::json!({}),
6130 meta: None,
6131 task: Some(TaskRequestParams { ttl: None }),
6132 }),
6133 extensions: Extensions::new(),
6134 };
6135
6136 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6137
6138 match resp.inner {
6139 Err(e) => {
6140 assert!(e.message.contains("does not support async tasks"));
6141 }
6142 _ => panic!("Expected error response"),
6143 }
6144 }
6145
6146 #[tokio::test]
6147 async fn test_get_nonexistent_task() {
6148 let mut router = McpRouter::new();
6149 init_router(&mut router).await;
6150
6151 let req = RouterRequest {
6152 id: RequestId::Number(1),
6153 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6154 task_id: "task-999".to_string(),
6155 meta: None,
6156 }),
6157 extensions: Extensions::new(),
6158 };
6159
6160 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6161
6162 match resp.inner {
6163 Err(e) => {
6164 assert!(e.message.contains("not found"));
6165 }
6166 _ => panic!("Expected error response"),
6167 }
6168 }
6169
6170 #[tokio::test]
6175 async fn test_subscribe_to_resource() {
6176 use crate::resource::ResourceBuilder;
6177
6178 let resource = ResourceBuilder::new("file:///test.txt")
6179 .name("Test File")
6180 .text("Hello");
6181
6182 let mut router = McpRouter::new().resource(resource);
6183 init_router(&mut router).await;
6184
6185 let req = RouterRequest {
6187 id: RequestId::Number(1),
6188 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6189 uri: "file:///test.txt".to_string(),
6190 meta: None,
6191 }),
6192 extensions: Extensions::new(),
6193 };
6194
6195 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6196
6197 match resp.inner {
6198 Ok(McpResponse::SubscribeResource(_)) => {
6199 assert!(router.is_subscribed("file:///test.txt"));
6201 }
6202 _ => panic!("Expected SubscribeResource response"),
6203 }
6204 }
6205
6206 #[tokio::test]
6207 async fn test_unsubscribe_from_resource() {
6208 use crate::resource::ResourceBuilder;
6209
6210 let resource = ResourceBuilder::new("file:///test.txt")
6211 .name("Test File")
6212 .text("Hello");
6213
6214 let mut router = McpRouter::new().resource(resource);
6215 init_router(&mut router).await;
6216
6217 let req = RouterRequest {
6219 id: RequestId::Number(1),
6220 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6221 uri: "file:///test.txt".to_string(),
6222 meta: None,
6223 }),
6224 extensions: Extensions::new(),
6225 };
6226 let _ = router.ready().await.unwrap().call(req).await.unwrap();
6227 assert!(router.is_subscribed("file:///test.txt"));
6228
6229 let req = RouterRequest {
6231 id: RequestId::Number(2),
6232 inner: McpRequest::UnsubscribeResource(UnsubscribeResourceParams {
6233 uri: "file:///test.txt".to_string(),
6234 meta: None,
6235 }),
6236 extensions: Extensions::new(),
6237 };
6238
6239 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6240
6241 match resp.inner {
6242 Ok(McpResponse::UnsubscribeResource(_)) => {
6243 assert!(!router.is_subscribed("file:///test.txt"));
6245 }
6246 _ => panic!("Expected UnsubscribeResource response"),
6247 }
6248 }
6249
6250 #[tokio::test]
6251 async fn test_subscribe_nonexistent_resource() {
6252 let mut router = McpRouter::new();
6253 init_router(&mut router).await;
6254
6255 let req = RouterRequest {
6256 id: RequestId::Number(1),
6257 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6258 uri: "file:///nonexistent.txt".to_string(),
6259 meta: None,
6260 }),
6261 extensions: Extensions::new(),
6262 };
6263
6264 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6265
6266 match resp.inner {
6267 Err(e) => {
6268 assert!(e.message.contains("not found"));
6269 }
6270 _ => panic!("Expected error response"),
6271 }
6272 }
6273
6274 #[tokio::test]
6275 async fn test_notify_resource_updated() {
6276 use crate::context::notification_channel;
6277 use crate::resource::ResourceBuilder;
6278
6279 let (tx, mut rx) = notification_channel(10);
6280
6281 let resource = ResourceBuilder::new("file:///test.txt")
6282 .name("Test File")
6283 .text("Hello");
6284
6285 let router = McpRouter::new()
6286 .resource(resource)
6287 .with_notification_sender(tx);
6288
6289 router.subscribe("file:///test.txt");
6291
6292 let sent = router.notify_resource_updated("file:///test.txt");
6294 assert!(sent);
6295
6296 let notification = rx.try_recv().unwrap();
6298 match notification {
6299 ServerNotification::ResourceUpdated { uri } => {
6300 assert_eq!(uri, "file:///test.txt");
6301 }
6302 _ => panic!("Expected ResourceUpdated notification"),
6303 }
6304 }
6305
6306 #[tokio::test]
6307 async fn test_notify_resource_updated_not_subscribed() {
6308 use crate::context::notification_channel;
6309 use crate::resource::ResourceBuilder;
6310
6311 let (tx, mut rx) = notification_channel(10);
6312
6313 let resource = ResourceBuilder::new("file:///test.txt")
6314 .name("Test File")
6315 .text("Hello");
6316
6317 let router = McpRouter::new()
6318 .resource(resource)
6319 .with_notification_sender(tx);
6320
6321 let sent = router.notify_resource_updated("file:///test.txt");
6323 assert!(!sent); assert!(rx.try_recv().is_err());
6327 }
6328
6329 #[tokio::test]
6330 async fn test_notify_resources_list_changed() {
6331 use crate::context::notification_channel;
6332
6333 let (tx, mut rx) = notification_channel(10);
6334 let router = McpRouter::new().with_notification_sender(tx);
6335
6336 let sent = router.notify_resources_list_changed();
6337 assert!(sent);
6338
6339 let notification = rx.try_recv().unwrap();
6340 match notification {
6341 ServerNotification::ResourcesListChanged => {}
6342 _ => panic!("Expected ResourcesListChanged notification"),
6343 }
6344 }
6345
6346 #[tokio::test]
6347 async fn test_subscribed_uris() {
6348 use crate::resource::ResourceBuilder;
6349
6350 let resource1 = ResourceBuilder::new("file:///a.txt").name("A").text("A");
6351
6352 let resource2 = ResourceBuilder::new("file:///b.txt").name("B").text("B");
6353
6354 let router = McpRouter::new().resource(resource1).resource(resource2);
6355
6356 router.subscribe("file:///a.txt");
6358 router.subscribe("file:///b.txt");
6359
6360 let uris = router.subscribed_uris();
6361 assert_eq!(uris.len(), 2);
6362 assert!(uris.contains(&"file:///a.txt".to_string()));
6363 assert!(uris.contains(&"file:///b.txt".to_string()));
6364 }
6365
6366 #[tokio::test]
6367 async fn test_subscription_capability_advertised() {
6368 use crate::resource::ResourceBuilder;
6369
6370 let resource = ResourceBuilder::new("file:///test.txt")
6371 .name("Test")
6372 .text("Hello");
6373
6374 let mut router = McpRouter::new().resource(resource);
6375
6376 let init_req = RouterRequest {
6378 id: RequestId::Number(0),
6379 inner: McpRequest::Initialize(InitializeParams {
6380 protocol_version: "2025-11-25".to_string(),
6381 capabilities: ClientCapabilities {
6382 roots: None,
6383 sampling: None,
6384 elicitation: None,
6385 tasks: None,
6386 experimental: None,
6387 extensions: None,
6388 },
6389 client_info: Implementation {
6390 name: "test".to_string(),
6391 version: "1.0".to_string(),
6392 ..Default::default()
6393 },
6394 meta: None,
6395 }),
6396 extensions: Extensions::new(),
6397 };
6398 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
6399
6400 match resp.inner {
6401 Ok(McpResponse::Initialize(result)) => {
6402 let resources_cap = result.capabilities.resources.unwrap();
6404 assert!(resources_cap.subscribe);
6405 }
6406 _ => panic!("Expected Initialize response"),
6407 }
6408 }
6409
6410 #[tokio::test]
6411 async fn test_completion_handler() {
6412 let router = McpRouter::new()
6413 .server_info("test", "1.0")
6414 .completion_handler(|params: CompleteParams| async move {
6415 let prefix = ¶ms.argument.value;
6417 let suggestions: Vec<String> = vec!["alpha", "beta", "gamma"]
6418 .into_iter()
6419 .filter(|s| s.starts_with(prefix))
6420 .map(String::from)
6421 .collect();
6422 Ok(CompleteResult::new(suggestions))
6423 });
6424
6425 let init_req = RouterRequest {
6427 id: RequestId::Number(0),
6428 inner: McpRequest::Initialize(InitializeParams {
6429 protocol_version: "2025-11-25".to_string(),
6430 capabilities: ClientCapabilities::default(),
6431 client_info: Implementation {
6432 name: "test".to_string(),
6433 version: "1.0".to_string(),
6434 ..Default::default()
6435 },
6436 meta: None,
6437 }),
6438 extensions: Extensions::new(),
6439 };
6440 let resp = router
6441 .clone()
6442 .ready()
6443 .await
6444 .unwrap()
6445 .call(init_req)
6446 .await
6447 .unwrap();
6448
6449 match resp.inner {
6451 Ok(McpResponse::Initialize(result)) => {
6452 assert!(result.capabilities.completions.is_some());
6453 }
6454 _ => panic!("Expected Initialize response"),
6455 }
6456
6457 router.handle_notification(McpNotification::Initialized);
6459
6460 let complete_req = RouterRequest {
6462 id: RequestId::Number(1),
6463 inner: McpRequest::Complete(CompleteParams {
6464 reference: CompletionReference::prompt("test-prompt"),
6465 argument: CompletionArgument::new("query", "al"),
6466 context: None,
6467 meta: None,
6468 }),
6469 extensions: Extensions::new(),
6470 };
6471 let resp = router
6472 .clone()
6473 .ready()
6474 .await
6475 .unwrap()
6476 .call(complete_req)
6477 .await
6478 .unwrap();
6479
6480 match resp.inner {
6481 Ok(McpResponse::Complete(result)) => {
6482 assert_eq!(result.completion.values, vec!["alpha"]);
6483 }
6484 _ => panic!("Expected Complete response"),
6485 }
6486 }
6487
6488 #[tokio::test]
6489 async fn test_completion_without_handler_returns_empty() {
6490 let router = McpRouter::new().server_info("test", "1.0");
6491
6492 let init_req = RouterRequest {
6494 id: RequestId::Number(0),
6495 inner: McpRequest::Initialize(InitializeParams {
6496 protocol_version: "2025-11-25".to_string(),
6497 capabilities: ClientCapabilities::default(),
6498 client_info: Implementation {
6499 name: "test".to_string(),
6500 version: "1.0".to_string(),
6501 ..Default::default()
6502 },
6503 meta: None,
6504 }),
6505 extensions: Extensions::new(),
6506 };
6507 let resp = router
6508 .clone()
6509 .ready()
6510 .await
6511 .unwrap()
6512 .call(init_req)
6513 .await
6514 .unwrap();
6515
6516 match resp.inner {
6518 Ok(McpResponse::Initialize(result)) => {
6519 assert!(result.capabilities.completions.is_none());
6520 }
6521 _ => panic!("Expected Initialize response"),
6522 }
6523
6524 router.handle_notification(McpNotification::Initialized);
6526
6527 let complete_req = RouterRequest {
6529 id: RequestId::Number(1),
6530 inner: McpRequest::Complete(CompleteParams {
6531 reference: CompletionReference::prompt("test-prompt"),
6532 argument: CompletionArgument::new("query", "al"),
6533 context: None,
6534 meta: None,
6535 }),
6536 extensions: Extensions::new(),
6537 };
6538 let resp = router
6539 .clone()
6540 .ready()
6541 .await
6542 .unwrap()
6543 .call(complete_req)
6544 .await
6545 .unwrap();
6546
6547 match resp.inner {
6548 Ok(McpResponse::Complete(result)) => {
6549 assert!(result.completion.values.is_empty());
6550 }
6551 _ => panic!("Expected Complete response"),
6552 }
6553 }
6554
6555 #[tokio::test]
6556 async fn test_tool_filter_list() {
6557 use crate::filter::CapabilityFilter;
6558 use crate::tool::Tool;
6559
6560 let public_tool = ToolBuilder::new("public")
6561 .description("Public tool")
6562 .handler(|_: AddInput| async move { Ok(CallToolResult::text("public")) })
6563 .build();
6564
6565 let admin_tool = ToolBuilder::new("admin")
6566 .description("Admin tool")
6567 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6568 .build();
6569
6570 let mut router = McpRouter::new()
6571 .tool(public_tool)
6572 .tool(admin_tool)
6573 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| tool.name != "admin"));
6574
6575 init_router(&mut router).await;
6577
6578 let req = RouterRequest {
6579 id: RequestId::Number(1),
6580 inner: McpRequest::ListTools(ListToolsParams::default()),
6581 extensions: Extensions::new(),
6582 };
6583
6584 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6585
6586 match resp.inner {
6587 Ok(McpResponse::ListTools(result)) => {
6588 assert_eq!(result.tools.len(), 1);
6590 assert_eq!(result.tools[0].name, "public");
6591 }
6592 _ => panic!("Expected ListTools response"),
6593 }
6594 }
6595
6596 #[tokio::test]
6597 async fn test_tool_filter_call_denied() {
6598 use crate::filter::CapabilityFilter;
6599 use crate::tool::Tool;
6600
6601 let admin_tool = ToolBuilder::new("admin")
6602 .description("Admin tool")
6603 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6604 .build();
6605
6606 let mut router = McpRouter::new()
6607 .tool(admin_tool)
6608 .tool_filter(CapabilityFilter::new(|_, _: &Tool| false)); init_router(&mut router).await;
6612
6613 let req = RouterRequest {
6614 id: RequestId::Number(1),
6615 inner: McpRequest::CallTool(CallToolParams {
6616 input_responses: None,
6617 request_state: None,
6618 name: "admin".to_string(),
6619 arguments: serde_json::json!({"a": 1, "b": 2}),
6620 meta: None,
6621 task: None,
6622 }),
6623 extensions: Extensions::new(),
6624 };
6625
6626 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6627
6628 match resp.inner {
6630 Err(e) => {
6631 assert_eq!(e.code, -32601); }
6633 _ => panic!("Expected JsonRpc error"),
6634 }
6635 }
6636
6637 #[tokio::test]
6638 async fn test_tool_filter_call_allowed() {
6639 use crate::filter::CapabilityFilter;
6640 use crate::tool::Tool;
6641
6642 let public_tool = ToolBuilder::new("public")
6643 .description("Public tool")
6644 .handler(|input: AddInput| async move {
6645 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6646 })
6647 .build();
6648
6649 let mut router = McpRouter::new()
6650 .tool(public_tool)
6651 .tool_filter(CapabilityFilter::new(|_, _: &Tool| true)); init_router(&mut router).await;
6655
6656 let req = RouterRequest {
6657 id: RequestId::Number(1),
6658 inner: McpRequest::CallTool(CallToolParams {
6659 input_responses: None,
6660 request_state: None,
6661 name: "public".to_string(),
6662 arguments: serde_json::json!({"a": 1, "b": 2}),
6663 meta: None,
6664 task: None,
6665 }),
6666 extensions: Extensions::new(),
6667 };
6668
6669 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6670
6671 match resp.inner {
6672 Ok(McpResponse::CallTool(result)) => {
6673 assert!(!result.is_error);
6674 }
6675 _ => panic!("Expected CallTool response"),
6676 }
6677 }
6678
6679 #[tokio::test]
6680 async fn test_tool_filter_custom_denial() {
6681 use crate::filter::{CapabilityFilter, DenialBehavior};
6682 use crate::tool::Tool;
6683
6684 let admin_tool = ToolBuilder::new("admin")
6685 .description("Admin tool")
6686 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6687 .build();
6688
6689 let mut router = McpRouter::new().tool(admin_tool).tool_filter(
6690 CapabilityFilter::new(|_, _: &Tool| false)
6691 .denial_behavior(DenialBehavior::Unauthorized),
6692 );
6693
6694 init_router(&mut router).await;
6696
6697 let req = RouterRequest {
6698 id: RequestId::Number(1),
6699 inner: McpRequest::CallTool(CallToolParams {
6700 input_responses: None,
6701 request_state: None,
6702 name: "admin".to_string(),
6703 arguments: serde_json::json!({"a": 1, "b": 2}),
6704 meta: None,
6705 task: None,
6706 }),
6707 extensions: Extensions::new(),
6708 };
6709
6710 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6711
6712 match resp.inner {
6714 Err(e) => {
6715 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
6717 }
6718 _ => panic!("Expected JsonRpc error"),
6719 }
6720 }
6721
6722 #[tokio::test]
6723 async fn test_resource_filter_list() {
6724 use crate::filter::CapabilityFilter;
6725 use crate::resource::{Resource, ResourceBuilder};
6726
6727 let public_resource = ResourceBuilder::new("file:///public.txt")
6728 .name("Public File")
6729 .text("public content");
6730
6731 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6732 .name("Secret File")
6733 .text("secret content");
6734
6735 let mut router = McpRouter::new()
6736 .resource(public_resource)
6737 .resource(secret_resource)
6738 .resource_filter(CapabilityFilter::new(|_, r: &Resource| {
6739 !r.name.contains("Secret")
6740 }));
6741
6742 init_router(&mut router).await;
6744
6745 let req = RouterRequest {
6746 id: RequestId::Number(1),
6747 inner: McpRequest::ListResources(ListResourcesParams::default()),
6748 extensions: Extensions::new(),
6749 };
6750
6751 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6752
6753 match resp.inner {
6754 Ok(McpResponse::ListResources(result)) => {
6755 assert_eq!(result.resources.len(), 1);
6757 assert_eq!(result.resources[0].name, "Public File");
6758 }
6759 _ => panic!("Expected ListResources response"),
6760 }
6761 }
6762
6763 #[tokio::test]
6764 async fn test_resource_filter_read_denied() {
6765 use crate::filter::CapabilityFilter;
6766 use crate::resource::{Resource, ResourceBuilder};
6767
6768 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6769 .name("Secret File")
6770 .text("secret content");
6771
6772 let mut router = McpRouter::new()
6773 .resource(secret_resource)
6774 .resource_filter(CapabilityFilter::new(|_, _: &Resource| false)); init_router(&mut router).await;
6778
6779 let req = RouterRequest {
6780 id: RequestId::Number(1),
6781 inner: McpRequest::ReadResource(ReadResourceParams {
6782 input_responses: None,
6783 request_state: None,
6784 uri: "file:///secret.txt".to_string(),
6785 meta: None,
6786 }),
6787 extensions: Extensions::new(),
6788 };
6789
6790 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6791
6792 match resp.inner {
6794 Err(e) => {
6795 assert_eq!(e.code, -32601); }
6797 _ => panic!("Expected JsonRpc error"),
6798 }
6799 }
6800
6801 #[tokio::test]
6802 async fn test_resource_filter_read_allowed() {
6803 use crate::filter::CapabilityFilter;
6804 use crate::resource::{Resource, ResourceBuilder};
6805
6806 let public_resource = ResourceBuilder::new("file:///public.txt")
6807 .name("Public File")
6808 .text("public content");
6809
6810 let mut router = McpRouter::new()
6811 .resource(public_resource)
6812 .resource_filter(CapabilityFilter::new(|_, _: &Resource| true)); init_router(&mut router).await;
6816
6817 let req = RouterRequest {
6818 id: RequestId::Number(1),
6819 inner: McpRequest::ReadResource(ReadResourceParams {
6820 input_responses: None,
6821 request_state: None,
6822 uri: "file:///public.txt".to_string(),
6823 meta: None,
6824 }),
6825 extensions: Extensions::new(),
6826 };
6827
6828 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6829
6830 match resp.inner {
6831 Ok(McpResponse::ReadResource(result)) => {
6832 assert_eq!(result.contents.len(), 1);
6833 assert_eq!(result.contents[0].text.as_deref(), Some("public content"));
6834 }
6835 _ => panic!("Expected ReadResource response"),
6836 }
6837 }
6838
6839 #[tokio::test]
6840 async fn test_resource_filter_custom_denial() {
6841 use crate::filter::{CapabilityFilter, DenialBehavior};
6842 use crate::resource::{Resource, ResourceBuilder};
6843
6844 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6845 .name("Secret File")
6846 .text("secret content");
6847
6848 let mut router = McpRouter::new().resource(secret_resource).resource_filter(
6849 CapabilityFilter::new(|_, _: &Resource| false)
6850 .denial_behavior(DenialBehavior::Unauthorized),
6851 );
6852
6853 init_router(&mut router).await;
6855
6856 let req = RouterRequest {
6857 id: RequestId::Number(1),
6858 inner: McpRequest::ReadResource(ReadResourceParams {
6859 input_responses: None,
6860 request_state: None,
6861 uri: "file:///secret.txt".to_string(),
6862 meta: None,
6863 }),
6864 extensions: Extensions::new(),
6865 };
6866
6867 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6868
6869 match resp.inner {
6871 Err(e) => {
6872 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
6874 }
6875 _ => panic!("Expected JsonRpc error"),
6876 }
6877 }
6878
6879 #[tokio::test]
6880 async fn test_prompt_filter_list() {
6881 use crate::filter::CapabilityFilter;
6882 use crate::prompt::{Prompt, PromptBuilder};
6883
6884 let public_prompt = PromptBuilder::new("greeting")
6885 .description("A greeting")
6886 .user_message("Hello!");
6887
6888 let admin_prompt = PromptBuilder::new("system_debug")
6889 .description("Admin prompt")
6890 .user_message("Debug");
6891
6892 let mut router = McpRouter::new()
6893 .prompt(public_prompt)
6894 .prompt(admin_prompt)
6895 .prompt_filter(CapabilityFilter::new(|_, p: &Prompt| {
6896 !p.name.contains("system")
6897 }));
6898
6899 init_router(&mut router).await;
6901
6902 let req = RouterRequest {
6903 id: RequestId::Number(1),
6904 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
6905 extensions: Extensions::new(),
6906 };
6907
6908 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6909
6910 match resp.inner {
6911 Ok(McpResponse::ListPrompts(result)) => {
6912 assert_eq!(result.prompts.len(), 1);
6914 assert_eq!(result.prompts[0].name, "greeting");
6915 }
6916 _ => panic!("Expected ListPrompts response"),
6917 }
6918 }
6919
6920 #[tokio::test]
6921 async fn test_prompt_filter_get_denied() {
6922 use crate::filter::CapabilityFilter;
6923 use crate::prompt::{Prompt, PromptBuilder};
6924 use std::collections::HashMap;
6925
6926 let admin_prompt = PromptBuilder::new("system_debug")
6927 .description("Admin prompt")
6928 .user_message("Debug");
6929
6930 let mut router = McpRouter::new()
6931 .prompt(admin_prompt)
6932 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| false)); init_router(&mut router).await;
6936
6937 let req = RouterRequest {
6938 id: RequestId::Number(1),
6939 inner: McpRequest::GetPrompt(GetPromptParams {
6940 input_responses: None,
6941 request_state: None,
6942 name: "system_debug".to_string(),
6943 arguments: HashMap::new(),
6944 meta: None,
6945 }),
6946 extensions: Extensions::new(),
6947 };
6948
6949 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6950
6951 match resp.inner {
6953 Err(e) => {
6954 assert_eq!(e.code, -32601); }
6956 _ => panic!("Expected JsonRpc error"),
6957 }
6958 }
6959
6960 #[tokio::test]
6961 async fn test_prompt_filter_get_allowed() {
6962 use crate::filter::CapabilityFilter;
6963 use crate::prompt::{Prompt, PromptBuilder};
6964 use std::collections::HashMap;
6965
6966 let public_prompt = PromptBuilder::new("greeting")
6967 .description("A greeting")
6968 .user_message("Hello!");
6969
6970 let mut router = McpRouter::new()
6971 .prompt(public_prompt)
6972 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| true)); init_router(&mut router).await;
6976
6977 let req = RouterRequest {
6978 id: RequestId::Number(1),
6979 inner: McpRequest::GetPrompt(GetPromptParams {
6980 input_responses: None,
6981 request_state: None,
6982 name: "greeting".to_string(),
6983 arguments: HashMap::new(),
6984 meta: None,
6985 }),
6986 extensions: Extensions::new(),
6987 };
6988
6989 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6990
6991 match resp.inner {
6992 Ok(McpResponse::GetPrompt(result)) => {
6993 assert_eq!(result.messages.len(), 1);
6994 }
6995 _ => panic!("Expected GetPrompt response"),
6996 }
6997 }
6998
6999 #[tokio::test]
7000 async fn test_prompt_filter_custom_denial() {
7001 use crate::filter::{CapabilityFilter, DenialBehavior};
7002 use crate::prompt::{Prompt, PromptBuilder};
7003 use std::collections::HashMap;
7004
7005 let admin_prompt = PromptBuilder::new("system_debug")
7006 .description("Admin prompt")
7007 .user_message("Debug");
7008
7009 let mut router = McpRouter::new().prompt(admin_prompt).prompt_filter(
7010 CapabilityFilter::new(|_, _: &Prompt| false)
7011 .denial_behavior(DenialBehavior::Unauthorized),
7012 );
7013
7014 init_router(&mut router).await;
7016
7017 let req = RouterRequest {
7018 id: RequestId::Number(1),
7019 inner: McpRequest::GetPrompt(GetPromptParams {
7020 input_responses: None,
7021 request_state: None,
7022 name: "system_debug".to_string(),
7023 arguments: HashMap::new(),
7024 meta: None,
7025 }),
7026 extensions: Extensions::new(),
7027 };
7028
7029 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7030
7031 match resp.inner {
7033 Err(e) => {
7034 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7036 }
7037 _ => panic!("Expected JsonRpc error"),
7038 }
7039 }
7040
7041 #[derive(Debug, Deserialize, JsonSchema)]
7046 struct StringInput {
7047 value: String,
7048 }
7049
7050 #[tokio::test]
7051 async fn test_router_merge_tools() {
7052 let tool_a = ToolBuilder::new("tool_a")
7054 .description("Tool A")
7055 .handler(|_: StringInput| async move { Ok(CallToolResult::text("A")) })
7056 .build();
7057
7058 let router_a = McpRouter::new().tool(tool_a);
7059
7060 let tool_b = ToolBuilder::new("tool_b")
7062 .description("Tool B")
7063 .handler(|_: StringInput| async move { Ok(CallToolResult::text("B")) })
7064 .build();
7065 let tool_c = ToolBuilder::new("tool_c")
7066 .description("Tool C")
7067 .handler(|_: StringInput| async move { Ok(CallToolResult::text("C")) })
7068 .build();
7069
7070 let router_b = McpRouter::new().tool(tool_b).tool(tool_c);
7071
7072 let mut merged = McpRouter::new()
7074 .server_info("merged", "1.0")
7075 .merge(router_a)
7076 .merge(router_b);
7077
7078 init_router(&mut merged).await;
7079
7080 let req = RouterRequest {
7082 id: RequestId::Number(1),
7083 inner: McpRequest::ListTools(ListToolsParams::default()),
7084 extensions: Extensions::new(),
7085 };
7086
7087 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7088
7089 match resp.inner {
7090 Ok(McpResponse::ListTools(result)) => {
7091 assert_eq!(result.tools.len(), 3);
7092 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7093 assert!(names.contains(&"tool_a"));
7094 assert!(names.contains(&"tool_b"));
7095 assert!(names.contains(&"tool_c"));
7096 }
7097 _ => panic!("Expected ListTools response"),
7098 }
7099 }
7100
7101 #[tokio::test]
7102 async fn test_router_merge_overwrites_duplicates() {
7103 let tool_v1 = ToolBuilder::new("shared")
7105 .description("Version 1")
7106 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v1")) })
7107 .build();
7108
7109 let router_a = McpRouter::new().tool(tool_v1);
7110
7111 let tool_v2 = ToolBuilder::new("shared")
7113 .description("Version 2")
7114 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v2")) })
7115 .build();
7116
7117 let router_b = McpRouter::new().tool(tool_v2);
7118
7119 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7121
7122 init_router(&mut merged).await;
7123
7124 let req = RouterRequest {
7125 id: RequestId::Number(1),
7126 inner: McpRequest::ListTools(ListToolsParams::default()),
7127 extensions: Extensions::new(),
7128 };
7129
7130 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7131
7132 match resp.inner {
7133 Ok(McpResponse::ListTools(result)) => {
7134 assert_eq!(result.tools.len(), 1);
7135 assert_eq!(result.tools[0].name, "shared");
7136 assert_eq!(result.tools[0].description.as_deref(), Some("Version 2"));
7137 }
7138 _ => panic!("Expected ListTools response"),
7139 }
7140 }
7141
7142 #[tokio::test]
7143 async fn test_router_merge_resources() {
7144 use crate::resource::ResourceBuilder;
7145
7146 let router_a = McpRouter::new().resource(
7148 ResourceBuilder::new("file:///a.txt")
7149 .name("File A")
7150 .text("content a"),
7151 );
7152
7153 let router_b = McpRouter::new().resource(
7154 ResourceBuilder::new("file:///b.txt")
7155 .name("File B")
7156 .text("content b"),
7157 );
7158
7159 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7160
7161 init_router(&mut merged).await;
7162
7163 let req = RouterRequest {
7164 id: RequestId::Number(1),
7165 inner: McpRequest::ListResources(ListResourcesParams::default()),
7166 extensions: Extensions::new(),
7167 };
7168
7169 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7170
7171 match resp.inner {
7172 Ok(McpResponse::ListResources(result)) => {
7173 assert_eq!(result.resources.len(), 2);
7174 let uris: Vec<&str> = result.resources.iter().map(|r| r.uri.as_str()).collect();
7175 assert!(uris.contains(&"file:///a.txt"));
7176 assert!(uris.contains(&"file:///b.txt"));
7177 }
7178 _ => panic!("Expected ListResources response"),
7179 }
7180 }
7181
7182 #[tokio::test]
7183 async fn test_router_merge_prompts() {
7184 use crate::prompt::PromptBuilder;
7185
7186 let router_a =
7187 McpRouter::new().prompt(PromptBuilder::new("prompt_a").user_message("Hello A"));
7188
7189 let router_b =
7190 McpRouter::new().prompt(PromptBuilder::new("prompt_b").user_message("Hello B"));
7191
7192 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7193
7194 init_router(&mut merged).await;
7195
7196 let req = RouterRequest {
7197 id: RequestId::Number(1),
7198 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7199 extensions: Extensions::new(),
7200 };
7201
7202 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7203
7204 match resp.inner {
7205 Ok(McpResponse::ListPrompts(result)) => {
7206 assert_eq!(result.prompts.len(), 2);
7207 let names: Vec<&str> = result.prompts.iter().map(|p| p.name.as_str()).collect();
7208 assert!(names.contains(&"prompt_a"));
7209 assert!(names.contains(&"prompt_b"));
7210 }
7211 _ => panic!("Expected ListPrompts response"),
7212 }
7213 }
7214
7215 #[tokio::test]
7216 async fn test_router_nest_prefixes_tools() {
7217 let tool_query = ToolBuilder::new("query")
7219 .description("Query the database")
7220 .handler(|_: StringInput| async move { Ok(CallToolResult::text("query result")) })
7221 .build();
7222 let tool_insert = ToolBuilder::new("insert")
7223 .description("Insert into database")
7224 .handler(|_: StringInput| async move { Ok(CallToolResult::text("insert result")) })
7225 .build();
7226
7227 let db_router = McpRouter::new().tool(tool_query).tool(tool_insert);
7228
7229 let mut router = McpRouter::new()
7231 .server_info("nested", "1.0")
7232 .nest("db", db_router);
7233
7234 init_router(&mut router).await;
7235
7236 let req = RouterRequest {
7237 id: RequestId::Number(1),
7238 inner: McpRequest::ListTools(ListToolsParams::default()),
7239 extensions: Extensions::new(),
7240 };
7241
7242 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7243
7244 match resp.inner {
7245 Ok(McpResponse::ListTools(result)) => {
7246 assert_eq!(result.tools.len(), 2);
7247 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7248 assert!(names.contains(&"db.query"));
7249 assert!(names.contains(&"db.insert"));
7250 }
7251 _ => panic!("Expected ListTools response"),
7252 }
7253 }
7254
7255 #[tokio::test]
7256 async fn test_router_nest_call_prefixed_tool() {
7257 let tool = ToolBuilder::new("echo")
7258 .description("Echo input")
7259 .handler(|input: StringInput| async move { Ok(CallToolResult::text(&input.value)) })
7260 .build();
7261
7262 let nested_router = McpRouter::new().tool(tool);
7263
7264 let mut router = McpRouter::new().nest("api", nested_router);
7265
7266 init_router(&mut router).await;
7267
7268 let req = RouterRequest {
7270 id: RequestId::Number(1),
7271 inner: McpRequest::CallTool(CallToolParams {
7272 input_responses: None,
7273 request_state: None,
7274 name: "api.echo".to_string(),
7275 arguments: serde_json::json!({"value": "hello world"}),
7276 meta: None,
7277 task: None,
7278 }),
7279 extensions: Extensions::new(),
7280 };
7281
7282 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7283
7284 match resp.inner {
7285 Ok(McpResponse::CallTool(result)) => {
7286 assert!(!result.is_error);
7287 match &result.content[0] {
7288 Content::Text { text, .. } => assert_eq!(text, "hello world"),
7289 _ => panic!("Expected text content"),
7290 }
7291 }
7292 _ => panic!("Expected CallTool response"),
7293 }
7294 }
7295
7296 #[tokio::test]
7297 async fn test_router_multiple_nests() {
7298 let db_tool = ToolBuilder::new("query")
7299 .description("Database query")
7300 .handler(|_: StringInput| async move { Ok(CallToolResult::text("db")) })
7301 .build();
7302
7303 let api_tool = ToolBuilder::new("fetch")
7304 .description("API fetch")
7305 .handler(|_: StringInput| async move { Ok(CallToolResult::text("api")) })
7306 .build();
7307
7308 let db_router = McpRouter::new().tool(db_tool);
7309 let api_router = McpRouter::new().tool(api_tool);
7310
7311 let mut router = McpRouter::new()
7312 .nest("db", db_router)
7313 .nest("api", api_router);
7314
7315 init_router(&mut router).await;
7316
7317 let req = RouterRequest {
7318 id: RequestId::Number(1),
7319 inner: McpRequest::ListTools(ListToolsParams::default()),
7320 extensions: Extensions::new(),
7321 };
7322
7323 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7324
7325 match resp.inner {
7326 Ok(McpResponse::ListTools(result)) => {
7327 assert_eq!(result.tools.len(), 2);
7328 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7329 assert!(names.contains(&"db.query"));
7330 assert!(names.contains(&"api.fetch"));
7331 }
7332 _ => panic!("Expected ListTools response"),
7333 }
7334 }
7335
7336 #[tokio::test]
7337 async fn test_router_merge_and_nest_combined() {
7338 let tool_a = ToolBuilder::new("local")
7340 .description("Local tool")
7341 .handler(|_: StringInput| async move { Ok(CallToolResult::text("local")) })
7342 .build();
7343
7344 let nested_tool = ToolBuilder::new("remote")
7345 .description("Remote tool")
7346 .handler(|_: StringInput| async move { Ok(CallToolResult::text("remote")) })
7347 .build();
7348
7349 let nested_router = McpRouter::new().tool(nested_tool);
7350
7351 let mut router = McpRouter::new()
7352 .tool(tool_a)
7353 .nest("external", nested_router);
7354
7355 init_router(&mut router).await;
7356
7357 let req = RouterRequest {
7358 id: RequestId::Number(1),
7359 inner: McpRequest::ListTools(ListToolsParams::default()),
7360 extensions: Extensions::new(),
7361 };
7362
7363 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7364
7365 match resp.inner {
7366 Ok(McpResponse::ListTools(result)) => {
7367 assert_eq!(result.tools.len(), 2);
7368 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7369 assert!(names.contains(&"local"));
7370 assert!(names.contains(&"external.remote"));
7371 }
7372 _ => panic!("Expected ListTools response"),
7373 }
7374 }
7375
7376 #[tokio::test]
7377 async fn test_router_merge_preserves_server_info() {
7378 let child_router = McpRouter::new()
7379 .server_info("child", "2.0")
7380 .instructions("Child instructions");
7381
7382 let mut router = McpRouter::new()
7383 .server_info("parent", "1.0")
7384 .instructions("Parent instructions")
7385 .merge(child_router);
7386
7387 init_router(&mut router).await;
7388
7389 let init_req = RouterRequest {
7391 id: RequestId::Number(99),
7392 inner: McpRequest::Initialize(InitializeParams {
7393 protocol_version: "2025-11-25".to_string(),
7394 capabilities: ClientCapabilities::default(),
7395 client_info: Implementation {
7396 name: "test".to_string(),
7397 version: "1.0".to_string(),
7398 ..Default::default()
7399 },
7400 meta: None,
7401 }),
7402 extensions: Extensions::new(),
7403 };
7404
7405 let child_router2 = McpRouter::new().server_info("child", "2.0");
7407 let mut fresh_router = McpRouter::new()
7408 .server_info("parent", "1.0")
7409 .merge(child_router2);
7410
7411 let resp = fresh_router
7412 .ready()
7413 .await
7414 .unwrap()
7415 .call(init_req)
7416 .await
7417 .unwrap();
7418
7419 match resp.inner {
7420 Ok(McpResponse::Initialize(result)) => {
7421 assert_eq!(result.server_info.name, "parent");
7422 assert_eq!(result.server_info.version, "1.0");
7423 }
7424 _ => panic!("Expected Initialize response"),
7425 }
7426 }
7427
7428 #[tokio::test]
7433 async fn test_auto_instructions_tools_only() {
7434 let tool_a = ToolBuilder::new("alpha")
7435 .description("Alpha tool")
7436 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7437 .build();
7438 let tool_b = ToolBuilder::new("beta")
7439 .description("Beta tool")
7440 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7441 .build();
7442
7443 let mut router = McpRouter::new()
7444 .auto_instructions()
7445 .tool(tool_a)
7446 .tool(tool_b);
7447
7448 let resp = send_initialize(&mut router).await;
7449 let instructions = resp.instructions.expect("should have instructions");
7450
7451 assert!(instructions.contains("## Tools"));
7452 assert!(instructions.contains("- **alpha**: Alpha tool"));
7453 assert!(instructions.contains("- **beta**: Beta tool"));
7454 assert!(!instructions.contains("## Resources"));
7456 assert!(!instructions.contains("## Prompts"));
7457 }
7458
7459 #[tokio::test]
7460 async fn test_auto_instructions_with_annotations() {
7461 let read_only_tool = ToolBuilder::new("query")
7462 .description("Run a query")
7463 .read_only()
7464 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7465 .build();
7466 let destructive_tool = ToolBuilder::new("delete")
7467 .description("Delete a record")
7468 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7469 .build();
7470 let idempotent_tool = ToolBuilder::new("upsert")
7471 .description("Upsert a record")
7472 .non_destructive()
7473 .idempotent()
7474 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7475 .build();
7476
7477 let mut router = McpRouter::new()
7478 .auto_instructions()
7479 .tool(read_only_tool)
7480 .tool(destructive_tool)
7481 .tool(idempotent_tool);
7482
7483 let resp = send_initialize(&mut router).await;
7484 let instructions = resp.instructions.unwrap();
7485
7486 assert!(instructions.contains("- **query**: Run a query [read-only]"));
7487 assert!(instructions.contains("- **delete**: Delete a record\n"));
7489 assert!(instructions.contains("- **upsert**: Upsert a record [idempotent]"));
7490 }
7491
7492 #[tokio::test]
7493 async fn test_auto_instructions_with_resources() {
7494 use crate::resource::ResourceBuilder;
7495
7496 let resource = ResourceBuilder::new("file:///schema.sql")
7497 .name("Schema")
7498 .description("Database schema")
7499 .text("CREATE TABLE ...");
7500
7501 let mut router = McpRouter::new().auto_instructions().resource(resource);
7502
7503 let resp = send_initialize(&mut router).await;
7504 let instructions = resp.instructions.unwrap();
7505
7506 assert!(instructions.contains("## Resources"));
7507 assert!(instructions.contains("- **file:///schema.sql**: Database schema"));
7508 assert!(!instructions.contains("## Tools"));
7509 }
7510
7511 #[tokio::test]
7512 async fn test_auto_instructions_with_resource_templates() {
7513 use crate::resource::ResourceTemplateBuilder;
7514
7515 let template = ResourceTemplateBuilder::new("file:///{path}")
7516 .name("File")
7517 .description("Read a file by path")
7518 .handler(
7519 |_uri: String, _vars: std::collections::HashMap<String, String>| async move {
7520 Ok(crate::ReadResourceResult::text("content", "text/plain"))
7521 },
7522 );
7523
7524 let mut router = McpRouter::new()
7525 .auto_instructions()
7526 .resource_template(template);
7527
7528 let resp = send_initialize(&mut router).await;
7529 let instructions = resp.instructions.unwrap();
7530
7531 assert!(instructions.contains("## Resources"));
7532 assert!(instructions.contains("- **file:///{path}**: Read a file by path"));
7533 }
7534
7535 #[tokio::test]
7536 async fn test_auto_instructions_with_prompts() {
7537 use crate::prompt::PromptBuilder;
7538
7539 let prompt = PromptBuilder::new("write_query")
7540 .description("Help write a SQL query")
7541 .user_message("Write a query for: {task}");
7542
7543 let mut router = McpRouter::new().auto_instructions().prompt(prompt);
7544
7545 let resp = send_initialize(&mut router).await;
7546 let instructions = resp.instructions.unwrap();
7547
7548 assert!(instructions.contains("## Prompts"));
7549 assert!(instructions.contains("- **write_query**: Help write a SQL query"));
7550 assert!(!instructions.contains("## Tools"));
7551 }
7552
7553 #[tokio::test]
7554 async fn test_auto_instructions_all_sections() {
7555 use crate::prompt::PromptBuilder;
7556 use crate::resource::ResourceBuilder;
7557
7558 let tool = ToolBuilder::new("query")
7559 .description("Execute SQL")
7560 .read_only()
7561 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7562 .build();
7563 let resource = ResourceBuilder::new("db://schema")
7564 .name("Schema")
7565 .description("Full database schema")
7566 .text("schema");
7567 let prompt = PromptBuilder::new("write_query")
7568 .description("Help write a SQL query")
7569 .user_message("Write a query");
7570
7571 let mut router = McpRouter::new()
7572 .auto_instructions()
7573 .tool(tool)
7574 .resource(resource)
7575 .prompt(prompt);
7576
7577 let resp = send_initialize(&mut router).await;
7578 let instructions = resp.instructions.unwrap();
7579
7580 assert!(instructions.contains("## Tools"));
7582 assert!(instructions.contains("## Resources"));
7583 assert!(instructions.contains("## Prompts"));
7584
7585 let tools_pos = instructions.find("## Tools").unwrap();
7587 let resources_pos = instructions.find("## Resources").unwrap();
7588 let prompts_pos = instructions.find("## Prompts").unwrap();
7589 assert!(tools_pos < resources_pos);
7590 assert!(resources_pos < prompts_pos);
7591 }
7592
7593 #[tokio::test]
7594 async fn test_auto_instructions_with_prefix_and_suffix() {
7595 let tool = ToolBuilder::new("echo")
7596 .description("Echo input")
7597 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7598 .build();
7599
7600 let mut router = McpRouter::new()
7601 .auto_instructions_with(
7602 Some("This server provides echo capabilities."),
7603 Some("Contact admin@example.com for support."),
7604 )
7605 .tool(tool);
7606
7607 let resp = send_initialize(&mut router).await;
7608 let instructions = resp.instructions.unwrap();
7609
7610 assert!(instructions.starts_with("This server provides echo capabilities."));
7611 assert!(instructions.ends_with("Contact admin@example.com for support."));
7612 assert!(instructions.contains("## Tools"));
7613 assert!(instructions.contains("- **echo**: Echo input"));
7614 }
7615
7616 #[tokio::test]
7617 async fn test_auto_instructions_prefix_only() {
7618 let tool = ToolBuilder::new("echo")
7619 .description("Echo input")
7620 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7621 .build();
7622
7623 let mut router = McpRouter::new()
7624 .auto_instructions_with(Some("My server intro."), None::<String>)
7625 .tool(tool);
7626
7627 let resp = send_initialize(&mut router).await;
7628 let instructions = resp.instructions.unwrap();
7629
7630 assert!(instructions.starts_with("My server intro."));
7631 assert!(instructions.contains("- **echo**: Echo input"));
7632 }
7633
7634 #[tokio::test]
7635 async fn test_auto_instructions_empty_router() {
7636 let mut router = McpRouter::new().auto_instructions();
7637
7638 let resp = send_initialize(&mut router).await;
7639 let instructions = resp.instructions.expect("should have instructions");
7640
7641 assert!(!instructions.contains("## Tools"));
7643 assert!(!instructions.contains("## Resources"));
7644 assert!(!instructions.contains("## Prompts"));
7645 assert!(instructions.is_empty());
7646 }
7647
7648 #[tokio::test]
7649 async fn test_auto_instructions_overrides_manual() {
7650 let tool = ToolBuilder::new("echo")
7651 .description("Echo input")
7652 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7653 .build();
7654
7655 let mut router = McpRouter::new()
7656 .instructions("This will be overridden")
7657 .auto_instructions()
7658 .tool(tool);
7659
7660 let resp = send_initialize(&mut router).await;
7661 let instructions = resp.instructions.unwrap();
7662
7663 assert!(!instructions.contains("This will be overridden"));
7664 assert!(instructions.contains("- **echo**: Echo input"));
7665 }
7666
7667 #[tokio::test]
7668 async fn test_no_auto_instructions_returns_manual() {
7669 let tool = ToolBuilder::new("echo")
7670 .description("Echo input")
7671 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7672 .build();
7673
7674 let mut router = McpRouter::new()
7675 .instructions("Manual instructions here")
7676 .tool(tool);
7677
7678 let resp = send_initialize(&mut router).await;
7679 let instructions = resp.instructions.unwrap();
7680
7681 assert_eq!(instructions, "Manual instructions here");
7682 }
7683
7684 #[tokio::test]
7685 async fn test_auto_instructions_no_description_fallback() {
7686 let tool = ToolBuilder::new("mystery")
7687 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7688 .build();
7689
7690 let mut router = McpRouter::new().auto_instructions().tool(tool);
7691
7692 let resp = send_initialize(&mut router).await;
7693 let instructions = resp.instructions.unwrap();
7694
7695 assert!(instructions.contains("- **mystery**: No description"));
7696 }
7697
7698 #[tokio::test]
7699 async fn test_auto_instructions_sorted_alphabetically() {
7700 let tool_z = ToolBuilder::new("zebra")
7701 .description("Z tool")
7702 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7703 .build();
7704 let tool_a = ToolBuilder::new("alpha")
7705 .description("A tool")
7706 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7707 .build();
7708 let tool_m = ToolBuilder::new("middle")
7709 .description("M tool")
7710 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7711 .build();
7712
7713 let mut router = McpRouter::new()
7714 .auto_instructions()
7715 .tool(tool_z)
7716 .tool(tool_a)
7717 .tool(tool_m);
7718
7719 let resp = send_initialize(&mut router).await;
7720 let instructions = resp.instructions.unwrap();
7721
7722 let alpha_pos = instructions.find("**alpha**").unwrap();
7723 let middle_pos = instructions.find("**middle**").unwrap();
7724 let zebra_pos = instructions.find("**zebra**").unwrap();
7725 assert!(alpha_pos < middle_pos);
7726 assert!(middle_pos < zebra_pos);
7727 }
7728
7729 #[tokio::test]
7730 async fn test_auto_instructions_read_only_and_idempotent_tags() {
7731 let tool = ToolBuilder::new("safe_update")
7732 .description("Safe update operation")
7733 .idempotent()
7734 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7735 .build();
7736
7737 let mut router = McpRouter::new().auto_instructions().tool(tool);
7738
7739 let resp = send_initialize(&mut router).await;
7740 let instructions = resp.instructions.unwrap();
7741
7742 assert!(
7743 instructions.contains("[idempotent]"),
7744 "got: {}",
7745 instructions
7746 );
7747 }
7748
7749 #[tokio::test]
7750 async fn test_auto_instructions_lazy_generation() {
7751 let mut router = McpRouter::new().auto_instructions();
7754
7755 let tool = ToolBuilder::new("late_tool")
7756 .description("Added after auto_instructions")
7757 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7758 .build();
7759
7760 router = router.tool(tool);
7761
7762 let resp = send_initialize(&mut router).await;
7763 let instructions = resp.instructions.unwrap();
7764
7765 assert!(instructions.contains("- **late_tool**: Added after auto_instructions"));
7766 }
7767
7768 #[tokio::test]
7769 async fn test_auto_instructions_multiple_annotation_tags() {
7770 let tool = ToolBuilder::new("update")
7771 .description("Update a record")
7772 .annotations(ToolAnnotations {
7773 read_only_hint: true,
7774 idempotent_hint: true,
7775 ..Default::default()
7776 })
7777 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7778 .build();
7779
7780 let mut router = McpRouter::new().auto_instructions().tool(tool);
7781
7782 let resp = send_initialize(&mut router).await;
7783 let instructions = resp.instructions.unwrap();
7784
7785 assert!(
7786 instructions.contains("[read-only, idempotent]"),
7787 "got: {}",
7788 instructions
7789 );
7790 }
7791
7792 #[tokio::test]
7793 async fn test_auto_instructions_no_annotations_no_tags() {
7794 let tool = ToolBuilder::new("fetch")
7796 .description("Fetch data")
7797 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7798 .build();
7799
7800 let mut router = McpRouter::new().auto_instructions().tool(tool);
7801
7802 let resp = send_initialize(&mut router).await;
7803 let instructions = resp.instructions.unwrap();
7804
7805 assert!(
7807 !instructions.contains('['),
7808 "should have no tags, got: {}",
7809 instructions
7810 );
7811 assert!(instructions.contains("- **fetch**: Fetch data"));
7812 }
7813
7814 async fn send_initialize(router: &mut McpRouter) -> InitializeResult {
7816 let init_req = RouterRequest {
7817 id: RequestId::Number(0),
7818 inner: McpRequest::Initialize(InitializeParams {
7819 protocol_version: "2025-11-25".to_string(),
7820 capabilities: ClientCapabilities {
7821 roots: None,
7822 sampling: None,
7823 elicitation: None,
7824 tasks: None,
7825 experimental: None,
7826 extensions: None,
7827 },
7828 client_info: Implementation {
7829 name: "test".to_string(),
7830 version: "1.0".to_string(),
7831 ..Default::default()
7832 },
7833 meta: None,
7834 }),
7835 extensions: Extensions::new(),
7836 };
7837 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
7838 match resp.inner {
7839 Ok(McpResponse::Initialize(result)) => result,
7840 other => panic!("Expected Initialize response, got {:?}", other),
7841 }
7842 }
7843
7844 #[tokio::test]
7845 async fn test_notify_tools_list_changed() {
7846 let (tx, mut rx) = crate::context::notification_channel(16);
7847
7848 let router = McpRouter::new()
7849 .server_info("test", "1.0")
7850 .with_notification_sender(tx);
7851
7852 assert!(router.notify_tools_list_changed());
7853
7854 let notification = rx.recv().await.unwrap();
7855 assert!(matches!(notification, ServerNotification::ToolsListChanged));
7856 }
7857
7858 #[tokio::test]
7859 async fn test_notify_prompts_list_changed() {
7860 let (tx, mut rx) = crate::context::notification_channel(16);
7861
7862 let router = McpRouter::new()
7863 .server_info("test", "1.0")
7864 .with_notification_sender(tx);
7865
7866 assert!(router.notify_prompts_list_changed());
7867
7868 let notification = rx.recv().await.unwrap();
7869 assert!(matches!(
7870 notification,
7871 ServerNotification::PromptsListChanged
7872 ));
7873 }
7874
7875 #[tokio::test]
7876 async fn test_notify_without_sender_returns_false() {
7877 let router = McpRouter::new().server_info("test", "1.0");
7878
7879 assert!(!router.notify_tools_list_changed());
7880 assert!(!router.notify_prompts_list_changed());
7881 assert!(!router.notify_resources_list_changed());
7882 }
7883
7884 #[tokio::test]
7885 async fn test_list_changed_capabilities_with_notification_sender() {
7886 let (tx, _rx) = crate::context::notification_channel(16);
7887 let tool = ToolBuilder::new("test")
7888 .description("test")
7889 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
7890 .build();
7891
7892 let mut router = McpRouter::new()
7893 .server_info("test", "1.0")
7894 .tool(tool)
7895 .with_notification_sender(tx);
7896
7897 init_router(&mut router).await;
7898
7899 let caps = router.capabilities();
7900 let tools_cap = caps.tools.expect("tools capability should be present");
7901 assert!(
7902 tools_cap.list_changed,
7903 "tools.listChanged should be true when notification sender is configured"
7904 );
7905 }
7906
7907 #[tokio::test]
7908 async fn test_list_changed_capabilities_without_notification_sender() {
7909 let tool = ToolBuilder::new("test")
7910 .description("test")
7911 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
7912 .build();
7913
7914 let mut router = McpRouter::new().server_info("test", "1.0").tool(tool);
7915
7916 init_router(&mut router).await;
7917
7918 let caps = router.capabilities();
7919 let tools_cap = caps.tools.expect("tools capability should be present");
7920 assert!(
7921 !tools_cap.list_changed,
7922 "tools.listChanged should be false without notification sender"
7923 );
7924 }
7925
7926 #[tokio::test]
7927 async fn test_set_logging_level_filters_messages() {
7928 let (tx, mut rx) = crate::context::notification_channel(16);
7929
7930 let mut router = McpRouter::new()
7931 .server_info("test", "1.0")
7932 .with_notification_sender(tx);
7933
7934 init_router(&mut router).await;
7935
7936 let set_level_req = RouterRequest {
7938 id: RequestId::Number(99),
7939 inner: McpRequest::SetLoggingLevel(SetLogLevelParams {
7940 level: LogLevel::Warning,
7941 meta: None,
7942 }),
7943 extensions: crate::context::Extensions::new(),
7944 };
7945 let resp = router
7946 .ready()
7947 .await
7948 .unwrap()
7949 .call(set_level_req)
7950 .await
7951 .unwrap();
7952 assert!(matches!(resp.inner, Ok(McpResponse::SetLoggingLevel(_))));
7953
7954 let ctx = router.create_context(RequestId::Number(100), None);
7956
7957 ctx.send_log(LoggingMessageParams::new(
7959 LogLevel::Error,
7960 serde_json::Value::Null,
7961 ));
7962 assert!(
7963 rx.try_recv().is_ok(),
7964 "Error should pass through Warning filter"
7965 );
7966
7967 ctx.send_log(LoggingMessageParams::new(
7969 LogLevel::Info,
7970 serde_json::Value::Null,
7971 ));
7972 assert!(
7973 rx.try_recv().is_err(),
7974 "Info should be filtered at Warning level"
7975 );
7976 }
7977
7978 #[test]
7979 fn test_paginate_no_page_size() {
7980 let items = vec![1, 2, 3, 4, 5];
7981 let (page, cursor) = paginate(items.clone(), None, None).unwrap();
7982 assert_eq!(page, items);
7983 assert!(cursor.is_none());
7984 }
7985
7986 #[test]
7987 fn test_paginate_first_page() {
7988 let items = vec![1, 2, 3, 4, 5];
7989 let (page, cursor) = paginate(items, None, Some(2)).unwrap();
7990 assert_eq!(page, vec![1, 2]);
7991 assert!(cursor.is_some());
7992 }
7993
7994 #[test]
7995 fn test_paginate_middle_page() {
7996 let items = vec![1, 2, 3, 4, 5];
7997 let (page1, cursor1) = paginate(items.clone(), None, Some(2)).unwrap();
7998 assert_eq!(page1, vec![1, 2]);
7999
8000 let (page2, cursor2) = paginate(items, cursor1.as_deref(), Some(2)).unwrap();
8001 assert_eq!(page2, vec![3, 4]);
8002 assert!(cursor2.is_some());
8003 }
8004
8005 #[test]
8006 fn test_paginate_last_page() {
8007 let items = vec![1, 2, 3, 4, 5];
8008 let cursor = encode_cursor(4);
8010 let (page, next) = paginate(items, Some(&cursor), Some(2)).unwrap();
8011 assert_eq!(page, vec![5]);
8012 assert!(next.is_none());
8013 }
8014
8015 #[test]
8016 fn test_paginate_exact_boundary() {
8017 let items = vec![1, 2, 3, 4];
8018 let (page, cursor) = paginate(items, None, Some(4)).unwrap();
8019 assert_eq!(page, vec![1, 2, 3, 4]);
8020 assert!(cursor.is_none());
8021 }
8022
8023 #[test]
8024 fn test_paginate_invalid_cursor() {
8025 let items = vec![1, 2, 3];
8026 let result = paginate(items, Some("not-valid-base64!@#$"), Some(2));
8027 assert!(result.is_err());
8028 }
8029
8030 #[test]
8031 fn test_cursor_round_trip() {
8032 let offset = 42;
8033 let encoded = encode_cursor(offset);
8034 let decoded = decode_cursor(&encoded).unwrap();
8035 assert_eq!(decoded, offset);
8036 }
8037
8038 #[tokio::test]
8039 async fn test_list_tools_pagination() {
8040 let tool_a = ToolBuilder::new("alpha")
8041 .description("a")
8042 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8043 .build();
8044 let tool_b = ToolBuilder::new("beta")
8045 .description("b")
8046 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8047 .build();
8048 let tool_c = ToolBuilder::new("gamma")
8049 .description("c")
8050 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8051 .build();
8052
8053 let mut router = McpRouter::new()
8054 .server_info("test", "1.0")
8055 .page_size(2)
8056 .tool(tool_a)
8057 .tool(tool_b)
8058 .tool(tool_c);
8059
8060 init_router(&mut router).await;
8061
8062 let req = RouterRequest {
8064 id: RequestId::Number(1),
8065 inner: McpRequest::ListTools(ListToolsParams {
8066 cursor: None,
8067 meta: None,
8068 }),
8069 extensions: Extensions::new(),
8070 };
8071 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8072 let (tools, next_cursor) = match resp.inner {
8073 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8074 other => panic!("Expected ListTools, got {:?}", other),
8075 };
8076 assert_eq!(tools.len(), 2);
8077 assert_eq!(tools[0].name, "alpha");
8078 assert_eq!(tools[1].name, "beta");
8079 assert!(next_cursor.is_some());
8080
8081 let req = RouterRequest {
8083 id: RequestId::Number(2),
8084 inner: McpRequest::ListTools(ListToolsParams {
8085 cursor: next_cursor,
8086 meta: None,
8087 }),
8088 extensions: Extensions::new(),
8089 };
8090 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8091 let (tools, next_cursor) = match resp.inner {
8092 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8093 other => panic!("Expected ListTools, got {:?}", other),
8094 };
8095 assert_eq!(tools.len(), 1);
8096 assert_eq!(tools[0].name, "gamma");
8097 assert!(next_cursor.is_none());
8098 }
8099
8100 #[tokio::test]
8101 async fn test_list_tools_no_pagination_by_default() {
8102 let tool_a = ToolBuilder::new("alpha")
8103 .description("a")
8104 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8105 .build();
8106 let tool_b = ToolBuilder::new("beta")
8107 .description("b")
8108 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8109 .build();
8110
8111 let mut router = McpRouter::new()
8112 .server_info("test", "1.0")
8113 .tool(tool_a)
8114 .tool(tool_b);
8115
8116 init_router(&mut router).await;
8117
8118 let req = RouterRequest {
8119 id: RequestId::Number(1),
8120 inner: McpRequest::ListTools(ListToolsParams {
8121 cursor: None,
8122 meta: None,
8123 }),
8124 extensions: Extensions::new(),
8125 };
8126 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8127 match resp.inner {
8128 Ok(McpResponse::ListTools(result)) => {
8129 assert_eq!(result.tools.len(), 2);
8130 assert!(result.next_cursor.is_none());
8131 }
8132 other => panic!("Expected ListTools, got {:?}", other),
8133 }
8134 }
8135
8136 #[cfg(feature = "dynamic-tools")]
8141 mod dynamic_tools_tests {
8142 use super::*;
8143
8144 #[tokio::test]
8145 async fn test_dynamic_tools_register_and_list() {
8146 let (router, registry) = McpRouter::new()
8147 .server_info("test", "1.0")
8148 .with_dynamic_tools();
8149
8150 let tool = ToolBuilder::new("dynamic_echo")
8151 .description("Dynamic echo")
8152 .handler(|input: AddInput| async move {
8153 Ok(CallToolResult::text(format!("{}", input.a)))
8154 })
8155 .build();
8156
8157 registry.register(tool);
8158
8159 let mut router = router;
8160 init_router(&mut router).await;
8161
8162 let req = RouterRequest {
8163 id: RequestId::Number(1),
8164 inner: McpRequest::ListTools(ListToolsParams::default()),
8165 extensions: Extensions::new(),
8166 };
8167
8168 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8169 match resp.inner {
8170 Ok(McpResponse::ListTools(result)) => {
8171 assert_eq!(result.tools.len(), 1);
8172 assert_eq!(result.tools[0].name, "dynamic_echo");
8173 }
8174 _ => panic!("Expected ListTools response"),
8175 }
8176 }
8177
8178 #[tokio::test]
8179 async fn test_dynamic_tools_unregister() {
8180 let (router, registry) = McpRouter::new()
8181 .server_info("test", "1.0")
8182 .with_dynamic_tools();
8183
8184 let tool = ToolBuilder::new("temp")
8185 .description("Temporary")
8186 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8187 .build();
8188
8189 registry.register(tool);
8190 assert!(registry.contains("temp"));
8191
8192 let removed = registry.unregister("temp");
8193 assert!(removed);
8194 assert!(!registry.contains("temp"));
8195
8196 assert!(!registry.unregister("temp"));
8198
8199 let mut router = router;
8200 init_router(&mut router).await;
8201
8202 let req = RouterRequest {
8203 id: RequestId::Number(1),
8204 inner: McpRequest::ListTools(ListToolsParams::default()),
8205 extensions: Extensions::new(),
8206 };
8207
8208 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8209 match resp.inner {
8210 Ok(McpResponse::ListTools(result)) => {
8211 assert_eq!(result.tools.len(), 0);
8212 }
8213 _ => panic!("Expected ListTools response"),
8214 }
8215 }
8216
8217 #[tokio::test]
8218 async fn test_dynamic_tools_merged_with_static() {
8219 let static_tool = ToolBuilder::new("static_tool")
8220 .description("Static")
8221 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8222 .build();
8223
8224 let (router, registry) = McpRouter::new()
8225 .server_info("test", "1.0")
8226 .tool(static_tool)
8227 .with_dynamic_tools();
8228
8229 let dynamic_tool = ToolBuilder::new("dynamic_tool")
8230 .description("Dynamic")
8231 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8232 .build();
8233
8234 registry.register(dynamic_tool);
8235
8236 let mut router = router;
8237 init_router(&mut router).await;
8238
8239 let req = RouterRequest {
8240 id: RequestId::Number(1),
8241 inner: McpRequest::ListTools(ListToolsParams::default()),
8242 extensions: Extensions::new(),
8243 };
8244
8245 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8246 match resp.inner {
8247 Ok(McpResponse::ListTools(result)) => {
8248 assert_eq!(result.tools.len(), 2);
8249 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8250 assert!(names.contains(&"static_tool"));
8251 assert!(names.contains(&"dynamic_tool"));
8252 }
8253 _ => panic!("Expected ListTools response"),
8254 }
8255 }
8256
8257 #[tokio::test]
8258 async fn test_static_tools_shadow_dynamic() {
8259 let static_tool = ToolBuilder::new("shared")
8260 .description("Static version")
8261 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8262 .build();
8263
8264 let (router, registry) = McpRouter::new()
8265 .server_info("test", "1.0")
8266 .tool(static_tool)
8267 .with_dynamic_tools();
8268
8269 let dynamic_tool = ToolBuilder::new("shared")
8270 .description("Dynamic version")
8271 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8272 .build();
8273
8274 registry.register(dynamic_tool);
8275
8276 let mut router = router;
8277 init_router(&mut router).await;
8278
8279 let req = RouterRequest {
8281 id: RequestId::Number(1),
8282 inner: McpRequest::ListTools(ListToolsParams::default()),
8283 extensions: Extensions::new(),
8284 };
8285
8286 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8287 match resp.inner {
8288 Ok(McpResponse::ListTools(result)) => {
8289 assert_eq!(result.tools.len(), 1);
8290 assert_eq!(result.tools[0].name, "shared");
8291 assert_eq!(
8292 result.tools[0].description.as_deref(),
8293 Some("Static version")
8294 );
8295 }
8296 _ => panic!("Expected ListTools response"),
8297 }
8298
8299 let req = RouterRequest {
8301 id: RequestId::Number(2),
8302 inner: McpRequest::CallTool(CallToolParams {
8303 input_responses: None,
8304 request_state: None,
8305 name: "shared".to_string(),
8306 arguments: serde_json::json!({"a": 1, "b": 2}),
8307 meta: None,
8308 task: None,
8309 }),
8310 extensions: Extensions::new(),
8311 };
8312
8313 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8314 match resp.inner {
8315 Ok(McpResponse::CallTool(result)) => {
8316 assert!(!result.is_error);
8317 match &result.content[0] {
8318 Content::Text { text, .. } => assert_eq!(text, "static"),
8319 _ => panic!("Expected text content"),
8320 }
8321 }
8322 _ => panic!("Expected CallTool response"),
8323 }
8324 }
8325
8326 #[tokio::test]
8327 async fn test_dynamic_tools_call() {
8328 let (router, registry) = McpRouter::new()
8329 .server_info("test", "1.0")
8330 .with_dynamic_tools();
8331
8332 let tool = ToolBuilder::new("add")
8333 .description("Add two numbers")
8334 .handler(|input: AddInput| async move {
8335 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
8336 })
8337 .build();
8338
8339 registry.register(tool);
8340
8341 let mut router = router;
8342 init_router(&mut router).await;
8343
8344 let req = RouterRequest {
8345 id: RequestId::Number(1),
8346 inner: McpRequest::CallTool(CallToolParams {
8347 input_responses: None,
8348 request_state: None,
8349 name: "add".to_string(),
8350 arguments: serde_json::json!({"a": 3, "b": 4}),
8351 meta: None,
8352 task: None,
8353 }),
8354 extensions: Extensions::new(),
8355 };
8356
8357 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8358 match resp.inner {
8359 Ok(McpResponse::CallTool(result)) => {
8360 assert!(!result.is_error);
8361 match &result.content[0] {
8362 Content::Text { text, .. } => assert_eq!(text, "7"),
8363 _ => panic!("Expected text content"),
8364 }
8365 }
8366 _ => panic!("Expected CallTool response"),
8367 }
8368 }
8369
8370 #[tokio::test]
8371 async fn test_dynamic_tools_notification_on_register() {
8372 let (tx, mut rx) = crate::context::notification_channel(16);
8373 let (router, registry) = McpRouter::new()
8374 .server_info("test", "1.0")
8375 .with_dynamic_tools();
8376 let _router = router.with_notification_sender(tx);
8377
8378 let tool = ToolBuilder::new("notified")
8379 .description("Test")
8380 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8381 .build();
8382
8383 registry.register(tool);
8384
8385 let notification = rx.recv().await.unwrap();
8386 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8387 }
8388
8389 #[tokio::test]
8390 async fn test_dynamic_tools_notification_on_unregister() {
8391 let (tx, mut rx) = crate::context::notification_channel(16);
8392 let (router, registry) = McpRouter::new()
8393 .server_info("test", "1.0")
8394 .with_dynamic_tools();
8395 let _router = router.with_notification_sender(tx);
8396
8397 let tool = ToolBuilder::new("notified")
8398 .description("Test")
8399 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8400 .build();
8401
8402 registry.register(tool);
8403 let _ = rx.recv().await.unwrap();
8405
8406 registry.unregister("notified");
8407 let notification = rx.recv().await.unwrap();
8408 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8409 }
8410
8411 #[tokio::test]
8412 async fn test_dynamic_tools_no_notification_on_empty_unregister() {
8413 let (tx, mut rx) = crate::context::notification_channel(16);
8414 let (router, registry) = McpRouter::new()
8415 .server_info("test", "1.0")
8416 .with_dynamic_tools();
8417 let _router = router.with_notification_sender(tx);
8418
8419 assert!(!registry.unregister("nonexistent"));
8421
8422 assert!(rx.try_recv().is_err());
8424 }
8425
8426 #[tokio::test]
8427 async fn test_dynamic_tools_filter_applies() {
8428 use crate::filter::CapabilityFilter;
8429
8430 let (router, registry) = McpRouter::new()
8431 .server_info("test", "1.0")
8432 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| {
8433 tool.name != "hidden"
8434 }))
8435 .with_dynamic_tools();
8436
8437 let visible = ToolBuilder::new("visible")
8438 .description("Visible")
8439 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8440 .build();
8441
8442 let hidden = ToolBuilder::new("hidden")
8443 .description("Hidden")
8444 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8445 .build();
8446
8447 registry.register(visible);
8448 registry.register(hidden);
8449
8450 let mut router = router;
8451 init_router(&mut router).await;
8452
8453 let req = RouterRequest {
8455 id: RequestId::Number(1),
8456 inner: McpRequest::ListTools(ListToolsParams::default()),
8457 extensions: Extensions::new(),
8458 };
8459
8460 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8461 match resp.inner {
8462 Ok(McpResponse::ListTools(result)) => {
8463 assert_eq!(result.tools.len(), 1);
8464 assert_eq!(result.tools[0].name, "visible");
8465 }
8466 _ => panic!("Expected ListTools response"),
8467 }
8468
8469 let req = RouterRequest {
8471 id: RequestId::Number(2),
8472 inner: McpRequest::CallTool(CallToolParams {
8473 input_responses: None,
8474 request_state: None,
8475 name: "hidden".to_string(),
8476 arguments: serde_json::json!({"a": 1, "b": 2}),
8477 meta: None,
8478 task: None,
8479 }),
8480 extensions: Extensions::new(),
8481 };
8482
8483 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8484 match resp.inner {
8485 Err(e) => {
8486 assert_eq!(e.code, -32601); }
8488 _ => panic!("Expected JsonRpc error"),
8489 }
8490 }
8491
8492 #[tokio::test]
8493 async fn test_dynamic_tools_capabilities_advertised() {
8494 let (mut router, _registry) = McpRouter::new()
8496 .server_info("test", "1.0")
8497 .with_dynamic_tools();
8498
8499 let init_req = RouterRequest {
8500 id: RequestId::Number(1),
8501 inner: McpRequest::Initialize(InitializeParams {
8502 protocol_version: "2025-11-25".to_string(),
8503 capabilities: ClientCapabilities::default(),
8504 client_info: Implementation {
8505 name: "test".to_string(),
8506 version: "1.0".to_string(),
8507 ..Default::default()
8508 },
8509 meta: None,
8510 }),
8511 extensions: Extensions::new(),
8512 };
8513
8514 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8515 match resp.inner {
8516 Ok(McpResponse::Initialize(result)) => {
8517 assert!(result.capabilities.tools.is_some());
8518 }
8519 _ => panic!("Expected Initialize response"),
8520 }
8521 }
8522
8523 #[tokio::test]
8524 async fn test_dynamic_tools_multi_session_notification() {
8525 let (tx1, mut rx1) = crate::context::notification_channel(16);
8526 let (tx2, mut rx2) = crate::context::notification_channel(16);
8527
8528 let (router, registry) = McpRouter::new()
8529 .server_info("test", "1.0")
8530 .with_dynamic_tools();
8531
8532 let _session1 = router.clone().with_notification_sender(tx1);
8534 let _session2 = router.clone().with_notification_sender(tx2);
8535
8536 let tool = ToolBuilder::new("broadcast")
8537 .description("Test")
8538 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8539 .build();
8540
8541 registry.register(tool);
8542
8543 let n1 = rx1.recv().await.unwrap();
8545 let n2 = rx2.recv().await.unwrap();
8546 assert!(matches!(n1, ServerNotification::ToolsListChanged));
8547 assert!(matches!(n2, ServerNotification::ToolsListChanged));
8548 }
8549
8550 #[tokio::test]
8551 async fn test_dynamic_tools_call_not_found() {
8552 let (router, _registry) = McpRouter::new()
8553 .server_info("test", "1.0")
8554 .with_dynamic_tools();
8555
8556 let mut router = router;
8557 init_router(&mut router).await;
8558
8559 let req = RouterRequest {
8560 id: RequestId::Number(1),
8561 inner: McpRequest::CallTool(CallToolParams {
8562 input_responses: None,
8563 request_state: None,
8564 name: "nonexistent".to_string(),
8565 arguments: serde_json::json!({}),
8566 meta: None,
8567 task: None,
8568 }),
8569 extensions: Extensions::new(),
8570 };
8571
8572 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8573 match resp.inner {
8574 Err(e) => {
8575 assert_eq!(e.code, -32601);
8576 }
8577 _ => panic!("Expected method not found error"),
8578 }
8579 }
8580
8581 #[tokio::test]
8582 async fn test_dynamic_tools_registry_list() {
8583 let (_, registry) = McpRouter::new()
8584 .server_info("test", "1.0")
8585 .with_dynamic_tools();
8586
8587 assert!(registry.list().is_empty());
8588
8589 let tool = ToolBuilder::new("tool_a")
8590 .description("A")
8591 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8592 .build();
8593 registry.register(tool);
8594
8595 let tool = ToolBuilder::new("tool_b")
8596 .description("B")
8597 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8598 .build();
8599 registry.register(tool);
8600
8601 let tools = registry.list();
8602 assert_eq!(tools.len(), 2);
8603 let names: Vec<&str> = tools.iter().map(|t| t.name.as_str()).collect();
8604 assert!(names.contains(&"tool_a"));
8605 assert!(names.contains(&"tool_b"));
8606 }
8607 } #[tokio::test]
8610 async fn test_tool_if_true_registers() {
8611 let tool = ToolBuilder::new("conditional")
8612 .description("Conditional tool")
8613 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8614 .build();
8615
8616 let mut router = McpRouter::new().tool_if(true, tool);
8617 init_router(&mut router).await;
8618
8619 let req = RouterRequest {
8620 id: RequestId::Number(1),
8621 inner: McpRequest::ListTools(ListToolsParams::default()),
8622 extensions: Extensions::new(),
8623 };
8624 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8625 match resp.inner {
8626 Ok(McpResponse::ListTools(result)) => {
8627 assert_eq!(result.tools.len(), 1);
8628 assert_eq!(result.tools[0].name, "conditional");
8629 }
8630 _ => panic!("Expected ListTools response"),
8631 }
8632 }
8633
8634 #[tokio::test]
8635 async fn test_tool_if_false_skips() {
8636 let tool = ToolBuilder::new("conditional")
8637 .description("Conditional tool")
8638 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8639 .build();
8640
8641 let mut router = McpRouter::new().tool_if(false, tool);
8642 init_router(&mut router).await;
8643
8644 let req = RouterRequest {
8645 id: RequestId::Number(1),
8646 inner: McpRequest::ListTools(ListToolsParams::default()),
8647 extensions: Extensions::new(),
8648 };
8649 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8650 match resp.inner {
8651 Ok(McpResponse::ListTools(result)) => {
8652 assert_eq!(result.tools.len(), 0);
8653 }
8654 _ => panic!("Expected ListTools response"),
8655 }
8656 }
8657
8658 #[tokio::test]
8659 async fn test_tools_if_batch_conditional() {
8660 let tools = vec![
8661 ToolBuilder::new("a")
8662 .description("Tool A")
8663 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8664 .build(),
8665 ToolBuilder::new("b")
8666 .description("Tool B")
8667 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8668 .build(),
8669 ];
8670
8671 let mut router = McpRouter::new().tools_if(false, tools);
8672 init_router(&mut router).await;
8673
8674 let req = RouterRequest {
8675 id: RequestId::Number(1),
8676 inner: McpRequest::ListTools(ListToolsParams::default()),
8677 extensions: Extensions::new(),
8678 };
8679 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8680 match resp.inner {
8681 Ok(McpResponse::ListTools(result)) => {
8682 assert_eq!(result.tools.len(), 0);
8683 }
8684 _ => panic!("Expected ListTools response"),
8685 }
8686 }
8687
8688 #[test]
8689 fn test_resource_if_true_registers() {
8690 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8691 .name("test")
8692 .text("hello");
8693
8694 let router = McpRouter::new().resource_if(true, resource);
8695 assert_eq!(router.inner.resources.len(), 1);
8696 }
8697
8698 #[test]
8699 fn test_resource_if_false_skips() {
8700 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8701 .name("test")
8702 .text("hello");
8703
8704 let router = McpRouter::new().resource_if(false, resource);
8705 assert_eq!(router.inner.resources.len(), 0);
8706 }
8707
8708 #[test]
8709 fn test_prompt_if_true_registers() {
8710 let prompt = crate::prompt::PromptBuilder::new("greet")
8711 .description("Greeting")
8712 .user_message("Hello!");
8713
8714 let router = McpRouter::new().prompt_if(true, prompt);
8715 assert_eq!(router.inner.prompts.len(), 1);
8716 }
8717
8718 #[test]
8719 fn test_prompt_if_false_skips() {
8720 let prompt = crate::prompt::PromptBuilder::new("greet")
8721 .description("Greeting")
8722 .user_message("Hello!");
8723
8724 let router = McpRouter::new().prompt_if(false, prompt);
8725 assert_eq!(router.inner.prompts.len(), 0);
8726 }
8727
8728 #[tokio::test]
8729 async fn test_disable_tool_hides_from_list() {
8730 let safe = ToolBuilder::new("safe")
8731 .description("Safe tool")
8732 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8733 .build();
8734 let dangerous = ToolBuilder::new("dangerous")
8735 .description("Dangerous tool")
8736 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8737 .build();
8738 let mut router = McpRouter::new().tool(safe).tool(dangerous);
8739 init_router(&mut router).await;
8740
8741 router.disable_tool("dangerous");
8742 assert!(router.is_tool_enabled("safe"));
8743 assert!(!router.is_tool_enabled("dangerous"));
8744
8745 let req = RouterRequest {
8746 id: RequestId::Number(1),
8747 inner: McpRequest::ListTools(ListToolsParams::default()),
8748 extensions: Extensions::new(),
8749 };
8750 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8751 match resp.inner {
8752 Ok(McpResponse::ListTools(result)) => {
8753 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8754 assert_eq!(names, vec!["safe"]);
8755 }
8756 _ => panic!("Expected ListTools response"),
8757 }
8758 }
8759
8760 #[tokio::test]
8761 async fn test_disable_tool_blocks_call() {
8762 let dangerous = ToolBuilder::new("dangerous")
8763 .description("Dangerous tool")
8764 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
8765 .build();
8766 let mut router = McpRouter::new().tool(dangerous);
8767 init_router(&mut router).await;
8768
8769 router.disable_tool("dangerous");
8770
8771 let req = RouterRequest {
8772 id: RequestId::Number(2),
8773 inner: McpRequest::CallTool(CallToolParams {
8774 input_responses: None,
8775 request_state: None,
8776 name: "dangerous".to_string(),
8777 arguments: serde_json::json!({"a": 1, "b": 2}),
8778 meta: None,
8779 task: None,
8780 }),
8781 extensions: Extensions::new(),
8782 };
8783 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8784 let err = resp.inner.expect_err("disabled tool should error");
8785 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
8786 }
8787
8788 #[tokio::test]
8789 async fn test_enable_tool_restores_visibility() {
8790 let tool = ToolBuilder::new("flippy")
8791 .description("Toggleable tool")
8792 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
8793 .build();
8794 let mut router = McpRouter::new().tool(tool);
8795 init_router(&mut router).await;
8796
8797 router.disable_tool("flippy");
8798 router.enable_tool("flippy");
8799 assert!(router.is_tool_enabled("flippy"));
8800
8801 let req = RouterRequest {
8802 id: RequestId::Number(3),
8803 inner: McpRequest::CallTool(CallToolParams {
8804 input_responses: None,
8805 request_state: None,
8806 name: "flippy".to_string(),
8807 arguments: serde_json::json!({"a": 1, "b": 2}),
8808 meta: None,
8809 task: None,
8810 }),
8811 extensions: Extensions::new(),
8812 };
8813 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8814 match resp.inner {
8815 Ok(McpResponse::CallTool(result)) => {
8816 assert_eq!(result.first_text(), Some("ran"));
8817 }
8818 _ => panic!("Expected CallTool response"),
8819 }
8820 }
8821
8822 #[tokio::test]
8823 async fn test_disable_propagates_through_fresh_session() {
8824 let tool = ToolBuilder::new("shared")
8825 .description("Shared across sessions")
8826 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8827 .build();
8828 let router = McpRouter::new().tool(tool);
8829
8830 router.disable_tool("shared");
8832 let mut child = router.with_fresh_session();
8833 init_router(&mut child).await;
8834 assert!(!child.is_tool_enabled("shared"));
8835
8836 let req = RouterRequest {
8837 id: RequestId::Number(4),
8838 inner: McpRequest::ListTools(ListToolsParams::default()),
8839 extensions: Extensions::new(),
8840 };
8841 let resp = child.ready().await.unwrap().call(req).await.unwrap();
8842 match resp.inner {
8843 Ok(McpResponse::ListTools(result)) => {
8844 assert!(result.tools.is_empty());
8845 }
8846 _ => panic!("Expected ListTools response"),
8847 }
8848 }
8849
8850 #[tokio::test]
8851 async fn test_disable_resource_and_prompt() {
8852 let resource = crate::resource::ResourceBuilder::new("file:///hidden.txt")
8853 .name("hidden")
8854 .text("secret");
8855 let prompt = crate::prompt::PromptBuilder::new("hidden_prompt")
8856 .description("hidden")
8857 .user_message("hello");
8858
8859 let mut router = McpRouter::new().resource(resource).prompt(prompt);
8860 init_router(&mut router).await;
8861
8862 router.disable_resource("file:///hidden.txt");
8863 router.disable_prompt("hidden_prompt");
8864 assert!(!router.is_resource_enabled("file:///hidden.txt"));
8865 assert!(!router.is_prompt_enabled("hidden_prompt"));
8866
8867 let req = RouterRequest {
8869 id: RequestId::Number(5),
8870 inner: McpRequest::ListResources(ListResourcesParams::default()),
8871 extensions: Extensions::new(),
8872 };
8873 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8874 match resp.inner {
8875 Ok(McpResponse::ListResources(result)) => {
8876 assert!(result.resources.is_empty());
8877 }
8878 _ => panic!("Expected ListResources response"),
8879 }
8880
8881 let req = RouterRequest {
8883 id: RequestId::Number(6),
8884 inner: McpRequest::ReadResource(ReadResourceParams {
8885 input_responses: None,
8886 request_state: None,
8887 uri: "file:///hidden.txt".to_string(),
8888 meta: None,
8889 }),
8890 extensions: Extensions::new(),
8891 };
8892 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8893 let err = resp.inner.expect_err("disabled resource should error");
8894 assert_eq!(err.code, -32602); let req = RouterRequest {
8898 id: RequestId::Number(7),
8899 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
8900 extensions: Extensions::new(),
8901 };
8902 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8903 match resp.inner {
8904 Ok(McpResponse::ListPrompts(result)) => {
8905 assert!(result.prompts.is_empty());
8906 }
8907 _ => panic!("Expected ListPrompts response"),
8908 }
8909
8910 let req = RouterRequest {
8912 id: RequestId::Number(8),
8913 inner: McpRequest::GetPrompt(GetPromptParams {
8914 input_responses: None,
8915 request_state: None,
8916 name: "hidden_prompt".to_string(),
8917 arguments: Default::default(),
8918 meta: None,
8919 }),
8920 extensions: Extensions::new(),
8921 };
8922 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8923 let err = resp.inner.expect_err("disabled prompt should error");
8924 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
8925 }
8926
8927 #[test]
8928 fn test_router_request_new() {
8929 let req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
8930 assert_eq!(req.id, RequestId::Number(1));
8931 assert!(req.extensions.is_empty());
8932 }
8933
8934 #[test]
8935 fn test_with_inner_preserves_extensions() {
8936 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
8937 req.extensions.insert(42u32);
8938
8939 let rewritten = req.with_inner(McpRequest::ListTools(Default::default()));
8940 assert!(matches!(rewritten.inner, McpRequest::ListTools(_)));
8941 assert_eq!(rewritten.id, RequestId::Number(1));
8942 assert_eq!(rewritten.extensions.get::<u32>(), Some(&42));
8943 }
8944
8945 #[test]
8946 fn test_with_id_and_inner_preserves_extensions() {
8947 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
8948 req.extensions.insert(String::from("token-abc"));
8949
8950 let rewritten = req.with_id_and_inner(
8951 RequestId::Number(99),
8952 McpRequest::ListResources(Default::default()),
8953 );
8954 assert_eq!(rewritten.id, RequestId::Number(99));
8955 assert!(matches!(rewritten.inner, McpRequest::ListResources(_)));
8956 assert_eq!(
8957 rewritten.extensions.get::<String>(),
8958 Some(&String::from("token-abc"))
8959 );
8960 }
8961
8962 #[test]
8963 fn test_clone_with_inner_preserves_extensions() {
8964 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
8965 req.extensions.insert(true);
8966
8967 let cloned = req.clone_with_inner(McpRequest::ListTools(Default::default()));
8968
8969 assert!(matches!(req.inner, McpRequest::Ping));
8971 assert_eq!(req.extensions.get::<bool>(), Some(&true));
8972
8973 assert!(matches!(cloned.inner, McpRequest::ListTools(_)));
8975 assert_eq!(cloned.extensions.get::<bool>(), Some(&true));
8976 }
8977
8978 #[test]
8979 fn test_router_response_is_error() {
8980 let ok_resp = RouterResponse {
8981 id: RequestId::Number(1),
8982 inner: Ok(McpResponse::Pong(Default::default())),
8983 };
8984 assert!(!ok_resp.is_error());
8985
8986 let err_resp = RouterResponse {
8987 id: RequestId::Number(2),
8988 inner: Err(JsonRpcError::internal_error("boom")),
8989 };
8990 assert!(err_resp.is_error());
8991 }
8992
8993 #[test]
8994 fn test_extensions_len_and_is_empty() {
8995 let mut ext = Extensions::new();
8996 assert!(ext.is_empty());
8997 assert_eq!(ext.len(), 0);
8998
8999 ext.insert(42u32);
9000 assert!(!ext.is_empty());
9001 assert_eq!(ext.len(), 1);
9002
9003 ext.insert(String::from("hello"));
9004 assert_eq!(ext.len(), 2);
9005 }
9006
9007 #[test]
9008 fn test_router_response_serde_roundtrip() {
9009 let response = RouterResponse {
9011 id: RequestId::Number(1),
9012 inner: Ok(McpResponse::Empty(EmptyResult {})),
9013 };
9014 let json = serde_json::to_string(&response).unwrap();
9015 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9016 assert_eq!(deserialized.id, RequestId::Number(1));
9017 assert!(!deserialized.is_error());
9018
9019 let response = RouterResponse {
9021 id: RequestId::String("req-2".into()),
9022 inner: Err(JsonRpcError::method_not_found("unknown")),
9023 };
9024 let json = serde_json::to_string(&response).unwrap();
9025 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9026 assert_eq!(deserialized.id, RequestId::String("req-2".into()));
9027 assert!(deserialized.is_error());
9028 }
9029
9030 #[tokio::test]
9037 async fn test_discover_dispatch_via_jsonrpc_service() {
9038 let router = McpRouter::new().server_info("unit-test-server", "4.2.0");
9041 let mut service = JsonRpcService::new(router);
9042
9043 let req = JsonRpcRequest::new(1, "server/discover");
9044 let resp = service.call_single(req).await.unwrap();
9045
9046 match resp {
9047 JsonRpcResponse::Result(r) => {
9048 let versions = r
9050 .result
9051 .get("supportedVersions")
9052 .and_then(|v| v.as_array())
9053 .expect("result.supportedVersions must be an array");
9054 assert!(!versions.is_empty(), "supportedVersions must not be empty");
9055
9056 assert_eq!(
9058 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
9059 "unit-test-server",
9060 "serverInfo.name must match configured value"
9061 );
9062 assert_eq!(
9063 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["version"], "4.2.0",
9064 "serverInfo.version must match configured value"
9065 );
9066
9067 assert!(
9070 r.result.get("protocolVersion").is_none(),
9071 "server/discover must NOT include protocolVersion: {:?}",
9072 r.result
9073 );
9074 }
9075 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
9076 _ => panic!("unexpected response variant"),
9077 }
9078 }
9079
9080 #[tokio::test]
9081 async fn test_discover_does_not_require_initialization() {
9082 let router = McpRouter::new().server_info("fresh-router", "1.0.0");
9085 let mut service = JsonRpcService::new(router);
9086
9087 let req = JsonRpcRequest::new(2, "server/discover");
9088 let resp = service.call_single(req).await.unwrap();
9089
9090 assert!(
9092 !matches!(resp, JsonRpcResponse::Error(_)),
9093 "server/discover must not require initialization: {:?}",
9094 resp
9095 );
9096 }
9097}
9098
9099#[cfg(test)]
9100mod cursor_property_tests {
9101 use super::{decode_cursor, encode_cursor};
9102 use proptest::prelude::*;
9103
9104 fn arb_cursor_text() -> BoxedStrategy<String> {
9105 prop_oneof![
9106 8 => prop::collection::vec(any::<char>(), 0..512)
9107 .prop_map(|chars| chars.into_iter().collect()),
9108 1 => Just("\0\r\n\t\u{001b}\u{007f}".repeat(64)),
9109 1 => Just("A".repeat(16 * 1024)),
9110 ]
9111 .boxed()
9112 }
9113
9114 proptest! {
9115 #![proptest_config(ProptestConfig::with_cases(512))]
9116
9117 #[test]
9119 fn cursor_round_trips(offset in any::<usize>()) {
9120 prop_assert_eq!(decode_cursor(&encode_cursor(offset)).unwrap(), offset);
9121 }
9122
9123 #[test]
9125 fn decode_cursor_never_panics(s in arb_cursor_text()) {
9126 let _ = decode_cursor(&s);
9127 }
9128 }
9129}