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
68async fn discard_unprepared_task(store: &Arc<dyn TaskStore>, task_id: &str) {
69 if !matches!(store.discard_task(task_id).await, Ok(true)) {
70 let _ = store
71 .cancel_task(task_id, Some("task preparation failed"))
72 .await;
73 }
74}
75
76#[cfg(feature = "stateless")]
81fn is_final_protocol_request(extensions: &crate::context::Extensions) -> bool {
82 extensions
83 .get::<crate::stateless::StatelessRequestMeta>()
84 .and_then(|meta| meta.protocol_version.as_deref())
85 == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
86}
87
88#[cfg(not(feature = "stateless"))]
89fn is_final_protocol_request(_extensions: &crate::context::Extensions) -> bool {
90 false
91}
92
93#[cfg(feature = "stateless")]
98fn client_declares_tasks(extensions: &crate::context::Extensions) -> bool {
99 final_client_capabilities(extensions).is_some_and(|capabilities| {
100 capabilities.extensions.as_ref().is_some_and(|declared| {
101 declared.contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
102 })
103 })
104}
105
106#[cfg(not(feature = "stateless"))]
107fn client_declares_tasks(_extensions: &crate::context::Extensions) -> bool {
108 false
109}
110
111fn decode_input_responses(
118 responses: &std::collections::HashMap<String, serde_json::Value>,
119) -> crate::protocol::InputResponses {
120 responses
121 .iter()
122 .filter_map(|(key, value)| {
123 serde_json::from_value(value.clone())
124 .ok()
125 .map(|response| (key.clone(), response))
126 })
127 .collect()
128}
129
130#[cfg(feature = "oauth")]
137fn request_principal(extensions: &crate::context::Extensions) -> Option<String> {
138 extensions
139 .get::<crate::oauth::token::TokenClaims>()
140 .and_then(|claims| claims.sub.clone())
141}
142
143#[cfg(not(feature = "oauth"))]
144fn request_principal(_extensions: &crate::context::Extensions) -> Option<String> {
145 None
146}
147
148fn unknown_task_error(task_id: &str) -> JsonRpcError {
153 JsonRpcError::invalid_params(format!("Task not found: {task_id}"))
154}
155
156pub(crate) fn tasks_client_capabilities() -> crate::protocol::ClientCapabilities {
159 crate::protocol::ClientCapabilities {
160 extensions: Some(
161 [(
162 tower_mcp_types::protocol::TASKS_EXTENSION_ID.to_string(),
163 serde_json::json!({}),
164 )]
165 .into_iter()
166 .collect(),
167 ),
168 ..Default::default()
169 }
170}
171
172#[cfg(feature = "stateless")]
173fn final_client_capabilities(
174 extensions: &crate::context::Extensions,
175) -> Option<&ClientCapabilities> {
176 extensions
177 .get::<crate::stateless::StatelessRequestMeta>()
178 .and_then(|meta| meta.client_capabilities.as_ref())
179}
180
181#[cfg(not(feature = "stateless"))]
182fn final_client_capabilities(
183 _extensions: &crate::context::Extensions,
184) -> Option<&ClientCapabilities> {
185 None
186}
187
188#[cfg(feature = "stateless")]
193fn json_value_contains(actual: &serde_json::Value, required: &serde_json::Value) -> bool {
194 match (actual, required) {
195 (serde_json::Value::Object(actual), serde_json::Value::Object(required)) => {
196 required.iter().all(|(key, value)| {
197 actual
198 .get(key)
199 .is_some_and(|a| json_value_contains(a, value))
200 })
201 }
202 _ => actual == required,
203 }
204}
205
206#[cfg(feature = "stateless")]
207fn client_capabilities_satisfy(actual: &ClientCapabilities, required: &ClientCapabilities) -> bool {
208 let actual = serde_json::to_value(actual).expect("ClientCapabilities is always serializable");
209 let mut required =
210 serde_json::to_value(required).expect("ClientCapabilities is always serializable");
211 if required.pointer("/roots/listChanged") == Some(&serde_json::Value::Bool(false))
216 && let Some(roots) = required
217 .get_mut("roots")
218 .and_then(serde_json::Value::as_object_mut)
219 {
220 roots.remove("listChanged");
221 }
222 json_value_contains(&actual, &required)
223}
224
225#[cfg(feature = "stateless")]
226fn validate_input_required_result(
227 extensions: &crate::context::Extensions,
228 result: &InputRequiredResult,
229) -> Result<()> {
230 result.validate().map_err(|message| {
231 Error::invalid_params(format!("invalid InputRequiredResult: {message}"))
232 })?;
233
234 let meta = extensions
235 .get::<crate::stateless::StatelessRequestMeta>()
236 .filter(|meta| {
237 meta.protocol_version.as_deref() == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
238 })
239 .ok_or_else(|| {
240 Error::invalid_params(
241 "InputRequiredResult is only supported by the 2026-07-28 request lifecycle",
242 )
243 })?;
244 let actual = meta.client_capabilities.as_ref().ok_or_else(|| {
245 Error::invalid_params("clientCapabilities is required for InputRequiredResult")
246 })?;
247
248 if let Some(requests) = &result.input_requests {
249 for request in requests.values() {
250 let (supported, required) = match request {
251 InputRequest::CreateMessage(params) => {
252 let requires_tools = params.tools.is_some();
253 let requires_context = params
254 .include_context
255 .is_some_and(|mode| mode != IncludeContext::None);
256 let required_sampling = SamplingCapability {
257 tools: requires_tools.then(SamplingToolsCapability::default),
258 context: requires_context.then(SamplingContextCapability::default),
259 ..SamplingCapability::default()
260 };
261 let supported = actual.sampling.as_ref().is_some_and(|sampling| {
262 (!requires_tools || sampling.tools.is_some())
263 && (!requires_context || sampling.context.is_some())
264 });
265 (
266 supported,
267 ClientCapabilities {
268 sampling: Some(required_sampling),
269 ..ClientCapabilities::default()
270 },
271 )
272 }
273 InputRequest::ListRoots(_) => (
274 actual.roots.is_some(),
275 ClientCapabilities {
276 roots: Some(RootsCapability::default()),
277 ..ClientCapabilities::default()
278 },
279 ),
280 InputRequest::Elicit(ElicitRequestParams::Form(_)) => {
281 let supported = actual.elicitation.as_ref().is_some_and(|elicitation| {
282 elicitation.form.is_some()
283 || (elicitation.form.is_none() && elicitation.url.is_none())
284 });
285 (
286 supported,
287 ClientCapabilities {
288 elicitation: Some(ElicitationCapability {
289 form: Some(ElicitationFormCapability::default()),
290 ..ElicitationCapability::default()
291 }),
292 ..ClientCapabilities::default()
293 },
294 )
295 }
296 InputRequest::Elicit(ElicitRequestParams::Url(_)) => (
297 actual
298 .elicitation
299 .as_ref()
300 .is_some_and(|elicitation| elicitation.url.is_some()),
301 ClientCapabilities {
302 elicitation: Some(ElicitationCapability {
303 url: Some(ElicitationUrlCapability::default()),
304 ..ElicitationCapability::default()
305 }),
306 ..ClientCapabilities::default()
307 },
308 ),
309 _ => {
310 return Err(Error::invalid_params(
311 "unsupported input request method in InputRequiredResult",
312 ));
313 }
314 };
315 if !supported {
316 return Err(Error::JsonRpc(
317 JsonRpcError::missing_required_client_capability(required),
318 ));
319 }
320 }
321 }
322 Ok(())
323}
324
325fn paginate<T>(
329 items: Vec<T>,
330 cursor: Option<&str>,
331 page_size: Option<usize>,
332) -> Result<(Vec<T>, Option<String>)> {
333 let Some(page_size) = page_size else {
334 return Ok((items, None));
335 };
336
337 let offset = match cursor {
338 Some(c) => decode_cursor(c)?,
339 None => 0,
340 };
341
342 if offset >= items.len() {
343 return Ok((Vec::new(), None));
344 }
345
346 let end = (offset + page_size).min(items.len());
347 let next_cursor = if end < items.len() {
348 Some(encode_cursor(end))
349 } else {
350 None
351 };
352
353 let mut items = items;
354 let page = items.drain(offset..end).collect();
355 Ok((page, next_cursor))
356}
357
358#[derive(Clone)]
382pub struct McpRouter {
383 inner: Arc<McpRouterInner>,
384 session: SessionState,
385}
386
387impl std::fmt::Debug for McpRouter {
388 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
389 f.debug_struct("McpRouter")
390 .field("server_name", &self.inner.server_name)
391 .field("server_version", &self.inner.server_version)
392 .field("tools_count", &self.inner.tools.len())
393 .field("resources_count", &self.inner.resources.len())
394 .field("prompts_count", &self.inner.prompts.len())
395 .field("session_phase", &self.session.phase())
396 .finish()
397 }
398}
399
400#[derive(Clone, Debug)]
402struct AutoInstructionsConfig {
403 prefix: Option<String>,
404 suffix: Option<String>,
405}
406
407#[cfg(all(feature = "http", feature = "stateless"))]
408type ModernNotificationSink = Arc<dyn Fn(&ServerNotification) -> bool + Send + Sync + 'static>;
409
410#[cfg(feature = "dynamic-tools")]
411type PromptInitializer = Arc<dyn Fn() -> Result<()> + Send + Sync + 'static>;
412
413#[derive(Clone)]
415struct McpRouterInner {
416 server_name: String,
417 server_version: String,
418 server_title: Option<String>,
420 server_description: Option<String>,
422 server_icons: Option<Vec<ToolIcon>>,
424 server_website_url: Option<String>,
426 instructions: Option<String>,
427 auto_instructions: Option<AutoInstructionsConfig>,
428 tools: HashMap<String, Arc<Tool>>,
429 resources: HashMap<String, Arc<Resource>>,
430 resource_templates: Vec<Arc<ResourceTemplate>>,
432 prompts: HashMap<String, Arc<Prompt>>,
433 in_flight: Arc<RwLock<HashMap<RequestId, CancellationToken>>>,
435 notification_tx: Option<NotificationSender>,
437 #[cfg(all(feature = "http", feature = "stateless"))]
442 modern_notification_sink: Arc<RwLock<Option<ModernNotificationSink>>>,
443 #[cfg(feature = "stateless")]
444 subscription_observer:
445 Arc<RwLock<Option<Arc<dyn crate::transport::subscriptions::SubscriptionObserver>>>>,
446 client_requester: Option<ClientRequesterHandle>,
448 task_store: Arc<dyn TaskStore>,
450 subscriptions: Arc<RwLock<HashSet<String>>>,
452 completion_handler: Option<CompletionHandler>,
454 tool_filter: Option<ToolFilter>,
456 resource_filter: Option<ResourceFilter>,
458 prompt_filter: Option<PromptFilter>,
460 extensions: Arc<crate::context::Extensions>,
462 protocol_extensions: HashMap<String, serde_json::Value>,
464 min_log_level: Arc<RwLock<LogLevel>>,
466 page_size: Option<usize>,
468 list_ttl_ms: Option<u64>,
472 read_ttl_ms: Option<u64>,
476 cache_scope: Option<CacheScope>,
481 logging_deprecated: Option<tower_mcp_types::protocol::DeprecationInfo>,
484 disabled_tools: Arc<RwLock<HashSet<String>>>,
486 disabled_resources: Arc<RwLock<HashSet<String>>>,
488 disabled_prompts: Arc<RwLock<HashSet<String>>>,
490 #[cfg(feature = "dynamic-tools")]
492 dynamic_tools: Option<Arc<DynamicToolsInner>>,
493 #[cfg(feature = "dynamic-tools")]
495 dynamic_prompts: Option<Arc<DynamicPromptsInner>>,
496 #[cfg(feature = "dynamic-tools")]
498 prompt_initializer: Option<PromptInitializer>,
499 #[cfg(feature = "dynamic-tools")]
501 dynamic_resources: Option<Arc<DynamicResourcesInner>>,
502 #[cfg(feature = "dynamic-tools")]
504 dynamic_resource_templates: Option<Arc<DynamicResourceTemplatesInner>>,
505}
506
507impl McpRouterInner {
508 fn generate_instructions(&self, config: &AutoInstructionsConfig) -> String {
510 let mut parts = Vec::new();
511
512 if let Some(prefix) = &config.prefix {
513 parts.push(prefix.clone());
514 }
515
516 if !self.tools.is_empty() {
518 let mut lines = vec!["## Tools".to_string(), String::new()];
519 let mut tools: Vec<_> = self.tools.values().collect();
520 tools.sort_by(|a, b| a.name.cmp(&b.name));
521 for tool in tools {
522 let desc = tool.description.as_deref().unwrap_or("No description");
523 let tags = annotation_tags(tool.annotations.as_ref());
524 if tags.is_empty() {
525 lines.push(format!("- **{}**: {}", tool.name, desc));
526 } else {
527 lines.push(format!("- **{}**: {} [{}]", tool.name, desc, tags));
528 }
529 }
530 parts.push(lines.join("\n"));
531 }
532
533 if !self.resources.is_empty() || !self.resource_templates.is_empty() {
535 let mut lines = vec!["## Resources".to_string(), String::new()];
536 let mut resources: Vec<_> = self.resources.values().collect();
537 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
538 for resource in resources {
539 let desc = resource.description.as_deref().unwrap_or("No description");
540 lines.push(format!("- **{}**: {}", resource.uri, desc));
541 }
542 let mut templates: Vec<_> = self.resource_templates.iter().collect();
543 templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
544 for template in templates {
545 let desc = template.description.as_deref().unwrap_or("No description");
546 lines.push(format!("- **{}**: {}", template.uri_template, desc));
547 }
548 parts.push(lines.join("\n"));
549 }
550
551 if !self.prompts.is_empty() {
553 let mut lines = vec!["## Prompts".to_string(), String::new()];
554 let mut prompts: Vec<_> = self.prompts.values().collect();
555 prompts.sort_by(|a, b| a.name.cmp(&b.name));
556 for prompt in prompts {
557 let desc = prompt.description.as_deref().unwrap_or("No description");
558 lines.push(format!("- **{}**: {}", prompt.name, desc));
559 }
560 parts.push(lines.join("\n"));
561 }
562
563 if let Some(suffix) = &config.suffix {
564 parts.push(suffix.clone());
565 }
566
567 parts.join("\n\n")
568 }
569}
570
571fn annotation_tags(annotations: Option<&crate::protocol::ToolAnnotations>) -> String {
577 let Some(ann) = annotations else {
578 return String::new();
579 };
580 let mut tags = Vec::new();
581 if ann.is_read_only() {
582 tags.push("read-only");
583 }
584 if ann.is_idempotent() {
585 tags.push("idempotent");
586 }
587 tags.join(", ")
588}
589
590impl McpRouter {
591 pub fn new() -> Self {
593 Self {
594 inner: Arc::new(McpRouterInner {
595 server_name: "tower-mcp".to_string(),
596 server_version: env!("CARGO_PKG_VERSION").to_string(),
597 server_title: None,
598 server_description: None,
599 server_icons: None,
600 server_website_url: None,
601 instructions: None,
602 auto_instructions: None,
603 tools: HashMap::new(),
604 resources: HashMap::new(),
605 resource_templates: Vec::new(),
606 prompts: HashMap::new(),
607 in_flight: Arc::new(RwLock::new(HashMap::new())),
608 notification_tx: None,
609 #[cfg(all(feature = "http", feature = "stateless"))]
610 modern_notification_sink: Arc::new(RwLock::new(None)),
611 #[cfg(feature = "stateless")]
612 subscription_observer: Arc::new(RwLock::new(None)),
613 client_requester: None,
614 task_store: Arc::new(MemoryTaskStore::new()),
615 subscriptions: Arc::new(RwLock::new(HashSet::new())),
616 extensions: Arc::new(crate::context::Extensions::new()),
617 protocol_extensions: HashMap::new(),
618 completion_handler: None,
619 tool_filter: None,
620 resource_filter: None,
621 prompt_filter: None,
622 min_log_level: Arc::new(RwLock::new(LogLevel::Debug)),
623 page_size: None,
624 list_ttl_ms: None,
625 read_ttl_ms: None,
626 cache_scope: None,
627 logging_deprecated: None,
628 disabled_tools: Arc::new(RwLock::new(HashSet::new())),
629 disabled_resources: Arc::new(RwLock::new(HashSet::new())),
630 disabled_prompts: Arc::new(RwLock::new(HashSet::new())),
631 #[cfg(feature = "dynamic-tools")]
632 dynamic_tools: None,
633 #[cfg(feature = "dynamic-tools")]
634 dynamic_prompts: None,
635 #[cfg(feature = "dynamic-tools")]
636 prompt_initializer: None,
637 #[cfg(feature = "dynamic-tools")]
638 dynamic_resources: None,
639 #[cfg(feature = "dynamic-tools")]
640 dynamic_resource_templates: None,
641 }),
642 session: SessionState::new(),
643 }
644 }
645
646 pub fn with_fresh_session(&self) -> Self {
654 Self {
655 inner: self.inner.clone(),
656 session: SessionState::new(),
657 }
658 }
659
660 pub fn tool_annotations_map(&self) -> ToolAnnotationsMap {
670 let disabled = self.inner.disabled_tools.read().unwrap();
671 let mut map = HashMap::new();
672 for (name, tool) in &self.inner.tools {
673 if disabled.contains(name) {
674 continue;
675 }
676 if let Some(annotations) = &tool.annotations {
677 map.insert(name.clone(), annotations.clone());
678 }
679 }
680 #[cfg(feature = "dynamic-tools")]
681 if let Some(dynamic) = &self.inner.dynamic_tools {
682 for tool in dynamic.list() {
683 if disabled.contains(&tool.name) {
684 continue;
685 }
686 if !map.contains_key(&tool.name)
688 && let Some(ref annotations) = tool.annotations
689 {
690 map.insert(tool.name.clone(), annotations.clone());
691 }
692 }
693 }
694 ToolAnnotationsMap { map: Arc::new(map) }
695 }
696
697 pub fn task_store(mut self, store: Arc<dyn TaskStore>) -> Self {
715 Arc::make_mut(&mut self.inner).task_store = store;
716 self
717 }
718
719 #[cfg(feature = "dynamic-tools")]
749 pub fn with_dynamic_tools(mut self) -> (Self, DynamicToolRegistry) {
750 let inner_dyn = Arc::new(DynamicToolsInner::new());
751 Arc::make_mut(&mut self.inner).dynamic_tools = Some(inner_dyn.clone());
752 (self, DynamicToolRegistry::new(inner_dyn))
753 }
754
755 #[cfg(feature = "dynamic-tools")]
778 pub fn with_dynamic_prompts(mut self) -> (Self, DynamicPromptRegistry) {
779 let inner_dyn = Arc::new(DynamicPromptsInner::new());
780 Arc::make_mut(&mut self.inner).dynamic_prompts = Some(inner_dyn.clone());
781 (self, DynamicPromptRegistry::new(inner_dyn))
782 }
783
784 #[cfg(feature = "dynamic-tools")]
790 pub fn dynamic_prompt_initializer<F>(mut self, initializer: F) -> Self
791 where
792 F: Fn() -> Result<()> + Send + Sync + 'static,
793 {
794 Arc::make_mut(&mut self.inner).prompt_initializer = Some(Arc::new(initializer));
795 self
796 }
797
798 #[cfg(feature = "dynamic-tools")]
821 pub fn with_dynamic_resources(mut self) -> (Self, DynamicResourceRegistry) {
822 let inner_dyn = Arc::new(DynamicResourcesInner::new());
823 Arc::make_mut(&mut self.inner).dynamic_resources = Some(inner_dyn.clone());
824 (self, DynamicResourceRegistry::new(inner_dyn))
825 }
826
827 #[cfg(feature = "dynamic-tools")]
849 pub fn with_dynamic_resource_templates(mut self) -> (Self, DynamicResourceTemplateRegistry) {
850 let inner_dyn = Arc::new(DynamicResourceTemplatesInner::new());
851 Arc::make_mut(&mut self.inner).dynamic_resource_templates = Some(inner_dyn.clone());
852 (self, DynamicResourceTemplateRegistry::new(inner_dyn))
853 }
854
855 #[cfg(feature = "stateless")]
863 #[cfg(feature = "http")]
864 pub(crate) fn with_request_notification_sender(mut self, tx: NotificationSender) -> Self {
865 Arc::make_mut(&mut self.inner).notification_tx = Some(tx);
866 self
867 }
868
869 pub fn with_notification_sender(mut self, tx: NotificationSender) -> Self {
873 let inner = Arc::make_mut(&mut self.inner);
874 #[cfg(feature = "dynamic-tools")]
877 if let Some(ref dynamic_tools) = inner.dynamic_tools {
878 dynamic_tools.add_notification_sender(tx.clone());
879 }
880 #[cfg(feature = "dynamic-tools")]
881 if let Some(ref dynamic_prompts) = inner.dynamic_prompts {
882 dynamic_prompts.add_notification_sender(tx.clone());
883 }
884 #[cfg(feature = "dynamic-tools")]
885 if let Some(ref dynamic_resources) = inner.dynamic_resources {
886 dynamic_resources.add_notification_sender(tx.clone());
887 }
888 #[cfg(feature = "dynamic-tools")]
889 if let Some(ref dynamic_resource_templates) = inner.dynamic_resource_templates {
890 dynamic_resource_templates.add_notification_sender(tx.clone());
891 }
892 inner.notification_tx = Some(tx);
893 self
894 }
895
896 #[cfg(feature = "stateless")]
903 pub fn with_subscription_observer(
904 self,
905 observer: Arc<dyn crate::transport::subscriptions::SubscriptionObserver>,
906 ) -> Self {
907 if let Ok(mut slot) = self.inner.subscription_observer.write() {
908 *slot = Some(observer);
909 }
910 self
911 }
912
913 #[cfg(feature = "stateless")]
915 pub(crate) fn subscription_observer(
916 &self,
917 ) -> Option<Arc<dyn crate::transport::subscriptions::SubscriptionObserver>> {
918 self.inner
919 .subscription_observer
920 .read()
921 .ok()
922 .and_then(|slot| slot.clone())
923 }
924
925 #[cfg(all(feature = "http", feature = "stateless"))]
927 pub(crate) fn attach_modern_notification_sink(&self, sink: ModernNotificationSink) {
928 if let Ok(mut active) = self.inner.modern_notification_sink.write() {
929 *active = Some(sink);
930 }
931 }
932
933 pub fn notification_sender(&self) -> Option<&NotificationSender> {
935 self.inner.notification_tx.as_ref()
936 }
937
938 pub fn with_client_requester(mut self, requester: ClientRequesterHandle) -> Self {
943 Arc::make_mut(&mut self.inner).client_requester = Some(requester);
944 self
945 }
946
947 pub fn client_requester(&self) -> Option<&ClientRequesterHandle> {
949 self.inner.client_requester.as_ref()
950 }
951
952 pub fn with_state<T: Clone + Send + Sync + 'static>(mut self, state: T) -> Self {
995 let inner = Arc::make_mut(&mut self.inner);
996 Arc::make_mut(&mut inner.extensions).insert(state);
997 self
998 }
999
1000 pub fn with_extension<T: Clone + Send + Sync + 'static>(self, value: T) -> Self {
1005 self.with_state(value)
1006 }
1007
1008 pub fn with_protocol_extension(mut self, extension: crate::ExtensionDeclaration) -> Self {
1015 let (identifier, settings) = extension.into_parts();
1016 Arc::make_mut(&mut self.inner)
1017 .protocol_extensions
1018 .insert(identifier, settings);
1019 self
1020 }
1021
1022 pub fn extensions(&self) -> &crate::context::Extensions {
1024 &self.inner.extensions
1025 }
1026
1027 pub fn create_context(
1032 &self,
1033 request_id: RequestId,
1034 progress_token: Option<ProgressToken>,
1035 ) -> RequestContext {
1036 self.create_context_with_extensions(request_id, progress_token, &Extensions::new())
1037 }
1038
1039 pub(crate) fn create_context_with_extensions(
1044 &self,
1045 request_id: RequestId,
1046 progress_token: Option<ProgressToken>,
1047 per_request: &Extensions,
1048 ) -> RequestContext {
1049 let ctx = RequestContext::new(request_id.clone());
1050
1051 let ctx = if let Some(token) = progress_token {
1053 ctx.with_progress_token(token)
1054 } else {
1055 ctx
1056 };
1057
1058 let ctx = if let Some(tx) = &self.inner.notification_tx {
1060 ctx.with_notification_sender(tx.clone())
1061 } else {
1062 ctx
1063 };
1064
1065 let mut merged = (*self.inner.extensions).clone();
1069 merged.merge(per_request);
1070 let negotiated_extensions = if is_final_protocol_request(per_request) {
1071 let server_capabilities =
1072 self.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
1073 final_client_capabilities(per_request)
1074 .map(|client_capabilities| {
1075 crate::NegotiatedExtensions::from_capabilities(
1076 client_capabilities,
1077 &server_capabilities,
1078 )
1079 })
1080 .unwrap_or_default()
1081 } else {
1082 self.session
1083 .get::<crate::NegotiatedExtensions>()
1084 .unwrap_or_default()
1085 };
1086 merged.insert(negotiated_extensions);
1087
1088 let final_lifecycle = is_final_protocol_request(per_request);
1093 let ctx = ctx.with_final_lifecycle(final_lifecycle);
1094 let ctx = if !final_lifecycle
1095 && let Some(requester) = merged
1096 .get::<ClientRequesterHandle>()
1097 .cloned()
1098 .or_else(|| self.inner.client_requester.clone())
1099 {
1100 ctx.with_client_requester(requester)
1101 } else {
1102 ctx
1103 };
1104
1105 let ctx = if let Some(token) = merged.get::<CancellationToken>() {
1109 ctx.with_cancellation_token(token.clone())
1110 } else {
1111 ctx
1112 };
1113
1114 let ctx = ctx.with_extensions(Arc::new(merged));
1115
1116 let ctx = ctx.with_min_log_level(self.inner.min_log_level.clone());
1118
1119 let token = ctx.cancellation_token();
1121 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1122 in_flight.insert(request_id, token);
1123 }
1124
1125 ctx
1126 }
1127
1128 pub fn complete_request(&self, request_id: &RequestId) {
1130 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1131 in_flight.remove(request_id);
1132 }
1133 }
1134
1135 fn cancel_request(&self, request_id: &RequestId) -> bool {
1137 let Ok(in_flight) = self.inner.in_flight.read() else {
1138 return false;
1139 };
1140 let Some(token) = in_flight.get(request_id) else {
1141 return false;
1142 };
1143 token.cancel();
1144 true
1145 }
1146
1147 pub fn server_info(mut self, name: impl Into<String>, version: impl Into<String>) -> Self {
1149 let inner = Arc::make_mut(&mut self.inner);
1150 inner.server_name = name.into();
1151 inner.server_version = version.into();
1152 self
1153 }
1154
1155 pub fn page_size(mut self, size: usize) -> Self {
1162 Arc::make_mut(&mut self.inner).page_size = Some(size);
1163 self
1164 }
1165
1166 pub fn list_ttl(mut self, ms: u64) -> Self {
1172 Arc::make_mut(&mut self.inner).list_ttl_ms = Some(ms);
1173 self
1174 }
1175
1176 pub fn read_ttl(mut self, ms: u64) -> Self {
1183 Arc::make_mut(&mut self.inner).read_ttl_ms = Some(ms);
1184 self
1185 }
1186
1187 pub fn cache_scope(mut self, scope: CacheScope) -> Self {
1196 Arc::make_mut(&mut self.inner).cache_scope = Some(scope);
1197 self
1198 }
1199
1200 pub fn logging_deprecated(mut self, info: tower_mcp_types::protocol::DeprecationInfo) -> Self {
1206 Arc::make_mut(&mut self.inner).logging_deprecated = Some(info);
1207 self
1208 }
1209
1210 pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
1212 Arc::make_mut(&mut self.inner).instructions = Some(instructions.into());
1213 self
1214 }
1215
1216 pub fn auto_instructions(mut self) -> Self {
1248 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1249 prefix: None,
1250 suffix: None,
1251 });
1252 self
1253 }
1254
1255 pub fn auto_instructions_with(
1272 mut self,
1273 prefix: Option<impl Into<String>>,
1274 suffix: Option<impl Into<String>>,
1275 ) -> Self {
1276 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1277 prefix: prefix.map(Into::into),
1278 suffix: suffix.map(Into::into),
1279 });
1280 self
1281 }
1282
1283 pub fn server_title(mut self, title: impl Into<String>) -> Self {
1285 Arc::make_mut(&mut self.inner).server_title = Some(title.into());
1286 self
1287 }
1288
1289 pub fn server_description(mut self, description: impl Into<String>) -> Self {
1291 Arc::make_mut(&mut self.inner).server_description = Some(description.into());
1292 self
1293 }
1294
1295 pub fn server_icons(mut self, icons: Vec<ToolIcon>) -> Self {
1297 Arc::make_mut(&mut self.inner).server_icons = Some(icons);
1298 self
1299 }
1300
1301 pub fn server_website_url(mut self, url: impl Into<String>) -> Self {
1303 Arc::make_mut(&mut self.inner).server_website_url = Some(url.into());
1304 self
1305 }
1306
1307 pub fn tool(mut self, tool: Tool) -> Self {
1309 Arc::make_mut(&mut self.inner)
1310 .tools
1311 .insert(tool.name.clone(), Arc::new(tool));
1312 self
1313 }
1314
1315 pub fn tool_if(self, condition: bool, tool: Tool) -> Self {
1341 if condition { self.tool(tool) } else { self }
1342 }
1343
1344 pub fn resource(mut self, resource: Resource) -> Self {
1346 Arc::make_mut(&mut self.inner)
1347 .resources
1348 .insert(resource.uri.clone(), Arc::new(resource));
1349 self
1350 }
1351
1352 pub fn resource_if(self, condition: bool, resource: Resource) -> Self {
1371 if condition {
1372 self.resource(resource)
1373 } else {
1374 self
1375 }
1376 }
1377
1378 pub fn resource_template(mut self, template: ResourceTemplate) -> Self {
1412 Arc::make_mut(&mut self.inner)
1413 .resource_templates
1414 .push(Arc::new(template));
1415 self
1416 }
1417
1418 pub fn prompt(mut self, prompt: Prompt) -> Self {
1420 Arc::make_mut(&mut self.inner)
1421 .prompts
1422 .insert(prompt.name.clone(), Arc::new(prompt));
1423 self
1424 }
1425
1426 pub fn prompt_if(self, condition: bool, prompt: Prompt) -> Self {
1445 if condition { self.prompt(prompt) } else { self }
1446 }
1447
1448 pub fn tools(self, tools: impl IntoIterator<Item = Tool>) -> Self {
1474 tools
1475 .into_iter()
1476 .fold(self, |router, tool| router.tool(tool))
1477 }
1478
1479 pub fn tools_if(self, condition: bool, tools: impl IntoIterator<Item = Tool>) -> Self {
1483 if condition { self.tools(tools) } else { self }
1484 }
1485
1486 pub fn resources(self, resources: impl IntoIterator<Item = Resource>) -> Self {
1505 resources
1506 .into_iter()
1507 .fold(self, |router, resource| router.resource(resource))
1508 }
1509
1510 pub fn resources_if(
1514 self,
1515 condition: bool,
1516 resources: impl IntoIterator<Item = Resource>,
1517 ) -> Self {
1518 if condition {
1519 self.resources(resources)
1520 } else {
1521 self
1522 }
1523 }
1524
1525 pub fn prompts(self, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1544 prompts
1545 .into_iter()
1546 .fold(self, |router, prompt| router.prompt(prompt))
1547 }
1548
1549 pub fn prompts_if(self, condition: bool, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1553 if condition {
1554 self.prompts(prompts)
1555 } else {
1556 self
1557 }
1558 }
1559
1560 pub fn merge(mut self, other: McpRouter) -> Self {
1605 let inner = Arc::make_mut(&mut self.inner);
1606 let other_inner = other.inner;
1607
1608 for (name, tool) in &other_inner.tools {
1610 inner.tools.insert(name.clone(), tool.clone());
1611 }
1612
1613 for (uri, resource) in &other_inner.resources {
1615 inner.resources.insert(uri.clone(), resource.clone());
1616 }
1617
1618 for template in &other_inner.resource_templates {
1621 inner.resource_templates.push(template.clone());
1622 }
1623
1624 for (name, prompt) in &other_inner.prompts {
1626 inner.prompts.insert(name.clone(), prompt.clone());
1627 }
1628
1629 for (identifier, settings) in &other_inner.protocol_extensions {
1631 inner
1632 .protocol_extensions
1633 .insert(identifier.clone(), settings.clone());
1634 }
1635
1636 self
1637 }
1638
1639 pub fn nest(mut self, prefix: impl Into<String>, other: McpRouter) -> Self {
1679 let prefix = prefix.into();
1680 let inner = Arc::make_mut(&mut self.inner);
1681 let other_inner = other.inner;
1682
1683 for tool in other_inner.tools.values() {
1685 let prefixed_tool = tool.with_name_prefix(&prefix);
1686 inner
1687 .tools
1688 .insert(prefixed_tool.name.clone(), Arc::new(prefixed_tool));
1689 }
1690
1691 for (uri, resource) in &other_inner.resources {
1693 inner.resources.insert(uri.clone(), resource.clone());
1694 }
1695
1696 for template in &other_inner.resource_templates {
1698 inner.resource_templates.push(template.clone());
1699 }
1700
1701 for (name, prompt) in &other_inner.prompts {
1703 inner.prompts.insert(name.clone(), prompt.clone());
1704 }
1705
1706 for (identifier, settings) in &other_inner.protocol_extensions {
1709 inner
1710 .protocol_extensions
1711 .insert(identifier.clone(), settings.clone());
1712 }
1713
1714 self
1715 }
1716
1717 pub fn completion_handler<F, Fut>(mut self, handler: F) -> Self
1745 where
1746 F: Fn(CompleteParams) -> Fut + Send + Sync + 'static,
1747 Fut: Future<Output = Result<CompleteResult>> + Send + 'static,
1748 {
1749 Arc::make_mut(&mut self.inner).completion_handler =
1750 Some(Arc::new(move |params| Box::pin(handler(params))));
1751 self
1752 }
1753
1754 pub fn tool_filter(mut self, filter: ToolFilter) -> Self {
1789 Arc::make_mut(&mut self.inner).tool_filter = Some(filter);
1790 self
1791 }
1792
1793 pub fn resource_filter(mut self, filter: ResourceFilter) -> Self {
1824 Arc::make_mut(&mut self.inner).resource_filter = Some(filter);
1825 self
1826 }
1827
1828 pub fn prompt_filter(mut self, filter: PromptFilter) -> Self {
1857 Arc::make_mut(&mut self.inner).prompt_filter = Some(filter);
1858 self
1859 }
1860
1861 pub fn session(&self) -> &SessionState {
1863 &self.session
1864 }
1865
1866 pub fn log(&self, params: LoggingMessageParams) -> bool {
1888 let Some(tx) = &self.inner.notification_tx else {
1889 return false;
1890 };
1891 tx.try_send(ServerNotification::LogMessage(params)).is_ok()
1892 }
1893
1894 pub fn log_info(&self, message: &str) -> bool {
1898 self.log(LoggingMessageParams::new(
1899 LogLevel::Info,
1900 serde_json::json!({ "message": message }),
1901 ))
1902 }
1903
1904 pub fn log_warning(&self, message: &str) -> bool {
1906 self.log(LoggingMessageParams::new(
1907 LogLevel::Warning,
1908 serde_json::json!({ "message": message }),
1909 ))
1910 }
1911
1912 pub fn log_error(&self, message: &str) -> bool {
1914 self.log(LoggingMessageParams::new(
1915 LogLevel::Error,
1916 serde_json::json!({ "message": message }),
1917 ))
1918 }
1919
1920 pub fn log_debug(&self, message: &str) -> bool {
1922 self.log(LoggingMessageParams::new(
1923 LogLevel::Debug,
1924 serde_json::json!({ "message": message }),
1925 ))
1926 }
1927
1928 pub fn is_subscribed(&self, uri: &str) -> bool {
1930 if let Ok(subs) = self.inner.subscriptions.read() {
1931 return subs.contains(uri);
1932 }
1933 false
1934 }
1935
1936 pub fn subscribed_uris(&self) -> Vec<String> {
1938 if let Ok(subs) = self.inner.subscriptions.read() {
1939 return subs.iter().cloned().collect();
1940 }
1941 Vec::new()
1942 }
1943
1944 fn subscribe(&self, uri: &str) -> bool {
1946 if let Ok(mut subs) = self.inner.subscriptions.write() {
1947 return subs.insert(uri.to_string());
1948 }
1949 false
1950 }
1951
1952 fn unsubscribe(&self, uri: &str) -> bool {
1954 if let Ok(mut subs) = self.inner.subscriptions.write() {
1955 return subs.remove(uri);
1956 }
1957 false
1958 }
1959
1960 pub fn notify_resource_updated(&self, uri: &str) -> bool {
1967 let notification = ServerNotification::ResourceUpdated {
1968 uri: uri.to_string(),
1969 };
1970 let mut sent = false;
1971
1972 if self.is_subscribed(uri)
1973 && let Some(tx) = &self.inner.notification_tx
1974 {
1975 sent |= tx.try_send(notification.clone()).is_ok();
1976 }
1977
1978 #[cfg(all(feature = "http", feature = "stateless"))]
1979 if let Ok(active) = self.inner.modern_notification_sink.read()
1980 && let Some(sink) = active.as_ref()
1981 {
1982 sent |= sink(¬ification);
1983 }
1984
1985 sent
1986 }
1987
1988 pub async fn notify_task_status_changed(&self, task_id: &str) {
2003 self.notify_task_state(task_id).await;
2004 }
2005
2006 pub fn notify_resources_list_changed(&self) -> bool {
2010 let Some(tx) = &self.inner.notification_tx else {
2011 return false;
2012 };
2013 tx.try_send(ServerNotification::ResourcesListChanged)
2014 .is_ok()
2015 }
2016
2017 pub fn notify_tools_list_changed(&self) -> bool {
2021 let Some(tx) = &self.inner.notification_tx else {
2022 return false;
2023 };
2024 tx.try_send(ServerNotification::ToolsListChanged).is_ok()
2025 }
2026
2027 pub fn notify_prompts_list_changed(&self) -> bool {
2031 let Some(tx) = &self.inner.notification_tx else {
2032 return false;
2033 };
2034 tx.try_send(ServerNotification::PromptsListChanged).is_ok()
2035 }
2036
2037 pub fn disable_tool(&self, name: impl Into<String>) {
2048 let mut set = self.inner.disabled_tools.write().unwrap();
2049 set.insert(name.into());
2050 }
2051
2052 pub fn enable_tool(&self, name: &str) {
2055 let mut set = self.inner.disabled_tools.write().unwrap();
2056 set.remove(name);
2057 }
2058
2059 pub fn is_tool_enabled(&self, name: &str) -> bool {
2063 !self.inner.disabled_tools.read().unwrap().contains(name)
2064 }
2065
2066 pub fn disable_resource(&self, uri: impl Into<String>) {
2069 let mut set = self.inner.disabled_resources.write().unwrap();
2070 set.insert(uri.into());
2071 }
2072
2073 pub fn enable_resource(&self, uri: &str) {
2075 let mut set = self.inner.disabled_resources.write().unwrap();
2076 set.remove(uri);
2077 }
2078
2079 pub fn is_resource_enabled(&self, uri: &str) -> bool {
2081 !self.inner.disabled_resources.read().unwrap().contains(uri)
2082 }
2083
2084 pub fn disable_prompt(&self, name: impl Into<String>) {
2087 let mut set = self.inner.disabled_prompts.write().unwrap();
2088 set.insert(name.into());
2089 }
2090
2091 pub fn enable_prompt(&self, name: &str) {
2093 let mut set = self.inner.disabled_prompts.write().unwrap();
2094 set.remove(name);
2095 }
2096
2097 pub fn is_prompt_enabled(&self, name: &str) -> bool {
2099 !self.inner.disabled_prompts.read().unwrap().contains(name)
2100 }
2101
2102 pub(crate) fn implementation(&self) -> Implementation {
2112 Implementation {
2113 name: self.inner.server_name.clone(),
2114 version: self.inner.server_version.clone(),
2115 title: self.inner.server_title.clone(),
2116 description: self.inner.server_description.clone(),
2117 icons: self.inner.server_icons.clone(),
2118 website_url: self.inner.server_website_url.clone(),
2119 meta: None,
2120 }
2121 }
2122
2123 #[cfg(feature = "http")]
2129 pub(crate) fn tool_input_schema(&self, name: &str) -> Option<serde_json::Value> {
2130 if let Some(tool) = self.inner.tools.get(name) {
2131 return Some(tool.input_schema.clone());
2132 }
2133 #[cfg(feature = "dynamic-tools")]
2134 if let Some(tool) = self
2135 .inner
2136 .dynamic_tools
2137 .as_ref()
2138 .and_then(|tools| tools.get(name))
2139 {
2140 return Some(tool.input_schema.clone());
2141 }
2142 None
2143 }
2144
2145 fn capabilities(&self) -> ServerCapabilities {
2146 let has_resources =
2147 !self.inner.resources.is_empty() || !self.inner.resource_templates.is_empty();
2148 let has_notifications = self.inner.notification_tx.is_some();
2149
2150 #[cfg(feature = "dynamic-tools")]
2151 let has_dynamic_tools = self.inner.dynamic_tools.is_some();
2152 #[cfg(not(feature = "dynamic-tools"))]
2153 let has_dynamic_tools = false;
2154
2155 #[cfg(feature = "dynamic-tools")]
2156 let has_dynamic_prompts = self.inner.dynamic_prompts.is_some();
2157 #[cfg(not(feature = "dynamic-tools"))]
2158 let has_dynamic_prompts = false;
2159
2160 #[cfg(feature = "dynamic-tools")]
2161 let has_dynamic_resources = self.inner.dynamic_resources.is_some()
2162 || self.inner.dynamic_resource_templates.is_some();
2163 #[cfg(not(feature = "dynamic-tools"))]
2164 let has_dynamic_resources = false;
2165
2166 ServerCapabilities {
2167 tools: if self.inner.tools.is_empty() && !has_dynamic_tools {
2168 None
2169 } else {
2170 Some(ToolsCapability {
2171 list_changed: has_notifications,
2172 })
2173 },
2174 resources: if has_resources || has_dynamic_resources {
2175 Some(ResourcesCapability {
2176 subscribe: true,
2177 list_changed: has_notifications,
2178 })
2179 } else {
2180 None
2181 },
2182 prompts: if self.inner.prompts.is_empty() && !has_dynamic_prompts {
2183 None
2184 } else {
2185 Some(PromptsCapability {
2186 list_changed: has_notifications,
2187 })
2188 },
2189 logging: if self.inner.notification_tx.is_some() {
2191 Some(LoggingCapability {
2192 deprecated: self.inner.logging_deprecated.clone(),
2193 })
2194 } else {
2195 None
2196 },
2197 tasks: {
2203 let has_task_support = self
2204 .inner
2205 .tools
2206 .values()
2207 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2208 if has_task_support {
2209 Some(TasksCapability {
2210 list: None,
2214 cancel: Some(TasksCancelCapability {}),
2215 requests: Some(TasksRequestsCapability {
2216 tools: Some(TasksToolsRequestsCapability {
2217 call: Some(TasksToolsCallCapability {}),
2218 }),
2219 }),
2220 })
2221 } else {
2222 None
2223 }
2224 },
2225 completions: if self.inner.completion_handler.is_some() {
2227 Some(CompletionsCapability::default())
2228 } else {
2229 None
2230 },
2231 experimental: None,
2232 extensions: {
2233 let mut map = self.inner.protocol_extensions.clone();
2234 let has_task_support = self
2235 .inner
2236 .tools
2237 .values()
2238 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2239 if has_task_support {
2240 map.insert(
2241 tower_mcp_types::protocol::TASKS_EXTENSION_ID.to_string(),
2242 serde_json::json!({}),
2243 );
2244 }
2245 (!map.is_empty()).then_some(map)
2246 },
2247 }
2248 }
2249
2250 fn capabilities_for_protocol(&self, protocol_version: Option<&str>) -> ServerCapabilities {
2258 let mut capabilities = self.capabilities();
2259 if protocol_version == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
2260 capabilities.tasks = None;
2261 if !self.final_tasks_enabled()
2262 && let Some(extensions) = capabilities.extensions.as_mut()
2263 {
2264 extensions.remove(tower_mcp_types::protocol::TASKS_EXTENSION_ID);
2265 if extensions.is_empty() {
2266 capabilities.extensions = None;
2267 }
2268 }
2269 }
2270 capabilities
2271 }
2272
2273 pub(crate) fn final_tasks_enabled(&self) -> bool {
2278 self.inner
2279 .protocol_extensions
2280 .contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
2281 }
2282
2283 fn require_negotiated_tasks(
2288 &self,
2289 extensions: &crate::context::Extensions,
2290 method: &str,
2291 ) -> Result<()> {
2292 if !self.final_tasks_enabled() {
2293 return Err(Error::JsonRpc(JsonRpcError::method_not_found(method)));
2294 }
2295 if client_declares_tasks(extensions) {
2296 return Ok(());
2297 }
2298 Err(Error::JsonRpc(
2299 JsonRpcError::missing_required_client_capability(tasks_client_capabilities()),
2300 ))
2301 }
2302
2303 async fn authorize_task(
2309 &self,
2310 task_id: &str,
2311 extensions: &crate::context::Extensions,
2312 ) -> Result<()> {
2313 let owner = self
2314 .inner
2315 .task_store
2316 .task_owner(task_id)
2317 .await
2318 .map_err(task_store_error)?
2319 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2320
2321 if crate::async_task::owner_matches(&owner, request_principal(extensions).as_deref()) {
2322 Ok(())
2323 } else {
2324 tracing::debug!(
2325 target: "mcp::tasks",
2326 task_id = %task_id,
2327 "task operation refused: principal does not own the task"
2328 );
2329 Err(Error::JsonRpc(unknown_task_error(task_id)))
2330 }
2331 }
2332
2333 async fn final_get_task(&self, task_id: &str) -> Result<McpResponse> {
2335 let (detailed, meta) = self.detailed_task(task_id).await?;
2336 let mut result = crate::tasks::GetTaskResult::new(detailed);
2337 result.meta = meta;
2338 Ok(McpResponse::FinalGetTask(result))
2339 }
2340
2341 async fn detailed_task(
2347 &self,
2348 task_id: &str,
2349 ) -> Result<(
2350 crate::tasks::DetailedTask,
2351 Option<serde_json::Map<String, serde_json::Value>>,
2352 )> {
2353 let (task, result, error) = self
2354 .inner
2355 .task_store
2356 .get_task_result(task_id)
2357 .await
2358 .map_err(task_store_error)?
2359 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2360
2361 let mut metadata = crate::tasks::TaskMetadata::new(
2362 task.task_id.clone(),
2363 task.created_at.clone(),
2364 task.last_updated_at.clone(),
2365 task.ttl,
2366 );
2367 metadata.status_message = task.status_message.clone();
2368 metadata.poll_interval_ms = task.poll_interval;
2369
2370 let meta = task.meta.and_then(|value| value.as_object().cloned());
2371 let detailed = match task.status {
2372 TaskStatus::Working => crate::tasks::DetailedTask::working(metadata),
2373 TaskStatus::InputRequired => {
2374 let outstanding = self
2377 .inner
2378 .task_store
2379 .outstanding_input_requests(task_id)
2380 .await
2381 .map_err(task_store_error)?
2382 .unwrap_or_default();
2383 crate::tasks::DetailedTask::input_required(metadata, outstanding)
2384 }
2385 TaskStatus::Completed => {
2386 let mut object = result
2389 .map(serde_json::to_value)
2390 .transpose()
2391 .map_err(|e| {
2392 Error::JsonRpc(JsonRpcError::internal_error(format!(
2393 "failed to encode task result: {e}"
2394 )))
2395 })?
2396 .and_then(|value| value.as_object().cloned())
2397 .unwrap_or_default();
2398 object.insert(
2402 "resultType".to_string(),
2403 serde_json::Value::String("complete".to_string()),
2404 );
2405 crate::tasks::DetailedTask::completed(metadata, object)
2406 }
2407 TaskStatus::Failed => crate::tasks::DetailedTask::failed(
2408 metadata,
2409 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
2410 ),
2411 TaskStatus::Cancelled => crate::tasks::DetailedTask::cancelled(metadata),
2412 _ => crate::tasks::DetailedTask::working(metadata),
2415 };
2416 Ok((detailed, meta))
2417 }
2418
2419 async fn park_task_for_input(
2425 &self,
2426 task_id: &str,
2427 input_required: crate::protocol::InputRequiredResult,
2428 ) {
2429 let requests = input_required.input_requests.unwrap_or_default();
2430 if requests.is_empty() {
2431 let error = JsonRpcError::internal_error(
2434 "handler asked for input without naming any requests, so the task has \
2435 nothing to wait for",
2436 );
2437 if let Err(e) = self.inner.task_store.fail_task(task_id, error).await {
2438 tracing::warn!(task_id = %task_id, error = %e, "failed to record task failure");
2439 }
2440 self.notify_task_state(task_id).await;
2441 return;
2442 }
2443
2444 if let Err(e) = self
2445 .inner
2446 .task_store
2447 .require_input(task_id, requests, input_required.request_state.as_deref())
2448 .await
2449 {
2450 tracing::warn!(task_id = %task_id, error = %e, "failed to park task for input");
2451 }
2452 self.notify_task_state(task_id).await;
2453 }
2454
2455 async fn resume_task(&self, task_id: &str) {
2463 let resume = match self.inner.task_store.resume_context(task_id).await {
2464 Ok(Some(resume)) => resume,
2465 Ok(None) => {
2466 let error = JsonRpcError::internal_error(
2470 "this task store cannot resume a task after input was provided; \
2471 implement TaskStore::resume_context to support handlers that ask \
2472 for input",
2473 );
2474 if let Err(e) = self.inner.task_store.fail_task(task_id, error).await {
2475 tracing::warn!(task_id = %task_id, error = %e, "failed to record task failure");
2476 }
2477 self.notify_task_state(task_id).await;
2478 return;
2479 }
2480 Err(e) => {
2481 tracing::warn!(task_id = %task_id, error = %e, "failed to read resume context");
2482 return;
2483 }
2484 };
2485
2486 let tool = self.inner.tools.get(&resume.tool_name).cloned();
2488 #[cfg(feature = "dynamic-tools")]
2489 let tool = tool.or_else(|| {
2490 self.inner
2491 .dynamic_tools
2492 .as_ref()
2493 .and_then(|d| d.get(&resume.tool_name))
2494 });
2495 let Some(tool) = tool else {
2496 let error = JsonRpcError::internal_error(format!(
2497 "tool '{}' is no longer registered, so the task cannot resume",
2498 resume.tool_name
2499 ));
2500 if let Err(e) = self.inner.task_store.fail_task(task_id, error).await {
2501 tracing::warn!(task_id = %task_id, error = %e, "failed to record task failure");
2502 }
2503 self.notify_task_state(task_id).await;
2504 return;
2505 };
2506
2507 let mut ctx = RequestContext::new(RequestId::String(task_id.to_string()));
2508 #[cfg(feature = "stateless")]
2513 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
2514 Some(resume.input_responses),
2515 None,
2516 ));
2517 if let Some(tx) = &self.inner.notification_tx {
2518 ctx = ctx.with_notification_sender(tx.clone());
2519 }
2520
2521 let task_id = task_id.to_string();
2522 let notifier = self.clone();
2523 let task_store = self.inner.task_store.clone();
2524 tokio::spawn(async move {
2525 let outcome = tool.call_outcome_with_context(ctx, resume.arguments).await;
2526 let result = match outcome {
2527 Ok(crate::protocol::RequestOutcome::Complete(result)) => result,
2528 Ok(crate::protocol::RequestOutcome::InputRequired(input_required)) => {
2531 notifier.park_task_for_input(&task_id, input_required).await;
2532 return;
2533 }
2534 Err(error) => CallToolResult::error(error.to_string()),
2535 };
2536
2537 if let Err(e) = task_store.complete_task(&task_id, result).await {
2538 tracing::warn!(task_id = %task_id, error = %e, "failed to record task completion");
2539 }
2540 notifier.notify_task_state(&task_id).await;
2541 });
2542 }
2543
2544 async fn notify_task_state(&self, task_id: &str) {
2553 if !self.final_tasks_enabled() {
2554 return;
2555 }
2556
2557 let (detailed, meta) = match self.detailed_task(task_id).await {
2558 Ok(detailed) => detailed,
2559 Err(error) => {
2560 tracing::debug!(
2561 target: "mcp::tasks",
2562 task_id = %task_id,
2563 %error,
2564 "skipping task notification: task state unavailable"
2565 );
2566 return;
2567 }
2568 };
2569
2570 let notification = ServerNotification::FinalTaskStatusChanged(
2571 crate::tasks::TaskStatusNotificationParams {
2572 task: detailed,
2573 meta,
2574 },
2575 );
2576
2577 #[cfg(all(feature = "http", feature = "stateless"))]
2582 if let Ok(active) = self.inner.modern_notification_sink.read()
2583 && let Some(sink) = active.as_ref()
2584 {
2585 sink(¬ification);
2586 return;
2587 }
2588
2589 if let Some(tx) = &self.inner.notification_tx {
2590 let _ = tx.try_send(notification);
2591 }
2592 }
2593
2594 fn effective_cache_scope(&self, ttl_ms: Option<u64>) -> Option<CacheScope> {
2601 self.inner
2602 .cache_scope
2603 .or_else(|| ttl_ms.map(|_| CacheScope::Private))
2604 }
2605
2606 fn apply_read_cache_hints(&self, mut result: ReadResourceResult) -> ReadResourceResult {
2611 if result.ttl_ms.is_none() {
2612 result.ttl_ms = self.inner.read_ttl_ms;
2613 }
2614 if result.cache_scope.is_none() {
2615 result.cache_scope = self.effective_cache_scope(result.ttl_ms);
2616 }
2617 result
2618 }
2619
2620 async fn handle(
2622 &self,
2623 request_id: RequestId,
2624 request: McpRequest,
2625 extensions: Extensions,
2626 ) -> Result<McpResponse> {
2627 let method = request.method_name();
2629 if !is_final_protocol_request(&extensions) && !self.session.is_request_allowed(method) {
2630 tracing::warn!(
2631 method = %method,
2632 phase = ?self.session.phase(),
2633 "Request rejected: session not initialized"
2634 );
2635 return Err(Error::JsonRpc(JsonRpcError::invalid_request(format!(
2636 "Session not initialized. Only 'initialize' and 'ping' are allowed before initialization. Got: {}",
2637 method
2638 ))));
2639 }
2640
2641 match request {
2642 McpRequest::Initialize(params) => {
2643 tracing::info!(
2644 client = %params.client_info.name,
2645 version = %params.client_info.version,
2646 "Client initializing"
2647 );
2648
2649 let protocol_support = extensions.get::<crate::ProtocolSupport>();
2653 let requested_is_legacy = crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2654 .contains(¶ms.protocol_version.as_str());
2655 let requested_is_supported = requested_is_legacy
2656 && protocol_support
2657 .is_none_or(|support| support.contains(¶ms.protocol_version));
2658 let protocol_version = if requested_is_supported {
2659 params.protocol_version
2660 } else {
2661 match protocol_support {
2662 None => crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
2663 Some(support) => support
2664 .versions()
2665 .iter()
2666 .find(|version| {
2667 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2668 .contains(&version.as_str())
2669 })
2670 .cloned()
2671 .ok_or_else(|| {
2672 Error::JsonRpc(JsonRpcError::unsupported_protocol_version(
2673 params.protocol_version,
2674 support.versions().iter().map(String::as_str),
2675 ))
2676 })?,
2677 }
2678 };
2679
2680 self.session.mark_initializing();
2682 let capabilities = self.capabilities_for_protocol(Some(&protocol_version));
2683 self.session.insert(params.capabilities.clone());
2684 self.session
2685 .insert(crate::NegotiatedExtensions::from_capabilities(
2686 ¶ms.capabilities,
2687 &capabilities,
2688 ));
2689
2690 Ok(McpResponse::Initialize(InitializeResult {
2691 protocol_version,
2692 capabilities,
2693 server_info: self.implementation(),
2694 instructions: if let Some(config) = &self.inner.auto_instructions {
2695 Some(self.inner.generate_instructions(config))
2696 } else {
2697 self.inner.instructions.clone()
2698 },
2699 meta: None,
2700 }))
2701 }
2702
2703 McpRequest::Discover(_) => {
2704 tracing::debug!("Stateless server/discover request");
2711 let server_info = self.implementation();
2712 let supported_versions = extensions.get::<crate::ProtocolSupport>().map_or_else(
2713 || {
2714 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2715 .iter()
2716 .map(|version| (*version).to_string())
2717 .collect()
2718 },
2719 |support| support.versions().to_vec(),
2720 );
2721 let capabilities = self
2726 .capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
2727 Ok(McpResponse::Discover(DiscoverResult {
2728 supported_versions,
2729 capabilities,
2730 ttl_ms: None,
2731 cache_scope: None,
2732 instructions: if let Some(config) = &self.inner.auto_instructions {
2733 Some(self.inner.generate_instructions(config))
2734 } else {
2735 self.inner.instructions.clone()
2736 },
2737 meta: Some(crate::protocol::ResultMeta {
2738 server_info: Some(server_info),
2739 }),
2740 }))
2741 }
2742
2743 McpRequest::ListTools(params) => {
2744 let final_protocol = is_final_protocol_request(&extensions);
2745 let final_tasks_negotiated = final_protocol
2746 && self.final_tasks_enabled()
2747 && client_declares_tasks(&extensions);
2748 let filter = self.inner.tool_filter.as_ref();
2749 let disabled = self.inner.disabled_tools.read().unwrap().clone();
2750 let is_visible = |t: &Tool| {
2751 !disabled.contains(&t.name)
2752 && !(final_protocol
2753 && matches!(t.task_support, TaskSupportMode::Required)
2754 && !final_tasks_negotiated)
2755 && filter
2756 .map(|f| f.is_visible(&self.session, t))
2757 .unwrap_or(true)
2758 };
2759 let definition = |t: &Tool| {
2760 let mut definition = t.definition();
2761 if final_protocol {
2762 definition.execution = None;
2763 }
2764 definition
2765 };
2766
2767 let mut tools: Vec<ToolDefinition> = self
2769 .inner
2770 .tools
2771 .values()
2772 .filter(|t| is_visible(t))
2773 .map(|t| definition(t))
2774 .collect();
2775
2776 #[cfg(feature = "dynamic-tools")]
2778 if let Some(ref dynamic) = self.inner.dynamic_tools {
2779 let static_names: HashSet<String> =
2780 tools.iter().map(|t| t.name.clone()).collect();
2781 for t in dynamic.list() {
2782 if !static_names.contains(&t.name) && is_visible(&t) {
2783 tools.push(definition(&t));
2784 }
2785 }
2786 }
2787
2788 tools.sort_by(|a, b| a.name.cmp(&b.name));
2789
2790 let (tools, next_cursor) =
2791 paginate(tools, params.cursor.as_deref(), self.inner.page_size)?;
2792
2793 Ok(McpResponse::ListTools(ListToolsResult {
2794 tools,
2795 next_cursor,
2796 ttl_ms: self.inner.list_ttl_ms,
2797 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
2798 meta: None,
2799 }))
2800 }
2801
2802 McpRequest::CallTool(params) => {
2803 if self
2805 .inner
2806 .disabled_tools
2807 .read()
2808 .unwrap()
2809 .contains(¶ms.name)
2810 {
2811 tracing::info!(
2812 target: "mcp::tools",
2813 tool = %params.name,
2814 status = "disabled",
2815 "tool call completed"
2816 );
2817 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2818 }
2819
2820 let tool = self.inner.tools.get(¶ms.name).cloned();
2822 #[cfg(feature = "dynamic-tools")]
2823 let tool = tool.or_else(|| {
2824 self.inner
2825 .dynamic_tools
2826 .as_ref()
2827 .and_then(|d| d.get(¶ms.name))
2828 });
2829
2830 let tool = match tool {
2831 Some(t) => t,
2832 None => {
2833 tracing::info!(
2834 target: "mcp::tools",
2835 tool = %params.name,
2836 status = "not_found",
2837 "tool call completed"
2838 );
2839 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2840 }
2841 };
2842
2843 if let Some(filter) = &self.inner.tool_filter
2845 && !filter.is_visible(&self.session, &tool)
2846 {
2847 tracing::info!(
2848 target: "mcp::tools",
2849 tool = %params.name,
2850 status = "denied",
2851 "tool call completed"
2852 );
2853 return Err(filter.denial_error(¶ms.name));
2854 }
2855
2856 let final_protocol = is_final_protocol_request(&extensions);
2860 let task_ttl = if final_protocol {
2861 if params.task.is_some() {
2862 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
2863 "The final Tasks extension does not allow a 'task' request parameter",
2864 )));
2865 }
2866
2867 let server_enabled = self.final_tasks_enabled();
2868 let tasks_negotiated = server_enabled && client_declares_tasks(&extensions);
2869 match tool.task_support {
2870 TaskSupportMode::Required if !server_enabled => {
2871 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
2874 ¶ms.name,
2875 )));
2876 }
2877 TaskSupportMode::Required if !tasks_negotiated => {
2878 return Err(Error::JsonRpc(
2879 JsonRpcError::missing_required_client_capability(
2880 tasks_client_capabilities(),
2881 ),
2882 ));
2883 }
2884 TaskSupportMode::Required | TaskSupportMode::Optional
2885 if tasks_negotiated =>
2886 {
2887 Some(None)
2888 }
2889 _ => None,
2890 }
2891 } else {
2892 match (¶ms.task, tool.task_support) {
2893 (Some(_), TaskSupportMode::Forbidden) => {
2894 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2895 "Tool '{}' does not support async tasks",
2896 params.name
2897 ))));
2898 }
2899 (None, TaskSupportMode::Required) => {
2900 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2901 "Tool '{}' requires async task execution (include 'task' in params)",
2902 params.name
2903 ))));
2904 }
2905 (Some(task), _) => Some(task.ttl),
2906 (None, _) => None,
2907 }
2908 };
2909
2910 #[cfg(feature = "stateless")]
2914 if let Some(required) = tool.required_client_capabilities()
2915 && let Some(meta) = extensions.get::<crate::stateless::StatelessRequestMeta>()
2916 && meta.protocol_version.as_deref()
2917 == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
2918 && !meta
2919 .client_capabilities
2920 .as_ref()
2921 .is_some_and(|actual| client_capabilities_satisfy(actual, required))
2922 {
2923 return Err(Error::JsonRpc(
2924 JsonRpcError::missing_required_client_capability(required.clone()),
2925 ));
2926 }
2927
2928 if let Some(task_ttl) = task_ttl {
2929 let (task_id, cancellation_token) = self
2931 .inner
2932 .task_store
2933 .create_task(
2934 ¶ms.name,
2935 params.arguments.clone(),
2936 task_ttl,
2937 request_principal(&extensions),
2938 )
2939 .await
2940 .map_err(task_store_error)?;
2941
2942 tracing::info!(task_id = %task_id, tool = %params.name, "Created async task");
2943
2944 let progress_token = params.meta.and_then(|m| m.progress_token);
2946 let ctx = self.create_context_with_extensions(
2947 request_id,
2948 progress_token,
2949 &extensions,
2950 );
2951
2952 let task_store = self.inner.task_store.clone();
2953 let task_context = crate::tool::TaskContext::new(task_id.clone());
2954 let mut ctx = ctx;
2955 ctx.extensions_mut().insert(task_context.clone());
2956 let preparation = match tool
2957 .prepare_task(task_context, params.arguments.clone())
2958 .await
2959 {
2960 Ok(preparation) => preparation,
2961 Err(error) => {
2962 discard_unprepared_task(&task_store, &task_id).await;
2963 return Err(error);
2964 }
2965 };
2966 if let Some(meta) = preparation.meta {
2967 let value = serde_json::Value::Object(meta);
2968 if let Err(error) = crate::protocol::validate_meta_object(&value) {
2969 discard_unprepared_task(&task_store, &task_id).await;
2970 return Err(Error::invalid_params(format!(
2971 "Invalid task metadata: {error}"
2972 )));
2973 }
2974 let persisted = match task_store.set_task_meta(&task_id, value).await {
2975 Ok(persisted) => persisted,
2976 Err(error) => {
2977 discard_unprepared_task(&task_store, &task_id).await;
2978 return Err(task_store_error(error));
2979 }
2980 };
2981 if !persisted {
2982 discard_unprepared_task(&task_store, &task_id).await;
2983 return Err(Error::JsonRpc(JsonRpcError::internal_error(
2984 "Task store could not persist preparation metadata",
2985 )));
2986 }
2987 }
2988 ctx.extensions_mut().merge(&preparation.extensions);
2989
2990 let tool = tool.clone();
2992 let arguments = params.arguments;
2993 let task_id_clone = task_id.clone();
2994
2995 let tool_name = params.name.clone();
2996 let notifier = self.clone();
2997 tokio::spawn(async move {
2998 if cancellation_token.is_cancelled() {
3000 tracing::debug!(task_id = %task_id_clone, "Task cancelled before execution");
3001 notifier.notify_task_state(&task_id_clone).await;
3002 return;
3003 }
3004
3005 let start = std::time::Instant::now();
3012 let outcome = tool.call_outcome_with_context(ctx, arguments).await;
3013 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
3014
3015 let result = match outcome {
3016 Ok(crate::protocol::RequestOutcome::Complete(result)) => result,
3017 Ok(crate::protocol::RequestOutcome::InputRequired(input_required)) => {
3018 notifier
3019 .park_task_for_input(&task_id_clone, input_required)
3020 .await;
3021 return;
3022 }
3023 Err(error) => CallToolResult::error(error.to_string()),
3027 };
3028
3029 if cancellation_token.is_cancelled() {
3030 tracing::debug!(task_id = %task_id_clone, "Task cancelled during execution");
3031 notifier.notify_task_state(&task_id_clone).await;
3032 } else {
3033 let status = if result.is_error { "error" } else { "success" };
3038 let error_msg = result
3039 .is_error
3040 .then(|| result.first_text().unwrap_or("Tool execution failed"))
3041 .map(str::to_string);
3042 if let Err(e) = task_store.complete_task(&task_id_clone, result).await {
3043 tracing::warn!(task_id = %task_id_clone, error = %e, "failed to record task completion");
3044 }
3045 tracing::info!(
3046 target: "mcp::tools",
3047 tool = %tool_name,
3048 task_id = %task_id_clone,
3049 duration_ms,
3050 status,
3051 error = error_msg.as_deref().unwrap_or_default(),
3052 "tool call completed"
3053 );
3054 notifier.notify_task_state(&task_id_clone).await;
3055 }
3056 });
3057
3058 let task = self
3059 .inner
3060 .task_store
3061 .get_task(&task_id)
3062 .await
3063 .map_err(task_store_error)?
3064 .ok_or_else(|| {
3065 Error::JsonRpc(JsonRpcError::internal_error(
3066 "Failed to retrieve created task",
3067 ))
3068 })?;
3069
3070 if is_final_protocol_request(&extensions) {
3074 let mut metadata = crate::tasks::TaskMetadata::new(
3075 task.task_id.clone(),
3076 task.created_at.clone(),
3077 task.last_updated_at.clone(),
3078 task.ttl,
3079 );
3080 metadata.status_message = task.status_message.clone();
3081 metadata.poll_interval_ms = task.poll_interval;
3082 let mut result = crate::tasks::CreateTaskResult::new(
3083 crate::tasks::Task::new(metadata, task.status),
3084 );
3085 result.meta = task.meta.and_then(|value| value.as_object().cloned());
3086 return Ok(McpResponse::FinalCreateTask(result));
3087 }
3088 Ok(McpResponse::CreateTask(CreateTaskResult::new(task)))
3089 } else {
3090 let progress_token = params.meta.and_then(|m| m.progress_token);
3092 let ctx = self.create_context_with_extensions(
3093 request_id,
3094 progress_token,
3095 &extensions,
3096 );
3097 #[cfg(feature = "stateless")]
3098 let ctx = {
3099 let mut ctx = ctx;
3100 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3101 params.input_responses,
3102 params.request_state,
3103 ));
3104 ctx
3105 };
3106
3107 let start = std::time::Instant::now();
3108 let outcome = tool
3109 .call_outcome_with_context(ctx, params.arguments)
3110 .await?;
3111 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
3112
3113 match outcome {
3114 RequestOutcome::Complete(result) => {
3115 let status = if result.is_error { "error" } else { "success" };
3116 tracing::info!(
3117 target: "mcp::tools",
3118 tool = %params.name,
3119 duration_ms,
3120 status,
3121 "tool call completed"
3122 );
3123 Ok(McpResponse::CallTool(result))
3124 }
3125 RequestOutcome::InputRequired(result) => {
3126 #[cfg(feature = "stateless")]
3127 {
3128 validate_input_required_result(&extensions, &result)?;
3129 tracing::info!(
3130 target: "mcp::tools",
3131 tool = %params.name,
3132 duration_ms,
3133 status = "input_required",
3134 "tool call requires client input"
3135 );
3136 Ok(McpResponse::InputRequired(result))
3137 }
3138 #[cfg(not(feature = "stateless"))]
3139 {
3140 let _ = result;
3141 Err(Error::invalid_params(
3142 "InputRequiredResult support was not compiled",
3143 ))
3144 }
3145 }
3146 }
3147 }
3148 }
3149
3150 McpRequest::ListResources(params) => {
3151 let disabled = self.inner.disabled_resources.read().unwrap().clone();
3152 let is_visible = |r: &Resource| -> bool {
3153 !disabled.contains(&r.uri)
3154 && self
3155 .inner
3156 .resource_filter
3157 .as_ref()
3158 .map(|f| f.is_visible(&self.session, r))
3159 .unwrap_or(true)
3160 };
3161
3162 let mut resources: Vec<ResourceDefinition> = self
3163 .inner
3164 .resources
3165 .values()
3166 .filter(|r| is_visible(r))
3167 .map(|r| r.definition())
3168 .collect();
3169
3170 #[cfg(feature = "dynamic-tools")]
3172 if let Some(ref dynamic) = self.inner.dynamic_resources {
3173 let static_uris: HashSet<String> =
3174 resources.iter().map(|r| r.uri.clone()).collect();
3175 for r in dynamic.list() {
3176 if !static_uris.contains(&r.uri) && is_visible(&r) {
3177 resources.push(r.definition());
3178 }
3179 }
3180 }
3181
3182 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
3183
3184 let (resources, next_cursor) =
3185 paginate(resources, params.cursor.as_deref(), self.inner.page_size)?;
3186
3187 Ok(McpResponse::ListResources(ListResourcesResult {
3188 resources,
3189 next_cursor,
3190 ttl_ms: self.inner.list_ttl_ms,
3191 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3192 meta: None,
3193 }))
3194 }
3195
3196 McpRequest::ListResourceTemplates(params) => {
3197 let mut resource_templates: Vec<ResourceTemplateDefinition> = self
3198 .inner
3199 .resource_templates
3200 .iter()
3201 .map(|t| t.definition())
3202 .collect();
3203
3204 #[cfg(feature = "dynamic-tools")]
3206 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3207 let static_patterns: HashSet<String> = resource_templates
3208 .iter()
3209 .map(|t| t.uri_template.clone())
3210 .collect();
3211 for t in dynamic.list() {
3212 if !static_patterns.contains(&t.uri_template) {
3213 resource_templates.push(t.definition());
3214 }
3215 }
3216 }
3217
3218 resource_templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
3219
3220 let (resource_templates, next_cursor) = paginate(
3221 resource_templates,
3222 params.cursor.as_deref(),
3223 self.inner.page_size,
3224 )?;
3225
3226 Ok(McpResponse::ListResourceTemplates(
3227 ListResourceTemplatesResult {
3228 resource_templates,
3229 next_cursor,
3230 ttl_ms: self.inner.list_ttl_ms,
3231 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3232 meta: None,
3233 },
3234 ))
3235 }
3236
3237 McpRequest::ReadResource(params) => {
3238 if self
3240 .inner
3241 .disabled_resources
3242 .read()
3243 .unwrap()
3244 .contains(¶ms.uri)
3245 {
3246 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3247 ¶ms.uri,
3248 )));
3249 }
3250
3251 if let Some(resource) = self.inner.resources.get(¶ms.uri) {
3253 if let Some(filter) = &self.inner.resource_filter
3255 && !filter.is_visible(&self.session, resource)
3256 {
3257 return Err(filter.denial_error(¶ms.uri));
3258 }
3259
3260 tracing::debug!(uri = %params.uri, "Reading static resource");
3261 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3262 #[cfg(feature = "stateless")]
3263 let ctx = {
3264 let mut ctx = ctx;
3265 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3266 params.input_responses.clone(),
3267 params.request_state.clone(),
3268 ));
3269 ctx
3270 };
3271 return match resource.read_outcome_with_context(ctx).await? {
3272 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3273 self.apply_read_cache_hints(result),
3274 )),
3275 RequestOutcome::InputRequired(result) => {
3276 #[cfg(feature = "stateless")]
3277 {
3278 validate_input_required_result(&extensions, &result)?;
3279 Ok(McpResponse::InputRequired(result))
3280 }
3281 #[cfg(not(feature = "stateless"))]
3282 {
3283 let _ = result;
3284 Err(Error::invalid_params(
3285 "InputRequiredResult support was not compiled",
3286 ))
3287 }
3288 }
3289 };
3290 }
3291
3292 #[cfg(feature = "dynamic-tools")]
3294 #[allow(clippy::collapsible_if)]
3295 if let Some(ref dynamic) = self.inner.dynamic_resources {
3296 if let Some(resource) = dynamic.get(¶ms.uri) {
3297 if let Some(filter) = &self.inner.resource_filter
3298 && !filter.is_visible(&self.session, &resource)
3299 {
3300 return Err(filter.denial_error(¶ms.uri));
3301 }
3302 tracing::debug!(uri = %params.uri, "Reading dynamic resource");
3303 let ctx =
3304 self.create_context_with_extensions(request_id, None, &extensions);
3305 #[cfg(feature = "stateless")]
3306 let ctx = {
3307 let mut ctx = ctx;
3308 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3309 params.input_responses.clone(),
3310 params.request_state.clone(),
3311 ));
3312 ctx
3313 };
3314 return match resource.read_outcome_with_context(ctx).await? {
3315 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3316 self.apply_read_cache_hints(result),
3317 )),
3318 RequestOutcome::InputRequired(result) => {
3319 #[cfg(feature = "stateless")]
3320 {
3321 validate_input_required_result(&extensions, &result)?;
3322 Ok(McpResponse::InputRequired(result))
3323 }
3324 #[cfg(not(feature = "stateless"))]
3325 {
3326 let _ = result;
3327 Err(Error::invalid_params(
3328 "InputRequiredResult support was not compiled",
3329 ))
3330 }
3331 }
3332 };
3333 }
3334 }
3335
3336 for template in &self.inner.resource_templates {
3338 if let Some(variables) = template.match_uri(¶ms.uri) {
3339 tracing::debug!(
3340 uri = %params.uri,
3341 template = %template.uri_template,
3342 "Reading resource via template"
3343 );
3344 let ctx =
3345 self.create_context_with_extensions(request_id, None, &extensions);
3346 #[cfg(feature = "stateless")]
3347 let ctx = {
3348 let mut ctx = ctx;
3349 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3350 params.input_responses.clone(),
3351 params.request_state.clone(),
3352 ));
3353 ctx
3354 };
3355 return match template
3356 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3357 .await?
3358 {
3359 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3360 self.apply_read_cache_hints(result),
3361 )),
3362 RequestOutcome::InputRequired(result) => {
3363 #[cfg(feature = "stateless")]
3364 {
3365 validate_input_required_result(&extensions, &result)?;
3366 Ok(McpResponse::InputRequired(result))
3367 }
3368 #[cfg(not(feature = "stateless"))]
3369 {
3370 let _ = result;
3371 Err(Error::invalid_params(
3372 "InputRequiredResult support was not compiled",
3373 ))
3374 }
3375 }
3376 };
3377 }
3378 }
3379
3380 #[cfg(feature = "dynamic-tools")]
3382 #[allow(clippy::collapsible_if)]
3383 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3384 if let Some((template, variables)) = dynamic.match_uri(¶ms.uri) {
3385 tracing::debug!(
3386 uri = %params.uri,
3387 template = %template.uri_template,
3388 "Reading resource via dynamic template"
3389 );
3390 let ctx =
3391 self.create_context_with_extensions(request_id, None, &extensions);
3392 #[cfg(feature = "stateless")]
3393 let ctx = {
3394 let mut ctx = ctx;
3395 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3396 params.input_responses.clone(),
3397 params.request_state.clone(),
3398 ));
3399 ctx
3400 };
3401 return match template
3402 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3403 .await?
3404 {
3405 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3406 self.apply_read_cache_hints(result),
3407 )),
3408 RequestOutcome::InputRequired(result) => {
3409 #[cfg(feature = "stateless")]
3410 {
3411 validate_input_required_result(&extensions, &result)?;
3412 Ok(McpResponse::InputRequired(result))
3413 }
3414 #[cfg(not(feature = "stateless"))]
3415 {
3416 let _ = result;
3417 Err(Error::invalid_params(
3418 "InputRequiredResult support was not compiled",
3419 ))
3420 }
3421 }
3422 };
3423 }
3424 }
3425
3426 Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3428 ¶ms.uri,
3429 )))
3430 }
3431
3432 McpRequest::SubscribeResource(params) => {
3433 if !self.inner.resources.contains_key(¶ms.uri) {
3435 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3436 ¶ms.uri,
3437 )));
3438 }
3439
3440 tracing::debug!(uri = %params.uri, "Subscribing to resource");
3441 self.subscribe(¶ms.uri);
3442
3443 Ok(McpResponse::SubscribeResource(EmptyResult {}))
3444 }
3445
3446 McpRequest::UnsubscribeResource(params) => {
3447 if !self.inner.resources.contains_key(¶ms.uri) {
3449 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3450 ¶ms.uri,
3451 )));
3452 }
3453
3454 tracing::debug!(uri = %params.uri, "Unsubscribing from resource");
3455 self.unsubscribe(¶ms.uri);
3456
3457 Ok(McpResponse::UnsubscribeResource(EmptyResult {}))
3458 }
3459
3460 McpRequest::ListPrompts(params) => {
3461 #[cfg(feature = "dynamic-tools")]
3462 if let Some(initializer) = &self.inner.prompt_initializer {
3463 initializer()?;
3464 }
3465 let disabled = self.inner.disabled_prompts.read().unwrap().clone();
3466 let is_visible = |p: &Prompt| -> bool {
3467 !disabled.contains(&p.name)
3468 && self
3469 .inner
3470 .prompt_filter
3471 .as_ref()
3472 .map(|f| f.is_visible(&self.session, p))
3473 .unwrap_or(true)
3474 };
3475
3476 let mut prompts: Vec<PromptDefinition> = self
3477 .inner
3478 .prompts
3479 .values()
3480 .filter(|p| is_visible(p))
3481 .map(|p| p.definition())
3482 .collect();
3483
3484 #[cfg(feature = "dynamic-tools")]
3486 if let Some(ref dynamic) = self.inner.dynamic_prompts {
3487 let static_names: HashSet<String> =
3488 prompts.iter().map(|p| p.name.clone()).collect();
3489 for p in dynamic.list() {
3490 if !static_names.contains(&p.name) && is_visible(&p) {
3491 prompts.push(p.definition());
3492 }
3493 }
3494 }
3495
3496 prompts.sort_by(|a, b| a.name.cmp(&b.name));
3497
3498 let (prompts, next_cursor) =
3499 paginate(prompts, params.cursor.as_deref(), self.inner.page_size)?;
3500
3501 Ok(McpResponse::ListPrompts(ListPromptsResult {
3502 prompts,
3503 next_cursor,
3504 ttl_ms: self.inner.list_ttl_ms,
3505 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3506 meta: None,
3507 }))
3508 }
3509
3510 McpRequest::GetPrompt(params) => {
3511 #[cfg(feature = "dynamic-tools")]
3512 if let Some(initializer) = &self.inner.prompt_initializer {
3513 initializer()?;
3514 }
3515 if self
3517 .inner
3518 .disabled_prompts
3519 .read()
3520 .unwrap()
3521 .contains(¶ms.name)
3522 {
3523 return Err(Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3524 "Prompt not found: {}",
3525 params.name
3526 ))));
3527 }
3528
3529 let prompt = self.inner.prompts.get(¶ms.name).cloned();
3531 #[cfg(feature = "dynamic-tools")]
3532 let prompt = prompt.or_else(|| {
3533 self.inner
3534 .dynamic_prompts
3535 .as_ref()
3536 .and_then(|d| d.get(¶ms.name))
3537 });
3538 let prompt = prompt.ok_or_else(|| {
3539 Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3540 "Prompt not found: {}",
3541 params.name
3542 )))
3543 })?;
3544
3545 if let Some(filter) = &self.inner.prompt_filter
3547 && !filter.is_visible(&self.session, &prompt)
3548 {
3549 return Err(filter.denial_error(¶ms.name));
3550 }
3551
3552 tracing::debug!(name = %params.name, "Getting prompt");
3553 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3554 #[cfg(feature = "stateless")]
3555 let ctx = {
3556 let mut ctx = ctx;
3557 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3558 params.input_responses,
3559 params.request_state,
3560 ));
3561 ctx
3562 };
3563 let outcome = prompt
3564 .get_outcome_with_context(ctx, params.arguments)
3565 .await?;
3566
3567 match outcome {
3568 RequestOutcome::Complete(result) => Ok(McpResponse::GetPrompt(result)),
3569 RequestOutcome::InputRequired(result) => {
3570 #[cfg(feature = "stateless")]
3571 {
3572 validate_input_required_result(&extensions, &result)?;
3573 Ok(McpResponse::InputRequired(result))
3574 }
3575 #[cfg(not(feature = "stateless"))]
3576 {
3577 let _ = result;
3578 Err(Error::invalid_params(
3579 "InputRequiredResult support was not compiled",
3580 ))
3581 }
3582 }
3583 }
3584 }
3585
3586 McpRequest::Ping => Ok(McpResponse::Pong(EmptyResult {})),
3587
3588 McpRequest::GetTaskInfo(params) => {
3589 if is_final_protocol_request(&extensions) {
3590 self.require_negotiated_tasks(&extensions, "tasks/get")?;
3591 self.authorize_task(¶ms.task_id, &extensions).await?;
3592 return self.final_get_task(¶ms.task_id).await;
3593 }
3594 self.authorize_task(¶ms.task_id, &extensions).await?;
3595
3596 let (mut task, result, error) = self
3603 .inner
3604 .task_store
3605 .get_task_result(¶ms.task_id)
3606 .await
3607 .map_err(task_store_error)?
3608 .ok_or_else(|| {
3609 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3610 "Task not found: {}",
3611 params.task_id
3612 )))
3613 })?;
3614
3615 match task.status {
3616 TaskStatus::Completed => task.result = result,
3617 TaskStatus::Failed => {
3618 task.error = Some(
3622 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
3623 );
3624 }
3625 _ => {}
3626 }
3627
3628 Ok(McpResponse::GetTaskInfo(task))
3629 }
3630
3631 McpRequest::UpdateTask(params) => {
3632 if is_final_protocol_request(&extensions) {
3633 self.require_negotiated_tasks(&extensions, "tasks/update")?;
3634 self.authorize_task(¶ms.task_id, &extensions).await?;
3635 let applied = self
3639 .inner
3640 .task_store
3641 .apply_input_responses(
3642 ¶ms.task_id,
3643 decode_input_responses(¶ms.input_responses),
3644 )
3645 .await
3646 .map_err(task_store_error)?
3647 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3648 self.notify_task_state(¶ms.task_id).await;
3652 if applied.is_complete() {
3656 self.resume_task(¶ms.task_id).await;
3657 }
3658 return Ok(McpResponse::FinalTaskAck(
3659 crate::tasks::TaskAcknowledgement::new(),
3660 ));
3661 }
3662
3663 self.authorize_task(¶ms.task_id, &extensions).await?;
3664
3665 self.inner
3672 .task_store
3673 .apply_input_responses(
3674 ¶ms.task_id,
3675 decode_input_responses(¶ms.input_responses),
3676 )
3677 .await
3678 .map_err(task_store_error)?
3679 .ok_or_else(|| {
3680 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3681 "Task not found: {}",
3682 params.task_id
3683 )))
3684 })?;
3685 self.notify_task_state(¶ms.task_id).await;
3689 Ok(McpResponse::UpdateTask(EmptyResult {}))
3690 }
3691
3692 McpRequest::CancelTask(params) => {
3693 if is_final_protocol_request(&extensions) {
3694 self.require_negotiated_tasks(&extensions, "tasks/cancel")?;
3695 self.authorize_task(¶ms.task_id, &extensions).await?;
3696 self.inner
3700 .task_store
3701 .cancel_task(¶ms.task_id, params.reason.as_deref())
3702 .await
3703 .map_err(task_store_error)?
3704 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3705 self.notify_task_state(¶ms.task_id).await;
3706 return Ok(McpResponse::FinalTaskAck(
3707 crate::tasks::TaskAcknowledgement::new(),
3708 ));
3709 }
3710
3711 self.authorize_task(¶ms.task_id, &extensions).await?;
3712
3713 let current = self
3715 .inner
3716 .task_store
3717 .get_task(¶ms.task_id)
3718 .await
3719 .map_err(task_store_error)?
3720 .ok_or_else(|| {
3721 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3722 "Task not found: {}",
3723 params.task_id
3724 )))
3725 })?;
3726
3727 if current.status.is_terminal() {
3728 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
3729 "Task {} is already in terminal state: {}",
3730 params.task_id, current.status
3731 ))));
3732 }
3733
3734 self.inner
3735 .task_store
3736 .cancel_task(¶ms.task_id, params.reason.as_deref())
3737 .await
3738 .map_err(task_store_error)?
3739 .ok_or_else(|| {
3740 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3741 "Task not found: {}",
3742 params.task_id
3743 )))
3744 })?;
3745
3746 Ok(McpResponse::CancelTask(EmptyResult {}))
3750 }
3751
3752 McpRequest::SetLoggingLevel(params) => {
3753 tracing::debug!(level = ?params.level, "Client set logging level");
3754 if let Ok(mut level) = self.inner.min_log_level.write() {
3755 *level = params.level;
3756 }
3757 Ok(McpResponse::SetLoggingLevel(EmptyResult {}))
3758 }
3759
3760 McpRequest::Complete(params) => {
3761 tracing::debug!(
3762 reference = ?params.reference,
3763 argument = %params.argument.name,
3764 "Completion request"
3765 );
3766
3767 if let Some(ref handler) = self.inner.completion_handler {
3769 let result = handler(params).await?;
3770 Ok(McpResponse::Complete(result))
3771 } else {
3772 Ok(McpResponse::Complete(CompleteResult::new(vec![])))
3774 }
3775 }
3776
3777 #[cfg(feature = "stateless")]
3778 McpRequest::SubscriptionsListen(params) => {
3779 if !is_final_protocol_request(&extensions) {
3786 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
3789 "subscriptions/listen",
3790 )));
3791 }
3792 let Some(requested) = params.notifications else {
3793 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
3794 "subscriptions/listen requires a notifications filter",
3795 )));
3796 };
3797 if requested.task_ids.is_some() && !client_declares_tasks(&extensions) {
3800 return Err(Error::JsonRpc(
3801 JsonRpcError::missing_required_client_capability(
3802 tasks_client_capabilities(),
3803 ),
3804 ));
3805 }
3806 let notifications = crate::transport::subscriptions::accepted_subscription_filter(
3807 requested,
3808 self.final_tasks_enabled(),
3809 );
3810 Ok(McpResponse::SubscriptionsAccepted(
3811 crate::protocol::SubscriptionsAcceptedResult { notifications },
3812 ))
3813 }
3814
3815 McpRequest::Unknown { method, .. } => {
3816 Err(Error::JsonRpc(JsonRpcError::method_not_found(&method)))
3817 }
3818 _ => Err(Error::JsonRpc(JsonRpcError::method_not_found(
3819 "unknown method",
3820 ))),
3821 }
3822 }
3823
3824 pub fn handle_notification(&self, notification: McpNotification) {
3826 match notification {
3827 McpNotification::Initialized => {
3828 let phase_before = self.session.phase();
3829 if self.session.mark_initialized() {
3830 if phase_before == crate::session::SessionPhase::Uninitialized {
3831 tracing::info!(
3832 "Session initialized from uninitialized state (race resolved)"
3833 );
3834 } else {
3835 tracing::info!("Session initialized, entering operation phase");
3836 }
3837 } else {
3838 tracing::warn!(
3839 phase = ?self.session.phase(),
3840 "Received initialized notification in unexpected state"
3841 );
3842 }
3843 }
3844 McpNotification::Cancelled(params) => {
3845 if let Some(ref request_id) = params.request_id {
3846 if self.cancel_request(request_id) {
3847 tracing::info!(
3848 request_id = ?request_id,
3849 reason = ?params.reason,
3850 "Request cancelled"
3851 );
3852 } else {
3853 tracing::debug!(
3854 request_id = ?request_id,
3855 reason = ?params.reason,
3856 "Cancellation requested for unknown request"
3857 );
3858 }
3859 } else {
3860 tracing::debug!(
3861 reason = ?params.reason,
3862 "Cancellation notification received without request_id"
3863 );
3864 }
3865 }
3866 McpNotification::Progress(params) => {
3867 tracing::trace!(
3868 token = ?params.progress_token,
3869 progress = params.progress,
3870 total = ?params.total,
3871 "Progress notification"
3872 );
3873 }
3881 McpNotification::RootsListChanged => {
3882 tracing::info!("Client roots list changed");
3883 }
3886 McpNotification::Unknown { method, .. } => {
3887 tracing::debug!(method = %method, "Unknown notification received");
3888 }
3889 _ => {
3890 tracing::debug!("Unrecognized notification variant received");
3891 }
3892 }
3893 }
3894}
3895
3896impl Default for McpRouter {
3897 fn default() -> Self {
3898 Self::new()
3899 }
3900}
3901
3902pub use crate::context::Extensions;
3908
3909#[derive(Debug, Clone)]
3934pub struct ToolAnnotationsMap {
3935 map: Arc<HashMap<String, ToolAnnotations>>,
3936}
3937
3938impl ToolAnnotationsMap {
3939 pub fn get(&self, tool_name: &str) -> Option<&ToolAnnotations> {
3943 self.map.get(tool_name)
3944 }
3945
3946 pub fn is_read_only(&self, tool_name: &str) -> bool {
3951 self.map.get(tool_name).is_some_and(|a| a.read_only_hint)
3952 }
3953
3954 pub fn is_destructive(&self, tool_name: &str) -> bool {
3959 self.map.get(tool_name).is_none_or(|a| a.destructive_hint)
3960 }
3961
3962 pub fn is_idempotent(&self, tool_name: &str) -> bool {
3967 self.map.get(tool_name).is_some_and(|a| a.idempotent_hint)
3968 }
3969}
3970
3971#[derive(Debug, Clone)]
3993pub struct RouterRequest {
3994 pub id: RequestId,
3996 pub inner: McpRequest,
3998 pub extensions: Extensions,
4000}
4001
4002impl RouterRequest {
4003 pub fn new(id: RequestId, inner: McpRequest) -> Self {
4005 Self {
4006 id,
4007 inner,
4008 extensions: Extensions::new(),
4009 }
4010 }
4011
4012 pub fn with_inner(self, inner: McpRequest) -> Self {
4018 Self {
4019 id: self.id,
4020 inner,
4021 extensions: self.extensions,
4022 }
4023 }
4024
4025 pub fn with_id_and_inner(self, id: RequestId, inner: McpRequest) -> Self {
4031 Self {
4032 id,
4033 inner,
4034 extensions: self.extensions,
4035 }
4036 }
4037
4038 pub fn clone_with_inner(&self, inner: McpRequest) -> Self {
4046 Self {
4047 id: self.id.clone(),
4048 inner,
4049 extensions: self.extensions.clone(),
4050 }
4051 }
4052}
4053
4054#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
4056pub struct RouterResponse {
4057 pub id: RequestId,
4059 pub inner: std::result::Result<McpResponse, JsonRpcError>,
4061}
4062
4063impl RouterResponse {
4064 pub fn is_error(&self) -> bool {
4080 self.inner.is_err()
4081 }
4082
4083 pub fn into_jsonrpc(self) -> JsonRpcResponse {
4085 match self.inner {
4086 Ok(response) => match serde_json::to_value(response) {
4087 Ok(result) => JsonRpcResponse::result(self.id, result),
4088 Err(e) => {
4089 tracing::error!(error = %e, "Failed to serialize response");
4090 JsonRpcResponse::error(
4091 Some(self.id),
4092 JsonRpcError::internal_error(format!("Serialization error: {}", e)),
4093 )
4094 }
4095 },
4096 Err(error) => JsonRpcResponse::error(Some(self.id), error),
4097 }
4098 }
4099}
4100
4101impl Service<RouterRequest> for McpRouter {
4102 type Response = RouterResponse;
4103 type Error = std::convert::Infallible; type Future =
4105 Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
4106
4107 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
4108 Poll::Ready(Ok(()))
4109 }
4110
4111 fn call(&mut self, req: RouterRequest) -> Self::Future {
4112 let router = self.clone();
4113 let request_id = req.id.clone();
4114 Box::pin(async move {
4115 let result = router.handle(req.id, req.inner, req.extensions).await;
4116 router.complete_request(&request_id);
4118 Ok(RouterResponse {
4119 id: request_id,
4120 inner: result.map_err(|e| match e {
4125 Error::JsonRpc(err) => err,
4126 Error::Tool(err) => JsonRpcError::internal_error(err.to_string()),
4127 e => JsonRpcError::internal_error(e.to_string()),
4128 }),
4129 })
4130 })
4131 }
4132}
4133
4134#[cfg(test)]
4135mod tests {
4136 use super::*;
4137 use crate::extract::{Context, Json};
4138 use crate::jsonrpc::JsonRpcService;
4139 use crate::tool::ToolBuilder;
4140 use schemars::JsonSchema;
4141 use serde::Deserialize;
4142 use tower::ServiceExt;
4143
4144 #[derive(Debug, Deserialize, JsonSchema)]
4145 struct AddInput {
4146 a: i64,
4147 b: i64,
4148 }
4149
4150 #[cfg(feature = "stateless")]
4151 fn final_extensions(client_capabilities: ClientCapabilities) -> Extensions {
4152 let mut extensions = Extensions::new();
4153 extensions.insert(crate::stateless::StatelessRequestMeta {
4154 protocol_version: Some(PROTOCOL_VERSION_2026_07_28.to_string()),
4155 client_capabilities: Some(client_capabilities),
4156 ..Default::default()
4157 });
4158 extensions
4159 }
4160
4161 #[cfg(feature = "stateless")]
4162 fn tasks_client_extensions() -> Extensions {
4163 final_extensions(ClientCapabilities {
4164 extensions: Some(
4165 [(TASKS_EXTENSION_ID.to_string(), serde_json::json!({}))]
4166 .into_iter()
4167 .collect(),
4168 ),
4169 ..Default::default()
4170 })
4171 }
4172
4173 #[cfg(feature = "stateless")]
4174 #[tokio::test]
4175 async fn final_tasks_require_server_opt_in_and_client_declaration() {
4176 let tool = || {
4177 ToolBuilder::new("optional_task")
4178 .task_support(TaskSupportMode::Optional)
4179 .handler(|input: AddInput| async move {
4180 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4181 })
4182 .build()
4183 };
4184 let task_params = |task| CallToolParams {
4185 name: "optional_task".to_string(),
4186 arguments: serde_json::json!({"a": 1, "b": 2}),
4187 input_responses: None,
4188 request_state: None,
4189 meta: None,
4190 task,
4191 };
4192
4193 let implicit = McpRouter::new().tool(tool());
4197 let McpResponse::Discover(result) = implicit
4198 .handle(
4199 RequestId::Number(1),
4200 McpRequest::Discover(DiscoverParams::default()),
4201 Extensions::new(),
4202 )
4203 .await
4204 .unwrap()
4205 else {
4206 panic!("Expected Discover response");
4207 };
4208 assert!(
4209 result
4210 .capabilities
4211 .extensions
4212 .as_ref()
4213 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4214 );
4215 let error = implicit
4216 .handle(
4217 RequestId::Number(2),
4218 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4219 tasks_client_extensions(),
4220 )
4221 .await
4222 .unwrap_err();
4223 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4224
4225 let router = McpRouter::new().tool(tool()).with_tasks();
4227 let McpResponse::Discover(result) = router
4228 .handle(
4229 RequestId::Number(3),
4230 McpRequest::Discover(DiscoverParams::default()),
4231 Extensions::new(),
4232 )
4233 .await
4234 .unwrap()
4235 else {
4236 panic!("Expected Discover response");
4237 };
4238 assert!(
4239 result
4240 .capabilities
4241 .extensions
4242 .as_ref()
4243 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID)),
4244 "with_tasks() must advertise the extension on the final path"
4245 );
4246 assert!(
4247 result.capabilities.tasks.is_none(),
4248 "the legacy capability shape is never advertised on the final path"
4249 );
4250
4251 let response = router
4254 .handle(
4255 RequestId::Number(4),
4256 McpRequest::CallTool(task_params(None)),
4257 final_extensions(ClientCapabilities::default()),
4258 )
4259 .await
4260 .unwrap();
4261 assert!(matches!(response, McpResponse::CallTool(_)));
4262
4263 let response = router
4266 .handle(
4267 RequestId::Number(5),
4268 McpRequest::CallTool(task_params(None)),
4269 tasks_client_extensions(),
4270 )
4271 .await
4272 .unwrap();
4273 assert!(
4274 matches!(response, McpResponse::FinalCreateTask(_)),
4275 "a negotiated request must receive a task, got {response:?}"
4276 );
4277
4278 let error = router
4281 .handle(
4282 RequestId::Number(6),
4283 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4284 tasks_client_extensions(),
4285 )
4286 .await
4287 .unwrap_err();
4288 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4289 }
4290
4291 #[cfg(feature = "stateless")]
4292 #[tokio::test]
4293 async fn final_task_methods_serve_the_negotiated_wire_shapes() {
4294 let router = McpRouter::new()
4295 .tool(
4296 ToolBuilder::new("optional_task")
4297 .task_support(TaskSupportMode::Optional)
4298 .handler(|input: AddInput| async move {
4299 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4300 })
4301 .task_preparation(|task, _input| async move {
4302 let mut meta = serde_json::Map::new();
4303 meta.insert(
4304 "dev.tower-mcp/owner-test".to_string(),
4305 serde_json::json!({"taskId": task.task_id()}),
4306 );
4307 Ok(crate::TaskPreparation::new().with_meta(meta))
4308 })
4309 .build(),
4310 )
4311 .with_tasks();
4312
4313 let McpResponse::FinalCreateTask(created) = router
4314 .handle(
4315 RequestId::Number(1),
4316 McpRequest::CallTool(CallToolParams {
4317 name: "optional_task".to_string(),
4318 arguments: serde_json::json!({"a": 1, "b": 2}),
4319 input_responses: None,
4320 request_state: None,
4321 meta: None,
4322 task: None,
4323 }),
4324 tasks_client_extensions(),
4325 )
4326 .await
4327 .unwrap()
4328 else {
4329 panic!("Expected a final create-task response");
4330 };
4331
4332 let wire = serde_json::to_value(&created).unwrap();
4334 assert_eq!(wire["resultType"], "task");
4335 assert!(wire.get("task").is_none(), "final results are flat: {wire}");
4336 assert!(wire["ttlMs"].is_number() || wire["ttlMs"].is_null());
4337 assert!(wire.get("ttl").is_none(), "legacy field name leaked");
4338 let task_id = created.task.metadata.task_id.clone();
4339 assert_eq!(
4340 created.meta.as_ref().unwrap()["dev.tower-mcp/owner-test"]["taskId"],
4341 task_id
4342 );
4343
4344 let McpResponse::FinalGetTask(fetched) = router
4346 .handle(
4347 RequestId::Number(2),
4348 McpRequest::GetTaskInfo(GetTaskInfoParams {
4349 task_id: task_id.clone(),
4350 meta: None,
4351 }),
4352 tasks_client_extensions(),
4353 )
4354 .await
4355 .unwrap()
4356 else {
4357 panic!("Expected a final get-task response");
4358 };
4359 let wire = serde_json::to_value(&fetched).unwrap();
4360 assert_eq!(wire["resultType"], "complete");
4361 assert_eq!(wire["taskId"], serde_json::json!(task_id));
4362 assert!(wire["status"].is_string());
4363
4364 for (id, request) in [
4366 (
4367 3,
4368 McpRequest::UpdateTask(UpdateTaskParams {
4369 task_id: task_id.clone(),
4370 input_responses: HashMap::new(),
4371 meta: None,
4372 }),
4373 ),
4374 (
4375 4,
4376 McpRequest::CancelTask(CancelTaskParams {
4377 task_id: task_id.clone(),
4378 reason: None,
4379 meta: None,
4380 }),
4381 ),
4382 ] {
4383 let response = router
4384 .handle(RequestId::Number(id), request, tasks_client_extensions())
4385 .await
4386 .unwrap();
4387 let McpResponse::FinalTaskAck(ack) = response else {
4388 panic!("Expected a final ack for request {id}");
4389 };
4390 assert_eq!(
4391 serde_json::to_value(&ack).unwrap(),
4392 serde_json::json!({"resultType": "complete"})
4393 );
4394 }
4395
4396 let error = router
4398 .handle(
4399 RequestId::Number(5),
4400 McpRequest::GetTaskInfo(GetTaskInfoParams {
4401 task_id: "does-not-exist".to_string(),
4402 meta: None,
4403 }),
4404 tasks_client_extensions(),
4405 )
4406 .await
4407 .unwrap_err();
4408 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4409
4410 let error = router
4412 .handle(
4413 RequestId::Number(6),
4414 McpRequest::GetTaskInfo(GetTaskInfoParams {
4415 task_id: task_id.clone(),
4416 meta: None,
4417 }),
4418 final_extensions(ClientCapabilities::default()),
4419 )
4420 .await
4421 .unwrap_err();
4422 let Error::JsonRpc(error) = error else {
4423 panic!("expected a JSON-RPC error");
4424 };
4425 assert_eq!(error.code, -32021);
4426 assert_eq!(
4427 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4428 serde_json::json!({})
4429 );
4430 }
4431
4432 #[cfg(feature = "stateless")]
4433 #[tokio::test]
4434 async fn final_required_task_tools_follow_per_request_capabilities() {
4435 let router = McpRouter::new()
4436 .tool(
4437 ToolBuilder::new("required_task")
4438 .task_support(TaskSupportMode::Required)
4439 .handler(|input: AddInput| async move {
4440 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4441 })
4442 .build(),
4443 )
4444 .with_tasks();
4445 let params = || CallToolParams {
4446 name: "required_task".to_string(),
4447 arguments: serde_json::json!({"a": 1, "b": 2}),
4448 input_responses: None,
4449 request_state: None,
4450 meta: None,
4451 task: None,
4452 };
4453
4454 let McpResponse::ListTools(without_tasks) = router
4455 .handle(
4456 RequestId::Number(1),
4457 McpRequest::ListTools(ListToolsParams::default()),
4458 final_extensions(ClientCapabilities::default()),
4459 )
4460 .await
4461 .unwrap()
4462 else {
4463 panic!("expected tools/list")
4464 };
4465 assert!(without_tasks.tools.is_empty());
4466
4467 let McpResponse::ListTools(with_tasks) = router
4468 .handle(
4469 RequestId::Number(2),
4470 McpRequest::ListTools(ListToolsParams::default()),
4471 tasks_client_extensions(),
4472 )
4473 .await
4474 .unwrap()
4475 else {
4476 panic!("expected tools/list")
4477 };
4478 assert_eq!(with_tasks.tools.len(), 1);
4479 assert!(with_tasks.tools[0].execution.is_none());
4480
4481 let error = router
4482 .handle(
4483 RequestId::Number(3),
4484 McpRequest::CallTool(params()),
4485 final_extensions(ClientCapabilities::default()),
4486 )
4487 .await
4488 .unwrap_err();
4489 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32021));
4490
4491 let response = router
4492 .handle(
4493 RequestId::Number(4),
4494 McpRequest::CallTool(params()),
4495 tasks_client_extensions(),
4496 )
4497 .await
4498 .unwrap();
4499 assert!(matches!(response, McpResponse::FinalCreateTask(_)));
4500 }
4501
4502 #[cfg(all(feature = "oauth", feature = "stateless"))]
4503 #[tokio::test]
4504 async fn task_operations_are_bound_to_the_creating_principal() {
4505 fn as_principal(subject: &str) -> Extensions {
4506 let mut extensions = tasks_client_extensions();
4507 extensions.insert(crate::oauth::token::TokenClaims {
4508 sub: Some(subject.to_string()),
4509 iss: None,
4510 aud: None,
4511 exp: None,
4512 scope: None,
4513 client_id: None,
4514 extra: HashMap::new(),
4515 });
4516 extensions
4517 }
4518
4519 let router = McpRouter::new()
4520 .tool(
4521 ToolBuilder::new("optional_task")
4522 .task_support(TaskSupportMode::Optional)
4523 .handler(|input: AddInput| async move {
4524 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4525 })
4526 .build(),
4527 )
4528 .with_tasks();
4529
4530 let McpResponse::FinalCreateTask(created) = router
4531 .handle(
4532 RequestId::Number(1),
4533 McpRequest::CallTool(CallToolParams {
4534 name: "optional_task".to_string(),
4535 arguments: serde_json::json!({"a": 1, "b": 2}),
4536 input_responses: None,
4537 request_state: None,
4538 meta: None,
4539 task: None,
4540 }),
4541 as_principal("alice"),
4542 )
4543 .await
4544 .unwrap()
4545 else {
4546 panic!("Expected a final create-task response");
4547 };
4548 let task_id = created.task.metadata.task_id.clone();
4549
4550 assert!(
4552 router
4553 .handle(
4554 RequestId::Number(2),
4555 McpRequest::GetTaskInfo(GetTaskInfoParams {
4556 task_id: task_id.clone(),
4557 meta: None,
4558 }),
4559 as_principal("alice"),
4560 )
4561 .await
4562 .is_ok()
4563 );
4564
4565 for (id, label, context) in [
4568 (3, "another principal", as_principal("bob")),
4569 (4, "no principal", tasks_client_extensions()),
4570 ] {
4571 for (offset, request) in [
4572 McpRequest::GetTaskInfo(GetTaskInfoParams {
4573 task_id: task_id.clone(),
4574 meta: None,
4575 }),
4576 McpRequest::UpdateTask(UpdateTaskParams {
4577 task_id: task_id.clone(),
4578 input_responses: HashMap::new(),
4579 meta: None,
4580 }),
4581 McpRequest::CancelTask(CancelTaskParams {
4582 task_id: task_id.clone(),
4583 reason: None,
4584 meta: None,
4585 }),
4586 ]
4587 .into_iter()
4588 .enumerate()
4589 {
4590 let error = router
4591 .handle(
4592 RequestId::Number(id * 10 + offset as i64),
4593 request,
4594 context.clone(),
4595 )
4596 .await
4597 .unwrap_err();
4598 assert!(
4599 matches!(error, Error::JsonRpc(ref e) if e.code == -32602),
4600 "{label} was served: {error:?}"
4601 );
4602 let Error::JsonRpc(error) = error else {
4605 unreachable!()
4606 };
4607 assert!(
4608 error.message.contains("not found"),
4609 "refusal leaked that the task exists: {}",
4610 error.message
4611 );
4612 }
4613 }
4614
4615 assert!(
4617 router
4618 .handle(
4619 RequestId::Number(9),
4620 McpRequest::GetTaskInfo(GetTaskInfoParams {
4621 task_id: task_id.clone(),
4622 meta: None,
4623 }),
4624 as_principal("alice"),
4625 )
4626 .await
4627 .is_ok(),
4628 "a refused cancel must not have cancelled the task"
4629 );
4630 }
4631
4632 #[cfg(all(feature = "oauth", feature = "stateless"))]
4633 #[tokio::test]
4634 async fn final_tasks_work_across_independent_routers_with_a_shared_store() {
4635 fn as_principal(subject: &str) -> Extensions {
4636 let mut extensions = tasks_client_extensions();
4637 extensions.insert(crate::oauth::token::TokenClaims {
4638 sub: Some(subject.to_string()),
4639 iss: None,
4640 aud: None,
4641 exp: None,
4642 scope: None,
4643 client_id: None,
4644 extra: HashMap::new(),
4645 });
4646 extensions
4647 }
4648
4649 fn router_with_store(store: Arc<dyn TaskStore>) -> McpRouter {
4650 McpRouter::new()
4651 .tool(
4652 ToolBuilder::new("shared_task")
4653 .task_support(TaskSupportMode::Optional)
4654 .handler(|_input: serde_json::Value| async move {
4655 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
4656 Ok(CallToolResult::text("done"))
4657 })
4658 .build(),
4659 )
4660 .task_store(store)
4661 .with_tasks()
4662 }
4663
4664 let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
4665 let router_a = router_with_store(store.clone());
4666 let router_b = router_with_store(store);
4667
4668 let McpResponse::FinalCreateTask(created) = router_a
4669 .handle(
4670 RequestId::Number(1),
4671 McpRequest::CallTool(CallToolParams {
4672 name: "shared_task".to_string(),
4673 arguments: serde_json::json!({}),
4674 input_responses: None,
4675 request_state: None,
4676 meta: None,
4677 task: None,
4678 }),
4679 as_principal("alice"),
4680 )
4681 .await
4682 .unwrap()
4683 else {
4684 panic!("router A did not create a final task")
4685 };
4686 let task_id = created.task.metadata.task_id;
4687
4688 assert!(
4690 router_b
4691 .handle(
4692 RequestId::Number(2),
4693 McpRequest::GetTaskInfo(GetTaskInfoParams {
4694 task_id: task_id.clone(),
4695 meta: None,
4696 }),
4697 as_principal("alice"),
4698 )
4699 .await
4700 .is_ok()
4701 );
4702
4703 let denied = router_b
4705 .handle(
4706 RequestId::Number(3),
4707 McpRequest::GetTaskInfo(GetTaskInfoParams {
4708 task_id: task_id.clone(),
4709 meta: None,
4710 }),
4711 as_principal("bob"),
4712 )
4713 .await
4714 .unwrap_err();
4715 let unknown = router_b
4716 .handle(
4717 RequestId::Number(4),
4718 McpRequest::GetTaskInfo(GetTaskInfoParams {
4719 task_id: "unknown-task".to_string(),
4720 meta: None,
4721 }),
4722 as_principal("bob"),
4723 )
4724 .await
4725 .unwrap_err();
4726 let (Error::JsonRpc(denied), Error::JsonRpc(unknown)) = (denied, unknown) else {
4727 panic!("expected JSON-RPC task denials")
4728 };
4729 assert_eq!(denied.code, unknown.code);
4730 assert_eq!(
4731 denied.message.replace(&task_id, "<task-id>"),
4732 unknown.message.replace("unknown-task", "<task-id>")
4733 );
4734 assert_eq!(denied.data, unknown.data);
4735
4736 assert!(matches!(
4739 router_b
4740 .handle(
4741 RequestId::Number(5),
4742 McpRequest::CancelTask(CancelTaskParams {
4743 task_id: task_id.clone(),
4744 reason: None,
4745 meta: None,
4746 }),
4747 as_principal("alice"),
4748 )
4749 .await
4750 .unwrap(),
4751 McpResponse::FinalTaskAck(_)
4752 ));
4753 let McpResponse::FinalGetTask(fetched) = router_a
4754 .handle(
4755 RequestId::Number(6),
4756 McpRequest::GetTaskInfo(GetTaskInfoParams {
4757 task_id,
4758 meta: None,
4759 }),
4760 as_principal("alice"),
4761 )
4762 .await
4763 .unwrap()
4764 else {
4765 panic!("router A did not read the shared task")
4766 };
4767 assert_eq!(fetched.task.status(), TaskStatus::Cancelled);
4768 }
4769
4770 #[test]
4771 fn router_advertises_only_locally_declared_protocol_extensions() {
4772 let router = McpRouter::new().with_protocol_extension(
4773 crate::ExtensionDeclaration::new(
4774 "com.example/rendering",
4775 serde_json::json!({"formats": ["html"]}),
4776 )
4777 .unwrap(),
4778 );
4779
4780 let stable = router.capabilities();
4781 let final_capabilities =
4782 router.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
4783 for capabilities in [stable, final_capabilities] {
4784 let extensions = capabilities.extensions.unwrap();
4785 assert_eq!(extensions.len(), 1);
4786 assert_eq!(extensions["com.example/rendering"]["formats"][0], "html");
4787 assert!(!extensions.contains_key("com.example/client-only"));
4788 }
4789 }
4790
4791 #[tokio::test]
4792 async fn initialize_persists_negotiated_extensions_for_legacy_contexts() {
4793 let router = McpRouter::new().with_protocol_extension(
4794 crate::ExtensionDeclaration::new(
4795 "com.example/shared",
4796 serde_json::json!({"server": true}),
4797 )
4798 .unwrap(),
4799 );
4800 let client_capabilities = ClientCapabilities {
4801 extensions: Some(HashMap::from([
4802 (
4803 "com.example/shared".to_string(),
4804 serde_json::json!({"client": true}),
4805 ),
4806 ("com.example/client-only".to_string(), serde_json::json!({})),
4807 ])),
4808 ..ClientCapabilities::default()
4809 };
4810
4811 router
4812 .handle(
4813 RequestId::Number(1),
4814 McpRequest::Initialize(InitializeParams {
4815 protocol_version: crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
4816 capabilities: client_capabilities,
4817 client_info: Implementation {
4818 name: "extension-test".to_string(),
4819 version: "1.0.0".to_string(),
4820 title: None,
4821 description: None,
4822 icons: None,
4823 website_url: None,
4824 meta: None,
4825 },
4826 meta: None,
4827 }),
4828 Extensions::new(),
4829 )
4830 .await
4831 .unwrap();
4832
4833 let context = router.create_context(RequestId::Number(2), None);
4834 let negotiated = context.negotiated_extensions().unwrap();
4835 assert!(negotiated.contains("com.example/shared"));
4836 assert!(!negotiated.contains("com.example/client-only"));
4837 }
4838
4839 #[cfg(feature = "stateless")]
4840 #[test]
4841 fn final_request_context_exposes_only_negotiated_extensions() {
4842 let router = McpRouter::new().with_protocol_extension(
4843 crate::ExtensionDeclaration::new(
4844 "com.example/shared",
4845 serde_json::json!({"server": true}),
4846 )
4847 .unwrap(),
4848 );
4849 let per_request = final_extensions(ClientCapabilities {
4850 extensions: Some(HashMap::from([
4851 (
4852 "com.example/shared".to_string(),
4853 serde_json::json!({"client": true}),
4854 ),
4855 ("com.example/client-only".to_string(), serde_json::json!({})),
4856 ])),
4857 ..ClientCapabilities::default()
4858 });
4859
4860 let context =
4861 router.create_context_with_extensions(RequestId::Number(1), None, &per_request);
4862 let negotiated = context.negotiated_extensions().unwrap();
4863
4864 assert_eq!(negotiated.len(), 1);
4865 assert_eq!(
4866 negotiated
4867 .get("com.example/shared")
4868 .unwrap()
4869 .client_settings()["client"],
4870 true
4871 );
4872 assert!(!negotiated.contains("com.example/client-only"));
4873 }
4874
4875 #[cfg(feature = "stateless")]
4876 #[tokio::test]
4877 async fn final_protocol_withholds_incomplete_tasks_advertisement() {
4878 let optional = ToolBuilder::new("optional_task")
4879 .task_support(TaskSupportMode::Optional)
4880 .handler(|input: AddInput| async move {
4881 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4882 })
4883 .build();
4884 let required = ToolBuilder::new("required_task")
4885 .task_support(TaskSupportMode::Required)
4886 .handler(|input: AddInput| async move {
4887 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4888 })
4889 .build();
4890 let mut router = McpRouter::new().tool(optional).tool(required);
4891
4892 let stable_capabilities = router.capabilities();
4894 assert!(stable_capabilities.tasks.is_some());
4895 assert!(
4896 stable_capabilities
4897 .extensions
4898 .as_ref()
4899 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID))
4900 );
4901
4902 let response = router
4904 .handle(
4905 RequestId::Number(1),
4906 McpRequest::Discover(DiscoverParams::default()),
4907 Extensions::new(),
4908 )
4909 .await
4910 .unwrap();
4911 let McpResponse::Discover(result) = response else {
4912 panic!("Expected Discover response");
4913 };
4914 assert!(result.capabilities.tasks.is_none());
4915 assert!(
4916 result
4917 .capabilities
4918 .extensions
4919 .as_ref()
4920 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4921 );
4922
4923 init_router(&mut router).await;
4924
4925 let response = router
4927 .handle(
4928 RequestId::Number(2),
4929 McpRequest::ListTools(ListToolsParams::default()),
4930 Extensions::new(),
4931 )
4932 .await
4933 .unwrap();
4934 let McpResponse::ListTools(result) = response else {
4935 panic!("Expected ListTools response");
4936 };
4937 assert_eq!(result.tools.len(), 2);
4938 assert!(result.tools.iter().all(|tool| tool.execution.is_some()));
4939
4940 let response = router
4943 .handle(
4944 RequestId::Number(3),
4945 McpRequest::ListTools(ListToolsParams::default()),
4946 final_extensions(ClientCapabilities::default()),
4947 )
4948 .await
4949 .unwrap();
4950 let McpResponse::ListTools(result) = response else {
4951 panic!("Expected ListTools response");
4952 };
4953 assert_eq!(result.tools.len(), 1);
4954 assert_eq!(result.tools[0].name, "optional_task");
4955 assert!(result.tools[0].execution.is_none());
4956 }
4957
4958 #[cfg(feature = "stateless")]
4959 #[tokio::test]
4960 async fn final_protocol_enforces_tasks_negotiation() {
4961 let optional = ToolBuilder::new("optional_task")
4962 .task_support(TaskSupportMode::Optional)
4963 .handler(|input: AddInput| async move {
4964 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4965 })
4966 .build();
4967 let required = ToolBuilder::new("required_task")
4968 .task_support(TaskSupportMode::Required)
4969 .handler(|input: AddInput| async move {
4970 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4971 })
4972 .build();
4973 let mut router = McpRouter::new().tool(optional).tool(required).with_tasks();
4974 init_router(&mut router).await;
4975
4976 let response = router
4978 .handle(
4979 RequestId::Number(1),
4980 McpRequest::CallTool(CallToolParams {
4981 name: "optional_task".to_string(),
4982 arguments: serde_json::json!({"a": 1, "b": 2}),
4983 input_responses: None,
4984 request_state: None,
4985 meta: None,
4986 task: None,
4987 }),
4988 final_extensions(ClientCapabilities::default()),
4989 )
4990 .await
4991 .unwrap();
4992 assert!(matches!(response, McpResponse::CallTool(_)));
4993
4994 let error = router
4996 .handle(
4997 RequestId::Number(2),
4998 McpRequest::CallTool(CallToolParams {
4999 name: "optional_task".to_string(),
5000 arguments: serde_json::json!({"a": 1, "b": 2}),
5001 input_responses: None,
5002 request_state: None,
5003 meta: None,
5004 task: Some(TaskRequestParams { ttl: None }),
5005 }),
5006 final_extensions(ClientCapabilities::default()),
5007 )
5008 .await
5009 .unwrap_err();
5010 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32602));
5011
5012 let error = router
5016 .handle(
5017 RequestId::Number(3),
5018 McpRequest::CallTool(CallToolParams {
5019 name: "required_task".to_string(),
5020 arguments: serde_json::json!({"a": 1, "b": 2}),
5021 input_responses: None,
5022 request_state: None,
5023 meta: None,
5024 task: None,
5025 }),
5026 final_extensions(ClientCapabilities::default()),
5027 )
5028 .await
5029 .unwrap_err();
5030 let Error::JsonRpc(error) = error else {
5031 panic!("expected a JSON-RPC error");
5032 };
5033 assert_eq!(error.code, -32021);
5034 assert_eq!(
5035 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
5036 serde_json::json!({}),
5037 "the error must name the extension the client needs to declare"
5038 );
5039
5040 let task_requests = [
5041 McpRequest::GetTaskInfo(GetTaskInfoParams {
5042 task_id: "task-unknown".to_string(),
5043 meta: None,
5044 }),
5045 McpRequest::UpdateTask(UpdateTaskParams {
5046 task_id: "task-unknown".to_string(),
5047 input_responses: HashMap::new(),
5048 meta: None,
5049 }),
5050 McpRequest::CancelTask(CancelTaskParams {
5051 task_id: "task-unknown".to_string(),
5052 reason: None,
5053 meta: None,
5054 }),
5055 ];
5056 for (index, request) in task_requests.into_iter().enumerate() {
5057 let error = router
5058 .handle(
5059 RequestId::Number(4 + index as i64),
5060 request,
5061 final_extensions(ClientCapabilities::default()),
5062 )
5063 .await
5064 .unwrap_err();
5065 let Error::JsonRpc(error) = error else {
5066 panic!("expected a JSON-RPC error");
5067 };
5068 assert_eq!(error.code, -32021);
5069 assert_eq!(
5070 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
5071 serde_json::json!({})
5072 );
5073 }
5074
5075 let router_without_tasks = McpRouter::new();
5078 let error = router_without_tasks
5079 .handle(
5080 RequestId::Number(7),
5081 McpRequest::GetTaskInfo(GetTaskInfoParams {
5082 task_id: "task-unknown".to_string(),
5083 meta: None,
5084 }),
5085 final_extensions(tasks_client_capabilities()),
5086 )
5087 .await
5088 .unwrap_err();
5089 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32601));
5090 }
5091
5092 #[cfg(feature = "stateless")]
5093 #[test]
5094 fn input_required_capability_validation_uses_capability_semantics() {
5095 let roots = InputRequiredResult::with_requests(
5096 [(
5097 "roots".to_string(),
5098 InputRequest::ListRoots(ListRootsParams::default()),
5099 )]
5100 .into_iter()
5101 .collect(),
5102 );
5103 let extensions = final_extensions(ClientCapabilities {
5104 roots: Some(RootsCapability {
5105 list_changed: true,
5106 deprecated: None,
5107 }),
5108 ..Default::default()
5109 });
5110 validate_input_required_result(&extensions, &roots).unwrap();
5111 assert!(client_capabilities_satisfy(
5112 extensions
5113 .get::<crate::stateless::StatelessRequestMeta>()
5114 .and_then(|meta| meta.client_capabilities.as_ref())
5115 .unwrap(),
5116 &ClientCapabilities {
5117 roots: Some(RootsCapability::default()),
5118 ..Default::default()
5119 }
5120 ));
5121
5122 let sampling_with_tools = InputRequiredResult::with_requests(
5123 [(
5124 "sample".to_string(),
5125 InputRequest::CreateMessage(CreateMessageParams {
5126 tools: Some(Vec::new()),
5127 ..CreateMessageParams::new(vec![SamplingMessage::user("hello")], 10)
5128 }),
5129 )]
5130 .into_iter()
5131 .collect(),
5132 );
5133 let extensions = final_extensions(ClientCapabilities {
5134 sampling: Some(SamplingCapability::default()),
5135 ..Default::default()
5136 });
5137 assert!(validate_input_required_result(&extensions, &sampling_with_tools).is_err());
5138
5139 let form = InputRequiredResult::with_requests(
5140 [(
5141 "form".to_string(),
5142 InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
5143 mode: Some(ElicitMode::Form),
5144 message: "name".into(),
5145 requested_schema: ElicitFormSchema::new(),
5146 meta: None,
5147 })),
5148 )]
5149 .into_iter()
5150 .collect(),
5151 );
5152 let extensions = final_extensions(ClientCapabilities {
5153 elicitation: Some(ElicitationCapability::default()),
5154 ..Default::default()
5155 });
5156 validate_input_required_result(&extensions, &form).unwrap();
5157 }
5158
5159 async fn init_router(router: &mut McpRouter) {
5161 let init_req = RouterRequest {
5163 id: RequestId::Number(0),
5164 inner: McpRequest::Initialize(InitializeParams {
5165 protocol_version: "2025-11-25".to_string(),
5166 capabilities: ClientCapabilities {
5167 roots: None,
5168 sampling: None,
5169 elicitation: None,
5170 tasks: None,
5171 experimental: None,
5172 extensions: None,
5173 },
5174 client_info: Implementation {
5175 name: "test".to_string(),
5176 version: "1.0".to_string(),
5177 ..Default::default()
5178 },
5179 meta: None,
5180 }),
5181 extensions: Extensions::new(),
5182 };
5183 let _ = router.ready().await.unwrap().call(init_req).await.unwrap();
5184 router.handle_notification(McpNotification::Initialized);
5186 }
5187
5188 #[tokio::test]
5189 async fn test_router_list_tools() {
5190 let add_tool = ToolBuilder::new("add")
5191 .description("Add two numbers")
5192 .handler(|input: AddInput| async move {
5193 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5194 })
5195 .build();
5196
5197 let mut router = McpRouter::new().tool(add_tool);
5198
5199 init_router(&mut router).await;
5201
5202 let req = RouterRequest {
5203 id: RequestId::Number(1),
5204 inner: McpRequest::ListTools(ListToolsParams::default()),
5205 extensions: Extensions::new(),
5206 };
5207
5208 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5209
5210 match resp.inner {
5211 Ok(McpResponse::ListTools(result)) => {
5212 assert_eq!(result.tools.len(), 1);
5213 assert_eq!(result.tools[0].name, "add");
5214 }
5215 _ => panic!("Expected ListTools response"),
5216 }
5217 }
5218
5219 #[tokio::test]
5220 async fn test_router_call_tool() {
5221 let add_tool = ToolBuilder::new("add")
5222 .description("Add two numbers")
5223 .handler(|input: AddInput| async move {
5224 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5225 })
5226 .build();
5227
5228 let mut router = McpRouter::new().tool(add_tool);
5229
5230 init_router(&mut router).await;
5232
5233 let req = RouterRequest {
5234 id: RequestId::Number(1),
5235 inner: McpRequest::CallTool(CallToolParams {
5236 input_responses: None,
5237 request_state: None,
5238 name: "add".to_string(),
5239 arguments: serde_json::json!({"a": 2, "b": 3}),
5240 meta: None,
5241 task: None,
5242 }),
5243 extensions: Extensions::new(),
5244 };
5245
5246 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5247
5248 match resp.inner {
5249 Ok(McpResponse::CallTool(result)) => {
5250 assert!(!result.is_error);
5251 match &result.content[0] {
5253 Content::Text { text, .. } => assert_eq!(text, "5"),
5254 _ => panic!("Expected text content"),
5255 }
5256 }
5257 _ => panic!("Expected CallTool response"),
5258 }
5259 }
5260
5261 async fn init_jsonrpc_service(service: &mut JsonRpcService<McpRouter>, router: &McpRouter) {
5263 let init_req = JsonRpcRequest::new(0, "initialize").with_params(serde_json::json!({
5264 "protocolVersion": "2025-11-25",
5265 "capabilities": {},
5266 "clientInfo": { "name": "test", "version": "1.0" }
5267 }));
5268 let _ = service.call_single(init_req).await.unwrap();
5269 router.handle_notification(McpNotification::Initialized);
5270 }
5271
5272 #[tokio::test]
5273 async fn test_jsonrpc_service() {
5274 let add_tool = ToolBuilder::new("add")
5275 .description("Add two numbers")
5276 .handler(|input: AddInput| async move {
5277 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5278 })
5279 .build();
5280
5281 let router = McpRouter::new().tool(add_tool);
5282 let mut service = JsonRpcService::new(router.clone());
5283
5284 init_jsonrpc_service(&mut service, &router).await;
5286
5287 let req = JsonRpcRequest::new(1, "tools/list");
5288
5289 let resp = service.call_single(req).await.unwrap();
5290
5291 match resp {
5292 JsonRpcResponse::Result(r) => {
5293 assert_eq!(r.id, RequestId::Number(1));
5294 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5295 assert_eq!(tools.len(), 1);
5296 }
5297 JsonRpcResponse::Error(_) => panic!("Expected success response"),
5298 _ => panic!("unexpected response variant"),
5299 }
5300 }
5301
5302 #[tokio::test]
5303 async fn test_batch_request() {
5304 let add_tool = ToolBuilder::new("add")
5305 .description("Add two numbers")
5306 .handler(|input: AddInput| async move {
5307 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5308 })
5309 .build();
5310
5311 let router = McpRouter::new().tool(add_tool);
5312 let mut service = JsonRpcService::new(router.clone())
5313 .protocol_versions(["2025-03-26"])
5314 .unwrap();
5315
5316 init_jsonrpc_service(&mut service, &router).await;
5318
5319 let requests = vec![
5321 JsonRpcRequest::new(1, "tools/list"),
5322 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5323 "name": "add",
5324 "arguments": {"a": 10, "b": 20}
5325 })),
5326 JsonRpcRequest::new(3, "ping"),
5327 ];
5328
5329 let responses = service.call_batch(requests).await.unwrap();
5330
5331 assert_eq!(responses.len(), 3);
5332
5333 match &responses[0] {
5335 JsonRpcResponse::Result(r) => {
5336 assert_eq!(r.id, RequestId::Number(1));
5337 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5338 assert_eq!(tools.len(), 1);
5339 }
5340 JsonRpcResponse::Error(_) => panic!("Expected success for tools/list"),
5341 _ => panic!("unexpected response variant"),
5342 }
5343
5344 match &responses[1] {
5346 JsonRpcResponse::Result(r) => {
5347 assert_eq!(r.id, RequestId::Number(2));
5348 let content = r.result.get("content").unwrap().as_array().unwrap();
5349 let text = content[0].get("text").unwrap().as_str().unwrap();
5350 assert_eq!(text, "30");
5351 }
5352 JsonRpcResponse::Error(_) => panic!("Expected success for tools/call"),
5353 _ => panic!("unexpected response variant"),
5354 }
5355
5356 match &responses[2] {
5358 JsonRpcResponse::Result(r) => {
5359 assert_eq!(r.id, RequestId::Number(3));
5360 }
5361 JsonRpcResponse::Error(_) => panic!("Expected success for ping"),
5362 _ => panic!("unexpected response variant"),
5363 }
5364 }
5365
5366 #[tokio::test]
5367 async fn test_empty_batch_error() {
5368 let router = McpRouter::new();
5369 let mut service = JsonRpcService::new(router);
5370
5371 let result = service.call_batch(vec![]).await;
5372 assert!(result.is_err());
5373 }
5374
5375 #[tokio::test]
5380 async fn test_progress_token_extraction() {
5381 use crate::context::{ServerNotification, notification_channel};
5382 use crate::protocol::ProgressToken;
5383 use std::sync::Arc;
5384 use std::sync::atomic::{AtomicBool, Ordering};
5385
5386 let progress_reported = Arc::new(AtomicBool::new(false));
5388 let progress_ref = progress_reported.clone();
5389
5390 let tool = ToolBuilder::new("progress_tool")
5392 .description("Tool that reports progress")
5393 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5394 let reported = progress_ref.clone();
5395 async move {
5396 ctx.report_progress(50.0, Some(100.0), Some("Halfway"))
5398 .await;
5399 reported.store(true, Ordering::SeqCst);
5400 Ok(CallToolResult::text("done"))
5401 }
5402 })
5403 .build();
5404
5405 let (tx, mut rx) = notification_channel(10);
5407 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5408 let mut service = JsonRpcService::new(router.clone());
5409
5410 init_jsonrpc_service(&mut service, &router).await;
5412
5413 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5415 "name": "progress_tool",
5416 "arguments": {"a": 1, "b": 2},
5417 "_meta": {
5418 "progressToken": "test-token-123"
5419 }
5420 }));
5421
5422 let resp = service.call_single(req).await.unwrap();
5423
5424 match resp {
5426 JsonRpcResponse::Result(_) => {}
5427 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
5428 _ => panic!("unexpected response variant"),
5429 }
5430
5431 assert!(progress_reported.load(Ordering::SeqCst));
5433
5434 let notification = rx.try_recv().expect("Expected progress notification");
5436 match notification {
5437 ServerNotification::Progress(params) => {
5438 assert_eq!(
5439 params.progress_token,
5440 ProgressToken::String("test-token-123".to_string())
5441 );
5442 assert_eq!(params.progress, 50.0);
5443 assert_eq!(params.total, Some(100.0));
5444 assert_eq!(params.message.as_deref(), Some("Halfway"));
5445 }
5446 _ => panic!("Expected Progress notification"),
5447 }
5448 }
5449
5450 #[tokio::test]
5451 async fn test_tool_call_without_progress_token() {
5452 use crate::context::notification_channel;
5453 use std::sync::Arc;
5454 use std::sync::atomic::{AtomicBool, Ordering};
5455
5456 let progress_attempted = Arc::new(AtomicBool::new(false));
5457 let progress_ref = progress_attempted.clone();
5458
5459 let tool = ToolBuilder::new("no_token_tool")
5460 .description("Tool that tries to report progress without token")
5461 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5462 let attempted = progress_ref.clone();
5463 async move {
5464 ctx.report_progress(50.0, Some(100.0), None).await;
5466 attempted.store(true, Ordering::SeqCst);
5467 Ok(CallToolResult::text("done"))
5468 }
5469 })
5470 .build();
5471
5472 let (tx, mut rx) = notification_channel(10);
5473 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5474 let mut service = JsonRpcService::new(router.clone());
5475
5476 init_jsonrpc_service(&mut service, &router).await;
5477
5478 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5480 "name": "no_token_tool",
5481 "arguments": {"a": 1, "b": 2}
5482 }));
5483
5484 let resp = service.call_single(req).await.unwrap();
5485 assert!(matches!(resp, JsonRpcResponse::Result(_)));
5486
5487 assert!(progress_attempted.load(Ordering::SeqCst));
5489
5490 assert!(rx.try_recv().is_err());
5492 }
5493
5494 #[tokio::test]
5495 async fn test_batch_errors_returned_not_dropped() {
5496 let add_tool = ToolBuilder::new("add")
5497 .description("Add two numbers")
5498 .handler(|input: AddInput| async move {
5499 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5500 })
5501 .build();
5502
5503 let router = McpRouter::new().tool(add_tool);
5504 let mut service = JsonRpcService::new(router.clone())
5505 .protocol_versions(["2025-03-26"])
5506 .unwrap();
5507
5508 init_jsonrpc_service(&mut service, &router).await;
5509
5510 let requests = vec![
5512 JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5514 "name": "add",
5515 "arguments": {"a": 10, "b": 20}
5516 })),
5517 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5519 "name": "nonexistent_tool",
5520 "arguments": {}
5521 })),
5522 JsonRpcRequest::new(3, "ping"),
5524 ];
5525
5526 let responses = service.call_batch(requests).await.unwrap();
5527
5528 assert_eq!(responses.len(), 3);
5530
5531 match &responses[0] {
5533 JsonRpcResponse::Result(r) => {
5534 assert_eq!(r.id, RequestId::Number(1));
5535 }
5536 JsonRpcResponse::Error(_) => panic!("Expected success for first request"),
5537 _ => panic!("unexpected response variant"),
5538 }
5539
5540 match &responses[1] {
5542 JsonRpcResponse::Error(e) => {
5543 assert_eq!(e.id, Some(RequestId::Number(2)));
5544 assert!(e.error.message.contains("not found") || e.error.code == -32601);
5546 }
5547 JsonRpcResponse::Result(_) => panic!("Expected error for second request"),
5548 _ => panic!("unexpected response variant"),
5549 }
5550
5551 match &responses[2] {
5553 JsonRpcResponse::Result(r) => {
5554 assert_eq!(r.id, RequestId::Number(3));
5555 }
5556 JsonRpcResponse::Error(_) => panic!("Expected success for third request"),
5557 _ => panic!("unexpected response variant"),
5558 }
5559 }
5560
5561 #[tokio::test]
5566 async fn test_list_resource_templates() {
5567 use crate::resource::ResourceTemplateBuilder;
5568 use std::collections::HashMap;
5569
5570 let template = ResourceTemplateBuilder::new("file:///{path}")
5571 .name("Project Files")
5572 .description("Access project files")
5573 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5574 Ok(ReadResourceResult {
5575 contents: vec![ResourceContent {
5576 uri,
5577 mime_type: None,
5578 text: None,
5579 blob: None,
5580 meta: None,
5581 }],
5582 meta: None,
5583 ..Default::default()
5584 })
5585 });
5586
5587 let mut router = McpRouter::new().resource_template(template);
5588
5589 init_router(&mut router).await;
5591
5592 let req = RouterRequest {
5593 id: RequestId::Number(1),
5594 inner: McpRequest::ListResourceTemplates(ListResourceTemplatesParams::default()),
5595 extensions: Extensions::new(),
5596 };
5597
5598 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5599
5600 match resp.inner {
5601 Ok(McpResponse::ListResourceTemplates(result)) => {
5602 assert_eq!(result.resource_templates.len(), 1);
5603 assert_eq!(result.resource_templates[0].uri_template, "file:///{path}");
5604 assert_eq!(result.resource_templates[0].name, "Project Files");
5605 }
5606 _ => panic!("Expected ListResourceTemplates response"),
5607 }
5608 }
5609
5610 #[tokio::test]
5611 async fn test_read_resource_via_template() {
5612 use crate::resource::ResourceTemplateBuilder;
5613 use std::collections::HashMap;
5614
5615 let template = ResourceTemplateBuilder::new("db://users/{id}")
5616 .name("User Records")
5617 .handler(|uri: String, vars: HashMap<String, String>| async move {
5618 let id = vars.get("id").unwrap().clone();
5619 Ok(ReadResourceResult {
5620 contents: vec![ResourceContent {
5621 uri,
5622 mime_type: Some("application/json".to_string()),
5623 text: Some(format!(r#"{{"id": "{}"}}"#, id)),
5624 blob: None,
5625 meta: None,
5626 }],
5627 meta: None,
5628 ..Default::default()
5629 })
5630 });
5631
5632 let mut router = McpRouter::new().resource_template(template);
5633
5634 init_router(&mut router).await;
5636
5637 let req = RouterRequest {
5639 id: RequestId::Number(1),
5640 inner: McpRequest::ReadResource(ReadResourceParams {
5641 input_responses: None,
5642 request_state: None,
5643 uri: "db://users/123".to_string(),
5644 meta: None,
5645 }),
5646 extensions: Extensions::new(),
5647 };
5648
5649 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5650
5651 match resp.inner {
5652 Ok(McpResponse::ReadResource(result)) => {
5653 assert_eq!(result.contents.len(), 1);
5654 assert_eq!(result.contents[0].uri, "db://users/123");
5655 assert!(result.contents[0].text.as_ref().unwrap().contains("123"));
5656 }
5657 _ => panic!("Expected ReadResource response"),
5658 }
5659 }
5660
5661 #[tokio::test]
5662 async fn test_static_resource_takes_precedence_over_template() {
5663 use crate::resource::{ResourceBuilder, ResourceTemplateBuilder};
5664 use std::collections::HashMap;
5665
5666 let template = ResourceTemplateBuilder::new("file:///{path}")
5668 .name("Files Template")
5669 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5670 Ok(ReadResourceResult {
5671 contents: vec![ResourceContent {
5672 uri,
5673 mime_type: None,
5674 text: Some("from template".to_string()),
5675 blob: None,
5676 meta: None,
5677 }],
5678 meta: None,
5679 ..Default::default()
5680 })
5681 });
5682
5683 let static_resource = ResourceBuilder::new("file:///README.md")
5685 .name("README")
5686 .text("from static resource");
5687
5688 let mut router = McpRouter::new()
5689 .resource_template(template)
5690 .resource(static_resource);
5691
5692 init_router(&mut router).await;
5694
5695 let req = RouterRequest {
5697 id: RequestId::Number(1),
5698 inner: McpRequest::ReadResource(ReadResourceParams {
5699 input_responses: None,
5700 request_state: None,
5701 uri: "file:///README.md".to_string(),
5702 meta: None,
5703 }),
5704 extensions: Extensions::new(),
5705 };
5706
5707 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5708
5709 match resp.inner {
5710 Ok(McpResponse::ReadResource(result)) => {
5711 assert_eq!(
5713 result.contents[0].text.as_deref(),
5714 Some("from static resource")
5715 );
5716 }
5717 _ => panic!("Expected ReadResource response"),
5718 }
5719 }
5720
5721 #[tokio::test]
5722 async fn test_resource_not_found_when_no_match() {
5723 use crate::resource::ResourceTemplateBuilder;
5724 use std::collections::HashMap;
5725
5726 let template = ResourceTemplateBuilder::new("db://users/{id}")
5727 .name("Users")
5728 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5729 Ok(ReadResourceResult {
5730 contents: vec![ResourceContent {
5731 uri,
5732 mime_type: None,
5733 text: None,
5734 blob: None,
5735 meta: None,
5736 }],
5737 meta: None,
5738 ..Default::default()
5739 })
5740 });
5741
5742 let mut router = McpRouter::new().resource_template(template);
5743
5744 init_router(&mut router).await;
5746
5747 let req = RouterRequest {
5749 id: RequestId::Number(1),
5750 inner: McpRequest::ReadResource(ReadResourceParams {
5751 input_responses: None,
5752 request_state: None,
5753 uri: "db://posts/123".to_string(),
5754 meta: None,
5755 }),
5756 extensions: Extensions::new(),
5757 };
5758
5759 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5760
5761 match resp.inner {
5762 Err(err) => {
5763 assert!(err.message.contains("not found"));
5764 }
5765 Ok(_) => panic!("Expected error for non-matching URI"),
5766 }
5767 }
5768
5769 #[tokio::test]
5770 async fn test_capabilities_include_resources_with_only_templates() {
5771 use crate::resource::ResourceTemplateBuilder;
5772 use std::collections::HashMap;
5773
5774 let template = ResourceTemplateBuilder::new("file:///{path}")
5775 .name("Files")
5776 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5777 Ok(ReadResourceResult {
5778 contents: vec![ResourceContent {
5779 uri,
5780 mime_type: None,
5781 text: None,
5782 blob: None,
5783 meta: None,
5784 }],
5785 meta: None,
5786 ..Default::default()
5787 })
5788 });
5789
5790 let mut router = McpRouter::new().resource_template(template);
5791
5792 let init_req = RouterRequest {
5794 id: RequestId::Number(0),
5795 inner: McpRequest::Initialize(InitializeParams {
5796 protocol_version: "2025-11-25".to_string(),
5797 capabilities: ClientCapabilities {
5798 roots: None,
5799 sampling: None,
5800 elicitation: None,
5801 tasks: None,
5802 experimental: None,
5803 extensions: None,
5804 },
5805 client_info: Implementation {
5806 name: "test".to_string(),
5807 version: "1.0".to_string(),
5808 ..Default::default()
5809 },
5810 meta: None,
5811 }),
5812 extensions: Extensions::new(),
5813 };
5814 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5815
5816 match resp.inner {
5817 Ok(McpResponse::Initialize(result)) => {
5818 assert!(result.capabilities.resources.is_some());
5820 }
5821 _ => panic!("Expected Initialize response"),
5822 }
5823 }
5824
5825 #[tokio::test]
5830 async fn test_log_sends_notification() {
5831 use crate::context::notification_channel;
5832
5833 let (tx, mut rx) = notification_channel(10);
5834 let router = McpRouter::new().with_notification_sender(tx);
5835
5836 let sent = router.log_info("Test message");
5838 assert!(sent);
5839
5840 let notification = rx.try_recv().unwrap();
5842 match notification {
5843 ServerNotification::LogMessage(params) => {
5844 assert_eq!(params.level, LogLevel::Info);
5845 let data = params.data;
5846 assert_eq!(
5847 data.get("message").unwrap().as_str().unwrap(),
5848 "Test message"
5849 );
5850 }
5851 _ => panic!("Expected LogMessage notification"),
5852 }
5853 }
5854
5855 #[tokio::test]
5856 async fn test_log_with_custom_params() {
5857 use crate::context::notification_channel;
5858
5859 let (tx, mut rx) = notification_channel(10);
5860 let router = McpRouter::new().with_notification_sender(tx);
5861
5862 let params = LoggingMessageParams::new(
5864 LogLevel::Error,
5865 serde_json::json!({
5866 "error": "Connection failed",
5867 "host": "localhost"
5868 }),
5869 )
5870 .with_logger("database");
5871
5872 let sent = router.log(params);
5873 assert!(sent);
5874
5875 let notification = rx.try_recv().unwrap();
5876 match notification {
5877 ServerNotification::LogMessage(params) => {
5878 assert_eq!(params.level, LogLevel::Error);
5879 assert_eq!(params.logger.as_deref(), Some("database"));
5880 let data = params.data;
5881 assert_eq!(
5882 data.get("error").unwrap().as_str().unwrap(),
5883 "Connection failed"
5884 );
5885 }
5886 _ => panic!("Expected LogMessage notification"),
5887 }
5888 }
5889
5890 #[tokio::test]
5891 async fn test_log_without_channel_returns_false() {
5892 let router = McpRouter::new();
5894
5895 assert!(!router.log_info("Test"));
5897 assert!(!router.log_warning("Test"));
5898 assert!(!router.log_error("Test"));
5899 assert!(!router.log_debug("Test"));
5900 }
5901
5902 #[tokio::test]
5903 async fn test_logging_capability_with_channel() {
5904 use crate::context::notification_channel;
5905
5906 let (tx, _rx) = notification_channel(10);
5907 let mut router = McpRouter::new().with_notification_sender(tx);
5908
5909 let init_req = RouterRequest {
5911 id: RequestId::Number(0),
5912 inner: McpRequest::Initialize(InitializeParams {
5913 protocol_version: "2025-11-25".to_string(),
5914 capabilities: ClientCapabilities {
5915 roots: None,
5916 sampling: None,
5917 elicitation: None,
5918 tasks: None,
5919 experimental: None,
5920 extensions: None,
5921 },
5922 client_info: Implementation {
5923 name: "test".to_string(),
5924 version: "1.0".to_string(),
5925 ..Default::default()
5926 },
5927 meta: None,
5928 }),
5929 extensions: Extensions::new(),
5930 };
5931 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5932
5933 match resp.inner {
5934 Ok(McpResponse::Initialize(result)) => {
5935 assert!(result.capabilities.logging.is_some());
5937 }
5938 _ => panic!("Expected Initialize response"),
5939 }
5940 }
5941
5942 #[tokio::test]
5943 async fn test_no_logging_capability_without_channel() {
5944 let mut router = McpRouter::new();
5945
5946 let init_req = RouterRequest {
5948 id: RequestId::Number(0),
5949 inner: McpRequest::Initialize(InitializeParams {
5950 protocol_version: "2025-11-25".to_string(),
5951 capabilities: ClientCapabilities {
5952 roots: None,
5953 sampling: None,
5954 elicitation: None,
5955 tasks: None,
5956 experimental: None,
5957 extensions: None,
5958 },
5959 client_info: Implementation {
5960 name: "test".to_string(),
5961 version: "1.0".to_string(),
5962 ..Default::default()
5963 },
5964 meta: None,
5965 }),
5966 extensions: Extensions::new(),
5967 };
5968 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5969
5970 match resp.inner {
5971 Ok(McpResponse::Initialize(result)) => {
5972 assert!(result.capabilities.logging.is_none());
5974 }
5975 _ => panic!("Expected Initialize response"),
5976 }
5977 }
5978
5979 #[cfg(feature = "stateless")]
5988 #[tokio::test]
5989 async fn a_task_resumes_after_its_input_is_answered() {
5990 use crate::async_task::{MemoryTaskStore, TaskStore};
5991 use crate::protocol::{
5992 ElicitFormParams, ElicitFormSchema, ElicitRequestParams, InputRequest, InputRequests,
5993 InputRequiredResult, RequestOutcome,
5994 };
5995
5996 let asks = ToolBuilder::new("asks")
5997 .description("Needs a decision")
5998 .task_support(TaskSupportMode::Optional)
5999 .mrtr_handler::<serde_json::Value, _, _>(|ctx, _input| async move {
6000 if let Some(responses) = ctx.input_responses()
6003 && responses.contains_key("decision")
6004 {
6005 return Ok(RequestOutcome::Complete(CallToolResult::text("approved")));
6006 }
6007 let mut requests: InputRequests = Default::default();
6008 requests.insert(
6009 "decision".to_string(),
6010 InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
6011 mode: None,
6012 message: "approve?".to_string(),
6013 requested_schema: ElicitFormSchema::new(),
6014 meta: None,
6015 })),
6016 );
6017 Ok(RequestOutcome::input_required(
6018 InputRequiredResult::with_requests(requests),
6019 ))
6020 })
6021 .build();
6022
6023 let store = std::sync::Arc::new(MemoryTaskStore::new());
6024 let mut router = McpRouter::new()
6025 .task_store(store.clone())
6026 .tool(asks)
6027 .with_tasks();
6028 init_router(&mut router).await;
6029
6030 let resp = router
6031 .ready()
6032 .await
6033 .unwrap()
6034 .call(RouterRequest {
6035 id: RequestId::Number(1),
6036 inner: McpRequest::CallTool(CallToolParams {
6037 input_responses: None,
6038 request_state: None,
6039 name: "asks".to_string(),
6040 arguments: serde_json::json!({}),
6041 meta: None,
6042 task: None,
6043 }),
6044 extensions: tasks_client_extensions(),
6045 })
6046 .await
6047 .unwrap();
6048 let task_id = match resp.inner {
6049 Ok(McpResponse::FinalCreateTask(result)) => result.task.metadata.task_id,
6050 other => panic!("expected a created task, got {other:?}"),
6051 };
6052
6053 let mut parked = false;
6055 for _ in 0..50 {
6056 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
6057 let task = store.get_task(&task_id).await.unwrap().unwrap();
6058 if task.status == TaskStatus::InputRequired {
6059 parked = true;
6060 break;
6061 }
6062 assert_ne!(task.status, TaskStatus::Failed, "must park, not fail");
6063 }
6064 assert!(parked, "the task must reach input_required");
6065
6066 let outstanding = store
6068 .outstanding_input_requests(&task_id)
6069 .await
6070 .unwrap()
6071 .unwrap();
6072 assert!(outstanding.contains_key("decision"));
6073
6074 let update = router
6076 .ready()
6077 .await
6078 .unwrap()
6079 .call(RouterRequest {
6080 id: RequestId::Number(2),
6081 inner: McpRequest::UpdateTask(UpdateTaskParams {
6082 task_id: task_id.clone(),
6083 input_responses: [(
6084 "decision".to_string(),
6085 serde_json::json!({"action": "accept"}),
6086 )]
6087 .into_iter()
6088 .collect(),
6089 meta: None,
6090 }),
6091 extensions: tasks_client_extensions(),
6092 })
6093 .await
6094 .unwrap();
6095 assert!(update.inner.is_ok(), "update must be acknowledged");
6096
6097 let mut completed = None;
6099 for _ in 0..50 {
6100 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
6101 let (task, result, _) = store.get_task_result(&task_id).await.unwrap().unwrap();
6102 if task.status == TaskStatus::Completed {
6103 completed = result;
6104 break;
6105 }
6106 assert_ne!(
6107 task.status,
6108 TaskStatus::Failed,
6109 "the resumed handler must not fail"
6110 );
6111 }
6112 let completed = completed.expect("the task must complete after resuming");
6113 assert_eq!(completed.all_text(), "approved");
6114 }
6115
6116 #[cfg(feature = "stateless")]
6123 #[tokio::test]
6124 async fn a_store_without_resume_support_fails_the_task() {
6125 let store = Arc::new(CountingTaskStore::new());
6126 assert!(
6127 store.resume_context("anything").await.unwrap().is_none(),
6128 "the trait default must report no resume support"
6129 );
6130
6131 let (task_id, _cancel) = store
6132 .create_task("t", serde_json::json!({}), None, None)
6133 .await
6134 .unwrap();
6135 let requests: crate::protocol::InputRequests = [(
6136 "k".to_string(),
6137 crate::protocol::InputRequest::ListRoots(crate::protocol::ListRootsParams {
6138 meta: None,
6139 }),
6140 )]
6141 .into_iter()
6142 .collect();
6143 store.require_input(&task_id, requests, None).await.unwrap();
6144
6145 let mut router = McpRouter::new()
6146 .task_store(store.clone() as Arc<dyn TaskStore>)
6147 .with_tasks();
6148 init_router(&mut router).await;
6149 router
6150 .ready()
6151 .await
6152 .unwrap()
6153 .call(RouterRequest {
6154 id: RequestId::Number(1),
6155 inner: McpRequest::UpdateTask(UpdateTaskParams {
6156 task_id: task_id.clone(),
6157 input_responses: [("k".to_string(), serde_json::json!({"roots": []}))]
6158 .into_iter()
6159 .collect(),
6160 meta: None,
6161 }),
6162 extensions: tasks_client_extensions(),
6163 })
6164 .await
6165 .unwrap();
6166
6167 let (task, _, error) = store.get_task_result(&task_id).await.unwrap().unwrap();
6168 assert_eq!(
6169 task.status,
6170 TaskStatus::Failed,
6171 "a store that cannot resume must fail the task, not strand it"
6172 );
6173 assert!(
6174 error.unwrap().message.contains("resume_context"),
6175 "the failure must name what to implement"
6176 );
6177 }
6178
6179 #[cfg(feature = "stateless")]
6187 #[tokio::test]
6188 async fn asking_for_input_without_requests_fails_rather_than_stranding() {
6189 use crate::async_task::{MemoryTaskStore, TaskStore};
6190 use crate::protocol::{InputRequiredResult, RequestOutcome};
6191
6192 let asks = ToolBuilder::new("asks")
6193 .description("Wants input")
6194 .task_support(TaskSupportMode::Optional)
6195 .mrtr_handler::<serde_json::Value, _, _>(|_ctx, _input| async move {
6196 Ok(RequestOutcome::input_required(
6197 InputRequiredResult::new().with_request_state("state"),
6198 ))
6199 })
6200 .build();
6201
6202 let store = std::sync::Arc::new(MemoryTaskStore::new());
6203 let mut router = McpRouter::new()
6204 .task_store(store.clone())
6205 .tool(asks)
6206 .with_tasks();
6207 init_router(&mut router).await;
6208
6209 let resp = router
6210 .ready()
6211 .await
6212 .unwrap()
6213 .call(RouterRequest {
6214 id: RequestId::Number(1),
6215 inner: McpRequest::CallTool(CallToolParams {
6216 input_responses: None,
6217 request_state: None,
6218 name: "asks".to_string(),
6219 arguments: serde_json::json!({}),
6220 meta: None,
6221 task: None,
6222 }),
6223 extensions: tasks_client_extensions(),
6224 })
6225 .await
6226 .unwrap();
6227
6228 let task_id = match resp.inner {
6229 Ok(McpResponse::FinalCreateTask(result)) => result.task.metadata.task_id,
6230 other => panic!("expected a created task, got {other:?}"),
6231 };
6232
6233 let mut error = None;
6235 for _ in 0..50 {
6236 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
6237 let (task, _, err) = store.get_task_result(&task_id).await.unwrap().unwrap();
6238 if task.status == TaskStatus::Failed {
6239 error = err;
6240 break;
6241 }
6242 }
6243
6244 let error = error.expect("the task must reach a terminal failure");
6245 assert!(
6246 error.message.contains("nothing to wait for"),
6247 "must explain why the task cannot park: {}",
6248 error.message
6249 );
6250 assert!(
6251 !error.message.contains("call_outcome_with_context"),
6252 "must not name an internal Rust API: {}",
6253 error.message
6254 );
6255 }
6256
6257 #[tokio::test]
6262 async fn stable_task_update_applies_input_responses() {
6263 use crate::async_task::{MemoryTaskStore, TaskStore};
6264 use crate::protocol::{InputRequest, ListRootsParams};
6265
6266 let store = std::sync::Arc::new(MemoryTaskStore::new());
6267 let mut router = McpRouter::new().task_store(store.clone());
6268 init_router(&mut router).await;
6269
6270 let (task_id, _cancel) = store
6271 .create_task("permission_gate", serde_json::json!({}), None, None)
6272 .await
6273 .expect("create task");
6274 let requests: crate::protocol::InputRequests = [(
6275 "permission".to_string(),
6276 InputRequest::ListRoots(ListRootsParams { meta: None }),
6277 )]
6278 .into_iter()
6279 .collect();
6280 store
6281 .require_input(&task_id, requests, Some("need a decision"))
6282 .await
6283 .expect("require input");
6284 assert_eq!(
6285 store.get_task(&task_id).await.unwrap().unwrap().status,
6286 TaskStatus::InputRequired
6287 );
6288
6289 let resp = router
6291 .ready()
6292 .await
6293 .unwrap()
6294 .call(RouterRequest {
6295 id: RequestId::Number(1),
6296 inner: McpRequest::UpdateTask(UpdateTaskParams {
6297 task_id: task_id.clone(),
6298 input_responses: [("permission".to_string(), serde_json::json!({"roots": []}))]
6299 .into_iter()
6300 .collect(),
6301 meta: None,
6302 }),
6303 extensions: Extensions::new(),
6304 })
6305 .await
6306 .unwrap();
6307 assert!(
6308 matches!(resp.inner, Ok(McpResponse::UpdateTask(_))),
6309 "the empty-result acknowledgment shape is unchanged: {:?}",
6310 resp.inner
6311 );
6312
6313 assert!(
6316 store
6317 .outstanding_input_requests(&task_id)
6318 .await
6319 .unwrap()
6320 .unwrap()
6321 .is_empty(),
6322 "the outstanding request must be consumed"
6323 );
6324 assert_eq!(
6325 store.get_task(&task_id).await.unwrap().unwrap().status,
6326 TaskStatus::Working,
6327 "answering the last outstanding request resumes the task"
6328 );
6329 }
6330
6331 #[tokio::test]
6334 async fn stable_task_update_ignores_unmatched_keys() {
6335 use crate::async_task::{MemoryTaskStore, TaskStore};
6336
6337 let store = std::sync::Arc::new(MemoryTaskStore::new());
6338 let mut router = McpRouter::new().task_store(store.clone());
6339 init_router(&mut router).await;
6340
6341 let (task_id, _cancel) = store
6342 .create_task("noop", serde_json::json!({}), None, None)
6343 .await
6344 .expect("create task");
6345
6346 let resp = router
6347 .ready()
6348 .await
6349 .unwrap()
6350 .call(RouterRequest {
6351 id: RequestId::Number(1),
6352 inner: McpRequest::UpdateTask(UpdateTaskParams {
6353 task_id: task_id.clone(),
6354 input_responses: [(
6355 "never-issued".to_string(),
6356 serde_json::json!({"roots": []}),
6357 )]
6358 .into_iter()
6359 .collect(),
6360 meta: None,
6361 }),
6362 extensions: Extensions::new(),
6363 })
6364 .await
6365 .unwrap();
6366 assert!(
6367 matches!(resp.inner, Ok(McpResponse::UpdateTask(_))),
6368 "an unmatched key is ignored, not rejected: {:?}",
6369 resp.inner
6370 );
6371
6372 let unknown = router
6373 .ready()
6374 .await
6375 .unwrap()
6376 .call(RouterRequest {
6377 id: RequestId::Number(2),
6378 inner: McpRequest::UpdateTask(UpdateTaskParams {
6379 task_id: "no-such-task".to_string(),
6380 input_responses: HashMap::new(),
6381 meta: None,
6382 }),
6383 extensions: Extensions::new(),
6384 })
6385 .await
6386 .unwrap();
6387 match unknown.inner {
6388 Err(error) => assert_eq!(error.code, -32602),
6389 other => panic!("expected -32602 for an unknown task, got {other:?}"),
6390 }
6391 }
6392
6393 #[tokio::test]
6394 async fn test_create_task_via_call_tool() {
6395 let add_tool = ToolBuilder::new("add")
6396 .description("Add two numbers")
6397 .task_support(TaskSupportMode::Optional)
6398 .handler(|input: AddInput| async move {
6399 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6400 })
6401 .build();
6402
6403 let mut router = McpRouter::new().tool(add_tool);
6404 init_router(&mut router).await;
6405
6406 let req = RouterRequest {
6407 id: RequestId::Number(1),
6408 inner: McpRequest::CallTool(CallToolParams {
6409 input_responses: None,
6410 request_state: None,
6411 name: "add".to_string(),
6412 arguments: serde_json::json!({"a": 5, "b": 10}),
6413 meta: None,
6414 task: Some(TaskRequestParams { ttl: None }),
6415 }),
6416 extensions: Extensions::new(),
6417 };
6418
6419 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6420
6421 match resp.inner {
6422 Ok(McpResponse::CreateTask(result)) => {
6423 assert!(!result.task.task_id.is_empty());
6424 assert_eq!(result.task.status, TaskStatus::Working);
6425 }
6426 _ => panic!("Expected CreateTask response"),
6427 }
6428 }
6429
6430 struct CountingTaskStore {
6433 inner: MemoryTaskStore,
6434 creates: std::sync::atomic::AtomicUsize,
6435 gets: std::sync::atomic::AtomicUsize,
6436 completes: std::sync::atomic::AtomicUsize,
6437 }
6438
6439 impl CountingTaskStore {
6440 fn new() -> Self {
6441 Self {
6442 inner: MemoryTaskStore::new(),
6443 creates: std::sync::atomic::AtomicUsize::new(0),
6444 gets: std::sync::atomic::AtomicUsize::new(0),
6445 completes: std::sync::atomic::AtomicUsize::new(0),
6446 }
6447 }
6448 }
6449
6450 #[async_trait::async_trait]
6451 impl TaskStore for CountingTaskStore {
6452 async fn create_task(
6453 &self,
6454 tool_name: &str,
6455 arguments: serde_json::Value,
6456 ttl: Option<u64>,
6457 owner: crate::async_task::TaskOwner,
6458 ) -> crate::async_task::Result<(String, crate::async_task::CancellationToken)> {
6459 self.creates
6460 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6461 self.inner
6462 .create_task(tool_name, arguments, ttl, owner)
6463 .await
6464 }
6465
6466 async fn get_task(&self, task_id: &str) -> crate::async_task::Result<Option<TaskObject>> {
6467 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6468 self.inner.get_task(task_id).await
6469 }
6470
6471 async fn task_owner(
6472 &self,
6473 task_id: &str,
6474 ) -> crate::async_task::Result<Option<crate::async_task::TaskOwner>> {
6475 self.inner.task_owner(task_id).await
6476 }
6477
6478 async fn get_task_result(
6479 &self,
6480 task_id: &str,
6481 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
6482 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6485 self.inner.get_task_result(task_id).await
6486 }
6487
6488 async fn wait_for_completion(
6489 &self,
6490 task_id: &str,
6491 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
6492 self.inner.wait_for_completion(task_id).await
6493 }
6494
6495 async fn list_tasks(
6496 &self,
6497 status_filter: Option<TaskStatus>,
6498 ) -> crate::async_task::Result<Vec<TaskObject>> {
6499 self.inner.list_tasks(status_filter).await
6500 }
6501
6502 async fn require_input(
6503 &self,
6504 task_id: &str,
6505 requests: crate::protocol::InputRequests,
6506 message: Option<&str>,
6507 ) -> crate::async_task::Result<bool> {
6508 self.inner.require_input(task_id, requests, message).await
6509 }
6510
6511 async fn outstanding_input_requests(
6512 &self,
6513 task_id: &str,
6514 ) -> crate::async_task::Result<Option<crate::protocol::InputRequests>> {
6515 self.inner.outstanding_input_requests(task_id).await
6516 }
6517
6518 async fn apply_input_responses(
6519 &self,
6520 task_id: &str,
6521 responses: crate::protocol::InputResponses,
6522 ) -> crate::async_task::Result<Option<crate::async_task::AppliedInputResponses>> {
6523 self.inner.apply_input_responses(task_id, responses).await
6524 }
6525
6526 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> crate::async_task::Result<bool> {
6527 self.inner.set_ttl(task_id, ttl_ms).await
6528 }
6529
6530 async fn complete_task(
6531 &self,
6532 task_id: &str,
6533 result: CallToolResult,
6534 ) -> crate::async_task::Result<bool> {
6535 self.completes
6536 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6537 self.inner.complete_task(task_id, result).await
6538 }
6539
6540 async fn fail_task(
6541 &self,
6542 task_id: &str,
6543 error: JsonRpcError,
6544 ) -> crate::async_task::Result<bool> {
6545 self.inner.fail_task(task_id, error).await
6546 }
6547
6548 async fn cancel_task(
6549 &self,
6550 task_id: &str,
6551 reason: Option<&str>,
6552 ) -> crate::async_task::Result<Option<TaskObject>> {
6553 self.inner.cancel_task(task_id, reason).await
6554 }
6555 }
6556
6557 #[tokio::test]
6558 async fn test_injected_task_store_used_by_dispatch() {
6559 let store = Arc::new(CountingTaskStore::new());
6560
6561 let add_tool = ToolBuilder::new("add")
6562 .description("Add two numbers")
6563 .task_support(TaskSupportMode::Optional)
6564 .handler(|input: AddInput| async move {
6565 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6566 })
6567 .build();
6568
6569 let mut router = McpRouter::new()
6570 .tool(add_tool)
6571 .task_store(store.clone() as Arc<dyn TaskStore>);
6572 init_router(&mut router).await;
6573
6574 let req = RouterRequest {
6576 id: RequestId::Number(1),
6577 inner: McpRequest::CallTool(CallToolParams {
6578 input_responses: None,
6579 request_state: None,
6580 name: "add".to_string(),
6581 arguments: serde_json::json!({"a": 2, "b": 3}),
6582 meta: None,
6583 task: Some(TaskRequestParams { ttl: None }),
6584 }),
6585 extensions: Extensions::new(),
6586 };
6587 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6588 let task_id = match resp.inner {
6589 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6590 other => panic!("Expected CreateTask response, got {other:?}"),
6591 };
6592
6593 assert_eq!(
6594 store.creates.load(std::sync::atomic::Ordering::Relaxed),
6595 1,
6596 "create_task must go through the injected store"
6597 );
6598
6599 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6601 assert_eq!(
6602 store.completes.load(std::sync::atomic::Ordering::Relaxed),
6603 1,
6604 "complete_task must go through the injected store"
6605 );
6606
6607 let gets_before = store.gets.load(std::sync::atomic::Ordering::Relaxed);
6609 let req = RouterRequest {
6610 id: RequestId::Number(2),
6611 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6612 task_id: task_id.clone(),
6613 meta: None,
6614 }),
6615 extensions: Extensions::new(),
6616 };
6617 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6618 match resp.inner {
6619 Ok(McpResponse::GetTaskInfo(info)) => {
6620 assert_eq!(info.task_id, task_id);
6621 assert_eq!(info.status, TaskStatus::Completed);
6622 }
6623 other => panic!("Expected GetTaskInfo response, got {other:?}"),
6624 }
6625 assert!(
6626 store.gets.load(std::sync::atomic::Ordering::Relaxed) > gets_before,
6627 "tasks/get must go through the injected store"
6628 );
6629 }
6630
6631 #[tokio::test]
6632 async fn test_removed_tasks_methods_get_method_not_found() {
6633 let mut router = McpRouter::new();
6637 init_router(&mut router).await;
6638
6639 for method in ["tasks/list", "tasks/result"] {
6640 let req = RouterRequest {
6641 id: RequestId::Number(1),
6642 inner: McpRequest::Unknown {
6643 method: method.to_string(),
6644 params: None,
6645 },
6646 extensions: Extensions::new(),
6647 };
6648
6649 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6650
6651 match resp.inner {
6652 Err(err) => {
6653 assert_eq!(err.code, -32601, "{method} must be MethodNotFound");
6654 }
6655 other => panic!("Expected MethodNotFound error for {method}, got {other:?}"),
6656 }
6657 }
6658 }
6659
6660 #[tokio::test]
6661 async fn test_task_lifecycle_complete() {
6662 let add_tool = ToolBuilder::new("add")
6663 .description("Add two numbers")
6664 .task_support(TaskSupportMode::Optional)
6665 .handler(|input: AddInput| async move {
6666 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6667 })
6668 .build();
6669
6670 let mut router = McpRouter::new().tool(add_tool);
6671 init_router(&mut router).await;
6672
6673 let req = RouterRequest {
6675 id: RequestId::Number(1),
6676 inner: McpRequest::CallTool(CallToolParams {
6677 input_responses: None,
6678 request_state: None,
6679 name: "add".to_string(),
6680 arguments: serde_json::json!({"a": 7, "b": 8}),
6681 meta: None,
6682 task: Some(TaskRequestParams { ttl: None }),
6683 }),
6684 extensions: Extensions::new(),
6685 };
6686
6687 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6688 let task_id = match resp.inner {
6689 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6690 _ => panic!("Expected CreateTask response"),
6691 };
6692
6693 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6695
6696 let req = RouterRequest {
6700 id: RequestId::Number(2),
6701 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6702 task_id: task_id.clone(),
6703 meta: None,
6704 }),
6705 extensions: Extensions::new(),
6706 };
6707
6708 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6709
6710 match resp.inner {
6711 Ok(McpResponse::GetTaskInfo(info)) => {
6712 assert_eq!(info.task_id, task_id);
6713 assert_eq!(info.status, TaskStatus::Completed);
6714 }
6715 _ => panic!("Expected GetTaskInfo response"),
6716 }
6717 }
6718
6719 #[tokio::test]
6720 async fn test_task_cancellation() {
6721 let slow_tool = ToolBuilder::new("slow")
6723 .description("Slow tool")
6724 .task_support(TaskSupportMode::Optional)
6725 .handler(|_input: serde_json::Value| async move {
6726 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
6727 Ok(CallToolResult::text("done"))
6728 })
6729 .build();
6730
6731 let mut router = McpRouter::new().tool(slow_tool);
6732 init_router(&mut router).await;
6733
6734 let req = RouterRequest {
6736 id: RequestId::Number(1),
6737 inner: McpRequest::CallTool(CallToolParams {
6738 input_responses: None,
6739 request_state: None,
6740 name: "slow".to_string(),
6741 arguments: serde_json::json!({}),
6742 meta: None,
6743 task: Some(TaskRequestParams { ttl: None }),
6744 }),
6745 extensions: Extensions::new(),
6746 };
6747
6748 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6749 let task_id = match resp.inner {
6750 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6751 _ => panic!("Expected CreateTask response"),
6752 };
6753
6754 let req = RouterRequest {
6756 id: RequestId::Number(2),
6757 inner: McpRequest::CancelTask(CancelTaskParams {
6758 task_id: task_id.clone(),
6759 reason: Some("Test cancellation".to_string()),
6760 meta: None,
6761 }),
6762 extensions: Extensions::new(),
6763 };
6764
6765 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6766
6767 match resp.inner {
6769 Ok(McpResponse::CancelTask(EmptyResult {})) => {}
6770 other => panic!("Expected empty CancelTask ack, got {other:?}"),
6771 }
6772
6773 let req = RouterRequest {
6775 id: RequestId::Number(3),
6776 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6777 task_id: task_id.clone(),
6778 meta: None,
6779 }),
6780 extensions: Extensions::new(),
6781 };
6782 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6783 match resp.inner {
6784 Ok(McpResponse::GetTaskInfo(info)) => {
6785 assert_eq!(info.status, TaskStatus::Cancelled);
6786 }
6787 _ => panic!("Expected GetTaskInfo response"),
6788 }
6789 }
6790
6791 #[tokio::test]
6792 async fn test_get_task_info() {
6793 let add_tool = ToolBuilder::new("add")
6794 .description("Add two numbers")
6795 .task_support(TaskSupportMode::Optional)
6796 .handler(|input: AddInput| async move {
6797 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6798 })
6799 .build();
6800
6801 let mut router = McpRouter::new().tool(add_tool);
6802 init_router(&mut router).await;
6803
6804 let req = RouterRequest {
6806 id: RequestId::Number(1),
6807 inner: McpRequest::CallTool(CallToolParams {
6808 input_responses: None,
6809 request_state: None,
6810 name: "add".to_string(),
6811 arguments: serde_json::json!({"a": 1, "b": 2}),
6812 meta: None,
6813 task: Some(TaskRequestParams { ttl: Some(600_000) }),
6814 }),
6815 extensions: Extensions::new(),
6816 };
6817
6818 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6819 let task_id = match resp.inner {
6820 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6821 _ => panic!("Expected CreateTask response"),
6822 };
6823
6824 let req = RouterRequest {
6826 id: RequestId::Number(2),
6827 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6828 task_id: task_id.clone(),
6829 meta: None,
6830 }),
6831 extensions: Extensions::new(),
6832 };
6833
6834 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6835
6836 match resp.inner {
6837 Ok(McpResponse::GetTaskInfo(info)) => {
6838 assert_eq!(info.task_id, task_id);
6839 assert!(info.created_at.contains('T')); assert_eq!(info.ttl, Some(600_000));
6841 }
6842 _ => panic!("Expected GetTaskInfo response"),
6843 }
6844 }
6845
6846 #[tokio::test]
6847 async fn test_task_forbidden_tool_rejects_task_params() {
6848 let tool = ToolBuilder::new("sync_only")
6849 .description("Sync only tool")
6850 .handler(|_input: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
6851 .build();
6852
6853 let mut router = McpRouter::new().tool(tool);
6854 init_router(&mut router).await;
6855
6856 let req = RouterRequest {
6858 id: RequestId::Number(1),
6859 inner: McpRequest::CallTool(CallToolParams {
6860 input_responses: None,
6861 request_state: None,
6862 name: "sync_only".to_string(),
6863 arguments: serde_json::json!({}),
6864 meta: None,
6865 task: Some(TaskRequestParams { ttl: None }),
6866 }),
6867 extensions: Extensions::new(),
6868 };
6869
6870 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6871
6872 match resp.inner {
6873 Err(e) => {
6874 assert!(e.message.contains("does not support async tasks"));
6875 }
6876 _ => panic!("Expected error response"),
6877 }
6878 }
6879
6880 #[tokio::test]
6881 async fn test_get_nonexistent_task() {
6882 let mut router = McpRouter::new();
6883 init_router(&mut router).await;
6884
6885 let req = RouterRequest {
6886 id: RequestId::Number(1),
6887 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6888 task_id: "task-999".to_string(),
6889 meta: None,
6890 }),
6891 extensions: Extensions::new(),
6892 };
6893
6894 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6895
6896 match resp.inner {
6897 Err(e) => {
6898 assert!(e.message.contains("not found"));
6899 }
6900 _ => panic!("Expected error response"),
6901 }
6902 }
6903
6904 #[tokio::test]
6909 async fn test_subscribe_to_resource() {
6910 use crate::resource::ResourceBuilder;
6911
6912 let resource = ResourceBuilder::new("file:///test.txt")
6913 .name("Test File")
6914 .text("Hello");
6915
6916 let mut router = McpRouter::new().resource(resource);
6917 init_router(&mut router).await;
6918
6919 let req = RouterRequest {
6921 id: RequestId::Number(1),
6922 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6923 uri: "file:///test.txt".to_string(),
6924 meta: None,
6925 }),
6926 extensions: Extensions::new(),
6927 };
6928
6929 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6930
6931 match resp.inner {
6932 Ok(McpResponse::SubscribeResource(_)) => {
6933 assert!(router.is_subscribed("file:///test.txt"));
6935 }
6936 _ => panic!("Expected SubscribeResource response"),
6937 }
6938 }
6939
6940 #[tokio::test]
6941 async fn test_unsubscribe_from_resource() {
6942 use crate::resource::ResourceBuilder;
6943
6944 let resource = ResourceBuilder::new("file:///test.txt")
6945 .name("Test File")
6946 .text("Hello");
6947
6948 let mut router = McpRouter::new().resource(resource);
6949 init_router(&mut router).await;
6950
6951 let req = RouterRequest {
6953 id: RequestId::Number(1),
6954 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6955 uri: "file:///test.txt".to_string(),
6956 meta: None,
6957 }),
6958 extensions: Extensions::new(),
6959 };
6960 let _ = router.ready().await.unwrap().call(req).await.unwrap();
6961 assert!(router.is_subscribed("file:///test.txt"));
6962
6963 let req = RouterRequest {
6965 id: RequestId::Number(2),
6966 inner: McpRequest::UnsubscribeResource(UnsubscribeResourceParams {
6967 uri: "file:///test.txt".to_string(),
6968 meta: None,
6969 }),
6970 extensions: Extensions::new(),
6971 };
6972
6973 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6974
6975 match resp.inner {
6976 Ok(McpResponse::UnsubscribeResource(_)) => {
6977 assert!(!router.is_subscribed("file:///test.txt"));
6979 }
6980 _ => panic!("Expected UnsubscribeResource response"),
6981 }
6982 }
6983
6984 #[tokio::test]
6985 async fn test_subscribe_nonexistent_resource() {
6986 let mut router = McpRouter::new();
6987 init_router(&mut router).await;
6988
6989 let req = RouterRequest {
6990 id: RequestId::Number(1),
6991 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6992 uri: "file:///nonexistent.txt".to_string(),
6993 meta: None,
6994 }),
6995 extensions: Extensions::new(),
6996 };
6997
6998 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6999
7000 match resp.inner {
7001 Err(e) => {
7002 assert!(e.message.contains("not found"));
7003 }
7004 _ => panic!("Expected error response"),
7005 }
7006 }
7007
7008 #[tokio::test]
7009 async fn test_notify_resource_updated() {
7010 use crate::context::notification_channel;
7011 use crate::resource::ResourceBuilder;
7012
7013 let (tx, mut rx) = notification_channel(10);
7014
7015 let resource = ResourceBuilder::new("file:///test.txt")
7016 .name("Test File")
7017 .text("Hello");
7018
7019 let router = McpRouter::new()
7020 .resource(resource)
7021 .with_notification_sender(tx);
7022
7023 router.subscribe("file:///test.txt");
7025
7026 let sent = router.notify_resource_updated("file:///test.txt");
7028 assert!(sent);
7029
7030 let notification = rx.try_recv().unwrap();
7032 match notification {
7033 ServerNotification::ResourceUpdated { uri } => {
7034 assert_eq!(uri, "file:///test.txt");
7035 }
7036 _ => panic!("Expected ResourceUpdated notification"),
7037 }
7038 }
7039
7040 #[tokio::test]
7041 async fn test_notify_resource_updated_not_subscribed() {
7042 use crate::context::notification_channel;
7043 use crate::resource::ResourceBuilder;
7044
7045 let (tx, mut rx) = notification_channel(10);
7046
7047 let resource = ResourceBuilder::new("file:///test.txt")
7048 .name("Test File")
7049 .text("Hello");
7050
7051 let router = McpRouter::new()
7052 .resource(resource)
7053 .with_notification_sender(tx);
7054
7055 let sent = router.notify_resource_updated("file:///test.txt");
7057 assert!(!sent); assert!(rx.try_recv().is_err());
7061 }
7062
7063 #[tokio::test]
7064 async fn test_notify_resources_list_changed() {
7065 use crate::context::notification_channel;
7066
7067 let (tx, mut rx) = notification_channel(10);
7068 let router = McpRouter::new().with_notification_sender(tx);
7069
7070 let sent = router.notify_resources_list_changed();
7071 assert!(sent);
7072
7073 let notification = rx.try_recv().unwrap();
7074 match notification {
7075 ServerNotification::ResourcesListChanged => {}
7076 _ => panic!("Expected ResourcesListChanged notification"),
7077 }
7078 }
7079
7080 #[tokio::test]
7081 async fn test_subscribed_uris() {
7082 use crate::resource::ResourceBuilder;
7083
7084 let resource1 = ResourceBuilder::new("file:///a.txt").name("A").text("A");
7085
7086 let resource2 = ResourceBuilder::new("file:///b.txt").name("B").text("B");
7087
7088 let router = McpRouter::new().resource(resource1).resource(resource2);
7089
7090 router.subscribe("file:///a.txt");
7092 router.subscribe("file:///b.txt");
7093
7094 let uris = router.subscribed_uris();
7095 assert_eq!(uris.len(), 2);
7096 assert!(uris.contains(&"file:///a.txt".to_string()));
7097 assert!(uris.contains(&"file:///b.txt".to_string()));
7098 }
7099
7100 #[tokio::test]
7101 async fn test_subscription_capability_advertised() {
7102 use crate::resource::ResourceBuilder;
7103
7104 let resource = ResourceBuilder::new("file:///test.txt")
7105 .name("Test")
7106 .text("Hello");
7107
7108 let mut router = McpRouter::new().resource(resource);
7109
7110 let init_req = RouterRequest {
7112 id: RequestId::Number(0),
7113 inner: McpRequest::Initialize(InitializeParams {
7114 protocol_version: "2025-11-25".to_string(),
7115 capabilities: ClientCapabilities {
7116 roots: None,
7117 sampling: None,
7118 elicitation: None,
7119 tasks: None,
7120 experimental: None,
7121 extensions: None,
7122 },
7123 client_info: Implementation {
7124 name: "test".to_string(),
7125 version: "1.0".to_string(),
7126 ..Default::default()
7127 },
7128 meta: None,
7129 }),
7130 extensions: Extensions::new(),
7131 };
7132 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
7133
7134 match resp.inner {
7135 Ok(McpResponse::Initialize(result)) => {
7136 let resources_cap = result.capabilities.resources.unwrap();
7138 assert!(resources_cap.subscribe);
7139 }
7140 _ => panic!("Expected Initialize response"),
7141 }
7142 }
7143
7144 #[tokio::test]
7145 async fn test_completion_handler() {
7146 let router = McpRouter::new()
7147 .server_info("test", "1.0")
7148 .completion_handler(|params: CompleteParams| async move {
7149 let prefix = ¶ms.argument.value;
7151 let suggestions: Vec<String> = vec!["alpha", "beta", "gamma"]
7152 .into_iter()
7153 .filter(|s| s.starts_with(prefix))
7154 .map(String::from)
7155 .collect();
7156 Ok(CompleteResult::new(suggestions))
7157 });
7158
7159 let init_req = RouterRequest {
7161 id: RequestId::Number(0),
7162 inner: McpRequest::Initialize(InitializeParams {
7163 protocol_version: "2025-11-25".to_string(),
7164 capabilities: ClientCapabilities::default(),
7165 client_info: Implementation {
7166 name: "test".to_string(),
7167 version: "1.0".to_string(),
7168 ..Default::default()
7169 },
7170 meta: None,
7171 }),
7172 extensions: Extensions::new(),
7173 };
7174 let resp = router
7175 .clone()
7176 .ready()
7177 .await
7178 .unwrap()
7179 .call(init_req)
7180 .await
7181 .unwrap();
7182
7183 match resp.inner {
7185 Ok(McpResponse::Initialize(result)) => {
7186 assert!(result.capabilities.completions.is_some());
7187 }
7188 _ => panic!("Expected Initialize response"),
7189 }
7190
7191 router.handle_notification(McpNotification::Initialized);
7193
7194 let complete_req = RouterRequest {
7196 id: RequestId::Number(1),
7197 inner: McpRequest::Complete(CompleteParams {
7198 reference: CompletionReference::prompt("test-prompt"),
7199 argument: CompletionArgument::new("query", "al"),
7200 context: None,
7201 meta: None,
7202 }),
7203 extensions: Extensions::new(),
7204 };
7205 let resp = router
7206 .clone()
7207 .ready()
7208 .await
7209 .unwrap()
7210 .call(complete_req)
7211 .await
7212 .unwrap();
7213
7214 match resp.inner {
7215 Ok(McpResponse::Complete(result)) => {
7216 assert_eq!(result.completion.values, vec!["alpha"]);
7217 }
7218 _ => panic!("Expected Complete response"),
7219 }
7220 }
7221
7222 #[tokio::test]
7223 async fn test_completion_without_handler_returns_empty() {
7224 let router = McpRouter::new().server_info("test", "1.0");
7225
7226 let init_req = RouterRequest {
7228 id: RequestId::Number(0),
7229 inner: McpRequest::Initialize(InitializeParams {
7230 protocol_version: "2025-11-25".to_string(),
7231 capabilities: ClientCapabilities::default(),
7232 client_info: Implementation {
7233 name: "test".to_string(),
7234 version: "1.0".to_string(),
7235 ..Default::default()
7236 },
7237 meta: None,
7238 }),
7239 extensions: Extensions::new(),
7240 };
7241 let resp = router
7242 .clone()
7243 .ready()
7244 .await
7245 .unwrap()
7246 .call(init_req)
7247 .await
7248 .unwrap();
7249
7250 match resp.inner {
7252 Ok(McpResponse::Initialize(result)) => {
7253 assert!(result.capabilities.completions.is_none());
7254 }
7255 _ => panic!("Expected Initialize response"),
7256 }
7257
7258 router.handle_notification(McpNotification::Initialized);
7260
7261 let complete_req = RouterRequest {
7263 id: RequestId::Number(1),
7264 inner: McpRequest::Complete(CompleteParams {
7265 reference: CompletionReference::prompt("test-prompt"),
7266 argument: CompletionArgument::new("query", "al"),
7267 context: None,
7268 meta: None,
7269 }),
7270 extensions: Extensions::new(),
7271 };
7272 let resp = router
7273 .clone()
7274 .ready()
7275 .await
7276 .unwrap()
7277 .call(complete_req)
7278 .await
7279 .unwrap();
7280
7281 match resp.inner {
7282 Ok(McpResponse::Complete(result)) => {
7283 assert!(result.completion.values.is_empty());
7284 }
7285 _ => panic!("Expected Complete response"),
7286 }
7287 }
7288
7289 #[tokio::test]
7290 async fn test_tool_filter_list() {
7291 use crate::filter::CapabilityFilter;
7292 use crate::tool::Tool;
7293
7294 let public_tool = ToolBuilder::new("public")
7295 .description("Public tool")
7296 .handler(|_: AddInput| async move { Ok(CallToolResult::text("public")) })
7297 .build();
7298
7299 let admin_tool = ToolBuilder::new("admin")
7300 .description("Admin tool")
7301 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
7302 .build();
7303
7304 let mut router = McpRouter::new()
7305 .tool(public_tool)
7306 .tool(admin_tool)
7307 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| tool.name != "admin"));
7308
7309 init_router(&mut router).await;
7311
7312 let req = RouterRequest {
7313 id: RequestId::Number(1),
7314 inner: McpRequest::ListTools(ListToolsParams::default()),
7315 extensions: Extensions::new(),
7316 };
7317
7318 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7319
7320 match resp.inner {
7321 Ok(McpResponse::ListTools(result)) => {
7322 assert_eq!(result.tools.len(), 1);
7324 assert_eq!(result.tools[0].name, "public");
7325 }
7326 _ => panic!("Expected ListTools response"),
7327 }
7328 }
7329
7330 #[tokio::test]
7331 async fn test_tool_filter_call_denied() {
7332 use crate::filter::CapabilityFilter;
7333 use crate::tool::Tool;
7334
7335 let admin_tool = ToolBuilder::new("admin")
7336 .description("Admin tool")
7337 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
7338 .build();
7339
7340 let mut router = McpRouter::new()
7341 .tool(admin_tool)
7342 .tool_filter(CapabilityFilter::new(|_, _: &Tool| false)); init_router(&mut router).await;
7346
7347 let req = RouterRequest {
7348 id: RequestId::Number(1),
7349 inner: McpRequest::CallTool(CallToolParams {
7350 input_responses: None,
7351 request_state: None,
7352 name: "admin".to_string(),
7353 arguments: serde_json::json!({"a": 1, "b": 2}),
7354 meta: None,
7355 task: None,
7356 }),
7357 extensions: Extensions::new(),
7358 };
7359
7360 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7361
7362 match resp.inner {
7364 Err(e) => {
7365 assert_eq!(e.code, -32601); }
7367 _ => panic!("Expected JsonRpc error"),
7368 }
7369 }
7370
7371 #[tokio::test]
7372 async fn test_tool_filter_call_allowed() {
7373 use crate::filter::CapabilityFilter;
7374 use crate::tool::Tool;
7375
7376 let public_tool = ToolBuilder::new("public")
7377 .description("Public tool")
7378 .handler(|input: AddInput| async move {
7379 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
7380 })
7381 .build();
7382
7383 let mut router = McpRouter::new()
7384 .tool(public_tool)
7385 .tool_filter(CapabilityFilter::new(|_, _: &Tool| true)); init_router(&mut router).await;
7389
7390 let req = RouterRequest {
7391 id: RequestId::Number(1),
7392 inner: McpRequest::CallTool(CallToolParams {
7393 input_responses: None,
7394 request_state: None,
7395 name: "public".to_string(),
7396 arguments: serde_json::json!({"a": 1, "b": 2}),
7397 meta: None,
7398 task: None,
7399 }),
7400 extensions: Extensions::new(),
7401 };
7402
7403 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7404
7405 match resp.inner {
7406 Ok(McpResponse::CallTool(result)) => {
7407 assert!(!result.is_error);
7408 }
7409 _ => panic!("Expected CallTool response"),
7410 }
7411 }
7412
7413 #[tokio::test]
7414 async fn test_tool_filter_custom_denial() {
7415 use crate::filter::{CapabilityFilter, DenialBehavior};
7416 use crate::tool::Tool;
7417
7418 let admin_tool = ToolBuilder::new("admin")
7419 .description("Admin tool")
7420 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
7421 .build();
7422
7423 let mut router = McpRouter::new().tool(admin_tool).tool_filter(
7424 CapabilityFilter::new(|_, _: &Tool| false)
7425 .denial_behavior(DenialBehavior::Unauthorized),
7426 );
7427
7428 init_router(&mut router).await;
7430
7431 let req = RouterRequest {
7432 id: RequestId::Number(1),
7433 inner: McpRequest::CallTool(CallToolParams {
7434 input_responses: None,
7435 request_state: None,
7436 name: "admin".to_string(),
7437 arguments: serde_json::json!({"a": 1, "b": 2}),
7438 meta: None,
7439 task: None,
7440 }),
7441 extensions: Extensions::new(),
7442 };
7443
7444 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7445
7446 match resp.inner {
7448 Err(e) => {
7449 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7451 }
7452 _ => panic!("Expected JsonRpc error"),
7453 }
7454 }
7455
7456 #[tokio::test]
7457 async fn test_resource_filter_list() {
7458 use crate::filter::CapabilityFilter;
7459 use crate::resource::{Resource, ResourceBuilder};
7460
7461 let public_resource = ResourceBuilder::new("file:///public.txt")
7462 .name("Public File")
7463 .text("public content");
7464
7465 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7466 .name("Secret File")
7467 .text("secret content");
7468
7469 let mut router = McpRouter::new()
7470 .resource(public_resource)
7471 .resource(secret_resource)
7472 .resource_filter(CapabilityFilter::new(|_, r: &Resource| {
7473 !r.name.contains("Secret")
7474 }));
7475
7476 init_router(&mut router).await;
7478
7479 let req = RouterRequest {
7480 id: RequestId::Number(1),
7481 inner: McpRequest::ListResources(ListResourcesParams::default()),
7482 extensions: Extensions::new(),
7483 };
7484
7485 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7486
7487 match resp.inner {
7488 Ok(McpResponse::ListResources(result)) => {
7489 assert_eq!(result.resources.len(), 1);
7491 assert_eq!(result.resources[0].name, "Public File");
7492 }
7493 _ => panic!("Expected ListResources response"),
7494 }
7495 }
7496
7497 #[tokio::test]
7498 async fn test_resource_filter_read_denied() {
7499 use crate::filter::CapabilityFilter;
7500 use crate::resource::{Resource, ResourceBuilder};
7501
7502 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7503 .name("Secret File")
7504 .text("secret content");
7505
7506 let mut router = McpRouter::new()
7507 .resource(secret_resource)
7508 .resource_filter(CapabilityFilter::new(|_, _: &Resource| false)); init_router(&mut router).await;
7512
7513 let req = RouterRequest {
7514 id: RequestId::Number(1),
7515 inner: McpRequest::ReadResource(ReadResourceParams {
7516 input_responses: None,
7517 request_state: None,
7518 uri: "file:///secret.txt".to_string(),
7519 meta: None,
7520 }),
7521 extensions: Extensions::new(),
7522 };
7523
7524 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7525
7526 match resp.inner {
7528 Err(e) => {
7529 assert_eq!(e.code, -32601); }
7531 _ => panic!("Expected JsonRpc error"),
7532 }
7533 }
7534
7535 #[tokio::test]
7536 async fn test_resource_filter_read_allowed() {
7537 use crate::filter::CapabilityFilter;
7538 use crate::resource::{Resource, ResourceBuilder};
7539
7540 let public_resource = ResourceBuilder::new("file:///public.txt")
7541 .name("Public File")
7542 .text("public content");
7543
7544 let mut router = McpRouter::new()
7545 .resource(public_resource)
7546 .resource_filter(CapabilityFilter::new(|_, _: &Resource| true)); init_router(&mut router).await;
7550
7551 let req = RouterRequest {
7552 id: RequestId::Number(1),
7553 inner: McpRequest::ReadResource(ReadResourceParams {
7554 input_responses: None,
7555 request_state: None,
7556 uri: "file:///public.txt".to_string(),
7557 meta: None,
7558 }),
7559 extensions: Extensions::new(),
7560 };
7561
7562 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7563
7564 match resp.inner {
7565 Ok(McpResponse::ReadResource(result)) => {
7566 assert_eq!(result.contents.len(), 1);
7567 assert_eq!(result.contents[0].text.as_deref(), Some("public content"));
7568 }
7569 _ => panic!("Expected ReadResource response"),
7570 }
7571 }
7572
7573 #[tokio::test]
7574 async fn test_resource_filter_custom_denial() {
7575 use crate::filter::{CapabilityFilter, DenialBehavior};
7576 use crate::resource::{Resource, ResourceBuilder};
7577
7578 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7579 .name("Secret File")
7580 .text("secret content");
7581
7582 let mut router = McpRouter::new().resource(secret_resource).resource_filter(
7583 CapabilityFilter::new(|_, _: &Resource| false)
7584 .denial_behavior(DenialBehavior::Unauthorized),
7585 );
7586
7587 init_router(&mut router).await;
7589
7590 let req = RouterRequest {
7591 id: RequestId::Number(1),
7592 inner: McpRequest::ReadResource(ReadResourceParams {
7593 input_responses: None,
7594 request_state: None,
7595 uri: "file:///secret.txt".to_string(),
7596 meta: None,
7597 }),
7598 extensions: Extensions::new(),
7599 };
7600
7601 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7602
7603 match resp.inner {
7605 Err(e) => {
7606 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7608 }
7609 _ => panic!("Expected JsonRpc error"),
7610 }
7611 }
7612
7613 #[tokio::test]
7614 async fn test_prompt_filter_list() {
7615 use crate::filter::CapabilityFilter;
7616 use crate::prompt::{Prompt, PromptBuilder};
7617
7618 let public_prompt = PromptBuilder::new("greeting")
7619 .description("A greeting")
7620 .user_message("Hello!");
7621
7622 let admin_prompt = PromptBuilder::new("system_debug")
7623 .description("Admin prompt")
7624 .user_message("Debug");
7625
7626 let mut router = McpRouter::new()
7627 .prompt(public_prompt)
7628 .prompt(admin_prompt)
7629 .prompt_filter(CapabilityFilter::new(|_, p: &Prompt| {
7630 !p.name.contains("system")
7631 }));
7632
7633 init_router(&mut router).await;
7635
7636 let req = RouterRequest {
7637 id: RequestId::Number(1),
7638 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7639 extensions: Extensions::new(),
7640 };
7641
7642 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7643
7644 match resp.inner {
7645 Ok(McpResponse::ListPrompts(result)) => {
7646 assert_eq!(result.prompts.len(), 1);
7648 assert_eq!(result.prompts[0].name, "greeting");
7649 }
7650 _ => panic!("Expected ListPrompts response"),
7651 }
7652 }
7653
7654 #[tokio::test]
7655 async fn test_prompt_filter_get_denied() {
7656 use crate::filter::CapabilityFilter;
7657 use crate::prompt::{Prompt, PromptBuilder};
7658 use std::collections::HashMap;
7659
7660 let admin_prompt = PromptBuilder::new("system_debug")
7661 .description("Admin prompt")
7662 .user_message("Debug");
7663
7664 let mut router = McpRouter::new()
7665 .prompt(admin_prompt)
7666 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| false)); init_router(&mut router).await;
7670
7671 let req = RouterRequest {
7672 id: RequestId::Number(1),
7673 inner: McpRequest::GetPrompt(GetPromptParams {
7674 input_responses: None,
7675 request_state: None,
7676 name: "system_debug".to_string(),
7677 arguments: HashMap::new(),
7678 meta: None,
7679 }),
7680 extensions: Extensions::new(),
7681 };
7682
7683 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7684
7685 match resp.inner {
7687 Err(e) => {
7688 assert_eq!(e.code, -32601); }
7690 _ => panic!("Expected JsonRpc error"),
7691 }
7692 }
7693
7694 #[tokio::test]
7695 async fn test_prompt_filter_get_allowed() {
7696 use crate::filter::CapabilityFilter;
7697 use crate::prompt::{Prompt, PromptBuilder};
7698 use std::collections::HashMap;
7699
7700 let public_prompt = PromptBuilder::new("greeting")
7701 .description("A greeting")
7702 .user_message("Hello!");
7703
7704 let mut router = McpRouter::new()
7705 .prompt(public_prompt)
7706 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| true)); init_router(&mut router).await;
7710
7711 let req = RouterRequest {
7712 id: RequestId::Number(1),
7713 inner: McpRequest::GetPrompt(GetPromptParams {
7714 input_responses: None,
7715 request_state: None,
7716 name: "greeting".to_string(),
7717 arguments: HashMap::new(),
7718 meta: None,
7719 }),
7720 extensions: Extensions::new(),
7721 };
7722
7723 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7724
7725 match resp.inner {
7726 Ok(McpResponse::GetPrompt(result)) => {
7727 assert_eq!(result.messages.len(), 1);
7728 }
7729 _ => panic!("Expected GetPrompt response"),
7730 }
7731 }
7732
7733 #[tokio::test]
7734 async fn test_prompt_filter_custom_denial() {
7735 use crate::filter::{CapabilityFilter, DenialBehavior};
7736 use crate::prompt::{Prompt, PromptBuilder};
7737 use std::collections::HashMap;
7738
7739 let admin_prompt = PromptBuilder::new("system_debug")
7740 .description("Admin prompt")
7741 .user_message("Debug");
7742
7743 let mut router = McpRouter::new().prompt(admin_prompt).prompt_filter(
7744 CapabilityFilter::new(|_, _: &Prompt| false)
7745 .denial_behavior(DenialBehavior::Unauthorized),
7746 );
7747
7748 init_router(&mut router).await;
7750
7751 let req = RouterRequest {
7752 id: RequestId::Number(1),
7753 inner: McpRequest::GetPrompt(GetPromptParams {
7754 input_responses: None,
7755 request_state: None,
7756 name: "system_debug".to_string(),
7757 arguments: HashMap::new(),
7758 meta: None,
7759 }),
7760 extensions: Extensions::new(),
7761 };
7762
7763 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7764
7765 match resp.inner {
7767 Err(e) => {
7768 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7770 }
7771 _ => panic!("Expected JsonRpc error"),
7772 }
7773 }
7774
7775 #[derive(Debug, Deserialize, JsonSchema)]
7780 struct StringInput {
7781 value: String,
7782 }
7783
7784 #[tokio::test]
7785 async fn test_router_merge_tools() {
7786 let tool_a = ToolBuilder::new("tool_a")
7788 .description("Tool A")
7789 .handler(|_: StringInput| async move { Ok(CallToolResult::text("A")) })
7790 .build();
7791
7792 let router_a = McpRouter::new().tool(tool_a);
7793
7794 let tool_b = ToolBuilder::new("tool_b")
7796 .description("Tool B")
7797 .handler(|_: StringInput| async move { Ok(CallToolResult::text("B")) })
7798 .build();
7799 let tool_c = ToolBuilder::new("tool_c")
7800 .description("Tool C")
7801 .handler(|_: StringInput| async move { Ok(CallToolResult::text("C")) })
7802 .build();
7803
7804 let router_b = McpRouter::new().tool(tool_b).tool(tool_c);
7805
7806 let mut merged = McpRouter::new()
7808 .server_info("merged", "1.0")
7809 .merge(router_a)
7810 .merge(router_b);
7811
7812 init_router(&mut merged).await;
7813
7814 let req = RouterRequest {
7816 id: RequestId::Number(1),
7817 inner: McpRequest::ListTools(ListToolsParams::default()),
7818 extensions: Extensions::new(),
7819 };
7820
7821 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7822
7823 match resp.inner {
7824 Ok(McpResponse::ListTools(result)) => {
7825 assert_eq!(result.tools.len(), 3);
7826 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7827 assert!(names.contains(&"tool_a"));
7828 assert!(names.contains(&"tool_b"));
7829 assert!(names.contains(&"tool_c"));
7830 }
7831 _ => panic!("Expected ListTools response"),
7832 }
7833 }
7834
7835 #[tokio::test]
7836 async fn test_router_merge_overwrites_duplicates() {
7837 let tool_v1 = ToolBuilder::new("shared")
7839 .description("Version 1")
7840 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v1")) })
7841 .build();
7842
7843 let router_a = McpRouter::new().tool(tool_v1);
7844
7845 let tool_v2 = ToolBuilder::new("shared")
7847 .description("Version 2")
7848 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v2")) })
7849 .build();
7850
7851 let router_b = McpRouter::new().tool(tool_v2);
7852
7853 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7855
7856 init_router(&mut merged).await;
7857
7858 let req = RouterRequest {
7859 id: RequestId::Number(1),
7860 inner: McpRequest::ListTools(ListToolsParams::default()),
7861 extensions: Extensions::new(),
7862 };
7863
7864 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7865
7866 match resp.inner {
7867 Ok(McpResponse::ListTools(result)) => {
7868 assert_eq!(result.tools.len(), 1);
7869 assert_eq!(result.tools[0].name, "shared");
7870 assert_eq!(result.tools[0].description.as_deref(), Some("Version 2"));
7871 }
7872 _ => panic!("Expected ListTools response"),
7873 }
7874 }
7875
7876 #[tokio::test]
7877 async fn test_router_merge_resources() {
7878 use crate::resource::ResourceBuilder;
7879
7880 let router_a = McpRouter::new().resource(
7882 ResourceBuilder::new("file:///a.txt")
7883 .name("File A")
7884 .text("content a"),
7885 );
7886
7887 let router_b = McpRouter::new().resource(
7888 ResourceBuilder::new("file:///b.txt")
7889 .name("File B")
7890 .text("content b"),
7891 );
7892
7893 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7894
7895 init_router(&mut merged).await;
7896
7897 let req = RouterRequest {
7898 id: RequestId::Number(1),
7899 inner: McpRequest::ListResources(ListResourcesParams::default()),
7900 extensions: Extensions::new(),
7901 };
7902
7903 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7904
7905 match resp.inner {
7906 Ok(McpResponse::ListResources(result)) => {
7907 assert_eq!(result.resources.len(), 2);
7908 let uris: Vec<&str> = result.resources.iter().map(|r| r.uri.as_str()).collect();
7909 assert!(uris.contains(&"file:///a.txt"));
7910 assert!(uris.contains(&"file:///b.txt"));
7911 }
7912 _ => panic!("Expected ListResources response"),
7913 }
7914 }
7915
7916 #[tokio::test]
7917 async fn test_router_merge_prompts() {
7918 use crate::prompt::PromptBuilder;
7919
7920 let router_a =
7921 McpRouter::new().prompt(PromptBuilder::new("prompt_a").user_message("Hello A"));
7922
7923 let router_b =
7924 McpRouter::new().prompt(PromptBuilder::new("prompt_b").user_message("Hello B"));
7925
7926 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7927
7928 init_router(&mut merged).await;
7929
7930 let req = RouterRequest {
7931 id: RequestId::Number(1),
7932 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7933 extensions: Extensions::new(),
7934 };
7935
7936 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7937
7938 match resp.inner {
7939 Ok(McpResponse::ListPrompts(result)) => {
7940 assert_eq!(result.prompts.len(), 2);
7941 let names: Vec<&str> = result.prompts.iter().map(|p| p.name.as_str()).collect();
7942 assert!(names.contains(&"prompt_a"));
7943 assert!(names.contains(&"prompt_b"));
7944 }
7945 _ => panic!("Expected ListPrompts response"),
7946 }
7947 }
7948
7949 #[tokio::test]
7950 async fn test_router_nest_prefixes_tools() {
7951 let tool_query = ToolBuilder::new("query")
7953 .description("Query the database")
7954 .handler(|_: StringInput| async move { Ok(CallToolResult::text("query result")) })
7955 .build();
7956 let tool_insert = ToolBuilder::new("insert")
7957 .description("Insert into database")
7958 .handler(|_: StringInput| async move { Ok(CallToolResult::text("insert result")) })
7959 .build();
7960
7961 let db_router = McpRouter::new().tool(tool_query).tool(tool_insert);
7962
7963 let mut router = McpRouter::new()
7965 .server_info("nested", "1.0")
7966 .nest("db", db_router);
7967
7968 init_router(&mut router).await;
7969
7970 let req = RouterRequest {
7971 id: RequestId::Number(1),
7972 inner: McpRequest::ListTools(ListToolsParams::default()),
7973 extensions: Extensions::new(),
7974 };
7975
7976 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7977
7978 match resp.inner {
7979 Ok(McpResponse::ListTools(result)) => {
7980 assert_eq!(result.tools.len(), 2);
7981 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7982 assert!(names.contains(&"db.query"));
7983 assert!(names.contains(&"db.insert"));
7984 }
7985 _ => panic!("Expected ListTools response"),
7986 }
7987 }
7988
7989 #[tokio::test]
7990 async fn test_router_nest_call_prefixed_tool() {
7991 let tool = ToolBuilder::new("echo")
7992 .description("Echo input")
7993 .handler(|input: StringInput| async move { Ok(CallToolResult::text(&input.value)) })
7994 .build();
7995
7996 let nested_router = McpRouter::new().tool(tool);
7997
7998 let mut router = McpRouter::new().nest("api", nested_router);
7999
8000 init_router(&mut router).await;
8001
8002 let req = RouterRequest {
8004 id: RequestId::Number(1),
8005 inner: McpRequest::CallTool(CallToolParams {
8006 input_responses: None,
8007 request_state: None,
8008 name: "api.echo".to_string(),
8009 arguments: serde_json::json!({"value": "hello world"}),
8010 meta: None,
8011 task: None,
8012 }),
8013 extensions: Extensions::new(),
8014 };
8015
8016 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8017
8018 match resp.inner {
8019 Ok(McpResponse::CallTool(result)) => {
8020 assert!(!result.is_error);
8021 match &result.content[0] {
8022 Content::Text { text, .. } => assert_eq!(text, "hello world"),
8023 _ => panic!("Expected text content"),
8024 }
8025 }
8026 _ => panic!("Expected CallTool response"),
8027 }
8028 }
8029
8030 #[tokio::test]
8031 async fn test_router_multiple_nests() {
8032 let db_tool = ToolBuilder::new("query")
8033 .description("Database query")
8034 .handler(|_: StringInput| async move { Ok(CallToolResult::text("db")) })
8035 .build();
8036
8037 let api_tool = ToolBuilder::new("fetch")
8038 .description("API fetch")
8039 .handler(|_: StringInput| async move { Ok(CallToolResult::text("api")) })
8040 .build();
8041
8042 let db_router = McpRouter::new().tool(db_tool);
8043 let api_router = McpRouter::new().tool(api_tool);
8044
8045 let mut router = McpRouter::new()
8046 .nest("db", db_router)
8047 .nest("api", api_router);
8048
8049 init_router(&mut router).await;
8050
8051 let req = RouterRequest {
8052 id: RequestId::Number(1),
8053 inner: McpRequest::ListTools(ListToolsParams::default()),
8054 extensions: Extensions::new(),
8055 };
8056
8057 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8058
8059 match resp.inner {
8060 Ok(McpResponse::ListTools(result)) => {
8061 assert_eq!(result.tools.len(), 2);
8062 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8063 assert!(names.contains(&"db.query"));
8064 assert!(names.contains(&"api.fetch"));
8065 }
8066 _ => panic!("Expected ListTools response"),
8067 }
8068 }
8069
8070 #[tokio::test]
8071 async fn test_router_merge_and_nest_combined() {
8072 let tool_a = ToolBuilder::new("local")
8074 .description("Local tool")
8075 .handler(|_: StringInput| async move { Ok(CallToolResult::text("local")) })
8076 .build();
8077
8078 let nested_tool = ToolBuilder::new("remote")
8079 .description("Remote tool")
8080 .handler(|_: StringInput| async move { Ok(CallToolResult::text("remote")) })
8081 .build();
8082
8083 let nested_router = McpRouter::new().tool(nested_tool);
8084
8085 let mut router = McpRouter::new()
8086 .tool(tool_a)
8087 .nest("external", nested_router);
8088
8089 init_router(&mut router).await;
8090
8091 let req = RouterRequest {
8092 id: RequestId::Number(1),
8093 inner: McpRequest::ListTools(ListToolsParams::default()),
8094 extensions: Extensions::new(),
8095 };
8096
8097 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8098
8099 match resp.inner {
8100 Ok(McpResponse::ListTools(result)) => {
8101 assert_eq!(result.tools.len(), 2);
8102 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8103 assert!(names.contains(&"local"));
8104 assert!(names.contains(&"external.remote"));
8105 }
8106 _ => panic!("Expected ListTools response"),
8107 }
8108 }
8109
8110 #[tokio::test]
8111 async fn test_router_merge_preserves_server_info() {
8112 let child_router = McpRouter::new()
8113 .server_info("child", "2.0")
8114 .instructions("Child instructions");
8115
8116 let mut router = McpRouter::new()
8117 .server_info("parent", "1.0")
8118 .instructions("Parent instructions")
8119 .merge(child_router);
8120
8121 init_router(&mut router).await;
8122
8123 let init_req = RouterRequest {
8125 id: RequestId::Number(99),
8126 inner: McpRequest::Initialize(InitializeParams {
8127 protocol_version: "2025-11-25".to_string(),
8128 capabilities: ClientCapabilities::default(),
8129 client_info: Implementation {
8130 name: "test".to_string(),
8131 version: "1.0".to_string(),
8132 ..Default::default()
8133 },
8134 meta: None,
8135 }),
8136 extensions: Extensions::new(),
8137 };
8138
8139 let child_router2 = McpRouter::new().server_info("child", "2.0");
8141 let mut fresh_router = McpRouter::new()
8142 .server_info("parent", "1.0")
8143 .merge(child_router2);
8144
8145 let resp = fresh_router
8146 .ready()
8147 .await
8148 .unwrap()
8149 .call(init_req)
8150 .await
8151 .unwrap();
8152
8153 match resp.inner {
8154 Ok(McpResponse::Initialize(result)) => {
8155 assert_eq!(result.server_info.name, "parent");
8156 assert_eq!(result.server_info.version, "1.0");
8157 }
8158 _ => panic!("Expected Initialize response"),
8159 }
8160 }
8161
8162 #[tokio::test]
8167 async fn test_auto_instructions_tools_only() {
8168 let tool_a = ToolBuilder::new("alpha")
8169 .description("Alpha tool")
8170 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8171 .build();
8172 let tool_b = ToolBuilder::new("beta")
8173 .description("Beta tool")
8174 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8175 .build();
8176
8177 let mut router = McpRouter::new()
8178 .auto_instructions()
8179 .tool(tool_a)
8180 .tool(tool_b);
8181
8182 let resp = send_initialize(&mut router).await;
8183 let instructions = resp.instructions.expect("should have instructions");
8184
8185 assert!(instructions.contains("## Tools"));
8186 assert!(instructions.contains("- **alpha**: Alpha tool"));
8187 assert!(instructions.contains("- **beta**: Beta tool"));
8188 assert!(!instructions.contains("## Resources"));
8190 assert!(!instructions.contains("## Prompts"));
8191 }
8192
8193 #[tokio::test]
8194 async fn test_auto_instructions_with_annotations() {
8195 let read_only_tool = ToolBuilder::new("query")
8196 .description("Run a query")
8197 .read_only()
8198 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8199 .build();
8200 let destructive_tool = ToolBuilder::new("delete")
8201 .description("Delete a record")
8202 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8203 .build();
8204 let idempotent_tool = ToolBuilder::new("upsert")
8205 .description("Upsert a record")
8206 .non_destructive()
8207 .idempotent()
8208 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8209 .build();
8210
8211 let mut router = McpRouter::new()
8212 .auto_instructions()
8213 .tool(read_only_tool)
8214 .tool(destructive_tool)
8215 .tool(idempotent_tool);
8216
8217 let resp = send_initialize(&mut router).await;
8218 let instructions = resp.instructions.unwrap();
8219
8220 assert!(instructions.contains("- **query**: Run a query [read-only]"));
8221 assert!(instructions.contains("- **delete**: Delete a record\n"));
8223 assert!(instructions.contains("- **upsert**: Upsert a record [idempotent]"));
8224 }
8225
8226 #[tokio::test]
8227 async fn test_auto_instructions_with_resources() {
8228 use crate::resource::ResourceBuilder;
8229
8230 let resource = ResourceBuilder::new("file:///schema.sql")
8231 .name("Schema")
8232 .description("Database schema")
8233 .text("CREATE TABLE ...");
8234
8235 let mut router = McpRouter::new().auto_instructions().resource(resource);
8236
8237 let resp = send_initialize(&mut router).await;
8238 let instructions = resp.instructions.unwrap();
8239
8240 assert!(instructions.contains("## Resources"));
8241 assert!(instructions.contains("- **file:///schema.sql**: Database schema"));
8242 assert!(!instructions.contains("## Tools"));
8243 }
8244
8245 #[tokio::test]
8246 async fn test_auto_instructions_with_resource_templates() {
8247 use crate::resource::ResourceTemplateBuilder;
8248
8249 let template = ResourceTemplateBuilder::new("file:///{path}")
8250 .name("File")
8251 .description("Read a file by path")
8252 .handler(
8253 |_uri: String, _vars: std::collections::HashMap<String, String>| async move {
8254 Ok(crate::ReadResourceResult::text("content", "text/plain"))
8255 },
8256 );
8257
8258 let mut router = McpRouter::new()
8259 .auto_instructions()
8260 .resource_template(template);
8261
8262 let resp = send_initialize(&mut router).await;
8263 let instructions = resp.instructions.unwrap();
8264
8265 assert!(instructions.contains("## Resources"));
8266 assert!(instructions.contains("- **file:///{path}**: Read a file by path"));
8267 }
8268
8269 #[tokio::test]
8270 async fn test_auto_instructions_with_prompts() {
8271 use crate::prompt::PromptBuilder;
8272
8273 let prompt = PromptBuilder::new("write_query")
8274 .description("Help write a SQL query")
8275 .user_message("Write a query for: {task}");
8276
8277 let mut router = McpRouter::new().auto_instructions().prompt(prompt);
8278
8279 let resp = send_initialize(&mut router).await;
8280 let instructions = resp.instructions.unwrap();
8281
8282 assert!(instructions.contains("## Prompts"));
8283 assert!(instructions.contains("- **write_query**: Help write a SQL query"));
8284 assert!(!instructions.contains("## Tools"));
8285 }
8286
8287 #[tokio::test]
8288 async fn test_auto_instructions_all_sections() {
8289 use crate::prompt::PromptBuilder;
8290 use crate::resource::ResourceBuilder;
8291
8292 let tool = ToolBuilder::new("query")
8293 .description("Execute SQL")
8294 .read_only()
8295 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8296 .build();
8297 let resource = ResourceBuilder::new("db://schema")
8298 .name("Schema")
8299 .description("Full database schema")
8300 .text("schema");
8301 let prompt = PromptBuilder::new("write_query")
8302 .description("Help write a SQL query")
8303 .user_message("Write a query");
8304
8305 let mut router = McpRouter::new()
8306 .auto_instructions()
8307 .tool(tool)
8308 .resource(resource)
8309 .prompt(prompt);
8310
8311 let resp = send_initialize(&mut router).await;
8312 let instructions = resp.instructions.unwrap();
8313
8314 assert!(instructions.contains("## Tools"));
8316 assert!(instructions.contains("## Resources"));
8317 assert!(instructions.contains("## Prompts"));
8318
8319 let tools_pos = instructions.find("## Tools").unwrap();
8321 let resources_pos = instructions.find("## Resources").unwrap();
8322 let prompts_pos = instructions.find("## Prompts").unwrap();
8323 assert!(tools_pos < resources_pos);
8324 assert!(resources_pos < prompts_pos);
8325 }
8326
8327 #[tokio::test]
8328 async fn test_auto_instructions_with_prefix_and_suffix() {
8329 let tool = ToolBuilder::new("echo")
8330 .description("Echo input")
8331 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8332 .build();
8333
8334 let mut router = McpRouter::new()
8335 .auto_instructions_with(
8336 Some("This server provides echo capabilities."),
8337 Some("Contact admin@example.com for support."),
8338 )
8339 .tool(tool);
8340
8341 let resp = send_initialize(&mut router).await;
8342 let instructions = resp.instructions.unwrap();
8343
8344 assert!(instructions.starts_with("This server provides echo capabilities."));
8345 assert!(instructions.ends_with("Contact admin@example.com for support."));
8346 assert!(instructions.contains("## Tools"));
8347 assert!(instructions.contains("- **echo**: Echo input"));
8348 }
8349
8350 #[tokio::test]
8351 async fn test_auto_instructions_prefix_only() {
8352 let tool = ToolBuilder::new("echo")
8353 .description("Echo input")
8354 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8355 .build();
8356
8357 let mut router = McpRouter::new()
8358 .auto_instructions_with(Some("My server intro."), None::<String>)
8359 .tool(tool);
8360
8361 let resp = send_initialize(&mut router).await;
8362 let instructions = resp.instructions.unwrap();
8363
8364 assert!(instructions.starts_with("My server intro."));
8365 assert!(instructions.contains("- **echo**: Echo input"));
8366 }
8367
8368 #[tokio::test]
8369 async fn test_auto_instructions_empty_router() {
8370 let mut router = McpRouter::new().auto_instructions();
8371
8372 let resp = send_initialize(&mut router).await;
8373 let instructions = resp.instructions.expect("should have instructions");
8374
8375 assert!(!instructions.contains("## Tools"));
8377 assert!(!instructions.contains("## Resources"));
8378 assert!(!instructions.contains("## Prompts"));
8379 assert!(instructions.is_empty());
8380 }
8381
8382 #[tokio::test]
8383 async fn test_auto_instructions_overrides_manual() {
8384 let tool = ToolBuilder::new("echo")
8385 .description("Echo input")
8386 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8387 .build();
8388
8389 let mut router = McpRouter::new()
8390 .instructions("This will be overridden")
8391 .auto_instructions()
8392 .tool(tool);
8393
8394 let resp = send_initialize(&mut router).await;
8395 let instructions = resp.instructions.unwrap();
8396
8397 assert!(!instructions.contains("This will be overridden"));
8398 assert!(instructions.contains("- **echo**: Echo input"));
8399 }
8400
8401 #[tokio::test]
8402 async fn test_no_auto_instructions_returns_manual() {
8403 let tool = ToolBuilder::new("echo")
8404 .description("Echo input")
8405 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8406 .build();
8407
8408 let mut router = McpRouter::new()
8409 .instructions("Manual instructions here")
8410 .tool(tool);
8411
8412 let resp = send_initialize(&mut router).await;
8413 let instructions = resp.instructions.unwrap();
8414
8415 assert_eq!(instructions, "Manual instructions here");
8416 }
8417
8418 #[tokio::test]
8419 async fn test_auto_instructions_no_description_fallback() {
8420 let tool = ToolBuilder::new("mystery")
8421 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8422 .build();
8423
8424 let mut router = McpRouter::new().auto_instructions().tool(tool);
8425
8426 let resp = send_initialize(&mut router).await;
8427 let instructions = resp.instructions.unwrap();
8428
8429 assert!(instructions.contains("- **mystery**: No description"));
8430 }
8431
8432 #[tokio::test]
8433 async fn test_auto_instructions_sorted_alphabetically() {
8434 let tool_z = ToolBuilder::new("zebra")
8435 .description("Z tool")
8436 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8437 .build();
8438 let tool_a = ToolBuilder::new("alpha")
8439 .description("A tool")
8440 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8441 .build();
8442 let tool_m = ToolBuilder::new("middle")
8443 .description("M tool")
8444 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8445 .build();
8446
8447 let mut router = McpRouter::new()
8448 .auto_instructions()
8449 .tool(tool_z)
8450 .tool(tool_a)
8451 .tool(tool_m);
8452
8453 let resp = send_initialize(&mut router).await;
8454 let instructions = resp.instructions.unwrap();
8455
8456 let alpha_pos = instructions.find("**alpha**").unwrap();
8457 let middle_pos = instructions.find("**middle**").unwrap();
8458 let zebra_pos = instructions.find("**zebra**").unwrap();
8459 assert!(alpha_pos < middle_pos);
8460 assert!(middle_pos < zebra_pos);
8461 }
8462
8463 #[tokio::test]
8464 async fn test_auto_instructions_read_only_and_idempotent_tags() {
8465 let tool = ToolBuilder::new("safe_update")
8466 .description("Safe update operation")
8467 .idempotent()
8468 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8469 .build();
8470
8471 let mut router = McpRouter::new().auto_instructions().tool(tool);
8472
8473 let resp = send_initialize(&mut router).await;
8474 let instructions = resp.instructions.unwrap();
8475
8476 assert!(
8477 instructions.contains("[idempotent]"),
8478 "got: {}",
8479 instructions
8480 );
8481 }
8482
8483 #[tokio::test]
8484 async fn test_auto_instructions_lazy_generation() {
8485 let mut router = McpRouter::new().auto_instructions();
8488
8489 let tool = ToolBuilder::new("late_tool")
8490 .description("Added after auto_instructions")
8491 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8492 .build();
8493
8494 router = router.tool(tool);
8495
8496 let resp = send_initialize(&mut router).await;
8497 let instructions = resp.instructions.unwrap();
8498
8499 assert!(instructions.contains("- **late_tool**: Added after auto_instructions"));
8500 }
8501
8502 #[tokio::test]
8503 async fn test_auto_instructions_multiple_annotation_tags() {
8504 let tool = ToolBuilder::new("update")
8505 .description("Update a record")
8506 .annotations(ToolAnnotations {
8507 read_only_hint: true,
8508 idempotent_hint: true,
8509 ..Default::default()
8510 })
8511 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8512 .build();
8513
8514 let mut router = McpRouter::new().auto_instructions().tool(tool);
8515
8516 let resp = send_initialize(&mut router).await;
8517 let instructions = resp.instructions.unwrap();
8518
8519 assert!(
8520 instructions.contains("[read-only, idempotent]"),
8521 "got: {}",
8522 instructions
8523 );
8524 }
8525
8526 #[tokio::test]
8527 async fn test_auto_instructions_no_annotations_no_tags() {
8528 let tool = ToolBuilder::new("fetch")
8530 .description("Fetch data")
8531 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8532 .build();
8533
8534 let mut router = McpRouter::new().auto_instructions().tool(tool);
8535
8536 let resp = send_initialize(&mut router).await;
8537 let instructions = resp.instructions.unwrap();
8538
8539 assert!(
8541 !instructions.contains('['),
8542 "should have no tags, got: {}",
8543 instructions
8544 );
8545 assert!(instructions.contains("- **fetch**: Fetch data"));
8546 }
8547
8548 async fn send_initialize(router: &mut McpRouter) -> InitializeResult {
8550 let init_req = RouterRequest {
8551 id: RequestId::Number(0),
8552 inner: McpRequest::Initialize(InitializeParams {
8553 protocol_version: "2025-11-25".to_string(),
8554 capabilities: ClientCapabilities {
8555 roots: None,
8556 sampling: None,
8557 elicitation: None,
8558 tasks: None,
8559 experimental: None,
8560 extensions: None,
8561 },
8562 client_info: Implementation {
8563 name: "test".to_string(),
8564 version: "1.0".to_string(),
8565 ..Default::default()
8566 },
8567 meta: None,
8568 }),
8569 extensions: Extensions::new(),
8570 };
8571 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8572 match resp.inner {
8573 Ok(McpResponse::Initialize(result)) => result,
8574 other => panic!("Expected Initialize response, got {:?}", other),
8575 }
8576 }
8577
8578 #[tokio::test]
8579 async fn test_notify_tools_list_changed() {
8580 let (tx, mut rx) = crate::context::notification_channel(16);
8581
8582 let router = McpRouter::new()
8583 .server_info("test", "1.0")
8584 .with_notification_sender(tx);
8585
8586 assert!(router.notify_tools_list_changed());
8587
8588 let notification = rx.recv().await.unwrap();
8589 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8590 }
8591
8592 #[tokio::test]
8593 async fn test_notify_prompts_list_changed() {
8594 let (tx, mut rx) = crate::context::notification_channel(16);
8595
8596 let router = McpRouter::new()
8597 .server_info("test", "1.0")
8598 .with_notification_sender(tx);
8599
8600 assert!(router.notify_prompts_list_changed());
8601
8602 let notification = rx.recv().await.unwrap();
8603 assert!(matches!(
8604 notification,
8605 ServerNotification::PromptsListChanged
8606 ));
8607 }
8608
8609 #[tokio::test]
8610 async fn test_notify_without_sender_returns_false() {
8611 let router = McpRouter::new().server_info("test", "1.0");
8612
8613 assert!(!router.notify_tools_list_changed());
8614 assert!(!router.notify_prompts_list_changed());
8615 assert!(!router.notify_resources_list_changed());
8616 }
8617
8618 #[tokio::test]
8619 async fn test_list_changed_capabilities_with_notification_sender() {
8620 let (tx, _rx) = crate::context::notification_channel(16);
8621 let tool = ToolBuilder::new("test")
8622 .description("test")
8623 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8624 .build();
8625
8626 let mut router = McpRouter::new()
8627 .server_info("test", "1.0")
8628 .tool(tool)
8629 .with_notification_sender(tx);
8630
8631 init_router(&mut router).await;
8632
8633 let caps = router.capabilities();
8634 let tools_cap = caps.tools.expect("tools capability should be present");
8635 assert!(
8636 tools_cap.list_changed,
8637 "tools.listChanged should be true when notification sender is configured"
8638 );
8639 }
8640
8641 #[tokio::test]
8642 async fn test_list_changed_capabilities_without_notification_sender() {
8643 let tool = ToolBuilder::new("test")
8644 .description("test")
8645 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8646 .build();
8647
8648 let mut router = McpRouter::new().server_info("test", "1.0").tool(tool);
8649
8650 init_router(&mut router).await;
8651
8652 let caps = router.capabilities();
8653 let tools_cap = caps.tools.expect("tools capability should be present");
8654 assert!(
8655 !tools_cap.list_changed,
8656 "tools.listChanged should be false without notification sender"
8657 );
8658 }
8659
8660 #[tokio::test]
8661 async fn test_set_logging_level_filters_messages() {
8662 let (tx, mut rx) = crate::context::notification_channel(16);
8663
8664 let mut router = McpRouter::new()
8665 .server_info("test", "1.0")
8666 .with_notification_sender(tx);
8667
8668 init_router(&mut router).await;
8669
8670 let set_level_req = RouterRequest {
8672 id: RequestId::Number(99),
8673 inner: McpRequest::SetLoggingLevel(SetLogLevelParams {
8674 level: LogLevel::Warning,
8675 meta: None,
8676 }),
8677 extensions: crate::context::Extensions::new(),
8678 };
8679 let resp = router
8680 .ready()
8681 .await
8682 .unwrap()
8683 .call(set_level_req)
8684 .await
8685 .unwrap();
8686 assert!(matches!(resp.inner, Ok(McpResponse::SetLoggingLevel(_))));
8687
8688 let ctx = router.create_context(RequestId::Number(100), None);
8690
8691 ctx.send_log(LoggingMessageParams::new(
8693 LogLevel::Error,
8694 serde_json::Value::Null,
8695 ));
8696 assert!(
8697 rx.try_recv().is_ok(),
8698 "Error should pass through Warning filter"
8699 );
8700
8701 ctx.send_log(LoggingMessageParams::new(
8703 LogLevel::Info,
8704 serde_json::Value::Null,
8705 ));
8706 assert!(
8707 rx.try_recv().is_err(),
8708 "Info should be filtered at Warning level"
8709 );
8710 }
8711
8712 #[test]
8713 fn test_paginate_no_page_size() {
8714 let items = vec![1, 2, 3, 4, 5];
8715 let (page, cursor) = paginate(items.clone(), None, None).unwrap();
8716 assert_eq!(page, items);
8717 assert!(cursor.is_none());
8718 }
8719
8720 #[test]
8721 fn test_paginate_first_page() {
8722 let items = vec![1, 2, 3, 4, 5];
8723 let (page, cursor) = paginate(items, None, Some(2)).unwrap();
8724 assert_eq!(page, vec![1, 2]);
8725 assert!(cursor.is_some());
8726 }
8727
8728 #[test]
8729 fn test_paginate_middle_page() {
8730 let items = vec![1, 2, 3, 4, 5];
8731 let (page1, cursor1) = paginate(items.clone(), None, Some(2)).unwrap();
8732 assert_eq!(page1, vec![1, 2]);
8733
8734 let (page2, cursor2) = paginate(items, cursor1.as_deref(), Some(2)).unwrap();
8735 assert_eq!(page2, vec![3, 4]);
8736 assert!(cursor2.is_some());
8737 }
8738
8739 #[test]
8740 fn test_paginate_last_page() {
8741 let items = vec![1, 2, 3, 4, 5];
8742 let cursor = encode_cursor(4);
8744 let (page, next) = paginate(items, Some(&cursor), Some(2)).unwrap();
8745 assert_eq!(page, vec![5]);
8746 assert!(next.is_none());
8747 }
8748
8749 #[test]
8750 fn test_paginate_exact_boundary() {
8751 let items = vec![1, 2, 3, 4];
8752 let (page, cursor) = paginate(items, None, Some(4)).unwrap();
8753 assert_eq!(page, vec![1, 2, 3, 4]);
8754 assert!(cursor.is_none());
8755 }
8756
8757 #[test]
8758 fn test_paginate_invalid_cursor() {
8759 let items = vec![1, 2, 3];
8760 let result = paginate(items, Some("not-valid-base64!@#$"), Some(2));
8761 assert!(result.is_err());
8762 }
8763
8764 #[test]
8765 fn test_cursor_round_trip() {
8766 let offset = 42;
8767 let encoded = encode_cursor(offset);
8768 let decoded = decode_cursor(&encoded).unwrap();
8769 assert_eq!(decoded, offset);
8770 }
8771
8772 #[tokio::test]
8773 async fn test_list_tools_pagination() {
8774 let tool_a = ToolBuilder::new("alpha")
8775 .description("a")
8776 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8777 .build();
8778 let tool_b = ToolBuilder::new("beta")
8779 .description("b")
8780 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8781 .build();
8782 let tool_c = ToolBuilder::new("gamma")
8783 .description("c")
8784 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8785 .build();
8786
8787 let mut router = McpRouter::new()
8788 .server_info("test", "1.0")
8789 .page_size(2)
8790 .tool(tool_a)
8791 .tool(tool_b)
8792 .tool(tool_c);
8793
8794 init_router(&mut router).await;
8795
8796 let req = RouterRequest {
8798 id: RequestId::Number(1),
8799 inner: McpRequest::ListTools(ListToolsParams {
8800 cursor: None,
8801 meta: None,
8802 }),
8803 extensions: Extensions::new(),
8804 };
8805 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8806 let (tools, next_cursor) = match resp.inner {
8807 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8808 other => panic!("Expected ListTools, got {:?}", other),
8809 };
8810 assert_eq!(tools.len(), 2);
8811 assert_eq!(tools[0].name, "alpha");
8812 assert_eq!(tools[1].name, "beta");
8813 assert!(next_cursor.is_some());
8814
8815 let req = RouterRequest {
8817 id: RequestId::Number(2),
8818 inner: McpRequest::ListTools(ListToolsParams {
8819 cursor: next_cursor,
8820 meta: None,
8821 }),
8822 extensions: Extensions::new(),
8823 };
8824 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8825 let (tools, next_cursor) = match resp.inner {
8826 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8827 other => panic!("Expected ListTools, got {:?}", other),
8828 };
8829 assert_eq!(tools.len(), 1);
8830 assert_eq!(tools[0].name, "gamma");
8831 assert!(next_cursor.is_none());
8832 }
8833
8834 #[tokio::test]
8835 async fn test_list_tools_no_pagination_by_default() {
8836 let tool_a = ToolBuilder::new("alpha")
8837 .description("a")
8838 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8839 .build();
8840 let tool_b = ToolBuilder::new("beta")
8841 .description("b")
8842 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8843 .build();
8844
8845 let mut router = McpRouter::new()
8846 .server_info("test", "1.0")
8847 .tool(tool_a)
8848 .tool(tool_b);
8849
8850 init_router(&mut router).await;
8851
8852 let req = RouterRequest {
8853 id: RequestId::Number(1),
8854 inner: McpRequest::ListTools(ListToolsParams {
8855 cursor: None,
8856 meta: None,
8857 }),
8858 extensions: Extensions::new(),
8859 };
8860 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8861 match resp.inner {
8862 Ok(McpResponse::ListTools(result)) => {
8863 assert_eq!(result.tools.len(), 2);
8864 assert!(result.next_cursor.is_none());
8865 }
8866 other => panic!("Expected ListTools, got {:?}", other),
8867 }
8868 }
8869
8870 #[cfg(feature = "dynamic-tools")]
8875 mod dynamic_tools_tests {
8876 use super::*;
8877
8878 #[tokio::test]
8879 async fn test_dynamic_tools_register_and_list() {
8880 let (router, registry) = McpRouter::new()
8881 .server_info("test", "1.0")
8882 .with_dynamic_tools();
8883
8884 let tool = ToolBuilder::new("dynamic_echo")
8885 .description("Dynamic echo")
8886 .handler(|input: AddInput| async move {
8887 Ok(CallToolResult::text(format!("{}", input.a)))
8888 })
8889 .build();
8890
8891 registry.register(tool);
8892
8893 let mut router = router;
8894 init_router(&mut router).await;
8895
8896 let req = RouterRequest {
8897 id: RequestId::Number(1),
8898 inner: McpRequest::ListTools(ListToolsParams::default()),
8899 extensions: Extensions::new(),
8900 };
8901
8902 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8903 match resp.inner {
8904 Ok(McpResponse::ListTools(result)) => {
8905 assert_eq!(result.tools.len(), 1);
8906 assert_eq!(result.tools[0].name, "dynamic_echo");
8907 }
8908 _ => panic!("Expected ListTools response"),
8909 }
8910 }
8911
8912 #[tokio::test]
8913 async fn test_dynamic_tools_unregister() {
8914 let (router, registry) = McpRouter::new()
8915 .server_info("test", "1.0")
8916 .with_dynamic_tools();
8917
8918 let tool = ToolBuilder::new("temp")
8919 .description("Temporary")
8920 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8921 .build();
8922
8923 registry.register(tool);
8924 assert!(registry.contains("temp"));
8925
8926 let removed = registry.unregister("temp");
8927 assert!(removed);
8928 assert!(!registry.contains("temp"));
8929
8930 assert!(!registry.unregister("temp"));
8932
8933 let mut router = router;
8934 init_router(&mut router).await;
8935
8936 let req = RouterRequest {
8937 id: RequestId::Number(1),
8938 inner: McpRequest::ListTools(ListToolsParams::default()),
8939 extensions: Extensions::new(),
8940 };
8941
8942 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8943 match resp.inner {
8944 Ok(McpResponse::ListTools(result)) => {
8945 assert_eq!(result.tools.len(), 0);
8946 }
8947 _ => panic!("Expected ListTools response"),
8948 }
8949 }
8950
8951 #[tokio::test]
8952 async fn test_dynamic_tools_merged_with_static() {
8953 let static_tool = ToolBuilder::new("static_tool")
8954 .description("Static")
8955 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8956 .build();
8957
8958 let (router, registry) = McpRouter::new()
8959 .server_info("test", "1.0")
8960 .tool(static_tool)
8961 .with_dynamic_tools();
8962
8963 let dynamic_tool = ToolBuilder::new("dynamic_tool")
8964 .description("Dynamic")
8965 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8966 .build();
8967
8968 registry.register(dynamic_tool);
8969
8970 let mut router = router;
8971 init_router(&mut router).await;
8972
8973 let req = RouterRequest {
8974 id: RequestId::Number(1),
8975 inner: McpRequest::ListTools(ListToolsParams::default()),
8976 extensions: Extensions::new(),
8977 };
8978
8979 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8980 match resp.inner {
8981 Ok(McpResponse::ListTools(result)) => {
8982 assert_eq!(result.tools.len(), 2);
8983 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8984 assert!(names.contains(&"static_tool"));
8985 assert!(names.contains(&"dynamic_tool"));
8986 }
8987 _ => panic!("Expected ListTools response"),
8988 }
8989 }
8990
8991 #[tokio::test]
8992 async fn test_static_tools_shadow_dynamic() {
8993 let static_tool = ToolBuilder::new("shared")
8994 .description("Static version")
8995 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8996 .build();
8997
8998 let (router, registry) = McpRouter::new()
8999 .server_info("test", "1.0")
9000 .tool(static_tool)
9001 .with_dynamic_tools();
9002
9003 let dynamic_tool = ToolBuilder::new("shared")
9004 .description("Dynamic version")
9005 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
9006 .build();
9007
9008 registry.register(dynamic_tool);
9009
9010 let mut router = router;
9011 init_router(&mut router).await;
9012
9013 let req = RouterRequest {
9015 id: RequestId::Number(1),
9016 inner: McpRequest::ListTools(ListToolsParams::default()),
9017 extensions: Extensions::new(),
9018 };
9019
9020 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9021 match resp.inner {
9022 Ok(McpResponse::ListTools(result)) => {
9023 assert_eq!(result.tools.len(), 1);
9024 assert_eq!(result.tools[0].name, "shared");
9025 assert_eq!(
9026 result.tools[0].description.as_deref(),
9027 Some("Static version")
9028 );
9029 }
9030 _ => panic!("Expected ListTools response"),
9031 }
9032
9033 let req = RouterRequest {
9035 id: RequestId::Number(2),
9036 inner: McpRequest::CallTool(CallToolParams {
9037 input_responses: None,
9038 request_state: None,
9039 name: "shared".to_string(),
9040 arguments: serde_json::json!({"a": 1, "b": 2}),
9041 meta: None,
9042 task: None,
9043 }),
9044 extensions: Extensions::new(),
9045 };
9046
9047 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9048 match resp.inner {
9049 Ok(McpResponse::CallTool(result)) => {
9050 assert!(!result.is_error);
9051 match &result.content[0] {
9052 Content::Text { text, .. } => assert_eq!(text, "static"),
9053 _ => panic!("Expected text content"),
9054 }
9055 }
9056 _ => panic!("Expected CallTool response"),
9057 }
9058 }
9059
9060 #[tokio::test]
9061 async fn test_dynamic_tools_call() {
9062 let (router, registry) = McpRouter::new()
9063 .server_info("test", "1.0")
9064 .with_dynamic_tools();
9065
9066 let tool = ToolBuilder::new("add")
9067 .description("Add two numbers")
9068 .handler(|input: AddInput| async move {
9069 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
9070 })
9071 .build();
9072
9073 registry.register(tool);
9074
9075 let mut router = router;
9076 init_router(&mut router).await;
9077
9078 let req = RouterRequest {
9079 id: RequestId::Number(1),
9080 inner: McpRequest::CallTool(CallToolParams {
9081 input_responses: None,
9082 request_state: None,
9083 name: "add".to_string(),
9084 arguments: serde_json::json!({"a": 3, "b": 4}),
9085 meta: None,
9086 task: None,
9087 }),
9088 extensions: Extensions::new(),
9089 };
9090
9091 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9092 match resp.inner {
9093 Ok(McpResponse::CallTool(result)) => {
9094 assert!(!result.is_error);
9095 match &result.content[0] {
9096 Content::Text { text, .. } => assert_eq!(text, "7"),
9097 _ => panic!("Expected text content"),
9098 }
9099 }
9100 _ => panic!("Expected CallTool response"),
9101 }
9102 }
9103
9104 #[tokio::test]
9105 async fn test_dynamic_tools_notification_on_register() {
9106 let (tx, mut rx) = crate::context::notification_channel(16);
9107 let (router, registry) = McpRouter::new()
9108 .server_info("test", "1.0")
9109 .with_dynamic_tools();
9110 let _router = router.with_notification_sender(tx);
9111
9112 let tool = ToolBuilder::new("notified")
9113 .description("Test")
9114 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9115 .build();
9116
9117 registry.register(tool);
9118
9119 let notification = rx.recv().await.unwrap();
9120 assert!(matches!(notification, ServerNotification::ToolsListChanged));
9121 }
9122
9123 #[tokio::test]
9124 async fn test_dynamic_tools_notification_on_unregister() {
9125 let (tx, mut rx) = crate::context::notification_channel(16);
9126 let (router, registry) = McpRouter::new()
9127 .server_info("test", "1.0")
9128 .with_dynamic_tools();
9129 let _router = router.with_notification_sender(tx);
9130
9131 let tool = ToolBuilder::new("notified")
9132 .description("Test")
9133 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9134 .build();
9135
9136 registry.register(tool);
9137 let _ = rx.recv().await.unwrap();
9139
9140 registry.unregister("notified");
9141 let notification = rx.recv().await.unwrap();
9142 assert!(matches!(notification, ServerNotification::ToolsListChanged));
9143 }
9144
9145 #[tokio::test]
9146 async fn test_dynamic_tools_no_notification_on_empty_unregister() {
9147 let (tx, mut rx) = crate::context::notification_channel(16);
9148 let (router, registry) = McpRouter::new()
9149 .server_info("test", "1.0")
9150 .with_dynamic_tools();
9151 let _router = router.with_notification_sender(tx);
9152
9153 assert!(!registry.unregister("nonexistent"));
9155
9156 assert!(rx.try_recv().is_err());
9158 }
9159
9160 #[tokio::test]
9161 async fn test_dynamic_tools_filter_applies() {
9162 use crate::filter::CapabilityFilter;
9163
9164 let (router, registry) = McpRouter::new()
9165 .server_info("test", "1.0")
9166 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| {
9167 tool.name != "hidden"
9168 }))
9169 .with_dynamic_tools();
9170
9171 let visible = ToolBuilder::new("visible")
9172 .description("Visible")
9173 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9174 .build();
9175
9176 let hidden = ToolBuilder::new("hidden")
9177 .description("Hidden")
9178 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9179 .build();
9180
9181 registry.register(visible);
9182 registry.register(hidden);
9183
9184 let mut router = router;
9185 init_router(&mut router).await;
9186
9187 let req = RouterRequest {
9189 id: RequestId::Number(1),
9190 inner: McpRequest::ListTools(ListToolsParams::default()),
9191 extensions: Extensions::new(),
9192 };
9193
9194 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9195 match resp.inner {
9196 Ok(McpResponse::ListTools(result)) => {
9197 assert_eq!(result.tools.len(), 1);
9198 assert_eq!(result.tools[0].name, "visible");
9199 }
9200 _ => panic!("Expected ListTools response"),
9201 }
9202
9203 let req = RouterRequest {
9205 id: RequestId::Number(2),
9206 inner: McpRequest::CallTool(CallToolParams {
9207 input_responses: None,
9208 request_state: None,
9209 name: "hidden".to_string(),
9210 arguments: serde_json::json!({"a": 1, "b": 2}),
9211 meta: None,
9212 task: None,
9213 }),
9214 extensions: Extensions::new(),
9215 };
9216
9217 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9218 match resp.inner {
9219 Err(e) => {
9220 assert_eq!(e.code, -32601); }
9222 _ => panic!("Expected JsonRpc error"),
9223 }
9224 }
9225
9226 #[tokio::test]
9227 async fn test_dynamic_tools_capabilities_advertised() {
9228 let (mut router, _registry) = McpRouter::new()
9230 .server_info("test", "1.0")
9231 .with_dynamic_tools();
9232
9233 let init_req = RouterRequest {
9234 id: RequestId::Number(1),
9235 inner: McpRequest::Initialize(InitializeParams {
9236 protocol_version: "2025-11-25".to_string(),
9237 capabilities: ClientCapabilities::default(),
9238 client_info: Implementation {
9239 name: "test".to_string(),
9240 version: "1.0".to_string(),
9241 ..Default::default()
9242 },
9243 meta: None,
9244 }),
9245 extensions: Extensions::new(),
9246 };
9247
9248 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
9249 match resp.inner {
9250 Ok(McpResponse::Initialize(result)) => {
9251 assert!(result.capabilities.tools.is_some());
9252 }
9253 _ => panic!("Expected Initialize response"),
9254 }
9255 }
9256
9257 #[tokio::test]
9258 async fn test_dynamic_tools_multi_session_notification() {
9259 let (tx1, mut rx1) = crate::context::notification_channel(16);
9260 let (tx2, mut rx2) = crate::context::notification_channel(16);
9261
9262 let (router, registry) = McpRouter::new()
9263 .server_info("test", "1.0")
9264 .with_dynamic_tools();
9265
9266 let _session1 = router.clone().with_notification_sender(tx1);
9268 let _session2 = router.clone().with_notification_sender(tx2);
9269
9270 let tool = ToolBuilder::new("broadcast")
9271 .description("Test")
9272 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9273 .build();
9274
9275 registry.register(tool);
9276
9277 let n1 = rx1.recv().await.unwrap();
9279 let n2 = rx2.recv().await.unwrap();
9280 assert!(matches!(n1, ServerNotification::ToolsListChanged));
9281 assert!(matches!(n2, ServerNotification::ToolsListChanged));
9282 }
9283
9284 #[tokio::test]
9285 async fn test_dynamic_tools_call_not_found() {
9286 let (router, _registry) = McpRouter::new()
9287 .server_info("test", "1.0")
9288 .with_dynamic_tools();
9289
9290 let mut router = router;
9291 init_router(&mut router).await;
9292
9293 let req = RouterRequest {
9294 id: RequestId::Number(1),
9295 inner: McpRequest::CallTool(CallToolParams {
9296 input_responses: None,
9297 request_state: None,
9298 name: "nonexistent".to_string(),
9299 arguments: serde_json::json!({}),
9300 meta: None,
9301 task: None,
9302 }),
9303 extensions: Extensions::new(),
9304 };
9305
9306 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9307 match resp.inner {
9308 Err(e) => {
9309 assert_eq!(e.code, -32601);
9310 }
9311 _ => panic!("Expected method not found error"),
9312 }
9313 }
9314
9315 #[tokio::test]
9316 async fn test_dynamic_tools_registry_list() {
9317 let (_, registry) = McpRouter::new()
9318 .server_info("test", "1.0")
9319 .with_dynamic_tools();
9320
9321 assert!(registry.list().is_empty());
9322
9323 let tool = ToolBuilder::new("tool_a")
9324 .description("A")
9325 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9326 .build();
9327 registry.register(tool);
9328
9329 let tool = ToolBuilder::new("tool_b")
9330 .description("B")
9331 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9332 .build();
9333 registry.register(tool);
9334
9335 let tools = registry.list();
9336 assert_eq!(tools.len(), 2);
9337 let names: Vec<&str> = tools.iter().map(|t| t.name.as_str()).collect();
9338 assert!(names.contains(&"tool_a"));
9339 assert!(names.contains(&"tool_b"));
9340 }
9341 } #[tokio::test]
9344 async fn test_tool_if_true_registers() {
9345 let tool = ToolBuilder::new("conditional")
9346 .description("Conditional tool")
9347 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9348 .build();
9349
9350 let mut router = McpRouter::new().tool_if(true, tool);
9351 init_router(&mut router).await;
9352
9353 let req = RouterRequest {
9354 id: RequestId::Number(1),
9355 inner: McpRequest::ListTools(ListToolsParams::default()),
9356 extensions: Extensions::new(),
9357 };
9358 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9359 match resp.inner {
9360 Ok(McpResponse::ListTools(result)) => {
9361 assert_eq!(result.tools.len(), 1);
9362 assert_eq!(result.tools[0].name, "conditional");
9363 }
9364 _ => panic!("Expected ListTools response"),
9365 }
9366 }
9367
9368 #[tokio::test]
9369 async fn test_tool_if_false_skips() {
9370 let tool = ToolBuilder::new("conditional")
9371 .description("Conditional tool")
9372 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9373 .build();
9374
9375 let mut router = McpRouter::new().tool_if(false, tool);
9376 init_router(&mut router).await;
9377
9378 let req = RouterRequest {
9379 id: RequestId::Number(1),
9380 inner: McpRequest::ListTools(ListToolsParams::default()),
9381 extensions: Extensions::new(),
9382 };
9383 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9384 match resp.inner {
9385 Ok(McpResponse::ListTools(result)) => {
9386 assert_eq!(result.tools.len(), 0);
9387 }
9388 _ => panic!("Expected ListTools response"),
9389 }
9390 }
9391
9392 #[tokio::test]
9393 async fn test_tools_if_batch_conditional() {
9394 let tools = vec![
9395 ToolBuilder::new("a")
9396 .description("Tool A")
9397 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9398 .build(),
9399 ToolBuilder::new("b")
9400 .description("Tool B")
9401 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9402 .build(),
9403 ];
9404
9405 let mut router = McpRouter::new().tools_if(false, tools);
9406 init_router(&mut router).await;
9407
9408 let req = RouterRequest {
9409 id: RequestId::Number(1),
9410 inner: McpRequest::ListTools(ListToolsParams::default()),
9411 extensions: Extensions::new(),
9412 };
9413 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9414 match resp.inner {
9415 Ok(McpResponse::ListTools(result)) => {
9416 assert_eq!(result.tools.len(), 0);
9417 }
9418 _ => panic!("Expected ListTools response"),
9419 }
9420 }
9421
9422 #[test]
9423 fn test_resource_if_true_registers() {
9424 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
9425 .name("test")
9426 .text("hello");
9427
9428 let router = McpRouter::new().resource_if(true, resource);
9429 assert_eq!(router.inner.resources.len(), 1);
9430 }
9431
9432 #[test]
9433 fn test_resource_if_false_skips() {
9434 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
9435 .name("test")
9436 .text("hello");
9437
9438 let router = McpRouter::new().resource_if(false, resource);
9439 assert_eq!(router.inner.resources.len(), 0);
9440 }
9441
9442 #[test]
9443 fn test_prompt_if_true_registers() {
9444 let prompt = crate::prompt::PromptBuilder::new("greet")
9445 .description("Greeting")
9446 .user_message("Hello!");
9447
9448 let router = McpRouter::new().prompt_if(true, prompt);
9449 assert_eq!(router.inner.prompts.len(), 1);
9450 }
9451
9452 #[test]
9453 fn test_prompt_if_false_skips() {
9454 let prompt = crate::prompt::PromptBuilder::new("greet")
9455 .description("Greeting")
9456 .user_message("Hello!");
9457
9458 let router = McpRouter::new().prompt_if(false, prompt);
9459 assert_eq!(router.inner.prompts.len(), 0);
9460 }
9461
9462 #[tokio::test]
9463 async fn test_disable_tool_hides_from_list() {
9464 let safe = ToolBuilder::new("safe")
9465 .description("Safe tool")
9466 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9467 .build();
9468 let dangerous = ToolBuilder::new("dangerous")
9469 .description("Dangerous tool")
9470 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9471 .build();
9472 let mut router = McpRouter::new().tool(safe).tool(dangerous);
9473 init_router(&mut router).await;
9474
9475 router.disable_tool("dangerous");
9476 assert!(router.is_tool_enabled("safe"));
9477 assert!(!router.is_tool_enabled("dangerous"));
9478
9479 let req = RouterRequest {
9480 id: RequestId::Number(1),
9481 inner: McpRequest::ListTools(ListToolsParams::default()),
9482 extensions: Extensions::new(),
9483 };
9484 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9485 match resp.inner {
9486 Ok(McpResponse::ListTools(result)) => {
9487 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
9488 assert_eq!(names, vec!["safe"]);
9489 }
9490 _ => panic!("Expected ListTools response"),
9491 }
9492 }
9493
9494 #[tokio::test]
9495 async fn test_disable_tool_blocks_call() {
9496 let dangerous = ToolBuilder::new("dangerous")
9497 .description("Dangerous tool")
9498 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
9499 .build();
9500 let mut router = McpRouter::new().tool(dangerous);
9501 init_router(&mut router).await;
9502
9503 router.disable_tool("dangerous");
9504
9505 let req = RouterRequest {
9506 id: RequestId::Number(2),
9507 inner: McpRequest::CallTool(CallToolParams {
9508 input_responses: None,
9509 request_state: None,
9510 name: "dangerous".to_string(),
9511 arguments: serde_json::json!({"a": 1, "b": 2}),
9512 meta: None,
9513 task: None,
9514 }),
9515 extensions: Extensions::new(),
9516 };
9517 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9518 let err = resp.inner.expect_err("disabled tool should error");
9519 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
9520 }
9521
9522 #[tokio::test]
9523 async fn test_enable_tool_restores_visibility() {
9524 let tool = ToolBuilder::new("flippy")
9525 .description("Toggleable tool")
9526 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
9527 .build();
9528 let mut router = McpRouter::new().tool(tool);
9529 init_router(&mut router).await;
9530
9531 router.disable_tool("flippy");
9532 router.enable_tool("flippy");
9533 assert!(router.is_tool_enabled("flippy"));
9534
9535 let req = RouterRequest {
9536 id: RequestId::Number(3),
9537 inner: McpRequest::CallTool(CallToolParams {
9538 input_responses: None,
9539 request_state: None,
9540 name: "flippy".to_string(),
9541 arguments: serde_json::json!({"a": 1, "b": 2}),
9542 meta: None,
9543 task: None,
9544 }),
9545 extensions: Extensions::new(),
9546 };
9547 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9548 match resp.inner {
9549 Ok(McpResponse::CallTool(result)) => {
9550 assert_eq!(result.first_text(), Some("ran"));
9551 }
9552 _ => panic!("Expected CallTool response"),
9553 }
9554 }
9555
9556 #[tokio::test]
9557 async fn test_disable_propagates_through_fresh_session() {
9558 let tool = ToolBuilder::new("shared")
9559 .description("Shared across sessions")
9560 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9561 .build();
9562 let router = McpRouter::new().tool(tool);
9563
9564 router.disable_tool("shared");
9566 let mut child = router.with_fresh_session();
9567 init_router(&mut child).await;
9568 assert!(!child.is_tool_enabled("shared"));
9569
9570 let req = RouterRequest {
9571 id: RequestId::Number(4),
9572 inner: McpRequest::ListTools(ListToolsParams::default()),
9573 extensions: Extensions::new(),
9574 };
9575 let resp = child.ready().await.unwrap().call(req).await.unwrap();
9576 match resp.inner {
9577 Ok(McpResponse::ListTools(result)) => {
9578 assert!(result.tools.is_empty());
9579 }
9580 _ => panic!("Expected ListTools response"),
9581 }
9582 }
9583
9584 #[tokio::test]
9585 async fn test_disable_resource_and_prompt() {
9586 let resource = crate::resource::ResourceBuilder::new("file:///hidden.txt")
9587 .name("hidden")
9588 .text("secret");
9589 let prompt = crate::prompt::PromptBuilder::new("hidden_prompt")
9590 .description("hidden")
9591 .user_message("hello");
9592
9593 let mut router = McpRouter::new().resource(resource).prompt(prompt);
9594 init_router(&mut router).await;
9595
9596 router.disable_resource("file:///hidden.txt");
9597 router.disable_prompt("hidden_prompt");
9598 assert!(!router.is_resource_enabled("file:///hidden.txt"));
9599 assert!(!router.is_prompt_enabled("hidden_prompt"));
9600
9601 let req = RouterRequest {
9603 id: RequestId::Number(5),
9604 inner: McpRequest::ListResources(ListResourcesParams::default()),
9605 extensions: Extensions::new(),
9606 };
9607 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9608 match resp.inner {
9609 Ok(McpResponse::ListResources(result)) => {
9610 assert!(result.resources.is_empty());
9611 }
9612 _ => panic!("Expected ListResources response"),
9613 }
9614
9615 let req = RouterRequest {
9617 id: RequestId::Number(6),
9618 inner: McpRequest::ReadResource(ReadResourceParams {
9619 input_responses: None,
9620 request_state: None,
9621 uri: "file:///hidden.txt".to_string(),
9622 meta: None,
9623 }),
9624 extensions: Extensions::new(),
9625 };
9626 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9627 let err = resp.inner.expect_err("disabled resource should error");
9628 assert_eq!(err.code, -32602); let req = RouterRequest {
9632 id: RequestId::Number(7),
9633 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
9634 extensions: Extensions::new(),
9635 };
9636 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9637 match resp.inner {
9638 Ok(McpResponse::ListPrompts(result)) => {
9639 assert!(result.prompts.is_empty());
9640 }
9641 _ => panic!("Expected ListPrompts response"),
9642 }
9643
9644 let req = RouterRequest {
9646 id: RequestId::Number(8),
9647 inner: McpRequest::GetPrompt(GetPromptParams {
9648 input_responses: None,
9649 request_state: None,
9650 name: "hidden_prompt".to_string(),
9651 arguments: Default::default(),
9652 meta: None,
9653 }),
9654 extensions: Extensions::new(),
9655 };
9656 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9657 let err = resp.inner.expect_err("disabled prompt should error");
9658 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
9659 }
9660
9661 #[test]
9662 fn test_router_request_new() {
9663 let req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9664 assert_eq!(req.id, RequestId::Number(1));
9665 assert!(req.extensions.is_empty());
9666 }
9667
9668 #[test]
9669 fn test_with_inner_preserves_extensions() {
9670 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9671 req.extensions.insert(42u32);
9672
9673 let rewritten = req.with_inner(McpRequest::ListTools(Default::default()));
9674 assert!(matches!(rewritten.inner, McpRequest::ListTools(_)));
9675 assert_eq!(rewritten.id, RequestId::Number(1));
9676 assert_eq!(rewritten.extensions.get::<u32>(), Some(&42));
9677 }
9678
9679 #[test]
9680 fn test_with_id_and_inner_preserves_extensions() {
9681 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9682 req.extensions.insert(String::from("token-abc"));
9683
9684 let rewritten = req.with_id_and_inner(
9685 RequestId::Number(99),
9686 McpRequest::ListResources(Default::default()),
9687 );
9688 assert_eq!(rewritten.id, RequestId::Number(99));
9689 assert!(matches!(rewritten.inner, McpRequest::ListResources(_)));
9690 assert_eq!(
9691 rewritten.extensions.get::<String>(),
9692 Some(&String::from("token-abc"))
9693 );
9694 }
9695
9696 #[test]
9697 fn test_clone_with_inner_preserves_extensions() {
9698 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9699 req.extensions.insert(true);
9700
9701 let cloned = req.clone_with_inner(McpRequest::ListTools(Default::default()));
9702
9703 assert!(matches!(req.inner, McpRequest::Ping));
9705 assert_eq!(req.extensions.get::<bool>(), Some(&true));
9706
9707 assert!(matches!(cloned.inner, McpRequest::ListTools(_)));
9709 assert_eq!(cloned.extensions.get::<bool>(), Some(&true));
9710 }
9711
9712 #[test]
9713 fn test_router_response_is_error() {
9714 let ok_resp = RouterResponse {
9715 id: RequestId::Number(1),
9716 inner: Ok(McpResponse::Pong(Default::default())),
9717 };
9718 assert!(!ok_resp.is_error());
9719
9720 let err_resp = RouterResponse {
9721 id: RequestId::Number(2),
9722 inner: Err(JsonRpcError::internal_error("boom")),
9723 };
9724 assert!(err_resp.is_error());
9725 }
9726
9727 #[test]
9728 fn test_extensions_len_and_is_empty() {
9729 let mut ext = Extensions::new();
9730 assert!(ext.is_empty());
9731 assert_eq!(ext.len(), 0);
9732
9733 ext.insert(42u32);
9734 assert!(!ext.is_empty());
9735 assert_eq!(ext.len(), 1);
9736
9737 ext.insert(String::from("hello"));
9738 assert_eq!(ext.len(), 2);
9739 }
9740
9741 #[test]
9742 fn test_router_response_serde_roundtrip() {
9743 let response = RouterResponse {
9745 id: RequestId::Number(1),
9746 inner: Ok(McpResponse::Empty(EmptyResult {})),
9747 };
9748 let json = serde_json::to_string(&response).unwrap();
9749 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9750 assert_eq!(deserialized.id, RequestId::Number(1));
9751 assert!(!deserialized.is_error());
9752
9753 let response = RouterResponse {
9755 id: RequestId::String("req-2".into()),
9756 inner: Err(JsonRpcError::method_not_found("unknown")),
9757 };
9758 let json = serde_json::to_string(&response).unwrap();
9759 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9760 assert_eq!(deserialized.id, RequestId::String("req-2".into()));
9761 assert!(deserialized.is_error());
9762 }
9763
9764 #[tokio::test]
9771 async fn test_discover_dispatch_via_jsonrpc_service() {
9772 let router = McpRouter::new().server_info("unit-test-server", "4.2.0");
9775 let mut service = JsonRpcService::new(router);
9776
9777 let req = JsonRpcRequest::new(1, "server/discover");
9778 let resp = service.call_single(req).await.unwrap();
9779
9780 match resp {
9781 JsonRpcResponse::Result(r) => {
9782 let versions = r
9784 .result
9785 .get("supportedVersions")
9786 .and_then(|v| v.as_array())
9787 .expect("result.supportedVersions must be an array");
9788 assert!(!versions.is_empty(), "supportedVersions must not be empty");
9789
9790 assert_eq!(
9792 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
9793 "unit-test-server",
9794 "serverInfo.name must match configured value"
9795 );
9796 assert_eq!(
9797 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["version"], "4.2.0",
9798 "serverInfo.version must match configured value"
9799 );
9800
9801 assert!(
9804 r.result.get("protocolVersion").is_none(),
9805 "server/discover must NOT include protocolVersion: {:?}",
9806 r.result
9807 );
9808 }
9809 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
9810 _ => panic!("unexpected response variant"),
9811 }
9812 }
9813
9814 #[tokio::test]
9815 async fn test_discover_does_not_require_initialization() {
9816 let router = McpRouter::new().server_info("fresh-router", "1.0.0");
9819 let mut service = JsonRpcService::new(router);
9820
9821 let req = JsonRpcRequest::new(2, "server/discover");
9822 let resp = service.call_single(req).await.unwrap();
9823
9824 assert!(
9826 !matches!(resp, JsonRpcResponse::Error(_)),
9827 "server/discover must not require initialization: {:?}",
9828 resp
9829 );
9830 }
9831}
9832
9833#[cfg(test)]
9834mod cursor_property_tests {
9835 use super::{decode_cursor, encode_cursor};
9836 use proptest::prelude::*;
9837
9838 fn arb_cursor_text() -> BoxedStrategy<String> {
9839 prop_oneof![
9840 8 => prop::collection::vec(any::<char>(), 0..512)
9841 .prop_map(|chars| chars.into_iter().collect()),
9842 1 => Just("\0\r\n\t\u{001b}\u{007f}".repeat(64)),
9843 1 => Just("A".repeat(16 * 1024)),
9844 ]
9845 .boxed()
9846 }
9847
9848 proptest! {
9849 #![proptest_config(ProptestConfig::with_cases(512))]
9850
9851 #[test]
9853 fn cursor_round_trips(offset in any::<usize>()) {
9854 prop_assert_eq!(decode_cursor(&encode_cursor(offset)).unwrap(), offset);
9855 }
9856
9857 #[test]
9859 fn decode_cursor_never_panics(s in arb_cursor_text()) {
9860 let _ = decode_cursor(&s);
9861 }
9862 }
9863}