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 let _ = self
3520 .inner
3521 .task_store
3522 .get_task(¶ms.task_id)
3523 .await
3524 .map_err(task_store_error)?
3525 .ok_or_else(|| {
3526 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3527 "Task not found: {}",
3528 params.task_id
3529 )))
3530 })?;
3531 Ok(McpResponse::UpdateTask(EmptyResult {}))
3532 }
3533
3534 McpRequest::CancelTask(params) => {
3535 if is_final_protocol_request(&extensions) {
3536 self.require_negotiated_tasks(&extensions, "tasks/cancel")?;
3537 self.authorize_task(¶ms.task_id, &extensions).await?;
3538 self.inner
3542 .task_store
3543 .cancel_task(¶ms.task_id, params.reason.as_deref())
3544 .await
3545 .map_err(task_store_error)?
3546 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3547 self.notify_task_state(¶ms.task_id).await;
3548 return Ok(McpResponse::FinalTaskAck(
3549 crate::tasks::TaskAcknowledgement::new(),
3550 ));
3551 }
3552
3553 self.authorize_task(¶ms.task_id, &extensions).await?;
3554
3555 let current = self
3557 .inner
3558 .task_store
3559 .get_task(¶ms.task_id)
3560 .await
3561 .map_err(task_store_error)?
3562 .ok_or_else(|| {
3563 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3564 "Task not found: {}",
3565 params.task_id
3566 )))
3567 })?;
3568
3569 if current.status.is_terminal() {
3570 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
3571 "Task {} is already in terminal state: {}",
3572 params.task_id, current.status
3573 ))));
3574 }
3575
3576 self.inner
3577 .task_store
3578 .cancel_task(¶ms.task_id, params.reason.as_deref())
3579 .await
3580 .map_err(task_store_error)?
3581 .ok_or_else(|| {
3582 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3583 "Task not found: {}",
3584 params.task_id
3585 )))
3586 })?;
3587
3588 Ok(McpResponse::CancelTask(EmptyResult {}))
3592 }
3593
3594 McpRequest::SetLoggingLevel(params) => {
3595 tracing::debug!(level = ?params.level, "Client set logging level");
3596 if let Ok(mut level) = self.inner.min_log_level.write() {
3597 *level = params.level;
3598 }
3599 Ok(McpResponse::SetLoggingLevel(EmptyResult {}))
3600 }
3601
3602 McpRequest::Complete(params) => {
3603 tracing::debug!(
3604 reference = ?params.reference,
3605 argument = %params.argument.name,
3606 "Completion request"
3607 );
3608
3609 if let Some(ref handler) = self.inner.completion_handler {
3611 let result = handler(params).await?;
3612 Ok(McpResponse::Complete(result))
3613 } else {
3614 Ok(McpResponse::Complete(CompleteResult::new(vec![])))
3616 }
3617 }
3618
3619 #[cfg(feature = "stateless")]
3620 McpRequest::SubscriptionsListen(params) => {
3621 if !is_final_protocol_request(&extensions) {
3628 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
3631 "subscriptions/listen",
3632 )));
3633 }
3634 let Some(requested) = params.notifications else {
3635 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
3636 "subscriptions/listen requires a notifications filter",
3637 )));
3638 };
3639 if requested.task_ids.is_some() && !client_declares_tasks(&extensions) {
3642 return Err(Error::JsonRpc(
3643 JsonRpcError::missing_required_client_capability(
3644 tasks_client_capabilities(),
3645 ),
3646 ));
3647 }
3648 let notifications = crate::transport::subscriptions::accepted_subscription_filter(
3649 requested,
3650 self.final_tasks_enabled(),
3651 );
3652 Ok(McpResponse::SubscriptionsAccepted(
3653 crate::protocol::SubscriptionsAcceptedResult { notifications },
3654 ))
3655 }
3656
3657 McpRequest::Unknown { method, .. } => {
3658 Err(Error::JsonRpc(JsonRpcError::method_not_found(&method)))
3659 }
3660 _ => Err(Error::JsonRpc(JsonRpcError::method_not_found(
3661 "unknown method",
3662 ))),
3663 }
3664 }
3665
3666 pub fn handle_notification(&self, notification: McpNotification) {
3668 match notification {
3669 McpNotification::Initialized => {
3670 let phase_before = self.session.phase();
3671 if self.session.mark_initialized() {
3672 if phase_before == crate::session::SessionPhase::Uninitialized {
3673 tracing::info!(
3674 "Session initialized from uninitialized state (race resolved)"
3675 );
3676 } else {
3677 tracing::info!("Session initialized, entering operation phase");
3678 }
3679 } else {
3680 tracing::warn!(
3681 phase = ?self.session.phase(),
3682 "Received initialized notification in unexpected state"
3683 );
3684 }
3685 }
3686 McpNotification::Cancelled(params) => {
3687 if let Some(ref request_id) = params.request_id {
3688 if self.cancel_request(request_id) {
3689 tracing::info!(
3690 request_id = ?request_id,
3691 reason = ?params.reason,
3692 "Request cancelled"
3693 );
3694 } else {
3695 tracing::debug!(
3696 request_id = ?request_id,
3697 reason = ?params.reason,
3698 "Cancellation requested for unknown request"
3699 );
3700 }
3701 } else {
3702 tracing::debug!(
3703 reason = ?params.reason,
3704 "Cancellation notification received without request_id"
3705 );
3706 }
3707 }
3708 McpNotification::Progress(params) => {
3709 tracing::trace!(
3710 token = ?params.progress_token,
3711 progress = params.progress,
3712 total = ?params.total,
3713 "Progress notification"
3714 );
3715 }
3723 McpNotification::RootsListChanged => {
3724 tracing::info!("Client roots list changed");
3725 }
3728 McpNotification::Unknown { method, .. } => {
3729 tracing::debug!(method = %method, "Unknown notification received");
3730 }
3731 _ => {
3732 tracing::debug!("Unrecognized notification variant received");
3733 }
3734 }
3735 }
3736}
3737
3738impl Default for McpRouter {
3739 fn default() -> Self {
3740 Self::new()
3741 }
3742}
3743
3744pub use crate::context::Extensions;
3750
3751#[derive(Debug, Clone)]
3776pub struct ToolAnnotationsMap {
3777 map: Arc<HashMap<String, ToolAnnotations>>,
3778}
3779
3780impl ToolAnnotationsMap {
3781 pub fn get(&self, tool_name: &str) -> Option<&ToolAnnotations> {
3785 self.map.get(tool_name)
3786 }
3787
3788 pub fn is_read_only(&self, tool_name: &str) -> bool {
3793 self.map.get(tool_name).is_some_and(|a| a.read_only_hint)
3794 }
3795
3796 pub fn is_destructive(&self, tool_name: &str) -> bool {
3801 self.map.get(tool_name).is_none_or(|a| a.destructive_hint)
3802 }
3803
3804 pub fn is_idempotent(&self, tool_name: &str) -> bool {
3809 self.map.get(tool_name).is_some_and(|a| a.idempotent_hint)
3810 }
3811}
3812
3813#[derive(Debug, Clone)]
3835pub struct RouterRequest {
3836 pub id: RequestId,
3838 pub inner: McpRequest,
3840 pub extensions: Extensions,
3842}
3843
3844impl RouterRequest {
3845 pub fn new(id: RequestId, inner: McpRequest) -> Self {
3847 Self {
3848 id,
3849 inner,
3850 extensions: Extensions::new(),
3851 }
3852 }
3853
3854 pub fn with_inner(self, inner: McpRequest) -> Self {
3860 Self {
3861 id: self.id,
3862 inner,
3863 extensions: self.extensions,
3864 }
3865 }
3866
3867 pub fn with_id_and_inner(self, id: RequestId, inner: McpRequest) -> Self {
3873 Self {
3874 id,
3875 inner,
3876 extensions: self.extensions,
3877 }
3878 }
3879
3880 pub fn clone_with_inner(&self, inner: McpRequest) -> Self {
3888 Self {
3889 id: self.id.clone(),
3890 inner,
3891 extensions: self.extensions.clone(),
3892 }
3893 }
3894}
3895
3896#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
3898pub struct RouterResponse {
3899 pub id: RequestId,
3901 pub inner: std::result::Result<McpResponse, JsonRpcError>,
3903}
3904
3905impl RouterResponse {
3906 pub fn is_error(&self) -> bool {
3922 self.inner.is_err()
3923 }
3924
3925 pub fn into_jsonrpc(self) -> JsonRpcResponse {
3927 match self.inner {
3928 Ok(response) => match serde_json::to_value(response) {
3929 Ok(result) => JsonRpcResponse::result(self.id, result),
3930 Err(e) => {
3931 tracing::error!(error = %e, "Failed to serialize response");
3932 JsonRpcResponse::error(
3933 Some(self.id),
3934 JsonRpcError::internal_error(format!("Serialization error: {}", e)),
3935 )
3936 }
3937 },
3938 Err(error) => JsonRpcResponse::error(Some(self.id), error),
3939 }
3940 }
3941}
3942
3943impl Service<RouterRequest> for McpRouter {
3944 type Response = RouterResponse;
3945 type Error = std::convert::Infallible; type Future =
3947 Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
3948
3949 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
3950 Poll::Ready(Ok(()))
3951 }
3952
3953 fn call(&mut self, req: RouterRequest) -> Self::Future {
3954 let router = self.clone();
3955 let request_id = req.id.clone();
3956 Box::pin(async move {
3957 let result = router.handle(req.id, req.inner, req.extensions).await;
3958 router.complete_request(&request_id);
3960 Ok(RouterResponse {
3961 id: request_id,
3962 inner: result.map_err(|e| match e {
3967 Error::JsonRpc(err) => err,
3968 Error::Tool(err) => JsonRpcError::internal_error(err.to_string()),
3969 e => JsonRpcError::internal_error(e.to_string()),
3970 }),
3971 })
3972 })
3973 }
3974}
3975
3976#[cfg(test)]
3977mod tests {
3978 use super::*;
3979 use crate::extract::{Context, Json};
3980 use crate::jsonrpc::JsonRpcService;
3981 use crate::tool::ToolBuilder;
3982 use schemars::JsonSchema;
3983 use serde::Deserialize;
3984 use tower::ServiceExt;
3985
3986 #[derive(Debug, Deserialize, JsonSchema)]
3987 struct AddInput {
3988 a: i64,
3989 b: i64,
3990 }
3991
3992 #[cfg(feature = "stateless")]
3993 fn final_extensions(client_capabilities: ClientCapabilities) -> Extensions {
3994 let mut extensions = Extensions::new();
3995 extensions.insert(crate::stateless::StatelessRequestMeta {
3996 protocol_version: Some(PROTOCOL_VERSION_2026_07_28.to_string()),
3997 client_capabilities: Some(client_capabilities),
3998 ..Default::default()
3999 });
4000 extensions
4001 }
4002
4003 #[cfg(feature = "stateless")]
4004 fn tasks_client_extensions() -> Extensions {
4005 final_extensions(ClientCapabilities {
4006 extensions: Some(
4007 [(TASKS_EXTENSION_ID.to_string(), serde_json::json!({}))]
4008 .into_iter()
4009 .collect(),
4010 ),
4011 ..Default::default()
4012 })
4013 }
4014
4015 #[cfg(feature = "stateless")]
4016 #[tokio::test]
4017 async fn final_tasks_require_server_opt_in_and_client_declaration() {
4018 let tool = || {
4019 ToolBuilder::new("optional_task")
4020 .task_support(TaskSupportMode::Optional)
4021 .handler(|input: AddInput| async move {
4022 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4023 })
4024 .build()
4025 };
4026 let task_params = |task| CallToolParams {
4027 name: "optional_task".to_string(),
4028 arguments: serde_json::json!({"a": 1, "b": 2}),
4029 input_responses: None,
4030 request_state: None,
4031 meta: None,
4032 task,
4033 };
4034
4035 let implicit = McpRouter::new().tool(tool());
4039 let McpResponse::Discover(result) = implicit
4040 .handle(
4041 RequestId::Number(1),
4042 McpRequest::Discover(DiscoverParams::default()),
4043 Extensions::new(),
4044 )
4045 .await
4046 .unwrap()
4047 else {
4048 panic!("Expected Discover response");
4049 };
4050 assert!(
4051 result
4052 .capabilities
4053 .extensions
4054 .as_ref()
4055 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4056 );
4057 let error = implicit
4058 .handle(
4059 RequestId::Number(2),
4060 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4061 tasks_client_extensions(),
4062 )
4063 .await
4064 .unwrap_err();
4065 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4066
4067 let router = McpRouter::new().tool(tool()).with_tasks();
4069 let McpResponse::Discover(result) = router
4070 .handle(
4071 RequestId::Number(3),
4072 McpRequest::Discover(DiscoverParams::default()),
4073 Extensions::new(),
4074 )
4075 .await
4076 .unwrap()
4077 else {
4078 panic!("Expected Discover response");
4079 };
4080 assert!(
4081 result
4082 .capabilities
4083 .extensions
4084 .as_ref()
4085 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID)),
4086 "with_tasks() must advertise the extension on the final path"
4087 );
4088 assert!(
4089 result.capabilities.tasks.is_none(),
4090 "the legacy capability shape is never advertised on the final path"
4091 );
4092
4093 let response = router
4096 .handle(
4097 RequestId::Number(4),
4098 McpRequest::CallTool(task_params(None)),
4099 final_extensions(ClientCapabilities::default()),
4100 )
4101 .await
4102 .unwrap();
4103 assert!(matches!(response, McpResponse::CallTool(_)));
4104
4105 let response = router
4108 .handle(
4109 RequestId::Number(5),
4110 McpRequest::CallTool(task_params(None)),
4111 tasks_client_extensions(),
4112 )
4113 .await
4114 .unwrap();
4115 assert!(
4116 matches!(response, McpResponse::FinalCreateTask(_)),
4117 "a negotiated request must receive a task, got {response:?}"
4118 );
4119
4120 let error = router
4123 .handle(
4124 RequestId::Number(6),
4125 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4126 tasks_client_extensions(),
4127 )
4128 .await
4129 .unwrap_err();
4130 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4131 }
4132
4133 #[cfg(feature = "stateless")]
4134 #[tokio::test]
4135 async fn final_task_methods_serve_the_negotiated_wire_shapes() {
4136 let router = McpRouter::new()
4137 .tool(
4138 ToolBuilder::new("optional_task")
4139 .task_support(TaskSupportMode::Optional)
4140 .handler(|input: AddInput| async move {
4141 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4142 })
4143 .task_preparation(|task, _input| async move {
4144 let mut meta = serde_json::Map::new();
4145 meta.insert(
4146 "dev.tower-mcp/owner-test".to_string(),
4147 serde_json::json!({"taskId": task.task_id()}),
4148 );
4149 Ok(crate::TaskPreparation::new().with_meta(meta))
4150 })
4151 .build(),
4152 )
4153 .with_tasks();
4154
4155 let McpResponse::FinalCreateTask(created) = router
4156 .handle(
4157 RequestId::Number(1),
4158 McpRequest::CallTool(CallToolParams {
4159 name: "optional_task".to_string(),
4160 arguments: serde_json::json!({"a": 1, "b": 2}),
4161 input_responses: None,
4162 request_state: None,
4163 meta: None,
4164 task: None,
4165 }),
4166 tasks_client_extensions(),
4167 )
4168 .await
4169 .unwrap()
4170 else {
4171 panic!("Expected a final create-task response");
4172 };
4173
4174 let wire = serde_json::to_value(&created).unwrap();
4176 assert_eq!(wire["resultType"], "task");
4177 assert!(wire.get("task").is_none(), "final results are flat: {wire}");
4178 assert!(wire["ttlMs"].is_number() || wire["ttlMs"].is_null());
4179 assert!(wire.get("ttl").is_none(), "legacy field name leaked");
4180 let task_id = created.task.metadata.task_id.clone();
4181 assert_eq!(
4182 created.meta.as_ref().unwrap()["dev.tower-mcp/owner-test"]["taskId"],
4183 task_id
4184 );
4185
4186 let McpResponse::FinalGetTask(fetched) = router
4188 .handle(
4189 RequestId::Number(2),
4190 McpRequest::GetTaskInfo(GetTaskInfoParams {
4191 task_id: task_id.clone(),
4192 meta: None,
4193 }),
4194 tasks_client_extensions(),
4195 )
4196 .await
4197 .unwrap()
4198 else {
4199 panic!("Expected a final get-task response");
4200 };
4201 let wire = serde_json::to_value(&fetched).unwrap();
4202 assert_eq!(wire["resultType"], "complete");
4203 assert_eq!(wire["taskId"], serde_json::json!(task_id));
4204 assert!(wire["status"].is_string());
4205
4206 for (id, request) in [
4208 (
4209 3,
4210 McpRequest::UpdateTask(UpdateTaskParams {
4211 task_id: task_id.clone(),
4212 input_responses: HashMap::new(),
4213 meta: None,
4214 }),
4215 ),
4216 (
4217 4,
4218 McpRequest::CancelTask(CancelTaskParams {
4219 task_id: task_id.clone(),
4220 reason: None,
4221 meta: None,
4222 }),
4223 ),
4224 ] {
4225 let response = router
4226 .handle(RequestId::Number(id), request, tasks_client_extensions())
4227 .await
4228 .unwrap();
4229 let McpResponse::FinalTaskAck(ack) = response else {
4230 panic!("Expected a final ack for request {id}");
4231 };
4232 assert_eq!(
4233 serde_json::to_value(&ack).unwrap(),
4234 serde_json::json!({"resultType": "complete"})
4235 );
4236 }
4237
4238 let error = router
4240 .handle(
4241 RequestId::Number(5),
4242 McpRequest::GetTaskInfo(GetTaskInfoParams {
4243 task_id: "does-not-exist".to_string(),
4244 meta: None,
4245 }),
4246 tasks_client_extensions(),
4247 )
4248 .await
4249 .unwrap_err();
4250 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4251
4252 let error = router
4254 .handle(
4255 RequestId::Number(6),
4256 McpRequest::GetTaskInfo(GetTaskInfoParams {
4257 task_id: task_id.clone(),
4258 meta: None,
4259 }),
4260 final_extensions(ClientCapabilities::default()),
4261 )
4262 .await
4263 .unwrap_err();
4264 let Error::JsonRpc(error) = error else {
4265 panic!("expected a JSON-RPC error");
4266 };
4267 assert_eq!(error.code, -32021);
4268 assert_eq!(
4269 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4270 serde_json::json!({})
4271 );
4272 }
4273
4274 #[cfg(feature = "stateless")]
4275 #[tokio::test]
4276 async fn final_required_task_tools_follow_per_request_capabilities() {
4277 let router = McpRouter::new()
4278 .tool(
4279 ToolBuilder::new("required_task")
4280 .task_support(TaskSupportMode::Required)
4281 .handler(|input: AddInput| async move {
4282 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4283 })
4284 .build(),
4285 )
4286 .with_tasks();
4287 let params = || CallToolParams {
4288 name: "required_task".to_string(),
4289 arguments: serde_json::json!({"a": 1, "b": 2}),
4290 input_responses: None,
4291 request_state: None,
4292 meta: None,
4293 task: None,
4294 };
4295
4296 let McpResponse::ListTools(without_tasks) = router
4297 .handle(
4298 RequestId::Number(1),
4299 McpRequest::ListTools(ListToolsParams::default()),
4300 final_extensions(ClientCapabilities::default()),
4301 )
4302 .await
4303 .unwrap()
4304 else {
4305 panic!("expected tools/list")
4306 };
4307 assert!(without_tasks.tools.is_empty());
4308
4309 let McpResponse::ListTools(with_tasks) = router
4310 .handle(
4311 RequestId::Number(2),
4312 McpRequest::ListTools(ListToolsParams::default()),
4313 tasks_client_extensions(),
4314 )
4315 .await
4316 .unwrap()
4317 else {
4318 panic!("expected tools/list")
4319 };
4320 assert_eq!(with_tasks.tools.len(), 1);
4321 assert!(with_tasks.tools[0].execution.is_none());
4322
4323 let error = router
4324 .handle(
4325 RequestId::Number(3),
4326 McpRequest::CallTool(params()),
4327 final_extensions(ClientCapabilities::default()),
4328 )
4329 .await
4330 .unwrap_err();
4331 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32021));
4332
4333 let response = router
4334 .handle(
4335 RequestId::Number(4),
4336 McpRequest::CallTool(params()),
4337 tasks_client_extensions(),
4338 )
4339 .await
4340 .unwrap();
4341 assert!(matches!(response, McpResponse::FinalCreateTask(_)));
4342 }
4343
4344 #[cfg(all(feature = "oauth", feature = "stateless"))]
4345 #[tokio::test]
4346 async fn task_operations_are_bound_to_the_creating_principal() {
4347 fn as_principal(subject: &str) -> Extensions {
4348 let mut extensions = tasks_client_extensions();
4349 extensions.insert(crate::oauth::token::TokenClaims {
4350 sub: Some(subject.to_string()),
4351 iss: None,
4352 aud: None,
4353 exp: None,
4354 scope: None,
4355 client_id: None,
4356 extra: HashMap::new(),
4357 });
4358 extensions
4359 }
4360
4361 let router = McpRouter::new()
4362 .tool(
4363 ToolBuilder::new("optional_task")
4364 .task_support(TaskSupportMode::Optional)
4365 .handler(|input: AddInput| async move {
4366 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4367 })
4368 .build(),
4369 )
4370 .with_tasks();
4371
4372 let McpResponse::FinalCreateTask(created) = router
4373 .handle(
4374 RequestId::Number(1),
4375 McpRequest::CallTool(CallToolParams {
4376 name: "optional_task".to_string(),
4377 arguments: serde_json::json!({"a": 1, "b": 2}),
4378 input_responses: None,
4379 request_state: None,
4380 meta: None,
4381 task: None,
4382 }),
4383 as_principal("alice"),
4384 )
4385 .await
4386 .unwrap()
4387 else {
4388 panic!("Expected a final create-task response");
4389 };
4390 let task_id = created.task.metadata.task_id.clone();
4391
4392 assert!(
4394 router
4395 .handle(
4396 RequestId::Number(2),
4397 McpRequest::GetTaskInfo(GetTaskInfoParams {
4398 task_id: task_id.clone(),
4399 meta: None,
4400 }),
4401 as_principal("alice"),
4402 )
4403 .await
4404 .is_ok()
4405 );
4406
4407 for (id, label, context) in [
4410 (3, "another principal", as_principal("bob")),
4411 (4, "no principal", tasks_client_extensions()),
4412 ] {
4413 for (offset, request) in [
4414 McpRequest::GetTaskInfo(GetTaskInfoParams {
4415 task_id: task_id.clone(),
4416 meta: None,
4417 }),
4418 McpRequest::UpdateTask(UpdateTaskParams {
4419 task_id: task_id.clone(),
4420 input_responses: HashMap::new(),
4421 meta: None,
4422 }),
4423 McpRequest::CancelTask(CancelTaskParams {
4424 task_id: task_id.clone(),
4425 reason: None,
4426 meta: None,
4427 }),
4428 ]
4429 .into_iter()
4430 .enumerate()
4431 {
4432 let error = router
4433 .handle(
4434 RequestId::Number(id * 10 + offset as i64),
4435 request,
4436 context.clone(),
4437 )
4438 .await
4439 .unwrap_err();
4440 assert!(
4441 matches!(error, Error::JsonRpc(ref e) if e.code == -32602),
4442 "{label} was served: {error:?}"
4443 );
4444 let Error::JsonRpc(error) = error else {
4447 unreachable!()
4448 };
4449 assert!(
4450 error.message.contains("not found"),
4451 "refusal leaked that the task exists: {}",
4452 error.message
4453 );
4454 }
4455 }
4456
4457 assert!(
4459 router
4460 .handle(
4461 RequestId::Number(9),
4462 McpRequest::GetTaskInfo(GetTaskInfoParams {
4463 task_id: task_id.clone(),
4464 meta: None,
4465 }),
4466 as_principal("alice"),
4467 )
4468 .await
4469 .is_ok(),
4470 "a refused cancel must not have cancelled the task"
4471 );
4472 }
4473
4474 #[cfg(all(feature = "oauth", feature = "stateless"))]
4475 #[tokio::test]
4476 async fn final_tasks_work_across_independent_routers_with_a_shared_store() {
4477 fn as_principal(subject: &str) -> Extensions {
4478 let mut extensions = tasks_client_extensions();
4479 extensions.insert(crate::oauth::token::TokenClaims {
4480 sub: Some(subject.to_string()),
4481 iss: None,
4482 aud: None,
4483 exp: None,
4484 scope: None,
4485 client_id: None,
4486 extra: HashMap::new(),
4487 });
4488 extensions
4489 }
4490
4491 fn router_with_store(store: Arc<dyn TaskStore>) -> McpRouter {
4492 McpRouter::new()
4493 .tool(
4494 ToolBuilder::new("shared_task")
4495 .task_support(TaskSupportMode::Optional)
4496 .handler(|_input: serde_json::Value| async move {
4497 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
4498 Ok(CallToolResult::text("done"))
4499 })
4500 .build(),
4501 )
4502 .task_store(store)
4503 .with_tasks()
4504 }
4505
4506 let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
4507 let router_a = router_with_store(store.clone());
4508 let router_b = router_with_store(store);
4509
4510 let McpResponse::FinalCreateTask(created) = router_a
4511 .handle(
4512 RequestId::Number(1),
4513 McpRequest::CallTool(CallToolParams {
4514 name: "shared_task".to_string(),
4515 arguments: serde_json::json!({}),
4516 input_responses: None,
4517 request_state: None,
4518 meta: None,
4519 task: None,
4520 }),
4521 as_principal("alice"),
4522 )
4523 .await
4524 .unwrap()
4525 else {
4526 panic!("router A did not create a final task")
4527 };
4528 let task_id = created.task.metadata.task_id;
4529
4530 assert!(
4532 router_b
4533 .handle(
4534 RequestId::Number(2),
4535 McpRequest::GetTaskInfo(GetTaskInfoParams {
4536 task_id: task_id.clone(),
4537 meta: None,
4538 }),
4539 as_principal("alice"),
4540 )
4541 .await
4542 .is_ok()
4543 );
4544
4545 let denied = router_b
4547 .handle(
4548 RequestId::Number(3),
4549 McpRequest::GetTaskInfo(GetTaskInfoParams {
4550 task_id: task_id.clone(),
4551 meta: None,
4552 }),
4553 as_principal("bob"),
4554 )
4555 .await
4556 .unwrap_err();
4557 let unknown = router_b
4558 .handle(
4559 RequestId::Number(4),
4560 McpRequest::GetTaskInfo(GetTaskInfoParams {
4561 task_id: "unknown-task".to_string(),
4562 meta: None,
4563 }),
4564 as_principal("bob"),
4565 )
4566 .await
4567 .unwrap_err();
4568 let (Error::JsonRpc(denied), Error::JsonRpc(unknown)) = (denied, unknown) else {
4569 panic!("expected JSON-RPC task denials")
4570 };
4571 assert_eq!(denied.code, unknown.code);
4572 assert_eq!(
4573 denied.message.replace(&task_id, "<task-id>"),
4574 unknown.message.replace("unknown-task", "<task-id>")
4575 );
4576 assert_eq!(denied.data, unknown.data);
4577
4578 assert!(matches!(
4581 router_b
4582 .handle(
4583 RequestId::Number(5),
4584 McpRequest::CancelTask(CancelTaskParams {
4585 task_id: task_id.clone(),
4586 reason: None,
4587 meta: None,
4588 }),
4589 as_principal("alice"),
4590 )
4591 .await
4592 .unwrap(),
4593 McpResponse::FinalTaskAck(_)
4594 ));
4595 let McpResponse::FinalGetTask(fetched) = router_a
4596 .handle(
4597 RequestId::Number(6),
4598 McpRequest::GetTaskInfo(GetTaskInfoParams {
4599 task_id,
4600 meta: None,
4601 }),
4602 as_principal("alice"),
4603 )
4604 .await
4605 .unwrap()
4606 else {
4607 panic!("router A did not read the shared task")
4608 };
4609 assert_eq!(fetched.task.status(), TaskStatus::Cancelled);
4610 }
4611
4612 #[test]
4613 fn router_advertises_only_locally_declared_protocol_extensions() {
4614 let router = McpRouter::new().with_protocol_extension(
4615 crate::ExtensionDeclaration::new(
4616 "com.example/rendering",
4617 serde_json::json!({"formats": ["html"]}),
4618 )
4619 .unwrap(),
4620 );
4621
4622 let stable = router.capabilities();
4623 let final_capabilities =
4624 router.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
4625 for capabilities in [stable, final_capabilities] {
4626 let extensions = capabilities.extensions.unwrap();
4627 assert_eq!(extensions.len(), 1);
4628 assert_eq!(extensions["com.example/rendering"]["formats"][0], "html");
4629 assert!(!extensions.contains_key("com.example/client-only"));
4630 }
4631 }
4632
4633 #[tokio::test]
4634 async fn initialize_persists_negotiated_extensions_for_legacy_contexts() {
4635 let router = McpRouter::new().with_protocol_extension(
4636 crate::ExtensionDeclaration::new(
4637 "com.example/shared",
4638 serde_json::json!({"server": true}),
4639 )
4640 .unwrap(),
4641 );
4642 let client_capabilities = ClientCapabilities {
4643 extensions: Some(HashMap::from([
4644 (
4645 "com.example/shared".to_string(),
4646 serde_json::json!({"client": true}),
4647 ),
4648 ("com.example/client-only".to_string(), serde_json::json!({})),
4649 ])),
4650 ..ClientCapabilities::default()
4651 };
4652
4653 router
4654 .handle(
4655 RequestId::Number(1),
4656 McpRequest::Initialize(InitializeParams {
4657 protocol_version: crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
4658 capabilities: client_capabilities,
4659 client_info: Implementation {
4660 name: "extension-test".to_string(),
4661 version: "1.0.0".to_string(),
4662 title: None,
4663 description: None,
4664 icons: None,
4665 website_url: None,
4666 meta: None,
4667 },
4668 meta: None,
4669 }),
4670 Extensions::new(),
4671 )
4672 .await
4673 .unwrap();
4674
4675 let context = router.create_context(RequestId::Number(2), None);
4676 let negotiated = context.negotiated_extensions().unwrap();
4677 assert!(negotiated.contains("com.example/shared"));
4678 assert!(!negotiated.contains("com.example/client-only"));
4679 }
4680
4681 #[cfg(feature = "stateless")]
4682 #[test]
4683 fn final_request_context_exposes_only_negotiated_extensions() {
4684 let router = McpRouter::new().with_protocol_extension(
4685 crate::ExtensionDeclaration::new(
4686 "com.example/shared",
4687 serde_json::json!({"server": true}),
4688 )
4689 .unwrap(),
4690 );
4691 let per_request = final_extensions(ClientCapabilities {
4692 extensions: Some(HashMap::from([
4693 (
4694 "com.example/shared".to_string(),
4695 serde_json::json!({"client": true}),
4696 ),
4697 ("com.example/client-only".to_string(), serde_json::json!({})),
4698 ])),
4699 ..ClientCapabilities::default()
4700 });
4701
4702 let context =
4703 router.create_context_with_extensions(RequestId::Number(1), None, &per_request);
4704 let negotiated = context.negotiated_extensions().unwrap();
4705
4706 assert_eq!(negotiated.len(), 1);
4707 assert_eq!(
4708 negotiated
4709 .get("com.example/shared")
4710 .unwrap()
4711 .client_settings()["client"],
4712 true
4713 );
4714 assert!(!negotiated.contains("com.example/client-only"));
4715 }
4716
4717 #[cfg(feature = "stateless")]
4718 #[tokio::test]
4719 async fn final_protocol_withholds_incomplete_tasks_advertisement() {
4720 let optional = ToolBuilder::new("optional_task")
4721 .task_support(TaskSupportMode::Optional)
4722 .handler(|input: AddInput| async move {
4723 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4724 })
4725 .build();
4726 let required = ToolBuilder::new("required_task")
4727 .task_support(TaskSupportMode::Required)
4728 .handler(|input: AddInput| async move {
4729 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4730 })
4731 .build();
4732 let mut router = McpRouter::new().tool(optional).tool(required);
4733
4734 let stable_capabilities = router.capabilities();
4736 assert!(stable_capabilities.tasks.is_some());
4737 assert!(
4738 stable_capabilities
4739 .extensions
4740 .as_ref()
4741 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID))
4742 );
4743
4744 let response = router
4746 .handle(
4747 RequestId::Number(1),
4748 McpRequest::Discover(DiscoverParams::default()),
4749 Extensions::new(),
4750 )
4751 .await
4752 .unwrap();
4753 let McpResponse::Discover(result) = response else {
4754 panic!("Expected Discover response");
4755 };
4756 assert!(result.capabilities.tasks.is_none());
4757 assert!(
4758 result
4759 .capabilities
4760 .extensions
4761 .as_ref()
4762 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4763 );
4764
4765 init_router(&mut router).await;
4766
4767 let response = router
4769 .handle(
4770 RequestId::Number(2),
4771 McpRequest::ListTools(ListToolsParams::default()),
4772 Extensions::new(),
4773 )
4774 .await
4775 .unwrap();
4776 let McpResponse::ListTools(result) = response else {
4777 panic!("Expected ListTools response");
4778 };
4779 assert_eq!(result.tools.len(), 2);
4780 assert!(result.tools.iter().all(|tool| tool.execution.is_some()));
4781
4782 let response = router
4785 .handle(
4786 RequestId::Number(3),
4787 McpRequest::ListTools(ListToolsParams::default()),
4788 final_extensions(ClientCapabilities::default()),
4789 )
4790 .await
4791 .unwrap();
4792 let McpResponse::ListTools(result) = response else {
4793 panic!("Expected ListTools response");
4794 };
4795 assert_eq!(result.tools.len(), 1);
4796 assert_eq!(result.tools[0].name, "optional_task");
4797 assert!(result.tools[0].execution.is_none());
4798 }
4799
4800 #[cfg(feature = "stateless")]
4801 #[tokio::test]
4802 async fn final_protocol_enforces_tasks_negotiation() {
4803 let optional = ToolBuilder::new("optional_task")
4804 .task_support(TaskSupportMode::Optional)
4805 .handler(|input: AddInput| async move {
4806 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4807 })
4808 .build();
4809 let required = ToolBuilder::new("required_task")
4810 .task_support(TaskSupportMode::Required)
4811 .handler(|input: AddInput| async move {
4812 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4813 })
4814 .build();
4815 let mut router = McpRouter::new().tool(optional).tool(required).with_tasks();
4816 init_router(&mut router).await;
4817
4818 let response = router
4820 .handle(
4821 RequestId::Number(1),
4822 McpRequest::CallTool(CallToolParams {
4823 name: "optional_task".to_string(),
4824 arguments: serde_json::json!({"a": 1, "b": 2}),
4825 input_responses: None,
4826 request_state: None,
4827 meta: None,
4828 task: None,
4829 }),
4830 final_extensions(ClientCapabilities::default()),
4831 )
4832 .await
4833 .unwrap();
4834 assert!(matches!(response, McpResponse::CallTool(_)));
4835
4836 let error = router
4838 .handle(
4839 RequestId::Number(2),
4840 McpRequest::CallTool(CallToolParams {
4841 name: "optional_task".to_string(),
4842 arguments: serde_json::json!({"a": 1, "b": 2}),
4843 input_responses: None,
4844 request_state: None,
4845 meta: None,
4846 task: Some(TaskRequestParams { ttl: None }),
4847 }),
4848 final_extensions(ClientCapabilities::default()),
4849 )
4850 .await
4851 .unwrap_err();
4852 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32602));
4853
4854 let error = router
4858 .handle(
4859 RequestId::Number(3),
4860 McpRequest::CallTool(CallToolParams {
4861 name: "required_task".to_string(),
4862 arguments: serde_json::json!({"a": 1, "b": 2}),
4863 input_responses: None,
4864 request_state: None,
4865 meta: None,
4866 task: None,
4867 }),
4868 final_extensions(ClientCapabilities::default()),
4869 )
4870 .await
4871 .unwrap_err();
4872 let Error::JsonRpc(error) = error else {
4873 panic!("expected a JSON-RPC error");
4874 };
4875 assert_eq!(error.code, -32021);
4876 assert_eq!(
4877 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4878 serde_json::json!({}),
4879 "the error must name the extension the client needs to declare"
4880 );
4881
4882 let task_requests = [
4883 McpRequest::GetTaskInfo(GetTaskInfoParams {
4884 task_id: "task-unknown".to_string(),
4885 meta: None,
4886 }),
4887 McpRequest::UpdateTask(UpdateTaskParams {
4888 task_id: "task-unknown".to_string(),
4889 input_responses: HashMap::new(),
4890 meta: None,
4891 }),
4892 McpRequest::CancelTask(CancelTaskParams {
4893 task_id: "task-unknown".to_string(),
4894 reason: None,
4895 meta: None,
4896 }),
4897 ];
4898 for (index, request) in task_requests.into_iter().enumerate() {
4899 let error = router
4900 .handle(
4901 RequestId::Number(4 + index as i64),
4902 request,
4903 final_extensions(ClientCapabilities::default()),
4904 )
4905 .await
4906 .unwrap_err();
4907 let Error::JsonRpc(error) = error else {
4908 panic!("expected a JSON-RPC error");
4909 };
4910 assert_eq!(error.code, -32021);
4911 assert_eq!(
4912 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4913 serde_json::json!({})
4914 );
4915 }
4916
4917 let router_without_tasks = McpRouter::new();
4920 let error = router_without_tasks
4921 .handle(
4922 RequestId::Number(7),
4923 McpRequest::GetTaskInfo(GetTaskInfoParams {
4924 task_id: "task-unknown".to_string(),
4925 meta: None,
4926 }),
4927 final_extensions(tasks_client_capabilities()),
4928 )
4929 .await
4930 .unwrap_err();
4931 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32601));
4932 }
4933
4934 #[cfg(feature = "stateless")]
4935 #[test]
4936 fn input_required_capability_validation_uses_capability_semantics() {
4937 let roots = InputRequiredResult::with_requests(
4938 [(
4939 "roots".to_string(),
4940 InputRequest::ListRoots(ListRootsParams::default()),
4941 )]
4942 .into_iter()
4943 .collect(),
4944 );
4945 let extensions = final_extensions(ClientCapabilities {
4946 roots: Some(RootsCapability {
4947 list_changed: true,
4948 deprecated: None,
4949 }),
4950 ..Default::default()
4951 });
4952 validate_input_required_result(&extensions, &roots).unwrap();
4953 assert!(client_capabilities_satisfy(
4954 extensions
4955 .get::<crate::stateless::StatelessRequestMeta>()
4956 .and_then(|meta| meta.client_capabilities.as_ref())
4957 .unwrap(),
4958 &ClientCapabilities {
4959 roots: Some(RootsCapability::default()),
4960 ..Default::default()
4961 }
4962 ));
4963
4964 let sampling_with_tools = InputRequiredResult::with_requests(
4965 [(
4966 "sample".to_string(),
4967 InputRequest::CreateMessage(CreateMessageParams {
4968 tools: Some(Vec::new()),
4969 ..CreateMessageParams::new(vec![SamplingMessage::user("hello")], 10)
4970 }),
4971 )]
4972 .into_iter()
4973 .collect(),
4974 );
4975 let extensions = final_extensions(ClientCapabilities {
4976 sampling: Some(SamplingCapability::default()),
4977 ..Default::default()
4978 });
4979 assert!(validate_input_required_result(&extensions, &sampling_with_tools).is_err());
4980
4981 let form = InputRequiredResult::with_requests(
4982 [(
4983 "form".to_string(),
4984 InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
4985 mode: Some(ElicitMode::Form),
4986 message: "name".into(),
4987 requested_schema: ElicitFormSchema::new(),
4988 meta: None,
4989 })),
4990 )]
4991 .into_iter()
4992 .collect(),
4993 );
4994 let extensions = final_extensions(ClientCapabilities {
4995 elicitation: Some(ElicitationCapability::default()),
4996 ..Default::default()
4997 });
4998 validate_input_required_result(&extensions, &form).unwrap();
4999 }
5000
5001 async fn init_router(router: &mut McpRouter) {
5003 let init_req = RouterRequest {
5005 id: RequestId::Number(0),
5006 inner: McpRequest::Initialize(InitializeParams {
5007 protocol_version: "2025-11-25".to_string(),
5008 capabilities: ClientCapabilities {
5009 roots: None,
5010 sampling: None,
5011 elicitation: None,
5012 tasks: None,
5013 experimental: None,
5014 extensions: None,
5015 },
5016 client_info: Implementation {
5017 name: "test".to_string(),
5018 version: "1.0".to_string(),
5019 ..Default::default()
5020 },
5021 meta: None,
5022 }),
5023 extensions: Extensions::new(),
5024 };
5025 let _ = router.ready().await.unwrap().call(init_req).await.unwrap();
5026 router.handle_notification(McpNotification::Initialized);
5028 }
5029
5030 #[tokio::test]
5031 async fn test_router_list_tools() {
5032 let add_tool = ToolBuilder::new("add")
5033 .description("Add two numbers")
5034 .handler(|input: AddInput| async move {
5035 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5036 })
5037 .build();
5038
5039 let mut router = McpRouter::new().tool(add_tool);
5040
5041 init_router(&mut router).await;
5043
5044 let req = RouterRequest {
5045 id: RequestId::Number(1),
5046 inner: McpRequest::ListTools(ListToolsParams::default()),
5047 extensions: Extensions::new(),
5048 };
5049
5050 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5051
5052 match resp.inner {
5053 Ok(McpResponse::ListTools(result)) => {
5054 assert_eq!(result.tools.len(), 1);
5055 assert_eq!(result.tools[0].name, "add");
5056 }
5057 _ => panic!("Expected ListTools response"),
5058 }
5059 }
5060
5061 #[tokio::test]
5062 async fn test_router_call_tool() {
5063 let add_tool = ToolBuilder::new("add")
5064 .description("Add two numbers")
5065 .handler(|input: AddInput| async move {
5066 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5067 })
5068 .build();
5069
5070 let mut router = McpRouter::new().tool(add_tool);
5071
5072 init_router(&mut router).await;
5074
5075 let req = RouterRequest {
5076 id: RequestId::Number(1),
5077 inner: McpRequest::CallTool(CallToolParams {
5078 input_responses: None,
5079 request_state: None,
5080 name: "add".to_string(),
5081 arguments: serde_json::json!({"a": 2, "b": 3}),
5082 meta: None,
5083 task: None,
5084 }),
5085 extensions: Extensions::new(),
5086 };
5087
5088 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5089
5090 match resp.inner {
5091 Ok(McpResponse::CallTool(result)) => {
5092 assert!(!result.is_error);
5093 match &result.content[0] {
5095 Content::Text { text, .. } => assert_eq!(text, "5"),
5096 _ => panic!("Expected text content"),
5097 }
5098 }
5099 _ => panic!("Expected CallTool response"),
5100 }
5101 }
5102
5103 async fn init_jsonrpc_service(service: &mut JsonRpcService<McpRouter>, router: &McpRouter) {
5105 let init_req = JsonRpcRequest::new(0, "initialize").with_params(serde_json::json!({
5106 "protocolVersion": "2025-11-25",
5107 "capabilities": {},
5108 "clientInfo": { "name": "test", "version": "1.0" }
5109 }));
5110 let _ = service.call_single(init_req).await.unwrap();
5111 router.handle_notification(McpNotification::Initialized);
5112 }
5113
5114 #[tokio::test]
5115 async fn test_jsonrpc_service() {
5116 let add_tool = ToolBuilder::new("add")
5117 .description("Add two numbers")
5118 .handler(|input: AddInput| async move {
5119 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5120 })
5121 .build();
5122
5123 let router = McpRouter::new().tool(add_tool);
5124 let mut service = JsonRpcService::new(router.clone());
5125
5126 init_jsonrpc_service(&mut service, &router).await;
5128
5129 let req = JsonRpcRequest::new(1, "tools/list");
5130
5131 let resp = service.call_single(req).await.unwrap();
5132
5133 match resp {
5134 JsonRpcResponse::Result(r) => {
5135 assert_eq!(r.id, RequestId::Number(1));
5136 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5137 assert_eq!(tools.len(), 1);
5138 }
5139 JsonRpcResponse::Error(_) => panic!("Expected success response"),
5140 _ => panic!("unexpected response variant"),
5141 }
5142 }
5143
5144 #[tokio::test]
5145 async fn test_batch_request() {
5146 let add_tool = ToolBuilder::new("add")
5147 .description("Add two numbers")
5148 .handler(|input: AddInput| async move {
5149 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5150 })
5151 .build();
5152
5153 let router = McpRouter::new().tool(add_tool);
5154 let mut service = JsonRpcService::new(router.clone())
5155 .protocol_versions(["2025-03-26"])
5156 .unwrap();
5157
5158 init_jsonrpc_service(&mut service, &router).await;
5160
5161 let requests = vec![
5163 JsonRpcRequest::new(1, "tools/list"),
5164 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5165 "name": "add",
5166 "arguments": {"a": 10, "b": 20}
5167 })),
5168 JsonRpcRequest::new(3, "ping"),
5169 ];
5170
5171 let responses = service.call_batch(requests).await.unwrap();
5172
5173 assert_eq!(responses.len(), 3);
5174
5175 match &responses[0] {
5177 JsonRpcResponse::Result(r) => {
5178 assert_eq!(r.id, RequestId::Number(1));
5179 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5180 assert_eq!(tools.len(), 1);
5181 }
5182 JsonRpcResponse::Error(_) => panic!("Expected success for tools/list"),
5183 _ => panic!("unexpected response variant"),
5184 }
5185
5186 match &responses[1] {
5188 JsonRpcResponse::Result(r) => {
5189 assert_eq!(r.id, RequestId::Number(2));
5190 let content = r.result.get("content").unwrap().as_array().unwrap();
5191 let text = content[0].get("text").unwrap().as_str().unwrap();
5192 assert_eq!(text, "30");
5193 }
5194 JsonRpcResponse::Error(_) => panic!("Expected success for tools/call"),
5195 _ => panic!("unexpected response variant"),
5196 }
5197
5198 match &responses[2] {
5200 JsonRpcResponse::Result(r) => {
5201 assert_eq!(r.id, RequestId::Number(3));
5202 }
5203 JsonRpcResponse::Error(_) => panic!("Expected success for ping"),
5204 _ => panic!("unexpected response variant"),
5205 }
5206 }
5207
5208 #[tokio::test]
5209 async fn test_empty_batch_error() {
5210 let router = McpRouter::new();
5211 let mut service = JsonRpcService::new(router);
5212
5213 let result = service.call_batch(vec![]).await;
5214 assert!(result.is_err());
5215 }
5216
5217 #[tokio::test]
5222 async fn test_progress_token_extraction() {
5223 use crate::context::{ServerNotification, notification_channel};
5224 use crate::protocol::ProgressToken;
5225 use std::sync::Arc;
5226 use std::sync::atomic::{AtomicBool, Ordering};
5227
5228 let progress_reported = Arc::new(AtomicBool::new(false));
5230 let progress_ref = progress_reported.clone();
5231
5232 let tool = ToolBuilder::new("progress_tool")
5234 .description("Tool that reports progress")
5235 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5236 let reported = progress_ref.clone();
5237 async move {
5238 ctx.report_progress(50.0, Some(100.0), Some("Halfway"))
5240 .await;
5241 reported.store(true, Ordering::SeqCst);
5242 Ok(CallToolResult::text("done"))
5243 }
5244 })
5245 .build();
5246
5247 let (tx, mut rx) = notification_channel(10);
5249 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5250 let mut service = JsonRpcService::new(router.clone());
5251
5252 init_jsonrpc_service(&mut service, &router).await;
5254
5255 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5257 "name": "progress_tool",
5258 "arguments": {"a": 1, "b": 2},
5259 "_meta": {
5260 "progressToken": "test-token-123"
5261 }
5262 }));
5263
5264 let resp = service.call_single(req).await.unwrap();
5265
5266 match resp {
5268 JsonRpcResponse::Result(_) => {}
5269 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
5270 _ => panic!("unexpected response variant"),
5271 }
5272
5273 assert!(progress_reported.load(Ordering::SeqCst));
5275
5276 let notification = rx.try_recv().expect("Expected progress notification");
5278 match notification {
5279 ServerNotification::Progress(params) => {
5280 assert_eq!(
5281 params.progress_token,
5282 ProgressToken::String("test-token-123".to_string())
5283 );
5284 assert_eq!(params.progress, 50.0);
5285 assert_eq!(params.total, Some(100.0));
5286 assert_eq!(params.message.as_deref(), Some("Halfway"));
5287 }
5288 _ => panic!("Expected Progress notification"),
5289 }
5290 }
5291
5292 #[tokio::test]
5293 async fn test_tool_call_without_progress_token() {
5294 use crate::context::notification_channel;
5295 use std::sync::Arc;
5296 use std::sync::atomic::{AtomicBool, Ordering};
5297
5298 let progress_attempted = Arc::new(AtomicBool::new(false));
5299 let progress_ref = progress_attempted.clone();
5300
5301 let tool = ToolBuilder::new("no_token_tool")
5302 .description("Tool that tries to report progress without token")
5303 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5304 let attempted = progress_ref.clone();
5305 async move {
5306 ctx.report_progress(50.0, Some(100.0), None).await;
5308 attempted.store(true, Ordering::SeqCst);
5309 Ok(CallToolResult::text("done"))
5310 }
5311 })
5312 .build();
5313
5314 let (tx, mut rx) = notification_channel(10);
5315 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5316 let mut service = JsonRpcService::new(router.clone());
5317
5318 init_jsonrpc_service(&mut service, &router).await;
5319
5320 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5322 "name": "no_token_tool",
5323 "arguments": {"a": 1, "b": 2}
5324 }));
5325
5326 let resp = service.call_single(req).await.unwrap();
5327 assert!(matches!(resp, JsonRpcResponse::Result(_)));
5328
5329 assert!(progress_attempted.load(Ordering::SeqCst));
5331
5332 assert!(rx.try_recv().is_err());
5334 }
5335
5336 #[tokio::test]
5337 async fn test_batch_errors_returned_not_dropped() {
5338 let add_tool = ToolBuilder::new("add")
5339 .description("Add two numbers")
5340 .handler(|input: AddInput| async move {
5341 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5342 })
5343 .build();
5344
5345 let router = McpRouter::new().tool(add_tool);
5346 let mut service = JsonRpcService::new(router.clone())
5347 .protocol_versions(["2025-03-26"])
5348 .unwrap();
5349
5350 init_jsonrpc_service(&mut service, &router).await;
5351
5352 let requests = vec![
5354 JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5356 "name": "add",
5357 "arguments": {"a": 10, "b": 20}
5358 })),
5359 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5361 "name": "nonexistent_tool",
5362 "arguments": {}
5363 })),
5364 JsonRpcRequest::new(3, "ping"),
5366 ];
5367
5368 let responses = service.call_batch(requests).await.unwrap();
5369
5370 assert_eq!(responses.len(), 3);
5372
5373 match &responses[0] {
5375 JsonRpcResponse::Result(r) => {
5376 assert_eq!(r.id, RequestId::Number(1));
5377 }
5378 JsonRpcResponse::Error(_) => panic!("Expected success for first request"),
5379 _ => panic!("unexpected response variant"),
5380 }
5381
5382 match &responses[1] {
5384 JsonRpcResponse::Error(e) => {
5385 assert_eq!(e.id, Some(RequestId::Number(2)));
5386 assert!(e.error.message.contains("not found") || e.error.code == -32601);
5388 }
5389 JsonRpcResponse::Result(_) => panic!("Expected error for second request"),
5390 _ => panic!("unexpected response variant"),
5391 }
5392
5393 match &responses[2] {
5395 JsonRpcResponse::Result(r) => {
5396 assert_eq!(r.id, RequestId::Number(3));
5397 }
5398 JsonRpcResponse::Error(_) => panic!("Expected success for third request"),
5399 _ => panic!("unexpected response variant"),
5400 }
5401 }
5402
5403 #[tokio::test]
5408 async fn test_list_resource_templates() {
5409 use crate::resource::ResourceTemplateBuilder;
5410 use std::collections::HashMap;
5411
5412 let template = ResourceTemplateBuilder::new("file:///{path}")
5413 .name("Project Files")
5414 .description("Access project files")
5415 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5416 Ok(ReadResourceResult {
5417 contents: vec![ResourceContent {
5418 uri,
5419 mime_type: None,
5420 text: None,
5421 blob: None,
5422 meta: None,
5423 }],
5424 meta: None,
5425 ..Default::default()
5426 })
5427 });
5428
5429 let mut router = McpRouter::new().resource_template(template);
5430
5431 init_router(&mut router).await;
5433
5434 let req = RouterRequest {
5435 id: RequestId::Number(1),
5436 inner: McpRequest::ListResourceTemplates(ListResourceTemplatesParams::default()),
5437 extensions: Extensions::new(),
5438 };
5439
5440 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5441
5442 match resp.inner {
5443 Ok(McpResponse::ListResourceTemplates(result)) => {
5444 assert_eq!(result.resource_templates.len(), 1);
5445 assert_eq!(result.resource_templates[0].uri_template, "file:///{path}");
5446 assert_eq!(result.resource_templates[0].name, "Project Files");
5447 }
5448 _ => panic!("Expected ListResourceTemplates response"),
5449 }
5450 }
5451
5452 #[tokio::test]
5453 async fn test_read_resource_via_template() {
5454 use crate::resource::ResourceTemplateBuilder;
5455 use std::collections::HashMap;
5456
5457 let template = ResourceTemplateBuilder::new("db://users/{id}")
5458 .name("User Records")
5459 .handler(|uri: String, vars: HashMap<String, String>| async move {
5460 let id = vars.get("id").unwrap().clone();
5461 Ok(ReadResourceResult {
5462 contents: vec![ResourceContent {
5463 uri,
5464 mime_type: Some("application/json".to_string()),
5465 text: Some(format!(r#"{{"id": "{}"}}"#, id)),
5466 blob: None,
5467 meta: None,
5468 }],
5469 meta: None,
5470 ..Default::default()
5471 })
5472 });
5473
5474 let mut router = McpRouter::new().resource_template(template);
5475
5476 init_router(&mut router).await;
5478
5479 let req = RouterRequest {
5481 id: RequestId::Number(1),
5482 inner: McpRequest::ReadResource(ReadResourceParams {
5483 input_responses: None,
5484 request_state: None,
5485 uri: "db://users/123".to_string(),
5486 meta: None,
5487 }),
5488 extensions: Extensions::new(),
5489 };
5490
5491 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5492
5493 match resp.inner {
5494 Ok(McpResponse::ReadResource(result)) => {
5495 assert_eq!(result.contents.len(), 1);
5496 assert_eq!(result.contents[0].uri, "db://users/123");
5497 assert!(result.contents[0].text.as_ref().unwrap().contains("123"));
5498 }
5499 _ => panic!("Expected ReadResource response"),
5500 }
5501 }
5502
5503 #[tokio::test]
5504 async fn test_static_resource_takes_precedence_over_template() {
5505 use crate::resource::{ResourceBuilder, ResourceTemplateBuilder};
5506 use std::collections::HashMap;
5507
5508 let template = ResourceTemplateBuilder::new("file:///{path}")
5510 .name("Files Template")
5511 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5512 Ok(ReadResourceResult {
5513 contents: vec![ResourceContent {
5514 uri,
5515 mime_type: None,
5516 text: Some("from template".to_string()),
5517 blob: None,
5518 meta: None,
5519 }],
5520 meta: None,
5521 ..Default::default()
5522 })
5523 });
5524
5525 let static_resource = ResourceBuilder::new("file:///README.md")
5527 .name("README")
5528 .text("from static resource");
5529
5530 let mut router = McpRouter::new()
5531 .resource_template(template)
5532 .resource(static_resource);
5533
5534 init_router(&mut router).await;
5536
5537 let req = RouterRequest {
5539 id: RequestId::Number(1),
5540 inner: McpRequest::ReadResource(ReadResourceParams {
5541 input_responses: None,
5542 request_state: None,
5543 uri: "file:///README.md".to_string(),
5544 meta: None,
5545 }),
5546 extensions: Extensions::new(),
5547 };
5548
5549 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5550
5551 match resp.inner {
5552 Ok(McpResponse::ReadResource(result)) => {
5553 assert_eq!(
5555 result.contents[0].text.as_deref(),
5556 Some("from static resource")
5557 );
5558 }
5559 _ => panic!("Expected ReadResource response"),
5560 }
5561 }
5562
5563 #[tokio::test]
5564 async fn test_resource_not_found_when_no_match() {
5565 use crate::resource::ResourceTemplateBuilder;
5566 use std::collections::HashMap;
5567
5568 let template = ResourceTemplateBuilder::new("db://users/{id}")
5569 .name("Users")
5570 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5571 Ok(ReadResourceResult {
5572 contents: vec![ResourceContent {
5573 uri,
5574 mime_type: None,
5575 text: None,
5576 blob: None,
5577 meta: None,
5578 }],
5579 meta: None,
5580 ..Default::default()
5581 })
5582 });
5583
5584 let mut router = McpRouter::new().resource_template(template);
5585
5586 init_router(&mut router).await;
5588
5589 let req = RouterRequest {
5591 id: RequestId::Number(1),
5592 inner: McpRequest::ReadResource(ReadResourceParams {
5593 input_responses: None,
5594 request_state: None,
5595 uri: "db://posts/123".to_string(),
5596 meta: None,
5597 }),
5598 extensions: Extensions::new(),
5599 };
5600
5601 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5602
5603 match resp.inner {
5604 Err(err) => {
5605 assert!(err.message.contains("not found"));
5606 }
5607 Ok(_) => panic!("Expected error for non-matching URI"),
5608 }
5609 }
5610
5611 #[tokio::test]
5612 async fn test_capabilities_include_resources_with_only_templates() {
5613 use crate::resource::ResourceTemplateBuilder;
5614 use std::collections::HashMap;
5615
5616 let template = ResourceTemplateBuilder::new("file:///{path}")
5617 .name("Files")
5618 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5619 Ok(ReadResourceResult {
5620 contents: vec![ResourceContent {
5621 uri,
5622 mime_type: None,
5623 text: None,
5624 blob: None,
5625 meta: None,
5626 }],
5627 meta: None,
5628 ..Default::default()
5629 })
5630 });
5631
5632 let mut router = McpRouter::new().resource_template(template);
5633
5634 let init_req = RouterRequest {
5636 id: RequestId::Number(0),
5637 inner: McpRequest::Initialize(InitializeParams {
5638 protocol_version: "2025-11-25".to_string(),
5639 capabilities: ClientCapabilities {
5640 roots: None,
5641 sampling: None,
5642 elicitation: None,
5643 tasks: None,
5644 experimental: None,
5645 extensions: None,
5646 },
5647 client_info: Implementation {
5648 name: "test".to_string(),
5649 version: "1.0".to_string(),
5650 ..Default::default()
5651 },
5652 meta: None,
5653 }),
5654 extensions: Extensions::new(),
5655 };
5656 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5657
5658 match resp.inner {
5659 Ok(McpResponse::Initialize(result)) => {
5660 assert!(result.capabilities.resources.is_some());
5662 }
5663 _ => panic!("Expected Initialize response"),
5664 }
5665 }
5666
5667 #[tokio::test]
5672 async fn test_log_sends_notification() {
5673 use crate::context::notification_channel;
5674
5675 let (tx, mut rx) = notification_channel(10);
5676 let router = McpRouter::new().with_notification_sender(tx);
5677
5678 let sent = router.log_info("Test message");
5680 assert!(sent);
5681
5682 let notification = rx.try_recv().unwrap();
5684 match notification {
5685 ServerNotification::LogMessage(params) => {
5686 assert_eq!(params.level, LogLevel::Info);
5687 let data = params.data;
5688 assert_eq!(
5689 data.get("message").unwrap().as_str().unwrap(),
5690 "Test message"
5691 );
5692 }
5693 _ => panic!("Expected LogMessage notification"),
5694 }
5695 }
5696
5697 #[tokio::test]
5698 async fn test_log_with_custom_params() {
5699 use crate::context::notification_channel;
5700
5701 let (tx, mut rx) = notification_channel(10);
5702 let router = McpRouter::new().with_notification_sender(tx);
5703
5704 let params = LoggingMessageParams::new(
5706 LogLevel::Error,
5707 serde_json::json!({
5708 "error": "Connection failed",
5709 "host": "localhost"
5710 }),
5711 )
5712 .with_logger("database");
5713
5714 let sent = router.log(params);
5715 assert!(sent);
5716
5717 let notification = rx.try_recv().unwrap();
5718 match notification {
5719 ServerNotification::LogMessage(params) => {
5720 assert_eq!(params.level, LogLevel::Error);
5721 assert_eq!(params.logger.as_deref(), Some("database"));
5722 let data = params.data;
5723 assert_eq!(
5724 data.get("error").unwrap().as_str().unwrap(),
5725 "Connection failed"
5726 );
5727 }
5728 _ => panic!("Expected LogMessage notification"),
5729 }
5730 }
5731
5732 #[tokio::test]
5733 async fn test_log_without_channel_returns_false() {
5734 let router = McpRouter::new();
5736
5737 assert!(!router.log_info("Test"));
5739 assert!(!router.log_warning("Test"));
5740 assert!(!router.log_error("Test"));
5741 assert!(!router.log_debug("Test"));
5742 }
5743
5744 #[tokio::test]
5745 async fn test_logging_capability_with_channel() {
5746 use crate::context::notification_channel;
5747
5748 let (tx, _rx) = notification_channel(10);
5749 let mut router = McpRouter::new().with_notification_sender(tx);
5750
5751 let init_req = RouterRequest {
5753 id: RequestId::Number(0),
5754 inner: McpRequest::Initialize(InitializeParams {
5755 protocol_version: "2025-11-25".to_string(),
5756 capabilities: ClientCapabilities {
5757 roots: None,
5758 sampling: None,
5759 elicitation: None,
5760 tasks: None,
5761 experimental: None,
5762 extensions: None,
5763 },
5764 client_info: Implementation {
5765 name: "test".to_string(),
5766 version: "1.0".to_string(),
5767 ..Default::default()
5768 },
5769 meta: None,
5770 }),
5771 extensions: Extensions::new(),
5772 };
5773 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5774
5775 match resp.inner {
5776 Ok(McpResponse::Initialize(result)) => {
5777 assert!(result.capabilities.logging.is_some());
5779 }
5780 _ => panic!("Expected Initialize response"),
5781 }
5782 }
5783
5784 #[tokio::test]
5785 async fn test_no_logging_capability_without_channel() {
5786 let mut router = McpRouter::new();
5787
5788 let init_req = RouterRequest {
5790 id: RequestId::Number(0),
5791 inner: McpRequest::Initialize(InitializeParams {
5792 protocol_version: "2025-11-25".to_string(),
5793 capabilities: ClientCapabilities {
5794 roots: None,
5795 sampling: None,
5796 elicitation: None,
5797 tasks: None,
5798 experimental: None,
5799 extensions: None,
5800 },
5801 client_info: Implementation {
5802 name: "test".to_string(),
5803 version: "1.0".to_string(),
5804 ..Default::default()
5805 },
5806 meta: None,
5807 }),
5808 extensions: Extensions::new(),
5809 };
5810 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5811
5812 match resp.inner {
5813 Ok(McpResponse::Initialize(result)) => {
5814 assert!(result.capabilities.logging.is_none());
5816 }
5817 _ => panic!("Expected Initialize response"),
5818 }
5819 }
5820
5821 #[tokio::test]
5826 async fn test_create_task_via_call_tool() {
5827 let add_tool = ToolBuilder::new("add")
5828 .description("Add two numbers")
5829 .task_support(TaskSupportMode::Optional)
5830 .handler(|input: AddInput| async move {
5831 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5832 })
5833 .build();
5834
5835 let mut router = McpRouter::new().tool(add_tool);
5836 init_router(&mut router).await;
5837
5838 let req = RouterRequest {
5839 id: RequestId::Number(1),
5840 inner: McpRequest::CallTool(CallToolParams {
5841 input_responses: None,
5842 request_state: None,
5843 name: "add".to_string(),
5844 arguments: serde_json::json!({"a": 5, "b": 10}),
5845 meta: None,
5846 task: Some(TaskRequestParams { ttl: None }),
5847 }),
5848 extensions: Extensions::new(),
5849 };
5850
5851 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5852
5853 match resp.inner {
5854 Ok(McpResponse::CreateTask(result)) => {
5855 assert!(!result.task.task_id.is_empty());
5856 assert_eq!(result.task.status, TaskStatus::Working);
5857 }
5858 _ => panic!("Expected CreateTask response"),
5859 }
5860 }
5861
5862 struct CountingTaskStore {
5865 inner: MemoryTaskStore,
5866 creates: std::sync::atomic::AtomicUsize,
5867 gets: std::sync::atomic::AtomicUsize,
5868 completes: std::sync::atomic::AtomicUsize,
5869 }
5870
5871 impl CountingTaskStore {
5872 fn new() -> Self {
5873 Self {
5874 inner: MemoryTaskStore::new(),
5875 creates: std::sync::atomic::AtomicUsize::new(0),
5876 gets: std::sync::atomic::AtomicUsize::new(0),
5877 completes: std::sync::atomic::AtomicUsize::new(0),
5878 }
5879 }
5880 }
5881
5882 #[async_trait::async_trait]
5883 impl TaskStore for CountingTaskStore {
5884 async fn create_task(
5885 &self,
5886 tool_name: &str,
5887 arguments: serde_json::Value,
5888 ttl: Option<u64>,
5889 owner: crate::async_task::TaskOwner,
5890 ) -> crate::async_task::Result<(String, crate::async_task::CancellationToken)> {
5891 self.creates
5892 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5893 self.inner
5894 .create_task(tool_name, arguments, ttl, owner)
5895 .await
5896 }
5897
5898 async fn get_task(&self, task_id: &str) -> crate::async_task::Result<Option<TaskObject>> {
5899 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5900 self.inner.get_task(task_id).await
5901 }
5902
5903 async fn task_owner(
5904 &self,
5905 task_id: &str,
5906 ) -> crate::async_task::Result<Option<crate::async_task::TaskOwner>> {
5907 self.inner.task_owner(task_id).await
5908 }
5909
5910 async fn get_task_result(
5911 &self,
5912 task_id: &str,
5913 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
5914 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5917 self.inner.get_task_result(task_id).await
5918 }
5919
5920 async fn wait_for_completion(
5921 &self,
5922 task_id: &str,
5923 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
5924 self.inner.wait_for_completion(task_id).await
5925 }
5926
5927 async fn list_tasks(
5928 &self,
5929 status_filter: Option<TaskStatus>,
5930 ) -> crate::async_task::Result<Vec<TaskObject>> {
5931 self.inner.list_tasks(status_filter).await
5932 }
5933
5934 async fn require_input(
5935 &self,
5936 task_id: &str,
5937 requests: crate::protocol::InputRequests,
5938 message: Option<&str>,
5939 ) -> crate::async_task::Result<bool> {
5940 self.inner.require_input(task_id, requests, message).await
5941 }
5942
5943 async fn outstanding_input_requests(
5944 &self,
5945 task_id: &str,
5946 ) -> crate::async_task::Result<Option<crate::protocol::InputRequests>> {
5947 self.inner.outstanding_input_requests(task_id).await
5948 }
5949
5950 async fn apply_input_responses(
5951 &self,
5952 task_id: &str,
5953 responses: crate::protocol::InputResponses,
5954 ) -> crate::async_task::Result<Option<crate::async_task::AppliedInputResponses>> {
5955 self.inner.apply_input_responses(task_id, responses).await
5956 }
5957
5958 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> crate::async_task::Result<bool> {
5959 self.inner.set_ttl(task_id, ttl_ms).await
5960 }
5961
5962 async fn complete_task(
5963 &self,
5964 task_id: &str,
5965 result: CallToolResult,
5966 ) -> crate::async_task::Result<bool> {
5967 self.completes
5968 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5969 self.inner.complete_task(task_id, result).await
5970 }
5971
5972 async fn fail_task(
5973 &self,
5974 task_id: &str,
5975 error: JsonRpcError,
5976 ) -> crate::async_task::Result<bool> {
5977 self.inner.fail_task(task_id, error).await
5978 }
5979
5980 async fn cancel_task(
5981 &self,
5982 task_id: &str,
5983 reason: Option<&str>,
5984 ) -> crate::async_task::Result<Option<TaskObject>> {
5985 self.inner.cancel_task(task_id, reason).await
5986 }
5987 }
5988
5989 #[tokio::test]
5990 async fn test_injected_task_store_used_by_dispatch() {
5991 let store = Arc::new(CountingTaskStore::new());
5992
5993 let add_tool = ToolBuilder::new("add")
5994 .description("Add two numbers")
5995 .task_support(TaskSupportMode::Optional)
5996 .handler(|input: AddInput| async move {
5997 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5998 })
5999 .build();
6000
6001 let mut router = McpRouter::new()
6002 .tool(add_tool)
6003 .task_store(store.clone() as Arc<dyn TaskStore>);
6004 init_router(&mut router).await;
6005
6006 let req = RouterRequest {
6008 id: RequestId::Number(1),
6009 inner: McpRequest::CallTool(CallToolParams {
6010 input_responses: None,
6011 request_state: None,
6012 name: "add".to_string(),
6013 arguments: serde_json::json!({"a": 2, "b": 3}),
6014 meta: None,
6015 task: Some(TaskRequestParams { ttl: None }),
6016 }),
6017 extensions: Extensions::new(),
6018 };
6019 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6020 let task_id = match resp.inner {
6021 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6022 other => panic!("Expected CreateTask response, got {other:?}"),
6023 };
6024
6025 assert_eq!(
6026 store.creates.load(std::sync::atomic::Ordering::Relaxed),
6027 1,
6028 "create_task must go through the injected store"
6029 );
6030
6031 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6033 assert_eq!(
6034 store.completes.load(std::sync::atomic::Ordering::Relaxed),
6035 1,
6036 "complete_task must go through the injected store"
6037 );
6038
6039 let gets_before = store.gets.load(std::sync::atomic::Ordering::Relaxed);
6041 let req = RouterRequest {
6042 id: RequestId::Number(2),
6043 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6044 task_id: task_id.clone(),
6045 meta: None,
6046 }),
6047 extensions: Extensions::new(),
6048 };
6049 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6050 match resp.inner {
6051 Ok(McpResponse::GetTaskInfo(info)) => {
6052 assert_eq!(info.task_id, task_id);
6053 assert_eq!(info.status, TaskStatus::Completed);
6054 }
6055 other => panic!("Expected GetTaskInfo response, got {other:?}"),
6056 }
6057 assert!(
6058 store.gets.load(std::sync::atomic::Ordering::Relaxed) > gets_before,
6059 "tasks/get must go through the injected store"
6060 );
6061 }
6062
6063 #[tokio::test]
6064 async fn test_removed_tasks_methods_get_method_not_found() {
6065 let mut router = McpRouter::new();
6069 init_router(&mut router).await;
6070
6071 for method in ["tasks/list", "tasks/result"] {
6072 let req = RouterRequest {
6073 id: RequestId::Number(1),
6074 inner: McpRequest::Unknown {
6075 method: method.to_string(),
6076 params: None,
6077 },
6078 extensions: Extensions::new(),
6079 };
6080
6081 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6082
6083 match resp.inner {
6084 Err(err) => {
6085 assert_eq!(err.code, -32601, "{method} must be MethodNotFound");
6086 }
6087 other => panic!("Expected MethodNotFound error for {method}, got {other:?}"),
6088 }
6089 }
6090 }
6091
6092 #[tokio::test]
6093 async fn test_task_lifecycle_complete() {
6094 let add_tool = ToolBuilder::new("add")
6095 .description("Add two numbers")
6096 .task_support(TaskSupportMode::Optional)
6097 .handler(|input: AddInput| async move {
6098 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6099 })
6100 .build();
6101
6102 let mut router = McpRouter::new().tool(add_tool);
6103 init_router(&mut router).await;
6104
6105 let req = RouterRequest {
6107 id: RequestId::Number(1),
6108 inner: McpRequest::CallTool(CallToolParams {
6109 input_responses: None,
6110 request_state: None,
6111 name: "add".to_string(),
6112 arguments: serde_json::json!({"a": 7, "b": 8}),
6113 meta: None,
6114 task: Some(TaskRequestParams { ttl: None }),
6115 }),
6116 extensions: Extensions::new(),
6117 };
6118
6119 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6120 let task_id = match resp.inner {
6121 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6122 _ => panic!("Expected CreateTask response"),
6123 };
6124
6125 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6127
6128 let req = RouterRequest {
6132 id: RequestId::Number(2),
6133 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6134 task_id: task_id.clone(),
6135 meta: None,
6136 }),
6137 extensions: Extensions::new(),
6138 };
6139
6140 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6141
6142 match resp.inner {
6143 Ok(McpResponse::GetTaskInfo(info)) => {
6144 assert_eq!(info.task_id, task_id);
6145 assert_eq!(info.status, TaskStatus::Completed);
6146 }
6147 _ => panic!("Expected GetTaskInfo response"),
6148 }
6149 }
6150
6151 #[tokio::test]
6152 async fn test_task_cancellation() {
6153 let slow_tool = ToolBuilder::new("slow")
6155 .description("Slow tool")
6156 .task_support(TaskSupportMode::Optional)
6157 .handler(|_input: serde_json::Value| async move {
6158 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
6159 Ok(CallToolResult::text("done"))
6160 })
6161 .build();
6162
6163 let mut router = McpRouter::new().tool(slow_tool);
6164 init_router(&mut router).await;
6165
6166 let req = RouterRequest {
6168 id: RequestId::Number(1),
6169 inner: McpRequest::CallTool(CallToolParams {
6170 input_responses: None,
6171 request_state: None,
6172 name: "slow".to_string(),
6173 arguments: serde_json::json!({}),
6174 meta: None,
6175 task: Some(TaskRequestParams { ttl: None }),
6176 }),
6177 extensions: Extensions::new(),
6178 };
6179
6180 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6181 let task_id = match resp.inner {
6182 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6183 _ => panic!("Expected CreateTask response"),
6184 };
6185
6186 let req = RouterRequest {
6188 id: RequestId::Number(2),
6189 inner: McpRequest::CancelTask(CancelTaskParams {
6190 task_id: task_id.clone(),
6191 reason: Some("Test cancellation".to_string()),
6192 meta: None,
6193 }),
6194 extensions: Extensions::new(),
6195 };
6196
6197 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6198
6199 match resp.inner {
6201 Ok(McpResponse::CancelTask(EmptyResult {})) => {}
6202 other => panic!("Expected empty CancelTask ack, got {other:?}"),
6203 }
6204
6205 let req = RouterRequest {
6207 id: RequestId::Number(3),
6208 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6209 task_id: task_id.clone(),
6210 meta: None,
6211 }),
6212 extensions: Extensions::new(),
6213 };
6214 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6215 match resp.inner {
6216 Ok(McpResponse::GetTaskInfo(info)) => {
6217 assert_eq!(info.status, TaskStatus::Cancelled);
6218 }
6219 _ => panic!("Expected GetTaskInfo response"),
6220 }
6221 }
6222
6223 #[tokio::test]
6224 async fn test_get_task_info() {
6225 let add_tool = ToolBuilder::new("add")
6226 .description("Add two numbers")
6227 .task_support(TaskSupportMode::Optional)
6228 .handler(|input: AddInput| async move {
6229 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6230 })
6231 .build();
6232
6233 let mut router = McpRouter::new().tool(add_tool);
6234 init_router(&mut router).await;
6235
6236 let req = RouterRequest {
6238 id: RequestId::Number(1),
6239 inner: McpRequest::CallTool(CallToolParams {
6240 input_responses: None,
6241 request_state: None,
6242 name: "add".to_string(),
6243 arguments: serde_json::json!({"a": 1, "b": 2}),
6244 meta: None,
6245 task: Some(TaskRequestParams { ttl: Some(600_000) }),
6246 }),
6247 extensions: Extensions::new(),
6248 };
6249
6250 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6251 let task_id = match resp.inner {
6252 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6253 _ => panic!("Expected CreateTask response"),
6254 };
6255
6256 let req = RouterRequest {
6258 id: RequestId::Number(2),
6259 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6260 task_id: task_id.clone(),
6261 meta: None,
6262 }),
6263 extensions: Extensions::new(),
6264 };
6265
6266 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6267
6268 match resp.inner {
6269 Ok(McpResponse::GetTaskInfo(info)) => {
6270 assert_eq!(info.task_id, task_id);
6271 assert!(info.created_at.contains('T')); assert_eq!(info.ttl, Some(600_000));
6273 }
6274 _ => panic!("Expected GetTaskInfo response"),
6275 }
6276 }
6277
6278 #[tokio::test]
6279 async fn test_task_forbidden_tool_rejects_task_params() {
6280 let tool = ToolBuilder::new("sync_only")
6281 .description("Sync only tool")
6282 .handler(|_input: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
6283 .build();
6284
6285 let mut router = McpRouter::new().tool(tool);
6286 init_router(&mut router).await;
6287
6288 let req = RouterRequest {
6290 id: RequestId::Number(1),
6291 inner: McpRequest::CallTool(CallToolParams {
6292 input_responses: None,
6293 request_state: None,
6294 name: "sync_only".to_string(),
6295 arguments: serde_json::json!({}),
6296 meta: None,
6297 task: Some(TaskRequestParams { ttl: None }),
6298 }),
6299 extensions: Extensions::new(),
6300 };
6301
6302 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6303
6304 match resp.inner {
6305 Err(e) => {
6306 assert!(e.message.contains("does not support async tasks"));
6307 }
6308 _ => panic!("Expected error response"),
6309 }
6310 }
6311
6312 #[tokio::test]
6313 async fn test_get_nonexistent_task() {
6314 let mut router = McpRouter::new();
6315 init_router(&mut router).await;
6316
6317 let req = RouterRequest {
6318 id: RequestId::Number(1),
6319 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6320 task_id: "task-999".to_string(),
6321 meta: None,
6322 }),
6323 extensions: Extensions::new(),
6324 };
6325
6326 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6327
6328 match resp.inner {
6329 Err(e) => {
6330 assert!(e.message.contains("not found"));
6331 }
6332 _ => panic!("Expected error response"),
6333 }
6334 }
6335
6336 #[tokio::test]
6341 async fn test_subscribe_to_resource() {
6342 use crate::resource::ResourceBuilder;
6343
6344 let resource = ResourceBuilder::new("file:///test.txt")
6345 .name("Test File")
6346 .text("Hello");
6347
6348 let mut router = McpRouter::new().resource(resource);
6349 init_router(&mut router).await;
6350
6351 let req = RouterRequest {
6353 id: RequestId::Number(1),
6354 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6355 uri: "file:///test.txt".to_string(),
6356 meta: None,
6357 }),
6358 extensions: Extensions::new(),
6359 };
6360
6361 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6362
6363 match resp.inner {
6364 Ok(McpResponse::SubscribeResource(_)) => {
6365 assert!(router.is_subscribed("file:///test.txt"));
6367 }
6368 _ => panic!("Expected SubscribeResource response"),
6369 }
6370 }
6371
6372 #[tokio::test]
6373 async fn test_unsubscribe_from_resource() {
6374 use crate::resource::ResourceBuilder;
6375
6376 let resource = ResourceBuilder::new("file:///test.txt")
6377 .name("Test File")
6378 .text("Hello");
6379
6380 let mut router = McpRouter::new().resource(resource);
6381 init_router(&mut router).await;
6382
6383 let req = RouterRequest {
6385 id: RequestId::Number(1),
6386 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6387 uri: "file:///test.txt".to_string(),
6388 meta: None,
6389 }),
6390 extensions: Extensions::new(),
6391 };
6392 let _ = router.ready().await.unwrap().call(req).await.unwrap();
6393 assert!(router.is_subscribed("file:///test.txt"));
6394
6395 let req = RouterRequest {
6397 id: RequestId::Number(2),
6398 inner: McpRequest::UnsubscribeResource(UnsubscribeResourceParams {
6399 uri: "file:///test.txt".to_string(),
6400 meta: None,
6401 }),
6402 extensions: Extensions::new(),
6403 };
6404
6405 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6406
6407 match resp.inner {
6408 Ok(McpResponse::UnsubscribeResource(_)) => {
6409 assert!(!router.is_subscribed("file:///test.txt"));
6411 }
6412 _ => panic!("Expected UnsubscribeResource response"),
6413 }
6414 }
6415
6416 #[tokio::test]
6417 async fn test_subscribe_nonexistent_resource() {
6418 let mut router = McpRouter::new();
6419 init_router(&mut router).await;
6420
6421 let req = RouterRequest {
6422 id: RequestId::Number(1),
6423 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6424 uri: "file:///nonexistent.txt".to_string(),
6425 meta: None,
6426 }),
6427 extensions: Extensions::new(),
6428 };
6429
6430 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6431
6432 match resp.inner {
6433 Err(e) => {
6434 assert!(e.message.contains("not found"));
6435 }
6436 _ => panic!("Expected error response"),
6437 }
6438 }
6439
6440 #[tokio::test]
6441 async fn test_notify_resource_updated() {
6442 use crate::context::notification_channel;
6443 use crate::resource::ResourceBuilder;
6444
6445 let (tx, mut rx) = notification_channel(10);
6446
6447 let resource = ResourceBuilder::new("file:///test.txt")
6448 .name("Test File")
6449 .text("Hello");
6450
6451 let router = McpRouter::new()
6452 .resource(resource)
6453 .with_notification_sender(tx);
6454
6455 router.subscribe("file:///test.txt");
6457
6458 let sent = router.notify_resource_updated("file:///test.txt");
6460 assert!(sent);
6461
6462 let notification = rx.try_recv().unwrap();
6464 match notification {
6465 ServerNotification::ResourceUpdated { uri } => {
6466 assert_eq!(uri, "file:///test.txt");
6467 }
6468 _ => panic!("Expected ResourceUpdated notification"),
6469 }
6470 }
6471
6472 #[tokio::test]
6473 async fn test_notify_resource_updated_not_subscribed() {
6474 use crate::context::notification_channel;
6475 use crate::resource::ResourceBuilder;
6476
6477 let (tx, mut rx) = notification_channel(10);
6478
6479 let resource = ResourceBuilder::new("file:///test.txt")
6480 .name("Test File")
6481 .text("Hello");
6482
6483 let router = McpRouter::new()
6484 .resource(resource)
6485 .with_notification_sender(tx);
6486
6487 let sent = router.notify_resource_updated("file:///test.txt");
6489 assert!(!sent); assert!(rx.try_recv().is_err());
6493 }
6494
6495 #[tokio::test]
6496 async fn test_notify_resources_list_changed() {
6497 use crate::context::notification_channel;
6498
6499 let (tx, mut rx) = notification_channel(10);
6500 let router = McpRouter::new().with_notification_sender(tx);
6501
6502 let sent = router.notify_resources_list_changed();
6503 assert!(sent);
6504
6505 let notification = rx.try_recv().unwrap();
6506 match notification {
6507 ServerNotification::ResourcesListChanged => {}
6508 _ => panic!("Expected ResourcesListChanged notification"),
6509 }
6510 }
6511
6512 #[tokio::test]
6513 async fn test_subscribed_uris() {
6514 use crate::resource::ResourceBuilder;
6515
6516 let resource1 = ResourceBuilder::new("file:///a.txt").name("A").text("A");
6517
6518 let resource2 = ResourceBuilder::new("file:///b.txt").name("B").text("B");
6519
6520 let router = McpRouter::new().resource(resource1).resource(resource2);
6521
6522 router.subscribe("file:///a.txt");
6524 router.subscribe("file:///b.txt");
6525
6526 let uris = router.subscribed_uris();
6527 assert_eq!(uris.len(), 2);
6528 assert!(uris.contains(&"file:///a.txt".to_string()));
6529 assert!(uris.contains(&"file:///b.txt".to_string()));
6530 }
6531
6532 #[tokio::test]
6533 async fn test_subscription_capability_advertised() {
6534 use crate::resource::ResourceBuilder;
6535
6536 let resource = ResourceBuilder::new("file:///test.txt")
6537 .name("Test")
6538 .text("Hello");
6539
6540 let mut router = McpRouter::new().resource(resource);
6541
6542 let init_req = RouterRequest {
6544 id: RequestId::Number(0),
6545 inner: McpRequest::Initialize(InitializeParams {
6546 protocol_version: "2025-11-25".to_string(),
6547 capabilities: ClientCapabilities {
6548 roots: None,
6549 sampling: None,
6550 elicitation: None,
6551 tasks: None,
6552 experimental: None,
6553 extensions: None,
6554 },
6555 client_info: Implementation {
6556 name: "test".to_string(),
6557 version: "1.0".to_string(),
6558 ..Default::default()
6559 },
6560 meta: None,
6561 }),
6562 extensions: Extensions::new(),
6563 };
6564 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
6565
6566 match resp.inner {
6567 Ok(McpResponse::Initialize(result)) => {
6568 let resources_cap = result.capabilities.resources.unwrap();
6570 assert!(resources_cap.subscribe);
6571 }
6572 _ => panic!("Expected Initialize response"),
6573 }
6574 }
6575
6576 #[tokio::test]
6577 async fn test_completion_handler() {
6578 let router = McpRouter::new()
6579 .server_info("test", "1.0")
6580 .completion_handler(|params: CompleteParams| async move {
6581 let prefix = ¶ms.argument.value;
6583 let suggestions: Vec<String> = vec!["alpha", "beta", "gamma"]
6584 .into_iter()
6585 .filter(|s| s.starts_with(prefix))
6586 .map(String::from)
6587 .collect();
6588 Ok(CompleteResult::new(suggestions))
6589 });
6590
6591 let init_req = RouterRequest {
6593 id: RequestId::Number(0),
6594 inner: McpRequest::Initialize(InitializeParams {
6595 protocol_version: "2025-11-25".to_string(),
6596 capabilities: ClientCapabilities::default(),
6597 client_info: Implementation {
6598 name: "test".to_string(),
6599 version: "1.0".to_string(),
6600 ..Default::default()
6601 },
6602 meta: None,
6603 }),
6604 extensions: Extensions::new(),
6605 };
6606 let resp = router
6607 .clone()
6608 .ready()
6609 .await
6610 .unwrap()
6611 .call(init_req)
6612 .await
6613 .unwrap();
6614
6615 match resp.inner {
6617 Ok(McpResponse::Initialize(result)) => {
6618 assert!(result.capabilities.completions.is_some());
6619 }
6620 _ => panic!("Expected Initialize response"),
6621 }
6622
6623 router.handle_notification(McpNotification::Initialized);
6625
6626 let complete_req = RouterRequest {
6628 id: RequestId::Number(1),
6629 inner: McpRequest::Complete(CompleteParams {
6630 reference: CompletionReference::prompt("test-prompt"),
6631 argument: CompletionArgument::new("query", "al"),
6632 context: None,
6633 meta: None,
6634 }),
6635 extensions: Extensions::new(),
6636 };
6637 let resp = router
6638 .clone()
6639 .ready()
6640 .await
6641 .unwrap()
6642 .call(complete_req)
6643 .await
6644 .unwrap();
6645
6646 match resp.inner {
6647 Ok(McpResponse::Complete(result)) => {
6648 assert_eq!(result.completion.values, vec!["alpha"]);
6649 }
6650 _ => panic!("Expected Complete response"),
6651 }
6652 }
6653
6654 #[tokio::test]
6655 async fn test_completion_without_handler_returns_empty() {
6656 let router = McpRouter::new().server_info("test", "1.0");
6657
6658 let init_req = RouterRequest {
6660 id: RequestId::Number(0),
6661 inner: McpRequest::Initialize(InitializeParams {
6662 protocol_version: "2025-11-25".to_string(),
6663 capabilities: ClientCapabilities::default(),
6664 client_info: Implementation {
6665 name: "test".to_string(),
6666 version: "1.0".to_string(),
6667 ..Default::default()
6668 },
6669 meta: None,
6670 }),
6671 extensions: Extensions::new(),
6672 };
6673 let resp = router
6674 .clone()
6675 .ready()
6676 .await
6677 .unwrap()
6678 .call(init_req)
6679 .await
6680 .unwrap();
6681
6682 match resp.inner {
6684 Ok(McpResponse::Initialize(result)) => {
6685 assert!(result.capabilities.completions.is_none());
6686 }
6687 _ => panic!("Expected Initialize response"),
6688 }
6689
6690 router.handle_notification(McpNotification::Initialized);
6692
6693 let complete_req = RouterRequest {
6695 id: RequestId::Number(1),
6696 inner: McpRequest::Complete(CompleteParams {
6697 reference: CompletionReference::prompt("test-prompt"),
6698 argument: CompletionArgument::new("query", "al"),
6699 context: None,
6700 meta: None,
6701 }),
6702 extensions: Extensions::new(),
6703 };
6704 let resp = router
6705 .clone()
6706 .ready()
6707 .await
6708 .unwrap()
6709 .call(complete_req)
6710 .await
6711 .unwrap();
6712
6713 match resp.inner {
6714 Ok(McpResponse::Complete(result)) => {
6715 assert!(result.completion.values.is_empty());
6716 }
6717 _ => panic!("Expected Complete response"),
6718 }
6719 }
6720
6721 #[tokio::test]
6722 async fn test_tool_filter_list() {
6723 use crate::filter::CapabilityFilter;
6724 use crate::tool::Tool;
6725
6726 let public_tool = ToolBuilder::new("public")
6727 .description("Public tool")
6728 .handler(|_: AddInput| async move { Ok(CallToolResult::text("public")) })
6729 .build();
6730
6731 let admin_tool = ToolBuilder::new("admin")
6732 .description("Admin tool")
6733 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6734 .build();
6735
6736 let mut router = McpRouter::new()
6737 .tool(public_tool)
6738 .tool(admin_tool)
6739 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| tool.name != "admin"));
6740
6741 init_router(&mut router).await;
6743
6744 let req = RouterRequest {
6745 id: RequestId::Number(1),
6746 inner: McpRequest::ListTools(ListToolsParams::default()),
6747 extensions: Extensions::new(),
6748 };
6749
6750 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6751
6752 match resp.inner {
6753 Ok(McpResponse::ListTools(result)) => {
6754 assert_eq!(result.tools.len(), 1);
6756 assert_eq!(result.tools[0].name, "public");
6757 }
6758 _ => panic!("Expected ListTools response"),
6759 }
6760 }
6761
6762 #[tokio::test]
6763 async fn test_tool_filter_call_denied() {
6764 use crate::filter::CapabilityFilter;
6765 use crate::tool::Tool;
6766
6767 let admin_tool = ToolBuilder::new("admin")
6768 .description("Admin tool")
6769 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6770 .build();
6771
6772 let mut router = McpRouter::new()
6773 .tool(admin_tool)
6774 .tool_filter(CapabilityFilter::new(|_, _: &Tool| false)); init_router(&mut router).await;
6778
6779 let req = RouterRequest {
6780 id: RequestId::Number(1),
6781 inner: McpRequest::CallTool(CallToolParams {
6782 input_responses: None,
6783 request_state: None,
6784 name: "admin".to_string(),
6785 arguments: serde_json::json!({"a": 1, "b": 2}),
6786 meta: None,
6787 task: None,
6788 }),
6789 extensions: Extensions::new(),
6790 };
6791
6792 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6793
6794 match resp.inner {
6796 Err(e) => {
6797 assert_eq!(e.code, -32601); }
6799 _ => panic!("Expected JsonRpc error"),
6800 }
6801 }
6802
6803 #[tokio::test]
6804 async fn test_tool_filter_call_allowed() {
6805 use crate::filter::CapabilityFilter;
6806 use crate::tool::Tool;
6807
6808 let public_tool = ToolBuilder::new("public")
6809 .description("Public tool")
6810 .handler(|input: AddInput| async move {
6811 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6812 })
6813 .build();
6814
6815 let mut router = McpRouter::new()
6816 .tool(public_tool)
6817 .tool_filter(CapabilityFilter::new(|_, _: &Tool| true)); init_router(&mut router).await;
6821
6822 let req = RouterRequest {
6823 id: RequestId::Number(1),
6824 inner: McpRequest::CallTool(CallToolParams {
6825 input_responses: None,
6826 request_state: None,
6827 name: "public".to_string(),
6828 arguments: serde_json::json!({"a": 1, "b": 2}),
6829 meta: None,
6830 task: None,
6831 }),
6832 extensions: Extensions::new(),
6833 };
6834
6835 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6836
6837 match resp.inner {
6838 Ok(McpResponse::CallTool(result)) => {
6839 assert!(!result.is_error);
6840 }
6841 _ => panic!("Expected CallTool response"),
6842 }
6843 }
6844
6845 #[tokio::test]
6846 async fn test_tool_filter_custom_denial() {
6847 use crate::filter::{CapabilityFilter, DenialBehavior};
6848 use crate::tool::Tool;
6849
6850 let admin_tool = ToolBuilder::new("admin")
6851 .description("Admin tool")
6852 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6853 .build();
6854
6855 let mut router = McpRouter::new().tool(admin_tool).tool_filter(
6856 CapabilityFilter::new(|_, _: &Tool| false)
6857 .denial_behavior(DenialBehavior::Unauthorized),
6858 );
6859
6860 init_router(&mut router).await;
6862
6863 let req = RouterRequest {
6864 id: RequestId::Number(1),
6865 inner: McpRequest::CallTool(CallToolParams {
6866 input_responses: None,
6867 request_state: None,
6868 name: "admin".to_string(),
6869 arguments: serde_json::json!({"a": 1, "b": 2}),
6870 meta: None,
6871 task: None,
6872 }),
6873 extensions: Extensions::new(),
6874 };
6875
6876 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6877
6878 match resp.inner {
6880 Err(e) => {
6881 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
6883 }
6884 _ => panic!("Expected JsonRpc error"),
6885 }
6886 }
6887
6888 #[tokio::test]
6889 async fn test_resource_filter_list() {
6890 use crate::filter::CapabilityFilter;
6891 use crate::resource::{Resource, ResourceBuilder};
6892
6893 let public_resource = ResourceBuilder::new("file:///public.txt")
6894 .name("Public File")
6895 .text("public content");
6896
6897 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6898 .name("Secret File")
6899 .text("secret content");
6900
6901 let mut router = McpRouter::new()
6902 .resource(public_resource)
6903 .resource(secret_resource)
6904 .resource_filter(CapabilityFilter::new(|_, r: &Resource| {
6905 !r.name.contains("Secret")
6906 }));
6907
6908 init_router(&mut router).await;
6910
6911 let req = RouterRequest {
6912 id: RequestId::Number(1),
6913 inner: McpRequest::ListResources(ListResourcesParams::default()),
6914 extensions: Extensions::new(),
6915 };
6916
6917 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6918
6919 match resp.inner {
6920 Ok(McpResponse::ListResources(result)) => {
6921 assert_eq!(result.resources.len(), 1);
6923 assert_eq!(result.resources[0].name, "Public File");
6924 }
6925 _ => panic!("Expected ListResources response"),
6926 }
6927 }
6928
6929 #[tokio::test]
6930 async fn test_resource_filter_read_denied() {
6931 use crate::filter::CapabilityFilter;
6932 use crate::resource::{Resource, ResourceBuilder};
6933
6934 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6935 .name("Secret File")
6936 .text("secret content");
6937
6938 let mut router = McpRouter::new()
6939 .resource(secret_resource)
6940 .resource_filter(CapabilityFilter::new(|_, _: &Resource| false)); init_router(&mut router).await;
6944
6945 let req = RouterRequest {
6946 id: RequestId::Number(1),
6947 inner: McpRequest::ReadResource(ReadResourceParams {
6948 input_responses: None,
6949 request_state: None,
6950 uri: "file:///secret.txt".to_string(),
6951 meta: None,
6952 }),
6953 extensions: Extensions::new(),
6954 };
6955
6956 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6957
6958 match resp.inner {
6960 Err(e) => {
6961 assert_eq!(e.code, -32601); }
6963 _ => panic!("Expected JsonRpc error"),
6964 }
6965 }
6966
6967 #[tokio::test]
6968 async fn test_resource_filter_read_allowed() {
6969 use crate::filter::CapabilityFilter;
6970 use crate::resource::{Resource, ResourceBuilder};
6971
6972 let public_resource = ResourceBuilder::new("file:///public.txt")
6973 .name("Public File")
6974 .text("public content");
6975
6976 let mut router = McpRouter::new()
6977 .resource(public_resource)
6978 .resource_filter(CapabilityFilter::new(|_, _: &Resource| true)); init_router(&mut router).await;
6982
6983 let req = RouterRequest {
6984 id: RequestId::Number(1),
6985 inner: McpRequest::ReadResource(ReadResourceParams {
6986 input_responses: None,
6987 request_state: None,
6988 uri: "file:///public.txt".to_string(),
6989 meta: None,
6990 }),
6991 extensions: Extensions::new(),
6992 };
6993
6994 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6995
6996 match resp.inner {
6997 Ok(McpResponse::ReadResource(result)) => {
6998 assert_eq!(result.contents.len(), 1);
6999 assert_eq!(result.contents[0].text.as_deref(), Some("public content"));
7000 }
7001 _ => panic!("Expected ReadResource response"),
7002 }
7003 }
7004
7005 #[tokio::test]
7006 async fn test_resource_filter_custom_denial() {
7007 use crate::filter::{CapabilityFilter, DenialBehavior};
7008 use crate::resource::{Resource, ResourceBuilder};
7009
7010 let secret_resource = ResourceBuilder::new("file:///secret.txt")
7011 .name("Secret File")
7012 .text("secret content");
7013
7014 let mut router = McpRouter::new().resource(secret_resource).resource_filter(
7015 CapabilityFilter::new(|_, _: &Resource| false)
7016 .denial_behavior(DenialBehavior::Unauthorized),
7017 );
7018
7019 init_router(&mut router).await;
7021
7022 let req = RouterRequest {
7023 id: RequestId::Number(1),
7024 inner: McpRequest::ReadResource(ReadResourceParams {
7025 input_responses: None,
7026 request_state: None,
7027 uri: "file:///secret.txt".to_string(),
7028 meta: None,
7029 }),
7030 extensions: Extensions::new(),
7031 };
7032
7033 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7034
7035 match resp.inner {
7037 Err(e) => {
7038 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7040 }
7041 _ => panic!("Expected JsonRpc error"),
7042 }
7043 }
7044
7045 #[tokio::test]
7046 async fn test_prompt_filter_list() {
7047 use crate::filter::CapabilityFilter;
7048 use crate::prompt::{Prompt, PromptBuilder};
7049
7050 let public_prompt = PromptBuilder::new("greeting")
7051 .description("A greeting")
7052 .user_message("Hello!");
7053
7054 let admin_prompt = PromptBuilder::new("system_debug")
7055 .description("Admin prompt")
7056 .user_message("Debug");
7057
7058 let mut router = McpRouter::new()
7059 .prompt(public_prompt)
7060 .prompt(admin_prompt)
7061 .prompt_filter(CapabilityFilter::new(|_, p: &Prompt| {
7062 !p.name.contains("system")
7063 }));
7064
7065 init_router(&mut router).await;
7067
7068 let req = RouterRequest {
7069 id: RequestId::Number(1),
7070 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7071 extensions: Extensions::new(),
7072 };
7073
7074 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7075
7076 match resp.inner {
7077 Ok(McpResponse::ListPrompts(result)) => {
7078 assert_eq!(result.prompts.len(), 1);
7080 assert_eq!(result.prompts[0].name, "greeting");
7081 }
7082 _ => panic!("Expected ListPrompts response"),
7083 }
7084 }
7085
7086 #[tokio::test]
7087 async fn test_prompt_filter_get_denied() {
7088 use crate::filter::CapabilityFilter;
7089 use crate::prompt::{Prompt, PromptBuilder};
7090 use std::collections::HashMap;
7091
7092 let admin_prompt = PromptBuilder::new("system_debug")
7093 .description("Admin prompt")
7094 .user_message("Debug");
7095
7096 let mut router = McpRouter::new()
7097 .prompt(admin_prompt)
7098 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| false)); init_router(&mut router).await;
7102
7103 let req = RouterRequest {
7104 id: RequestId::Number(1),
7105 inner: McpRequest::GetPrompt(GetPromptParams {
7106 input_responses: None,
7107 request_state: None,
7108 name: "system_debug".to_string(),
7109 arguments: HashMap::new(),
7110 meta: None,
7111 }),
7112 extensions: Extensions::new(),
7113 };
7114
7115 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7116
7117 match resp.inner {
7119 Err(e) => {
7120 assert_eq!(e.code, -32601); }
7122 _ => panic!("Expected JsonRpc error"),
7123 }
7124 }
7125
7126 #[tokio::test]
7127 async fn test_prompt_filter_get_allowed() {
7128 use crate::filter::CapabilityFilter;
7129 use crate::prompt::{Prompt, PromptBuilder};
7130 use std::collections::HashMap;
7131
7132 let public_prompt = PromptBuilder::new("greeting")
7133 .description("A greeting")
7134 .user_message("Hello!");
7135
7136 let mut router = McpRouter::new()
7137 .prompt(public_prompt)
7138 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| true)); init_router(&mut router).await;
7142
7143 let req = RouterRequest {
7144 id: RequestId::Number(1),
7145 inner: McpRequest::GetPrompt(GetPromptParams {
7146 input_responses: None,
7147 request_state: None,
7148 name: "greeting".to_string(),
7149 arguments: HashMap::new(),
7150 meta: None,
7151 }),
7152 extensions: Extensions::new(),
7153 };
7154
7155 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7156
7157 match resp.inner {
7158 Ok(McpResponse::GetPrompt(result)) => {
7159 assert_eq!(result.messages.len(), 1);
7160 }
7161 _ => panic!("Expected GetPrompt response"),
7162 }
7163 }
7164
7165 #[tokio::test]
7166 async fn test_prompt_filter_custom_denial() {
7167 use crate::filter::{CapabilityFilter, DenialBehavior};
7168 use crate::prompt::{Prompt, PromptBuilder};
7169 use std::collections::HashMap;
7170
7171 let admin_prompt = PromptBuilder::new("system_debug")
7172 .description("Admin prompt")
7173 .user_message("Debug");
7174
7175 let mut router = McpRouter::new().prompt(admin_prompt).prompt_filter(
7176 CapabilityFilter::new(|_, _: &Prompt| false)
7177 .denial_behavior(DenialBehavior::Unauthorized),
7178 );
7179
7180 init_router(&mut router).await;
7182
7183 let req = RouterRequest {
7184 id: RequestId::Number(1),
7185 inner: McpRequest::GetPrompt(GetPromptParams {
7186 input_responses: None,
7187 request_state: None,
7188 name: "system_debug".to_string(),
7189 arguments: HashMap::new(),
7190 meta: None,
7191 }),
7192 extensions: Extensions::new(),
7193 };
7194
7195 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7196
7197 match resp.inner {
7199 Err(e) => {
7200 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7202 }
7203 _ => panic!("Expected JsonRpc error"),
7204 }
7205 }
7206
7207 #[derive(Debug, Deserialize, JsonSchema)]
7212 struct StringInput {
7213 value: String,
7214 }
7215
7216 #[tokio::test]
7217 async fn test_router_merge_tools() {
7218 let tool_a = ToolBuilder::new("tool_a")
7220 .description("Tool A")
7221 .handler(|_: StringInput| async move { Ok(CallToolResult::text("A")) })
7222 .build();
7223
7224 let router_a = McpRouter::new().tool(tool_a);
7225
7226 let tool_b = ToolBuilder::new("tool_b")
7228 .description("Tool B")
7229 .handler(|_: StringInput| async move { Ok(CallToolResult::text("B")) })
7230 .build();
7231 let tool_c = ToolBuilder::new("tool_c")
7232 .description("Tool C")
7233 .handler(|_: StringInput| async move { Ok(CallToolResult::text("C")) })
7234 .build();
7235
7236 let router_b = McpRouter::new().tool(tool_b).tool(tool_c);
7237
7238 let mut merged = McpRouter::new()
7240 .server_info("merged", "1.0")
7241 .merge(router_a)
7242 .merge(router_b);
7243
7244 init_router(&mut merged).await;
7245
7246 let req = RouterRequest {
7248 id: RequestId::Number(1),
7249 inner: McpRequest::ListTools(ListToolsParams::default()),
7250 extensions: Extensions::new(),
7251 };
7252
7253 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7254
7255 match resp.inner {
7256 Ok(McpResponse::ListTools(result)) => {
7257 assert_eq!(result.tools.len(), 3);
7258 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7259 assert!(names.contains(&"tool_a"));
7260 assert!(names.contains(&"tool_b"));
7261 assert!(names.contains(&"tool_c"));
7262 }
7263 _ => panic!("Expected ListTools response"),
7264 }
7265 }
7266
7267 #[tokio::test]
7268 async fn test_router_merge_overwrites_duplicates() {
7269 let tool_v1 = ToolBuilder::new("shared")
7271 .description("Version 1")
7272 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v1")) })
7273 .build();
7274
7275 let router_a = McpRouter::new().tool(tool_v1);
7276
7277 let tool_v2 = ToolBuilder::new("shared")
7279 .description("Version 2")
7280 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v2")) })
7281 .build();
7282
7283 let router_b = McpRouter::new().tool(tool_v2);
7284
7285 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7287
7288 init_router(&mut merged).await;
7289
7290 let req = RouterRequest {
7291 id: RequestId::Number(1),
7292 inner: McpRequest::ListTools(ListToolsParams::default()),
7293 extensions: Extensions::new(),
7294 };
7295
7296 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7297
7298 match resp.inner {
7299 Ok(McpResponse::ListTools(result)) => {
7300 assert_eq!(result.tools.len(), 1);
7301 assert_eq!(result.tools[0].name, "shared");
7302 assert_eq!(result.tools[0].description.as_deref(), Some("Version 2"));
7303 }
7304 _ => panic!("Expected ListTools response"),
7305 }
7306 }
7307
7308 #[tokio::test]
7309 async fn test_router_merge_resources() {
7310 use crate::resource::ResourceBuilder;
7311
7312 let router_a = McpRouter::new().resource(
7314 ResourceBuilder::new("file:///a.txt")
7315 .name("File A")
7316 .text("content a"),
7317 );
7318
7319 let router_b = McpRouter::new().resource(
7320 ResourceBuilder::new("file:///b.txt")
7321 .name("File B")
7322 .text("content b"),
7323 );
7324
7325 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7326
7327 init_router(&mut merged).await;
7328
7329 let req = RouterRequest {
7330 id: RequestId::Number(1),
7331 inner: McpRequest::ListResources(ListResourcesParams::default()),
7332 extensions: Extensions::new(),
7333 };
7334
7335 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7336
7337 match resp.inner {
7338 Ok(McpResponse::ListResources(result)) => {
7339 assert_eq!(result.resources.len(), 2);
7340 let uris: Vec<&str> = result.resources.iter().map(|r| r.uri.as_str()).collect();
7341 assert!(uris.contains(&"file:///a.txt"));
7342 assert!(uris.contains(&"file:///b.txt"));
7343 }
7344 _ => panic!("Expected ListResources response"),
7345 }
7346 }
7347
7348 #[tokio::test]
7349 async fn test_router_merge_prompts() {
7350 use crate::prompt::PromptBuilder;
7351
7352 let router_a =
7353 McpRouter::new().prompt(PromptBuilder::new("prompt_a").user_message("Hello A"));
7354
7355 let router_b =
7356 McpRouter::new().prompt(PromptBuilder::new("prompt_b").user_message("Hello B"));
7357
7358 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7359
7360 init_router(&mut merged).await;
7361
7362 let req = RouterRequest {
7363 id: RequestId::Number(1),
7364 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
7365 extensions: Extensions::new(),
7366 };
7367
7368 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7369
7370 match resp.inner {
7371 Ok(McpResponse::ListPrompts(result)) => {
7372 assert_eq!(result.prompts.len(), 2);
7373 let names: Vec<&str> = result.prompts.iter().map(|p| p.name.as_str()).collect();
7374 assert!(names.contains(&"prompt_a"));
7375 assert!(names.contains(&"prompt_b"));
7376 }
7377 _ => panic!("Expected ListPrompts response"),
7378 }
7379 }
7380
7381 #[tokio::test]
7382 async fn test_router_nest_prefixes_tools() {
7383 let tool_query = ToolBuilder::new("query")
7385 .description("Query the database")
7386 .handler(|_: StringInput| async move { Ok(CallToolResult::text("query result")) })
7387 .build();
7388 let tool_insert = ToolBuilder::new("insert")
7389 .description("Insert into database")
7390 .handler(|_: StringInput| async move { Ok(CallToolResult::text("insert result")) })
7391 .build();
7392
7393 let db_router = McpRouter::new().tool(tool_query).tool(tool_insert);
7394
7395 let mut router = McpRouter::new()
7397 .server_info("nested", "1.0")
7398 .nest("db", db_router);
7399
7400 init_router(&mut router).await;
7401
7402 let req = RouterRequest {
7403 id: RequestId::Number(1),
7404 inner: McpRequest::ListTools(ListToolsParams::default()),
7405 extensions: Extensions::new(),
7406 };
7407
7408 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7409
7410 match resp.inner {
7411 Ok(McpResponse::ListTools(result)) => {
7412 assert_eq!(result.tools.len(), 2);
7413 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7414 assert!(names.contains(&"db.query"));
7415 assert!(names.contains(&"db.insert"));
7416 }
7417 _ => panic!("Expected ListTools response"),
7418 }
7419 }
7420
7421 #[tokio::test]
7422 async fn test_router_nest_call_prefixed_tool() {
7423 let tool = ToolBuilder::new("echo")
7424 .description("Echo input")
7425 .handler(|input: StringInput| async move { Ok(CallToolResult::text(&input.value)) })
7426 .build();
7427
7428 let nested_router = McpRouter::new().tool(tool);
7429
7430 let mut router = McpRouter::new().nest("api", nested_router);
7431
7432 init_router(&mut router).await;
7433
7434 let req = RouterRequest {
7436 id: RequestId::Number(1),
7437 inner: McpRequest::CallTool(CallToolParams {
7438 input_responses: None,
7439 request_state: None,
7440 name: "api.echo".to_string(),
7441 arguments: serde_json::json!({"value": "hello world"}),
7442 meta: None,
7443 task: None,
7444 }),
7445 extensions: Extensions::new(),
7446 };
7447
7448 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7449
7450 match resp.inner {
7451 Ok(McpResponse::CallTool(result)) => {
7452 assert!(!result.is_error);
7453 match &result.content[0] {
7454 Content::Text { text, .. } => assert_eq!(text, "hello world"),
7455 _ => panic!("Expected text content"),
7456 }
7457 }
7458 _ => panic!("Expected CallTool response"),
7459 }
7460 }
7461
7462 #[tokio::test]
7463 async fn test_router_multiple_nests() {
7464 let db_tool = ToolBuilder::new("query")
7465 .description("Database query")
7466 .handler(|_: StringInput| async move { Ok(CallToolResult::text("db")) })
7467 .build();
7468
7469 let api_tool = ToolBuilder::new("fetch")
7470 .description("API fetch")
7471 .handler(|_: StringInput| async move { Ok(CallToolResult::text("api")) })
7472 .build();
7473
7474 let db_router = McpRouter::new().tool(db_tool);
7475 let api_router = McpRouter::new().tool(api_tool);
7476
7477 let mut router = McpRouter::new()
7478 .nest("db", db_router)
7479 .nest("api", api_router);
7480
7481 init_router(&mut router).await;
7482
7483 let req = RouterRequest {
7484 id: RequestId::Number(1),
7485 inner: McpRequest::ListTools(ListToolsParams::default()),
7486 extensions: Extensions::new(),
7487 };
7488
7489 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7490
7491 match resp.inner {
7492 Ok(McpResponse::ListTools(result)) => {
7493 assert_eq!(result.tools.len(), 2);
7494 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7495 assert!(names.contains(&"db.query"));
7496 assert!(names.contains(&"api.fetch"));
7497 }
7498 _ => panic!("Expected ListTools response"),
7499 }
7500 }
7501
7502 #[tokio::test]
7503 async fn test_router_merge_and_nest_combined() {
7504 let tool_a = ToolBuilder::new("local")
7506 .description("Local tool")
7507 .handler(|_: StringInput| async move { Ok(CallToolResult::text("local")) })
7508 .build();
7509
7510 let nested_tool = ToolBuilder::new("remote")
7511 .description("Remote tool")
7512 .handler(|_: StringInput| async move { Ok(CallToolResult::text("remote")) })
7513 .build();
7514
7515 let nested_router = McpRouter::new().tool(nested_tool);
7516
7517 let mut router = McpRouter::new()
7518 .tool(tool_a)
7519 .nest("external", nested_router);
7520
7521 init_router(&mut router).await;
7522
7523 let req = RouterRequest {
7524 id: RequestId::Number(1),
7525 inner: McpRequest::ListTools(ListToolsParams::default()),
7526 extensions: Extensions::new(),
7527 };
7528
7529 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7530
7531 match resp.inner {
7532 Ok(McpResponse::ListTools(result)) => {
7533 assert_eq!(result.tools.len(), 2);
7534 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7535 assert!(names.contains(&"local"));
7536 assert!(names.contains(&"external.remote"));
7537 }
7538 _ => panic!("Expected ListTools response"),
7539 }
7540 }
7541
7542 #[tokio::test]
7543 async fn test_router_merge_preserves_server_info() {
7544 let child_router = McpRouter::new()
7545 .server_info("child", "2.0")
7546 .instructions("Child instructions");
7547
7548 let mut router = McpRouter::new()
7549 .server_info("parent", "1.0")
7550 .instructions("Parent instructions")
7551 .merge(child_router);
7552
7553 init_router(&mut router).await;
7554
7555 let init_req = RouterRequest {
7557 id: RequestId::Number(99),
7558 inner: McpRequest::Initialize(InitializeParams {
7559 protocol_version: "2025-11-25".to_string(),
7560 capabilities: ClientCapabilities::default(),
7561 client_info: Implementation {
7562 name: "test".to_string(),
7563 version: "1.0".to_string(),
7564 ..Default::default()
7565 },
7566 meta: None,
7567 }),
7568 extensions: Extensions::new(),
7569 };
7570
7571 let child_router2 = McpRouter::new().server_info("child", "2.0");
7573 let mut fresh_router = McpRouter::new()
7574 .server_info("parent", "1.0")
7575 .merge(child_router2);
7576
7577 let resp = fresh_router
7578 .ready()
7579 .await
7580 .unwrap()
7581 .call(init_req)
7582 .await
7583 .unwrap();
7584
7585 match resp.inner {
7586 Ok(McpResponse::Initialize(result)) => {
7587 assert_eq!(result.server_info.name, "parent");
7588 assert_eq!(result.server_info.version, "1.0");
7589 }
7590 _ => panic!("Expected Initialize response"),
7591 }
7592 }
7593
7594 #[tokio::test]
7599 async fn test_auto_instructions_tools_only() {
7600 let tool_a = ToolBuilder::new("alpha")
7601 .description("Alpha tool")
7602 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7603 .build();
7604 let tool_b = ToolBuilder::new("beta")
7605 .description("Beta tool")
7606 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7607 .build();
7608
7609 let mut router = McpRouter::new()
7610 .auto_instructions()
7611 .tool(tool_a)
7612 .tool(tool_b);
7613
7614 let resp = send_initialize(&mut router).await;
7615 let instructions = resp.instructions.expect("should have instructions");
7616
7617 assert!(instructions.contains("## Tools"));
7618 assert!(instructions.contains("- **alpha**: Alpha tool"));
7619 assert!(instructions.contains("- **beta**: Beta tool"));
7620 assert!(!instructions.contains("## Resources"));
7622 assert!(!instructions.contains("## Prompts"));
7623 }
7624
7625 #[tokio::test]
7626 async fn test_auto_instructions_with_annotations() {
7627 let read_only_tool = ToolBuilder::new("query")
7628 .description("Run a query")
7629 .read_only()
7630 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7631 .build();
7632 let destructive_tool = ToolBuilder::new("delete")
7633 .description("Delete a record")
7634 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7635 .build();
7636 let idempotent_tool = ToolBuilder::new("upsert")
7637 .description("Upsert a record")
7638 .non_destructive()
7639 .idempotent()
7640 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7641 .build();
7642
7643 let mut router = McpRouter::new()
7644 .auto_instructions()
7645 .tool(read_only_tool)
7646 .tool(destructive_tool)
7647 .tool(idempotent_tool);
7648
7649 let resp = send_initialize(&mut router).await;
7650 let instructions = resp.instructions.unwrap();
7651
7652 assert!(instructions.contains("- **query**: Run a query [read-only]"));
7653 assert!(instructions.contains("- **delete**: Delete a record\n"));
7655 assert!(instructions.contains("- **upsert**: Upsert a record [idempotent]"));
7656 }
7657
7658 #[tokio::test]
7659 async fn test_auto_instructions_with_resources() {
7660 use crate::resource::ResourceBuilder;
7661
7662 let resource = ResourceBuilder::new("file:///schema.sql")
7663 .name("Schema")
7664 .description("Database schema")
7665 .text("CREATE TABLE ...");
7666
7667 let mut router = McpRouter::new().auto_instructions().resource(resource);
7668
7669 let resp = send_initialize(&mut router).await;
7670 let instructions = resp.instructions.unwrap();
7671
7672 assert!(instructions.contains("## Resources"));
7673 assert!(instructions.contains("- **file:///schema.sql**: Database schema"));
7674 assert!(!instructions.contains("## Tools"));
7675 }
7676
7677 #[tokio::test]
7678 async fn test_auto_instructions_with_resource_templates() {
7679 use crate::resource::ResourceTemplateBuilder;
7680
7681 let template = ResourceTemplateBuilder::new("file:///{path}")
7682 .name("File")
7683 .description("Read a file by path")
7684 .handler(
7685 |_uri: String, _vars: std::collections::HashMap<String, String>| async move {
7686 Ok(crate::ReadResourceResult::text("content", "text/plain"))
7687 },
7688 );
7689
7690 let mut router = McpRouter::new()
7691 .auto_instructions()
7692 .resource_template(template);
7693
7694 let resp = send_initialize(&mut router).await;
7695 let instructions = resp.instructions.unwrap();
7696
7697 assert!(instructions.contains("## Resources"));
7698 assert!(instructions.contains("- **file:///{path}**: Read a file by path"));
7699 }
7700
7701 #[tokio::test]
7702 async fn test_auto_instructions_with_prompts() {
7703 use crate::prompt::PromptBuilder;
7704
7705 let prompt = PromptBuilder::new("write_query")
7706 .description("Help write a SQL query")
7707 .user_message("Write a query for: {task}");
7708
7709 let mut router = McpRouter::new().auto_instructions().prompt(prompt);
7710
7711 let resp = send_initialize(&mut router).await;
7712 let instructions = resp.instructions.unwrap();
7713
7714 assert!(instructions.contains("## Prompts"));
7715 assert!(instructions.contains("- **write_query**: Help write a SQL query"));
7716 assert!(!instructions.contains("## Tools"));
7717 }
7718
7719 #[tokio::test]
7720 async fn test_auto_instructions_all_sections() {
7721 use crate::prompt::PromptBuilder;
7722 use crate::resource::ResourceBuilder;
7723
7724 let tool = ToolBuilder::new("query")
7725 .description("Execute SQL")
7726 .read_only()
7727 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7728 .build();
7729 let resource = ResourceBuilder::new("db://schema")
7730 .name("Schema")
7731 .description("Full database schema")
7732 .text("schema");
7733 let prompt = PromptBuilder::new("write_query")
7734 .description("Help write a SQL query")
7735 .user_message("Write a query");
7736
7737 let mut router = McpRouter::new()
7738 .auto_instructions()
7739 .tool(tool)
7740 .resource(resource)
7741 .prompt(prompt);
7742
7743 let resp = send_initialize(&mut router).await;
7744 let instructions = resp.instructions.unwrap();
7745
7746 assert!(instructions.contains("## Tools"));
7748 assert!(instructions.contains("## Resources"));
7749 assert!(instructions.contains("## Prompts"));
7750
7751 let tools_pos = instructions.find("## Tools").unwrap();
7753 let resources_pos = instructions.find("## Resources").unwrap();
7754 let prompts_pos = instructions.find("## Prompts").unwrap();
7755 assert!(tools_pos < resources_pos);
7756 assert!(resources_pos < prompts_pos);
7757 }
7758
7759 #[tokio::test]
7760 async fn test_auto_instructions_with_prefix_and_suffix() {
7761 let tool = ToolBuilder::new("echo")
7762 .description("Echo input")
7763 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7764 .build();
7765
7766 let mut router = McpRouter::new()
7767 .auto_instructions_with(
7768 Some("This server provides echo capabilities."),
7769 Some("Contact admin@example.com for support."),
7770 )
7771 .tool(tool);
7772
7773 let resp = send_initialize(&mut router).await;
7774 let instructions = resp.instructions.unwrap();
7775
7776 assert!(instructions.starts_with("This server provides echo capabilities."));
7777 assert!(instructions.ends_with("Contact admin@example.com for support."));
7778 assert!(instructions.contains("## Tools"));
7779 assert!(instructions.contains("- **echo**: Echo input"));
7780 }
7781
7782 #[tokio::test]
7783 async fn test_auto_instructions_prefix_only() {
7784 let tool = ToolBuilder::new("echo")
7785 .description("Echo input")
7786 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7787 .build();
7788
7789 let mut router = McpRouter::new()
7790 .auto_instructions_with(Some("My server intro."), None::<String>)
7791 .tool(tool);
7792
7793 let resp = send_initialize(&mut router).await;
7794 let instructions = resp.instructions.unwrap();
7795
7796 assert!(instructions.starts_with("My server intro."));
7797 assert!(instructions.contains("- **echo**: Echo input"));
7798 }
7799
7800 #[tokio::test]
7801 async fn test_auto_instructions_empty_router() {
7802 let mut router = McpRouter::new().auto_instructions();
7803
7804 let resp = send_initialize(&mut router).await;
7805 let instructions = resp.instructions.expect("should have instructions");
7806
7807 assert!(!instructions.contains("## Tools"));
7809 assert!(!instructions.contains("## Resources"));
7810 assert!(!instructions.contains("## Prompts"));
7811 assert!(instructions.is_empty());
7812 }
7813
7814 #[tokio::test]
7815 async fn test_auto_instructions_overrides_manual() {
7816 let tool = ToolBuilder::new("echo")
7817 .description("Echo input")
7818 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7819 .build();
7820
7821 let mut router = McpRouter::new()
7822 .instructions("This will be overridden")
7823 .auto_instructions()
7824 .tool(tool);
7825
7826 let resp = send_initialize(&mut router).await;
7827 let instructions = resp.instructions.unwrap();
7828
7829 assert!(!instructions.contains("This will be overridden"));
7830 assert!(instructions.contains("- **echo**: Echo input"));
7831 }
7832
7833 #[tokio::test]
7834 async fn test_no_auto_instructions_returns_manual() {
7835 let tool = ToolBuilder::new("echo")
7836 .description("Echo input")
7837 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7838 .build();
7839
7840 let mut router = McpRouter::new()
7841 .instructions("Manual instructions here")
7842 .tool(tool);
7843
7844 let resp = send_initialize(&mut router).await;
7845 let instructions = resp.instructions.unwrap();
7846
7847 assert_eq!(instructions, "Manual instructions here");
7848 }
7849
7850 #[tokio::test]
7851 async fn test_auto_instructions_no_description_fallback() {
7852 let tool = ToolBuilder::new("mystery")
7853 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7854 .build();
7855
7856 let mut router = McpRouter::new().auto_instructions().tool(tool);
7857
7858 let resp = send_initialize(&mut router).await;
7859 let instructions = resp.instructions.unwrap();
7860
7861 assert!(instructions.contains("- **mystery**: No description"));
7862 }
7863
7864 #[tokio::test]
7865 async fn test_auto_instructions_sorted_alphabetically() {
7866 let tool_z = ToolBuilder::new("zebra")
7867 .description("Z tool")
7868 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7869 .build();
7870 let tool_a = ToolBuilder::new("alpha")
7871 .description("A tool")
7872 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7873 .build();
7874 let tool_m = ToolBuilder::new("middle")
7875 .description("M tool")
7876 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7877 .build();
7878
7879 let mut router = McpRouter::new()
7880 .auto_instructions()
7881 .tool(tool_z)
7882 .tool(tool_a)
7883 .tool(tool_m);
7884
7885 let resp = send_initialize(&mut router).await;
7886 let instructions = resp.instructions.unwrap();
7887
7888 let alpha_pos = instructions.find("**alpha**").unwrap();
7889 let middle_pos = instructions.find("**middle**").unwrap();
7890 let zebra_pos = instructions.find("**zebra**").unwrap();
7891 assert!(alpha_pos < middle_pos);
7892 assert!(middle_pos < zebra_pos);
7893 }
7894
7895 #[tokio::test]
7896 async fn test_auto_instructions_read_only_and_idempotent_tags() {
7897 let tool = ToolBuilder::new("safe_update")
7898 .description("Safe update operation")
7899 .idempotent()
7900 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7901 .build();
7902
7903 let mut router = McpRouter::new().auto_instructions().tool(tool);
7904
7905 let resp = send_initialize(&mut router).await;
7906 let instructions = resp.instructions.unwrap();
7907
7908 assert!(
7909 instructions.contains("[idempotent]"),
7910 "got: {}",
7911 instructions
7912 );
7913 }
7914
7915 #[tokio::test]
7916 async fn test_auto_instructions_lazy_generation() {
7917 let mut router = McpRouter::new().auto_instructions();
7920
7921 let tool = ToolBuilder::new("late_tool")
7922 .description("Added after auto_instructions")
7923 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7924 .build();
7925
7926 router = router.tool(tool);
7927
7928 let resp = send_initialize(&mut router).await;
7929 let instructions = resp.instructions.unwrap();
7930
7931 assert!(instructions.contains("- **late_tool**: Added after auto_instructions"));
7932 }
7933
7934 #[tokio::test]
7935 async fn test_auto_instructions_multiple_annotation_tags() {
7936 let tool = ToolBuilder::new("update")
7937 .description("Update a record")
7938 .annotations(ToolAnnotations {
7939 read_only_hint: true,
7940 idempotent_hint: true,
7941 ..Default::default()
7942 })
7943 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7944 .build();
7945
7946 let mut router = McpRouter::new().auto_instructions().tool(tool);
7947
7948 let resp = send_initialize(&mut router).await;
7949 let instructions = resp.instructions.unwrap();
7950
7951 assert!(
7952 instructions.contains("[read-only, idempotent]"),
7953 "got: {}",
7954 instructions
7955 );
7956 }
7957
7958 #[tokio::test]
7959 async fn test_auto_instructions_no_annotations_no_tags() {
7960 let tool = ToolBuilder::new("fetch")
7962 .description("Fetch data")
7963 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7964 .build();
7965
7966 let mut router = McpRouter::new().auto_instructions().tool(tool);
7967
7968 let resp = send_initialize(&mut router).await;
7969 let instructions = resp.instructions.unwrap();
7970
7971 assert!(
7973 !instructions.contains('['),
7974 "should have no tags, got: {}",
7975 instructions
7976 );
7977 assert!(instructions.contains("- **fetch**: Fetch data"));
7978 }
7979
7980 async fn send_initialize(router: &mut McpRouter) -> InitializeResult {
7982 let init_req = RouterRequest {
7983 id: RequestId::Number(0),
7984 inner: McpRequest::Initialize(InitializeParams {
7985 protocol_version: "2025-11-25".to_string(),
7986 capabilities: ClientCapabilities {
7987 roots: None,
7988 sampling: None,
7989 elicitation: None,
7990 tasks: None,
7991 experimental: None,
7992 extensions: None,
7993 },
7994 client_info: Implementation {
7995 name: "test".to_string(),
7996 version: "1.0".to_string(),
7997 ..Default::default()
7998 },
7999 meta: None,
8000 }),
8001 extensions: Extensions::new(),
8002 };
8003 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8004 match resp.inner {
8005 Ok(McpResponse::Initialize(result)) => result,
8006 other => panic!("Expected Initialize response, got {:?}", other),
8007 }
8008 }
8009
8010 #[tokio::test]
8011 async fn test_notify_tools_list_changed() {
8012 let (tx, mut rx) = crate::context::notification_channel(16);
8013
8014 let router = McpRouter::new()
8015 .server_info("test", "1.0")
8016 .with_notification_sender(tx);
8017
8018 assert!(router.notify_tools_list_changed());
8019
8020 let notification = rx.recv().await.unwrap();
8021 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8022 }
8023
8024 #[tokio::test]
8025 async fn test_notify_prompts_list_changed() {
8026 let (tx, mut rx) = crate::context::notification_channel(16);
8027
8028 let router = McpRouter::new()
8029 .server_info("test", "1.0")
8030 .with_notification_sender(tx);
8031
8032 assert!(router.notify_prompts_list_changed());
8033
8034 let notification = rx.recv().await.unwrap();
8035 assert!(matches!(
8036 notification,
8037 ServerNotification::PromptsListChanged
8038 ));
8039 }
8040
8041 #[tokio::test]
8042 async fn test_notify_without_sender_returns_false() {
8043 let router = McpRouter::new().server_info("test", "1.0");
8044
8045 assert!(!router.notify_tools_list_changed());
8046 assert!(!router.notify_prompts_list_changed());
8047 assert!(!router.notify_resources_list_changed());
8048 }
8049
8050 #[tokio::test]
8051 async fn test_list_changed_capabilities_with_notification_sender() {
8052 let (tx, _rx) = crate::context::notification_channel(16);
8053 let tool = ToolBuilder::new("test")
8054 .description("test")
8055 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8056 .build();
8057
8058 let mut router = McpRouter::new()
8059 .server_info("test", "1.0")
8060 .tool(tool)
8061 .with_notification_sender(tx);
8062
8063 init_router(&mut router).await;
8064
8065 let caps = router.capabilities();
8066 let tools_cap = caps.tools.expect("tools capability should be present");
8067 assert!(
8068 tools_cap.list_changed,
8069 "tools.listChanged should be true when notification sender is configured"
8070 );
8071 }
8072
8073 #[tokio::test]
8074 async fn test_list_changed_capabilities_without_notification_sender() {
8075 let tool = ToolBuilder::new("test")
8076 .description("test")
8077 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8078 .build();
8079
8080 let mut router = McpRouter::new().server_info("test", "1.0").tool(tool);
8081
8082 init_router(&mut router).await;
8083
8084 let caps = router.capabilities();
8085 let tools_cap = caps.tools.expect("tools capability should be present");
8086 assert!(
8087 !tools_cap.list_changed,
8088 "tools.listChanged should be false without notification sender"
8089 );
8090 }
8091
8092 #[tokio::test]
8093 async fn test_set_logging_level_filters_messages() {
8094 let (tx, mut rx) = crate::context::notification_channel(16);
8095
8096 let mut router = McpRouter::new()
8097 .server_info("test", "1.0")
8098 .with_notification_sender(tx);
8099
8100 init_router(&mut router).await;
8101
8102 let set_level_req = RouterRequest {
8104 id: RequestId::Number(99),
8105 inner: McpRequest::SetLoggingLevel(SetLogLevelParams {
8106 level: LogLevel::Warning,
8107 meta: None,
8108 }),
8109 extensions: crate::context::Extensions::new(),
8110 };
8111 let resp = router
8112 .ready()
8113 .await
8114 .unwrap()
8115 .call(set_level_req)
8116 .await
8117 .unwrap();
8118 assert!(matches!(resp.inner, Ok(McpResponse::SetLoggingLevel(_))));
8119
8120 let ctx = router.create_context(RequestId::Number(100), None);
8122
8123 ctx.send_log(LoggingMessageParams::new(
8125 LogLevel::Error,
8126 serde_json::Value::Null,
8127 ));
8128 assert!(
8129 rx.try_recv().is_ok(),
8130 "Error should pass through Warning filter"
8131 );
8132
8133 ctx.send_log(LoggingMessageParams::new(
8135 LogLevel::Info,
8136 serde_json::Value::Null,
8137 ));
8138 assert!(
8139 rx.try_recv().is_err(),
8140 "Info should be filtered at Warning level"
8141 );
8142 }
8143
8144 #[test]
8145 fn test_paginate_no_page_size() {
8146 let items = vec![1, 2, 3, 4, 5];
8147 let (page, cursor) = paginate(items.clone(), None, None).unwrap();
8148 assert_eq!(page, items);
8149 assert!(cursor.is_none());
8150 }
8151
8152 #[test]
8153 fn test_paginate_first_page() {
8154 let items = vec![1, 2, 3, 4, 5];
8155 let (page, cursor) = paginate(items, None, Some(2)).unwrap();
8156 assert_eq!(page, vec![1, 2]);
8157 assert!(cursor.is_some());
8158 }
8159
8160 #[test]
8161 fn test_paginate_middle_page() {
8162 let items = vec![1, 2, 3, 4, 5];
8163 let (page1, cursor1) = paginate(items.clone(), None, Some(2)).unwrap();
8164 assert_eq!(page1, vec![1, 2]);
8165
8166 let (page2, cursor2) = paginate(items, cursor1.as_deref(), Some(2)).unwrap();
8167 assert_eq!(page2, vec![3, 4]);
8168 assert!(cursor2.is_some());
8169 }
8170
8171 #[test]
8172 fn test_paginate_last_page() {
8173 let items = vec![1, 2, 3, 4, 5];
8174 let cursor = encode_cursor(4);
8176 let (page, next) = paginate(items, Some(&cursor), Some(2)).unwrap();
8177 assert_eq!(page, vec![5]);
8178 assert!(next.is_none());
8179 }
8180
8181 #[test]
8182 fn test_paginate_exact_boundary() {
8183 let items = vec![1, 2, 3, 4];
8184 let (page, cursor) = paginate(items, None, Some(4)).unwrap();
8185 assert_eq!(page, vec![1, 2, 3, 4]);
8186 assert!(cursor.is_none());
8187 }
8188
8189 #[test]
8190 fn test_paginate_invalid_cursor() {
8191 let items = vec![1, 2, 3];
8192 let result = paginate(items, Some("not-valid-base64!@#$"), Some(2));
8193 assert!(result.is_err());
8194 }
8195
8196 #[test]
8197 fn test_cursor_round_trip() {
8198 let offset = 42;
8199 let encoded = encode_cursor(offset);
8200 let decoded = decode_cursor(&encoded).unwrap();
8201 assert_eq!(decoded, offset);
8202 }
8203
8204 #[tokio::test]
8205 async fn test_list_tools_pagination() {
8206 let tool_a = ToolBuilder::new("alpha")
8207 .description("a")
8208 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8209 .build();
8210 let tool_b = ToolBuilder::new("beta")
8211 .description("b")
8212 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8213 .build();
8214 let tool_c = ToolBuilder::new("gamma")
8215 .description("c")
8216 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8217 .build();
8218
8219 let mut router = McpRouter::new()
8220 .server_info("test", "1.0")
8221 .page_size(2)
8222 .tool(tool_a)
8223 .tool(tool_b)
8224 .tool(tool_c);
8225
8226 init_router(&mut router).await;
8227
8228 let req = RouterRequest {
8230 id: RequestId::Number(1),
8231 inner: McpRequest::ListTools(ListToolsParams {
8232 cursor: None,
8233 meta: None,
8234 }),
8235 extensions: Extensions::new(),
8236 };
8237 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8238 let (tools, next_cursor) = match resp.inner {
8239 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8240 other => panic!("Expected ListTools, got {:?}", other),
8241 };
8242 assert_eq!(tools.len(), 2);
8243 assert_eq!(tools[0].name, "alpha");
8244 assert_eq!(tools[1].name, "beta");
8245 assert!(next_cursor.is_some());
8246
8247 let req = RouterRequest {
8249 id: RequestId::Number(2),
8250 inner: McpRequest::ListTools(ListToolsParams {
8251 cursor: next_cursor,
8252 meta: None,
8253 }),
8254 extensions: Extensions::new(),
8255 };
8256 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8257 let (tools, next_cursor) = match resp.inner {
8258 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8259 other => panic!("Expected ListTools, got {:?}", other),
8260 };
8261 assert_eq!(tools.len(), 1);
8262 assert_eq!(tools[0].name, "gamma");
8263 assert!(next_cursor.is_none());
8264 }
8265
8266 #[tokio::test]
8267 async fn test_list_tools_no_pagination_by_default() {
8268 let tool_a = ToolBuilder::new("alpha")
8269 .description("a")
8270 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8271 .build();
8272 let tool_b = ToolBuilder::new("beta")
8273 .description("b")
8274 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8275 .build();
8276
8277 let mut router = McpRouter::new()
8278 .server_info("test", "1.0")
8279 .tool(tool_a)
8280 .tool(tool_b);
8281
8282 init_router(&mut router).await;
8283
8284 let req = RouterRequest {
8285 id: RequestId::Number(1),
8286 inner: McpRequest::ListTools(ListToolsParams {
8287 cursor: None,
8288 meta: None,
8289 }),
8290 extensions: Extensions::new(),
8291 };
8292 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8293 match resp.inner {
8294 Ok(McpResponse::ListTools(result)) => {
8295 assert_eq!(result.tools.len(), 2);
8296 assert!(result.next_cursor.is_none());
8297 }
8298 other => panic!("Expected ListTools, got {:?}", other),
8299 }
8300 }
8301
8302 #[cfg(feature = "dynamic-tools")]
8307 mod dynamic_tools_tests {
8308 use super::*;
8309
8310 #[tokio::test]
8311 async fn test_dynamic_tools_register_and_list() {
8312 let (router, registry) = McpRouter::new()
8313 .server_info("test", "1.0")
8314 .with_dynamic_tools();
8315
8316 let tool = ToolBuilder::new("dynamic_echo")
8317 .description("Dynamic echo")
8318 .handler(|input: AddInput| async move {
8319 Ok(CallToolResult::text(format!("{}", input.a)))
8320 })
8321 .build();
8322
8323 registry.register(tool);
8324
8325 let mut router = router;
8326 init_router(&mut router).await;
8327
8328 let req = RouterRequest {
8329 id: RequestId::Number(1),
8330 inner: McpRequest::ListTools(ListToolsParams::default()),
8331 extensions: Extensions::new(),
8332 };
8333
8334 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8335 match resp.inner {
8336 Ok(McpResponse::ListTools(result)) => {
8337 assert_eq!(result.tools.len(), 1);
8338 assert_eq!(result.tools[0].name, "dynamic_echo");
8339 }
8340 _ => panic!("Expected ListTools response"),
8341 }
8342 }
8343
8344 #[tokio::test]
8345 async fn test_dynamic_tools_unregister() {
8346 let (router, registry) = McpRouter::new()
8347 .server_info("test", "1.0")
8348 .with_dynamic_tools();
8349
8350 let tool = ToolBuilder::new("temp")
8351 .description("Temporary")
8352 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8353 .build();
8354
8355 registry.register(tool);
8356 assert!(registry.contains("temp"));
8357
8358 let removed = registry.unregister("temp");
8359 assert!(removed);
8360 assert!(!registry.contains("temp"));
8361
8362 assert!(!registry.unregister("temp"));
8364
8365 let mut router = router;
8366 init_router(&mut router).await;
8367
8368 let req = RouterRequest {
8369 id: RequestId::Number(1),
8370 inner: McpRequest::ListTools(ListToolsParams::default()),
8371 extensions: Extensions::new(),
8372 };
8373
8374 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8375 match resp.inner {
8376 Ok(McpResponse::ListTools(result)) => {
8377 assert_eq!(result.tools.len(), 0);
8378 }
8379 _ => panic!("Expected ListTools response"),
8380 }
8381 }
8382
8383 #[tokio::test]
8384 async fn test_dynamic_tools_merged_with_static() {
8385 let static_tool = ToolBuilder::new("static_tool")
8386 .description("Static")
8387 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8388 .build();
8389
8390 let (router, registry) = McpRouter::new()
8391 .server_info("test", "1.0")
8392 .tool(static_tool)
8393 .with_dynamic_tools();
8394
8395 let dynamic_tool = ToolBuilder::new("dynamic_tool")
8396 .description("Dynamic")
8397 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8398 .build();
8399
8400 registry.register(dynamic_tool);
8401
8402 let mut router = router;
8403 init_router(&mut router).await;
8404
8405 let req = RouterRequest {
8406 id: RequestId::Number(1),
8407 inner: McpRequest::ListTools(ListToolsParams::default()),
8408 extensions: Extensions::new(),
8409 };
8410
8411 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8412 match resp.inner {
8413 Ok(McpResponse::ListTools(result)) => {
8414 assert_eq!(result.tools.len(), 2);
8415 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8416 assert!(names.contains(&"static_tool"));
8417 assert!(names.contains(&"dynamic_tool"));
8418 }
8419 _ => panic!("Expected ListTools response"),
8420 }
8421 }
8422
8423 #[tokio::test]
8424 async fn test_static_tools_shadow_dynamic() {
8425 let static_tool = ToolBuilder::new("shared")
8426 .description("Static version")
8427 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8428 .build();
8429
8430 let (router, registry) = McpRouter::new()
8431 .server_info("test", "1.0")
8432 .tool(static_tool)
8433 .with_dynamic_tools();
8434
8435 let dynamic_tool = ToolBuilder::new("shared")
8436 .description("Dynamic version")
8437 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8438 .build();
8439
8440 registry.register(dynamic_tool);
8441
8442 let mut router = router;
8443 init_router(&mut router).await;
8444
8445 let req = RouterRequest {
8447 id: RequestId::Number(1),
8448 inner: McpRequest::ListTools(ListToolsParams::default()),
8449 extensions: Extensions::new(),
8450 };
8451
8452 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8453 match resp.inner {
8454 Ok(McpResponse::ListTools(result)) => {
8455 assert_eq!(result.tools.len(), 1);
8456 assert_eq!(result.tools[0].name, "shared");
8457 assert_eq!(
8458 result.tools[0].description.as_deref(),
8459 Some("Static version")
8460 );
8461 }
8462 _ => panic!("Expected ListTools response"),
8463 }
8464
8465 let req = RouterRequest {
8467 id: RequestId::Number(2),
8468 inner: McpRequest::CallTool(CallToolParams {
8469 input_responses: None,
8470 request_state: None,
8471 name: "shared".to_string(),
8472 arguments: serde_json::json!({"a": 1, "b": 2}),
8473 meta: None,
8474 task: None,
8475 }),
8476 extensions: Extensions::new(),
8477 };
8478
8479 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8480 match resp.inner {
8481 Ok(McpResponse::CallTool(result)) => {
8482 assert!(!result.is_error);
8483 match &result.content[0] {
8484 Content::Text { text, .. } => assert_eq!(text, "static"),
8485 _ => panic!("Expected text content"),
8486 }
8487 }
8488 _ => panic!("Expected CallTool response"),
8489 }
8490 }
8491
8492 #[tokio::test]
8493 async fn test_dynamic_tools_call() {
8494 let (router, registry) = McpRouter::new()
8495 .server_info("test", "1.0")
8496 .with_dynamic_tools();
8497
8498 let tool = ToolBuilder::new("add")
8499 .description("Add two numbers")
8500 .handler(|input: AddInput| async move {
8501 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
8502 })
8503 .build();
8504
8505 registry.register(tool);
8506
8507 let mut router = router;
8508 init_router(&mut router).await;
8509
8510 let req = RouterRequest {
8511 id: RequestId::Number(1),
8512 inner: McpRequest::CallTool(CallToolParams {
8513 input_responses: None,
8514 request_state: None,
8515 name: "add".to_string(),
8516 arguments: serde_json::json!({"a": 3, "b": 4}),
8517 meta: None,
8518 task: None,
8519 }),
8520 extensions: Extensions::new(),
8521 };
8522
8523 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8524 match resp.inner {
8525 Ok(McpResponse::CallTool(result)) => {
8526 assert!(!result.is_error);
8527 match &result.content[0] {
8528 Content::Text { text, .. } => assert_eq!(text, "7"),
8529 _ => panic!("Expected text content"),
8530 }
8531 }
8532 _ => panic!("Expected CallTool response"),
8533 }
8534 }
8535
8536 #[tokio::test]
8537 async fn test_dynamic_tools_notification_on_register() {
8538 let (tx, mut rx) = crate::context::notification_channel(16);
8539 let (router, registry) = McpRouter::new()
8540 .server_info("test", "1.0")
8541 .with_dynamic_tools();
8542 let _router = router.with_notification_sender(tx);
8543
8544 let tool = ToolBuilder::new("notified")
8545 .description("Test")
8546 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8547 .build();
8548
8549 registry.register(tool);
8550
8551 let notification = rx.recv().await.unwrap();
8552 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8553 }
8554
8555 #[tokio::test]
8556 async fn test_dynamic_tools_notification_on_unregister() {
8557 let (tx, mut rx) = crate::context::notification_channel(16);
8558 let (router, registry) = McpRouter::new()
8559 .server_info("test", "1.0")
8560 .with_dynamic_tools();
8561 let _router = router.with_notification_sender(tx);
8562
8563 let tool = ToolBuilder::new("notified")
8564 .description("Test")
8565 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8566 .build();
8567
8568 registry.register(tool);
8569 let _ = rx.recv().await.unwrap();
8571
8572 registry.unregister("notified");
8573 let notification = rx.recv().await.unwrap();
8574 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8575 }
8576
8577 #[tokio::test]
8578 async fn test_dynamic_tools_no_notification_on_empty_unregister() {
8579 let (tx, mut rx) = crate::context::notification_channel(16);
8580 let (router, registry) = McpRouter::new()
8581 .server_info("test", "1.0")
8582 .with_dynamic_tools();
8583 let _router = router.with_notification_sender(tx);
8584
8585 assert!(!registry.unregister("nonexistent"));
8587
8588 assert!(rx.try_recv().is_err());
8590 }
8591
8592 #[tokio::test]
8593 async fn test_dynamic_tools_filter_applies() {
8594 use crate::filter::CapabilityFilter;
8595
8596 let (router, registry) = McpRouter::new()
8597 .server_info("test", "1.0")
8598 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| {
8599 tool.name != "hidden"
8600 }))
8601 .with_dynamic_tools();
8602
8603 let visible = ToolBuilder::new("visible")
8604 .description("Visible")
8605 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8606 .build();
8607
8608 let hidden = ToolBuilder::new("hidden")
8609 .description("Hidden")
8610 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8611 .build();
8612
8613 registry.register(visible);
8614 registry.register(hidden);
8615
8616 let mut router = router;
8617 init_router(&mut router).await;
8618
8619 let req = RouterRequest {
8621 id: RequestId::Number(1),
8622 inner: McpRequest::ListTools(ListToolsParams::default()),
8623 extensions: Extensions::new(),
8624 };
8625
8626 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8627 match resp.inner {
8628 Ok(McpResponse::ListTools(result)) => {
8629 assert_eq!(result.tools.len(), 1);
8630 assert_eq!(result.tools[0].name, "visible");
8631 }
8632 _ => panic!("Expected ListTools response"),
8633 }
8634
8635 let req = RouterRequest {
8637 id: RequestId::Number(2),
8638 inner: McpRequest::CallTool(CallToolParams {
8639 input_responses: None,
8640 request_state: None,
8641 name: "hidden".to_string(),
8642 arguments: serde_json::json!({"a": 1, "b": 2}),
8643 meta: None,
8644 task: None,
8645 }),
8646 extensions: Extensions::new(),
8647 };
8648
8649 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8650 match resp.inner {
8651 Err(e) => {
8652 assert_eq!(e.code, -32601); }
8654 _ => panic!("Expected JsonRpc error"),
8655 }
8656 }
8657
8658 #[tokio::test]
8659 async fn test_dynamic_tools_capabilities_advertised() {
8660 let (mut router, _registry) = McpRouter::new()
8662 .server_info("test", "1.0")
8663 .with_dynamic_tools();
8664
8665 let init_req = RouterRequest {
8666 id: RequestId::Number(1),
8667 inner: McpRequest::Initialize(InitializeParams {
8668 protocol_version: "2025-11-25".to_string(),
8669 capabilities: ClientCapabilities::default(),
8670 client_info: Implementation {
8671 name: "test".to_string(),
8672 version: "1.0".to_string(),
8673 ..Default::default()
8674 },
8675 meta: None,
8676 }),
8677 extensions: Extensions::new(),
8678 };
8679
8680 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8681 match resp.inner {
8682 Ok(McpResponse::Initialize(result)) => {
8683 assert!(result.capabilities.tools.is_some());
8684 }
8685 _ => panic!("Expected Initialize response"),
8686 }
8687 }
8688
8689 #[tokio::test]
8690 async fn test_dynamic_tools_multi_session_notification() {
8691 let (tx1, mut rx1) = crate::context::notification_channel(16);
8692 let (tx2, mut rx2) = crate::context::notification_channel(16);
8693
8694 let (router, registry) = McpRouter::new()
8695 .server_info("test", "1.0")
8696 .with_dynamic_tools();
8697
8698 let _session1 = router.clone().with_notification_sender(tx1);
8700 let _session2 = router.clone().with_notification_sender(tx2);
8701
8702 let tool = ToolBuilder::new("broadcast")
8703 .description("Test")
8704 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8705 .build();
8706
8707 registry.register(tool);
8708
8709 let n1 = rx1.recv().await.unwrap();
8711 let n2 = rx2.recv().await.unwrap();
8712 assert!(matches!(n1, ServerNotification::ToolsListChanged));
8713 assert!(matches!(n2, ServerNotification::ToolsListChanged));
8714 }
8715
8716 #[tokio::test]
8717 async fn test_dynamic_tools_call_not_found() {
8718 let (router, _registry) = McpRouter::new()
8719 .server_info("test", "1.0")
8720 .with_dynamic_tools();
8721
8722 let mut router = router;
8723 init_router(&mut router).await;
8724
8725 let req = RouterRequest {
8726 id: RequestId::Number(1),
8727 inner: McpRequest::CallTool(CallToolParams {
8728 input_responses: None,
8729 request_state: None,
8730 name: "nonexistent".to_string(),
8731 arguments: serde_json::json!({}),
8732 meta: None,
8733 task: None,
8734 }),
8735 extensions: Extensions::new(),
8736 };
8737
8738 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8739 match resp.inner {
8740 Err(e) => {
8741 assert_eq!(e.code, -32601);
8742 }
8743 _ => panic!("Expected method not found error"),
8744 }
8745 }
8746
8747 #[tokio::test]
8748 async fn test_dynamic_tools_registry_list() {
8749 let (_, registry) = McpRouter::new()
8750 .server_info("test", "1.0")
8751 .with_dynamic_tools();
8752
8753 assert!(registry.list().is_empty());
8754
8755 let tool = ToolBuilder::new("tool_a")
8756 .description("A")
8757 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8758 .build();
8759 registry.register(tool);
8760
8761 let tool = ToolBuilder::new("tool_b")
8762 .description("B")
8763 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8764 .build();
8765 registry.register(tool);
8766
8767 let tools = registry.list();
8768 assert_eq!(tools.len(), 2);
8769 let names: Vec<&str> = tools.iter().map(|t| t.name.as_str()).collect();
8770 assert!(names.contains(&"tool_a"));
8771 assert!(names.contains(&"tool_b"));
8772 }
8773 } #[tokio::test]
8776 async fn test_tool_if_true_registers() {
8777 let tool = ToolBuilder::new("conditional")
8778 .description("Conditional tool")
8779 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8780 .build();
8781
8782 let mut router = McpRouter::new().tool_if(true, tool);
8783 init_router(&mut router).await;
8784
8785 let req = RouterRequest {
8786 id: RequestId::Number(1),
8787 inner: McpRequest::ListTools(ListToolsParams::default()),
8788 extensions: Extensions::new(),
8789 };
8790 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8791 match resp.inner {
8792 Ok(McpResponse::ListTools(result)) => {
8793 assert_eq!(result.tools.len(), 1);
8794 assert_eq!(result.tools[0].name, "conditional");
8795 }
8796 _ => panic!("Expected ListTools response"),
8797 }
8798 }
8799
8800 #[tokio::test]
8801 async fn test_tool_if_false_skips() {
8802 let tool = ToolBuilder::new("conditional")
8803 .description("Conditional tool")
8804 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8805 .build();
8806
8807 let mut router = McpRouter::new().tool_if(false, tool);
8808 init_router(&mut router).await;
8809
8810 let req = RouterRequest {
8811 id: RequestId::Number(1),
8812 inner: McpRequest::ListTools(ListToolsParams::default()),
8813 extensions: Extensions::new(),
8814 };
8815 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8816 match resp.inner {
8817 Ok(McpResponse::ListTools(result)) => {
8818 assert_eq!(result.tools.len(), 0);
8819 }
8820 _ => panic!("Expected ListTools response"),
8821 }
8822 }
8823
8824 #[tokio::test]
8825 async fn test_tools_if_batch_conditional() {
8826 let tools = vec![
8827 ToolBuilder::new("a")
8828 .description("Tool A")
8829 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8830 .build(),
8831 ToolBuilder::new("b")
8832 .description("Tool B")
8833 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8834 .build(),
8835 ];
8836
8837 let mut router = McpRouter::new().tools_if(false, tools);
8838 init_router(&mut router).await;
8839
8840 let req = RouterRequest {
8841 id: RequestId::Number(1),
8842 inner: McpRequest::ListTools(ListToolsParams::default()),
8843 extensions: Extensions::new(),
8844 };
8845 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8846 match resp.inner {
8847 Ok(McpResponse::ListTools(result)) => {
8848 assert_eq!(result.tools.len(), 0);
8849 }
8850 _ => panic!("Expected ListTools response"),
8851 }
8852 }
8853
8854 #[test]
8855 fn test_resource_if_true_registers() {
8856 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8857 .name("test")
8858 .text("hello");
8859
8860 let router = McpRouter::new().resource_if(true, resource);
8861 assert_eq!(router.inner.resources.len(), 1);
8862 }
8863
8864 #[test]
8865 fn test_resource_if_false_skips() {
8866 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8867 .name("test")
8868 .text("hello");
8869
8870 let router = McpRouter::new().resource_if(false, resource);
8871 assert_eq!(router.inner.resources.len(), 0);
8872 }
8873
8874 #[test]
8875 fn test_prompt_if_true_registers() {
8876 let prompt = crate::prompt::PromptBuilder::new("greet")
8877 .description("Greeting")
8878 .user_message("Hello!");
8879
8880 let router = McpRouter::new().prompt_if(true, prompt);
8881 assert_eq!(router.inner.prompts.len(), 1);
8882 }
8883
8884 #[test]
8885 fn test_prompt_if_false_skips() {
8886 let prompt = crate::prompt::PromptBuilder::new("greet")
8887 .description("Greeting")
8888 .user_message("Hello!");
8889
8890 let router = McpRouter::new().prompt_if(false, prompt);
8891 assert_eq!(router.inner.prompts.len(), 0);
8892 }
8893
8894 #[tokio::test]
8895 async fn test_disable_tool_hides_from_list() {
8896 let safe = ToolBuilder::new("safe")
8897 .description("Safe tool")
8898 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8899 .build();
8900 let dangerous = ToolBuilder::new("dangerous")
8901 .description("Dangerous tool")
8902 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8903 .build();
8904 let mut router = McpRouter::new().tool(safe).tool(dangerous);
8905 init_router(&mut router).await;
8906
8907 router.disable_tool("dangerous");
8908 assert!(router.is_tool_enabled("safe"));
8909 assert!(!router.is_tool_enabled("dangerous"));
8910
8911 let req = RouterRequest {
8912 id: RequestId::Number(1),
8913 inner: McpRequest::ListTools(ListToolsParams::default()),
8914 extensions: Extensions::new(),
8915 };
8916 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8917 match resp.inner {
8918 Ok(McpResponse::ListTools(result)) => {
8919 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8920 assert_eq!(names, vec!["safe"]);
8921 }
8922 _ => panic!("Expected ListTools response"),
8923 }
8924 }
8925
8926 #[tokio::test]
8927 async fn test_disable_tool_blocks_call() {
8928 let dangerous = ToolBuilder::new("dangerous")
8929 .description("Dangerous tool")
8930 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
8931 .build();
8932 let mut router = McpRouter::new().tool(dangerous);
8933 init_router(&mut router).await;
8934
8935 router.disable_tool("dangerous");
8936
8937 let req = RouterRequest {
8938 id: RequestId::Number(2),
8939 inner: McpRequest::CallTool(CallToolParams {
8940 input_responses: None,
8941 request_state: None,
8942 name: "dangerous".to_string(),
8943 arguments: serde_json::json!({"a": 1, "b": 2}),
8944 meta: None,
8945 task: None,
8946 }),
8947 extensions: Extensions::new(),
8948 };
8949 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8950 let err = resp.inner.expect_err("disabled tool should error");
8951 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
8952 }
8953
8954 #[tokio::test]
8955 async fn test_enable_tool_restores_visibility() {
8956 let tool = ToolBuilder::new("flippy")
8957 .description("Toggleable tool")
8958 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
8959 .build();
8960 let mut router = McpRouter::new().tool(tool);
8961 init_router(&mut router).await;
8962
8963 router.disable_tool("flippy");
8964 router.enable_tool("flippy");
8965 assert!(router.is_tool_enabled("flippy"));
8966
8967 let req = RouterRequest {
8968 id: RequestId::Number(3),
8969 inner: McpRequest::CallTool(CallToolParams {
8970 input_responses: None,
8971 request_state: None,
8972 name: "flippy".to_string(),
8973 arguments: serde_json::json!({"a": 1, "b": 2}),
8974 meta: None,
8975 task: None,
8976 }),
8977 extensions: Extensions::new(),
8978 };
8979 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8980 match resp.inner {
8981 Ok(McpResponse::CallTool(result)) => {
8982 assert_eq!(result.first_text(), Some("ran"));
8983 }
8984 _ => panic!("Expected CallTool response"),
8985 }
8986 }
8987
8988 #[tokio::test]
8989 async fn test_disable_propagates_through_fresh_session() {
8990 let tool = ToolBuilder::new("shared")
8991 .description("Shared across sessions")
8992 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8993 .build();
8994 let router = McpRouter::new().tool(tool);
8995
8996 router.disable_tool("shared");
8998 let mut child = router.with_fresh_session();
8999 init_router(&mut child).await;
9000 assert!(!child.is_tool_enabled("shared"));
9001
9002 let req = RouterRequest {
9003 id: RequestId::Number(4),
9004 inner: McpRequest::ListTools(ListToolsParams::default()),
9005 extensions: Extensions::new(),
9006 };
9007 let resp = child.ready().await.unwrap().call(req).await.unwrap();
9008 match resp.inner {
9009 Ok(McpResponse::ListTools(result)) => {
9010 assert!(result.tools.is_empty());
9011 }
9012 _ => panic!("Expected ListTools response"),
9013 }
9014 }
9015
9016 #[tokio::test]
9017 async fn test_disable_resource_and_prompt() {
9018 let resource = crate::resource::ResourceBuilder::new("file:///hidden.txt")
9019 .name("hidden")
9020 .text("secret");
9021 let prompt = crate::prompt::PromptBuilder::new("hidden_prompt")
9022 .description("hidden")
9023 .user_message("hello");
9024
9025 let mut router = McpRouter::new().resource(resource).prompt(prompt);
9026 init_router(&mut router).await;
9027
9028 router.disable_resource("file:///hidden.txt");
9029 router.disable_prompt("hidden_prompt");
9030 assert!(!router.is_resource_enabled("file:///hidden.txt"));
9031 assert!(!router.is_prompt_enabled("hidden_prompt"));
9032
9033 let req = RouterRequest {
9035 id: RequestId::Number(5),
9036 inner: McpRequest::ListResources(ListResourcesParams::default()),
9037 extensions: Extensions::new(),
9038 };
9039 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9040 match resp.inner {
9041 Ok(McpResponse::ListResources(result)) => {
9042 assert!(result.resources.is_empty());
9043 }
9044 _ => panic!("Expected ListResources response"),
9045 }
9046
9047 let req = RouterRequest {
9049 id: RequestId::Number(6),
9050 inner: McpRequest::ReadResource(ReadResourceParams {
9051 input_responses: None,
9052 request_state: None,
9053 uri: "file:///hidden.txt".to_string(),
9054 meta: None,
9055 }),
9056 extensions: Extensions::new(),
9057 };
9058 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9059 let err = resp.inner.expect_err("disabled resource should error");
9060 assert_eq!(err.code, -32602); let req = RouterRequest {
9064 id: RequestId::Number(7),
9065 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
9066 extensions: Extensions::new(),
9067 };
9068 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9069 match resp.inner {
9070 Ok(McpResponse::ListPrompts(result)) => {
9071 assert!(result.prompts.is_empty());
9072 }
9073 _ => panic!("Expected ListPrompts response"),
9074 }
9075
9076 let req = RouterRequest {
9078 id: RequestId::Number(8),
9079 inner: McpRequest::GetPrompt(GetPromptParams {
9080 input_responses: None,
9081 request_state: None,
9082 name: "hidden_prompt".to_string(),
9083 arguments: Default::default(),
9084 meta: None,
9085 }),
9086 extensions: Extensions::new(),
9087 };
9088 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9089 let err = resp.inner.expect_err("disabled prompt should error");
9090 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
9091 }
9092
9093 #[test]
9094 fn test_router_request_new() {
9095 let req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9096 assert_eq!(req.id, RequestId::Number(1));
9097 assert!(req.extensions.is_empty());
9098 }
9099
9100 #[test]
9101 fn test_with_inner_preserves_extensions() {
9102 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9103 req.extensions.insert(42u32);
9104
9105 let rewritten = req.with_inner(McpRequest::ListTools(Default::default()));
9106 assert!(matches!(rewritten.inner, McpRequest::ListTools(_)));
9107 assert_eq!(rewritten.id, RequestId::Number(1));
9108 assert_eq!(rewritten.extensions.get::<u32>(), Some(&42));
9109 }
9110
9111 #[test]
9112 fn test_with_id_and_inner_preserves_extensions() {
9113 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9114 req.extensions.insert(String::from("token-abc"));
9115
9116 let rewritten = req.with_id_and_inner(
9117 RequestId::Number(99),
9118 McpRequest::ListResources(Default::default()),
9119 );
9120 assert_eq!(rewritten.id, RequestId::Number(99));
9121 assert!(matches!(rewritten.inner, McpRequest::ListResources(_)));
9122 assert_eq!(
9123 rewritten.extensions.get::<String>(),
9124 Some(&String::from("token-abc"))
9125 );
9126 }
9127
9128 #[test]
9129 fn test_clone_with_inner_preserves_extensions() {
9130 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9131 req.extensions.insert(true);
9132
9133 let cloned = req.clone_with_inner(McpRequest::ListTools(Default::default()));
9134
9135 assert!(matches!(req.inner, McpRequest::Ping));
9137 assert_eq!(req.extensions.get::<bool>(), Some(&true));
9138
9139 assert!(matches!(cloned.inner, McpRequest::ListTools(_)));
9141 assert_eq!(cloned.extensions.get::<bool>(), Some(&true));
9142 }
9143
9144 #[test]
9145 fn test_router_response_is_error() {
9146 let ok_resp = RouterResponse {
9147 id: RequestId::Number(1),
9148 inner: Ok(McpResponse::Pong(Default::default())),
9149 };
9150 assert!(!ok_resp.is_error());
9151
9152 let err_resp = RouterResponse {
9153 id: RequestId::Number(2),
9154 inner: Err(JsonRpcError::internal_error("boom")),
9155 };
9156 assert!(err_resp.is_error());
9157 }
9158
9159 #[test]
9160 fn test_extensions_len_and_is_empty() {
9161 let mut ext = Extensions::new();
9162 assert!(ext.is_empty());
9163 assert_eq!(ext.len(), 0);
9164
9165 ext.insert(42u32);
9166 assert!(!ext.is_empty());
9167 assert_eq!(ext.len(), 1);
9168
9169 ext.insert(String::from("hello"));
9170 assert_eq!(ext.len(), 2);
9171 }
9172
9173 #[test]
9174 fn test_router_response_serde_roundtrip() {
9175 let response = RouterResponse {
9177 id: RequestId::Number(1),
9178 inner: Ok(McpResponse::Empty(EmptyResult {})),
9179 };
9180 let json = serde_json::to_string(&response).unwrap();
9181 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9182 assert_eq!(deserialized.id, RequestId::Number(1));
9183 assert!(!deserialized.is_error());
9184
9185 let response = RouterResponse {
9187 id: RequestId::String("req-2".into()),
9188 inner: Err(JsonRpcError::method_not_found("unknown")),
9189 };
9190 let json = serde_json::to_string(&response).unwrap();
9191 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9192 assert_eq!(deserialized.id, RequestId::String("req-2".into()));
9193 assert!(deserialized.is_error());
9194 }
9195
9196 #[tokio::test]
9203 async fn test_discover_dispatch_via_jsonrpc_service() {
9204 let router = McpRouter::new().server_info("unit-test-server", "4.2.0");
9207 let mut service = JsonRpcService::new(router);
9208
9209 let req = JsonRpcRequest::new(1, "server/discover");
9210 let resp = service.call_single(req).await.unwrap();
9211
9212 match resp {
9213 JsonRpcResponse::Result(r) => {
9214 let versions = r
9216 .result
9217 .get("supportedVersions")
9218 .and_then(|v| v.as_array())
9219 .expect("result.supportedVersions must be an array");
9220 assert!(!versions.is_empty(), "supportedVersions must not be empty");
9221
9222 assert_eq!(
9224 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
9225 "unit-test-server",
9226 "serverInfo.name must match configured value"
9227 );
9228 assert_eq!(
9229 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["version"], "4.2.0",
9230 "serverInfo.version must match configured value"
9231 );
9232
9233 assert!(
9236 r.result.get("protocolVersion").is_none(),
9237 "server/discover must NOT include protocolVersion: {:?}",
9238 r.result
9239 );
9240 }
9241 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
9242 _ => panic!("unexpected response variant"),
9243 }
9244 }
9245
9246 #[tokio::test]
9247 async fn test_discover_does_not_require_initialization() {
9248 let router = McpRouter::new().server_info("fresh-router", "1.0.0");
9251 let mut service = JsonRpcService::new(router);
9252
9253 let req = JsonRpcRequest::new(2, "server/discover");
9254 let resp = service.call_single(req).await.unwrap();
9255
9256 assert!(
9258 !matches!(resp, JsonRpcResponse::Error(_)),
9259 "server/discover must not require initialization: {:?}",
9260 resp
9261 );
9262 }
9263}
9264
9265#[cfg(test)]
9266mod cursor_property_tests {
9267 use super::{decode_cursor, encode_cursor};
9268 use proptest::prelude::*;
9269
9270 fn arb_cursor_text() -> BoxedStrategy<String> {
9271 prop_oneof![
9272 8 => prop::collection::vec(any::<char>(), 0..512)
9273 .prop_map(|chars| chars.into_iter().collect()),
9274 1 => Just("\0\r\n\t\u{001b}\u{007f}".repeat(64)),
9275 1 => Just("A".repeat(16 * 1024)),
9276 ]
9277 .boxed()
9278 }
9279
9280 proptest! {
9281 #![proptest_config(ProptestConfig::with_cases(512))]
9282
9283 #[test]
9285 fn cursor_round_trips(offset in any::<usize>()) {
9286 prop_assert_eq!(decode_cursor(&encode_cursor(offset)).unwrap(), offset);
9287 }
9288
9289 #[test]
9291 fn decode_cursor_never_panics(s in arb_cursor_text()) {
9292 let _ = decode_cursor(&s);
9293 }
9294 }
9295}