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 ctx = if !is_final_protocol_request(per_request)
1093 && let Some(requester) = merged
1094 .get::<ClientRequesterHandle>()
1095 .cloned()
1096 .or_else(|| self.inner.client_requester.clone())
1097 {
1098 ctx.with_client_requester(requester)
1099 } else {
1100 ctx
1101 };
1102
1103 let ctx = if let Some(token) = merged.get::<CancellationToken>() {
1107 ctx.with_cancellation_token(token.clone())
1108 } else {
1109 ctx
1110 };
1111
1112 let ctx = ctx.with_extensions(Arc::new(merged));
1113
1114 let ctx = ctx.with_min_log_level(self.inner.min_log_level.clone());
1116
1117 let token = ctx.cancellation_token();
1119 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1120 in_flight.insert(request_id, token);
1121 }
1122
1123 ctx
1124 }
1125
1126 pub fn complete_request(&self, request_id: &RequestId) {
1128 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1129 in_flight.remove(request_id);
1130 }
1131 }
1132
1133 fn cancel_request(&self, request_id: &RequestId) -> bool {
1135 let Ok(in_flight) = self.inner.in_flight.read() else {
1136 return false;
1137 };
1138 let Some(token) = in_flight.get(request_id) else {
1139 return false;
1140 };
1141 token.cancel();
1142 true
1143 }
1144
1145 pub fn server_info(mut self, name: impl Into<String>, version: impl Into<String>) -> Self {
1147 let inner = Arc::make_mut(&mut self.inner);
1148 inner.server_name = name.into();
1149 inner.server_version = version.into();
1150 self
1151 }
1152
1153 pub fn page_size(mut self, size: usize) -> Self {
1160 Arc::make_mut(&mut self.inner).page_size = Some(size);
1161 self
1162 }
1163
1164 pub fn list_ttl(mut self, ms: u64) -> Self {
1170 Arc::make_mut(&mut self.inner).list_ttl_ms = Some(ms);
1171 self
1172 }
1173
1174 pub fn read_ttl(mut self, ms: u64) -> Self {
1181 Arc::make_mut(&mut self.inner).read_ttl_ms = Some(ms);
1182 self
1183 }
1184
1185 pub fn cache_scope(mut self, scope: CacheScope) -> Self {
1194 Arc::make_mut(&mut self.inner).cache_scope = Some(scope);
1195 self
1196 }
1197
1198 pub fn logging_deprecated(mut self, info: tower_mcp_types::protocol::DeprecationInfo) -> Self {
1204 Arc::make_mut(&mut self.inner).logging_deprecated = Some(info);
1205 self
1206 }
1207
1208 pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
1210 Arc::make_mut(&mut self.inner).instructions = Some(instructions.into());
1211 self
1212 }
1213
1214 pub fn auto_instructions(mut self) -> Self {
1246 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1247 prefix: None,
1248 suffix: None,
1249 });
1250 self
1251 }
1252
1253 pub fn auto_instructions_with(
1270 mut self,
1271 prefix: Option<impl Into<String>>,
1272 suffix: Option<impl Into<String>>,
1273 ) -> Self {
1274 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1275 prefix: prefix.map(Into::into),
1276 suffix: suffix.map(Into::into),
1277 });
1278 self
1279 }
1280
1281 pub fn server_title(mut self, title: impl Into<String>) -> Self {
1283 Arc::make_mut(&mut self.inner).server_title = Some(title.into());
1284 self
1285 }
1286
1287 pub fn server_description(mut self, description: impl Into<String>) -> Self {
1289 Arc::make_mut(&mut self.inner).server_description = Some(description.into());
1290 self
1291 }
1292
1293 pub fn server_icons(mut self, icons: Vec<ToolIcon>) -> Self {
1295 Arc::make_mut(&mut self.inner).server_icons = Some(icons);
1296 self
1297 }
1298
1299 pub fn server_website_url(mut self, url: impl Into<String>) -> Self {
1301 Arc::make_mut(&mut self.inner).server_website_url = Some(url.into());
1302 self
1303 }
1304
1305 pub fn tool(mut self, tool: Tool) -> Self {
1307 Arc::make_mut(&mut self.inner)
1308 .tools
1309 .insert(tool.name.clone(), Arc::new(tool));
1310 self
1311 }
1312
1313 pub fn tool_if(self, condition: bool, tool: Tool) -> Self {
1339 if condition { self.tool(tool) } else { self }
1340 }
1341
1342 pub fn resource(mut self, resource: Resource) -> Self {
1344 Arc::make_mut(&mut self.inner)
1345 .resources
1346 .insert(resource.uri.clone(), Arc::new(resource));
1347 self
1348 }
1349
1350 pub fn resource_if(self, condition: bool, resource: Resource) -> Self {
1369 if condition {
1370 self.resource(resource)
1371 } else {
1372 self
1373 }
1374 }
1375
1376 pub fn resource_template(mut self, template: ResourceTemplate) -> Self {
1410 Arc::make_mut(&mut self.inner)
1411 .resource_templates
1412 .push(Arc::new(template));
1413 self
1414 }
1415
1416 pub fn prompt(mut self, prompt: Prompt) -> Self {
1418 Arc::make_mut(&mut self.inner)
1419 .prompts
1420 .insert(prompt.name.clone(), Arc::new(prompt));
1421 self
1422 }
1423
1424 pub fn prompt_if(self, condition: bool, prompt: Prompt) -> Self {
1443 if condition { self.prompt(prompt) } else { self }
1444 }
1445
1446 pub fn tools(self, tools: impl IntoIterator<Item = Tool>) -> Self {
1472 tools
1473 .into_iter()
1474 .fold(self, |router, tool| router.tool(tool))
1475 }
1476
1477 pub fn tools_if(self, condition: bool, tools: impl IntoIterator<Item = Tool>) -> Self {
1481 if condition { self.tools(tools) } else { self }
1482 }
1483
1484 pub fn resources(self, resources: impl IntoIterator<Item = Resource>) -> Self {
1503 resources
1504 .into_iter()
1505 .fold(self, |router, resource| router.resource(resource))
1506 }
1507
1508 pub fn resources_if(
1512 self,
1513 condition: bool,
1514 resources: impl IntoIterator<Item = Resource>,
1515 ) -> Self {
1516 if condition {
1517 self.resources(resources)
1518 } else {
1519 self
1520 }
1521 }
1522
1523 pub fn prompts(self, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1542 prompts
1543 .into_iter()
1544 .fold(self, |router, prompt| router.prompt(prompt))
1545 }
1546
1547 pub fn prompts_if(self, condition: bool, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1551 if condition {
1552 self.prompts(prompts)
1553 } else {
1554 self
1555 }
1556 }
1557
1558 pub fn merge(mut self, other: McpRouter) -> Self {
1603 let inner = Arc::make_mut(&mut self.inner);
1604 let other_inner = other.inner;
1605
1606 for (name, tool) in &other_inner.tools {
1608 inner.tools.insert(name.clone(), tool.clone());
1609 }
1610
1611 for (uri, resource) in &other_inner.resources {
1613 inner.resources.insert(uri.clone(), resource.clone());
1614 }
1615
1616 for template in &other_inner.resource_templates {
1619 inner.resource_templates.push(template.clone());
1620 }
1621
1622 for (name, prompt) in &other_inner.prompts {
1624 inner.prompts.insert(name.clone(), prompt.clone());
1625 }
1626
1627 for (identifier, settings) in &other_inner.protocol_extensions {
1629 inner
1630 .protocol_extensions
1631 .insert(identifier.clone(), settings.clone());
1632 }
1633
1634 self
1635 }
1636
1637 pub fn nest(mut self, prefix: impl Into<String>, other: McpRouter) -> Self {
1677 let prefix = prefix.into();
1678 let inner = Arc::make_mut(&mut self.inner);
1679 let other_inner = other.inner;
1680
1681 for tool in other_inner.tools.values() {
1683 let prefixed_tool = tool.with_name_prefix(&prefix);
1684 inner
1685 .tools
1686 .insert(prefixed_tool.name.clone(), Arc::new(prefixed_tool));
1687 }
1688
1689 for (uri, resource) in &other_inner.resources {
1691 inner.resources.insert(uri.clone(), resource.clone());
1692 }
1693
1694 for template in &other_inner.resource_templates {
1696 inner.resource_templates.push(template.clone());
1697 }
1698
1699 for (name, prompt) in &other_inner.prompts {
1701 inner.prompts.insert(name.clone(), prompt.clone());
1702 }
1703
1704 for (identifier, settings) in &other_inner.protocol_extensions {
1707 inner
1708 .protocol_extensions
1709 .insert(identifier.clone(), settings.clone());
1710 }
1711
1712 self
1713 }
1714
1715 pub fn completion_handler<F, Fut>(mut self, handler: F) -> Self
1743 where
1744 F: Fn(CompleteParams) -> Fut + Send + Sync + 'static,
1745 Fut: Future<Output = Result<CompleteResult>> + Send + 'static,
1746 {
1747 Arc::make_mut(&mut self.inner).completion_handler =
1748 Some(Arc::new(move |params| Box::pin(handler(params))));
1749 self
1750 }
1751
1752 pub fn tool_filter(mut self, filter: ToolFilter) -> Self {
1787 Arc::make_mut(&mut self.inner).tool_filter = Some(filter);
1788 self
1789 }
1790
1791 pub fn resource_filter(mut self, filter: ResourceFilter) -> Self {
1822 Arc::make_mut(&mut self.inner).resource_filter = Some(filter);
1823 self
1824 }
1825
1826 pub fn prompt_filter(mut self, filter: PromptFilter) -> Self {
1855 Arc::make_mut(&mut self.inner).prompt_filter = Some(filter);
1856 self
1857 }
1858
1859 pub fn session(&self) -> &SessionState {
1861 &self.session
1862 }
1863
1864 pub fn log(&self, params: LoggingMessageParams) -> bool {
1886 let Some(tx) = &self.inner.notification_tx else {
1887 return false;
1888 };
1889 tx.try_send(ServerNotification::LogMessage(params)).is_ok()
1890 }
1891
1892 pub fn log_info(&self, message: &str) -> bool {
1896 self.log(LoggingMessageParams::new(
1897 LogLevel::Info,
1898 serde_json::json!({ "message": message }),
1899 ))
1900 }
1901
1902 pub fn log_warning(&self, message: &str) -> bool {
1904 self.log(LoggingMessageParams::new(
1905 LogLevel::Warning,
1906 serde_json::json!({ "message": message }),
1907 ))
1908 }
1909
1910 pub fn log_error(&self, message: &str) -> bool {
1912 self.log(LoggingMessageParams::new(
1913 LogLevel::Error,
1914 serde_json::json!({ "message": message }),
1915 ))
1916 }
1917
1918 pub fn log_debug(&self, message: &str) -> bool {
1920 self.log(LoggingMessageParams::new(
1921 LogLevel::Debug,
1922 serde_json::json!({ "message": message }),
1923 ))
1924 }
1925
1926 pub fn is_subscribed(&self, uri: &str) -> bool {
1928 if let Ok(subs) = self.inner.subscriptions.read() {
1929 return subs.contains(uri);
1930 }
1931 false
1932 }
1933
1934 pub fn subscribed_uris(&self) -> Vec<String> {
1936 if let Ok(subs) = self.inner.subscriptions.read() {
1937 return subs.iter().cloned().collect();
1938 }
1939 Vec::new()
1940 }
1941
1942 fn subscribe(&self, uri: &str) -> bool {
1944 if let Ok(mut subs) = self.inner.subscriptions.write() {
1945 return subs.insert(uri.to_string());
1946 }
1947 false
1948 }
1949
1950 fn unsubscribe(&self, uri: &str) -> bool {
1952 if let Ok(mut subs) = self.inner.subscriptions.write() {
1953 return subs.remove(uri);
1954 }
1955 false
1956 }
1957
1958 pub fn notify_resource_updated(&self, uri: &str) -> bool {
1965 let notification = ServerNotification::ResourceUpdated {
1966 uri: uri.to_string(),
1967 };
1968 let mut sent = false;
1969
1970 if self.is_subscribed(uri)
1971 && let Some(tx) = &self.inner.notification_tx
1972 {
1973 sent |= tx.try_send(notification.clone()).is_ok();
1974 }
1975
1976 #[cfg(all(feature = "http", feature = "stateless"))]
1977 if let Ok(active) = self.inner.modern_notification_sink.read()
1978 && let Some(sink) = active.as_ref()
1979 {
1980 sent |= sink(¬ification);
1981 }
1982
1983 sent
1984 }
1985
1986 pub async fn notify_task_status_changed(&self, task_id: &str) {
2001 self.notify_task_state(task_id).await;
2002 }
2003
2004 pub fn notify_resources_list_changed(&self) -> bool {
2008 let Some(tx) = &self.inner.notification_tx else {
2009 return false;
2010 };
2011 tx.try_send(ServerNotification::ResourcesListChanged)
2012 .is_ok()
2013 }
2014
2015 pub fn notify_tools_list_changed(&self) -> bool {
2019 let Some(tx) = &self.inner.notification_tx else {
2020 return false;
2021 };
2022 tx.try_send(ServerNotification::ToolsListChanged).is_ok()
2023 }
2024
2025 pub fn notify_prompts_list_changed(&self) -> bool {
2029 let Some(tx) = &self.inner.notification_tx else {
2030 return false;
2031 };
2032 tx.try_send(ServerNotification::PromptsListChanged).is_ok()
2033 }
2034
2035 pub fn disable_tool(&self, name: impl Into<String>) {
2046 let mut set = self.inner.disabled_tools.write().unwrap();
2047 set.insert(name.into());
2048 }
2049
2050 pub fn enable_tool(&self, name: &str) {
2053 let mut set = self.inner.disabled_tools.write().unwrap();
2054 set.remove(name);
2055 }
2056
2057 pub fn is_tool_enabled(&self, name: &str) -> bool {
2061 !self.inner.disabled_tools.read().unwrap().contains(name)
2062 }
2063
2064 pub fn disable_resource(&self, uri: impl Into<String>) {
2067 let mut set = self.inner.disabled_resources.write().unwrap();
2068 set.insert(uri.into());
2069 }
2070
2071 pub fn enable_resource(&self, uri: &str) {
2073 let mut set = self.inner.disabled_resources.write().unwrap();
2074 set.remove(uri);
2075 }
2076
2077 pub fn is_resource_enabled(&self, uri: &str) -> bool {
2079 !self.inner.disabled_resources.read().unwrap().contains(uri)
2080 }
2081
2082 pub fn disable_prompt(&self, name: impl Into<String>) {
2085 let mut set = self.inner.disabled_prompts.write().unwrap();
2086 set.insert(name.into());
2087 }
2088
2089 pub fn enable_prompt(&self, name: &str) {
2091 let mut set = self.inner.disabled_prompts.write().unwrap();
2092 set.remove(name);
2093 }
2094
2095 pub fn is_prompt_enabled(&self, name: &str) -> bool {
2097 !self.inner.disabled_prompts.read().unwrap().contains(name)
2098 }
2099
2100 pub(crate) fn implementation(&self) -> Implementation {
2110 Implementation {
2111 name: self.inner.server_name.clone(),
2112 version: self.inner.server_version.clone(),
2113 title: self.inner.server_title.clone(),
2114 description: self.inner.server_description.clone(),
2115 icons: self.inner.server_icons.clone(),
2116 website_url: self.inner.server_website_url.clone(),
2117 meta: None,
2118 }
2119 }
2120
2121 #[cfg(feature = "http")]
2127 pub(crate) fn tool_input_schema(&self, name: &str) -> Option<serde_json::Value> {
2128 if let Some(tool) = self.inner.tools.get(name) {
2129 return Some(tool.input_schema.clone());
2130 }
2131 #[cfg(feature = "dynamic-tools")]
2132 if let Some(tool) = self
2133 .inner
2134 .dynamic_tools
2135 .as_ref()
2136 .and_then(|tools| tools.get(name))
2137 {
2138 return Some(tool.input_schema.clone());
2139 }
2140 None
2141 }
2142
2143 fn capabilities(&self) -> ServerCapabilities {
2144 let has_resources =
2145 !self.inner.resources.is_empty() || !self.inner.resource_templates.is_empty();
2146 let has_notifications = self.inner.notification_tx.is_some();
2147
2148 #[cfg(feature = "dynamic-tools")]
2149 let has_dynamic_tools = self.inner.dynamic_tools.is_some();
2150 #[cfg(not(feature = "dynamic-tools"))]
2151 let has_dynamic_tools = false;
2152
2153 #[cfg(feature = "dynamic-tools")]
2154 let has_dynamic_prompts = self.inner.dynamic_prompts.is_some();
2155 #[cfg(not(feature = "dynamic-tools"))]
2156 let has_dynamic_prompts = false;
2157
2158 #[cfg(feature = "dynamic-tools")]
2159 let has_dynamic_resources = self.inner.dynamic_resources.is_some()
2160 || self.inner.dynamic_resource_templates.is_some();
2161 #[cfg(not(feature = "dynamic-tools"))]
2162 let has_dynamic_resources = false;
2163
2164 ServerCapabilities {
2165 tools: if self.inner.tools.is_empty() && !has_dynamic_tools {
2166 None
2167 } else {
2168 Some(ToolsCapability {
2169 list_changed: has_notifications,
2170 })
2171 },
2172 resources: if has_resources || has_dynamic_resources {
2173 Some(ResourcesCapability {
2174 subscribe: true,
2175 list_changed: has_notifications,
2176 })
2177 } else {
2178 None
2179 },
2180 prompts: if self.inner.prompts.is_empty() && !has_dynamic_prompts {
2181 None
2182 } else {
2183 Some(PromptsCapability {
2184 list_changed: has_notifications,
2185 })
2186 },
2187 logging: if self.inner.notification_tx.is_some() {
2189 Some(LoggingCapability {
2190 deprecated: self.inner.logging_deprecated.clone(),
2191 })
2192 } else {
2193 None
2194 },
2195 tasks: {
2201 let has_task_support = self
2202 .inner
2203 .tools
2204 .values()
2205 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2206 if has_task_support {
2207 Some(TasksCapability {
2208 list: None,
2212 cancel: Some(TasksCancelCapability {}),
2213 requests: Some(TasksRequestsCapability {
2214 tools: Some(TasksToolsRequestsCapability {
2215 call: Some(TasksToolsCallCapability {}),
2216 }),
2217 }),
2218 })
2219 } else {
2220 None
2221 }
2222 },
2223 completions: if self.inner.completion_handler.is_some() {
2225 Some(CompletionsCapability::default())
2226 } else {
2227 None
2228 },
2229 experimental: None,
2230 extensions: {
2231 let mut map = self.inner.protocol_extensions.clone();
2232 let has_task_support = self
2233 .inner
2234 .tools
2235 .values()
2236 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2237 if has_task_support {
2238 map.insert(
2239 tower_mcp_types::protocol::TASKS_EXTENSION_ID.to_string(),
2240 serde_json::json!({}),
2241 );
2242 }
2243 (!map.is_empty()).then_some(map)
2244 },
2245 }
2246 }
2247
2248 fn capabilities_for_protocol(&self, protocol_version: Option<&str>) -> ServerCapabilities {
2256 let mut capabilities = self.capabilities();
2257 if protocol_version == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
2258 capabilities.tasks = None;
2259 if !self.final_tasks_enabled()
2260 && let Some(extensions) = capabilities.extensions.as_mut()
2261 {
2262 extensions.remove(tower_mcp_types::protocol::TASKS_EXTENSION_ID);
2263 if extensions.is_empty() {
2264 capabilities.extensions = None;
2265 }
2266 }
2267 }
2268 capabilities
2269 }
2270
2271 pub(crate) fn final_tasks_enabled(&self) -> bool {
2276 self.inner
2277 .protocol_extensions
2278 .contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
2279 }
2280
2281 fn require_negotiated_tasks(
2286 &self,
2287 extensions: &crate::context::Extensions,
2288 method: &str,
2289 ) -> Result<()> {
2290 if !self.final_tasks_enabled() {
2291 return Err(Error::JsonRpc(JsonRpcError::method_not_found(method)));
2292 }
2293 if client_declares_tasks(extensions) {
2294 return Ok(());
2295 }
2296 Err(Error::JsonRpc(
2297 JsonRpcError::missing_required_client_capability(tasks_client_capabilities()),
2298 ))
2299 }
2300
2301 async fn authorize_task(
2307 &self,
2308 task_id: &str,
2309 extensions: &crate::context::Extensions,
2310 ) -> Result<()> {
2311 let owner = self
2312 .inner
2313 .task_store
2314 .task_owner(task_id)
2315 .await
2316 .map_err(task_store_error)?
2317 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2318
2319 if crate::async_task::owner_matches(&owner, request_principal(extensions).as_deref()) {
2320 Ok(())
2321 } else {
2322 tracing::debug!(
2323 target: "mcp::tasks",
2324 task_id = %task_id,
2325 "task operation refused: principal does not own the task"
2326 );
2327 Err(Error::JsonRpc(unknown_task_error(task_id)))
2328 }
2329 }
2330
2331 async fn final_get_task(&self, task_id: &str) -> Result<McpResponse> {
2333 let (detailed, meta) = self.detailed_task(task_id).await?;
2334 let mut result = crate::tasks::GetTaskResult::new(detailed);
2335 result.meta = meta;
2336 Ok(McpResponse::FinalGetTask(result))
2337 }
2338
2339 async fn detailed_task(
2345 &self,
2346 task_id: &str,
2347 ) -> Result<(
2348 crate::tasks::DetailedTask,
2349 Option<serde_json::Map<String, serde_json::Value>>,
2350 )> {
2351 let (task, result, error) = self
2352 .inner
2353 .task_store
2354 .get_task_result(task_id)
2355 .await
2356 .map_err(task_store_error)?
2357 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2358
2359 let mut metadata = crate::tasks::TaskMetadata::new(
2360 task.task_id.clone(),
2361 task.created_at.clone(),
2362 task.last_updated_at.clone(),
2363 task.ttl,
2364 );
2365 metadata.status_message = task.status_message.clone();
2366 metadata.poll_interval_ms = task.poll_interval;
2367
2368 let meta = task.meta.and_then(|value| value.as_object().cloned());
2369 let detailed = match task.status {
2370 TaskStatus::Working => crate::tasks::DetailedTask::working(metadata),
2371 TaskStatus::InputRequired => {
2372 let outstanding = self
2375 .inner
2376 .task_store
2377 .outstanding_input_requests(task_id)
2378 .await
2379 .map_err(task_store_error)?
2380 .unwrap_or_default();
2381 crate::tasks::DetailedTask::input_required(metadata, outstanding)
2382 }
2383 TaskStatus::Completed => {
2384 let mut object = result
2387 .map(serde_json::to_value)
2388 .transpose()
2389 .map_err(|e| {
2390 Error::JsonRpc(JsonRpcError::internal_error(format!(
2391 "failed to encode task result: {e}"
2392 )))
2393 })?
2394 .and_then(|value| value.as_object().cloned())
2395 .unwrap_or_default();
2396 object.insert(
2400 "resultType".to_string(),
2401 serde_json::Value::String("complete".to_string()),
2402 );
2403 crate::tasks::DetailedTask::completed(metadata, object)
2404 }
2405 TaskStatus::Failed => crate::tasks::DetailedTask::failed(
2406 metadata,
2407 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
2408 ),
2409 TaskStatus::Cancelled => crate::tasks::DetailedTask::cancelled(metadata),
2410 _ => crate::tasks::DetailedTask::working(metadata),
2413 };
2414 Ok((detailed, meta))
2415 }
2416
2417 async fn notify_task_state(&self, task_id: &str) {
2426 if !self.final_tasks_enabled() {
2427 return;
2428 }
2429
2430 let (detailed, meta) = match self.detailed_task(task_id).await {
2431 Ok(detailed) => detailed,
2432 Err(error) => {
2433 tracing::debug!(
2434 target: "mcp::tasks",
2435 task_id = %task_id,
2436 %error,
2437 "skipping task notification: task state unavailable"
2438 );
2439 return;
2440 }
2441 };
2442
2443 let notification = ServerNotification::FinalTaskStatusChanged(
2444 crate::tasks::TaskStatusNotificationParams {
2445 task: detailed,
2446 meta,
2447 },
2448 );
2449
2450 #[cfg(all(feature = "http", feature = "stateless"))]
2455 if let Ok(active) = self.inner.modern_notification_sink.read()
2456 && let Some(sink) = active.as_ref()
2457 {
2458 sink(¬ification);
2459 return;
2460 }
2461
2462 if let Some(tx) = &self.inner.notification_tx {
2463 let _ = tx.try_send(notification);
2464 }
2465 }
2466
2467 fn effective_cache_scope(&self, ttl_ms: Option<u64>) -> Option<CacheScope> {
2474 self.inner
2475 .cache_scope
2476 .or_else(|| ttl_ms.map(|_| CacheScope::Private))
2477 }
2478
2479 fn apply_read_cache_hints(&self, mut result: ReadResourceResult) -> ReadResourceResult {
2484 if result.ttl_ms.is_none() {
2485 result.ttl_ms = self.inner.read_ttl_ms;
2486 }
2487 if result.cache_scope.is_none() {
2488 result.cache_scope = self.effective_cache_scope(result.ttl_ms);
2489 }
2490 result
2491 }
2492
2493 async fn handle(
2495 &self,
2496 request_id: RequestId,
2497 request: McpRequest,
2498 extensions: Extensions,
2499 ) -> Result<McpResponse> {
2500 let method = request.method_name();
2502 if !is_final_protocol_request(&extensions) && !self.session.is_request_allowed(method) {
2503 tracing::warn!(
2504 method = %method,
2505 phase = ?self.session.phase(),
2506 "Request rejected: session not initialized"
2507 );
2508 return Err(Error::JsonRpc(JsonRpcError::invalid_request(format!(
2509 "Session not initialized. Only 'initialize' and 'ping' are allowed before initialization. Got: {}",
2510 method
2511 ))));
2512 }
2513
2514 match request {
2515 McpRequest::Initialize(params) => {
2516 tracing::info!(
2517 client = %params.client_info.name,
2518 version = %params.client_info.version,
2519 "Client initializing"
2520 );
2521
2522 let protocol_support = extensions.get::<crate::ProtocolSupport>();
2526 let requested_is_legacy = crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2527 .contains(¶ms.protocol_version.as_str());
2528 let requested_is_supported = requested_is_legacy
2529 && protocol_support
2530 .is_none_or(|support| support.contains(¶ms.protocol_version));
2531 let protocol_version = if requested_is_supported {
2532 params.protocol_version
2533 } else {
2534 match protocol_support {
2535 None => crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
2536 Some(support) => support
2537 .versions()
2538 .iter()
2539 .find(|version| {
2540 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2541 .contains(&version.as_str())
2542 })
2543 .cloned()
2544 .ok_or_else(|| {
2545 Error::JsonRpc(JsonRpcError::unsupported_protocol_version(
2546 params.protocol_version,
2547 support.versions().iter().map(String::as_str),
2548 ))
2549 })?,
2550 }
2551 };
2552
2553 self.session.mark_initializing();
2555 let capabilities = self.capabilities_for_protocol(Some(&protocol_version));
2556 self.session.insert(params.capabilities.clone());
2557 self.session
2558 .insert(crate::NegotiatedExtensions::from_capabilities(
2559 ¶ms.capabilities,
2560 &capabilities,
2561 ));
2562
2563 Ok(McpResponse::Initialize(InitializeResult {
2564 protocol_version,
2565 capabilities,
2566 server_info: self.implementation(),
2567 instructions: if let Some(config) = &self.inner.auto_instructions {
2568 Some(self.inner.generate_instructions(config))
2569 } else {
2570 self.inner.instructions.clone()
2571 },
2572 meta: None,
2573 }))
2574 }
2575
2576 McpRequest::Discover(_) => {
2577 tracing::debug!("Stateless server/discover request");
2584 let server_info = self.implementation();
2585 let supported_versions = extensions.get::<crate::ProtocolSupport>().map_or_else(
2586 || {
2587 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2588 .iter()
2589 .map(|version| (*version).to_string())
2590 .collect()
2591 },
2592 |support| support.versions().to_vec(),
2593 );
2594 let capabilities = self
2599 .capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
2600 Ok(McpResponse::Discover(DiscoverResult {
2601 supported_versions,
2602 capabilities,
2603 ttl_ms: None,
2604 cache_scope: None,
2605 instructions: if let Some(config) = &self.inner.auto_instructions {
2606 Some(self.inner.generate_instructions(config))
2607 } else {
2608 self.inner.instructions.clone()
2609 },
2610 meta: Some(crate::protocol::ResultMeta {
2611 server_info: Some(server_info),
2612 }),
2613 }))
2614 }
2615
2616 McpRequest::ListTools(params) => {
2617 let final_protocol = is_final_protocol_request(&extensions);
2618 let final_tasks_negotiated = final_protocol
2619 && self.final_tasks_enabled()
2620 && client_declares_tasks(&extensions);
2621 let filter = self.inner.tool_filter.as_ref();
2622 let disabled = self.inner.disabled_tools.read().unwrap().clone();
2623 let is_visible = |t: &Tool| {
2624 !disabled.contains(&t.name)
2625 && !(final_protocol
2626 && matches!(t.task_support, TaskSupportMode::Required)
2627 && !final_tasks_negotiated)
2628 && filter
2629 .map(|f| f.is_visible(&self.session, t))
2630 .unwrap_or(true)
2631 };
2632 let definition = |t: &Tool| {
2633 let mut definition = t.definition();
2634 if final_protocol {
2635 definition.execution = None;
2636 }
2637 definition
2638 };
2639
2640 let mut tools: Vec<ToolDefinition> = self
2642 .inner
2643 .tools
2644 .values()
2645 .filter(|t| is_visible(t))
2646 .map(|t| definition(t))
2647 .collect();
2648
2649 #[cfg(feature = "dynamic-tools")]
2651 if let Some(ref dynamic) = self.inner.dynamic_tools {
2652 let static_names: HashSet<String> =
2653 tools.iter().map(|t| t.name.clone()).collect();
2654 for t in dynamic.list() {
2655 if !static_names.contains(&t.name) && is_visible(&t) {
2656 tools.push(definition(&t));
2657 }
2658 }
2659 }
2660
2661 tools.sort_by(|a, b| a.name.cmp(&b.name));
2662
2663 let (tools, next_cursor) =
2664 paginate(tools, params.cursor.as_deref(), self.inner.page_size)?;
2665
2666 Ok(McpResponse::ListTools(ListToolsResult {
2667 tools,
2668 next_cursor,
2669 ttl_ms: self.inner.list_ttl_ms,
2670 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
2671 meta: None,
2672 }))
2673 }
2674
2675 McpRequest::CallTool(params) => {
2676 if self
2678 .inner
2679 .disabled_tools
2680 .read()
2681 .unwrap()
2682 .contains(¶ms.name)
2683 {
2684 tracing::info!(
2685 target: "mcp::tools",
2686 tool = %params.name,
2687 status = "disabled",
2688 "tool call completed"
2689 );
2690 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2691 }
2692
2693 let tool = self.inner.tools.get(¶ms.name).cloned();
2695 #[cfg(feature = "dynamic-tools")]
2696 let tool = tool.or_else(|| {
2697 self.inner
2698 .dynamic_tools
2699 .as_ref()
2700 .and_then(|d| d.get(¶ms.name))
2701 });
2702
2703 let tool = match tool {
2704 Some(t) => t,
2705 None => {
2706 tracing::info!(
2707 target: "mcp::tools",
2708 tool = %params.name,
2709 status = "not_found",
2710 "tool call completed"
2711 );
2712 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2713 }
2714 };
2715
2716 if let Some(filter) = &self.inner.tool_filter
2718 && !filter.is_visible(&self.session, &tool)
2719 {
2720 tracing::info!(
2721 target: "mcp::tools",
2722 tool = %params.name,
2723 status = "denied",
2724 "tool call completed"
2725 );
2726 return Err(filter.denial_error(¶ms.name));
2727 }
2728
2729 let final_protocol = is_final_protocol_request(&extensions);
2733 let task_ttl = if final_protocol {
2734 if params.task.is_some() {
2735 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
2736 "The final Tasks extension does not allow a 'task' request parameter",
2737 )));
2738 }
2739
2740 let server_enabled = self.final_tasks_enabled();
2741 let tasks_negotiated = server_enabled && client_declares_tasks(&extensions);
2742 match tool.task_support {
2743 TaskSupportMode::Required if !server_enabled => {
2744 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
2747 ¶ms.name,
2748 )));
2749 }
2750 TaskSupportMode::Required if !tasks_negotiated => {
2751 return Err(Error::JsonRpc(
2752 JsonRpcError::missing_required_client_capability(
2753 tasks_client_capabilities(),
2754 ),
2755 ));
2756 }
2757 TaskSupportMode::Required | TaskSupportMode::Optional
2758 if tasks_negotiated =>
2759 {
2760 Some(None)
2761 }
2762 _ => None,
2763 }
2764 } else {
2765 match (¶ms.task, tool.task_support) {
2766 (Some(_), TaskSupportMode::Forbidden) => {
2767 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2768 "Tool '{}' does not support async tasks",
2769 params.name
2770 ))));
2771 }
2772 (None, TaskSupportMode::Required) => {
2773 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2774 "Tool '{}' requires async task execution (include 'task' in params)",
2775 params.name
2776 ))));
2777 }
2778 (Some(task), _) => Some(task.ttl),
2779 (None, _) => None,
2780 }
2781 };
2782
2783 #[cfg(feature = "stateless")]
2787 if let Some(required) = tool.required_client_capabilities()
2788 && let Some(meta) = extensions.get::<crate::stateless::StatelessRequestMeta>()
2789 && meta.protocol_version.as_deref()
2790 == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
2791 && !meta
2792 .client_capabilities
2793 .as_ref()
2794 .is_some_and(|actual| client_capabilities_satisfy(actual, required))
2795 {
2796 return Err(Error::JsonRpc(
2797 JsonRpcError::missing_required_client_capability(required.clone()),
2798 ));
2799 }
2800
2801 if let Some(task_ttl) = task_ttl {
2802 let (task_id, cancellation_token) = self
2804 .inner
2805 .task_store
2806 .create_task(
2807 ¶ms.name,
2808 params.arguments.clone(),
2809 task_ttl,
2810 request_principal(&extensions),
2811 )
2812 .await
2813 .map_err(task_store_error)?;
2814
2815 tracing::info!(task_id = %task_id, tool = %params.name, "Created async task");
2816
2817 let progress_token = params.meta.and_then(|m| m.progress_token);
2819 let ctx = self.create_context_with_extensions(
2820 request_id,
2821 progress_token,
2822 &extensions,
2823 );
2824
2825 let task_store = self.inner.task_store.clone();
2826 let task_context = crate::tool::TaskContext::new(task_id.clone());
2827 let mut ctx = ctx;
2828 ctx.extensions_mut().insert(task_context.clone());
2829 let preparation = match tool
2830 .prepare_task(task_context, params.arguments.clone())
2831 .await
2832 {
2833 Ok(preparation) => preparation,
2834 Err(error) => {
2835 discard_unprepared_task(&task_store, &task_id).await;
2836 return Err(error);
2837 }
2838 };
2839 if let Some(meta) = preparation.meta {
2840 let value = serde_json::Value::Object(meta);
2841 if let Err(error) = crate::protocol::validate_meta_object(&value) {
2842 discard_unprepared_task(&task_store, &task_id).await;
2843 return Err(Error::invalid_params(format!(
2844 "Invalid task metadata: {error}"
2845 )));
2846 }
2847 let persisted = match task_store.set_task_meta(&task_id, value).await {
2848 Ok(persisted) => persisted,
2849 Err(error) => {
2850 discard_unprepared_task(&task_store, &task_id).await;
2851 return Err(task_store_error(error));
2852 }
2853 };
2854 if !persisted {
2855 discard_unprepared_task(&task_store, &task_id).await;
2856 return Err(Error::JsonRpc(JsonRpcError::internal_error(
2857 "Task store could not persist preparation metadata",
2858 )));
2859 }
2860 }
2861 ctx.extensions_mut().merge(&preparation.extensions);
2862
2863 let tool = tool.clone();
2865 let arguments = params.arguments;
2866 let task_id_clone = task_id.clone();
2867
2868 let tool_name = params.name.clone();
2869 let notifier = self.clone();
2870 tokio::spawn(async move {
2871 if cancellation_token.is_cancelled() {
2873 tracing::debug!(task_id = %task_id_clone, "Task cancelled before execution");
2874 notifier.notify_task_state(&task_id_clone).await;
2875 return;
2876 }
2877
2878 let start = std::time::Instant::now();
2880 let result = tool.call_with_context(ctx, arguments).await;
2881 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
2882
2883 if cancellation_token.is_cancelled() {
2884 tracing::debug!(task_id = %task_id_clone, "Task cancelled during execution");
2885 notifier.notify_task_state(&task_id_clone).await;
2886 } else {
2887 let status = if result.is_error { "error" } else { "success" };
2892 let error_msg = result
2893 .is_error
2894 .then(|| result.first_text().unwrap_or("Tool execution failed"))
2895 .map(str::to_string);
2896 if let Err(e) = task_store.complete_task(&task_id_clone, result).await {
2897 tracing::warn!(task_id = %task_id_clone, error = %e, "failed to record task completion");
2898 }
2899 tracing::info!(
2900 target: "mcp::tools",
2901 tool = %tool_name,
2902 task_id = %task_id_clone,
2903 duration_ms,
2904 status,
2905 error = error_msg.as_deref().unwrap_or_default(),
2906 "tool call completed"
2907 );
2908 notifier.notify_task_state(&task_id_clone).await;
2909 }
2910 });
2911
2912 let task = self
2913 .inner
2914 .task_store
2915 .get_task(&task_id)
2916 .await
2917 .map_err(task_store_error)?
2918 .ok_or_else(|| {
2919 Error::JsonRpc(JsonRpcError::internal_error(
2920 "Failed to retrieve created task",
2921 ))
2922 })?;
2923
2924 if is_final_protocol_request(&extensions) {
2928 let mut metadata = crate::tasks::TaskMetadata::new(
2929 task.task_id.clone(),
2930 task.created_at.clone(),
2931 task.last_updated_at.clone(),
2932 task.ttl,
2933 );
2934 metadata.status_message = task.status_message.clone();
2935 metadata.poll_interval_ms = task.poll_interval;
2936 let mut result = crate::tasks::CreateTaskResult::new(
2937 crate::tasks::Task::new(metadata, task.status),
2938 );
2939 result.meta = task.meta.and_then(|value| value.as_object().cloned());
2940 return Ok(McpResponse::FinalCreateTask(result));
2941 }
2942 Ok(McpResponse::CreateTask(CreateTaskResult::new(task)))
2943 } else {
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 #[cfg(feature = "stateless")]
2952 let ctx = {
2953 let mut ctx = ctx;
2954 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
2955 params.input_responses,
2956 params.request_state,
2957 ));
2958 ctx
2959 };
2960
2961 let start = std::time::Instant::now();
2962 let outcome = tool
2963 .call_outcome_with_context(ctx, params.arguments)
2964 .await?;
2965 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
2966
2967 match outcome {
2968 RequestOutcome::Complete(result) => {
2969 let status = if result.is_error { "error" } else { "success" };
2970 tracing::info!(
2971 target: "mcp::tools",
2972 tool = %params.name,
2973 duration_ms,
2974 status,
2975 "tool call completed"
2976 );
2977 Ok(McpResponse::CallTool(result))
2978 }
2979 RequestOutcome::InputRequired(result) => {
2980 #[cfg(feature = "stateless")]
2981 {
2982 validate_input_required_result(&extensions, &result)?;
2983 tracing::info!(
2984 target: "mcp::tools",
2985 tool = %params.name,
2986 duration_ms,
2987 status = "input_required",
2988 "tool call requires client input"
2989 );
2990 Ok(McpResponse::InputRequired(result))
2991 }
2992 #[cfg(not(feature = "stateless"))]
2993 {
2994 let _ = result;
2995 Err(Error::invalid_params(
2996 "InputRequiredResult support was not compiled",
2997 ))
2998 }
2999 }
3000 }
3001 }
3002 }
3003
3004 McpRequest::ListResources(params) => {
3005 let disabled = self.inner.disabled_resources.read().unwrap().clone();
3006 let is_visible = |r: &Resource| -> bool {
3007 !disabled.contains(&r.uri)
3008 && self
3009 .inner
3010 .resource_filter
3011 .as_ref()
3012 .map(|f| f.is_visible(&self.session, r))
3013 .unwrap_or(true)
3014 };
3015
3016 let mut resources: Vec<ResourceDefinition> = self
3017 .inner
3018 .resources
3019 .values()
3020 .filter(|r| is_visible(r))
3021 .map(|r| r.definition())
3022 .collect();
3023
3024 #[cfg(feature = "dynamic-tools")]
3026 if let Some(ref dynamic) = self.inner.dynamic_resources {
3027 let static_uris: HashSet<String> =
3028 resources.iter().map(|r| r.uri.clone()).collect();
3029 for r in dynamic.list() {
3030 if !static_uris.contains(&r.uri) && is_visible(&r) {
3031 resources.push(r.definition());
3032 }
3033 }
3034 }
3035
3036 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
3037
3038 let (resources, next_cursor) =
3039 paginate(resources, params.cursor.as_deref(), self.inner.page_size)?;
3040
3041 Ok(McpResponse::ListResources(ListResourcesResult {
3042 resources,
3043 next_cursor,
3044 ttl_ms: self.inner.list_ttl_ms,
3045 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3046 meta: None,
3047 }))
3048 }
3049
3050 McpRequest::ListResourceTemplates(params) => {
3051 let mut resource_templates: Vec<ResourceTemplateDefinition> = self
3052 .inner
3053 .resource_templates
3054 .iter()
3055 .map(|t| t.definition())
3056 .collect();
3057
3058 #[cfg(feature = "dynamic-tools")]
3060 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3061 let static_patterns: HashSet<String> = resource_templates
3062 .iter()
3063 .map(|t| t.uri_template.clone())
3064 .collect();
3065 for t in dynamic.list() {
3066 if !static_patterns.contains(&t.uri_template) {
3067 resource_templates.push(t.definition());
3068 }
3069 }
3070 }
3071
3072 resource_templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
3073
3074 let (resource_templates, next_cursor) = paginate(
3075 resource_templates,
3076 params.cursor.as_deref(),
3077 self.inner.page_size,
3078 )?;
3079
3080 Ok(McpResponse::ListResourceTemplates(
3081 ListResourceTemplatesResult {
3082 resource_templates,
3083 next_cursor,
3084 ttl_ms: self.inner.list_ttl_ms,
3085 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3086 meta: None,
3087 },
3088 ))
3089 }
3090
3091 McpRequest::ReadResource(params) => {
3092 if self
3094 .inner
3095 .disabled_resources
3096 .read()
3097 .unwrap()
3098 .contains(¶ms.uri)
3099 {
3100 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3101 ¶ms.uri,
3102 )));
3103 }
3104
3105 if let Some(resource) = self.inner.resources.get(¶ms.uri) {
3107 if let Some(filter) = &self.inner.resource_filter
3109 && !filter.is_visible(&self.session, resource)
3110 {
3111 return Err(filter.denial_error(¶ms.uri));
3112 }
3113
3114 tracing::debug!(uri = %params.uri, "Reading static resource");
3115 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3116 #[cfg(feature = "stateless")]
3117 let ctx = {
3118 let mut ctx = ctx;
3119 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3120 params.input_responses.clone(),
3121 params.request_state.clone(),
3122 ));
3123 ctx
3124 };
3125 return match resource.read_outcome_with_context(ctx).await? {
3126 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3127 self.apply_read_cache_hints(result),
3128 )),
3129 RequestOutcome::InputRequired(result) => {
3130 #[cfg(feature = "stateless")]
3131 {
3132 validate_input_required_result(&extensions, &result)?;
3133 Ok(McpResponse::InputRequired(result))
3134 }
3135 #[cfg(not(feature = "stateless"))]
3136 {
3137 let _ = result;
3138 Err(Error::invalid_params(
3139 "InputRequiredResult support was not compiled",
3140 ))
3141 }
3142 }
3143 };
3144 }
3145
3146 #[cfg(feature = "dynamic-tools")]
3148 #[allow(clippy::collapsible_if)]
3149 if let Some(ref dynamic) = self.inner.dynamic_resources {
3150 if let Some(resource) = dynamic.get(¶ms.uri) {
3151 if let Some(filter) = &self.inner.resource_filter
3152 && !filter.is_visible(&self.session, &resource)
3153 {
3154 return Err(filter.denial_error(¶ms.uri));
3155 }
3156 tracing::debug!(uri = %params.uri, "Reading dynamic resource");
3157 let ctx =
3158 self.create_context_with_extensions(request_id, None, &extensions);
3159 #[cfg(feature = "stateless")]
3160 let ctx = {
3161 let mut ctx = ctx;
3162 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3163 params.input_responses.clone(),
3164 params.request_state.clone(),
3165 ));
3166 ctx
3167 };
3168 return match resource.read_outcome_with_context(ctx).await? {
3169 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3170 self.apply_read_cache_hints(result),
3171 )),
3172 RequestOutcome::InputRequired(result) => {
3173 #[cfg(feature = "stateless")]
3174 {
3175 validate_input_required_result(&extensions, &result)?;
3176 Ok(McpResponse::InputRequired(result))
3177 }
3178 #[cfg(not(feature = "stateless"))]
3179 {
3180 let _ = result;
3181 Err(Error::invalid_params(
3182 "InputRequiredResult support was not compiled",
3183 ))
3184 }
3185 }
3186 };
3187 }
3188 }
3189
3190 for template in &self.inner.resource_templates {
3192 if let Some(variables) = template.match_uri(¶ms.uri) {
3193 tracing::debug!(
3194 uri = %params.uri,
3195 template = %template.uri_template,
3196 "Reading resource via template"
3197 );
3198 let ctx =
3199 self.create_context_with_extensions(request_id, None, &extensions);
3200 #[cfg(feature = "stateless")]
3201 let ctx = {
3202 let mut ctx = ctx;
3203 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3204 params.input_responses.clone(),
3205 params.request_state.clone(),
3206 ));
3207 ctx
3208 };
3209 return match template
3210 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3211 .await?
3212 {
3213 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3214 self.apply_read_cache_hints(result),
3215 )),
3216 RequestOutcome::InputRequired(result) => {
3217 #[cfg(feature = "stateless")]
3218 {
3219 validate_input_required_result(&extensions, &result)?;
3220 Ok(McpResponse::InputRequired(result))
3221 }
3222 #[cfg(not(feature = "stateless"))]
3223 {
3224 let _ = result;
3225 Err(Error::invalid_params(
3226 "InputRequiredResult support was not compiled",
3227 ))
3228 }
3229 }
3230 };
3231 }
3232 }
3233
3234 #[cfg(feature = "dynamic-tools")]
3236 #[allow(clippy::collapsible_if)]
3237 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3238 if let Some((template, variables)) = dynamic.match_uri(¶ms.uri) {
3239 tracing::debug!(
3240 uri = %params.uri,
3241 template = %template.uri_template,
3242 "Reading resource via dynamic template"
3243 );
3244 let ctx =
3245 self.create_context_with_extensions(request_id, None, &extensions);
3246 #[cfg(feature = "stateless")]
3247 let ctx = {
3248 let mut ctx = ctx;
3249 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3250 params.input_responses.clone(),
3251 params.request_state.clone(),
3252 ));
3253 ctx
3254 };
3255 return match template
3256 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3257 .await?
3258 {
3259 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3260 self.apply_read_cache_hints(result),
3261 )),
3262 RequestOutcome::InputRequired(result) => {
3263 #[cfg(feature = "stateless")]
3264 {
3265 validate_input_required_result(&extensions, &result)?;
3266 Ok(McpResponse::InputRequired(result))
3267 }
3268 #[cfg(not(feature = "stateless"))]
3269 {
3270 let _ = result;
3271 Err(Error::invalid_params(
3272 "InputRequiredResult support was not compiled",
3273 ))
3274 }
3275 }
3276 };
3277 }
3278 }
3279
3280 Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3282 ¶ms.uri,
3283 )))
3284 }
3285
3286 McpRequest::SubscribeResource(params) => {
3287 if !self.inner.resources.contains_key(¶ms.uri) {
3289 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3290 ¶ms.uri,
3291 )));
3292 }
3293
3294 tracing::debug!(uri = %params.uri, "Subscribing to resource");
3295 self.subscribe(¶ms.uri);
3296
3297 Ok(McpResponse::SubscribeResource(EmptyResult {}))
3298 }
3299
3300 McpRequest::UnsubscribeResource(params) => {
3301 if !self.inner.resources.contains_key(¶ms.uri) {
3303 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3304 ¶ms.uri,
3305 )));
3306 }
3307
3308 tracing::debug!(uri = %params.uri, "Unsubscribing from resource");
3309 self.unsubscribe(¶ms.uri);
3310
3311 Ok(McpResponse::UnsubscribeResource(EmptyResult {}))
3312 }
3313
3314 McpRequest::ListPrompts(params) => {
3315 #[cfg(feature = "dynamic-tools")]
3316 if let Some(initializer) = &self.inner.prompt_initializer {
3317 initializer()?;
3318 }
3319 let disabled = self.inner.disabled_prompts.read().unwrap().clone();
3320 let is_visible = |p: &Prompt| -> bool {
3321 !disabled.contains(&p.name)
3322 && self
3323 .inner
3324 .prompt_filter
3325 .as_ref()
3326 .map(|f| f.is_visible(&self.session, p))
3327 .unwrap_or(true)
3328 };
3329
3330 let mut prompts: Vec<PromptDefinition> = self
3331 .inner
3332 .prompts
3333 .values()
3334 .filter(|p| is_visible(p))
3335 .map(|p| p.definition())
3336 .collect();
3337
3338 #[cfg(feature = "dynamic-tools")]
3340 if let Some(ref dynamic) = self.inner.dynamic_prompts {
3341 let static_names: HashSet<String> =
3342 prompts.iter().map(|p| p.name.clone()).collect();
3343 for p in dynamic.list() {
3344 if !static_names.contains(&p.name) && is_visible(&p) {
3345 prompts.push(p.definition());
3346 }
3347 }
3348 }
3349
3350 prompts.sort_by(|a, b| a.name.cmp(&b.name));
3351
3352 let (prompts, next_cursor) =
3353 paginate(prompts, params.cursor.as_deref(), self.inner.page_size)?;
3354
3355 Ok(McpResponse::ListPrompts(ListPromptsResult {
3356 prompts,
3357 next_cursor,
3358 ttl_ms: self.inner.list_ttl_ms,
3359 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3360 meta: None,
3361 }))
3362 }
3363
3364 McpRequest::GetPrompt(params) => {
3365 #[cfg(feature = "dynamic-tools")]
3366 if let Some(initializer) = &self.inner.prompt_initializer {
3367 initializer()?;
3368 }
3369 if self
3371 .inner
3372 .disabled_prompts
3373 .read()
3374 .unwrap()
3375 .contains(¶ms.name)
3376 {
3377 return Err(Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3378 "Prompt not found: {}",
3379 params.name
3380 ))));
3381 }
3382
3383 let prompt = self.inner.prompts.get(¶ms.name).cloned();
3385 #[cfg(feature = "dynamic-tools")]
3386 let prompt = prompt.or_else(|| {
3387 self.inner
3388 .dynamic_prompts
3389 .as_ref()
3390 .and_then(|d| d.get(¶ms.name))
3391 });
3392 let prompt = prompt.ok_or_else(|| {
3393 Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3394 "Prompt not found: {}",
3395 params.name
3396 )))
3397 })?;
3398
3399 if let Some(filter) = &self.inner.prompt_filter
3401 && !filter.is_visible(&self.session, &prompt)
3402 {
3403 return Err(filter.denial_error(¶ms.name));
3404 }
3405
3406 tracing::debug!(name = %params.name, "Getting prompt");
3407 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3408 #[cfg(feature = "stateless")]
3409 let ctx = {
3410 let mut ctx = ctx;
3411 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3412 params.input_responses,
3413 params.request_state,
3414 ));
3415 ctx
3416 };
3417 let outcome = prompt
3418 .get_outcome_with_context(ctx, params.arguments)
3419 .await?;
3420
3421 match outcome {
3422 RequestOutcome::Complete(result) => Ok(McpResponse::GetPrompt(result)),
3423 RequestOutcome::InputRequired(result) => {
3424 #[cfg(feature = "stateless")]
3425 {
3426 validate_input_required_result(&extensions, &result)?;
3427 Ok(McpResponse::InputRequired(result))
3428 }
3429 #[cfg(not(feature = "stateless"))]
3430 {
3431 let _ = result;
3432 Err(Error::invalid_params(
3433 "InputRequiredResult support was not compiled",
3434 ))
3435 }
3436 }
3437 }
3438 }
3439
3440 McpRequest::Ping => Ok(McpResponse::Pong(EmptyResult {})),
3441
3442 McpRequest::GetTaskInfo(params) => {
3443 if is_final_protocol_request(&extensions) {
3444 self.require_negotiated_tasks(&extensions, "tasks/get")?;
3445 self.authorize_task(¶ms.task_id, &extensions).await?;
3446 return self.final_get_task(¶ms.task_id).await;
3447 }
3448 self.authorize_task(¶ms.task_id, &extensions).await?;
3449
3450 let (mut task, result, error) = self
3457 .inner
3458 .task_store
3459 .get_task_result(¶ms.task_id)
3460 .await
3461 .map_err(task_store_error)?
3462 .ok_or_else(|| {
3463 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3464 "Task not found: {}",
3465 params.task_id
3466 )))
3467 })?;
3468
3469 match task.status {
3470 TaskStatus::Completed => task.result = result,
3471 TaskStatus::Failed => {
3472 task.error = Some(
3476 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
3477 );
3478 }
3479 _ => {}
3480 }
3481
3482 Ok(McpResponse::GetTaskInfo(task))
3483 }
3484
3485 McpRequest::UpdateTask(params) => {
3486 if is_final_protocol_request(&extensions) {
3487 self.require_negotiated_tasks(&extensions, "tasks/update")?;
3488 self.authorize_task(¶ms.task_id, &extensions).await?;
3489 self.inner
3493 .task_store
3494 .apply_input_responses(
3495 ¶ms.task_id,
3496 decode_input_responses(¶ms.input_responses),
3497 )
3498 .await
3499 .map_err(task_store_error)?
3500 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3501 self.notify_task_state(¶ms.task_id).await;
3505 return Ok(McpResponse::FinalTaskAck(
3506 crate::tasks::TaskAcknowledgement::new(),
3507 ));
3508 }
3509
3510 self.authorize_task(¶ms.task_id, &extensions).await?;
3511
3512 self.inner
3519 .task_store
3520 .apply_input_responses(
3521 ¶ms.task_id,
3522 decode_input_responses(¶ms.input_responses),
3523 )
3524 .await
3525 .map_err(task_store_error)?
3526 .ok_or_else(|| {
3527 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3528 "Task not found: {}",
3529 params.task_id
3530 )))
3531 })?;
3532 self.notify_task_state(¶ms.task_id).await;
3536 Ok(McpResponse::UpdateTask(EmptyResult {}))
3537 }
3538
3539 McpRequest::CancelTask(params) => {
3540 if is_final_protocol_request(&extensions) {
3541 self.require_negotiated_tasks(&extensions, "tasks/cancel")?;
3542 self.authorize_task(¶ms.task_id, &extensions).await?;
3543 self.inner
3547 .task_store
3548 .cancel_task(¶ms.task_id, params.reason.as_deref())
3549 .await
3550 .map_err(task_store_error)?
3551 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3552 self.notify_task_state(¶ms.task_id).await;
3553 return Ok(McpResponse::FinalTaskAck(
3554 crate::tasks::TaskAcknowledgement::new(),
3555 ));
3556 }
3557
3558 self.authorize_task(¶ms.task_id, &extensions).await?;
3559
3560 let current = self
3562 .inner
3563 .task_store
3564 .get_task(¶ms.task_id)
3565 .await
3566 .map_err(task_store_error)?
3567 .ok_or_else(|| {
3568 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3569 "Task not found: {}",
3570 params.task_id
3571 )))
3572 })?;
3573
3574 if current.status.is_terminal() {
3575 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
3576 "Task {} is already in terminal state: {}",
3577 params.task_id, current.status
3578 ))));
3579 }
3580
3581 self.inner
3582 .task_store
3583 .cancel_task(¶ms.task_id, params.reason.as_deref())
3584 .await
3585 .map_err(task_store_error)?
3586 .ok_or_else(|| {
3587 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3588 "Task not found: {}",
3589 params.task_id
3590 )))
3591 })?;
3592
3593 Ok(McpResponse::CancelTask(EmptyResult {}))
3597 }
3598
3599 McpRequest::SetLoggingLevel(params) => {
3600 tracing::debug!(level = ?params.level, "Client set logging level");
3601 if let Ok(mut level) = self.inner.min_log_level.write() {
3602 *level = params.level;
3603 }
3604 Ok(McpResponse::SetLoggingLevel(EmptyResult {}))
3605 }
3606
3607 McpRequest::Complete(params) => {
3608 tracing::debug!(
3609 reference = ?params.reference,
3610 argument = %params.argument.name,
3611 "Completion request"
3612 );
3613
3614 if let Some(ref handler) = self.inner.completion_handler {
3616 let result = handler(params).await?;
3617 Ok(McpResponse::Complete(result))
3618 } else {
3619 Ok(McpResponse::Complete(CompleteResult::new(vec![])))
3621 }
3622 }
3623
3624 #[cfg(feature = "stateless")]
3625 McpRequest::SubscriptionsListen(params) => {
3626 if !is_final_protocol_request(&extensions) {
3633 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
3636 "subscriptions/listen",
3637 )));
3638 }
3639 let Some(requested) = params.notifications else {
3640 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
3641 "subscriptions/listen requires a notifications filter",
3642 )));
3643 };
3644 if requested.task_ids.is_some() && !client_declares_tasks(&extensions) {
3647 return Err(Error::JsonRpc(
3648 JsonRpcError::missing_required_client_capability(
3649 tasks_client_capabilities(),
3650 ),
3651 ));
3652 }
3653 let notifications = crate::transport::subscriptions::accepted_subscription_filter(
3654 requested,
3655 self.final_tasks_enabled(),
3656 );
3657 Ok(McpResponse::SubscriptionsAccepted(
3658 crate::protocol::SubscriptionsAcceptedResult { notifications },
3659 ))
3660 }
3661
3662 McpRequest::Unknown { method, .. } => {
3663 Err(Error::JsonRpc(JsonRpcError::method_not_found(&method)))
3664 }
3665 _ => Err(Error::JsonRpc(JsonRpcError::method_not_found(
3666 "unknown method",
3667 ))),
3668 }
3669 }
3670
3671 pub fn handle_notification(&self, notification: McpNotification) {
3673 match notification {
3674 McpNotification::Initialized => {
3675 let phase_before = self.session.phase();
3676 if self.session.mark_initialized() {
3677 if phase_before == crate::session::SessionPhase::Uninitialized {
3678 tracing::info!(
3679 "Session initialized from uninitialized state (race resolved)"
3680 );
3681 } else {
3682 tracing::info!("Session initialized, entering operation phase");
3683 }
3684 } else {
3685 tracing::warn!(
3686 phase = ?self.session.phase(),
3687 "Received initialized notification in unexpected state"
3688 );
3689 }
3690 }
3691 McpNotification::Cancelled(params) => {
3692 if let Some(ref request_id) = params.request_id {
3693 if self.cancel_request(request_id) {
3694 tracing::info!(
3695 request_id = ?request_id,
3696 reason = ?params.reason,
3697 "Request cancelled"
3698 );
3699 } else {
3700 tracing::debug!(
3701 request_id = ?request_id,
3702 reason = ?params.reason,
3703 "Cancellation requested for unknown request"
3704 );
3705 }
3706 } else {
3707 tracing::debug!(
3708 reason = ?params.reason,
3709 "Cancellation notification received without request_id"
3710 );
3711 }
3712 }
3713 McpNotification::Progress(params) => {
3714 tracing::trace!(
3715 token = ?params.progress_token,
3716 progress = params.progress,
3717 total = ?params.total,
3718 "Progress notification"
3719 );
3720 }
3728 McpNotification::RootsListChanged => {
3729 tracing::info!("Client roots list changed");
3730 }
3733 McpNotification::Unknown { method, .. } => {
3734 tracing::debug!(method = %method, "Unknown notification received");
3735 }
3736 _ => {
3737 tracing::debug!("Unrecognized notification variant received");
3738 }
3739 }
3740 }
3741}
3742
3743impl Default for McpRouter {
3744 fn default() -> Self {
3745 Self::new()
3746 }
3747}
3748
3749pub use crate::context::Extensions;
3755
3756#[derive(Debug, Clone)]
3781pub struct ToolAnnotationsMap {
3782 map: Arc<HashMap<String, ToolAnnotations>>,
3783}
3784
3785impl ToolAnnotationsMap {
3786 pub fn get(&self, tool_name: &str) -> Option<&ToolAnnotations> {
3790 self.map.get(tool_name)
3791 }
3792
3793 pub fn is_read_only(&self, tool_name: &str) -> bool {
3798 self.map.get(tool_name).is_some_and(|a| a.read_only_hint)
3799 }
3800
3801 pub fn is_destructive(&self, tool_name: &str) -> bool {
3806 self.map.get(tool_name).is_none_or(|a| a.destructive_hint)
3807 }
3808
3809 pub fn is_idempotent(&self, tool_name: &str) -> bool {
3814 self.map.get(tool_name).is_some_and(|a| a.idempotent_hint)
3815 }
3816}
3817
3818#[derive(Debug, Clone)]
3840pub struct RouterRequest {
3841 pub id: RequestId,
3843 pub inner: McpRequest,
3845 pub extensions: Extensions,
3847}
3848
3849impl RouterRequest {
3850 pub fn new(id: RequestId, inner: McpRequest) -> Self {
3852 Self {
3853 id,
3854 inner,
3855 extensions: Extensions::new(),
3856 }
3857 }
3858
3859 pub fn with_inner(self, inner: McpRequest) -> Self {
3865 Self {
3866 id: self.id,
3867 inner,
3868 extensions: self.extensions,
3869 }
3870 }
3871
3872 pub fn with_id_and_inner(self, id: RequestId, inner: McpRequest) -> Self {
3878 Self {
3879 id,
3880 inner,
3881 extensions: self.extensions,
3882 }
3883 }
3884
3885 pub fn clone_with_inner(&self, inner: McpRequest) -> Self {
3893 Self {
3894 id: self.id.clone(),
3895 inner,
3896 extensions: self.extensions.clone(),
3897 }
3898 }
3899}
3900
3901#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
3903pub struct RouterResponse {
3904 pub id: RequestId,
3906 pub inner: std::result::Result<McpResponse, JsonRpcError>,
3908}
3909
3910impl RouterResponse {
3911 pub fn is_error(&self) -> bool {
3927 self.inner.is_err()
3928 }
3929
3930 pub fn into_jsonrpc(self) -> JsonRpcResponse {
3932 match self.inner {
3933 Ok(response) => match serde_json::to_value(response) {
3934 Ok(result) => JsonRpcResponse::result(self.id, result),
3935 Err(e) => {
3936 tracing::error!(error = %e, "Failed to serialize response");
3937 JsonRpcResponse::error(
3938 Some(self.id),
3939 JsonRpcError::internal_error(format!("Serialization error: {}", e)),
3940 )
3941 }
3942 },
3943 Err(error) => JsonRpcResponse::error(Some(self.id), error),
3944 }
3945 }
3946}
3947
3948impl Service<RouterRequest> for McpRouter {
3949 type Response = RouterResponse;
3950 type Error = std::convert::Infallible; type Future =
3952 Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
3953
3954 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
3955 Poll::Ready(Ok(()))
3956 }
3957
3958 fn call(&mut self, req: RouterRequest) -> Self::Future {
3959 let router = self.clone();
3960 let request_id = req.id.clone();
3961 Box::pin(async move {
3962 let result = router.handle(req.id, req.inner, req.extensions).await;
3963 router.complete_request(&request_id);
3965 Ok(RouterResponse {
3966 id: request_id,
3967 inner: result.map_err(|e| match e {
3972 Error::JsonRpc(err) => err,
3973 Error::Tool(err) => JsonRpcError::internal_error(err.to_string()),
3974 e => JsonRpcError::internal_error(e.to_string()),
3975 }),
3976 })
3977 })
3978 }
3979}
3980
3981#[cfg(test)]
3982mod tests {
3983 use super::*;
3984 use crate::extract::{Context, Json};
3985 use crate::jsonrpc::JsonRpcService;
3986 use crate::tool::ToolBuilder;
3987 use schemars::JsonSchema;
3988 use serde::Deserialize;
3989 use tower::ServiceExt;
3990
3991 #[derive(Debug, Deserialize, JsonSchema)]
3992 struct AddInput {
3993 a: i64,
3994 b: i64,
3995 }
3996
3997 #[cfg(feature = "stateless")]
3998 fn final_extensions(client_capabilities: ClientCapabilities) -> Extensions {
3999 let mut extensions = Extensions::new();
4000 extensions.insert(crate::stateless::StatelessRequestMeta {
4001 protocol_version: Some(PROTOCOL_VERSION_2026_07_28.to_string()),
4002 client_capabilities: Some(client_capabilities),
4003 ..Default::default()
4004 });
4005 extensions
4006 }
4007
4008 #[cfg(feature = "stateless")]
4009 fn tasks_client_extensions() -> Extensions {
4010 final_extensions(ClientCapabilities {
4011 extensions: Some(
4012 [(TASKS_EXTENSION_ID.to_string(), serde_json::json!({}))]
4013 .into_iter()
4014 .collect(),
4015 ),
4016 ..Default::default()
4017 })
4018 }
4019
4020 #[cfg(feature = "stateless")]
4021 #[tokio::test]
4022 async fn final_tasks_require_server_opt_in_and_client_declaration() {
4023 let tool = || {
4024 ToolBuilder::new("optional_task")
4025 .task_support(TaskSupportMode::Optional)
4026 .handler(|input: AddInput| async move {
4027 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4028 })
4029 .build()
4030 };
4031 let task_params = |task| CallToolParams {
4032 name: "optional_task".to_string(),
4033 arguments: serde_json::json!({"a": 1, "b": 2}),
4034 input_responses: None,
4035 request_state: None,
4036 meta: None,
4037 task,
4038 };
4039
4040 let implicit = McpRouter::new().tool(tool());
4044 let McpResponse::Discover(result) = implicit
4045 .handle(
4046 RequestId::Number(1),
4047 McpRequest::Discover(DiscoverParams::default()),
4048 Extensions::new(),
4049 )
4050 .await
4051 .unwrap()
4052 else {
4053 panic!("Expected Discover response");
4054 };
4055 assert!(
4056 result
4057 .capabilities
4058 .extensions
4059 .as_ref()
4060 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4061 );
4062 let error = implicit
4063 .handle(
4064 RequestId::Number(2),
4065 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4066 tasks_client_extensions(),
4067 )
4068 .await
4069 .unwrap_err();
4070 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4071
4072 let router = McpRouter::new().tool(tool()).with_tasks();
4074 let McpResponse::Discover(result) = router
4075 .handle(
4076 RequestId::Number(3),
4077 McpRequest::Discover(DiscoverParams::default()),
4078 Extensions::new(),
4079 )
4080 .await
4081 .unwrap()
4082 else {
4083 panic!("Expected Discover response");
4084 };
4085 assert!(
4086 result
4087 .capabilities
4088 .extensions
4089 .as_ref()
4090 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID)),
4091 "with_tasks() must advertise the extension on the final path"
4092 );
4093 assert!(
4094 result.capabilities.tasks.is_none(),
4095 "the legacy capability shape is never advertised on the final path"
4096 );
4097
4098 let response = router
4101 .handle(
4102 RequestId::Number(4),
4103 McpRequest::CallTool(task_params(None)),
4104 final_extensions(ClientCapabilities::default()),
4105 )
4106 .await
4107 .unwrap();
4108 assert!(matches!(response, McpResponse::CallTool(_)));
4109
4110 let response = router
4113 .handle(
4114 RequestId::Number(5),
4115 McpRequest::CallTool(task_params(None)),
4116 tasks_client_extensions(),
4117 )
4118 .await
4119 .unwrap();
4120 assert!(
4121 matches!(response, McpResponse::FinalCreateTask(_)),
4122 "a negotiated request must receive a task, got {response:?}"
4123 );
4124
4125 let error = router
4128 .handle(
4129 RequestId::Number(6),
4130 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4131 tasks_client_extensions(),
4132 )
4133 .await
4134 .unwrap_err();
4135 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4136 }
4137
4138 #[cfg(feature = "stateless")]
4139 #[tokio::test]
4140 async fn final_task_methods_serve_the_negotiated_wire_shapes() {
4141 let router = McpRouter::new()
4142 .tool(
4143 ToolBuilder::new("optional_task")
4144 .task_support(TaskSupportMode::Optional)
4145 .handler(|input: AddInput| async move {
4146 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4147 })
4148 .task_preparation(|task, _input| async move {
4149 let mut meta = serde_json::Map::new();
4150 meta.insert(
4151 "dev.tower-mcp/owner-test".to_string(),
4152 serde_json::json!({"taskId": task.task_id()}),
4153 );
4154 Ok(crate::TaskPreparation::new().with_meta(meta))
4155 })
4156 .build(),
4157 )
4158 .with_tasks();
4159
4160 let McpResponse::FinalCreateTask(created) = router
4161 .handle(
4162 RequestId::Number(1),
4163 McpRequest::CallTool(CallToolParams {
4164 name: "optional_task".to_string(),
4165 arguments: serde_json::json!({"a": 1, "b": 2}),
4166 input_responses: None,
4167 request_state: None,
4168 meta: None,
4169 task: None,
4170 }),
4171 tasks_client_extensions(),
4172 )
4173 .await
4174 .unwrap()
4175 else {
4176 panic!("Expected a final create-task response");
4177 };
4178
4179 let wire = serde_json::to_value(&created).unwrap();
4181 assert_eq!(wire["resultType"], "task");
4182 assert!(wire.get("task").is_none(), "final results are flat: {wire}");
4183 assert!(wire["ttlMs"].is_number() || wire["ttlMs"].is_null());
4184 assert!(wire.get("ttl").is_none(), "legacy field name leaked");
4185 let task_id = created.task.metadata.task_id.clone();
4186 assert_eq!(
4187 created.meta.as_ref().unwrap()["dev.tower-mcp/owner-test"]["taskId"],
4188 task_id
4189 );
4190
4191 let McpResponse::FinalGetTask(fetched) = router
4193 .handle(
4194 RequestId::Number(2),
4195 McpRequest::GetTaskInfo(GetTaskInfoParams {
4196 task_id: task_id.clone(),
4197 meta: None,
4198 }),
4199 tasks_client_extensions(),
4200 )
4201 .await
4202 .unwrap()
4203 else {
4204 panic!("Expected a final get-task response");
4205 };
4206 let wire = serde_json::to_value(&fetched).unwrap();
4207 assert_eq!(wire["resultType"], "complete");
4208 assert_eq!(wire["taskId"], serde_json::json!(task_id));
4209 assert!(wire["status"].is_string());
4210
4211 for (id, request) in [
4213 (
4214 3,
4215 McpRequest::UpdateTask(UpdateTaskParams {
4216 task_id: task_id.clone(),
4217 input_responses: HashMap::new(),
4218 meta: None,
4219 }),
4220 ),
4221 (
4222 4,
4223 McpRequest::CancelTask(CancelTaskParams {
4224 task_id: task_id.clone(),
4225 reason: None,
4226 meta: None,
4227 }),
4228 ),
4229 ] {
4230 let response = router
4231 .handle(RequestId::Number(id), request, tasks_client_extensions())
4232 .await
4233 .unwrap();
4234 let McpResponse::FinalTaskAck(ack) = response else {
4235 panic!("Expected a final ack for request {id}");
4236 };
4237 assert_eq!(
4238 serde_json::to_value(&ack).unwrap(),
4239 serde_json::json!({"resultType": "complete"})
4240 );
4241 }
4242
4243 let error = router
4245 .handle(
4246 RequestId::Number(5),
4247 McpRequest::GetTaskInfo(GetTaskInfoParams {
4248 task_id: "does-not-exist".to_string(),
4249 meta: None,
4250 }),
4251 tasks_client_extensions(),
4252 )
4253 .await
4254 .unwrap_err();
4255 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4256
4257 let error = router
4259 .handle(
4260 RequestId::Number(6),
4261 McpRequest::GetTaskInfo(GetTaskInfoParams {
4262 task_id: task_id.clone(),
4263 meta: None,
4264 }),
4265 final_extensions(ClientCapabilities::default()),
4266 )
4267 .await
4268 .unwrap_err();
4269 let Error::JsonRpc(error) = error else {
4270 panic!("expected a JSON-RPC error");
4271 };
4272 assert_eq!(error.code, -32021);
4273 assert_eq!(
4274 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4275 serde_json::json!({})
4276 );
4277 }
4278
4279 #[cfg(feature = "stateless")]
4280 #[tokio::test]
4281 async fn final_required_task_tools_follow_per_request_capabilities() {
4282 let router = McpRouter::new()
4283 .tool(
4284 ToolBuilder::new("required_task")
4285 .task_support(TaskSupportMode::Required)
4286 .handler(|input: AddInput| async move {
4287 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4288 })
4289 .build(),
4290 )
4291 .with_tasks();
4292 let params = || CallToolParams {
4293 name: "required_task".to_string(),
4294 arguments: serde_json::json!({"a": 1, "b": 2}),
4295 input_responses: None,
4296 request_state: None,
4297 meta: None,
4298 task: None,
4299 };
4300
4301 let McpResponse::ListTools(without_tasks) = router
4302 .handle(
4303 RequestId::Number(1),
4304 McpRequest::ListTools(ListToolsParams::default()),
4305 final_extensions(ClientCapabilities::default()),
4306 )
4307 .await
4308 .unwrap()
4309 else {
4310 panic!("expected tools/list")
4311 };
4312 assert!(without_tasks.tools.is_empty());
4313
4314 let McpResponse::ListTools(with_tasks) = router
4315 .handle(
4316 RequestId::Number(2),
4317 McpRequest::ListTools(ListToolsParams::default()),
4318 tasks_client_extensions(),
4319 )
4320 .await
4321 .unwrap()
4322 else {
4323 panic!("expected tools/list")
4324 };
4325 assert_eq!(with_tasks.tools.len(), 1);
4326 assert!(with_tasks.tools[0].execution.is_none());
4327
4328 let error = router
4329 .handle(
4330 RequestId::Number(3),
4331 McpRequest::CallTool(params()),
4332 final_extensions(ClientCapabilities::default()),
4333 )
4334 .await
4335 .unwrap_err();
4336 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32021));
4337
4338 let response = router
4339 .handle(
4340 RequestId::Number(4),
4341 McpRequest::CallTool(params()),
4342 tasks_client_extensions(),
4343 )
4344 .await
4345 .unwrap();
4346 assert!(matches!(response, McpResponse::FinalCreateTask(_)));
4347 }
4348
4349 #[cfg(all(feature = "oauth", feature = "stateless"))]
4350 #[tokio::test]
4351 async fn task_operations_are_bound_to_the_creating_principal() {
4352 fn as_principal(subject: &str) -> Extensions {
4353 let mut extensions = tasks_client_extensions();
4354 extensions.insert(crate::oauth::token::TokenClaims {
4355 sub: Some(subject.to_string()),
4356 iss: None,
4357 aud: None,
4358 exp: None,
4359 scope: None,
4360 client_id: None,
4361 extra: HashMap::new(),
4362 });
4363 extensions
4364 }
4365
4366 let router = McpRouter::new()
4367 .tool(
4368 ToolBuilder::new("optional_task")
4369 .task_support(TaskSupportMode::Optional)
4370 .handler(|input: AddInput| async move {
4371 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4372 })
4373 .build(),
4374 )
4375 .with_tasks();
4376
4377 let McpResponse::FinalCreateTask(created) = router
4378 .handle(
4379 RequestId::Number(1),
4380 McpRequest::CallTool(CallToolParams {
4381 name: "optional_task".to_string(),
4382 arguments: serde_json::json!({"a": 1, "b": 2}),
4383 input_responses: None,
4384 request_state: None,
4385 meta: None,
4386 task: None,
4387 }),
4388 as_principal("alice"),
4389 )
4390 .await
4391 .unwrap()
4392 else {
4393 panic!("Expected a final create-task response");
4394 };
4395 let task_id = created.task.metadata.task_id.clone();
4396
4397 assert!(
4399 router
4400 .handle(
4401 RequestId::Number(2),
4402 McpRequest::GetTaskInfo(GetTaskInfoParams {
4403 task_id: task_id.clone(),
4404 meta: None,
4405 }),
4406 as_principal("alice"),
4407 )
4408 .await
4409 .is_ok()
4410 );
4411
4412 for (id, label, context) in [
4415 (3, "another principal", as_principal("bob")),
4416 (4, "no principal", tasks_client_extensions()),
4417 ] {
4418 for (offset, request) in [
4419 McpRequest::GetTaskInfo(GetTaskInfoParams {
4420 task_id: task_id.clone(),
4421 meta: None,
4422 }),
4423 McpRequest::UpdateTask(UpdateTaskParams {
4424 task_id: task_id.clone(),
4425 input_responses: HashMap::new(),
4426 meta: None,
4427 }),
4428 McpRequest::CancelTask(CancelTaskParams {
4429 task_id: task_id.clone(),
4430 reason: None,
4431 meta: None,
4432 }),
4433 ]
4434 .into_iter()
4435 .enumerate()
4436 {
4437 let error = router
4438 .handle(
4439 RequestId::Number(id * 10 + offset as i64),
4440 request,
4441 context.clone(),
4442 )
4443 .await
4444 .unwrap_err();
4445 assert!(
4446 matches!(error, Error::JsonRpc(ref e) if e.code == -32602),
4447 "{label} was served: {error:?}"
4448 );
4449 let Error::JsonRpc(error) = error else {
4452 unreachable!()
4453 };
4454 assert!(
4455 error.message.contains("not found"),
4456 "refusal leaked that the task exists: {}",
4457 error.message
4458 );
4459 }
4460 }
4461
4462 assert!(
4464 router
4465 .handle(
4466 RequestId::Number(9),
4467 McpRequest::GetTaskInfo(GetTaskInfoParams {
4468 task_id: task_id.clone(),
4469 meta: None,
4470 }),
4471 as_principal("alice"),
4472 )
4473 .await
4474 .is_ok(),
4475 "a refused cancel must not have cancelled the task"
4476 );
4477 }
4478
4479 #[cfg(all(feature = "oauth", feature = "stateless"))]
4480 #[tokio::test]
4481 async fn final_tasks_work_across_independent_routers_with_a_shared_store() {
4482 fn as_principal(subject: &str) -> Extensions {
4483 let mut extensions = tasks_client_extensions();
4484 extensions.insert(crate::oauth::token::TokenClaims {
4485 sub: Some(subject.to_string()),
4486 iss: None,
4487 aud: None,
4488 exp: None,
4489 scope: None,
4490 client_id: None,
4491 extra: HashMap::new(),
4492 });
4493 extensions
4494 }
4495
4496 fn router_with_store(store: Arc<dyn TaskStore>) -> McpRouter {
4497 McpRouter::new()
4498 .tool(
4499 ToolBuilder::new("shared_task")
4500 .task_support(TaskSupportMode::Optional)
4501 .handler(|_input: serde_json::Value| async move {
4502 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
4503 Ok(CallToolResult::text("done"))
4504 })
4505 .build(),
4506 )
4507 .task_store(store)
4508 .with_tasks()
4509 }
4510
4511 let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
4512 let router_a = router_with_store(store.clone());
4513 let router_b = router_with_store(store);
4514
4515 let McpResponse::FinalCreateTask(created) = router_a
4516 .handle(
4517 RequestId::Number(1),
4518 McpRequest::CallTool(CallToolParams {
4519 name: "shared_task".to_string(),
4520 arguments: serde_json::json!({}),
4521 input_responses: None,
4522 request_state: None,
4523 meta: None,
4524 task: None,
4525 }),
4526 as_principal("alice"),
4527 )
4528 .await
4529 .unwrap()
4530 else {
4531 panic!("router A did not create a final task")
4532 };
4533 let task_id = created.task.metadata.task_id;
4534
4535 assert!(
4537 router_b
4538 .handle(
4539 RequestId::Number(2),
4540 McpRequest::GetTaskInfo(GetTaskInfoParams {
4541 task_id: task_id.clone(),
4542 meta: None,
4543 }),
4544 as_principal("alice"),
4545 )
4546 .await
4547 .is_ok()
4548 );
4549
4550 let denied = router_b
4552 .handle(
4553 RequestId::Number(3),
4554 McpRequest::GetTaskInfo(GetTaskInfoParams {
4555 task_id: task_id.clone(),
4556 meta: None,
4557 }),
4558 as_principal("bob"),
4559 )
4560 .await
4561 .unwrap_err();
4562 let unknown = router_b
4563 .handle(
4564 RequestId::Number(4),
4565 McpRequest::GetTaskInfo(GetTaskInfoParams {
4566 task_id: "unknown-task".to_string(),
4567 meta: None,
4568 }),
4569 as_principal("bob"),
4570 )
4571 .await
4572 .unwrap_err();
4573 let (Error::JsonRpc(denied), Error::JsonRpc(unknown)) = (denied, unknown) else {
4574 panic!("expected JSON-RPC task denials")
4575 };
4576 assert_eq!(denied.code, unknown.code);
4577 assert_eq!(
4578 denied.message.replace(&task_id, "<task-id>"),
4579 unknown.message.replace("unknown-task", "<task-id>")
4580 );
4581 assert_eq!(denied.data, unknown.data);
4582
4583 assert!(matches!(
4586 router_b
4587 .handle(
4588 RequestId::Number(5),
4589 McpRequest::CancelTask(CancelTaskParams {
4590 task_id: task_id.clone(),
4591 reason: None,
4592 meta: None,
4593 }),
4594 as_principal("alice"),
4595 )
4596 .await
4597 .unwrap(),
4598 McpResponse::FinalTaskAck(_)
4599 ));
4600 let McpResponse::FinalGetTask(fetched) = router_a
4601 .handle(
4602 RequestId::Number(6),
4603 McpRequest::GetTaskInfo(GetTaskInfoParams {
4604 task_id,
4605 meta: None,
4606 }),
4607 as_principal("alice"),
4608 )
4609 .await
4610 .unwrap()
4611 else {
4612 panic!("router A did not read the shared task")
4613 };
4614 assert_eq!(fetched.task.status(), TaskStatus::Cancelled);
4615 }
4616
4617 #[test]
4618 fn router_advertises_only_locally_declared_protocol_extensions() {
4619 let router = McpRouter::new().with_protocol_extension(
4620 crate::ExtensionDeclaration::new(
4621 "com.example/rendering",
4622 serde_json::json!({"formats": ["html"]}),
4623 )
4624 .unwrap(),
4625 );
4626
4627 let stable = router.capabilities();
4628 let final_capabilities =
4629 router.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
4630 for capabilities in [stable, final_capabilities] {
4631 let extensions = capabilities.extensions.unwrap();
4632 assert_eq!(extensions.len(), 1);
4633 assert_eq!(extensions["com.example/rendering"]["formats"][0], "html");
4634 assert!(!extensions.contains_key("com.example/client-only"));
4635 }
4636 }
4637
4638 #[tokio::test]
4639 async fn initialize_persists_negotiated_extensions_for_legacy_contexts() {
4640 let router = McpRouter::new().with_protocol_extension(
4641 crate::ExtensionDeclaration::new(
4642 "com.example/shared",
4643 serde_json::json!({"server": true}),
4644 )
4645 .unwrap(),
4646 );
4647 let client_capabilities = ClientCapabilities {
4648 extensions: Some(HashMap::from([
4649 (
4650 "com.example/shared".to_string(),
4651 serde_json::json!({"client": true}),
4652 ),
4653 ("com.example/client-only".to_string(), serde_json::json!({})),
4654 ])),
4655 ..ClientCapabilities::default()
4656 };
4657
4658 router
4659 .handle(
4660 RequestId::Number(1),
4661 McpRequest::Initialize(InitializeParams {
4662 protocol_version: crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
4663 capabilities: client_capabilities,
4664 client_info: Implementation {
4665 name: "extension-test".to_string(),
4666 version: "1.0.0".to_string(),
4667 title: None,
4668 description: None,
4669 icons: None,
4670 website_url: None,
4671 meta: None,
4672 },
4673 meta: None,
4674 }),
4675 Extensions::new(),
4676 )
4677 .await
4678 .unwrap();
4679
4680 let context = router.create_context(RequestId::Number(2), None);
4681 let negotiated = context.negotiated_extensions().unwrap();
4682 assert!(negotiated.contains("com.example/shared"));
4683 assert!(!negotiated.contains("com.example/client-only"));
4684 }
4685
4686 #[cfg(feature = "stateless")]
4687 #[test]
4688 fn final_request_context_exposes_only_negotiated_extensions() {
4689 let router = McpRouter::new().with_protocol_extension(
4690 crate::ExtensionDeclaration::new(
4691 "com.example/shared",
4692 serde_json::json!({"server": true}),
4693 )
4694 .unwrap(),
4695 );
4696 let per_request = final_extensions(ClientCapabilities {
4697 extensions: Some(HashMap::from([
4698 (
4699 "com.example/shared".to_string(),
4700 serde_json::json!({"client": true}),
4701 ),
4702 ("com.example/client-only".to_string(), serde_json::json!({})),
4703 ])),
4704 ..ClientCapabilities::default()
4705 });
4706
4707 let context =
4708 router.create_context_with_extensions(RequestId::Number(1), None, &per_request);
4709 let negotiated = context.negotiated_extensions().unwrap();
4710
4711 assert_eq!(negotiated.len(), 1);
4712 assert_eq!(
4713 negotiated
4714 .get("com.example/shared")
4715 .unwrap()
4716 .client_settings()["client"],
4717 true
4718 );
4719 assert!(!negotiated.contains("com.example/client-only"));
4720 }
4721
4722 #[cfg(feature = "stateless")]
4723 #[tokio::test]
4724 async fn final_protocol_withholds_incomplete_tasks_advertisement() {
4725 let optional = ToolBuilder::new("optional_task")
4726 .task_support(TaskSupportMode::Optional)
4727 .handler(|input: AddInput| async move {
4728 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4729 })
4730 .build();
4731 let required = ToolBuilder::new("required_task")
4732 .task_support(TaskSupportMode::Required)
4733 .handler(|input: AddInput| async move {
4734 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4735 })
4736 .build();
4737 let mut router = McpRouter::new().tool(optional).tool(required);
4738
4739 let stable_capabilities = router.capabilities();
4741 assert!(stable_capabilities.tasks.is_some());
4742 assert!(
4743 stable_capabilities
4744 .extensions
4745 .as_ref()
4746 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID))
4747 );
4748
4749 let response = router
4751 .handle(
4752 RequestId::Number(1),
4753 McpRequest::Discover(DiscoverParams::default()),
4754 Extensions::new(),
4755 )
4756 .await
4757 .unwrap();
4758 let McpResponse::Discover(result) = response else {
4759 panic!("Expected Discover response");
4760 };
4761 assert!(result.capabilities.tasks.is_none());
4762 assert!(
4763 result
4764 .capabilities
4765 .extensions
4766 .as_ref()
4767 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4768 );
4769
4770 init_router(&mut router).await;
4771
4772 let response = router
4774 .handle(
4775 RequestId::Number(2),
4776 McpRequest::ListTools(ListToolsParams::default()),
4777 Extensions::new(),
4778 )
4779 .await
4780 .unwrap();
4781 let McpResponse::ListTools(result) = response else {
4782 panic!("Expected ListTools response");
4783 };
4784 assert_eq!(result.tools.len(), 2);
4785 assert!(result.tools.iter().all(|tool| tool.execution.is_some()));
4786
4787 let response = router
4790 .handle(
4791 RequestId::Number(3),
4792 McpRequest::ListTools(ListToolsParams::default()),
4793 final_extensions(ClientCapabilities::default()),
4794 )
4795 .await
4796 .unwrap();
4797 let McpResponse::ListTools(result) = response else {
4798 panic!("Expected ListTools response");
4799 };
4800 assert_eq!(result.tools.len(), 1);
4801 assert_eq!(result.tools[0].name, "optional_task");
4802 assert!(result.tools[0].execution.is_none());
4803 }
4804
4805 #[cfg(feature = "stateless")]
4806 #[tokio::test]
4807 async fn final_protocol_enforces_tasks_negotiation() {
4808 let optional = ToolBuilder::new("optional_task")
4809 .task_support(TaskSupportMode::Optional)
4810 .handler(|input: AddInput| async move {
4811 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4812 })
4813 .build();
4814 let required = ToolBuilder::new("required_task")
4815 .task_support(TaskSupportMode::Required)
4816 .handler(|input: AddInput| async move {
4817 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4818 })
4819 .build();
4820 let mut router = McpRouter::new().tool(optional).tool(required).with_tasks();
4821 init_router(&mut router).await;
4822
4823 let response = router
4825 .handle(
4826 RequestId::Number(1),
4827 McpRequest::CallTool(CallToolParams {
4828 name: "optional_task".to_string(),
4829 arguments: serde_json::json!({"a": 1, "b": 2}),
4830 input_responses: None,
4831 request_state: None,
4832 meta: None,
4833 task: None,
4834 }),
4835 final_extensions(ClientCapabilities::default()),
4836 )
4837 .await
4838 .unwrap();
4839 assert!(matches!(response, McpResponse::CallTool(_)));
4840
4841 let error = router
4843 .handle(
4844 RequestId::Number(2),
4845 McpRequest::CallTool(CallToolParams {
4846 name: "optional_task".to_string(),
4847 arguments: serde_json::json!({"a": 1, "b": 2}),
4848 input_responses: None,
4849 request_state: None,
4850 meta: None,
4851 task: Some(TaskRequestParams { ttl: None }),
4852 }),
4853 final_extensions(ClientCapabilities::default()),
4854 )
4855 .await
4856 .unwrap_err();
4857 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32602));
4858
4859 let error = router
4863 .handle(
4864 RequestId::Number(3),
4865 McpRequest::CallTool(CallToolParams {
4866 name: "required_task".to_string(),
4867 arguments: serde_json::json!({"a": 1, "b": 2}),
4868 input_responses: None,
4869 request_state: None,
4870 meta: None,
4871 task: None,
4872 }),
4873 final_extensions(ClientCapabilities::default()),
4874 )
4875 .await
4876 .unwrap_err();
4877 let Error::JsonRpc(error) = error else {
4878 panic!("expected a JSON-RPC error");
4879 };
4880 assert_eq!(error.code, -32021);
4881 assert_eq!(
4882 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4883 serde_json::json!({}),
4884 "the error must name the extension the client needs to declare"
4885 );
4886
4887 let task_requests = [
4888 McpRequest::GetTaskInfo(GetTaskInfoParams {
4889 task_id: "task-unknown".to_string(),
4890 meta: None,
4891 }),
4892 McpRequest::UpdateTask(UpdateTaskParams {
4893 task_id: "task-unknown".to_string(),
4894 input_responses: HashMap::new(),
4895 meta: None,
4896 }),
4897 McpRequest::CancelTask(CancelTaskParams {
4898 task_id: "task-unknown".to_string(),
4899 reason: None,
4900 meta: None,
4901 }),
4902 ];
4903 for (index, request) in task_requests.into_iter().enumerate() {
4904 let error = router
4905 .handle(
4906 RequestId::Number(4 + index as i64),
4907 request,
4908 final_extensions(ClientCapabilities::default()),
4909 )
4910 .await
4911 .unwrap_err();
4912 let Error::JsonRpc(error) = error else {
4913 panic!("expected a JSON-RPC error");
4914 };
4915 assert_eq!(error.code, -32021);
4916 assert_eq!(
4917 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4918 serde_json::json!({})
4919 );
4920 }
4921
4922 let router_without_tasks = McpRouter::new();
4925 let error = router_without_tasks
4926 .handle(
4927 RequestId::Number(7),
4928 McpRequest::GetTaskInfo(GetTaskInfoParams {
4929 task_id: "task-unknown".to_string(),
4930 meta: None,
4931 }),
4932 final_extensions(tasks_client_capabilities()),
4933 )
4934 .await
4935 .unwrap_err();
4936 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32601));
4937 }
4938
4939 #[cfg(feature = "stateless")]
4940 #[test]
4941 fn input_required_capability_validation_uses_capability_semantics() {
4942 let roots = InputRequiredResult::with_requests(
4943 [(
4944 "roots".to_string(),
4945 InputRequest::ListRoots(ListRootsParams::default()),
4946 )]
4947 .into_iter()
4948 .collect(),
4949 );
4950 let extensions = final_extensions(ClientCapabilities {
4951 roots: Some(RootsCapability {
4952 list_changed: true,
4953 deprecated: None,
4954 }),
4955 ..Default::default()
4956 });
4957 validate_input_required_result(&extensions, &roots).unwrap();
4958 assert!(client_capabilities_satisfy(
4959 extensions
4960 .get::<crate::stateless::StatelessRequestMeta>()
4961 .and_then(|meta| meta.client_capabilities.as_ref())
4962 .unwrap(),
4963 &ClientCapabilities {
4964 roots: Some(RootsCapability::default()),
4965 ..Default::default()
4966 }
4967 ));
4968
4969 let sampling_with_tools = InputRequiredResult::with_requests(
4970 [(
4971 "sample".to_string(),
4972 InputRequest::CreateMessage(CreateMessageParams {
4973 tools: Some(Vec::new()),
4974 ..CreateMessageParams::new(vec![SamplingMessage::user("hello")], 10)
4975 }),
4976 )]
4977 .into_iter()
4978 .collect(),
4979 );
4980 let extensions = final_extensions(ClientCapabilities {
4981 sampling: Some(SamplingCapability::default()),
4982 ..Default::default()
4983 });
4984 assert!(validate_input_required_result(&extensions, &sampling_with_tools).is_err());
4985
4986 let form = InputRequiredResult::with_requests(
4987 [(
4988 "form".to_string(),
4989 InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
4990 mode: Some(ElicitMode::Form),
4991 message: "name".into(),
4992 requested_schema: ElicitFormSchema::new(),
4993 meta: None,
4994 })),
4995 )]
4996 .into_iter()
4997 .collect(),
4998 );
4999 let extensions = final_extensions(ClientCapabilities {
5000 elicitation: Some(ElicitationCapability::default()),
5001 ..Default::default()
5002 });
5003 validate_input_required_result(&extensions, &form).unwrap();
5004 }
5005
5006 async fn init_router(router: &mut McpRouter) {
5008 let init_req = RouterRequest {
5010 id: RequestId::Number(0),
5011 inner: McpRequest::Initialize(InitializeParams {
5012 protocol_version: "2025-11-25".to_string(),
5013 capabilities: ClientCapabilities {
5014 roots: None,
5015 sampling: None,
5016 elicitation: None,
5017 tasks: None,
5018 experimental: None,
5019 extensions: None,
5020 },
5021 client_info: Implementation {
5022 name: "test".to_string(),
5023 version: "1.0".to_string(),
5024 ..Default::default()
5025 },
5026 meta: None,
5027 }),
5028 extensions: Extensions::new(),
5029 };
5030 let _ = router.ready().await.unwrap().call(init_req).await.unwrap();
5031 router.handle_notification(McpNotification::Initialized);
5033 }
5034
5035 #[tokio::test]
5036 async fn test_router_list_tools() {
5037 let add_tool = ToolBuilder::new("add")
5038 .description("Add two numbers")
5039 .handler(|input: AddInput| async move {
5040 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5041 })
5042 .build();
5043
5044 let mut router = McpRouter::new().tool(add_tool);
5045
5046 init_router(&mut router).await;
5048
5049 let req = RouterRequest {
5050 id: RequestId::Number(1),
5051 inner: McpRequest::ListTools(ListToolsParams::default()),
5052 extensions: Extensions::new(),
5053 };
5054
5055 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5056
5057 match resp.inner {
5058 Ok(McpResponse::ListTools(result)) => {
5059 assert_eq!(result.tools.len(), 1);
5060 assert_eq!(result.tools[0].name, "add");
5061 }
5062 _ => panic!("Expected ListTools response"),
5063 }
5064 }
5065
5066 #[tokio::test]
5067 async fn test_router_call_tool() {
5068 let add_tool = ToolBuilder::new("add")
5069 .description("Add two numbers")
5070 .handler(|input: AddInput| async move {
5071 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5072 })
5073 .build();
5074
5075 let mut router = McpRouter::new().tool(add_tool);
5076
5077 init_router(&mut router).await;
5079
5080 let req = RouterRequest {
5081 id: RequestId::Number(1),
5082 inner: McpRequest::CallTool(CallToolParams {
5083 input_responses: None,
5084 request_state: None,
5085 name: "add".to_string(),
5086 arguments: serde_json::json!({"a": 2, "b": 3}),
5087 meta: None,
5088 task: None,
5089 }),
5090 extensions: Extensions::new(),
5091 };
5092
5093 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5094
5095 match resp.inner {
5096 Ok(McpResponse::CallTool(result)) => {
5097 assert!(!result.is_error);
5098 match &result.content[0] {
5100 Content::Text { text, .. } => assert_eq!(text, "5"),
5101 _ => panic!("Expected text content"),
5102 }
5103 }
5104 _ => panic!("Expected CallTool response"),
5105 }
5106 }
5107
5108 async fn init_jsonrpc_service(service: &mut JsonRpcService<McpRouter>, router: &McpRouter) {
5110 let init_req = JsonRpcRequest::new(0, "initialize").with_params(serde_json::json!({
5111 "protocolVersion": "2025-11-25",
5112 "capabilities": {},
5113 "clientInfo": { "name": "test", "version": "1.0" }
5114 }));
5115 let _ = service.call_single(init_req).await.unwrap();
5116 router.handle_notification(McpNotification::Initialized);
5117 }
5118
5119 #[tokio::test]
5120 async fn test_jsonrpc_service() {
5121 let add_tool = ToolBuilder::new("add")
5122 .description("Add two numbers")
5123 .handler(|input: AddInput| async move {
5124 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5125 })
5126 .build();
5127
5128 let router = McpRouter::new().tool(add_tool);
5129 let mut service = JsonRpcService::new(router.clone());
5130
5131 init_jsonrpc_service(&mut service, &router).await;
5133
5134 let req = JsonRpcRequest::new(1, "tools/list");
5135
5136 let resp = service.call_single(req).await.unwrap();
5137
5138 match resp {
5139 JsonRpcResponse::Result(r) => {
5140 assert_eq!(r.id, RequestId::Number(1));
5141 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5142 assert_eq!(tools.len(), 1);
5143 }
5144 JsonRpcResponse::Error(_) => panic!("Expected success response"),
5145 _ => panic!("unexpected response variant"),
5146 }
5147 }
5148
5149 #[tokio::test]
5150 async fn test_batch_request() {
5151 let add_tool = ToolBuilder::new("add")
5152 .description("Add two numbers")
5153 .handler(|input: AddInput| async move {
5154 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5155 })
5156 .build();
5157
5158 let router = McpRouter::new().tool(add_tool);
5159 let mut service = JsonRpcService::new(router.clone())
5160 .protocol_versions(["2025-03-26"])
5161 .unwrap();
5162
5163 init_jsonrpc_service(&mut service, &router).await;
5165
5166 let requests = vec![
5168 JsonRpcRequest::new(1, "tools/list"),
5169 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5170 "name": "add",
5171 "arguments": {"a": 10, "b": 20}
5172 })),
5173 JsonRpcRequest::new(3, "ping"),
5174 ];
5175
5176 let responses = service.call_batch(requests).await.unwrap();
5177
5178 assert_eq!(responses.len(), 3);
5179
5180 match &responses[0] {
5182 JsonRpcResponse::Result(r) => {
5183 assert_eq!(r.id, RequestId::Number(1));
5184 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5185 assert_eq!(tools.len(), 1);
5186 }
5187 JsonRpcResponse::Error(_) => panic!("Expected success for tools/list"),
5188 _ => panic!("unexpected response variant"),
5189 }
5190
5191 match &responses[1] {
5193 JsonRpcResponse::Result(r) => {
5194 assert_eq!(r.id, RequestId::Number(2));
5195 let content = r.result.get("content").unwrap().as_array().unwrap();
5196 let text = content[0].get("text").unwrap().as_str().unwrap();
5197 assert_eq!(text, "30");
5198 }
5199 JsonRpcResponse::Error(_) => panic!("Expected success for tools/call"),
5200 _ => panic!("unexpected response variant"),
5201 }
5202
5203 match &responses[2] {
5205 JsonRpcResponse::Result(r) => {
5206 assert_eq!(r.id, RequestId::Number(3));
5207 }
5208 JsonRpcResponse::Error(_) => panic!("Expected success for ping"),
5209 _ => panic!("unexpected response variant"),
5210 }
5211 }
5212
5213 #[tokio::test]
5214 async fn test_empty_batch_error() {
5215 let router = McpRouter::new();
5216 let mut service = JsonRpcService::new(router);
5217
5218 let result = service.call_batch(vec![]).await;
5219 assert!(result.is_err());
5220 }
5221
5222 #[tokio::test]
5227 async fn test_progress_token_extraction() {
5228 use crate::context::{ServerNotification, notification_channel};
5229 use crate::protocol::ProgressToken;
5230 use std::sync::Arc;
5231 use std::sync::atomic::{AtomicBool, Ordering};
5232
5233 let progress_reported = Arc::new(AtomicBool::new(false));
5235 let progress_ref = progress_reported.clone();
5236
5237 let tool = ToolBuilder::new("progress_tool")
5239 .description("Tool that reports progress")
5240 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5241 let reported = progress_ref.clone();
5242 async move {
5243 ctx.report_progress(50.0, Some(100.0), Some("Halfway"))
5245 .await;
5246 reported.store(true, Ordering::SeqCst);
5247 Ok(CallToolResult::text("done"))
5248 }
5249 })
5250 .build();
5251
5252 let (tx, mut rx) = notification_channel(10);
5254 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5255 let mut service = JsonRpcService::new(router.clone());
5256
5257 init_jsonrpc_service(&mut service, &router).await;
5259
5260 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5262 "name": "progress_tool",
5263 "arguments": {"a": 1, "b": 2},
5264 "_meta": {
5265 "progressToken": "test-token-123"
5266 }
5267 }));
5268
5269 let resp = service.call_single(req).await.unwrap();
5270
5271 match resp {
5273 JsonRpcResponse::Result(_) => {}
5274 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
5275 _ => panic!("unexpected response variant"),
5276 }
5277
5278 assert!(progress_reported.load(Ordering::SeqCst));
5280
5281 let notification = rx.try_recv().expect("Expected progress notification");
5283 match notification {
5284 ServerNotification::Progress(params) => {
5285 assert_eq!(
5286 params.progress_token,
5287 ProgressToken::String("test-token-123".to_string())
5288 );
5289 assert_eq!(params.progress, 50.0);
5290 assert_eq!(params.total, Some(100.0));
5291 assert_eq!(params.message.as_deref(), Some("Halfway"));
5292 }
5293 _ => panic!("Expected Progress notification"),
5294 }
5295 }
5296
5297 #[tokio::test]
5298 async fn test_tool_call_without_progress_token() {
5299 use crate::context::notification_channel;
5300 use std::sync::Arc;
5301 use std::sync::atomic::{AtomicBool, Ordering};
5302
5303 let progress_attempted = Arc::new(AtomicBool::new(false));
5304 let progress_ref = progress_attempted.clone();
5305
5306 let tool = ToolBuilder::new("no_token_tool")
5307 .description("Tool that tries to report progress without token")
5308 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5309 let attempted = progress_ref.clone();
5310 async move {
5311 ctx.report_progress(50.0, Some(100.0), None).await;
5313 attempted.store(true, Ordering::SeqCst);
5314 Ok(CallToolResult::text("done"))
5315 }
5316 })
5317 .build();
5318
5319 let (tx, mut rx) = notification_channel(10);
5320 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5321 let mut service = JsonRpcService::new(router.clone());
5322
5323 init_jsonrpc_service(&mut service, &router).await;
5324
5325 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5327 "name": "no_token_tool",
5328 "arguments": {"a": 1, "b": 2}
5329 }));
5330
5331 let resp = service.call_single(req).await.unwrap();
5332 assert!(matches!(resp, JsonRpcResponse::Result(_)));
5333
5334 assert!(progress_attempted.load(Ordering::SeqCst));
5336
5337 assert!(rx.try_recv().is_err());
5339 }
5340
5341 #[tokio::test]
5342 async fn test_batch_errors_returned_not_dropped() {
5343 let add_tool = ToolBuilder::new("add")
5344 .description("Add two numbers")
5345 .handler(|input: AddInput| async move {
5346 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5347 })
5348 .build();
5349
5350 let router = McpRouter::new().tool(add_tool);
5351 let mut service = JsonRpcService::new(router.clone())
5352 .protocol_versions(["2025-03-26"])
5353 .unwrap();
5354
5355 init_jsonrpc_service(&mut service, &router).await;
5356
5357 let requests = vec![
5359 JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5361 "name": "add",
5362 "arguments": {"a": 10, "b": 20}
5363 })),
5364 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5366 "name": "nonexistent_tool",
5367 "arguments": {}
5368 })),
5369 JsonRpcRequest::new(3, "ping"),
5371 ];
5372
5373 let responses = service.call_batch(requests).await.unwrap();
5374
5375 assert_eq!(responses.len(), 3);
5377
5378 match &responses[0] {
5380 JsonRpcResponse::Result(r) => {
5381 assert_eq!(r.id, RequestId::Number(1));
5382 }
5383 JsonRpcResponse::Error(_) => panic!("Expected success for first request"),
5384 _ => panic!("unexpected response variant"),
5385 }
5386
5387 match &responses[1] {
5389 JsonRpcResponse::Error(e) => {
5390 assert_eq!(e.id, Some(RequestId::Number(2)));
5391 assert!(e.error.message.contains("not found") || e.error.code == -32601);
5393 }
5394 JsonRpcResponse::Result(_) => panic!("Expected error for second request"),
5395 _ => panic!("unexpected response variant"),
5396 }
5397
5398 match &responses[2] {
5400 JsonRpcResponse::Result(r) => {
5401 assert_eq!(r.id, RequestId::Number(3));
5402 }
5403 JsonRpcResponse::Error(_) => panic!("Expected success for third request"),
5404 _ => panic!("unexpected response variant"),
5405 }
5406 }
5407
5408 #[tokio::test]
5413 async fn test_list_resource_templates() {
5414 use crate::resource::ResourceTemplateBuilder;
5415 use std::collections::HashMap;
5416
5417 let template = ResourceTemplateBuilder::new("file:///{path}")
5418 .name("Project Files")
5419 .description("Access project files")
5420 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5421 Ok(ReadResourceResult {
5422 contents: vec![ResourceContent {
5423 uri,
5424 mime_type: None,
5425 text: None,
5426 blob: None,
5427 meta: None,
5428 }],
5429 meta: None,
5430 ..Default::default()
5431 })
5432 });
5433
5434 let mut router = McpRouter::new().resource_template(template);
5435
5436 init_router(&mut router).await;
5438
5439 let req = RouterRequest {
5440 id: RequestId::Number(1),
5441 inner: McpRequest::ListResourceTemplates(ListResourceTemplatesParams::default()),
5442 extensions: Extensions::new(),
5443 };
5444
5445 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5446
5447 match resp.inner {
5448 Ok(McpResponse::ListResourceTemplates(result)) => {
5449 assert_eq!(result.resource_templates.len(), 1);
5450 assert_eq!(result.resource_templates[0].uri_template, "file:///{path}");
5451 assert_eq!(result.resource_templates[0].name, "Project Files");
5452 }
5453 _ => panic!("Expected ListResourceTemplates response"),
5454 }
5455 }
5456
5457 #[tokio::test]
5458 async fn test_read_resource_via_template() {
5459 use crate::resource::ResourceTemplateBuilder;
5460 use std::collections::HashMap;
5461
5462 let template = ResourceTemplateBuilder::new("db://users/{id}")
5463 .name("User Records")
5464 .handler(|uri: String, vars: HashMap<String, String>| async move {
5465 let id = vars.get("id").unwrap().clone();
5466 Ok(ReadResourceResult {
5467 contents: vec![ResourceContent {
5468 uri,
5469 mime_type: Some("application/json".to_string()),
5470 text: Some(format!(r#"{{"id": "{}"}}"#, id)),
5471 blob: None,
5472 meta: None,
5473 }],
5474 meta: None,
5475 ..Default::default()
5476 })
5477 });
5478
5479 let mut router = McpRouter::new().resource_template(template);
5480
5481 init_router(&mut router).await;
5483
5484 let req = RouterRequest {
5486 id: RequestId::Number(1),
5487 inner: McpRequest::ReadResource(ReadResourceParams {
5488 input_responses: None,
5489 request_state: None,
5490 uri: "db://users/123".to_string(),
5491 meta: None,
5492 }),
5493 extensions: Extensions::new(),
5494 };
5495
5496 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5497
5498 match resp.inner {
5499 Ok(McpResponse::ReadResource(result)) => {
5500 assert_eq!(result.contents.len(), 1);
5501 assert_eq!(result.contents[0].uri, "db://users/123");
5502 assert!(result.contents[0].text.as_ref().unwrap().contains("123"));
5503 }
5504 _ => panic!("Expected ReadResource response"),
5505 }
5506 }
5507
5508 #[tokio::test]
5509 async fn test_static_resource_takes_precedence_over_template() {
5510 use crate::resource::{ResourceBuilder, ResourceTemplateBuilder};
5511 use std::collections::HashMap;
5512
5513 let template = ResourceTemplateBuilder::new("file:///{path}")
5515 .name("Files Template")
5516 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5517 Ok(ReadResourceResult {
5518 contents: vec![ResourceContent {
5519 uri,
5520 mime_type: None,
5521 text: Some("from template".to_string()),
5522 blob: None,
5523 meta: None,
5524 }],
5525 meta: None,
5526 ..Default::default()
5527 })
5528 });
5529
5530 let static_resource = ResourceBuilder::new("file:///README.md")
5532 .name("README")
5533 .text("from static resource");
5534
5535 let mut router = McpRouter::new()
5536 .resource_template(template)
5537 .resource(static_resource);
5538
5539 init_router(&mut router).await;
5541
5542 let req = RouterRequest {
5544 id: RequestId::Number(1),
5545 inner: McpRequest::ReadResource(ReadResourceParams {
5546 input_responses: None,
5547 request_state: None,
5548 uri: "file:///README.md".to_string(),
5549 meta: None,
5550 }),
5551 extensions: Extensions::new(),
5552 };
5553
5554 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5555
5556 match resp.inner {
5557 Ok(McpResponse::ReadResource(result)) => {
5558 assert_eq!(
5560 result.contents[0].text.as_deref(),
5561 Some("from static resource")
5562 );
5563 }
5564 _ => panic!("Expected ReadResource response"),
5565 }
5566 }
5567
5568 #[tokio::test]
5569 async fn test_resource_not_found_when_no_match() {
5570 use crate::resource::ResourceTemplateBuilder;
5571 use std::collections::HashMap;
5572
5573 let template = ResourceTemplateBuilder::new("db://users/{id}")
5574 .name("Users")
5575 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5576 Ok(ReadResourceResult {
5577 contents: vec![ResourceContent {
5578 uri,
5579 mime_type: None,
5580 text: None,
5581 blob: None,
5582 meta: None,
5583 }],
5584 meta: None,
5585 ..Default::default()
5586 })
5587 });
5588
5589 let mut router = McpRouter::new().resource_template(template);
5590
5591 init_router(&mut router).await;
5593
5594 let req = RouterRequest {
5596 id: RequestId::Number(1),
5597 inner: McpRequest::ReadResource(ReadResourceParams {
5598 input_responses: None,
5599 request_state: None,
5600 uri: "db://posts/123".to_string(),
5601 meta: None,
5602 }),
5603 extensions: Extensions::new(),
5604 };
5605
5606 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5607
5608 match resp.inner {
5609 Err(err) => {
5610 assert!(err.message.contains("not found"));
5611 }
5612 Ok(_) => panic!("Expected error for non-matching URI"),
5613 }
5614 }
5615
5616 #[tokio::test]
5617 async fn test_capabilities_include_resources_with_only_templates() {
5618 use crate::resource::ResourceTemplateBuilder;
5619 use std::collections::HashMap;
5620
5621 let template = ResourceTemplateBuilder::new("file:///{path}")
5622 .name("Files")
5623 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5624 Ok(ReadResourceResult {
5625 contents: vec![ResourceContent {
5626 uri,
5627 mime_type: None,
5628 text: None,
5629 blob: None,
5630 meta: None,
5631 }],
5632 meta: None,
5633 ..Default::default()
5634 })
5635 });
5636
5637 let mut router = McpRouter::new().resource_template(template);
5638
5639 let init_req = RouterRequest {
5641 id: RequestId::Number(0),
5642 inner: McpRequest::Initialize(InitializeParams {
5643 protocol_version: "2025-11-25".to_string(),
5644 capabilities: ClientCapabilities {
5645 roots: None,
5646 sampling: None,
5647 elicitation: None,
5648 tasks: None,
5649 experimental: None,
5650 extensions: None,
5651 },
5652 client_info: Implementation {
5653 name: "test".to_string(),
5654 version: "1.0".to_string(),
5655 ..Default::default()
5656 },
5657 meta: None,
5658 }),
5659 extensions: Extensions::new(),
5660 };
5661 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5662
5663 match resp.inner {
5664 Ok(McpResponse::Initialize(result)) => {
5665 assert!(result.capabilities.resources.is_some());
5667 }
5668 _ => panic!("Expected Initialize response"),
5669 }
5670 }
5671
5672 #[tokio::test]
5677 async fn test_log_sends_notification() {
5678 use crate::context::notification_channel;
5679
5680 let (tx, mut rx) = notification_channel(10);
5681 let router = McpRouter::new().with_notification_sender(tx);
5682
5683 let sent = router.log_info("Test message");
5685 assert!(sent);
5686
5687 let notification = rx.try_recv().unwrap();
5689 match notification {
5690 ServerNotification::LogMessage(params) => {
5691 assert_eq!(params.level, LogLevel::Info);
5692 let data = params.data;
5693 assert_eq!(
5694 data.get("message").unwrap().as_str().unwrap(),
5695 "Test message"
5696 );
5697 }
5698 _ => panic!("Expected LogMessage notification"),
5699 }
5700 }
5701
5702 #[tokio::test]
5703 async fn test_log_with_custom_params() {
5704 use crate::context::notification_channel;
5705
5706 let (tx, mut rx) = notification_channel(10);
5707 let router = McpRouter::new().with_notification_sender(tx);
5708
5709 let params = LoggingMessageParams::new(
5711 LogLevel::Error,
5712 serde_json::json!({
5713 "error": "Connection failed",
5714 "host": "localhost"
5715 }),
5716 )
5717 .with_logger("database");
5718
5719 let sent = router.log(params);
5720 assert!(sent);
5721
5722 let notification = rx.try_recv().unwrap();
5723 match notification {
5724 ServerNotification::LogMessage(params) => {
5725 assert_eq!(params.level, LogLevel::Error);
5726 assert_eq!(params.logger.as_deref(), Some("database"));
5727 let data = params.data;
5728 assert_eq!(
5729 data.get("error").unwrap().as_str().unwrap(),
5730 "Connection failed"
5731 );
5732 }
5733 _ => panic!("Expected LogMessage notification"),
5734 }
5735 }
5736
5737 #[tokio::test]
5738 async fn test_log_without_channel_returns_false() {
5739 let router = McpRouter::new();
5741
5742 assert!(!router.log_info("Test"));
5744 assert!(!router.log_warning("Test"));
5745 assert!(!router.log_error("Test"));
5746 assert!(!router.log_debug("Test"));
5747 }
5748
5749 #[tokio::test]
5750 async fn test_logging_capability_with_channel() {
5751 use crate::context::notification_channel;
5752
5753 let (tx, _rx) = notification_channel(10);
5754 let mut router = McpRouter::new().with_notification_sender(tx);
5755
5756 let init_req = RouterRequest {
5758 id: RequestId::Number(0),
5759 inner: McpRequest::Initialize(InitializeParams {
5760 protocol_version: "2025-11-25".to_string(),
5761 capabilities: ClientCapabilities {
5762 roots: None,
5763 sampling: None,
5764 elicitation: None,
5765 tasks: None,
5766 experimental: None,
5767 extensions: None,
5768 },
5769 client_info: Implementation {
5770 name: "test".to_string(),
5771 version: "1.0".to_string(),
5772 ..Default::default()
5773 },
5774 meta: None,
5775 }),
5776 extensions: Extensions::new(),
5777 };
5778 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5779
5780 match resp.inner {
5781 Ok(McpResponse::Initialize(result)) => {
5782 assert!(result.capabilities.logging.is_some());
5784 }
5785 _ => panic!("Expected Initialize response"),
5786 }
5787 }
5788
5789 #[tokio::test]
5790 async fn test_no_logging_capability_without_channel() {
5791 let mut router = McpRouter::new();
5792
5793 let init_req = RouterRequest {
5795 id: RequestId::Number(0),
5796 inner: McpRequest::Initialize(InitializeParams {
5797 protocol_version: "2025-11-25".to_string(),
5798 capabilities: ClientCapabilities {
5799 roots: None,
5800 sampling: None,
5801 elicitation: None,
5802 tasks: None,
5803 experimental: None,
5804 extensions: None,
5805 },
5806 client_info: Implementation {
5807 name: "test".to_string(),
5808 version: "1.0".to_string(),
5809 ..Default::default()
5810 },
5811 meta: None,
5812 }),
5813 extensions: Extensions::new(),
5814 };
5815 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5816
5817 match resp.inner {
5818 Ok(McpResponse::Initialize(result)) => {
5819 assert!(result.capabilities.logging.is_none());
5821 }
5822 _ => panic!("Expected Initialize response"),
5823 }
5824 }
5825
5826 #[tokio::test]
5835 async fn stable_task_update_applies_input_responses() {
5836 use crate::async_task::{MemoryTaskStore, TaskStore};
5837 use crate::protocol::{InputRequest, ListRootsParams};
5838
5839 let store = std::sync::Arc::new(MemoryTaskStore::new());
5840 let mut router = McpRouter::new().task_store(store.clone());
5841 init_router(&mut router).await;
5842
5843 let (task_id, _cancel) = store
5844 .create_task("permission_gate", serde_json::json!({}), None, None)
5845 .await
5846 .expect("create task");
5847 let requests: crate::protocol::InputRequests = [(
5848 "permission".to_string(),
5849 InputRequest::ListRoots(ListRootsParams { meta: None }),
5850 )]
5851 .into_iter()
5852 .collect();
5853 store
5854 .require_input(&task_id, requests, Some("need a decision"))
5855 .await
5856 .expect("require input");
5857 assert_eq!(
5858 store.get_task(&task_id).await.unwrap().unwrap().status,
5859 TaskStatus::InputRequired
5860 );
5861
5862 let resp = router
5864 .ready()
5865 .await
5866 .unwrap()
5867 .call(RouterRequest {
5868 id: RequestId::Number(1),
5869 inner: McpRequest::UpdateTask(UpdateTaskParams {
5870 task_id: task_id.clone(),
5871 input_responses: [("permission".to_string(), serde_json::json!({"roots": []}))]
5872 .into_iter()
5873 .collect(),
5874 meta: None,
5875 }),
5876 extensions: Extensions::new(),
5877 })
5878 .await
5879 .unwrap();
5880 assert!(
5881 matches!(resp.inner, Ok(McpResponse::UpdateTask(_))),
5882 "the empty-result acknowledgment shape is unchanged: {:?}",
5883 resp.inner
5884 );
5885
5886 assert!(
5889 store
5890 .outstanding_input_requests(&task_id)
5891 .await
5892 .unwrap()
5893 .unwrap()
5894 .is_empty(),
5895 "the outstanding request must be consumed"
5896 );
5897 assert_eq!(
5898 store.get_task(&task_id).await.unwrap().unwrap().status,
5899 TaskStatus::Working,
5900 "answering the last outstanding request resumes the task"
5901 );
5902 }
5903
5904 #[tokio::test]
5907 async fn stable_task_update_ignores_unmatched_keys() {
5908 use crate::async_task::{MemoryTaskStore, TaskStore};
5909
5910 let store = std::sync::Arc::new(MemoryTaskStore::new());
5911 let mut router = McpRouter::new().task_store(store.clone());
5912 init_router(&mut router).await;
5913
5914 let (task_id, _cancel) = store
5915 .create_task("noop", serde_json::json!({}), None, None)
5916 .await
5917 .expect("create task");
5918
5919 let resp = router
5920 .ready()
5921 .await
5922 .unwrap()
5923 .call(RouterRequest {
5924 id: RequestId::Number(1),
5925 inner: McpRequest::UpdateTask(UpdateTaskParams {
5926 task_id: task_id.clone(),
5927 input_responses: [(
5928 "never-issued".to_string(),
5929 serde_json::json!({"roots": []}),
5930 )]
5931 .into_iter()
5932 .collect(),
5933 meta: None,
5934 }),
5935 extensions: Extensions::new(),
5936 })
5937 .await
5938 .unwrap();
5939 assert!(
5940 matches!(resp.inner, Ok(McpResponse::UpdateTask(_))),
5941 "an unmatched key is ignored, not rejected: {:?}",
5942 resp.inner
5943 );
5944
5945 let unknown = router
5946 .ready()
5947 .await
5948 .unwrap()
5949 .call(RouterRequest {
5950 id: RequestId::Number(2),
5951 inner: McpRequest::UpdateTask(UpdateTaskParams {
5952 task_id: "no-such-task".to_string(),
5953 input_responses: HashMap::new(),
5954 meta: None,
5955 }),
5956 extensions: Extensions::new(),
5957 })
5958 .await
5959 .unwrap();
5960 match unknown.inner {
5961 Err(error) => assert_eq!(error.code, -32602),
5962 other => panic!("expected -32602 for an unknown task, got {other:?}"),
5963 }
5964 }
5965
5966 #[tokio::test]
5967 async fn test_create_task_via_call_tool() {
5968 let add_tool = ToolBuilder::new("add")
5969 .description("Add two numbers")
5970 .task_support(TaskSupportMode::Optional)
5971 .handler(|input: AddInput| async move {
5972 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5973 })
5974 .build();
5975
5976 let mut router = McpRouter::new().tool(add_tool);
5977 init_router(&mut router).await;
5978
5979 let req = RouterRequest {
5980 id: RequestId::Number(1),
5981 inner: McpRequest::CallTool(CallToolParams {
5982 input_responses: None,
5983 request_state: None,
5984 name: "add".to_string(),
5985 arguments: serde_json::json!({"a": 5, "b": 10}),
5986 meta: None,
5987 task: Some(TaskRequestParams { ttl: None }),
5988 }),
5989 extensions: Extensions::new(),
5990 };
5991
5992 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5993
5994 match resp.inner {
5995 Ok(McpResponse::CreateTask(result)) => {
5996 assert!(!result.task.task_id.is_empty());
5997 assert_eq!(result.task.status, TaskStatus::Working);
5998 }
5999 _ => panic!("Expected CreateTask response"),
6000 }
6001 }
6002
6003 struct CountingTaskStore {
6006 inner: MemoryTaskStore,
6007 creates: std::sync::atomic::AtomicUsize,
6008 gets: std::sync::atomic::AtomicUsize,
6009 completes: std::sync::atomic::AtomicUsize,
6010 }
6011
6012 impl CountingTaskStore {
6013 fn new() -> Self {
6014 Self {
6015 inner: MemoryTaskStore::new(),
6016 creates: std::sync::atomic::AtomicUsize::new(0),
6017 gets: std::sync::atomic::AtomicUsize::new(0),
6018 completes: std::sync::atomic::AtomicUsize::new(0),
6019 }
6020 }
6021 }
6022
6023 #[async_trait::async_trait]
6024 impl TaskStore for CountingTaskStore {
6025 async fn create_task(
6026 &self,
6027 tool_name: &str,
6028 arguments: serde_json::Value,
6029 ttl: Option<u64>,
6030 owner: crate::async_task::TaskOwner,
6031 ) -> crate::async_task::Result<(String, crate::async_task::CancellationToken)> {
6032 self.creates
6033 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6034 self.inner
6035 .create_task(tool_name, arguments, ttl, owner)
6036 .await
6037 }
6038
6039 async fn get_task(&self, task_id: &str) -> crate::async_task::Result<Option<TaskObject>> {
6040 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6041 self.inner.get_task(task_id).await
6042 }
6043
6044 async fn task_owner(
6045 &self,
6046 task_id: &str,
6047 ) -> crate::async_task::Result<Option<crate::async_task::TaskOwner>> {
6048 self.inner.task_owner(task_id).await
6049 }
6050
6051 async fn get_task_result(
6052 &self,
6053 task_id: &str,
6054 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
6055 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6058 self.inner.get_task_result(task_id).await
6059 }
6060
6061 async fn wait_for_completion(
6062 &self,
6063 task_id: &str,
6064 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
6065 self.inner.wait_for_completion(task_id).await
6066 }
6067
6068 async fn list_tasks(
6069 &self,
6070 status_filter: Option<TaskStatus>,
6071 ) -> crate::async_task::Result<Vec<TaskObject>> {
6072 self.inner.list_tasks(status_filter).await
6073 }
6074
6075 async fn require_input(
6076 &self,
6077 task_id: &str,
6078 requests: crate::protocol::InputRequests,
6079 message: Option<&str>,
6080 ) -> crate::async_task::Result<bool> {
6081 self.inner.require_input(task_id, requests, message).await
6082 }
6083
6084 async fn outstanding_input_requests(
6085 &self,
6086 task_id: &str,
6087 ) -> crate::async_task::Result<Option<crate::protocol::InputRequests>> {
6088 self.inner.outstanding_input_requests(task_id).await
6089 }
6090
6091 async fn apply_input_responses(
6092 &self,
6093 task_id: &str,
6094 responses: crate::protocol::InputResponses,
6095 ) -> crate::async_task::Result<Option<crate::async_task::AppliedInputResponses>> {
6096 self.inner.apply_input_responses(task_id, responses).await
6097 }
6098
6099 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> crate::async_task::Result<bool> {
6100 self.inner.set_ttl(task_id, ttl_ms).await
6101 }
6102
6103 async fn complete_task(
6104 &self,
6105 task_id: &str,
6106 result: CallToolResult,
6107 ) -> crate::async_task::Result<bool> {
6108 self.completes
6109 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6110 self.inner.complete_task(task_id, result).await
6111 }
6112
6113 async fn fail_task(
6114 &self,
6115 task_id: &str,
6116 error: JsonRpcError,
6117 ) -> crate::async_task::Result<bool> {
6118 self.inner.fail_task(task_id, error).await
6119 }
6120
6121 async fn cancel_task(
6122 &self,
6123 task_id: &str,
6124 reason: Option<&str>,
6125 ) -> crate::async_task::Result<Option<TaskObject>> {
6126 self.inner.cancel_task(task_id, reason).await
6127 }
6128 }
6129
6130 #[tokio::test]
6131 async fn test_injected_task_store_used_by_dispatch() {
6132 let store = Arc::new(CountingTaskStore::new());
6133
6134 let add_tool = ToolBuilder::new("add")
6135 .description("Add two numbers")
6136 .task_support(TaskSupportMode::Optional)
6137 .handler(|input: AddInput| async move {
6138 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6139 })
6140 .build();
6141
6142 let mut router = McpRouter::new()
6143 .tool(add_tool)
6144 .task_store(store.clone() as Arc<dyn TaskStore>);
6145 init_router(&mut router).await;
6146
6147 let req = RouterRequest {
6149 id: RequestId::Number(1),
6150 inner: McpRequest::CallTool(CallToolParams {
6151 input_responses: None,
6152 request_state: None,
6153 name: "add".to_string(),
6154 arguments: serde_json::json!({"a": 2, "b": 3}),
6155 meta: None,
6156 task: Some(TaskRequestParams { ttl: None }),
6157 }),
6158 extensions: Extensions::new(),
6159 };
6160 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6161 let task_id = match resp.inner {
6162 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6163 other => panic!("Expected CreateTask response, got {other:?}"),
6164 };
6165
6166 assert_eq!(
6167 store.creates.load(std::sync::atomic::Ordering::Relaxed),
6168 1,
6169 "create_task must go through the injected store"
6170 );
6171
6172 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6174 assert_eq!(
6175 store.completes.load(std::sync::atomic::Ordering::Relaxed),
6176 1,
6177 "complete_task must go through the injected store"
6178 );
6179
6180 let gets_before = store.gets.load(std::sync::atomic::Ordering::Relaxed);
6182 let req = RouterRequest {
6183 id: RequestId::Number(2),
6184 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6185 task_id: task_id.clone(),
6186 meta: None,
6187 }),
6188 extensions: Extensions::new(),
6189 };
6190 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6191 match resp.inner {
6192 Ok(McpResponse::GetTaskInfo(info)) => {
6193 assert_eq!(info.task_id, task_id);
6194 assert_eq!(info.status, TaskStatus::Completed);
6195 }
6196 other => panic!("Expected GetTaskInfo response, got {other:?}"),
6197 }
6198 assert!(
6199 store.gets.load(std::sync::atomic::Ordering::Relaxed) > gets_before,
6200 "tasks/get must go through the injected store"
6201 );
6202 }
6203
6204 #[tokio::test]
6205 async fn test_removed_tasks_methods_get_method_not_found() {
6206 let mut router = McpRouter::new();
6210 init_router(&mut router).await;
6211
6212 for method in ["tasks/list", "tasks/result"] {
6213 let req = RouterRequest {
6214 id: RequestId::Number(1),
6215 inner: McpRequest::Unknown {
6216 method: method.to_string(),
6217 params: None,
6218 },
6219 extensions: Extensions::new(),
6220 };
6221
6222 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6223
6224 match resp.inner {
6225 Err(err) => {
6226 assert_eq!(err.code, -32601, "{method} must be MethodNotFound");
6227 }
6228 other => panic!("Expected MethodNotFound error for {method}, got {other:?}"),
6229 }
6230 }
6231 }
6232
6233 #[tokio::test]
6234 async fn test_task_lifecycle_complete() {
6235 let add_tool = ToolBuilder::new("add")
6236 .description("Add two numbers")
6237 .task_support(TaskSupportMode::Optional)
6238 .handler(|input: AddInput| async move {
6239 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6240 })
6241 .build();
6242
6243 let mut router = McpRouter::new().tool(add_tool);
6244 init_router(&mut router).await;
6245
6246 let req = RouterRequest {
6248 id: RequestId::Number(1),
6249 inner: McpRequest::CallTool(CallToolParams {
6250 input_responses: None,
6251 request_state: None,
6252 name: "add".to_string(),
6253 arguments: serde_json::json!({"a": 7, "b": 8}),
6254 meta: None,
6255 task: Some(TaskRequestParams { ttl: None }),
6256 }),
6257 extensions: Extensions::new(),
6258 };
6259
6260 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6261 let task_id = match resp.inner {
6262 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6263 _ => panic!("Expected CreateTask response"),
6264 };
6265
6266 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6268
6269 let req = RouterRequest {
6273 id: RequestId::Number(2),
6274 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6275 task_id: task_id.clone(),
6276 meta: None,
6277 }),
6278 extensions: Extensions::new(),
6279 };
6280
6281 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6282
6283 match resp.inner {
6284 Ok(McpResponse::GetTaskInfo(info)) => {
6285 assert_eq!(info.task_id, task_id);
6286 assert_eq!(info.status, TaskStatus::Completed);
6287 }
6288 _ => panic!("Expected GetTaskInfo response"),
6289 }
6290 }
6291
6292 #[tokio::test]
6293 async fn test_task_cancellation() {
6294 let slow_tool = ToolBuilder::new("slow")
6296 .description("Slow tool")
6297 .task_support(TaskSupportMode::Optional)
6298 .handler(|_input: serde_json::Value| async move {
6299 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
6300 Ok(CallToolResult::text("done"))
6301 })
6302 .build();
6303
6304 let mut router = McpRouter::new().tool(slow_tool);
6305 init_router(&mut router).await;
6306
6307 let req = RouterRequest {
6309 id: RequestId::Number(1),
6310 inner: McpRequest::CallTool(CallToolParams {
6311 input_responses: None,
6312 request_state: None,
6313 name: "slow".to_string(),
6314 arguments: serde_json::json!({}),
6315 meta: None,
6316 task: Some(TaskRequestParams { ttl: None }),
6317 }),
6318 extensions: Extensions::new(),
6319 };
6320
6321 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6322 let task_id = match resp.inner {
6323 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6324 _ => panic!("Expected CreateTask response"),
6325 };
6326
6327 let req = RouterRequest {
6329 id: RequestId::Number(2),
6330 inner: McpRequest::CancelTask(CancelTaskParams {
6331 task_id: task_id.clone(),
6332 reason: Some("Test cancellation".to_string()),
6333 meta: None,
6334 }),
6335 extensions: Extensions::new(),
6336 };
6337
6338 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6339
6340 match resp.inner {
6342 Ok(McpResponse::CancelTask(EmptyResult {})) => {}
6343 other => panic!("Expected empty CancelTask ack, got {other:?}"),
6344 }
6345
6346 let req = RouterRequest {
6348 id: RequestId::Number(3),
6349 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6350 task_id: task_id.clone(),
6351 meta: None,
6352 }),
6353 extensions: Extensions::new(),
6354 };
6355 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6356 match resp.inner {
6357 Ok(McpResponse::GetTaskInfo(info)) => {
6358 assert_eq!(info.status, TaskStatus::Cancelled);
6359 }
6360 _ => panic!("Expected GetTaskInfo response"),
6361 }
6362 }
6363
6364 #[tokio::test]
6365 async fn test_get_task_info() {
6366 let add_tool = ToolBuilder::new("add")
6367 .description("Add two numbers")
6368 .task_support(TaskSupportMode::Optional)
6369 .handler(|input: AddInput| async move {
6370 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6371 })
6372 .build();
6373
6374 let mut router = McpRouter::new().tool(add_tool);
6375 init_router(&mut router).await;
6376
6377 let req = RouterRequest {
6379 id: RequestId::Number(1),
6380 inner: McpRequest::CallTool(CallToolParams {
6381 input_responses: None,
6382 request_state: None,
6383 name: "add".to_string(),
6384 arguments: serde_json::json!({"a": 1, "b": 2}),
6385 meta: None,
6386 task: Some(TaskRequestParams { ttl: Some(600_000) }),
6387 }),
6388 extensions: Extensions::new(),
6389 };
6390
6391 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6392 let task_id = match resp.inner {
6393 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6394 _ => panic!("Expected CreateTask response"),
6395 };
6396
6397 let req = RouterRequest {
6399 id: RequestId::Number(2),
6400 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6401 task_id: task_id.clone(),
6402 meta: None,
6403 }),
6404 extensions: Extensions::new(),
6405 };
6406
6407 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6408
6409 match resp.inner {
6410 Ok(McpResponse::GetTaskInfo(info)) => {
6411 assert_eq!(info.task_id, task_id);
6412 assert!(info.created_at.contains('T')); assert_eq!(info.ttl, Some(600_000));
6414 }
6415 _ => panic!("Expected GetTaskInfo response"),
6416 }
6417 }
6418
6419 #[tokio::test]
6420 async fn test_task_forbidden_tool_rejects_task_params() {
6421 let tool = ToolBuilder::new("sync_only")
6422 .description("Sync only tool")
6423 .handler(|_input: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
6424 .build();
6425
6426 let mut router = McpRouter::new().tool(tool);
6427 init_router(&mut router).await;
6428
6429 let req = RouterRequest {
6431 id: RequestId::Number(1),
6432 inner: McpRequest::CallTool(CallToolParams {
6433 input_responses: None,
6434 request_state: None,
6435 name: "sync_only".to_string(),
6436 arguments: serde_json::json!({}),
6437 meta: None,
6438 task: Some(TaskRequestParams { ttl: None }),
6439 }),
6440 extensions: Extensions::new(),
6441 };
6442
6443 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6444
6445 match resp.inner {
6446 Err(e) => {
6447 assert!(e.message.contains("does not support async tasks"));
6448 }
6449 _ => panic!("Expected error response"),
6450 }
6451 }
6452
6453 #[tokio::test]
6454 async fn test_get_nonexistent_task() {
6455 let mut router = McpRouter::new();
6456 init_router(&mut router).await;
6457
6458 let req = RouterRequest {
6459 id: RequestId::Number(1),
6460 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6461 task_id: "task-999".to_string(),
6462 meta: None,
6463 }),
6464 extensions: Extensions::new(),
6465 };
6466
6467 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6468
6469 match resp.inner {
6470 Err(e) => {
6471 assert!(e.message.contains("not found"));
6472 }
6473 _ => panic!("Expected error response"),
6474 }
6475 }
6476
6477 #[tokio::test]
6482 async fn test_subscribe_to_resource() {
6483 use crate::resource::ResourceBuilder;
6484
6485 let resource = ResourceBuilder::new("file:///test.txt")
6486 .name("Test File")
6487 .text("Hello");
6488
6489 let mut router = McpRouter::new().resource(resource);
6490 init_router(&mut router).await;
6491
6492 let req = RouterRequest {
6494 id: RequestId::Number(1),
6495 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6496 uri: "file:///test.txt".to_string(),
6497 meta: None,
6498 }),
6499 extensions: Extensions::new(),
6500 };
6501
6502 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6503
6504 match resp.inner {
6505 Ok(McpResponse::SubscribeResource(_)) => {
6506 assert!(router.is_subscribed("file:///test.txt"));
6508 }
6509 _ => panic!("Expected SubscribeResource response"),
6510 }
6511 }
6512
6513 #[tokio::test]
6514 async fn test_unsubscribe_from_resource() {
6515 use crate::resource::ResourceBuilder;
6516
6517 let resource = ResourceBuilder::new("file:///test.txt")
6518 .name("Test File")
6519 .text("Hello");
6520
6521 let mut router = McpRouter::new().resource(resource);
6522 init_router(&mut router).await;
6523
6524 let req = RouterRequest {
6526 id: RequestId::Number(1),
6527 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6528 uri: "file:///test.txt".to_string(),
6529 meta: None,
6530 }),
6531 extensions: Extensions::new(),
6532 };
6533 let _ = router.ready().await.unwrap().call(req).await.unwrap();
6534 assert!(router.is_subscribed("file:///test.txt"));
6535
6536 let req = RouterRequest {
6538 id: RequestId::Number(2),
6539 inner: McpRequest::UnsubscribeResource(UnsubscribeResourceParams {
6540 uri: "file:///test.txt".to_string(),
6541 meta: None,
6542 }),
6543 extensions: Extensions::new(),
6544 };
6545
6546 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6547
6548 match resp.inner {
6549 Ok(McpResponse::UnsubscribeResource(_)) => {
6550 assert!(!router.is_subscribed("file:///test.txt"));
6552 }
6553 _ => panic!("Expected UnsubscribeResource response"),
6554 }
6555 }
6556
6557 #[tokio::test]
6558 async fn test_subscribe_nonexistent_resource() {
6559 let mut router = McpRouter::new();
6560 init_router(&mut router).await;
6561
6562 let req = RouterRequest {
6563 id: RequestId::Number(1),
6564 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6565 uri: "file:///nonexistent.txt".to_string(),
6566 meta: None,
6567 }),
6568 extensions: Extensions::new(),
6569 };
6570
6571 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6572
6573 match resp.inner {
6574 Err(e) => {
6575 assert!(e.message.contains("not found"));
6576 }
6577 _ => panic!("Expected error response"),
6578 }
6579 }
6580
6581 #[tokio::test]
6582 async fn test_notify_resource_updated() {
6583 use crate::context::notification_channel;
6584 use crate::resource::ResourceBuilder;
6585
6586 let (tx, mut rx) = notification_channel(10);
6587
6588 let resource = ResourceBuilder::new("file:///test.txt")
6589 .name("Test File")
6590 .text("Hello");
6591
6592 let router = McpRouter::new()
6593 .resource(resource)
6594 .with_notification_sender(tx);
6595
6596 router.subscribe("file:///test.txt");
6598
6599 let sent = router.notify_resource_updated("file:///test.txt");
6601 assert!(sent);
6602
6603 let notification = rx.try_recv().unwrap();
6605 match notification {
6606 ServerNotification::ResourceUpdated { uri } => {
6607 assert_eq!(uri, "file:///test.txt");
6608 }
6609 _ => panic!("Expected ResourceUpdated notification"),
6610 }
6611 }
6612
6613 #[tokio::test]
6614 async fn test_notify_resource_updated_not_subscribed() {
6615 use crate::context::notification_channel;
6616 use crate::resource::ResourceBuilder;
6617
6618 let (tx, mut rx) = notification_channel(10);
6619
6620 let resource = ResourceBuilder::new("file:///test.txt")
6621 .name("Test File")
6622 .text("Hello");
6623
6624 let router = McpRouter::new()
6625 .resource(resource)
6626 .with_notification_sender(tx);
6627
6628 let sent = router.notify_resource_updated("file:///test.txt");
6630 assert!(!sent); assert!(rx.try_recv().is_err());
6634 }
6635
6636 #[tokio::test]
6637 async fn test_notify_resources_list_changed() {
6638 use crate::context::notification_channel;
6639
6640 let (tx, mut rx) = notification_channel(10);
6641 let router = McpRouter::new().with_notification_sender(tx);
6642
6643 let sent = router.notify_resources_list_changed();
6644 assert!(sent);
6645
6646 let notification = rx.try_recv().unwrap();
6647 match notification {
6648 ServerNotification::ResourcesListChanged => {}
6649 _ => panic!("Expected ResourcesListChanged notification"),
6650 }
6651 }
6652
6653 #[tokio::test]
6654 async fn test_subscribed_uris() {
6655 use crate::resource::ResourceBuilder;
6656
6657 let resource1 = ResourceBuilder::new("file:///a.txt").name("A").text("A");
6658
6659 let resource2 = ResourceBuilder::new("file:///b.txt").name("B").text("B");
6660
6661 let router = McpRouter::new().resource(resource1).resource(resource2);
6662
6663 router.subscribe("file:///a.txt");
6665 router.subscribe("file:///b.txt");
6666
6667 let uris = router.subscribed_uris();
6668 assert_eq!(uris.len(), 2);
6669 assert!(uris.contains(&"file:///a.txt".to_string()));
6670 assert!(uris.contains(&"file:///b.txt".to_string()));
6671 }
6672
6673 #[tokio::test]
6674 async fn test_subscription_capability_advertised() {
6675 use crate::resource::ResourceBuilder;
6676
6677 let resource = ResourceBuilder::new("file:///test.txt")
6678 .name("Test")
6679 .text("Hello");
6680
6681 let mut router = McpRouter::new().resource(resource);
6682
6683 let init_req = RouterRequest {
6685 id: RequestId::Number(0),
6686 inner: McpRequest::Initialize(InitializeParams {
6687 protocol_version: "2025-11-25".to_string(),
6688 capabilities: ClientCapabilities {
6689 roots: None,
6690 sampling: None,
6691 elicitation: None,
6692 tasks: None,
6693 experimental: None,
6694 extensions: None,
6695 },
6696 client_info: Implementation {
6697 name: "test".to_string(),
6698 version: "1.0".to_string(),
6699 ..Default::default()
6700 },
6701 meta: None,
6702 }),
6703 extensions: Extensions::new(),
6704 };
6705 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
6706
6707 match resp.inner {
6708 Ok(McpResponse::Initialize(result)) => {
6709 let resources_cap = result.capabilities.resources.unwrap();
6711 assert!(resources_cap.subscribe);
6712 }
6713 _ => panic!("Expected Initialize response"),
6714 }
6715 }
6716
6717 #[tokio::test]
6718 async fn test_completion_handler() {
6719 let router = McpRouter::new()
6720 .server_info("test", "1.0")
6721 .completion_handler(|params: CompleteParams| async move {
6722 let prefix = ¶ms.argument.value;
6724 let suggestions: Vec<String> = vec!["alpha", "beta", "gamma"]
6725 .into_iter()
6726 .filter(|s| s.starts_with(prefix))
6727 .map(String::from)
6728 .collect();
6729 Ok(CompleteResult::new(suggestions))
6730 });
6731
6732 let init_req = RouterRequest {
6734 id: RequestId::Number(0),
6735 inner: McpRequest::Initialize(InitializeParams {
6736 protocol_version: "2025-11-25".to_string(),
6737 capabilities: ClientCapabilities::default(),
6738 client_info: Implementation {
6739 name: "test".to_string(),
6740 version: "1.0".to_string(),
6741 ..Default::default()
6742 },
6743 meta: None,
6744 }),
6745 extensions: Extensions::new(),
6746 };
6747 let resp = router
6748 .clone()
6749 .ready()
6750 .await
6751 .unwrap()
6752 .call(init_req)
6753 .await
6754 .unwrap();
6755
6756 match resp.inner {
6758 Ok(McpResponse::Initialize(result)) => {
6759 assert!(result.capabilities.completions.is_some());
6760 }
6761 _ => panic!("Expected Initialize response"),
6762 }
6763
6764 router.handle_notification(McpNotification::Initialized);
6766
6767 let complete_req = RouterRequest {
6769 id: RequestId::Number(1),
6770 inner: McpRequest::Complete(CompleteParams {
6771 reference: CompletionReference::prompt("test-prompt"),
6772 argument: CompletionArgument::new("query", "al"),
6773 context: None,
6774 meta: None,
6775 }),
6776 extensions: Extensions::new(),
6777 };
6778 let resp = router
6779 .clone()
6780 .ready()
6781 .await
6782 .unwrap()
6783 .call(complete_req)
6784 .await
6785 .unwrap();
6786
6787 match resp.inner {
6788 Ok(McpResponse::Complete(result)) => {
6789 assert_eq!(result.completion.values, vec!["alpha"]);
6790 }
6791 _ => panic!("Expected Complete response"),
6792 }
6793 }
6794
6795 #[tokio::test]
6796 async fn test_completion_without_handler_returns_empty() {
6797 let router = McpRouter::new().server_info("test", "1.0");
6798
6799 let init_req = RouterRequest {
6801 id: RequestId::Number(0),
6802 inner: McpRequest::Initialize(InitializeParams {
6803 protocol_version: "2025-11-25".to_string(),
6804 capabilities: ClientCapabilities::default(),
6805 client_info: Implementation {
6806 name: "test".to_string(),
6807 version: "1.0".to_string(),
6808 ..Default::default()
6809 },
6810 meta: None,
6811 }),
6812 extensions: Extensions::new(),
6813 };
6814 let resp = router
6815 .clone()
6816 .ready()
6817 .await
6818 .unwrap()
6819 .call(init_req)
6820 .await
6821 .unwrap();
6822
6823 match resp.inner {
6825 Ok(McpResponse::Initialize(result)) => {
6826 assert!(result.capabilities.completions.is_none());
6827 }
6828 _ => panic!("Expected Initialize response"),
6829 }
6830
6831 router.handle_notification(McpNotification::Initialized);
6833
6834 let complete_req = RouterRequest {
6836 id: RequestId::Number(1),
6837 inner: McpRequest::Complete(CompleteParams {
6838 reference: CompletionReference::prompt("test-prompt"),
6839 argument: CompletionArgument::new("query", "al"),
6840 context: None,
6841 meta: None,
6842 }),
6843 extensions: Extensions::new(),
6844 };
6845 let resp = router
6846 .clone()
6847 .ready()
6848 .await
6849 .unwrap()
6850 .call(complete_req)
6851 .await
6852 .unwrap();
6853
6854 match resp.inner {
6855 Ok(McpResponse::Complete(result)) => {
6856 assert!(result.completion.values.is_empty());
6857 }
6858 _ => panic!("Expected Complete response"),
6859 }
6860 }
6861
6862 #[tokio::test]
6863 async fn test_tool_filter_list() {
6864 use crate::filter::CapabilityFilter;
6865 use crate::tool::Tool;
6866
6867 let public_tool = ToolBuilder::new("public")
6868 .description("Public tool")
6869 .handler(|_: AddInput| async move { Ok(CallToolResult::text("public")) })
6870 .build();
6871
6872 let admin_tool = ToolBuilder::new("admin")
6873 .description("Admin tool")
6874 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6875 .build();
6876
6877 let mut router = McpRouter::new()
6878 .tool(public_tool)
6879 .tool(admin_tool)
6880 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| tool.name != "admin"));
6881
6882 init_router(&mut router).await;
6884
6885 let req = RouterRequest {
6886 id: RequestId::Number(1),
6887 inner: McpRequest::ListTools(ListToolsParams::default()),
6888 extensions: Extensions::new(),
6889 };
6890
6891 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6892
6893 match resp.inner {
6894 Ok(McpResponse::ListTools(result)) => {
6895 assert_eq!(result.tools.len(), 1);
6897 assert_eq!(result.tools[0].name, "public");
6898 }
6899 _ => panic!("Expected ListTools response"),
6900 }
6901 }
6902
6903 #[tokio::test]
6904 async fn test_tool_filter_call_denied() {
6905 use crate::filter::CapabilityFilter;
6906 use crate::tool::Tool;
6907
6908 let admin_tool = ToolBuilder::new("admin")
6909 .description("Admin tool")
6910 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6911 .build();
6912
6913 let mut router = McpRouter::new()
6914 .tool(admin_tool)
6915 .tool_filter(CapabilityFilter::new(|_, _: &Tool| false)); init_router(&mut router).await;
6919
6920 let req = RouterRequest {
6921 id: RequestId::Number(1),
6922 inner: McpRequest::CallTool(CallToolParams {
6923 input_responses: None,
6924 request_state: None,
6925 name: "admin".to_string(),
6926 arguments: serde_json::json!({"a": 1, "b": 2}),
6927 meta: None,
6928 task: None,
6929 }),
6930 extensions: Extensions::new(),
6931 };
6932
6933 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6934
6935 match resp.inner {
6937 Err(e) => {
6938 assert_eq!(e.code, -32601); }
6940 _ => panic!("Expected JsonRpc error"),
6941 }
6942 }
6943
6944 #[tokio::test]
6945 async fn test_tool_filter_call_allowed() {
6946 use crate::filter::CapabilityFilter;
6947 use crate::tool::Tool;
6948
6949 let public_tool = ToolBuilder::new("public")
6950 .description("Public tool")
6951 .handler(|input: AddInput| async move {
6952 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6953 })
6954 .build();
6955
6956 let mut router = McpRouter::new()
6957 .tool(public_tool)
6958 .tool_filter(CapabilityFilter::new(|_, _: &Tool| true)); init_router(&mut router).await;
6962
6963 let req = RouterRequest {
6964 id: RequestId::Number(1),
6965 inner: McpRequest::CallTool(CallToolParams {
6966 input_responses: None,
6967 request_state: None,
6968 name: "public".to_string(),
6969 arguments: serde_json::json!({"a": 1, "b": 2}),
6970 meta: None,
6971 task: None,
6972 }),
6973 extensions: Extensions::new(),
6974 };
6975
6976 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6977
6978 match resp.inner {
6979 Ok(McpResponse::CallTool(result)) => {
6980 assert!(!result.is_error);
6981 }
6982 _ => panic!("Expected CallTool response"),
6983 }
6984 }
6985
6986 #[tokio::test]
6987 async fn test_tool_filter_custom_denial() {
6988 use crate::filter::{CapabilityFilter, DenialBehavior};
6989 use crate::tool::Tool;
6990
6991 let admin_tool = ToolBuilder::new("admin")
6992 .description("Admin tool")
6993 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6994 .build();
6995
6996 let mut router = McpRouter::new().tool(admin_tool).tool_filter(
6997 CapabilityFilter::new(|_, _: &Tool| false)
6998 .denial_behavior(DenialBehavior::Unauthorized),
6999 );
7000
7001 init_router(&mut router).await;
7003
7004 let req = RouterRequest {
7005 id: RequestId::Number(1),
7006 inner: McpRequest::CallTool(CallToolParams {
7007 input_responses: None,
7008 request_state: None,
7009 name: "admin".to_string(),
7010 arguments: serde_json::json!({"a": 1, "b": 2}),
7011 meta: None,
7012 task: None,
7013 }),
7014 extensions: Extensions::new(),
7015 };
7016
7017 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7018
7019 match resp.inner {
7021 Err(e) => {
7022 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7024 }
7025 _ => panic!("Expected JsonRpc error"),
7026 }
7027 }
7028
7029 #[tokio::test]
7030 async fn test_resource_filter_list() {
7031 use crate::filter::CapabilityFilter;
7032 use crate::resource::{Resource, ResourceBuilder};
7033
7034 let public_resource = ResourceBuilder::new("file:///public.txt")
7035 .name("Public File")
7036 .text("public content");
7037
7038 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7039 .name("Secret File")
7040 .text("secret content");
7041
7042 let mut router = McpRouter::new()
7043 .resource(public_resource)
7044 .resource(secret_resource)
7045 .resource_filter(CapabilityFilter::new(|_, r: &Resource| {
7046 !r.name.contains("Secret")
7047 }));
7048
7049 init_router(&mut router).await;
7051
7052 let req = RouterRequest {
7053 id: RequestId::Number(1),
7054 inner: McpRequest::ListResources(ListResourcesParams::default()),
7055 extensions: Extensions::new(),
7056 };
7057
7058 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7059
7060 match resp.inner {
7061 Ok(McpResponse::ListResources(result)) => {
7062 assert_eq!(result.resources.len(), 1);
7064 assert_eq!(result.resources[0].name, "Public File");
7065 }
7066 _ => panic!("Expected ListResources response"),
7067 }
7068 }
7069
7070 #[tokio::test]
7071 async fn test_resource_filter_read_denied() {
7072 use crate::filter::CapabilityFilter;
7073 use crate::resource::{Resource, ResourceBuilder};
7074
7075 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7076 .name("Secret File")
7077 .text("secret content");
7078
7079 let mut router = McpRouter::new()
7080 .resource(secret_resource)
7081 .resource_filter(CapabilityFilter::new(|_, _: &Resource| false)); init_router(&mut router).await;
7085
7086 let req = RouterRequest {
7087 id: RequestId::Number(1),
7088 inner: McpRequest::ReadResource(ReadResourceParams {
7089 input_responses: None,
7090 request_state: None,
7091 uri: "file:///secret.txt".to_string(),
7092 meta: None,
7093 }),
7094 extensions: Extensions::new(),
7095 };
7096
7097 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7098
7099 match resp.inner {
7101 Err(e) => {
7102 assert_eq!(e.code, -32601); }
7104 _ => panic!("Expected JsonRpc error"),
7105 }
7106 }
7107
7108 #[tokio::test]
7109 async fn test_resource_filter_read_allowed() {
7110 use crate::filter::CapabilityFilter;
7111 use crate::resource::{Resource, ResourceBuilder};
7112
7113 let public_resource = ResourceBuilder::new("file:///public.txt")
7114 .name("Public File")
7115 .text("public content");
7116
7117 let mut router = McpRouter::new()
7118 .resource(public_resource)
7119 .resource_filter(CapabilityFilter::new(|_, _: &Resource| true)); init_router(&mut router).await;
7123
7124 let req = RouterRequest {
7125 id: RequestId::Number(1),
7126 inner: McpRequest::ReadResource(ReadResourceParams {
7127 input_responses: None,
7128 request_state: None,
7129 uri: "file:///public.txt".to_string(),
7130 meta: None,
7131 }),
7132 extensions: Extensions::new(),
7133 };
7134
7135 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7136
7137 match resp.inner {
7138 Ok(McpResponse::ReadResource(result)) => {
7139 assert_eq!(result.contents.len(), 1);
7140 assert_eq!(result.contents[0].text.as_deref(), Some("public content"));
7141 }
7142 _ => panic!("Expected ReadResource response"),
7143 }
7144 }
7145
7146 #[tokio::test]
7147 async fn test_resource_filter_custom_denial() {
7148 use crate::filter::{CapabilityFilter, DenialBehavior};
7149 use crate::resource::{Resource, ResourceBuilder};
7150
7151 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7152 .name("Secret File")
7153 .text("secret content");
7154
7155 let mut router = McpRouter::new().resource(secret_resource).resource_filter(
7156 CapabilityFilter::new(|_, _: &Resource| false)
7157 .denial_behavior(DenialBehavior::Unauthorized),
7158 );
7159
7160 init_router(&mut router).await;
7162
7163 let req = RouterRequest {
7164 id: RequestId::Number(1),
7165 inner: McpRequest::ReadResource(ReadResourceParams {
7166 input_responses: None,
7167 request_state: None,
7168 uri: "file:///secret.txt".to_string(),
7169 meta: None,
7170 }),
7171 extensions: Extensions::new(),
7172 };
7173
7174 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7175
7176 match resp.inner {
7178 Err(e) => {
7179 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7181 }
7182 _ => panic!("Expected JsonRpc error"),
7183 }
7184 }
7185
7186 #[tokio::test]
7187 async fn test_prompt_filter_list() {
7188 use crate::filter::CapabilityFilter;
7189 use crate::prompt::{Prompt, PromptBuilder};
7190
7191 let public_prompt = PromptBuilder::new("greeting")
7192 .description("A greeting")
7193 .user_message("Hello!");
7194
7195 let admin_prompt = PromptBuilder::new("system_debug")
7196 .description("Admin prompt")
7197 .user_message("Debug");
7198
7199 let mut router = McpRouter::new()
7200 .prompt(public_prompt)
7201 .prompt(admin_prompt)
7202 .prompt_filter(CapabilityFilter::new(|_, p: &Prompt| {
7203 !p.name.contains("system")
7204 }));
7205
7206 init_router(&mut router).await;
7208
7209 let req = RouterRequest {
7210 id: RequestId::Number(1),
7211 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7212 extensions: Extensions::new(),
7213 };
7214
7215 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7216
7217 match resp.inner {
7218 Ok(McpResponse::ListPrompts(result)) => {
7219 assert_eq!(result.prompts.len(), 1);
7221 assert_eq!(result.prompts[0].name, "greeting");
7222 }
7223 _ => panic!("Expected ListPrompts response"),
7224 }
7225 }
7226
7227 #[tokio::test]
7228 async fn test_prompt_filter_get_denied() {
7229 use crate::filter::CapabilityFilter;
7230 use crate::prompt::{Prompt, PromptBuilder};
7231 use std::collections::HashMap;
7232
7233 let admin_prompt = PromptBuilder::new("system_debug")
7234 .description("Admin prompt")
7235 .user_message("Debug");
7236
7237 let mut router = McpRouter::new()
7238 .prompt(admin_prompt)
7239 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| false)); init_router(&mut router).await;
7243
7244 let req = RouterRequest {
7245 id: RequestId::Number(1),
7246 inner: McpRequest::GetPrompt(GetPromptParams {
7247 input_responses: None,
7248 request_state: None,
7249 name: "system_debug".to_string(),
7250 arguments: HashMap::new(),
7251 meta: None,
7252 }),
7253 extensions: Extensions::new(),
7254 };
7255
7256 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7257
7258 match resp.inner {
7260 Err(e) => {
7261 assert_eq!(e.code, -32601); }
7263 _ => panic!("Expected JsonRpc error"),
7264 }
7265 }
7266
7267 #[tokio::test]
7268 async fn test_prompt_filter_get_allowed() {
7269 use crate::filter::CapabilityFilter;
7270 use crate::prompt::{Prompt, PromptBuilder};
7271 use std::collections::HashMap;
7272
7273 let public_prompt = PromptBuilder::new("greeting")
7274 .description("A greeting")
7275 .user_message("Hello!");
7276
7277 let mut router = McpRouter::new()
7278 .prompt(public_prompt)
7279 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| true)); init_router(&mut router).await;
7283
7284 let req = RouterRequest {
7285 id: RequestId::Number(1),
7286 inner: McpRequest::GetPrompt(GetPromptParams {
7287 input_responses: None,
7288 request_state: None,
7289 name: "greeting".to_string(),
7290 arguments: HashMap::new(),
7291 meta: None,
7292 }),
7293 extensions: Extensions::new(),
7294 };
7295
7296 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7297
7298 match resp.inner {
7299 Ok(McpResponse::GetPrompt(result)) => {
7300 assert_eq!(result.messages.len(), 1);
7301 }
7302 _ => panic!("Expected GetPrompt response"),
7303 }
7304 }
7305
7306 #[tokio::test]
7307 async fn test_prompt_filter_custom_denial() {
7308 use crate::filter::{CapabilityFilter, DenialBehavior};
7309 use crate::prompt::{Prompt, PromptBuilder};
7310 use std::collections::HashMap;
7311
7312 let admin_prompt = PromptBuilder::new("system_debug")
7313 .description("Admin prompt")
7314 .user_message("Debug");
7315
7316 let mut router = McpRouter::new().prompt(admin_prompt).prompt_filter(
7317 CapabilityFilter::new(|_, _: &Prompt| false)
7318 .denial_behavior(DenialBehavior::Unauthorized),
7319 );
7320
7321 init_router(&mut router).await;
7323
7324 let req = RouterRequest {
7325 id: RequestId::Number(1),
7326 inner: McpRequest::GetPrompt(GetPromptParams {
7327 input_responses: None,
7328 request_state: None,
7329 name: "system_debug".to_string(),
7330 arguments: HashMap::new(),
7331 meta: None,
7332 }),
7333 extensions: Extensions::new(),
7334 };
7335
7336 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7337
7338 match resp.inner {
7340 Err(e) => {
7341 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7343 }
7344 _ => panic!("Expected JsonRpc error"),
7345 }
7346 }
7347
7348 #[derive(Debug, Deserialize, JsonSchema)]
7353 struct StringInput {
7354 value: String,
7355 }
7356
7357 #[tokio::test]
7358 async fn test_router_merge_tools() {
7359 let tool_a = ToolBuilder::new("tool_a")
7361 .description("Tool A")
7362 .handler(|_: StringInput| async move { Ok(CallToolResult::text("A")) })
7363 .build();
7364
7365 let router_a = McpRouter::new().tool(tool_a);
7366
7367 let tool_b = ToolBuilder::new("tool_b")
7369 .description("Tool B")
7370 .handler(|_: StringInput| async move { Ok(CallToolResult::text("B")) })
7371 .build();
7372 let tool_c = ToolBuilder::new("tool_c")
7373 .description("Tool C")
7374 .handler(|_: StringInput| async move { Ok(CallToolResult::text("C")) })
7375 .build();
7376
7377 let router_b = McpRouter::new().tool(tool_b).tool(tool_c);
7378
7379 let mut merged = McpRouter::new()
7381 .server_info("merged", "1.0")
7382 .merge(router_a)
7383 .merge(router_b);
7384
7385 init_router(&mut merged).await;
7386
7387 let req = RouterRequest {
7389 id: RequestId::Number(1),
7390 inner: McpRequest::ListTools(ListToolsParams::default()),
7391 extensions: Extensions::new(),
7392 };
7393
7394 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7395
7396 match resp.inner {
7397 Ok(McpResponse::ListTools(result)) => {
7398 assert_eq!(result.tools.len(), 3);
7399 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7400 assert!(names.contains(&"tool_a"));
7401 assert!(names.contains(&"tool_b"));
7402 assert!(names.contains(&"tool_c"));
7403 }
7404 _ => panic!("Expected ListTools response"),
7405 }
7406 }
7407
7408 #[tokio::test]
7409 async fn test_router_merge_overwrites_duplicates() {
7410 let tool_v1 = ToolBuilder::new("shared")
7412 .description("Version 1")
7413 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v1")) })
7414 .build();
7415
7416 let router_a = McpRouter::new().tool(tool_v1);
7417
7418 let tool_v2 = ToolBuilder::new("shared")
7420 .description("Version 2")
7421 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v2")) })
7422 .build();
7423
7424 let router_b = McpRouter::new().tool(tool_v2);
7425
7426 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7428
7429 init_router(&mut merged).await;
7430
7431 let req = RouterRequest {
7432 id: RequestId::Number(1),
7433 inner: McpRequest::ListTools(ListToolsParams::default()),
7434 extensions: Extensions::new(),
7435 };
7436
7437 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7438
7439 match resp.inner {
7440 Ok(McpResponse::ListTools(result)) => {
7441 assert_eq!(result.tools.len(), 1);
7442 assert_eq!(result.tools[0].name, "shared");
7443 assert_eq!(result.tools[0].description.as_deref(), Some("Version 2"));
7444 }
7445 _ => panic!("Expected ListTools response"),
7446 }
7447 }
7448
7449 #[tokio::test]
7450 async fn test_router_merge_resources() {
7451 use crate::resource::ResourceBuilder;
7452
7453 let router_a = McpRouter::new().resource(
7455 ResourceBuilder::new("file:///a.txt")
7456 .name("File A")
7457 .text("content a"),
7458 );
7459
7460 let router_b = McpRouter::new().resource(
7461 ResourceBuilder::new("file:///b.txt")
7462 .name("File B")
7463 .text("content b"),
7464 );
7465
7466 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7467
7468 init_router(&mut merged).await;
7469
7470 let req = RouterRequest {
7471 id: RequestId::Number(1),
7472 inner: McpRequest::ListResources(ListResourcesParams::default()),
7473 extensions: Extensions::new(),
7474 };
7475
7476 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7477
7478 match resp.inner {
7479 Ok(McpResponse::ListResources(result)) => {
7480 assert_eq!(result.resources.len(), 2);
7481 let uris: Vec<&str> = result.resources.iter().map(|r| r.uri.as_str()).collect();
7482 assert!(uris.contains(&"file:///a.txt"));
7483 assert!(uris.contains(&"file:///b.txt"));
7484 }
7485 _ => panic!("Expected ListResources response"),
7486 }
7487 }
7488
7489 #[tokio::test]
7490 async fn test_router_merge_prompts() {
7491 use crate::prompt::PromptBuilder;
7492
7493 let router_a =
7494 McpRouter::new().prompt(PromptBuilder::new("prompt_a").user_message("Hello A"));
7495
7496 let router_b =
7497 McpRouter::new().prompt(PromptBuilder::new("prompt_b").user_message("Hello B"));
7498
7499 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7500
7501 init_router(&mut merged).await;
7502
7503 let req = RouterRequest {
7504 id: RequestId::Number(1),
7505 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7506 extensions: Extensions::new(),
7507 };
7508
7509 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7510
7511 match resp.inner {
7512 Ok(McpResponse::ListPrompts(result)) => {
7513 assert_eq!(result.prompts.len(), 2);
7514 let names: Vec<&str> = result.prompts.iter().map(|p| p.name.as_str()).collect();
7515 assert!(names.contains(&"prompt_a"));
7516 assert!(names.contains(&"prompt_b"));
7517 }
7518 _ => panic!("Expected ListPrompts response"),
7519 }
7520 }
7521
7522 #[tokio::test]
7523 async fn test_router_nest_prefixes_tools() {
7524 let tool_query = ToolBuilder::new("query")
7526 .description("Query the database")
7527 .handler(|_: StringInput| async move { Ok(CallToolResult::text("query result")) })
7528 .build();
7529 let tool_insert = ToolBuilder::new("insert")
7530 .description("Insert into database")
7531 .handler(|_: StringInput| async move { Ok(CallToolResult::text("insert result")) })
7532 .build();
7533
7534 let db_router = McpRouter::new().tool(tool_query).tool(tool_insert);
7535
7536 let mut router = McpRouter::new()
7538 .server_info("nested", "1.0")
7539 .nest("db", db_router);
7540
7541 init_router(&mut router).await;
7542
7543 let req = RouterRequest {
7544 id: RequestId::Number(1),
7545 inner: McpRequest::ListTools(ListToolsParams::default()),
7546 extensions: Extensions::new(),
7547 };
7548
7549 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7550
7551 match resp.inner {
7552 Ok(McpResponse::ListTools(result)) => {
7553 assert_eq!(result.tools.len(), 2);
7554 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7555 assert!(names.contains(&"db.query"));
7556 assert!(names.contains(&"db.insert"));
7557 }
7558 _ => panic!("Expected ListTools response"),
7559 }
7560 }
7561
7562 #[tokio::test]
7563 async fn test_router_nest_call_prefixed_tool() {
7564 let tool = ToolBuilder::new("echo")
7565 .description("Echo input")
7566 .handler(|input: StringInput| async move { Ok(CallToolResult::text(&input.value)) })
7567 .build();
7568
7569 let nested_router = McpRouter::new().tool(tool);
7570
7571 let mut router = McpRouter::new().nest("api", nested_router);
7572
7573 init_router(&mut router).await;
7574
7575 let req = RouterRequest {
7577 id: RequestId::Number(1),
7578 inner: McpRequest::CallTool(CallToolParams {
7579 input_responses: None,
7580 request_state: None,
7581 name: "api.echo".to_string(),
7582 arguments: serde_json::json!({"value": "hello world"}),
7583 meta: None,
7584 task: None,
7585 }),
7586 extensions: Extensions::new(),
7587 };
7588
7589 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7590
7591 match resp.inner {
7592 Ok(McpResponse::CallTool(result)) => {
7593 assert!(!result.is_error);
7594 match &result.content[0] {
7595 Content::Text { text, .. } => assert_eq!(text, "hello world"),
7596 _ => panic!("Expected text content"),
7597 }
7598 }
7599 _ => panic!("Expected CallTool response"),
7600 }
7601 }
7602
7603 #[tokio::test]
7604 async fn test_router_multiple_nests() {
7605 let db_tool = ToolBuilder::new("query")
7606 .description("Database query")
7607 .handler(|_: StringInput| async move { Ok(CallToolResult::text("db")) })
7608 .build();
7609
7610 let api_tool = ToolBuilder::new("fetch")
7611 .description("API fetch")
7612 .handler(|_: StringInput| async move { Ok(CallToolResult::text("api")) })
7613 .build();
7614
7615 let db_router = McpRouter::new().tool(db_tool);
7616 let api_router = McpRouter::new().tool(api_tool);
7617
7618 let mut router = McpRouter::new()
7619 .nest("db", db_router)
7620 .nest("api", api_router);
7621
7622 init_router(&mut router).await;
7623
7624 let req = RouterRequest {
7625 id: RequestId::Number(1),
7626 inner: McpRequest::ListTools(ListToolsParams::default()),
7627 extensions: Extensions::new(),
7628 };
7629
7630 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7631
7632 match resp.inner {
7633 Ok(McpResponse::ListTools(result)) => {
7634 assert_eq!(result.tools.len(), 2);
7635 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7636 assert!(names.contains(&"db.query"));
7637 assert!(names.contains(&"api.fetch"));
7638 }
7639 _ => panic!("Expected ListTools response"),
7640 }
7641 }
7642
7643 #[tokio::test]
7644 async fn test_router_merge_and_nest_combined() {
7645 let tool_a = ToolBuilder::new("local")
7647 .description("Local tool")
7648 .handler(|_: StringInput| async move { Ok(CallToolResult::text("local")) })
7649 .build();
7650
7651 let nested_tool = ToolBuilder::new("remote")
7652 .description("Remote tool")
7653 .handler(|_: StringInput| async move { Ok(CallToolResult::text("remote")) })
7654 .build();
7655
7656 let nested_router = McpRouter::new().tool(nested_tool);
7657
7658 let mut router = McpRouter::new()
7659 .tool(tool_a)
7660 .nest("external", nested_router);
7661
7662 init_router(&mut router).await;
7663
7664 let req = RouterRequest {
7665 id: RequestId::Number(1),
7666 inner: McpRequest::ListTools(ListToolsParams::default()),
7667 extensions: Extensions::new(),
7668 };
7669
7670 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7671
7672 match resp.inner {
7673 Ok(McpResponse::ListTools(result)) => {
7674 assert_eq!(result.tools.len(), 2);
7675 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7676 assert!(names.contains(&"local"));
7677 assert!(names.contains(&"external.remote"));
7678 }
7679 _ => panic!("Expected ListTools response"),
7680 }
7681 }
7682
7683 #[tokio::test]
7684 async fn test_router_merge_preserves_server_info() {
7685 let child_router = McpRouter::new()
7686 .server_info("child", "2.0")
7687 .instructions("Child instructions");
7688
7689 let mut router = McpRouter::new()
7690 .server_info("parent", "1.0")
7691 .instructions("Parent instructions")
7692 .merge(child_router);
7693
7694 init_router(&mut router).await;
7695
7696 let init_req = RouterRequest {
7698 id: RequestId::Number(99),
7699 inner: McpRequest::Initialize(InitializeParams {
7700 protocol_version: "2025-11-25".to_string(),
7701 capabilities: ClientCapabilities::default(),
7702 client_info: Implementation {
7703 name: "test".to_string(),
7704 version: "1.0".to_string(),
7705 ..Default::default()
7706 },
7707 meta: None,
7708 }),
7709 extensions: Extensions::new(),
7710 };
7711
7712 let child_router2 = McpRouter::new().server_info("child", "2.0");
7714 let mut fresh_router = McpRouter::new()
7715 .server_info("parent", "1.0")
7716 .merge(child_router2);
7717
7718 let resp = fresh_router
7719 .ready()
7720 .await
7721 .unwrap()
7722 .call(init_req)
7723 .await
7724 .unwrap();
7725
7726 match resp.inner {
7727 Ok(McpResponse::Initialize(result)) => {
7728 assert_eq!(result.server_info.name, "parent");
7729 assert_eq!(result.server_info.version, "1.0");
7730 }
7731 _ => panic!("Expected Initialize response"),
7732 }
7733 }
7734
7735 #[tokio::test]
7740 async fn test_auto_instructions_tools_only() {
7741 let tool_a = ToolBuilder::new("alpha")
7742 .description("Alpha tool")
7743 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7744 .build();
7745 let tool_b = ToolBuilder::new("beta")
7746 .description("Beta tool")
7747 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7748 .build();
7749
7750 let mut router = McpRouter::new()
7751 .auto_instructions()
7752 .tool(tool_a)
7753 .tool(tool_b);
7754
7755 let resp = send_initialize(&mut router).await;
7756 let instructions = resp.instructions.expect("should have instructions");
7757
7758 assert!(instructions.contains("## Tools"));
7759 assert!(instructions.contains("- **alpha**: Alpha tool"));
7760 assert!(instructions.contains("- **beta**: Beta tool"));
7761 assert!(!instructions.contains("## Resources"));
7763 assert!(!instructions.contains("## Prompts"));
7764 }
7765
7766 #[tokio::test]
7767 async fn test_auto_instructions_with_annotations() {
7768 let read_only_tool = ToolBuilder::new("query")
7769 .description("Run a query")
7770 .read_only()
7771 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7772 .build();
7773 let destructive_tool = ToolBuilder::new("delete")
7774 .description("Delete a record")
7775 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7776 .build();
7777 let idempotent_tool = ToolBuilder::new("upsert")
7778 .description("Upsert a record")
7779 .non_destructive()
7780 .idempotent()
7781 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7782 .build();
7783
7784 let mut router = McpRouter::new()
7785 .auto_instructions()
7786 .tool(read_only_tool)
7787 .tool(destructive_tool)
7788 .tool(idempotent_tool);
7789
7790 let resp = send_initialize(&mut router).await;
7791 let instructions = resp.instructions.unwrap();
7792
7793 assert!(instructions.contains("- **query**: Run a query [read-only]"));
7794 assert!(instructions.contains("- **delete**: Delete a record\n"));
7796 assert!(instructions.contains("- **upsert**: Upsert a record [idempotent]"));
7797 }
7798
7799 #[tokio::test]
7800 async fn test_auto_instructions_with_resources() {
7801 use crate::resource::ResourceBuilder;
7802
7803 let resource = ResourceBuilder::new("file:///schema.sql")
7804 .name("Schema")
7805 .description("Database schema")
7806 .text("CREATE TABLE ...");
7807
7808 let mut router = McpRouter::new().auto_instructions().resource(resource);
7809
7810 let resp = send_initialize(&mut router).await;
7811 let instructions = resp.instructions.unwrap();
7812
7813 assert!(instructions.contains("## Resources"));
7814 assert!(instructions.contains("- **file:///schema.sql**: Database schema"));
7815 assert!(!instructions.contains("## Tools"));
7816 }
7817
7818 #[tokio::test]
7819 async fn test_auto_instructions_with_resource_templates() {
7820 use crate::resource::ResourceTemplateBuilder;
7821
7822 let template = ResourceTemplateBuilder::new("file:///{path}")
7823 .name("File")
7824 .description("Read a file by path")
7825 .handler(
7826 |_uri: String, _vars: std::collections::HashMap<String, String>| async move {
7827 Ok(crate::ReadResourceResult::text("content", "text/plain"))
7828 },
7829 );
7830
7831 let mut router = McpRouter::new()
7832 .auto_instructions()
7833 .resource_template(template);
7834
7835 let resp = send_initialize(&mut router).await;
7836 let instructions = resp.instructions.unwrap();
7837
7838 assert!(instructions.contains("## Resources"));
7839 assert!(instructions.contains("- **file:///{path}**: Read a file by path"));
7840 }
7841
7842 #[tokio::test]
7843 async fn test_auto_instructions_with_prompts() {
7844 use crate::prompt::PromptBuilder;
7845
7846 let prompt = PromptBuilder::new("write_query")
7847 .description("Help write a SQL query")
7848 .user_message("Write a query for: {task}");
7849
7850 let mut router = McpRouter::new().auto_instructions().prompt(prompt);
7851
7852 let resp = send_initialize(&mut router).await;
7853 let instructions = resp.instructions.unwrap();
7854
7855 assert!(instructions.contains("## Prompts"));
7856 assert!(instructions.contains("- **write_query**: Help write a SQL query"));
7857 assert!(!instructions.contains("## Tools"));
7858 }
7859
7860 #[tokio::test]
7861 async fn test_auto_instructions_all_sections() {
7862 use crate::prompt::PromptBuilder;
7863 use crate::resource::ResourceBuilder;
7864
7865 let tool = ToolBuilder::new("query")
7866 .description("Execute SQL")
7867 .read_only()
7868 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7869 .build();
7870 let resource = ResourceBuilder::new("db://schema")
7871 .name("Schema")
7872 .description("Full database schema")
7873 .text("schema");
7874 let prompt = PromptBuilder::new("write_query")
7875 .description("Help write a SQL query")
7876 .user_message("Write a query");
7877
7878 let mut router = McpRouter::new()
7879 .auto_instructions()
7880 .tool(tool)
7881 .resource(resource)
7882 .prompt(prompt);
7883
7884 let resp = send_initialize(&mut router).await;
7885 let instructions = resp.instructions.unwrap();
7886
7887 assert!(instructions.contains("## Tools"));
7889 assert!(instructions.contains("## Resources"));
7890 assert!(instructions.contains("## Prompts"));
7891
7892 let tools_pos = instructions.find("## Tools").unwrap();
7894 let resources_pos = instructions.find("## Resources").unwrap();
7895 let prompts_pos = instructions.find("## Prompts").unwrap();
7896 assert!(tools_pos < resources_pos);
7897 assert!(resources_pos < prompts_pos);
7898 }
7899
7900 #[tokio::test]
7901 async fn test_auto_instructions_with_prefix_and_suffix() {
7902 let tool = ToolBuilder::new("echo")
7903 .description("Echo input")
7904 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7905 .build();
7906
7907 let mut router = McpRouter::new()
7908 .auto_instructions_with(
7909 Some("This server provides echo capabilities."),
7910 Some("Contact admin@example.com for support."),
7911 )
7912 .tool(tool);
7913
7914 let resp = send_initialize(&mut router).await;
7915 let instructions = resp.instructions.unwrap();
7916
7917 assert!(instructions.starts_with("This server provides echo capabilities."));
7918 assert!(instructions.ends_with("Contact admin@example.com for support."));
7919 assert!(instructions.contains("## Tools"));
7920 assert!(instructions.contains("- **echo**: Echo input"));
7921 }
7922
7923 #[tokio::test]
7924 async fn test_auto_instructions_prefix_only() {
7925 let tool = ToolBuilder::new("echo")
7926 .description("Echo input")
7927 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7928 .build();
7929
7930 let mut router = McpRouter::new()
7931 .auto_instructions_with(Some("My server intro."), None::<String>)
7932 .tool(tool);
7933
7934 let resp = send_initialize(&mut router).await;
7935 let instructions = resp.instructions.unwrap();
7936
7937 assert!(instructions.starts_with("My server intro."));
7938 assert!(instructions.contains("- **echo**: Echo input"));
7939 }
7940
7941 #[tokio::test]
7942 async fn test_auto_instructions_empty_router() {
7943 let mut router = McpRouter::new().auto_instructions();
7944
7945 let resp = send_initialize(&mut router).await;
7946 let instructions = resp.instructions.expect("should have instructions");
7947
7948 assert!(!instructions.contains("## Tools"));
7950 assert!(!instructions.contains("## Resources"));
7951 assert!(!instructions.contains("## Prompts"));
7952 assert!(instructions.is_empty());
7953 }
7954
7955 #[tokio::test]
7956 async fn test_auto_instructions_overrides_manual() {
7957 let tool = ToolBuilder::new("echo")
7958 .description("Echo input")
7959 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7960 .build();
7961
7962 let mut router = McpRouter::new()
7963 .instructions("This will be overridden")
7964 .auto_instructions()
7965 .tool(tool);
7966
7967 let resp = send_initialize(&mut router).await;
7968 let instructions = resp.instructions.unwrap();
7969
7970 assert!(!instructions.contains("This will be overridden"));
7971 assert!(instructions.contains("- **echo**: Echo input"));
7972 }
7973
7974 #[tokio::test]
7975 async fn test_no_auto_instructions_returns_manual() {
7976 let tool = ToolBuilder::new("echo")
7977 .description("Echo input")
7978 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7979 .build();
7980
7981 let mut router = McpRouter::new()
7982 .instructions("Manual instructions here")
7983 .tool(tool);
7984
7985 let resp = send_initialize(&mut router).await;
7986 let instructions = resp.instructions.unwrap();
7987
7988 assert_eq!(instructions, "Manual instructions here");
7989 }
7990
7991 #[tokio::test]
7992 async fn test_auto_instructions_no_description_fallback() {
7993 let tool = ToolBuilder::new("mystery")
7994 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7995 .build();
7996
7997 let mut router = McpRouter::new().auto_instructions().tool(tool);
7998
7999 let resp = send_initialize(&mut router).await;
8000 let instructions = resp.instructions.unwrap();
8001
8002 assert!(instructions.contains("- **mystery**: No description"));
8003 }
8004
8005 #[tokio::test]
8006 async fn test_auto_instructions_sorted_alphabetically() {
8007 let tool_z = ToolBuilder::new("zebra")
8008 .description("Z tool")
8009 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8010 .build();
8011 let tool_a = ToolBuilder::new("alpha")
8012 .description("A tool")
8013 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8014 .build();
8015 let tool_m = ToolBuilder::new("middle")
8016 .description("M tool")
8017 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8018 .build();
8019
8020 let mut router = McpRouter::new()
8021 .auto_instructions()
8022 .tool(tool_z)
8023 .tool(tool_a)
8024 .tool(tool_m);
8025
8026 let resp = send_initialize(&mut router).await;
8027 let instructions = resp.instructions.unwrap();
8028
8029 let alpha_pos = instructions.find("**alpha**").unwrap();
8030 let middle_pos = instructions.find("**middle**").unwrap();
8031 let zebra_pos = instructions.find("**zebra**").unwrap();
8032 assert!(alpha_pos < middle_pos);
8033 assert!(middle_pos < zebra_pos);
8034 }
8035
8036 #[tokio::test]
8037 async fn test_auto_instructions_read_only_and_idempotent_tags() {
8038 let tool = ToolBuilder::new("safe_update")
8039 .description("Safe update operation")
8040 .idempotent()
8041 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8042 .build();
8043
8044 let mut router = McpRouter::new().auto_instructions().tool(tool);
8045
8046 let resp = send_initialize(&mut router).await;
8047 let instructions = resp.instructions.unwrap();
8048
8049 assert!(
8050 instructions.contains("[idempotent]"),
8051 "got: {}",
8052 instructions
8053 );
8054 }
8055
8056 #[tokio::test]
8057 async fn test_auto_instructions_lazy_generation() {
8058 let mut router = McpRouter::new().auto_instructions();
8061
8062 let tool = ToolBuilder::new("late_tool")
8063 .description("Added after auto_instructions")
8064 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8065 .build();
8066
8067 router = router.tool(tool);
8068
8069 let resp = send_initialize(&mut router).await;
8070 let instructions = resp.instructions.unwrap();
8071
8072 assert!(instructions.contains("- **late_tool**: Added after auto_instructions"));
8073 }
8074
8075 #[tokio::test]
8076 async fn test_auto_instructions_multiple_annotation_tags() {
8077 let tool = ToolBuilder::new("update")
8078 .description("Update a record")
8079 .annotations(ToolAnnotations {
8080 read_only_hint: true,
8081 idempotent_hint: true,
8082 ..Default::default()
8083 })
8084 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8085 .build();
8086
8087 let mut router = McpRouter::new().auto_instructions().tool(tool);
8088
8089 let resp = send_initialize(&mut router).await;
8090 let instructions = resp.instructions.unwrap();
8091
8092 assert!(
8093 instructions.contains("[read-only, idempotent]"),
8094 "got: {}",
8095 instructions
8096 );
8097 }
8098
8099 #[tokio::test]
8100 async fn test_auto_instructions_no_annotations_no_tags() {
8101 let tool = ToolBuilder::new("fetch")
8103 .description("Fetch data")
8104 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
8105 .build();
8106
8107 let mut router = McpRouter::new().auto_instructions().tool(tool);
8108
8109 let resp = send_initialize(&mut router).await;
8110 let instructions = resp.instructions.unwrap();
8111
8112 assert!(
8114 !instructions.contains('['),
8115 "should have no tags, got: {}",
8116 instructions
8117 );
8118 assert!(instructions.contains("- **fetch**: Fetch data"));
8119 }
8120
8121 async fn send_initialize(router: &mut McpRouter) -> InitializeResult {
8123 let init_req = RouterRequest {
8124 id: RequestId::Number(0),
8125 inner: McpRequest::Initialize(InitializeParams {
8126 protocol_version: "2025-11-25".to_string(),
8127 capabilities: ClientCapabilities {
8128 roots: None,
8129 sampling: None,
8130 elicitation: None,
8131 tasks: None,
8132 experimental: None,
8133 extensions: None,
8134 },
8135 client_info: Implementation {
8136 name: "test".to_string(),
8137 version: "1.0".to_string(),
8138 ..Default::default()
8139 },
8140 meta: None,
8141 }),
8142 extensions: Extensions::new(),
8143 };
8144 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8145 match resp.inner {
8146 Ok(McpResponse::Initialize(result)) => result,
8147 other => panic!("Expected Initialize response, got {:?}", other),
8148 }
8149 }
8150
8151 #[tokio::test]
8152 async fn test_notify_tools_list_changed() {
8153 let (tx, mut rx) = crate::context::notification_channel(16);
8154
8155 let router = McpRouter::new()
8156 .server_info("test", "1.0")
8157 .with_notification_sender(tx);
8158
8159 assert!(router.notify_tools_list_changed());
8160
8161 let notification = rx.recv().await.unwrap();
8162 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8163 }
8164
8165 #[tokio::test]
8166 async fn test_notify_prompts_list_changed() {
8167 let (tx, mut rx) = crate::context::notification_channel(16);
8168
8169 let router = McpRouter::new()
8170 .server_info("test", "1.0")
8171 .with_notification_sender(tx);
8172
8173 assert!(router.notify_prompts_list_changed());
8174
8175 let notification = rx.recv().await.unwrap();
8176 assert!(matches!(
8177 notification,
8178 ServerNotification::PromptsListChanged
8179 ));
8180 }
8181
8182 #[tokio::test]
8183 async fn test_notify_without_sender_returns_false() {
8184 let router = McpRouter::new().server_info("test", "1.0");
8185
8186 assert!(!router.notify_tools_list_changed());
8187 assert!(!router.notify_prompts_list_changed());
8188 assert!(!router.notify_resources_list_changed());
8189 }
8190
8191 #[tokio::test]
8192 async fn test_list_changed_capabilities_with_notification_sender() {
8193 let (tx, _rx) = crate::context::notification_channel(16);
8194 let tool = ToolBuilder::new("test")
8195 .description("test")
8196 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8197 .build();
8198
8199 let mut router = McpRouter::new()
8200 .server_info("test", "1.0")
8201 .tool(tool)
8202 .with_notification_sender(tx);
8203
8204 init_router(&mut router).await;
8205
8206 let caps = router.capabilities();
8207 let tools_cap = caps.tools.expect("tools capability should be present");
8208 assert!(
8209 tools_cap.list_changed,
8210 "tools.listChanged should be true when notification sender is configured"
8211 );
8212 }
8213
8214 #[tokio::test]
8215 async fn test_list_changed_capabilities_without_notification_sender() {
8216 let tool = ToolBuilder::new("test")
8217 .description("test")
8218 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8219 .build();
8220
8221 let mut router = McpRouter::new().server_info("test", "1.0").tool(tool);
8222
8223 init_router(&mut router).await;
8224
8225 let caps = router.capabilities();
8226 let tools_cap = caps.tools.expect("tools capability should be present");
8227 assert!(
8228 !tools_cap.list_changed,
8229 "tools.listChanged should be false without notification sender"
8230 );
8231 }
8232
8233 #[tokio::test]
8234 async fn test_set_logging_level_filters_messages() {
8235 let (tx, mut rx) = crate::context::notification_channel(16);
8236
8237 let mut router = McpRouter::new()
8238 .server_info("test", "1.0")
8239 .with_notification_sender(tx);
8240
8241 init_router(&mut router).await;
8242
8243 let set_level_req = RouterRequest {
8245 id: RequestId::Number(99),
8246 inner: McpRequest::SetLoggingLevel(SetLogLevelParams {
8247 level: LogLevel::Warning,
8248 meta: None,
8249 }),
8250 extensions: crate::context::Extensions::new(),
8251 };
8252 let resp = router
8253 .ready()
8254 .await
8255 .unwrap()
8256 .call(set_level_req)
8257 .await
8258 .unwrap();
8259 assert!(matches!(resp.inner, Ok(McpResponse::SetLoggingLevel(_))));
8260
8261 let ctx = router.create_context(RequestId::Number(100), None);
8263
8264 ctx.send_log(LoggingMessageParams::new(
8266 LogLevel::Error,
8267 serde_json::Value::Null,
8268 ));
8269 assert!(
8270 rx.try_recv().is_ok(),
8271 "Error should pass through Warning filter"
8272 );
8273
8274 ctx.send_log(LoggingMessageParams::new(
8276 LogLevel::Info,
8277 serde_json::Value::Null,
8278 ));
8279 assert!(
8280 rx.try_recv().is_err(),
8281 "Info should be filtered at Warning level"
8282 );
8283 }
8284
8285 #[test]
8286 fn test_paginate_no_page_size() {
8287 let items = vec![1, 2, 3, 4, 5];
8288 let (page, cursor) = paginate(items.clone(), None, None).unwrap();
8289 assert_eq!(page, items);
8290 assert!(cursor.is_none());
8291 }
8292
8293 #[test]
8294 fn test_paginate_first_page() {
8295 let items = vec![1, 2, 3, 4, 5];
8296 let (page, cursor) = paginate(items, None, Some(2)).unwrap();
8297 assert_eq!(page, vec![1, 2]);
8298 assert!(cursor.is_some());
8299 }
8300
8301 #[test]
8302 fn test_paginate_middle_page() {
8303 let items = vec![1, 2, 3, 4, 5];
8304 let (page1, cursor1) = paginate(items.clone(), None, Some(2)).unwrap();
8305 assert_eq!(page1, vec![1, 2]);
8306
8307 let (page2, cursor2) = paginate(items, cursor1.as_deref(), Some(2)).unwrap();
8308 assert_eq!(page2, vec![3, 4]);
8309 assert!(cursor2.is_some());
8310 }
8311
8312 #[test]
8313 fn test_paginate_last_page() {
8314 let items = vec![1, 2, 3, 4, 5];
8315 let cursor = encode_cursor(4);
8317 let (page, next) = paginate(items, Some(&cursor), Some(2)).unwrap();
8318 assert_eq!(page, vec![5]);
8319 assert!(next.is_none());
8320 }
8321
8322 #[test]
8323 fn test_paginate_exact_boundary() {
8324 let items = vec![1, 2, 3, 4];
8325 let (page, cursor) = paginate(items, None, Some(4)).unwrap();
8326 assert_eq!(page, vec![1, 2, 3, 4]);
8327 assert!(cursor.is_none());
8328 }
8329
8330 #[test]
8331 fn test_paginate_invalid_cursor() {
8332 let items = vec![1, 2, 3];
8333 let result = paginate(items, Some("not-valid-base64!@#$"), Some(2));
8334 assert!(result.is_err());
8335 }
8336
8337 #[test]
8338 fn test_cursor_round_trip() {
8339 let offset = 42;
8340 let encoded = encode_cursor(offset);
8341 let decoded = decode_cursor(&encoded).unwrap();
8342 assert_eq!(decoded, offset);
8343 }
8344
8345 #[tokio::test]
8346 async fn test_list_tools_pagination() {
8347 let tool_a = ToolBuilder::new("alpha")
8348 .description("a")
8349 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8350 .build();
8351 let tool_b = ToolBuilder::new("beta")
8352 .description("b")
8353 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8354 .build();
8355 let tool_c = ToolBuilder::new("gamma")
8356 .description("c")
8357 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8358 .build();
8359
8360 let mut router = McpRouter::new()
8361 .server_info("test", "1.0")
8362 .page_size(2)
8363 .tool(tool_a)
8364 .tool(tool_b)
8365 .tool(tool_c);
8366
8367 init_router(&mut router).await;
8368
8369 let req = RouterRequest {
8371 id: RequestId::Number(1),
8372 inner: McpRequest::ListTools(ListToolsParams {
8373 cursor: None,
8374 meta: None,
8375 }),
8376 extensions: Extensions::new(),
8377 };
8378 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8379 let (tools, next_cursor) = match resp.inner {
8380 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8381 other => panic!("Expected ListTools, got {:?}", other),
8382 };
8383 assert_eq!(tools.len(), 2);
8384 assert_eq!(tools[0].name, "alpha");
8385 assert_eq!(tools[1].name, "beta");
8386 assert!(next_cursor.is_some());
8387
8388 let req = RouterRequest {
8390 id: RequestId::Number(2),
8391 inner: McpRequest::ListTools(ListToolsParams {
8392 cursor: next_cursor,
8393 meta: None,
8394 }),
8395 extensions: Extensions::new(),
8396 };
8397 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8398 let (tools, next_cursor) = match resp.inner {
8399 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8400 other => panic!("Expected ListTools, got {:?}", other),
8401 };
8402 assert_eq!(tools.len(), 1);
8403 assert_eq!(tools[0].name, "gamma");
8404 assert!(next_cursor.is_none());
8405 }
8406
8407 #[tokio::test]
8408 async fn test_list_tools_no_pagination_by_default() {
8409 let tool_a = ToolBuilder::new("alpha")
8410 .description("a")
8411 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8412 .build();
8413 let tool_b = ToolBuilder::new("beta")
8414 .description("b")
8415 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8416 .build();
8417
8418 let mut router = McpRouter::new()
8419 .server_info("test", "1.0")
8420 .tool(tool_a)
8421 .tool(tool_b);
8422
8423 init_router(&mut router).await;
8424
8425 let req = RouterRequest {
8426 id: RequestId::Number(1),
8427 inner: McpRequest::ListTools(ListToolsParams {
8428 cursor: None,
8429 meta: None,
8430 }),
8431 extensions: Extensions::new(),
8432 };
8433 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8434 match resp.inner {
8435 Ok(McpResponse::ListTools(result)) => {
8436 assert_eq!(result.tools.len(), 2);
8437 assert!(result.next_cursor.is_none());
8438 }
8439 other => panic!("Expected ListTools, got {:?}", other),
8440 }
8441 }
8442
8443 #[cfg(feature = "dynamic-tools")]
8448 mod dynamic_tools_tests {
8449 use super::*;
8450
8451 #[tokio::test]
8452 async fn test_dynamic_tools_register_and_list() {
8453 let (router, registry) = McpRouter::new()
8454 .server_info("test", "1.0")
8455 .with_dynamic_tools();
8456
8457 let tool = ToolBuilder::new("dynamic_echo")
8458 .description("Dynamic echo")
8459 .handler(|input: AddInput| async move {
8460 Ok(CallToolResult::text(format!("{}", input.a)))
8461 })
8462 .build();
8463
8464 registry.register(tool);
8465
8466 let mut router = router;
8467 init_router(&mut router).await;
8468
8469 let req = RouterRequest {
8470 id: RequestId::Number(1),
8471 inner: McpRequest::ListTools(ListToolsParams::default()),
8472 extensions: Extensions::new(),
8473 };
8474
8475 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8476 match resp.inner {
8477 Ok(McpResponse::ListTools(result)) => {
8478 assert_eq!(result.tools.len(), 1);
8479 assert_eq!(result.tools[0].name, "dynamic_echo");
8480 }
8481 _ => panic!("Expected ListTools response"),
8482 }
8483 }
8484
8485 #[tokio::test]
8486 async fn test_dynamic_tools_unregister() {
8487 let (router, registry) = McpRouter::new()
8488 .server_info("test", "1.0")
8489 .with_dynamic_tools();
8490
8491 let tool = ToolBuilder::new("temp")
8492 .description("Temporary")
8493 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8494 .build();
8495
8496 registry.register(tool);
8497 assert!(registry.contains("temp"));
8498
8499 let removed = registry.unregister("temp");
8500 assert!(removed);
8501 assert!(!registry.contains("temp"));
8502
8503 assert!(!registry.unregister("temp"));
8505
8506 let mut router = router;
8507 init_router(&mut router).await;
8508
8509 let req = RouterRequest {
8510 id: RequestId::Number(1),
8511 inner: McpRequest::ListTools(ListToolsParams::default()),
8512 extensions: Extensions::new(),
8513 };
8514
8515 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8516 match resp.inner {
8517 Ok(McpResponse::ListTools(result)) => {
8518 assert_eq!(result.tools.len(), 0);
8519 }
8520 _ => panic!("Expected ListTools response"),
8521 }
8522 }
8523
8524 #[tokio::test]
8525 async fn test_dynamic_tools_merged_with_static() {
8526 let static_tool = ToolBuilder::new("static_tool")
8527 .description("Static")
8528 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8529 .build();
8530
8531 let (router, registry) = McpRouter::new()
8532 .server_info("test", "1.0")
8533 .tool(static_tool)
8534 .with_dynamic_tools();
8535
8536 let dynamic_tool = ToolBuilder::new("dynamic_tool")
8537 .description("Dynamic")
8538 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8539 .build();
8540
8541 registry.register(dynamic_tool);
8542
8543 let mut router = router;
8544 init_router(&mut router).await;
8545
8546 let req = RouterRequest {
8547 id: RequestId::Number(1),
8548 inner: McpRequest::ListTools(ListToolsParams::default()),
8549 extensions: Extensions::new(),
8550 };
8551
8552 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8553 match resp.inner {
8554 Ok(McpResponse::ListTools(result)) => {
8555 assert_eq!(result.tools.len(), 2);
8556 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8557 assert!(names.contains(&"static_tool"));
8558 assert!(names.contains(&"dynamic_tool"));
8559 }
8560 _ => panic!("Expected ListTools response"),
8561 }
8562 }
8563
8564 #[tokio::test]
8565 async fn test_static_tools_shadow_dynamic() {
8566 let static_tool = ToolBuilder::new("shared")
8567 .description("Static version")
8568 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8569 .build();
8570
8571 let (router, registry) = McpRouter::new()
8572 .server_info("test", "1.0")
8573 .tool(static_tool)
8574 .with_dynamic_tools();
8575
8576 let dynamic_tool = ToolBuilder::new("shared")
8577 .description("Dynamic version")
8578 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8579 .build();
8580
8581 registry.register(dynamic_tool);
8582
8583 let mut router = router;
8584 init_router(&mut router).await;
8585
8586 let req = RouterRequest {
8588 id: RequestId::Number(1),
8589 inner: McpRequest::ListTools(ListToolsParams::default()),
8590 extensions: Extensions::new(),
8591 };
8592
8593 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8594 match resp.inner {
8595 Ok(McpResponse::ListTools(result)) => {
8596 assert_eq!(result.tools.len(), 1);
8597 assert_eq!(result.tools[0].name, "shared");
8598 assert_eq!(
8599 result.tools[0].description.as_deref(),
8600 Some("Static version")
8601 );
8602 }
8603 _ => panic!("Expected ListTools response"),
8604 }
8605
8606 let req = RouterRequest {
8608 id: RequestId::Number(2),
8609 inner: McpRequest::CallTool(CallToolParams {
8610 input_responses: None,
8611 request_state: None,
8612 name: "shared".to_string(),
8613 arguments: serde_json::json!({"a": 1, "b": 2}),
8614 meta: None,
8615 task: None,
8616 }),
8617 extensions: Extensions::new(),
8618 };
8619
8620 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8621 match resp.inner {
8622 Ok(McpResponse::CallTool(result)) => {
8623 assert!(!result.is_error);
8624 match &result.content[0] {
8625 Content::Text { text, .. } => assert_eq!(text, "static"),
8626 _ => panic!("Expected text content"),
8627 }
8628 }
8629 _ => panic!("Expected CallTool response"),
8630 }
8631 }
8632
8633 #[tokio::test]
8634 async fn test_dynamic_tools_call() {
8635 let (router, registry) = McpRouter::new()
8636 .server_info("test", "1.0")
8637 .with_dynamic_tools();
8638
8639 let tool = ToolBuilder::new("add")
8640 .description("Add two numbers")
8641 .handler(|input: AddInput| async move {
8642 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
8643 })
8644 .build();
8645
8646 registry.register(tool);
8647
8648 let mut router = router;
8649 init_router(&mut router).await;
8650
8651 let req = RouterRequest {
8652 id: RequestId::Number(1),
8653 inner: McpRequest::CallTool(CallToolParams {
8654 input_responses: None,
8655 request_state: None,
8656 name: "add".to_string(),
8657 arguments: serde_json::json!({"a": 3, "b": 4}),
8658 meta: None,
8659 task: None,
8660 }),
8661 extensions: Extensions::new(),
8662 };
8663
8664 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8665 match resp.inner {
8666 Ok(McpResponse::CallTool(result)) => {
8667 assert!(!result.is_error);
8668 match &result.content[0] {
8669 Content::Text { text, .. } => assert_eq!(text, "7"),
8670 _ => panic!("Expected text content"),
8671 }
8672 }
8673 _ => panic!("Expected CallTool response"),
8674 }
8675 }
8676
8677 #[tokio::test]
8678 async fn test_dynamic_tools_notification_on_register() {
8679 let (tx, mut rx) = crate::context::notification_channel(16);
8680 let (router, registry) = McpRouter::new()
8681 .server_info("test", "1.0")
8682 .with_dynamic_tools();
8683 let _router = router.with_notification_sender(tx);
8684
8685 let tool = ToolBuilder::new("notified")
8686 .description("Test")
8687 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8688 .build();
8689
8690 registry.register(tool);
8691
8692 let notification = rx.recv().await.unwrap();
8693 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8694 }
8695
8696 #[tokio::test]
8697 async fn test_dynamic_tools_notification_on_unregister() {
8698 let (tx, mut rx) = crate::context::notification_channel(16);
8699 let (router, registry) = McpRouter::new()
8700 .server_info("test", "1.0")
8701 .with_dynamic_tools();
8702 let _router = router.with_notification_sender(tx);
8703
8704 let tool = ToolBuilder::new("notified")
8705 .description("Test")
8706 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8707 .build();
8708
8709 registry.register(tool);
8710 let _ = rx.recv().await.unwrap();
8712
8713 registry.unregister("notified");
8714 let notification = rx.recv().await.unwrap();
8715 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8716 }
8717
8718 #[tokio::test]
8719 async fn test_dynamic_tools_no_notification_on_empty_unregister() {
8720 let (tx, mut rx) = crate::context::notification_channel(16);
8721 let (router, registry) = McpRouter::new()
8722 .server_info("test", "1.0")
8723 .with_dynamic_tools();
8724 let _router = router.with_notification_sender(tx);
8725
8726 assert!(!registry.unregister("nonexistent"));
8728
8729 assert!(rx.try_recv().is_err());
8731 }
8732
8733 #[tokio::test]
8734 async fn test_dynamic_tools_filter_applies() {
8735 use crate::filter::CapabilityFilter;
8736
8737 let (router, registry) = McpRouter::new()
8738 .server_info("test", "1.0")
8739 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| {
8740 tool.name != "hidden"
8741 }))
8742 .with_dynamic_tools();
8743
8744 let visible = ToolBuilder::new("visible")
8745 .description("Visible")
8746 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8747 .build();
8748
8749 let hidden = ToolBuilder::new("hidden")
8750 .description("Hidden")
8751 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8752 .build();
8753
8754 registry.register(visible);
8755 registry.register(hidden);
8756
8757 let mut router = router;
8758 init_router(&mut router).await;
8759
8760 let req = RouterRequest {
8762 id: RequestId::Number(1),
8763 inner: McpRequest::ListTools(ListToolsParams::default()),
8764 extensions: Extensions::new(),
8765 };
8766
8767 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8768 match resp.inner {
8769 Ok(McpResponse::ListTools(result)) => {
8770 assert_eq!(result.tools.len(), 1);
8771 assert_eq!(result.tools[0].name, "visible");
8772 }
8773 _ => panic!("Expected ListTools response"),
8774 }
8775
8776 let req = RouterRequest {
8778 id: RequestId::Number(2),
8779 inner: McpRequest::CallTool(CallToolParams {
8780 input_responses: None,
8781 request_state: None,
8782 name: "hidden".to_string(),
8783 arguments: serde_json::json!({"a": 1, "b": 2}),
8784 meta: None,
8785 task: None,
8786 }),
8787 extensions: Extensions::new(),
8788 };
8789
8790 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8791 match resp.inner {
8792 Err(e) => {
8793 assert_eq!(e.code, -32601); }
8795 _ => panic!("Expected JsonRpc error"),
8796 }
8797 }
8798
8799 #[tokio::test]
8800 async fn test_dynamic_tools_capabilities_advertised() {
8801 let (mut router, _registry) = McpRouter::new()
8803 .server_info("test", "1.0")
8804 .with_dynamic_tools();
8805
8806 let init_req = RouterRequest {
8807 id: RequestId::Number(1),
8808 inner: McpRequest::Initialize(InitializeParams {
8809 protocol_version: "2025-11-25".to_string(),
8810 capabilities: ClientCapabilities::default(),
8811 client_info: Implementation {
8812 name: "test".to_string(),
8813 version: "1.0".to_string(),
8814 ..Default::default()
8815 },
8816 meta: None,
8817 }),
8818 extensions: Extensions::new(),
8819 };
8820
8821 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8822 match resp.inner {
8823 Ok(McpResponse::Initialize(result)) => {
8824 assert!(result.capabilities.tools.is_some());
8825 }
8826 _ => panic!("Expected Initialize response"),
8827 }
8828 }
8829
8830 #[tokio::test]
8831 async fn test_dynamic_tools_multi_session_notification() {
8832 let (tx1, mut rx1) = crate::context::notification_channel(16);
8833 let (tx2, mut rx2) = crate::context::notification_channel(16);
8834
8835 let (router, registry) = McpRouter::new()
8836 .server_info("test", "1.0")
8837 .with_dynamic_tools();
8838
8839 let _session1 = router.clone().with_notification_sender(tx1);
8841 let _session2 = router.clone().with_notification_sender(tx2);
8842
8843 let tool = ToolBuilder::new("broadcast")
8844 .description("Test")
8845 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8846 .build();
8847
8848 registry.register(tool);
8849
8850 let n1 = rx1.recv().await.unwrap();
8852 let n2 = rx2.recv().await.unwrap();
8853 assert!(matches!(n1, ServerNotification::ToolsListChanged));
8854 assert!(matches!(n2, ServerNotification::ToolsListChanged));
8855 }
8856
8857 #[tokio::test]
8858 async fn test_dynamic_tools_call_not_found() {
8859 let (router, _registry) = McpRouter::new()
8860 .server_info("test", "1.0")
8861 .with_dynamic_tools();
8862
8863 let mut router = router;
8864 init_router(&mut router).await;
8865
8866 let req = RouterRequest {
8867 id: RequestId::Number(1),
8868 inner: McpRequest::CallTool(CallToolParams {
8869 input_responses: None,
8870 request_state: None,
8871 name: "nonexistent".to_string(),
8872 arguments: serde_json::json!({}),
8873 meta: None,
8874 task: None,
8875 }),
8876 extensions: Extensions::new(),
8877 };
8878
8879 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8880 match resp.inner {
8881 Err(e) => {
8882 assert_eq!(e.code, -32601);
8883 }
8884 _ => panic!("Expected method not found error"),
8885 }
8886 }
8887
8888 #[tokio::test]
8889 async fn test_dynamic_tools_registry_list() {
8890 let (_, registry) = McpRouter::new()
8891 .server_info("test", "1.0")
8892 .with_dynamic_tools();
8893
8894 assert!(registry.list().is_empty());
8895
8896 let tool = ToolBuilder::new("tool_a")
8897 .description("A")
8898 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8899 .build();
8900 registry.register(tool);
8901
8902 let tool = ToolBuilder::new("tool_b")
8903 .description("B")
8904 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8905 .build();
8906 registry.register(tool);
8907
8908 let tools = registry.list();
8909 assert_eq!(tools.len(), 2);
8910 let names: Vec<&str> = tools.iter().map(|t| t.name.as_str()).collect();
8911 assert!(names.contains(&"tool_a"));
8912 assert!(names.contains(&"tool_b"));
8913 }
8914 } #[tokio::test]
8917 async fn test_tool_if_true_registers() {
8918 let tool = ToolBuilder::new("conditional")
8919 .description("Conditional tool")
8920 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8921 .build();
8922
8923 let mut router = McpRouter::new().tool_if(true, tool);
8924 init_router(&mut router).await;
8925
8926 let req = RouterRequest {
8927 id: RequestId::Number(1),
8928 inner: McpRequest::ListTools(ListToolsParams::default()),
8929 extensions: Extensions::new(),
8930 };
8931 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8932 match resp.inner {
8933 Ok(McpResponse::ListTools(result)) => {
8934 assert_eq!(result.tools.len(), 1);
8935 assert_eq!(result.tools[0].name, "conditional");
8936 }
8937 _ => panic!("Expected ListTools response"),
8938 }
8939 }
8940
8941 #[tokio::test]
8942 async fn test_tool_if_false_skips() {
8943 let tool = ToolBuilder::new("conditional")
8944 .description("Conditional tool")
8945 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8946 .build();
8947
8948 let mut router = McpRouter::new().tool_if(false, tool);
8949 init_router(&mut router).await;
8950
8951 let req = RouterRequest {
8952 id: RequestId::Number(1),
8953 inner: McpRequest::ListTools(ListToolsParams::default()),
8954 extensions: Extensions::new(),
8955 };
8956 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8957 match resp.inner {
8958 Ok(McpResponse::ListTools(result)) => {
8959 assert_eq!(result.tools.len(), 0);
8960 }
8961 _ => panic!("Expected ListTools response"),
8962 }
8963 }
8964
8965 #[tokio::test]
8966 async fn test_tools_if_batch_conditional() {
8967 let tools = vec![
8968 ToolBuilder::new("a")
8969 .description("Tool A")
8970 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8971 .build(),
8972 ToolBuilder::new("b")
8973 .description("Tool B")
8974 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8975 .build(),
8976 ];
8977
8978 let mut router = McpRouter::new().tools_if(false, tools);
8979 init_router(&mut router).await;
8980
8981 let req = RouterRequest {
8982 id: RequestId::Number(1),
8983 inner: McpRequest::ListTools(ListToolsParams::default()),
8984 extensions: Extensions::new(),
8985 };
8986 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8987 match resp.inner {
8988 Ok(McpResponse::ListTools(result)) => {
8989 assert_eq!(result.tools.len(), 0);
8990 }
8991 _ => panic!("Expected ListTools response"),
8992 }
8993 }
8994
8995 #[test]
8996 fn test_resource_if_true_registers() {
8997 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8998 .name("test")
8999 .text("hello");
9000
9001 let router = McpRouter::new().resource_if(true, resource);
9002 assert_eq!(router.inner.resources.len(), 1);
9003 }
9004
9005 #[test]
9006 fn test_resource_if_false_skips() {
9007 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
9008 .name("test")
9009 .text("hello");
9010
9011 let router = McpRouter::new().resource_if(false, resource);
9012 assert_eq!(router.inner.resources.len(), 0);
9013 }
9014
9015 #[test]
9016 fn test_prompt_if_true_registers() {
9017 let prompt = crate::prompt::PromptBuilder::new("greet")
9018 .description("Greeting")
9019 .user_message("Hello!");
9020
9021 let router = McpRouter::new().prompt_if(true, prompt);
9022 assert_eq!(router.inner.prompts.len(), 1);
9023 }
9024
9025 #[test]
9026 fn test_prompt_if_false_skips() {
9027 let prompt = crate::prompt::PromptBuilder::new("greet")
9028 .description("Greeting")
9029 .user_message("Hello!");
9030
9031 let router = McpRouter::new().prompt_if(false, prompt);
9032 assert_eq!(router.inner.prompts.len(), 0);
9033 }
9034
9035 #[tokio::test]
9036 async fn test_disable_tool_hides_from_list() {
9037 let safe = ToolBuilder::new("safe")
9038 .description("Safe tool")
9039 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9040 .build();
9041 let dangerous = ToolBuilder::new("dangerous")
9042 .description("Dangerous tool")
9043 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9044 .build();
9045 let mut router = McpRouter::new().tool(safe).tool(dangerous);
9046 init_router(&mut router).await;
9047
9048 router.disable_tool("dangerous");
9049 assert!(router.is_tool_enabled("safe"));
9050 assert!(!router.is_tool_enabled("dangerous"));
9051
9052 let req = RouterRequest {
9053 id: RequestId::Number(1),
9054 inner: McpRequest::ListTools(ListToolsParams::default()),
9055 extensions: Extensions::new(),
9056 };
9057 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9058 match resp.inner {
9059 Ok(McpResponse::ListTools(result)) => {
9060 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
9061 assert_eq!(names, vec!["safe"]);
9062 }
9063 _ => panic!("Expected ListTools response"),
9064 }
9065 }
9066
9067 #[tokio::test]
9068 async fn test_disable_tool_blocks_call() {
9069 let dangerous = ToolBuilder::new("dangerous")
9070 .description("Dangerous tool")
9071 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
9072 .build();
9073 let mut router = McpRouter::new().tool(dangerous);
9074 init_router(&mut router).await;
9075
9076 router.disable_tool("dangerous");
9077
9078 let req = RouterRequest {
9079 id: RequestId::Number(2),
9080 inner: McpRequest::CallTool(CallToolParams {
9081 input_responses: None,
9082 request_state: None,
9083 name: "dangerous".to_string(),
9084 arguments: serde_json::json!({"a": 1, "b": 2}),
9085 meta: None,
9086 task: None,
9087 }),
9088 extensions: Extensions::new(),
9089 };
9090 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9091 let err = resp.inner.expect_err("disabled tool should error");
9092 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
9093 }
9094
9095 #[tokio::test]
9096 async fn test_enable_tool_restores_visibility() {
9097 let tool = ToolBuilder::new("flippy")
9098 .description("Toggleable tool")
9099 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
9100 .build();
9101 let mut router = McpRouter::new().tool(tool);
9102 init_router(&mut router).await;
9103
9104 router.disable_tool("flippy");
9105 router.enable_tool("flippy");
9106 assert!(router.is_tool_enabled("flippy"));
9107
9108 let req = RouterRequest {
9109 id: RequestId::Number(3),
9110 inner: McpRequest::CallTool(CallToolParams {
9111 input_responses: None,
9112 request_state: None,
9113 name: "flippy".to_string(),
9114 arguments: serde_json::json!({"a": 1, "b": 2}),
9115 meta: None,
9116 task: None,
9117 }),
9118 extensions: Extensions::new(),
9119 };
9120 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9121 match resp.inner {
9122 Ok(McpResponse::CallTool(result)) => {
9123 assert_eq!(result.first_text(), Some("ran"));
9124 }
9125 _ => panic!("Expected CallTool response"),
9126 }
9127 }
9128
9129 #[tokio::test]
9130 async fn test_disable_propagates_through_fresh_session() {
9131 let tool = ToolBuilder::new("shared")
9132 .description("Shared across sessions")
9133 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
9134 .build();
9135 let router = McpRouter::new().tool(tool);
9136
9137 router.disable_tool("shared");
9139 let mut child = router.with_fresh_session();
9140 init_router(&mut child).await;
9141 assert!(!child.is_tool_enabled("shared"));
9142
9143 let req = RouterRequest {
9144 id: RequestId::Number(4),
9145 inner: McpRequest::ListTools(ListToolsParams::default()),
9146 extensions: Extensions::new(),
9147 };
9148 let resp = child.ready().await.unwrap().call(req).await.unwrap();
9149 match resp.inner {
9150 Ok(McpResponse::ListTools(result)) => {
9151 assert!(result.tools.is_empty());
9152 }
9153 _ => panic!("Expected ListTools response"),
9154 }
9155 }
9156
9157 #[tokio::test]
9158 async fn test_disable_resource_and_prompt() {
9159 let resource = crate::resource::ResourceBuilder::new("file:///hidden.txt")
9160 .name("hidden")
9161 .text("secret");
9162 let prompt = crate::prompt::PromptBuilder::new("hidden_prompt")
9163 .description("hidden")
9164 .user_message("hello");
9165
9166 let mut router = McpRouter::new().resource(resource).prompt(prompt);
9167 init_router(&mut router).await;
9168
9169 router.disable_resource("file:///hidden.txt");
9170 router.disable_prompt("hidden_prompt");
9171 assert!(!router.is_resource_enabled("file:///hidden.txt"));
9172 assert!(!router.is_prompt_enabled("hidden_prompt"));
9173
9174 let req = RouterRequest {
9176 id: RequestId::Number(5),
9177 inner: McpRequest::ListResources(ListResourcesParams::default()),
9178 extensions: Extensions::new(),
9179 };
9180 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9181 match resp.inner {
9182 Ok(McpResponse::ListResources(result)) => {
9183 assert!(result.resources.is_empty());
9184 }
9185 _ => panic!("Expected ListResources response"),
9186 }
9187
9188 let req = RouterRequest {
9190 id: RequestId::Number(6),
9191 inner: McpRequest::ReadResource(ReadResourceParams {
9192 input_responses: None,
9193 request_state: None,
9194 uri: "file:///hidden.txt".to_string(),
9195 meta: None,
9196 }),
9197 extensions: Extensions::new(),
9198 };
9199 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9200 let err = resp.inner.expect_err("disabled resource should error");
9201 assert_eq!(err.code, -32602); let req = RouterRequest {
9205 id: RequestId::Number(7),
9206 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
9207 extensions: Extensions::new(),
9208 };
9209 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9210 match resp.inner {
9211 Ok(McpResponse::ListPrompts(result)) => {
9212 assert!(result.prompts.is_empty());
9213 }
9214 _ => panic!("Expected ListPrompts response"),
9215 }
9216
9217 let req = RouterRequest {
9219 id: RequestId::Number(8),
9220 inner: McpRequest::GetPrompt(GetPromptParams {
9221 input_responses: None,
9222 request_state: None,
9223 name: "hidden_prompt".to_string(),
9224 arguments: Default::default(),
9225 meta: None,
9226 }),
9227 extensions: Extensions::new(),
9228 };
9229 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9230 let err = resp.inner.expect_err("disabled prompt should error");
9231 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
9232 }
9233
9234 #[test]
9235 fn test_router_request_new() {
9236 let req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9237 assert_eq!(req.id, RequestId::Number(1));
9238 assert!(req.extensions.is_empty());
9239 }
9240
9241 #[test]
9242 fn test_with_inner_preserves_extensions() {
9243 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9244 req.extensions.insert(42u32);
9245
9246 let rewritten = req.with_inner(McpRequest::ListTools(Default::default()));
9247 assert!(matches!(rewritten.inner, McpRequest::ListTools(_)));
9248 assert_eq!(rewritten.id, RequestId::Number(1));
9249 assert_eq!(rewritten.extensions.get::<u32>(), Some(&42));
9250 }
9251
9252 #[test]
9253 fn test_with_id_and_inner_preserves_extensions() {
9254 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9255 req.extensions.insert(String::from("token-abc"));
9256
9257 let rewritten = req.with_id_and_inner(
9258 RequestId::Number(99),
9259 McpRequest::ListResources(Default::default()),
9260 );
9261 assert_eq!(rewritten.id, RequestId::Number(99));
9262 assert!(matches!(rewritten.inner, McpRequest::ListResources(_)));
9263 assert_eq!(
9264 rewritten.extensions.get::<String>(),
9265 Some(&String::from("token-abc"))
9266 );
9267 }
9268
9269 #[test]
9270 fn test_clone_with_inner_preserves_extensions() {
9271 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9272 req.extensions.insert(true);
9273
9274 let cloned = req.clone_with_inner(McpRequest::ListTools(Default::default()));
9275
9276 assert!(matches!(req.inner, McpRequest::Ping));
9278 assert_eq!(req.extensions.get::<bool>(), Some(&true));
9279
9280 assert!(matches!(cloned.inner, McpRequest::ListTools(_)));
9282 assert_eq!(cloned.extensions.get::<bool>(), Some(&true));
9283 }
9284
9285 #[test]
9286 fn test_router_response_is_error() {
9287 let ok_resp = RouterResponse {
9288 id: RequestId::Number(1),
9289 inner: Ok(McpResponse::Pong(Default::default())),
9290 };
9291 assert!(!ok_resp.is_error());
9292
9293 let err_resp = RouterResponse {
9294 id: RequestId::Number(2),
9295 inner: Err(JsonRpcError::internal_error("boom")),
9296 };
9297 assert!(err_resp.is_error());
9298 }
9299
9300 #[test]
9301 fn test_extensions_len_and_is_empty() {
9302 let mut ext = Extensions::new();
9303 assert!(ext.is_empty());
9304 assert_eq!(ext.len(), 0);
9305
9306 ext.insert(42u32);
9307 assert!(!ext.is_empty());
9308 assert_eq!(ext.len(), 1);
9309
9310 ext.insert(String::from("hello"));
9311 assert_eq!(ext.len(), 2);
9312 }
9313
9314 #[test]
9315 fn test_router_response_serde_roundtrip() {
9316 let response = RouterResponse {
9318 id: RequestId::Number(1),
9319 inner: Ok(McpResponse::Empty(EmptyResult {})),
9320 };
9321 let json = serde_json::to_string(&response).unwrap();
9322 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9323 assert_eq!(deserialized.id, RequestId::Number(1));
9324 assert!(!deserialized.is_error());
9325
9326 let response = RouterResponse {
9328 id: RequestId::String("req-2".into()),
9329 inner: Err(JsonRpcError::method_not_found("unknown")),
9330 };
9331 let json = serde_json::to_string(&response).unwrap();
9332 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9333 assert_eq!(deserialized.id, RequestId::String("req-2".into()));
9334 assert!(deserialized.is_error());
9335 }
9336
9337 #[tokio::test]
9344 async fn test_discover_dispatch_via_jsonrpc_service() {
9345 let router = McpRouter::new().server_info("unit-test-server", "4.2.0");
9348 let mut service = JsonRpcService::new(router);
9349
9350 let req = JsonRpcRequest::new(1, "server/discover");
9351 let resp = service.call_single(req).await.unwrap();
9352
9353 match resp {
9354 JsonRpcResponse::Result(r) => {
9355 let versions = r
9357 .result
9358 .get("supportedVersions")
9359 .and_then(|v| v.as_array())
9360 .expect("result.supportedVersions must be an array");
9361 assert!(!versions.is_empty(), "supportedVersions must not be empty");
9362
9363 assert_eq!(
9365 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
9366 "unit-test-server",
9367 "serverInfo.name must match configured value"
9368 );
9369 assert_eq!(
9370 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["version"], "4.2.0",
9371 "serverInfo.version must match configured value"
9372 );
9373
9374 assert!(
9377 r.result.get("protocolVersion").is_none(),
9378 "server/discover must NOT include protocolVersion: {:?}",
9379 r.result
9380 );
9381 }
9382 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
9383 _ => panic!("unexpected response variant"),
9384 }
9385 }
9386
9387 #[tokio::test]
9388 async fn test_discover_does_not_require_initialization() {
9389 let router = McpRouter::new().server_info("fresh-router", "1.0.0");
9392 let mut service = JsonRpcService::new(router);
9393
9394 let req = JsonRpcRequest::new(2, "server/discover");
9395 let resp = service.call_single(req).await.unwrap();
9396
9397 assert!(
9399 !matches!(resp, JsonRpcResponse::Error(_)),
9400 "server/discover must not require initialization: {:?}",
9401 resp
9402 );
9403 }
9404}
9405
9406#[cfg(test)]
9407mod cursor_property_tests {
9408 use super::{decode_cursor, encode_cursor};
9409 use proptest::prelude::*;
9410
9411 fn arb_cursor_text() -> BoxedStrategy<String> {
9412 prop_oneof![
9413 8 => prop::collection::vec(any::<char>(), 0..512)
9414 .prop_map(|chars| chars.into_iter().collect()),
9415 1 => Just("\0\r\n\t\u{001b}\u{007f}".repeat(64)),
9416 1 => Just("A".repeat(16 * 1024)),
9417 ]
9418 .boxed()
9419 }
9420
9421 proptest! {
9422 #![proptest_config(ProptestConfig::with_cases(512))]
9423
9424 #[test]
9426 fn cursor_round_trips(offset in any::<usize>()) {
9427 prop_assert_eq!(decode_cursor(&encode_cursor(offset)).unwrap(), offset);
9428 }
9429
9430 #[test]
9432 fn decode_cursor_never_panics(s in arb_cursor_text()) {
9433 let _ = decode_cursor(&s);
9434 }
9435 }
9436}