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 client_requester: Option<ClientRequesterHandle>,
445 task_store: Arc<dyn TaskStore>,
447 subscriptions: Arc<RwLock<HashSet<String>>>,
449 completion_handler: Option<CompletionHandler>,
451 tool_filter: Option<ToolFilter>,
453 resource_filter: Option<ResourceFilter>,
455 prompt_filter: Option<PromptFilter>,
457 extensions: Arc<crate::context::Extensions>,
459 protocol_extensions: HashMap<String, serde_json::Value>,
461 min_log_level: Arc<RwLock<LogLevel>>,
463 page_size: Option<usize>,
465 list_ttl_ms: Option<u64>,
469 read_ttl_ms: Option<u64>,
473 cache_scope: Option<CacheScope>,
478 logging_deprecated: Option<tower_mcp_types::protocol::DeprecationInfo>,
481 disabled_tools: Arc<RwLock<HashSet<String>>>,
483 disabled_resources: Arc<RwLock<HashSet<String>>>,
485 disabled_prompts: Arc<RwLock<HashSet<String>>>,
487 #[cfg(feature = "dynamic-tools")]
489 dynamic_tools: Option<Arc<DynamicToolsInner>>,
490 #[cfg(feature = "dynamic-tools")]
492 dynamic_prompts: Option<Arc<DynamicPromptsInner>>,
493 #[cfg(feature = "dynamic-tools")]
495 prompt_initializer: Option<PromptInitializer>,
496 #[cfg(feature = "dynamic-tools")]
498 dynamic_resources: Option<Arc<DynamicResourcesInner>>,
499 #[cfg(feature = "dynamic-tools")]
501 dynamic_resource_templates: Option<Arc<DynamicResourceTemplatesInner>>,
502}
503
504impl McpRouterInner {
505 fn generate_instructions(&self, config: &AutoInstructionsConfig) -> String {
507 let mut parts = Vec::new();
508
509 if let Some(prefix) = &config.prefix {
510 parts.push(prefix.clone());
511 }
512
513 if !self.tools.is_empty() {
515 let mut lines = vec!["## Tools".to_string(), String::new()];
516 let mut tools: Vec<_> = self.tools.values().collect();
517 tools.sort_by(|a, b| a.name.cmp(&b.name));
518 for tool in tools {
519 let desc = tool.description.as_deref().unwrap_or("No description");
520 let tags = annotation_tags(tool.annotations.as_ref());
521 if tags.is_empty() {
522 lines.push(format!("- **{}**: {}", tool.name, desc));
523 } else {
524 lines.push(format!("- **{}**: {} [{}]", tool.name, desc, tags));
525 }
526 }
527 parts.push(lines.join("\n"));
528 }
529
530 if !self.resources.is_empty() || !self.resource_templates.is_empty() {
532 let mut lines = vec!["## Resources".to_string(), String::new()];
533 let mut resources: Vec<_> = self.resources.values().collect();
534 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
535 for resource in resources {
536 let desc = resource.description.as_deref().unwrap_or("No description");
537 lines.push(format!("- **{}**: {}", resource.uri, desc));
538 }
539 let mut templates: Vec<_> = self.resource_templates.iter().collect();
540 templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
541 for template in templates {
542 let desc = template.description.as_deref().unwrap_or("No description");
543 lines.push(format!("- **{}**: {}", template.uri_template, desc));
544 }
545 parts.push(lines.join("\n"));
546 }
547
548 if !self.prompts.is_empty() {
550 let mut lines = vec!["## Prompts".to_string(), String::new()];
551 let mut prompts: Vec<_> = self.prompts.values().collect();
552 prompts.sort_by(|a, b| a.name.cmp(&b.name));
553 for prompt in prompts {
554 let desc = prompt.description.as_deref().unwrap_or("No description");
555 lines.push(format!("- **{}**: {}", prompt.name, desc));
556 }
557 parts.push(lines.join("\n"));
558 }
559
560 if let Some(suffix) = &config.suffix {
561 parts.push(suffix.clone());
562 }
563
564 parts.join("\n\n")
565 }
566}
567
568fn annotation_tags(annotations: Option<&crate::protocol::ToolAnnotations>) -> String {
574 let Some(ann) = annotations else {
575 return String::new();
576 };
577 let mut tags = Vec::new();
578 if ann.is_read_only() {
579 tags.push("read-only");
580 }
581 if ann.is_idempotent() {
582 tags.push("idempotent");
583 }
584 tags.join(", ")
585}
586
587impl McpRouter {
588 pub fn new() -> Self {
590 Self {
591 inner: Arc::new(McpRouterInner {
592 server_name: "tower-mcp".to_string(),
593 server_version: env!("CARGO_PKG_VERSION").to_string(),
594 server_title: None,
595 server_description: None,
596 server_icons: None,
597 server_website_url: None,
598 instructions: None,
599 auto_instructions: None,
600 tools: HashMap::new(),
601 resources: HashMap::new(),
602 resource_templates: Vec::new(),
603 prompts: HashMap::new(),
604 in_flight: Arc::new(RwLock::new(HashMap::new())),
605 notification_tx: None,
606 #[cfg(all(feature = "http", feature = "stateless"))]
607 modern_notification_sink: Arc::new(RwLock::new(None)),
608 client_requester: None,
609 task_store: Arc::new(MemoryTaskStore::new()),
610 subscriptions: Arc::new(RwLock::new(HashSet::new())),
611 extensions: Arc::new(crate::context::Extensions::new()),
612 protocol_extensions: HashMap::new(),
613 completion_handler: None,
614 tool_filter: None,
615 resource_filter: None,
616 prompt_filter: None,
617 min_log_level: Arc::new(RwLock::new(LogLevel::Debug)),
618 page_size: None,
619 list_ttl_ms: None,
620 read_ttl_ms: None,
621 cache_scope: None,
622 logging_deprecated: None,
623 disabled_tools: Arc::new(RwLock::new(HashSet::new())),
624 disabled_resources: Arc::new(RwLock::new(HashSet::new())),
625 disabled_prompts: Arc::new(RwLock::new(HashSet::new())),
626 #[cfg(feature = "dynamic-tools")]
627 dynamic_tools: None,
628 #[cfg(feature = "dynamic-tools")]
629 dynamic_prompts: None,
630 #[cfg(feature = "dynamic-tools")]
631 prompt_initializer: None,
632 #[cfg(feature = "dynamic-tools")]
633 dynamic_resources: None,
634 #[cfg(feature = "dynamic-tools")]
635 dynamic_resource_templates: None,
636 }),
637 session: SessionState::new(),
638 }
639 }
640
641 pub fn with_fresh_session(&self) -> Self {
649 Self {
650 inner: self.inner.clone(),
651 session: SessionState::new(),
652 }
653 }
654
655 pub fn tool_annotations_map(&self) -> ToolAnnotationsMap {
665 let disabled = self.inner.disabled_tools.read().unwrap();
666 let mut map = HashMap::new();
667 for (name, tool) in &self.inner.tools {
668 if disabled.contains(name) {
669 continue;
670 }
671 if let Some(annotations) = &tool.annotations {
672 map.insert(name.clone(), annotations.clone());
673 }
674 }
675 #[cfg(feature = "dynamic-tools")]
676 if let Some(dynamic) = &self.inner.dynamic_tools {
677 for tool in dynamic.list() {
678 if disabled.contains(&tool.name) {
679 continue;
680 }
681 if !map.contains_key(&tool.name)
683 && let Some(ref annotations) = tool.annotations
684 {
685 map.insert(tool.name.clone(), annotations.clone());
686 }
687 }
688 }
689 ToolAnnotationsMap { map: Arc::new(map) }
690 }
691
692 pub fn task_store(mut self, store: Arc<dyn TaskStore>) -> Self {
710 Arc::make_mut(&mut self.inner).task_store = store;
711 self
712 }
713
714 #[cfg(feature = "dynamic-tools")]
744 pub fn with_dynamic_tools(mut self) -> (Self, DynamicToolRegistry) {
745 let inner_dyn = Arc::new(DynamicToolsInner::new());
746 Arc::make_mut(&mut self.inner).dynamic_tools = Some(inner_dyn.clone());
747 (self, DynamicToolRegistry::new(inner_dyn))
748 }
749
750 #[cfg(feature = "dynamic-tools")]
773 pub fn with_dynamic_prompts(mut self) -> (Self, DynamicPromptRegistry) {
774 let inner_dyn = Arc::new(DynamicPromptsInner::new());
775 Arc::make_mut(&mut self.inner).dynamic_prompts = Some(inner_dyn.clone());
776 (self, DynamicPromptRegistry::new(inner_dyn))
777 }
778
779 #[cfg(feature = "dynamic-tools")]
785 pub fn dynamic_prompt_initializer<F>(mut self, initializer: F) -> Self
786 where
787 F: Fn() -> Result<()> + Send + Sync + 'static,
788 {
789 Arc::make_mut(&mut self.inner).prompt_initializer = Some(Arc::new(initializer));
790 self
791 }
792
793 #[cfg(feature = "dynamic-tools")]
816 pub fn with_dynamic_resources(mut self) -> (Self, DynamicResourceRegistry) {
817 let inner_dyn = Arc::new(DynamicResourcesInner::new());
818 Arc::make_mut(&mut self.inner).dynamic_resources = Some(inner_dyn.clone());
819 (self, DynamicResourceRegistry::new(inner_dyn))
820 }
821
822 #[cfg(feature = "dynamic-tools")]
844 pub fn with_dynamic_resource_templates(mut self) -> (Self, DynamicResourceTemplateRegistry) {
845 let inner_dyn = Arc::new(DynamicResourceTemplatesInner::new());
846 Arc::make_mut(&mut self.inner).dynamic_resource_templates = Some(inner_dyn.clone());
847 (self, DynamicResourceTemplateRegistry::new(inner_dyn))
848 }
849
850 #[cfg(feature = "stateless")]
858 #[cfg(feature = "http")]
859 pub(crate) fn with_request_notification_sender(mut self, tx: NotificationSender) -> Self {
860 Arc::make_mut(&mut self.inner).notification_tx = Some(tx);
861 self
862 }
863
864 pub fn with_notification_sender(mut self, tx: NotificationSender) -> Self {
868 let inner = Arc::make_mut(&mut self.inner);
869 #[cfg(feature = "dynamic-tools")]
872 if let Some(ref dynamic_tools) = inner.dynamic_tools {
873 dynamic_tools.add_notification_sender(tx.clone());
874 }
875 #[cfg(feature = "dynamic-tools")]
876 if let Some(ref dynamic_prompts) = inner.dynamic_prompts {
877 dynamic_prompts.add_notification_sender(tx.clone());
878 }
879 #[cfg(feature = "dynamic-tools")]
880 if let Some(ref dynamic_resources) = inner.dynamic_resources {
881 dynamic_resources.add_notification_sender(tx.clone());
882 }
883 #[cfg(feature = "dynamic-tools")]
884 if let Some(ref dynamic_resource_templates) = inner.dynamic_resource_templates {
885 dynamic_resource_templates.add_notification_sender(tx.clone());
886 }
887 inner.notification_tx = Some(tx);
888 self
889 }
890
891 #[cfg(all(feature = "http", feature = "stateless"))]
893 pub(crate) fn attach_modern_notification_sink(&self, sink: ModernNotificationSink) {
894 if let Ok(mut active) = self.inner.modern_notification_sink.write() {
895 *active = Some(sink);
896 }
897 }
898
899 pub fn notification_sender(&self) -> Option<&NotificationSender> {
901 self.inner.notification_tx.as_ref()
902 }
903
904 pub fn with_client_requester(mut self, requester: ClientRequesterHandle) -> Self {
909 Arc::make_mut(&mut self.inner).client_requester = Some(requester);
910 self
911 }
912
913 pub fn client_requester(&self) -> Option<&ClientRequesterHandle> {
915 self.inner.client_requester.as_ref()
916 }
917
918 pub fn with_state<T: Clone + Send + Sync + 'static>(mut self, state: T) -> Self {
961 let inner = Arc::make_mut(&mut self.inner);
962 Arc::make_mut(&mut inner.extensions).insert(state);
963 self
964 }
965
966 pub fn with_extension<T: Clone + Send + Sync + 'static>(self, value: T) -> Self {
971 self.with_state(value)
972 }
973
974 pub fn with_protocol_extension(mut self, extension: crate::ExtensionDeclaration) -> Self {
981 let (identifier, settings) = extension.into_parts();
982 Arc::make_mut(&mut self.inner)
983 .protocol_extensions
984 .insert(identifier, settings);
985 self
986 }
987
988 pub fn extensions(&self) -> &crate::context::Extensions {
990 &self.inner.extensions
991 }
992
993 pub fn create_context(
998 &self,
999 request_id: RequestId,
1000 progress_token: Option<ProgressToken>,
1001 ) -> RequestContext {
1002 self.create_context_with_extensions(request_id, progress_token, &Extensions::new())
1003 }
1004
1005 pub(crate) fn create_context_with_extensions(
1010 &self,
1011 request_id: RequestId,
1012 progress_token: Option<ProgressToken>,
1013 per_request: &Extensions,
1014 ) -> RequestContext {
1015 let ctx = RequestContext::new(request_id.clone());
1016
1017 let ctx = if let Some(token) = progress_token {
1019 ctx.with_progress_token(token)
1020 } else {
1021 ctx
1022 };
1023
1024 let ctx = if let Some(tx) = &self.inner.notification_tx {
1026 ctx.with_notification_sender(tx.clone())
1027 } else {
1028 ctx
1029 };
1030
1031 let mut merged = (*self.inner.extensions).clone();
1035 merged.merge(per_request);
1036 let negotiated_extensions = if is_final_protocol_request(per_request) {
1037 let server_capabilities =
1038 self.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
1039 final_client_capabilities(per_request)
1040 .map(|client_capabilities| {
1041 crate::NegotiatedExtensions::from_capabilities(
1042 client_capabilities,
1043 &server_capabilities,
1044 )
1045 })
1046 .unwrap_or_default()
1047 } else {
1048 self.session
1049 .get::<crate::NegotiatedExtensions>()
1050 .unwrap_or_default()
1051 };
1052 merged.insert(negotiated_extensions);
1053
1054 let ctx = if !is_final_protocol_request(per_request)
1059 && let Some(requester) = merged
1060 .get::<ClientRequesterHandle>()
1061 .cloned()
1062 .or_else(|| self.inner.client_requester.clone())
1063 {
1064 ctx.with_client_requester(requester)
1065 } else {
1066 ctx
1067 };
1068
1069 let ctx = if let Some(token) = merged.get::<CancellationToken>() {
1073 ctx.with_cancellation_token(token.clone())
1074 } else {
1075 ctx
1076 };
1077
1078 let ctx = ctx.with_extensions(Arc::new(merged));
1079
1080 let ctx = ctx.with_min_log_level(self.inner.min_log_level.clone());
1082
1083 let token = ctx.cancellation_token();
1085 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1086 in_flight.insert(request_id, token);
1087 }
1088
1089 ctx
1090 }
1091
1092 pub fn complete_request(&self, request_id: &RequestId) {
1094 if let Ok(mut in_flight) = self.inner.in_flight.write() {
1095 in_flight.remove(request_id);
1096 }
1097 }
1098
1099 fn cancel_request(&self, request_id: &RequestId) -> bool {
1101 let Ok(in_flight) = self.inner.in_flight.read() else {
1102 return false;
1103 };
1104 let Some(token) = in_flight.get(request_id) else {
1105 return false;
1106 };
1107 token.cancel();
1108 true
1109 }
1110
1111 pub fn server_info(mut self, name: impl Into<String>, version: impl Into<String>) -> Self {
1113 let inner = Arc::make_mut(&mut self.inner);
1114 inner.server_name = name.into();
1115 inner.server_version = version.into();
1116 self
1117 }
1118
1119 pub fn page_size(mut self, size: usize) -> Self {
1126 Arc::make_mut(&mut self.inner).page_size = Some(size);
1127 self
1128 }
1129
1130 pub fn list_ttl(mut self, ms: u64) -> Self {
1136 Arc::make_mut(&mut self.inner).list_ttl_ms = Some(ms);
1137 self
1138 }
1139
1140 pub fn read_ttl(mut self, ms: u64) -> Self {
1147 Arc::make_mut(&mut self.inner).read_ttl_ms = Some(ms);
1148 self
1149 }
1150
1151 pub fn cache_scope(mut self, scope: CacheScope) -> Self {
1160 Arc::make_mut(&mut self.inner).cache_scope = Some(scope);
1161 self
1162 }
1163
1164 pub fn logging_deprecated(mut self, info: tower_mcp_types::protocol::DeprecationInfo) -> Self {
1170 Arc::make_mut(&mut self.inner).logging_deprecated = Some(info);
1171 self
1172 }
1173
1174 pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
1176 Arc::make_mut(&mut self.inner).instructions = Some(instructions.into());
1177 self
1178 }
1179
1180 pub fn auto_instructions(mut self) -> Self {
1212 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1213 prefix: None,
1214 suffix: None,
1215 });
1216 self
1217 }
1218
1219 pub fn auto_instructions_with(
1236 mut self,
1237 prefix: Option<impl Into<String>>,
1238 suffix: Option<impl Into<String>>,
1239 ) -> Self {
1240 Arc::make_mut(&mut self.inner).auto_instructions = Some(AutoInstructionsConfig {
1241 prefix: prefix.map(Into::into),
1242 suffix: suffix.map(Into::into),
1243 });
1244 self
1245 }
1246
1247 pub fn server_title(mut self, title: impl Into<String>) -> Self {
1249 Arc::make_mut(&mut self.inner).server_title = Some(title.into());
1250 self
1251 }
1252
1253 pub fn server_description(mut self, description: impl Into<String>) -> Self {
1255 Arc::make_mut(&mut self.inner).server_description = Some(description.into());
1256 self
1257 }
1258
1259 pub fn server_icons(mut self, icons: Vec<ToolIcon>) -> Self {
1261 Arc::make_mut(&mut self.inner).server_icons = Some(icons);
1262 self
1263 }
1264
1265 pub fn server_website_url(mut self, url: impl Into<String>) -> Self {
1267 Arc::make_mut(&mut self.inner).server_website_url = Some(url.into());
1268 self
1269 }
1270
1271 pub fn tool(mut self, tool: Tool) -> Self {
1273 Arc::make_mut(&mut self.inner)
1274 .tools
1275 .insert(tool.name.clone(), Arc::new(tool));
1276 self
1277 }
1278
1279 pub fn tool_if(self, condition: bool, tool: Tool) -> Self {
1305 if condition { self.tool(tool) } else { self }
1306 }
1307
1308 pub fn resource(mut self, resource: Resource) -> Self {
1310 Arc::make_mut(&mut self.inner)
1311 .resources
1312 .insert(resource.uri.clone(), Arc::new(resource));
1313 self
1314 }
1315
1316 pub fn resource_if(self, condition: bool, resource: Resource) -> Self {
1335 if condition {
1336 self.resource(resource)
1337 } else {
1338 self
1339 }
1340 }
1341
1342 pub fn resource_template(mut self, template: ResourceTemplate) -> Self {
1376 Arc::make_mut(&mut self.inner)
1377 .resource_templates
1378 .push(Arc::new(template));
1379 self
1380 }
1381
1382 pub fn prompt(mut self, prompt: Prompt) -> Self {
1384 Arc::make_mut(&mut self.inner)
1385 .prompts
1386 .insert(prompt.name.clone(), Arc::new(prompt));
1387 self
1388 }
1389
1390 pub fn prompt_if(self, condition: bool, prompt: Prompt) -> Self {
1409 if condition { self.prompt(prompt) } else { self }
1410 }
1411
1412 pub fn tools(self, tools: impl IntoIterator<Item = Tool>) -> Self {
1438 tools
1439 .into_iter()
1440 .fold(self, |router, tool| router.tool(tool))
1441 }
1442
1443 pub fn tools_if(self, condition: bool, tools: impl IntoIterator<Item = Tool>) -> Self {
1447 if condition { self.tools(tools) } else { self }
1448 }
1449
1450 pub fn resources(self, resources: impl IntoIterator<Item = Resource>) -> Self {
1469 resources
1470 .into_iter()
1471 .fold(self, |router, resource| router.resource(resource))
1472 }
1473
1474 pub fn resources_if(
1478 self,
1479 condition: bool,
1480 resources: impl IntoIterator<Item = Resource>,
1481 ) -> Self {
1482 if condition {
1483 self.resources(resources)
1484 } else {
1485 self
1486 }
1487 }
1488
1489 pub fn prompts(self, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1508 prompts
1509 .into_iter()
1510 .fold(self, |router, prompt| router.prompt(prompt))
1511 }
1512
1513 pub fn prompts_if(self, condition: bool, prompts: impl IntoIterator<Item = Prompt>) -> Self {
1517 if condition {
1518 self.prompts(prompts)
1519 } else {
1520 self
1521 }
1522 }
1523
1524 pub fn merge(mut self, other: McpRouter) -> Self {
1569 let inner = Arc::make_mut(&mut self.inner);
1570 let other_inner = other.inner;
1571
1572 for (name, tool) in &other_inner.tools {
1574 inner.tools.insert(name.clone(), tool.clone());
1575 }
1576
1577 for (uri, resource) in &other_inner.resources {
1579 inner.resources.insert(uri.clone(), resource.clone());
1580 }
1581
1582 for template in &other_inner.resource_templates {
1585 inner.resource_templates.push(template.clone());
1586 }
1587
1588 for (name, prompt) in &other_inner.prompts {
1590 inner.prompts.insert(name.clone(), prompt.clone());
1591 }
1592
1593 for (identifier, settings) in &other_inner.protocol_extensions {
1595 inner
1596 .protocol_extensions
1597 .insert(identifier.clone(), settings.clone());
1598 }
1599
1600 self
1601 }
1602
1603 pub fn nest(mut self, prefix: impl Into<String>, other: McpRouter) -> Self {
1643 let prefix = prefix.into();
1644 let inner = Arc::make_mut(&mut self.inner);
1645 let other_inner = other.inner;
1646
1647 for tool in other_inner.tools.values() {
1649 let prefixed_tool = tool.with_name_prefix(&prefix);
1650 inner
1651 .tools
1652 .insert(prefixed_tool.name.clone(), Arc::new(prefixed_tool));
1653 }
1654
1655 for (uri, resource) in &other_inner.resources {
1657 inner.resources.insert(uri.clone(), resource.clone());
1658 }
1659
1660 for template in &other_inner.resource_templates {
1662 inner.resource_templates.push(template.clone());
1663 }
1664
1665 for (name, prompt) in &other_inner.prompts {
1667 inner.prompts.insert(name.clone(), prompt.clone());
1668 }
1669
1670 for (identifier, settings) in &other_inner.protocol_extensions {
1673 inner
1674 .protocol_extensions
1675 .insert(identifier.clone(), settings.clone());
1676 }
1677
1678 self
1679 }
1680
1681 pub fn completion_handler<F, Fut>(mut self, handler: F) -> Self
1709 where
1710 F: Fn(CompleteParams) -> Fut + Send + Sync + 'static,
1711 Fut: Future<Output = Result<CompleteResult>> + Send + 'static,
1712 {
1713 Arc::make_mut(&mut self.inner).completion_handler =
1714 Some(Arc::new(move |params| Box::pin(handler(params))));
1715 self
1716 }
1717
1718 pub fn tool_filter(mut self, filter: ToolFilter) -> Self {
1753 Arc::make_mut(&mut self.inner).tool_filter = Some(filter);
1754 self
1755 }
1756
1757 pub fn resource_filter(mut self, filter: ResourceFilter) -> Self {
1788 Arc::make_mut(&mut self.inner).resource_filter = Some(filter);
1789 self
1790 }
1791
1792 pub fn prompt_filter(mut self, filter: PromptFilter) -> Self {
1821 Arc::make_mut(&mut self.inner).prompt_filter = Some(filter);
1822 self
1823 }
1824
1825 pub fn session(&self) -> &SessionState {
1827 &self.session
1828 }
1829
1830 pub fn log(&self, params: LoggingMessageParams) -> bool {
1852 let Some(tx) = &self.inner.notification_tx else {
1853 return false;
1854 };
1855 tx.try_send(ServerNotification::LogMessage(params)).is_ok()
1856 }
1857
1858 pub fn log_info(&self, message: &str) -> bool {
1862 self.log(LoggingMessageParams::new(
1863 LogLevel::Info,
1864 serde_json::json!({ "message": message }),
1865 ))
1866 }
1867
1868 pub fn log_warning(&self, message: &str) -> bool {
1870 self.log(LoggingMessageParams::new(
1871 LogLevel::Warning,
1872 serde_json::json!({ "message": message }),
1873 ))
1874 }
1875
1876 pub fn log_error(&self, message: &str) -> bool {
1878 self.log(LoggingMessageParams::new(
1879 LogLevel::Error,
1880 serde_json::json!({ "message": message }),
1881 ))
1882 }
1883
1884 pub fn log_debug(&self, message: &str) -> bool {
1886 self.log(LoggingMessageParams::new(
1887 LogLevel::Debug,
1888 serde_json::json!({ "message": message }),
1889 ))
1890 }
1891
1892 pub fn is_subscribed(&self, uri: &str) -> bool {
1894 if let Ok(subs) = self.inner.subscriptions.read() {
1895 return subs.contains(uri);
1896 }
1897 false
1898 }
1899
1900 pub fn subscribed_uris(&self) -> Vec<String> {
1902 if let Ok(subs) = self.inner.subscriptions.read() {
1903 return subs.iter().cloned().collect();
1904 }
1905 Vec::new()
1906 }
1907
1908 fn subscribe(&self, uri: &str) -> bool {
1910 if let Ok(mut subs) = self.inner.subscriptions.write() {
1911 return subs.insert(uri.to_string());
1912 }
1913 false
1914 }
1915
1916 fn unsubscribe(&self, uri: &str) -> bool {
1918 if let Ok(mut subs) = self.inner.subscriptions.write() {
1919 return subs.remove(uri);
1920 }
1921 false
1922 }
1923
1924 pub fn notify_resource_updated(&self, uri: &str) -> bool {
1931 let notification = ServerNotification::ResourceUpdated {
1932 uri: uri.to_string(),
1933 };
1934 let mut sent = false;
1935
1936 if self.is_subscribed(uri)
1937 && let Some(tx) = &self.inner.notification_tx
1938 {
1939 sent |= tx.try_send(notification.clone()).is_ok();
1940 }
1941
1942 #[cfg(all(feature = "http", feature = "stateless"))]
1943 if let Ok(active) = self.inner.modern_notification_sink.read()
1944 && let Some(sink) = active.as_ref()
1945 {
1946 sent |= sink(¬ification);
1947 }
1948
1949 sent
1950 }
1951
1952 pub async fn notify_task_status_changed(&self, task_id: &str) {
1967 self.notify_task_state(task_id).await;
1968 }
1969
1970 pub fn notify_resources_list_changed(&self) -> bool {
1974 let Some(tx) = &self.inner.notification_tx else {
1975 return false;
1976 };
1977 tx.try_send(ServerNotification::ResourcesListChanged)
1978 .is_ok()
1979 }
1980
1981 pub fn notify_tools_list_changed(&self) -> bool {
1985 let Some(tx) = &self.inner.notification_tx else {
1986 return false;
1987 };
1988 tx.try_send(ServerNotification::ToolsListChanged).is_ok()
1989 }
1990
1991 pub fn notify_prompts_list_changed(&self) -> bool {
1995 let Some(tx) = &self.inner.notification_tx else {
1996 return false;
1997 };
1998 tx.try_send(ServerNotification::PromptsListChanged).is_ok()
1999 }
2000
2001 pub fn disable_tool(&self, name: impl Into<String>) {
2012 let mut set = self.inner.disabled_tools.write().unwrap();
2013 set.insert(name.into());
2014 }
2015
2016 pub fn enable_tool(&self, name: &str) {
2019 let mut set = self.inner.disabled_tools.write().unwrap();
2020 set.remove(name);
2021 }
2022
2023 pub fn is_tool_enabled(&self, name: &str) -> bool {
2027 !self.inner.disabled_tools.read().unwrap().contains(name)
2028 }
2029
2030 pub fn disable_resource(&self, uri: impl Into<String>) {
2033 let mut set = self.inner.disabled_resources.write().unwrap();
2034 set.insert(uri.into());
2035 }
2036
2037 pub fn enable_resource(&self, uri: &str) {
2039 let mut set = self.inner.disabled_resources.write().unwrap();
2040 set.remove(uri);
2041 }
2042
2043 pub fn is_resource_enabled(&self, uri: &str) -> bool {
2045 !self.inner.disabled_resources.read().unwrap().contains(uri)
2046 }
2047
2048 pub fn disable_prompt(&self, name: impl Into<String>) {
2051 let mut set = self.inner.disabled_prompts.write().unwrap();
2052 set.insert(name.into());
2053 }
2054
2055 pub fn enable_prompt(&self, name: &str) {
2057 let mut set = self.inner.disabled_prompts.write().unwrap();
2058 set.remove(name);
2059 }
2060
2061 pub fn is_prompt_enabled(&self, name: &str) -> bool {
2063 !self.inner.disabled_prompts.read().unwrap().contains(name)
2064 }
2065
2066 pub(crate) fn implementation(&self) -> Implementation {
2076 Implementation {
2077 name: self.inner.server_name.clone(),
2078 version: self.inner.server_version.clone(),
2079 title: self.inner.server_title.clone(),
2080 description: self.inner.server_description.clone(),
2081 icons: self.inner.server_icons.clone(),
2082 website_url: self.inner.server_website_url.clone(),
2083 meta: None,
2084 }
2085 }
2086
2087 #[cfg(feature = "http")]
2093 pub(crate) fn tool_input_schema(&self, name: &str) -> Option<serde_json::Value> {
2094 if let Some(tool) = self.inner.tools.get(name) {
2095 return Some(tool.input_schema.clone());
2096 }
2097 #[cfg(feature = "dynamic-tools")]
2098 if let Some(tool) = self
2099 .inner
2100 .dynamic_tools
2101 .as_ref()
2102 .and_then(|tools| tools.get(name))
2103 {
2104 return Some(tool.input_schema.clone());
2105 }
2106 None
2107 }
2108
2109 fn capabilities(&self) -> ServerCapabilities {
2110 let has_resources =
2111 !self.inner.resources.is_empty() || !self.inner.resource_templates.is_empty();
2112 let has_notifications = self.inner.notification_tx.is_some();
2113
2114 #[cfg(feature = "dynamic-tools")]
2115 let has_dynamic_tools = self.inner.dynamic_tools.is_some();
2116 #[cfg(not(feature = "dynamic-tools"))]
2117 let has_dynamic_tools = false;
2118
2119 #[cfg(feature = "dynamic-tools")]
2120 let has_dynamic_prompts = self.inner.dynamic_prompts.is_some();
2121 #[cfg(not(feature = "dynamic-tools"))]
2122 let has_dynamic_prompts = false;
2123
2124 #[cfg(feature = "dynamic-tools")]
2125 let has_dynamic_resources = self.inner.dynamic_resources.is_some()
2126 || self.inner.dynamic_resource_templates.is_some();
2127 #[cfg(not(feature = "dynamic-tools"))]
2128 let has_dynamic_resources = false;
2129
2130 ServerCapabilities {
2131 tools: if self.inner.tools.is_empty() && !has_dynamic_tools {
2132 None
2133 } else {
2134 Some(ToolsCapability {
2135 list_changed: has_notifications,
2136 })
2137 },
2138 resources: if has_resources || has_dynamic_resources {
2139 Some(ResourcesCapability {
2140 subscribe: true,
2141 list_changed: has_notifications,
2142 })
2143 } else {
2144 None
2145 },
2146 prompts: if self.inner.prompts.is_empty() && !has_dynamic_prompts {
2147 None
2148 } else {
2149 Some(PromptsCapability {
2150 list_changed: has_notifications,
2151 })
2152 },
2153 logging: if self.inner.notification_tx.is_some() {
2155 Some(LoggingCapability {
2156 deprecated: self.inner.logging_deprecated.clone(),
2157 })
2158 } else {
2159 None
2160 },
2161 tasks: {
2167 let has_task_support = self
2168 .inner
2169 .tools
2170 .values()
2171 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2172 if has_task_support {
2173 Some(TasksCapability {
2174 list: None,
2178 cancel: Some(TasksCancelCapability {}),
2179 requests: Some(TasksRequestsCapability {
2180 tools: Some(TasksToolsRequestsCapability {
2181 call: Some(TasksToolsCallCapability {}),
2182 }),
2183 }),
2184 })
2185 } else {
2186 None
2187 }
2188 },
2189 completions: if self.inner.completion_handler.is_some() {
2191 Some(CompletionsCapability::default())
2192 } else {
2193 None
2194 },
2195 experimental: None,
2196 extensions: {
2197 let mut map = self.inner.protocol_extensions.clone();
2198 let has_task_support = self
2199 .inner
2200 .tools
2201 .values()
2202 .any(|t| !matches!(t.task_support, TaskSupportMode::Forbidden));
2203 if has_task_support {
2204 map.insert(
2205 tower_mcp_types::protocol::TASKS_EXTENSION_ID.to_string(),
2206 serde_json::json!({}),
2207 );
2208 }
2209 (!map.is_empty()).then_some(map)
2210 },
2211 }
2212 }
2213
2214 fn capabilities_for_protocol(&self, protocol_version: Option<&str>) -> ServerCapabilities {
2222 let mut capabilities = self.capabilities();
2223 if protocol_version == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28) {
2224 capabilities.tasks = None;
2225 if !self.final_tasks_enabled()
2226 && let Some(extensions) = capabilities.extensions.as_mut()
2227 {
2228 extensions.remove(tower_mcp_types::protocol::TASKS_EXTENSION_ID);
2229 if extensions.is_empty() {
2230 capabilities.extensions = None;
2231 }
2232 }
2233 }
2234 capabilities
2235 }
2236
2237 pub(crate) fn final_tasks_enabled(&self) -> bool {
2242 self.inner
2243 .protocol_extensions
2244 .contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
2245 }
2246
2247 fn require_negotiated_tasks(
2252 &self,
2253 extensions: &crate::context::Extensions,
2254 method: &str,
2255 ) -> Result<()> {
2256 if !self.final_tasks_enabled() {
2257 return Err(Error::JsonRpc(JsonRpcError::method_not_found(method)));
2258 }
2259 if client_declares_tasks(extensions) {
2260 return Ok(());
2261 }
2262 Err(Error::JsonRpc(
2263 JsonRpcError::missing_required_client_capability(tasks_client_capabilities()),
2264 ))
2265 }
2266
2267 async fn authorize_task(
2273 &self,
2274 task_id: &str,
2275 extensions: &crate::context::Extensions,
2276 ) -> Result<()> {
2277 let owner = self
2278 .inner
2279 .task_store
2280 .task_owner(task_id)
2281 .await
2282 .map_err(task_store_error)?
2283 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2284
2285 if crate::async_task::owner_matches(&owner, request_principal(extensions).as_deref()) {
2286 Ok(())
2287 } else {
2288 tracing::debug!(
2289 target: "mcp::tasks",
2290 task_id = %task_id,
2291 "task operation refused: principal does not own the task"
2292 );
2293 Err(Error::JsonRpc(unknown_task_error(task_id)))
2294 }
2295 }
2296
2297 async fn final_get_task(&self, task_id: &str) -> Result<McpResponse> {
2299 let (detailed, meta) = self.detailed_task(task_id).await?;
2300 let mut result = crate::tasks::GetTaskResult::new(detailed);
2301 result.meta = meta;
2302 Ok(McpResponse::FinalGetTask(result))
2303 }
2304
2305 async fn detailed_task(
2311 &self,
2312 task_id: &str,
2313 ) -> Result<(
2314 crate::tasks::DetailedTask,
2315 Option<serde_json::Map<String, serde_json::Value>>,
2316 )> {
2317 let (task, result, error) = self
2318 .inner
2319 .task_store
2320 .get_task_result(task_id)
2321 .await
2322 .map_err(task_store_error)?
2323 .ok_or_else(|| Error::JsonRpc(unknown_task_error(task_id)))?;
2324
2325 let mut metadata = crate::tasks::TaskMetadata::new(
2326 task.task_id.clone(),
2327 task.created_at.clone(),
2328 task.last_updated_at.clone(),
2329 task.ttl,
2330 );
2331 metadata.status_message = task.status_message.clone();
2332 metadata.poll_interval_ms = task.poll_interval;
2333
2334 let meta = task.meta.and_then(|value| value.as_object().cloned());
2335 let detailed = match task.status {
2336 TaskStatus::Working => crate::tasks::DetailedTask::working(metadata),
2337 TaskStatus::InputRequired => {
2338 let outstanding = self
2341 .inner
2342 .task_store
2343 .outstanding_input_requests(task_id)
2344 .await
2345 .map_err(task_store_error)?
2346 .unwrap_or_default();
2347 crate::tasks::DetailedTask::input_required(metadata, outstanding)
2348 }
2349 TaskStatus::Completed => {
2350 let mut object = result
2353 .map(serde_json::to_value)
2354 .transpose()
2355 .map_err(|e| {
2356 Error::JsonRpc(JsonRpcError::internal_error(format!(
2357 "failed to encode task result: {e}"
2358 )))
2359 })?
2360 .and_then(|value| value.as_object().cloned())
2361 .unwrap_or_default();
2362 object.insert(
2366 "resultType".to_string(),
2367 serde_json::Value::String("complete".to_string()),
2368 );
2369 crate::tasks::DetailedTask::completed(metadata, object)
2370 }
2371 TaskStatus::Failed => crate::tasks::DetailedTask::failed(
2372 metadata,
2373 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
2374 ),
2375 TaskStatus::Cancelled => crate::tasks::DetailedTask::cancelled(metadata),
2376 _ => crate::tasks::DetailedTask::working(metadata),
2379 };
2380 Ok((detailed, meta))
2381 }
2382
2383 async fn notify_task_state(&self, task_id: &str) {
2392 if !self.final_tasks_enabled() {
2393 return;
2394 }
2395
2396 let (detailed, meta) = match self.detailed_task(task_id).await {
2397 Ok(detailed) => detailed,
2398 Err(error) => {
2399 tracing::debug!(
2400 target: "mcp::tasks",
2401 task_id = %task_id,
2402 %error,
2403 "skipping task notification: task state unavailable"
2404 );
2405 return;
2406 }
2407 };
2408
2409 let notification = ServerNotification::FinalTaskStatusChanged(
2410 crate::tasks::TaskStatusNotificationParams {
2411 task: detailed,
2412 meta,
2413 },
2414 );
2415
2416 #[cfg(all(feature = "http", feature = "stateless"))]
2421 if let Ok(active) = self.inner.modern_notification_sink.read()
2422 && let Some(sink) = active.as_ref()
2423 {
2424 sink(¬ification);
2425 return;
2426 }
2427
2428 if let Some(tx) = &self.inner.notification_tx {
2429 let _ = tx.try_send(notification);
2430 }
2431 }
2432
2433 fn effective_cache_scope(&self, ttl_ms: Option<u64>) -> Option<CacheScope> {
2440 self.inner
2441 .cache_scope
2442 .or_else(|| ttl_ms.map(|_| CacheScope::Private))
2443 }
2444
2445 fn apply_read_cache_hints(&self, mut result: ReadResourceResult) -> ReadResourceResult {
2450 if result.ttl_ms.is_none() {
2451 result.ttl_ms = self.inner.read_ttl_ms;
2452 }
2453 if result.cache_scope.is_none() {
2454 result.cache_scope = self.effective_cache_scope(result.ttl_ms);
2455 }
2456 result
2457 }
2458
2459 async fn handle(
2461 &self,
2462 request_id: RequestId,
2463 request: McpRequest,
2464 extensions: Extensions,
2465 ) -> Result<McpResponse> {
2466 let method = request.method_name();
2468 if !is_final_protocol_request(&extensions) && !self.session.is_request_allowed(method) {
2469 tracing::warn!(
2470 method = %method,
2471 phase = ?self.session.phase(),
2472 "Request rejected: session not initialized"
2473 );
2474 return Err(Error::JsonRpc(JsonRpcError::invalid_request(format!(
2475 "Session not initialized. Only 'initialize' and 'ping' are allowed before initialization. Got: {}",
2476 method
2477 ))));
2478 }
2479
2480 match request {
2481 McpRequest::Initialize(params) => {
2482 tracing::info!(
2483 client = %params.client_info.name,
2484 version = %params.client_info.version,
2485 "Client initializing"
2486 );
2487
2488 let protocol_support = extensions.get::<crate::ProtocolSupport>();
2492 let requested_is_legacy = crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2493 .contains(¶ms.protocol_version.as_str());
2494 let requested_is_supported = requested_is_legacy
2495 && protocol_support
2496 .is_none_or(|support| support.contains(¶ms.protocol_version));
2497 let protocol_version = if requested_is_supported {
2498 params.protocol_version
2499 } else {
2500 match protocol_support {
2501 None => crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
2502 Some(support) => support
2503 .versions()
2504 .iter()
2505 .find(|version| {
2506 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2507 .contains(&version.as_str())
2508 })
2509 .cloned()
2510 .ok_or_else(|| {
2511 Error::JsonRpc(JsonRpcError::unsupported_protocol_version(
2512 params.protocol_version,
2513 support.versions().iter().map(String::as_str),
2514 ))
2515 })?,
2516 }
2517 };
2518
2519 self.session.mark_initializing();
2521 let capabilities = self.capabilities_for_protocol(Some(&protocol_version));
2522 self.session.insert(params.capabilities.clone());
2523 self.session
2524 .insert(crate::NegotiatedExtensions::from_capabilities(
2525 ¶ms.capabilities,
2526 &capabilities,
2527 ));
2528
2529 Ok(McpResponse::Initialize(InitializeResult {
2530 protocol_version,
2531 capabilities,
2532 server_info: self.implementation(),
2533 instructions: if let Some(config) = &self.inner.auto_instructions {
2534 Some(self.inner.generate_instructions(config))
2535 } else {
2536 self.inner.instructions.clone()
2537 },
2538 meta: None,
2539 }))
2540 }
2541
2542 McpRequest::Discover(_) => {
2543 tracing::debug!("Stateless server/discover request");
2550 let server_info = self.implementation();
2551 let supported_versions = extensions.get::<crate::ProtocolSupport>().map_or_else(
2552 || {
2553 crate::protocol::SUPPORTED_PROTOCOL_VERSIONS
2554 .iter()
2555 .map(|version| (*version).to_string())
2556 .collect()
2557 },
2558 |support| support.versions().to_vec(),
2559 );
2560 let capabilities = self
2565 .capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
2566 Ok(McpResponse::Discover(DiscoverResult {
2567 supported_versions,
2568 capabilities,
2569 ttl_ms: None,
2570 cache_scope: None,
2571 instructions: if let Some(config) = &self.inner.auto_instructions {
2572 Some(self.inner.generate_instructions(config))
2573 } else {
2574 self.inner.instructions.clone()
2575 },
2576 meta: Some(crate::protocol::ResultMeta {
2577 server_info: Some(server_info),
2578 }),
2579 }))
2580 }
2581
2582 McpRequest::ListTools(params) => {
2583 let final_protocol = is_final_protocol_request(&extensions);
2584 let final_tasks_negotiated = final_protocol
2585 && self.final_tasks_enabled()
2586 && client_declares_tasks(&extensions);
2587 let filter = self.inner.tool_filter.as_ref();
2588 let disabled = self.inner.disabled_tools.read().unwrap().clone();
2589 let is_visible = |t: &Tool| {
2590 !disabled.contains(&t.name)
2591 && !(final_protocol
2592 && matches!(t.task_support, TaskSupportMode::Required)
2593 && !final_tasks_negotiated)
2594 && filter
2595 .map(|f| f.is_visible(&self.session, t))
2596 .unwrap_or(true)
2597 };
2598 let definition = |t: &Tool| {
2599 let mut definition = t.definition();
2600 if final_protocol {
2601 definition.execution = None;
2602 }
2603 definition
2604 };
2605
2606 let mut tools: Vec<ToolDefinition> = self
2608 .inner
2609 .tools
2610 .values()
2611 .filter(|t| is_visible(t))
2612 .map(|t| definition(t))
2613 .collect();
2614
2615 #[cfg(feature = "dynamic-tools")]
2617 if let Some(ref dynamic) = self.inner.dynamic_tools {
2618 let static_names: HashSet<String> =
2619 tools.iter().map(|t| t.name.clone()).collect();
2620 for t in dynamic.list() {
2621 if !static_names.contains(&t.name) && is_visible(&t) {
2622 tools.push(definition(&t));
2623 }
2624 }
2625 }
2626
2627 tools.sort_by(|a, b| a.name.cmp(&b.name));
2628
2629 let (tools, next_cursor) =
2630 paginate(tools, params.cursor.as_deref(), self.inner.page_size)?;
2631
2632 Ok(McpResponse::ListTools(ListToolsResult {
2633 tools,
2634 next_cursor,
2635 ttl_ms: self.inner.list_ttl_ms,
2636 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
2637 meta: None,
2638 }))
2639 }
2640
2641 McpRequest::CallTool(params) => {
2642 if self
2644 .inner
2645 .disabled_tools
2646 .read()
2647 .unwrap()
2648 .contains(¶ms.name)
2649 {
2650 tracing::info!(
2651 target: "mcp::tools",
2652 tool = %params.name,
2653 status = "disabled",
2654 "tool call completed"
2655 );
2656 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2657 }
2658
2659 let tool = self.inner.tools.get(¶ms.name).cloned();
2661 #[cfg(feature = "dynamic-tools")]
2662 let tool = tool.or_else(|| {
2663 self.inner
2664 .dynamic_tools
2665 .as_ref()
2666 .and_then(|d| d.get(¶ms.name))
2667 });
2668
2669 let tool = match tool {
2670 Some(t) => t,
2671 None => {
2672 tracing::info!(
2673 target: "mcp::tools",
2674 tool = %params.name,
2675 status = "not_found",
2676 "tool call completed"
2677 );
2678 return Err(Error::JsonRpc(JsonRpcError::method_not_found(¶ms.name)));
2679 }
2680 };
2681
2682 if let Some(filter) = &self.inner.tool_filter
2684 && !filter.is_visible(&self.session, &tool)
2685 {
2686 tracing::info!(
2687 target: "mcp::tools",
2688 tool = %params.name,
2689 status = "denied",
2690 "tool call completed"
2691 );
2692 return Err(filter.denial_error(¶ms.name));
2693 }
2694
2695 let final_protocol = is_final_protocol_request(&extensions);
2699 let task_ttl = if final_protocol {
2700 if params.task.is_some() {
2701 return Err(Error::JsonRpc(JsonRpcError::invalid_params(
2702 "The final Tasks extension does not allow a 'task' request parameter",
2703 )));
2704 }
2705
2706 let server_enabled = self.final_tasks_enabled();
2707 let tasks_negotiated = server_enabled && client_declares_tasks(&extensions);
2708 match tool.task_support {
2709 TaskSupportMode::Required if !server_enabled => {
2710 return Err(Error::JsonRpc(JsonRpcError::method_not_found(
2713 ¶ms.name,
2714 )));
2715 }
2716 TaskSupportMode::Required if !tasks_negotiated => {
2717 return Err(Error::JsonRpc(
2718 JsonRpcError::missing_required_client_capability(
2719 tasks_client_capabilities(),
2720 ),
2721 ));
2722 }
2723 TaskSupportMode::Required | TaskSupportMode::Optional
2724 if tasks_negotiated =>
2725 {
2726 Some(None)
2727 }
2728 _ => None,
2729 }
2730 } else {
2731 match (¶ms.task, tool.task_support) {
2732 (Some(_), TaskSupportMode::Forbidden) => {
2733 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2734 "Tool '{}' does not support async tasks",
2735 params.name
2736 ))));
2737 }
2738 (None, TaskSupportMode::Required) => {
2739 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
2740 "Tool '{}' requires async task execution (include 'task' in params)",
2741 params.name
2742 ))));
2743 }
2744 (Some(task), _) => Some(task.ttl),
2745 (None, _) => None,
2746 }
2747 };
2748
2749 #[cfg(feature = "stateless")]
2753 if let Some(required) = tool.required_client_capabilities()
2754 && let Some(meta) = extensions.get::<crate::stateless::StatelessRequestMeta>()
2755 && meta.protocol_version.as_deref()
2756 == Some(crate::protocol::PROTOCOL_VERSION_2026_07_28)
2757 && !meta
2758 .client_capabilities
2759 .as_ref()
2760 .is_some_and(|actual| client_capabilities_satisfy(actual, required))
2761 {
2762 return Err(Error::JsonRpc(
2763 JsonRpcError::missing_required_client_capability(required.clone()),
2764 ));
2765 }
2766
2767 if let Some(task_ttl) = task_ttl {
2768 let (task_id, cancellation_token) = self
2770 .inner
2771 .task_store
2772 .create_task(
2773 ¶ms.name,
2774 params.arguments.clone(),
2775 task_ttl,
2776 request_principal(&extensions),
2777 )
2778 .await
2779 .map_err(task_store_error)?;
2780
2781 tracing::info!(task_id = %task_id, tool = %params.name, "Created async task");
2782
2783 let progress_token = params.meta.and_then(|m| m.progress_token);
2785 let ctx = self.create_context_with_extensions(
2786 request_id,
2787 progress_token,
2788 &extensions,
2789 );
2790
2791 let task_store = self.inner.task_store.clone();
2792 let task_context = crate::tool::TaskContext::new(task_id.clone());
2793 let mut ctx = ctx;
2794 ctx.extensions_mut().insert(task_context.clone());
2795 let preparation = match tool
2796 .prepare_task(task_context, params.arguments.clone())
2797 .await
2798 {
2799 Ok(preparation) => preparation,
2800 Err(error) => {
2801 discard_unprepared_task(&task_store, &task_id).await;
2802 return Err(error);
2803 }
2804 };
2805 if let Some(meta) = preparation.meta {
2806 let value = serde_json::Value::Object(meta);
2807 if let Err(error) = crate::protocol::validate_meta_object(&value) {
2808 discard_unprepared_task(&task_store, &task_id).await;
2809 return Err(Error::invalid_params(format!(
2810 "Invalid task metadata: {error}"
2811 )));
2812 }
2813 let persisted = match task_store.set_task_meta(&task_id, value).await {
2814 Ok(persisted) => persisted,
2815 Err(error) => {
2816 discard_unprepared_task(&task_store, &task_id).await;
2817 return Err(task_store_error(error));
2818 }
2819 };
2820 if !persisted {
2821 discard_unprepared_task(&task_store, &task_id).await;
2822 return Err(Error::JsonRpc(JsonRpcError::internal_error(
2823 "Task store could not persist preparation metadata",
2824 )));
2825 }
2826 }
2827 ctx.extensions_mut().merge(&preparation.extensions);
2828
2829 let tool = tool.clone();
2831 let arguments = params.arguments;
2832 let task_id_clone = task_id.clone();
2833
2834 let tool_name = params.name.clone();
2835 let notifier = self.clone();
2836 tokio::spawn(async move {
2837 if cancellation_token.is_cancelled() {
2839 tracing::debug!(task_id = %task_id_clone, "Task cancelled before execution");
2840 notifier.notify_task_state(&task_id_clone).await;
2841 return;
2842 }
2843
2844 let start = std::time::Instant::now();
2846 let result = tool.call_with_context(ctx, arguments).await;
2847 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
2848
2849 if cancellation_token.is_cancelled() {
2850 tracing::debug!(task_id = %task_id_clone, "Task cancelled during execution");
2851 notifier.notify_task_state(&task_id_clone).await;
2852 } else {
2853 let status = if result.is_error { "error" } else { "success" };
2858 let error_msg = result
2859 .is_error
2860 .then(|| result.first_text().unwrap_or("Tool execution failed"))
2861 .map(str::to_string);
2862 if let Err(e) = task_store.complete_task(&task_id_clone, result).await {
2863 tracing::warn!(task_id = %task_id_clone, error = %e, "failed to record task completion");
2864 }
2865 tracing::info!(
2866 target: "mcp::tools",
2867 tool = %tool_name,
2868 task_id = %task_id_clone,
2869 duration_ms,
2870 status,
2871 error = error_msg.as_deref().unwrap_or_default(),
2872 "tool call completed"
2873 );
2874 notifier.notify_task_state(&task_id_clone).await;
2875 }
2876 });
2877
2878 let task = self
2879 .inner
2880 .task_store
2881 .get_task(&task_id)
2882 .await
2883 .map_err(task_store_error)?
2884 .ok_or_else(|| {
2885 Error::JsonRpc(JsonRpcError::internal_error(
2886 "Failed to retrieve created task",
2887 ))
2888 })?;
2889
2890 if is_final_protocol_request(&extensions) {
2894 let mut metadata = crate::tasks::TaskMetadata::new(
2895 task.task_id.clone(),
2896 task.created_at.clone(),
2897 task.last_updated_at.clone(),
2898 task.ttl,
2899 );
2900 metadata.status_message = task.status_message.clone();
2901 metadata.poll_interval_ms = task.poll_interval;
2902 let mut result = crate::tasks::CreateTaskResult::new(
2903 crate::tasks::Task::new(metadata, task.status),
2904 );
2905 result.meta = task.meta.and_then(|value| value.as_object().cloned());
2906 return Ok(McpResponse::FinalCreateTask(result));
2907 }
2908 Ok(McpResponse::CreateTask(CreateTaskResult::new(task)))
2909 } else {
2910 let progress_token = params.meta.and_then(|m| m.progress_token);
2912 let ctx = self.create_context_with_extensions(
2913 request_id,
2914 progress_token,
2915 &extensions,
2916 );
2917 #[cfg(feature = "stateless")]
2918 let ctx = {
2919 let mut ctx = ctx;
2920 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
2921 params.input_responses,
2922 params.request_state,
2923 ));
2924 ctx
2925 };
2926
2927 let start = std::time::Instant::now();
2928 let outcome = tool
2929 .call_outcome_with_context(ctx, params.arguments)
2930 .await?;
2931 let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
2932
2933 match outcome {
2934 RequestOutcome::Complete(result) => {
2935 let status = if result.is_error { "error" } else { "success" };
2936 tracing::info!(
2937 target: "mcp::tools",
2938 tool = %params.name,
2939 duration_ms,
2940 status,
2941 "tool call completed"
2942 );
2943 Ok(McpResponse::CallTool(result))
2944 }
2945 RequestOutcome::InputRequired(result) => {
2946 #[cfg(feature = "stateless")]
2947 {
2948 validate_input_required_result(&extensions, &result)?;
2949 tracing::info!(
2950 target: "mcp::tools",
2951 tool = %params.name,
2952 duration_ms,
2953 status = "input_required",
2954 "tool call requires client input"
2955 );
2956 Ok(McpResponse::InputRequired(result))
2957 }
2958 #[cfg(not(feature = "stateless"))]
2959 {
2960 let _ = result;
2961 Err(Error::invalid_params(
2962 "InputRequiredResult support was not compiled",
2963 ))
2964 }
2965 }
2966 }
2967 }
2968 }
2969
2970 McpRequest::ListResources(params) => {
2971 let disabled = self.inner.disabled_resources.read().unwrap().clone();
2972 let is_visible = |r: &Resource| -> bool {
2973 !disabled.contains(&r.uri)
2974 && self
2975 .inner
2976 .resource_filter
2977 .as_ref()
2978 .map(|f| f.is_visible(&self.session, r))
2979 .unwrap_or(true)
2980 };
2981
2982 let mut resources: Vec<ResourceDefinition> = self
2983 .inner
2984 .resources
2985 .values()
2986 .filter(|r| is_visible(r))
2987 .map(|r| r.definition())
2988 .collect();
2989
2990 #[cfg(feature = "dynamic-tools")]
2992 if let Some(ref dynamic) = self.inner.dynamic_resources {
2993 let static_uris: HashSet<String> =
2994 resources.iter().map(|r| r.uri.clone()).collect();
2995 for r in dynamic.list() {
2996 if !static_uris.contains(&r.uri) && is_visible(&r) {
2997 resources.push(r.definition());
2998 }
2999 }
3000 }
3001
3002 resources.sort_by(|a, b| a.uri.cmp(&b.uri));
3003
3004 let (resources, next_cursor) =
3005 paginate(resources, params.cursor.as_deref(), self.inner.page_size)?;
3006
3007 Ok(McpResponse::ListResources(ListResourcesResult {
3008 resources,
3009 next_cursor,
3010 ttl_ms: self.inner.list_ttl_ms,
3011 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3012 meta: None,
3013 }))
3014 }
3015
3016 McpRequest::ListResourceTemplates(params) => {
3017 let mut resource_templates: Vec<ResourceTemplateDefinition> = self
3018 .inner
3019 .resource_templates
3020 .iter()
3021 .map(|t| t.definition())
3022 .collect();
3023
3024 #[cfg(feature = "dynamic-tools")]
3026 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3027 let static_patterns: HashSet<String> = resource_templates
3028 .iter()
3029 .map(|t| t.uri_template.clone())
3030 .collect();
3031 for t in dynamic.list() {
3032 if !static_patterns.contains(&t.uri_template) {
3033 resource_templates.push(t.definition());
3034 }
3035 }
3036 }
3037
3038 resource_templates.sort_by(|a, b| a.uri_template.cmp(&b.uri_template));
3039
3040 let (resource_templates, next_cursor) = paginate(
3041 resource_templates,
3042 params.cursor.as_deref(),
3043 self.inner.page_size,
3044 )?;
3045
3046 Ok(McpResponse::ListResourceTemplates(
3047 ListResourceTemplatesResult {
3048 resource_templates,
3049 next_cursor,
3050 ttl_ms: self.inner.list_ttl_ms,
3051 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3052 meta: None,
3053 },
3054 ))
3055 }
3056
3057 McpRequest::ReadResource(params) => {
3058 if self
3060 .inner
3061 .disabled_resources
3062 .read()
3063 .unwrap()
3064 .contains(¶ms.uri)
3065 {
3066 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3067 ¶ms.uri,
3068 )));
3069 }
3070
3071 if let Some(resource) = self.inner.resources.get(¶ms.uri) {
3073 if let Some(filter) = &self.inner.resource_filter
3075 && !filter.is_visible(&self.session, resource)
3076 {
3077 return Err(filter.denial_error(¶ms.uri));
3078 }
3079
3080 tracing::debug!(uri = %params.uri, "Reading static resource");
3081 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3082 #[cfg(feature = "stateless")]
3083 let ctx = {
3084 let mut ctx = ctx;
3085 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3086 params.input_responses.clone(),
3087 params.request_state.clone(),
3088 ));
3089 ctx
3090 };
3091 return match resource.read_outcome_with_context(ctx).await? {
3092 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3093 self.apply_read_cache_hints(result),
3094 )),
3095 RequestOutcome::InputRequired(result) => {
3096 #[cfg(feature = "stateless")]
3097 {
3098 validate_input_required_result(&extensions, &result)?;
3099 Ok(McpResponse::InputRequired(result))
3100 }
3101 #[cfg(not(feature = "stateless"))]
3102 {
3103 let _ = result;
3104 Err(Error::invalid_params(
3105 "InputRequiredResult support was not compiled",
3106 ))
3107 }
3108 }
3109 };
3110 }
3111
3112 #[cfg(feature = "dynamic-tools")]
3114 #[allow(clippy::collapsible_if)]
3115 if let Some(ref dynamic) = self.inner.dynamic_resources {
3116 if let Some(resource) = dynamic.get(¶ms.uri) {
3117 if let Some(filter) = &self.inner.resource_filter
3118 && !filter.is_visible(&self.session, &resource)
3119 {
3120 return Err(filter.denial_error(¶ms.uri));
3121 }
3122 tracing::debug!(uri = %params.uri, "Reading dynamic resource");
3123 let ctx =
3124 self.create_context_with_extensions(request_id, None, &extensions);
3125 #[cfg(feature = "stateless")]
3126 let ctx = {
3127 let mut ctx = ctx;
3128 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3129 params.input_responses.clone(),
3130 params.request_state.clone(),
3131 ));
3132 ctx
3133 };
3134 return match resource.read_outcome_with_context(ctx).await? {
3135 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3136 self.apply_read_cache_hints(result),
3137 )),
3138 RequestOutcome::InputRequired(result) => {
3139 #[cfg(feature = "stateless")]
3140 {
3141 validate_input_required_result(&extensions, &result)?;
3142 Ok(McpResponse::InputRequired(result))
3143 }
3144 #[cfg(not(feature = "stateless"))]
3145 {
3146 let _ = result;
3147 Err(Error::invalid_params(
3148 "InputRequiredResult support was not compiled",
3149 ))
3150 }
3151 }
3152 };
3153 }
3154 }
3155
3156 for template in &self.inner.resource_templates {
3158 if let Some(variables) = template.match_uri(¶ms.uri) {
3159 tracing::debug!(
3160 uri = %params.uri,
3161 template = %template.uri_template,
3162 "Reading resource via template"
3163 );
3164 let ctx =
3165 self.create_context_with_extensions(request_id, None, &extensions);
3166 #[cfg(feature = "stateless")]
3167 let ctx = {
3168 let mut ctx = ctx;
3169 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3170 params.input_responses.clone(),
3171 params.request_state.clone(),
3172 ));
3173 ctx
3174 };
3175 return match template
3176 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3177 .await?
3178 {
3179 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3180 self.apply_read_cache_hints(result),
3181 )),
3182 RequestOutcome::InputRequired(result) => {
3183 #[cfg(feature = "stateless")]
3184 {
3185 validate_input_required_result(&extensions, &result)?;
3186 Ok(McpResponse::InputRequired(result))
3187 }
3188 #[cfg(not(feature = "stateless"))]
3189 {
3190 let _ = result;
3191 Err(Error::invalid_params(
3192 "InputRequiredResult support was not compiled",
3193 ))
3194 }
3195 }
3196 };
3197 }
3198 }
3199
3200 #[cfg(feature = "dynamic-tools")]
3202 #[allow(clippy::collapsible_if)]
3203 if let Some(ref dynamic) = self.inner.dynamic_resource_templates {
3204 if let Some((template, variables)) = dynamic.match_uri(¶ms.uri) {
3205 tracing::debug!(
3206 uri = %params.uri,
3207 template = %template.uri_template,
3208 "Reading resource via dynamic template"
3209 );
3210 let ctx =
3211 self.create_context_with_extensions(request_id, None, &extensions);
3212 #[cfg(feature = "stateless")]
3213 let ctx = {
3214 let mut ctx = ctx;
3215 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3216 params.input_responses.clone(),
3217 params.request_state.clone(),
3218 ));
3219 ctx
3220 };
3221 return match template
3222 .read_outcome_with_context(ctx, ¶ms.uri, variables)
3223 .await?
3224 {
3225 RequestOutcome::Complete(result) => Ok(McpResponse::ReadResource(
3226 self.apply_read_cache_hints(result),
3227 )),
3228 RequestOutcome::InputRequired(result) => {
3229 #[cfg(feature = "stateless")]
3230 {
3231 validate_input_required_result(&extensions, &result)?;
3232 Ok(McpResponse::InputRequired(result))
3233 }
3234 #[cfg(not(feature = "stateless"))]
3235 {
3236 let _ = result;
3237 Err(Error::invalid_params(
3238 "InputRequiredResult support was not compiled",
3239 ))
3240 }
3241 }
3242 };
3243 }
3244 }
3245
3246 Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3248 ¶ms.uri,
3249 )))
3250 }
3251
3252 McpRequest::SubscribeResource(params) => {
3253 if !self.inner.resources.contains_key(¶ms.uri) {
3255 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3256 ¶ms.uri,
3257 )));
3258 }
3259
3260 tracing::debug!(uri = %params.uri, "Subscribing to resource");
3261 self.subscribe(¶ms.uri);
3262
3263 Ok(McpResponse::SubscribeResource(EmptyResult {}))
3264 }
3265
3266 McpRequest::UnsubscribeResource(params) => {
3267 if !self.inner.resources.contains_key(¶ms.uri) {
3269 return Err(Error::JsonRpc(JsonRpcError::resource_not_found(
3270 ¶ms.uri,
3271 )));
3272 }
3273
3274 tracing::debug!(uri = %params.uri, "Unsubscribing from resource");
3275 self.unsubscribe(¶ms.uri);
3276
3277 Ok(McpResponse::UnsubscribeResource(EmptyResult {}))
3278 }
3279
3280 McpRequest::ListPrompts(params) => {
3281 #[cfg(feature = "dynamic-tools")]
3282 if let Some(initializer) = &self.inner.prompt_initializer {
3283 initializer()?;
3284 }
3285 let disabled = self.inner.disabled_prompts.read().unwrap().clone();
3286 let is_visible = |p: &Prompt| -> bool {
3287 !disabled.contains(&p.name)
3288 && self
3289 .inner
3290 .prompt_filter
3291 .as_ref()
3292 .map(|f| f.is_visible(&self.session, p))
3293 .unwrap_or(true)
3294 };
3295
3296 let mut prompts: Vec<PromptDefinition> = self
3297 .inner
3298 .prompts
3299 .values()
3300 .filter(|p| is_visible(p))
3301 .map(|p| p.definition())
3302 .collect();
3303
3304 #[cfg(feature = "dynamic-tools")]
3306 if let Some(ref dynamic) = self.inner.dynamic_prompts {
3307 let static_names: HashSet<String> =
3308 prompts.iter().map(|p| p.name.clone()).collect();
3309 for p in dynamic.list() {
3310 if !static_names.contains(&p.name) && is_visible(&p) {
3311 prompts.push(p.definition());
3312 }
3313 }
3314 }
3315
3316 prompts.sort_by(|a, b| a.name.cmp(&b.name));
3317
3318 let (prompts, next_cursor) =
3319 paginate(prompts, params.cursor.as_deref(), self.inner.page_size)?;
3320
3321 Ok(McpResponse::ListPrompts(ListPromptsResult {
3322 prompts,
3323 next_cursor,
3324 ttl_ms: self.inner.list_ttl_ms,
3325 cache_scope: self.effective_cache_scope(self.inner.list_ttl_ms),
3326 meta: None,
3327 }))
3328 }
3329
3330 McpRequest::GetPrompt(params) => {
3331 #[cfg(feature = "dynamic-tools")]
3332 if let Some(initializer) = &self.inner.prompt_initializer {
3333 initializer()?;
3334 }
3335 if self
3337 .inner
3338 .disabled_prompts
3339 .read()
3340 .unwrap()
3341 .contains(¶ms.name)
3342 {
3343 return Err(Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3344 "Prompt not found: {}",
3345 params.name
3346 ))));
3347 }
3348
3349 let prompt = self.inner.prompts.get(¶ms.name).cloned();
3351 #[cfg(feature = "dynamic-tools")]
3352 let prompt = prompt.or_else(|| {
3353 self.inner
3354 .dynamic_prompts
3355 .as_ref()
3356 .and_then(|d| d.get(¶ms.name))
3357 });
3358 let prompt = prompt.ok_or_else(|| {
3359 Error::JsonRpc(JsonRpcError::method_not_found(&format!(
3360 "Prompt not found: {}",
3361 params.name
3362 )))
3363 })?;
3364
3365 if let Some(filter) = &self.inner.prompt_filter
3367 && !filter.is_visible(&self.session, &prompt)
3368 {
3369 return Err(filter.denial_error(¶ms.name));
3370 }
3371
3372 tracing::debug!(name = %params.name, "Getting prompt");
3373 let ctx = self.create_context_with_extensions(request_id, None, &extensions);
3374 #[cfg(feature = "stateless")]
3375 let ctx = {
3376 let mut ctx = ctx;
3377 ctx.extensions_mut().insert(crate::mrtr::MrtrRequest::new(
3378 params.input_responses,
3379 params.request_state,
3380 ));
3381 ctx
3382 };
3383 let outcome = prompt
3384 .get_outcome_with_context(ctx, params.arguments)
3385 .await?;
3386
3387 match outcome {
3388 RequestOutcome::Complete(result) => Ok(McpResponse::GetPrompt(result)),
3389 RequestOutcome::InputRequired(result) => {
3390 #[cfg(feature = "stateless")]
3391 {
3392 validate_input_required_result(&extensions, &result)?;
3393 Ok(McpResponse::InputRequired(result))
3394 }
3395 #[cfg(not(feature = "stateless"))]
3396 {
3397 let _ = result;
3398 Err(Error::invalid_params(
3399 "InputRequiredResult support was not compiled",
3400 ))
3401 }
3402 }
3403 }
3404 }
3405
3406 McpRequest::Ping => Ok(McpResponse::Pong(EmptyResult {})),
3407
3408 McpRequest::GetTaskInfo(params) => {
3409 if is_final_protocol_request(&extensions) {
3410 self.require_negotiated_tasks(&extensions, "tasks/get")?;
3411 self.authorize_task(¶ms.task_id, &extensions).await?;
3412 return self.final_get_task(¶ms.task_id).await;
3413 }
3414 self.authorize_task(¶ms.task_id, &extensions).await?;
3415
3416 let (mut task, result, error) = self
3423 .inner
3424 .task_store
3425 .get_task_result(¶ms.task_id)
3426 .await
3427 .map_err(task_store_error)?
3428 .ok_or_else(|| {
3429 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3430 "Task not found: {}",
3431 params.task_id
3432 )))
3433 })?;
3434
3435 match task.status {
3436 TaskStatus::Completed => task.result = result,
3437 TaskStatus::Failed => {
3438 task.error = Some(
3442 error.unwrap_or_else(|| JsonRpcError::internal_error("Task failed")),
3443 );
3444 }
3445 _ => {}
3446 }
3447
3448 Ok(McpResponse::GetTaskInfo(task))
3449 }
3450
3451 McpRequest::UpdateTask(params) => {
3452 if is_final_protocol_request(&extensions) {
3453 self.require_negotiated_tasks(&extensions, "tasks/update")?;
3454 self.authorize_task(¶ms.task_id, &extensions).await?;
3455 self.inner
3459 .task_store
3460 .apply_input_responses(
3461 ¶ms.task_id,
3462 decode_input_responses(¶ms.input_responses),
3463 )
3464 .await
3465 .map_err(task_store_error)?
3466 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3467 self.notify_task_state(¶ms.task_id).await;
3471 return Ok(McpResponse::FinalTaskAck(
3472 crate::tasks::TaskAcknowledgement::new(),
3473 ));
3474 }
3475
3476 self.authorize_task(¶ms.task_id, &extensions).await?;
3477
3478 let _ = self
3486 .inner
3487 .task_store
3488 .get_task(¶ms.task_id)
3489 .await
3490 .map_err(task_store_error)?
3491 .ok_or_else(|| {
3492 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3493 "Task not found: {}",
3494 params.task_id
3495 )))
3496 })?;
3497 Ok(McpResponse::UpdateTask(EmptyResult {}))
3498 }
3499
3500 McpRequest::CancelTask(params) => {
3501 if is_final_protocol_request(&extensions) {
3502 self.require_negotiated_tasks(&extensions, "tasks/cancel")?;
3503 self.authorize_task(¶ms.task_id, &extensions).await?;
3504 self.inner
3508 .task_store
3509 .cancel_task(¶ms.task_id, params.reason.as_deref())
3510 .await
3511 .map_err(task_store_error)?
3512 .ok_or_else(|| Error::JsonRpc(unknown_task_error(¶ms.task_id)))?;
3513 self.notify_task_state(¶ms.task_id).await;
3514 return Ok(McpResponse::FinalTaskAck(
3515 crate::tasks::TaskAcknowledgement::new(),
3516 ));
3517 }
3518
3519 self.authorize_task(¶ms.task_id, &extensions).await?;
3520
3521 let current = self
3523 .inner
3524 .task_store
3525 .get_task(¶ms.task_id)
3526 .await
3527 .map_err(task_store_error)?
3528 .ok_or_else(|| {
3529 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3530 "Task not found: {}",
3531 params.task_id
3532 )))
3533 })?;
3534
3535 if current.status.is_terminal() {
3536 return Err(Error::JsonRpc(JsonRpcError::invalid_params(format!(
3537 "Task {} is already in terminal state: {}",
3538 params.task_id, current.status
3539 ))));
3540 }
3541
3542 self.inner
3543 .task_store
3544 .cancel_task(¶ms.task_id, params.reason.as_deref())
3545 .await
3546 .map_err(task_store_error)?
3547 .ok_or_else(|| {
3548 Error::JsonRpc(JsonRpcError::invalid_params(format!(
3549 "Task not found: {}",
3550 params.task_id
3551 )))
3552 })?;
3553
3554 Ok(McpResponse::CancelTask(EmptyResult {}))
3558 }
3559
3560 McpRequest::SetLoggingLevel(params) => {
3561 tracing::debug!(level = ?params.level, "Client set logging level");
3562 if let Ok(mut level) = self.inner.min_log_level.write() {
3563 *level = params.level;
3564 }
3565 Ok(McpResponse::SetLoggingLevel(EmptyResult {}))
3566 }
3567
3568 McpRequest::Complete(params) => {
3569 tracing::debug!(
3570 reference = ?params.reference,
3571 argument = %params.argument.name,
3572 "Completion request"
3573 );
3574
3575 if let Some(ref handler) = self.inner.completion_handler {
3577 let result = handler(params).await?;
3578 Ok(McpResponse::Complete(result))
3579 } else {
3580 Ok(McpResponse::Complete(CompleteResult::new(vec![])))
3582 }
3583 }
3584
3585 McpRequest::Unknown { method, .. } => {
3586 Err(Error::JsonRpc(JsonRpcError::method_not_found(&method)))
3587 }
3588 _ => Err(Error::JsonRpc(JsonRpcError::method_not_found(
3589 "unknown method",
3590 ))),
3591 }
3592 }
3593
3594 pub fn handle_notification(&self, notification: McpNotification) {
3596 match notification {
3597 McpNotification::Initialized => {
3598 let phase_before = self.session.phase();
3599 if self.session.mark_initialized() {
3600 if phase_before == crate::session::SessionPhase::Uninitialized {
3601 tracing::info!(
3602 "Session initialized from uninitialized state (race resolved)"
3603 );
3604 } else {
3605 tracing::info!("Session initialized, entering operation phase");
3606 }
3607 } else {
3608 tracing::warn!(
3609 phase = ?self.session.phase(),
3610 "Received initialized notification in unexpected state"
3611 );
3612 }
3613 }
3614 McpNotification::Cancelled(params) => {
3615 if let Some(ref request_id) = params.request_id {
3616 if self.cancel_request(request_id) {
3617 tracing::info!(
3618 request_id = ?request_id,
3619 reason = ?params.reason,
3620 "Request cancelled"
3621 );
3622 } else {
3623 tracing::debug!(
3624 request_id = ?request_id,
3625 reason = ?params.reason,
3626 "Cancellation requested for unknown request"
3627 );
3628 }
3629 } else {
3630 tracing::debug!(
3631 reason = ?params.reason,
3632 "Cancellation notification received without request_id"
3633 );
3634 }
3635 }
3636 McpNotification::Progress(params) => {
3637 tracing::trace!(
3638 token = ?params.progress_token,
3639 progress = params.progress,
3640 total = ?params.total,
3641 "Progress notification"
3642 );
3643 }
3651 McpNotification::RootsListChanged => {
3652 tracing::info!("Client roots list changed");
3653 }
3656 McpNotification::Unknown { method, .. } => {
3657 tracing::debug!(method = %method, "Unknown notification received");
3658 }
3659 _ => {
3660 tracing::debug!("Unrecognized notification variant received");
3661 }
3662 }
3663 }
3664}
3665
3666impl Default for McpRouter {
3667 fn default() -> Self {
3668 Self::new()
3669 }
3670}
3671
3672pub use crate::context::Extensions;
3678
3679#[derive(Debug, Clone)]
3704pub struct ToolAnnotationsMap {
3705 map: Arc<HashMap<String, ToolAnnotations>>,
3706}
3707
3708impl ToolAnnotationsMap {
3709 pub fn get(&self, tool_name: &str) -> Option<&ToolAnnotations> {
3713 self.map.get(tool_name)
3714 }
3715
3716 pub fn is_read_only(&self, tool_name: &str) -> bool {
3721 self.map.get(tool_name).is_some_and(|a| a.read_only_hint)
3722 }
3723
3724 pub fn is_destructive(&self, tool_name: &str) -> bool {
3729 self.map.get(tool_name).is_none_or(|a| a.destructive_hint)
3730 }
3731
3732 pub fn is_idempotent(&self, tool_name: &str) -> bool {
3737 self.map.get(tool_name).is_some_and(|a| a.idempotent_hint)
3738 }
3739}
3740
3741#[derive(Debug, Clone)]
3763pub struct RouterRequest {
3764 pub id: RequestId,
3766 pub inner: McpRequest,
3768 pub extensions: Extensions,
3770}
3771
3772impl RouterRequest {
3773 pub fn new(id: RequestId, inner: McpRequest) -> Self {
3775 Self {
3776 id,
3777 inner,
3778 extensions: Extensions::new(),
3779 }
3780 }
3781
3782 pub fn with_inner(self, inner: McpRequest) -> Self {
3788 Self {
3789 id: self.id,
3790 inner,
3791 extensions: self.extensions,
3792 }
3793 }
3794
3795 pub fn with_id_and_inner(self, id: RequestId, inner: McpRequest) -> Self {
3801 Self {
3802 id,
3803 inner,
3804 extensions: self.extensions,
3805 }
3806 }
3807
3808 pub fn clone_with_inner(&self, inner: McpRequest) -> Self {
3816 Self {
3817 id: self.id.clone(),
3818 inner,
3819 extensions: self.extensions.clone(),
3820 }
3821 }
3822}
3823
3824#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
3826pub struct RouterResponse {
3827 pub id: RequestId,
3829 pub inner: std::result::Result<McpResponse, JsonRpcError>,
3831}
3832
3833impl RouterResponse {
3834 pub fn is_error(&self) -> bool {
3850 self.inner.is_err()
3851 }
3852
3853 pub fn into_jsonrpc(self) -> JsonRpcResponse {
3855 match self.inner {
3856 Ok(response) => match serde_json::to_value(response) {
3857 Ok(result) => JsonRpcResponse::result(self.id, result),
3858 Err(e) => {
3859 tracing::error!(error = %e, "Failed to serialize response");
3860 JsonRpcResponse::error(
3861 Some(self.id),
3862 JsonRpcError::internal_error(format!("Serialization error: {}", e)),
3863 )
3864 }
3865 },
3866 Err(error) => JsonRpcResponse::error(Some(self.id), error),
3867 }
3868 }
3869}
3870
3871impl Service<RouterRequest> for McpRouter {
3872 type Response = RouterResponse;
3873 type Error = std::convert::Infallible; type Future =
3875 Pin<Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send>>;
3876
3877 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
3878 Poll::Ready(Ok(()))
3879 }
3880
3881 fn call(&mut self, req: RouterRequest) -> Self::Future {
3882 let router = self.clone();
3883 let request_id = req.id.clone();
3884 Box::pin(async move {
3885 let result = router.handle(req.id, req.inner, req.extensions).await;
3886 router.complete_request(&request_id);
3888 Ok(RouterResponse {
3889 id: request_id,
3890 inner: result.map_err(|e| match e {
3895 Error::JsonRpc(err) => err,
3896 Error::Tool(err) => JsonRpcError::internal_error(err.to_string()),
3897 e => JsonRpcError::internal_error(e.to_string()),
3898 }),
3899 })
3900 })
3901 }
3902}
3903
3904#[cfg(test)]
3905mod tests {
3906 use super::*;
3907 use crate::extract::{Context, Json};
3908 use crate::jsonrpc::JsonRpcService;
3909 use crate::tool::ToolBuilder;
3910 use schemars::JsonSchema;
3911 use serde::Deserialize;
3912 use tower::ServiceExt;
3913
3914 #[derive(Debug, Deserialize, JsonSchema)]
3915 struct AddInput {
3916 a: i64,
3917 b: i64,
3918 }
3919
3920 #[cfg(feature = "stateless")]
3921 fn final_extensions(client_capabilities: ClientCapabilities) -> Extensions {
3922 let mut extensions = Extensions::new();
3923 extensions.insert(crate::stateless::StatelessRequestMeta {
3924 protocol_version: Some(PROTOCOL_VERSION_2026_07_28.to_string()),
3925 client_capabilities: Some(client_capabilities),
3926 ..Default::default()
3927 });
3928 extensions
3929 }
3930
3931 #[cfg(feature = "stateless")]
3932 fn tasks_client_extensions() -> Extensions {
3933 final_extensions(ClientCapabilities {
3934 extensions: Some(
3935 [(TASKS_EXTENSION_ID.to_string(), serde_json::json!({}))]
3936 .into_iter()
3937 .collect(),
3938 ),
3939 ..Default::default()
3940 })
3941 }
3942
3943 #[cfg(feature = "stateless")]
3944 #[tokio::test]
3945 async fn final_tasks_require_server_opt_in_and_client_declaration() {
3946 let tool = || {
3947 ToolBuilder::new("optional_task")
3948 .task_support(TaskSupportMode::Optional)
3949 .handler(|input: AddInput| async move {
3950 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
3951 })
3952 .build()
3953 };
3954 let task_params = |task| CallToolParams {
3955 name: "optional_task".to_string(),
3956 arguments: serde_json::json!({"a": 1, "b": 2}),
3957 input_responses: None,
3958 request_state: None,
3959 meta: None,
3960 task,
3961 };
3962
3963 let implicit = McpRouter::new().tool(tool());
3967 let McpResponse::Discover(result) = implicit
3968 .handle(
3969 RequestId::Number(1),
3970 McpRequest::Discover(DiscoverParams::default()),
3971 Extensions::new(),
3972 )
3973 .await
3974 .unwrap()
3975 else {
3976 panic!("Expected Discover response");
3977 };
3978 assert!(
3979 result
3980 .capabilities
3981 .extensions
3982 .as_ref()
3983 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
3984 );
3985 let error = implicit
3986 .handle(
3987 RequestId::Number(2),
3988 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
3989 tasks_client_extensions(),
3990 )
3991 .await
3992 .unwrap_err();
3993 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
3994
3995 let router = McpRouter::new().tool(tool()).with_tasks();
3997 let McpResponse::Discover(result) = router
3998 .handle(
3999 RequestId::Number(3),
4000 McpRequest::Discover(DiscoverParams::default()),
4001 Extensions::new(),
4002 )
4003 .await
4004 .unwrap()
4005 else {
4006 panic!("Expected Discover response");
4007 };
4008 assert!(
4009 result
4010 .capabilities
4011 .extensions
4012 .as_ref()
4013 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID)),
4014 "with_tasks() must advertise the extension on the final path"
4015 );
4016 assert!(
4017 result.capabilities.tasks.is_none(),
4018 "the legacy capability shape is never advertised on the final path"
4019 );
4020
4021 let response = router
4024 .handle(
4025 RequestId::Number(4),
4026 McpRequest::CallTool(task_params(None)),
4027 final_extensions(ClientCapabilities::default()),
4028 )
4029 .await
4030 .unwrap();
4031 assert!(matches!(response, McpResponse::CallTool(_)));
4032
4033 let response = router
4036 .handle(
4037 RequestId::Number(5),
4038 McpRequest::CallTool(task_params(None)),
4039 tasks_client_extensions(),
4040 )
4041 .await
4042 .unwrap();
4043 assert!(
4044 matches!(response, McpResponse::FinalCreateTask(_)),
4045 "a negotiated request must receive a task, got {response:?}"
4046 );
4047
4048 let error = router
4051 .handle(
4052 RequestId::Number(6),
4053 McpRequest::CallTool(task_params(Some(TaskRequestParams { ttl: None }))),
4054 tasks_client_extensions(),
4055 )
4056 .await
4057 .unwrap_err();
4058 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4059 }
4060
4061 #[cfg(feature = "stateless")]
4062 #[tokio::test]
4063 async fn final_task_methods_serve_the_negotiated_wire_shapes() {
4064 let router = McpRouter::new()
4065 .tool(
4066 ToolBuilder::new("optional_task")
4067 .task_support(TaskSupportMode::Optional)
4068 .handler(|input: AddInput| async move {
4069 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4070 })
4071 .task_preparation(|task, _input| async move {
4072 let mut meta = serde_json::Map::new();
4073 meta.insert(
4074 "dev.tower-mcp/owner-test".to_string(),
4075 serde_json::json!({"taskId": task.task_id()}),
4076 );
4077 Ok(crate::TaskPreparation::new().with_meta(meta))
4078 })
4079 .build(),
4080 )
4081 .with_tasks();
4082
4083 let McpResponse::FinalCreateTask(created) = router
4084 .handle(
4085 RequestId::Number(1),
4086 McpRequest::CallTool(CallToolParams {
4087 name: "optional_task".to_string(),
4088 arguments: serde_json::json!({"a": 1, "b": 2}),
4089 input_responses: None,
4090 request_state: None,
4091 meta: None,
4092 task: None,
4093 }),
4094 tasks_client_extensions(),
4095 )
4096 .await
4097 .unwrap()
4098 else {
4099 panic!("Expected a final create-task response");
4100 };
4101
4102 let wire = serde_json::to_value(&created).unwrap();
4104 assert_eq!(wire["resultType"], "task");
4105 assert!(wire.get("task").is_none(), "final results are flat: {wire}");
4106 assert!(wire["ttlMs"].is_number() || wire["ttlMs"].is_null());
4107 assert!(wire.get("ttl").is_none(), "legacy field name leaked");
4108 let task_id = created.task.metadata.task_id.clone();
4109 assert_eq!(
4110 created.meta.as_ref().unwrap()["dev.tower-mcp/owner-test"]["taskId"],
4111 task_id
4112 );
4113
4114 let McpResponse::FinalGetTask(fetched) = router
4116 .handle(
4117 RequestId::Number(2),
4118 McpRequest::GetTaskInfo(GetTaskInfoParams {
4119 task_id: task_id.clone(),
4120 meta: None,
4121 }),
4122 tasks_client_extensions(),
4123 )
4124 .await
4125 .unwrap()
4126 else {
4127 panic!("Expected a final get-task response");
4128 };
4129 let wire = serde_json::to_value(&fetched).unwrap();
4130 assert_eq!(wire["resultType"], "complete");
4131 assert_eq!(wire["taskId"], serde_json::json!(task_id));
4132 assert!(wire["status"].is_string());
4133
4134 for (id, request) in [
4136 (
4137 3,
4138 McpRequest::UpdateTask(UpdateTaskParams {
4139 task_id: task_id.clone(),
4140 input_responses: HashMap::new(),
4141 meta: None,
4142 }),
4143 ),
4144 (
4145 4,
4146 McpRequest::CancelTask(CancelTaskParams {
4147 task_id: task_id.clone(),
4148 reason: None,
4149 meta: None,
4150 }),
4151 ),
4152 ] {
4153 let response = router
4154 .handle(RequestId::Number(id), request, tasks_client_extensions())
4155 .await
4156 .unwrap();
4157 let McpResponse::FinalTaskAck(ack) = response else {
4158 panic!("Expected a final ack for request {id}");
4159 };
4160 assert_eq!(
4161 serde_json::to_value(&ack).unwrap(),
4162 serde_json::json!({"resultType": "complete"})
4163 );
4164 }
4165
4166 let error = router
4168 .handle(
4169 RequestId::Number(5),
4170 McpRequest::GetTaskInfo(GetTaskInfoParams {
4171 task_id: "does-not-exist".to_string(),
4172 meta: None,
4173 }),
4174 tasks_client_extensions(),
4175 )
4176 .await
4177 .unwrap_err();
4178 assert!(matches!(error, Error::JsonRpc(e) if e.code == -32602));
4179
4180 let error = router
4182 .handle(
4183 RequestId::Number(6),
4184 McpRequest::GetTaskInfo(GetTaskInfoParams {
4185 task_id: task_id.clone(),
4186 meta: None,
4187 }),
4188 final_extensions(ClientCapabilities::default()),
4189 )
4190 .await
4191 .unwrap_err();
4192 let Error::JsonRpc(error) = error else {
4193 panic!("expected a JSON-RPC error");
4194 };
4195 assert_eq!(error.code, -32021);
4196 assert_eq!(
4197 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4198 serde_json::json!({})
4199 );
4200 }
4201
4202 #[cfg(feature = "stateless")]
4203 #[tokio::test]
4204 async fn final_required_task_tools_follow_per_request_capabilities() {
4205 let router = McpRouter::new()
4206 .tool(
4207 ToolBuilder::new("required_task")
4208 .task_support(TaskSupportMode::Required)
4209 .handler(|input: AddInput| async move {
4210 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4211 })
4212 .build(),
4213 )
4214 .with_tasks();
4215 let params = || CallToolParams {
4216 name: "required_task".to_string(),
4217 arguments: serde_json::json!({"a": 1, "b": 2}),
4218 input_responses: None,
4219 request_state: None,
4220 meta: None,
4221 task: None,
4222 };
4223
4224 let McpResponse::ListTools(without_tasks) = router
4225 .handle(
4226 RequestId::Number(1),
4227 McpRequest::ListTools(ListToolsParams::default()),
4228 final_extensions(ClientCapabilities::default()),
4229 )
4230 .await
4231 .unwrap()
4232 else {
4233 panic!("expected tools/list")
4234 };
4235 assert!(without_tasks.tools.is_empty());
4236
4237 let McpResponse::ListTools(with_tasks) = router
4238 .handle(
4239 RequestId::Number(2),
4240 McpRequest::ListTools(ListToolsParams::default()),
4241 tasks_client_extensions(),
4242 )
4243 .await
4244 .unwrap()
4245 else {
4246 panic!("expected tools/list")
4247 };
4248 assert_eq!(with_tasks.tools.len(), 1);
4249 assert!(with_tasks.tools[0].execution.is_none());
4250
4251 let error = router
4252 .handle(
4253 RequestId::Number(3),
4254 McpRequest::CallTool(params()),
4255 final_extensions(ClientCapabilities::default()),
4256 )
4257 .await
4258 .unwrap_err();
4259 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32021));
4260
4261 let response = router
4262 .handle(
4263 RequestId::Number(4),
4264 McpRequest::CallTool(params()),
4265 tasks_client_extensions(),
4266 )
4267 .await
4268 .unwrap();
4269 assert!(matches!(response, McpResponse::FinalCreateTask(_)));
4270 }
4271
4272 #[cfg(all(feature = "oauth", feature = "stateless"))]
4273 #[tokio::test]
4274 async fn task_operations_are_bound_to_the_creating_principal() {
4275 fn as_principal(subject: &str) -> Extensions {
4276 let mut extensions = tasks_client_extensions();
4277 extensions.insert(crate::oauth::token::TokenClaims {
4278 sub: Some(subject.to_string()),
4279 iss: None,
4280 aud: None,
4281 exp: None,
4282 scope: None,
4283 client_id: None,
4284 extra: HashMap::new(),
4285 });
4286 extensions
4287 }
4288
4289 let router = McpRouter::new()
4290 .tool(
4291 ToolBuilder::new("optional_task")
4292 .task_support(TaskSupportMode::Optional)
4293 .handler(|input: AddInput| async move {
4294 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4295 })
4296 .build(),
4297 )
4298 .with_tasks();
4299
4300 let McpResponse::FinalCreateTask(created) = router
4301 .handle(
4302 RequestId::Number(1),
4303 McpRequest::CallTool(CallToolParams {
4304 name: "optional_task".to_string(),
4305 arguments: serde_json::json!({"a": 1, "b": 2}),
4306 input_responses: None,
4307 request_state: None,
4308 meta: None,
4309 task: None,
4310 }),
4311 as_principal("alice"),
4312 )
4313 .await
4314 .unwrap()
4315 else {
4316 panic!("Expected a final create-task response");
4317 };
4318 let task_id = created.task.metadata.task_id.clone();
4319
4320 assert!(
4322 router
4323 .handle(
4324 RequestId::Number(2),
4325 McpRequest::GetTaskInfo(GetTaskInfoParams {
4326 task_id: task_id.clone(),
4327 meta: None,
4328 }),
4329 as_principal("alice"),
4330 )
4331 .await
4332 .is_ok()
4333 );
4334
4335 for (id, label, context) in [
4338 (3, "another principal", as_principal("bob")),
4339 (4, "no principal", tasks_client_extensions()),
4340 ] {
4341 for (offset, request) in [
4342 McpRequest::GetTaskInfo(GetTaskInfoParams {
4343 task_id: task_id.clone(),
4344 meta: None,
4345 }),
4346 McpRequest::UpdateTask(UpdateTaskParams {
4347 task_id: task_id.clone(),
4348 input_responses: HashMap::new(),
4349 meta: None,
4350 }),
4351 McpRequest::CancelTask(CancelTaskParams {
4352 task_id: task_id.clone(),
4353 reason: None,
4354 meta: None,
4355 }),
4356 ]
4357 .into_iter()
4358 .enumerate()
4359 {
4360 let error = router
4361 .handle(
4362 RequestId::Number(id * 10 + offset as i64),
4363 request,
4364 context.clone(),
4365 )
4366 .await
4367 .unwrap_err();
4368 assert!(
4369 matches!(error, Error::JsonRpc(ref e) if e.code == -32602),
4370 "{label} was served: {error:?}"
4371 );
4372 let Error::JsonRpc(error) = error else {
4375 unreachable!()
4376 };
4377 assert!(
4378 error.message.contains("not found"),
4379 "refusal leaked that the task exists: {}",
4380 error.message
4381 );
4382 }
4383 }
4384
4385 assert!(
4387 router
4388 .handle(
4389 RequestId::Number(9),
4390 McpRequest::GetTaskInfo(GetTaskInfoParams {
4391 task_id: task_id.clone(),
4392 meta: None,
4393 }),
4394 as_principal("alice"),
4395 )
4396 .await
4397 .is_ok(),
4398 "a refused cancel must not have cancelled the task"
4399 );
4400 }
4401
4402 #[cfg(all(feature = "oauth", feature = "stateless"))]
4403 #[tokio::test]
4404 async fn final_tasks_work_across_independent_routers_with_a_shared_store() {
4405 fn as_principal(subject: &str) -> Extensions {
4406 let mut extensions = tasks_client_extensions();
4407 extensions.insert(crate::oauth::token::TokenClaims {
4408 sub: Some(subject.to_string()),
4409 iss: None,
4410 aud: None,
4411 exp: None,
4412 scope: None,
4413 client_id: None,
4414 extra: HashMap::new(),
4415 });
4416 extensions
4417 }
4418
4419 fn router_with_store(store: Arc<dyn TaskStore>) -> McpRouter {
4420 McpRouter::new()
4421 .tool(
4422 ToolBuilder::new("shared_task")
4423 .task_support(TaskSupportMode::Optional)
4424 .handler(|_input: serde_json::Value| async move {
4425 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
4426 Ok(CallToolResult::text("done"))
4427 })
4428 .build(),
4429 )
4430 .task_store(store)
4431 .with_tasks()
4432 }
4433
4434 let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
4435 let router_a = router_with_store(store.clone());
4436 let router_b = router_with_store(store);
4437
4438 let McpResponse::FinalCreateTask(created) = router_a
4439 .handle(
4440 RequestId::Number(1),
4441 McpRequest::CallTool(CallToolParams {
4442 name: "shared_task".to_string(),
4443 arguments: serde_json::json!({}),
4444 input_responses: None,
4445 request_state: None,
4446 meta: None,
4447 task: None,
4448 }),
4449 as_principal("alice"),
4450 )
4451 .await
4452 .unwrap()
4453 else {
4454 panic!("router A did not create a final task")
4455 };
4456 let task_id = created.task.metadata.task_id;
4457
4458 assert!(
4460 router_b
4461 .handle(
4462 RequestId::Number(2),
4463 McpRequest::GetTaskInfo(GetTaskInfoParams {
4464 task_id: task_id.clone(),
4465 meta: None,
4466 }),
4467 as_principal("alice"),
4468 )
4469 .await
4470 .is_ok()
4471 );
4472
4473 let denied = router_b
4475 .handle(
4476 RequestId::Number(3),
4477 McpRequest::GetTaskInfo(GetTaskInfoParams {
4478 task_id: task_id.clone(),
4479 meta: None,
4480 }),
4481 as_principal("bob"),
4482 )
4483 .await
4484 .unwrap_err();
4485 let unknown = router_b
4486 .handle(
4487 RequestId::Number(4),
4488 McpRequest::GetTaskInfo(GetTaskInfoParams {
4489 task_id: "unknown-task".to_string(),
4490 meta: None,
4491 }),
4492 as_principal("bob"),
4493 )
4494 .await
4495 .unwrap_err();
4496 let (Error::JsonRpc(denied), Error::JsonRpc(unknown)) = (denied, unknown) else {
4497 panic!("expected JSON-RPC task denials")
4498 };
4499 assert_eq!(denied.code, unknown.code);
4500 assert_eq!(
4501 denied.message.replace(&task_id, "<task-id>"),
4502 unknown.message.replace("unknown-task", "<task-id>")
4503 );
4504 assert_eq!(denied.data, unknown.data);
4505
4506 assert!(matches!(
4509 router_b
4510 .handle(
4511 RequestId::Number(5),
4512 McpRequest::CancelTask(CancelTaskParams {
4513 task_id: task_id.clone(),
4514 reason: None,
4515 meta: None,
4516 }),
4517 as_principal("alice"),
4518 )
4519 .await
4520 .unwrap(),
4521 McpResponse::FinalTaskAck(_)
4522 ));
4523 let McpResponse::FinalGetTask(fetched) = router_a
4524 .handle(
4525 RequestId::Number(6),
4526 McpRequest::GetTaskInfo(GetTaskInfoParams {
4527 task_id,
4528 meta: None,
4529 }),
4530 as_principal("alice"),
4531 )
4532 .await
4533 .unwrap()
4534 else {
4535 panic!("router A did not read the shared task")
4536 };
4537 assert_eq!(fetched.task.status(), TaskStatus::Cancelled);
4538 }
4539
4540 #[test]
4541 fn router_advertises_only_locally_declared_protocol_extensions() {
4542 let router = McpRouter::new().with_protocol_extension(
4543 crate::ExtensionDeclaration::new(
4544 "com.example/rendering",
4545 serde_json::json!({"formats": ["html"]}),
4546 )
4547 .unwrap(),
4548 );
4549
4550 let stable = router.capabilities();
4551 let final_capabilities =
4552 router.capabilities_for_protocol(Some(crate::protocol::PROTOCOL_VERSION_2026_07_28));
4553 for capabilities in [stable, final_capabilities] {
4554 let extensions = capabilities.extensions.unwrap();
4555 assert_eq!(extensions.len(), 1);
4556 assert_eq!(extensions["com.example/rendering"]["formats"][0], "html");
4557 assert!(!extensions.contains_key("com.example/client-only"));
4558 }
4559 }
4560
4561 #[tokio::test]
4562 async fn initialize_persists_negotiated_extensions_for_legacy_contexts() {
4563 let router = McpRouter::new().with_protocol_extension(
4564 crate::ExtensionDeclaration::new(
4565 "com.example/shared",
4566 serde_json::json!({"server": true}),
4567 )
4568 .unwrap(),
4569 );
4570 let client_capabilities = ClientCapabilities {
4571 extensions: Some(HashMap::from([
4572 (
4573 "com.example/shared".to_string(),
4574 serde_json::json!({"client": true}),
4575 ),
4576 ("com.example/client-only".to_string(), serde_json::json!({})),
4577 ])),
4578 ..ClientCapabilities::default()
4579 };
4580
4581 router
4582 .handle(
4583 RequestId::Number(1),
4584 McpRequest::Initialize(InitializeParams {
4585 protocol_version: crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
4586 capabilities: client_capabilities,
4587 client_info: Implementation {
4588 name: "extension-test".to_string(),
4589 version: "1.0.0".to_string(),
4590 title: None,
4591 description: None,
4592 icons: None,
4593 website_url: None,
4594 meta: None,
4595 },
4596 meta: None,
4597 }),
4598 Extensions::new(),
4599 )
4600 .await
4601 .unwrap();
4602
4603 let context = router.create_context(RequestId::Number(2), None);
4604 let negotiated = context.negotiated_extensions().unwrap();
4605 assert!(negotiated.contains("com.example/shared"));
4606 assert!(!negotiated.contains("com.example/client-only"));
4607 }
4608
4609 #[cfg(feature = "stateless")]
4610 #[test]
4611 fn final_request_context_exposes_only_negotiated_extensions() {
4612 let router = McpRouter::new().with_protocol_extension(
4613 crate::ExtensionDeclaration::new(
4614 "com.example/shared",
4615 serde_json::json!({"server": true}),
4616 )
4617 .unwrap(),
4618 );
4619 let per_request = final_extensions(ClientCapabilities {
4620 extensions: Some(HashMap::from([
4621 (
4622 "com.example/shared".to_string(),
4623 serde_json::json!({"client": true}),
4624 ),
4625 ("com.example/client-only".to_string(), serde_json::json!({})),
4626 ])),
4627 ..ClientCapabilities::default()
4628 });
4629
4630 let context =
4631 router.create_context_with_extensions(RequestId::Number(1), None, &per_request);
4632 let negotiated = context.negotiated_extensions().unwrap();
4633
4634 assert_eq!(negotiated.len(), 1);
4635 assert_eq!(
4636 negotiated
4637 .get("com.example/shared")
4638 .unwrap()
4639 .client_settings()["client"],
4640 true
4641 );
4642 assert!(!negotiated.contains("com.example/client-only"));
4643 }
4644
4645 #[cfg(feature = "stateless")]
4646 #[tokio::test]
4647 async fn final_protocol_withholds_incomplete_tasks_advertisement() {
4648 let optional = ToolBuilder::new("optional_task")
4649 .task_support(TaskSupportMode::Optional)
4650 .handler(|input: AddInput| async move {
4651 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4652 })
4653 .build();
4654 let required = ToolBuilder::new("required_task")
4655 .task_support(TaskSupportMode::Required)
4656 .handler(|input: AddInput| async move {
4657 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4658 })
4659 .build();
4660 let mut router = McpRouter::new().tool(optional).tool(required);
4661
4662 let stable_capabilities = router.capabilities();
4664 assert!(stable_capabilities.tasks.is_some());
4665 assert!(
4666 stable_capabilities
4667 .extensions
4668 .as_ref()
4669 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID))
4670 );
4671
4672 let response = router
4674 .handle(
4675 RequestId::Number(1),
4676 McpRequest::Discover(DiscoverParams::default()),
4677 Extensions::new(),
4678 )
4679 .await
4680 .unwrap();
4681 let McpResponse::Discover(result) = response else {
4682 panic!("Expected Discover response");
4683 };
4684 assert!(result.capabilities.tasks.is_none());
4685 assert!(
4686 result
4687 .capabilities
4688 .extensions
4689 .as_ref()
4690 .is_none_or(|extensions| !extensions.contains_key(TASKS_EXTENSION_ID))
4691 );
4692
4693 init_router(&mut router).await;
4694
4695 let response = router
4697 .handle(
4698 RequestId::Number(2),
4699 McpRequest::ListTools(ListToolsParams::default()),
4700 Extensions::new(),
4701 )
4702 .await
4703 .unwrap();
4704 let McpResponse::ListTools(result) = response else {
4705 panic!("Expected ListTools response");
4706 };
4707 assert_eq!(result.tools.len(), 2);
4708 assert!(result.tools.iter().all(|tool| tool.execution.is_some()));
4709
4710 let response = router
4713 .handle(
4714 RequestId::Number(3),
4715 McpRequest::ListTools(ListToolsParams::default()),
4716 final_extensions(ClientCapabilities::default()),
4717 )
4718 .await
4719 .unwrap();
4720 let McpResponse::ListTools(result) = response else {
4721 panic!("Expected ListTools response");
4722 };
4723 assert_eq!(result.tools.len(), 1);
4724 assert_eq!(result.tools[0].name, "optional_task");
4725 assert!(result.tools[0].execution.is_none());
4726 }
4727
4728 #[cfg(feature = "stateless")]
4729 #[tokio::test]
4730 async fn final_protocol_enforces_tasks_negotiation() {
4731 let optional = ToolBuilder::new("optional_task")
4732 .task_support(TaskSupportMode::Optional)
4733 .handler(|input: AddInput| async move {
4734 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4735 })
4736 .build();
4737 let required = ToolBuilder::new("required_task")
4738 .task_support(TaskSupportMode::Required)
4739 .handler(|input: AddInput| async move {
4740 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4741 })
4742 .build();
4743 let mut router = McpRouter::new().tool(optional).tool(required).with_tasks();
4744 init_router(&mut router).await;
4745
4746 let response = router
4748 .handle(
4749 RequestId::Number(1),
4750 McpRequest::CallTool(CallToolParams {
4751 name: "optional_task".to_string(),
4752 arguments: serde_json::json!({"a": 1, "b": 2}),
4753 input_responses: None,
4754 request_state: None,
4755 meta: None,
4756 task: None,
4757 }),
4758 final_extensions(ClientCapabilities::default()),
4759 )
4760 .await
4761 .unwrap();
4762 assert!(matches!(response, McpResponse::CallTool(_)));
4763
4764 let error = router
4766 .handle(
4767 RequestId::Number(2),
4768 McpRequest::CallTool(CallToolParams {
4769 name: "optional_task".to_string(),
4770 arguments: serde_json::json!({"a": 1, "b": 2}),
4771 input_responses: None,
4772 request_state: None,
4773 meta: None,
4774 task: Some(TaskRequestParams { ttl: None }),
4775 }),
4776 final_extensions(ClientCapabilities::default()),
4777 )
4778 .await
4779 .unwrap_err();
4780 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32602));
4781
4782 let error = router
4786 .handle(
4787 RequestId::Number(3),
4788 McpRequest::CallTool(CallToolParams {
4789 name: "required_task".to_string(),
4790 arguments: serde_json::json!({"a": 1, "b": 2}),
4791 input_responses: None,
4792 request_state: None,
4793 meta: None,
4794 task: None,
4795 }),
4796 final_extensions(ClientCapabilities::default()),
4797 )
4798 .await
4799 .unwrap_err();
4800 let Error::JsonRpc(error) = error else {
4801 panic!("expected a JSON-RPC error");
4802 };
4803 assert_eq!(error.code, -32021);
4804 assert_eq!(
4805 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4806 serde_json::json!({}),
4807 "the error must name the extension the client needs to declare"
4808 );
4809
4810 let task_requests = [
4811 McpRequest::GetTaskInfo(GetTaskInfoParams {
4812 task_id: "task-unknown".to_string(),
4813 meta: None,
4814 }),
4815 McpRequest::UpdateTask(UpdateTaskParams {
4816 task_id: "task-unknown".to_string(),
4817 input_responses: HashMap::new(),
4818 meta: None,
4819 }),
4820 McpRequest::CancelTask(CancelTaskParams {
4821 task_id: "task-unknown".to_string(),
4822 reason: None,
4823 meta: None,
4824 }),
4825 ];
4826 for (index, request) in task_requests.into_iter().enumerate() {
4827 let error = router
4828 .handle(
4829 RequestId::Number(4 + index as i64),
4830 request,
4831 final_extensions(ClientCapabilities::default()),
4832 )
4833 .await
4834 .unwrap_err();
4835 let Error::JsonRpc(error) = error else {
4836 panic!("expected a JSON-RPC error");
4837 };
4838 assert_eq!(error.code, -32021);
4839 assert_eq!(
4840 error.data.as_ref().unwrap()["requiredCapabilities"]["extensions"]["io.modelcontextprotocol/tasks"],
4841 serde_json::json!({})
4842 );
4843 }
4844
4845 let router_without_tasks = McpRouter::new();
4848 let error = router_without_tasks
4849 .handle(
4850 RequestId::Number(7),
4851 McpRequest::GetTaskInfo(GetTaskInfoParams {
4852 task_id: "task-unknown".to_string(),
4853 meta: None,
4854 }),
4855 final_extensions(tasks_client_capabilities()),
4856 )
4857 .await
4858 .unwrap_err();
4859 assert!(matches!(error, Error::JsonRpc(error) if error.code == -32601));
4860 }
4861
4862 #[cfg(feature = "stateless")]
4863 #[test]
4864 fn input_required_capability_validation_uses_capability_semantics() {
4865 let roots = InputRequiredResult::with_requests(
4866 [(
4867 "roots".to_string(),
4868 InputRequest::ListRoots(ListRootsParams::default()),
4869 )]
4870 .into_iter()
4871 .collect(),
4872 );
4873 let extensions = final_extensions(ClientCapabilities {
4874 roots: Some(RootsCapability {
4875 list_changed: true,
4876 deprecated: None,
4877 }),
4878 ..Default::default()
4879 });
4880 validate_input_required_result(&extensions, &roots).unwrap();
4881 assert!(client_capabilities_satisfy(
4882 extensions
4883 .get::<crate::stateless::StatelessRequestMeta>()
4884 .and_then(|meta| meta.client_capabilities.as_ref())
4885 .unwrap(),
4886 &ClientCapabilities {
4887 roots: Some(RootsCapability::default()),
4888 ..Default::default()
4889 }
4890 ));
4891
4892 let sampling_with_tools = InputRequiredResult::with_requests(
4893 [(
4894 "sample".to_string(),
4895 InputRequest::CreateMessage(CreateMessageParams {
4896 tools: Some(Vec::new()),
4897 ..CreateMessageParams::new(vec![SamplingMessage::user("hello")], 10)
4898 }),
4899 )]
4900 .into_iter()
4901 .collect(),
4902 );
4903 let extensions = final_extensions(ClientCapabilities {
4904 sampling: Some(SamplingCapability::default()),
4905 ..Default::default()
4906 });
4907 assert!(validate_input_required_result(&extensions, &sampling_with_tools).is_err());
4908
4909 let form = InputRequiredResult::with_requests(
4910 [(
4911 "form".to_string(),
4912 InputRequest::Elicit(ElicitRequestParams::Form(ElicitFormParams {
4913 mode: Some(ElicitMode::Form),
4914 message: "name".into(),
4915 requested_schema: ElicitFormSchema::new(),
4916 meta: None,
4917 })),
4918 )]
4919 .into_iter()
4920 .collect(),
4921 );
4922 let extensions = final_extensions(ClientCapabilities {
4923 elicitation: Some(ElicitationCapability::default()),
4924 ..Default::default()
4925 });
4926 validate_input_required_result(&extensions, &form).unwrap();
4927 }
4928
4929 async fn init_router(router: &mut McpRouter) {
4931 let init_req = RouterRequest {
4933 id: RequestId::Number(0),
4934 inner: McpRequest::Initialize(InitializeParams {
4935 protocol_version: "2025-11-25".to_string(),
4936 capabilities: ClientCapabilities {
4937 roots: None,
4938 sampling: None,
4939 elicitation: None,
4940 tasks: None,
4941 experimental: None,
4942 extensions: None,
4943 },
4944 client_info: Implementation {
4945 name: "test".to_string(),
4946 version: "1.0".to_string(),
4947 ..Default::default()
4948 },
4949 meta: None,
4950 }),
4951 extensions: Extensions::new(),
4952 };
4953 let _ = router.ready().await.unwrap().call(init_req).await.unwrap();
4954 router.handle_notification(McpNotification::Initialized);
4956 }
4957
4958 #[tokio::test]
4959 async fn test_router_list_tools() {
4960 let add_tool = ToolBuilder::new("add")
4961 .description("Add two numbers")
4962 .handler(|input: AddInput| async move {
4963 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4964 })
4965 .build();
4966
4967 let mut router = McpRouter::new().tool(add_tool);
4968
4969 init_router(&mut router).await;
4971
4972 let req = RouterRequest {
4973 id: RequestId::Number(1),
4974 inner: McpRequest::ListTools(ListToolsParams::default()),
4975 extensions: Extensions::new(),
4976 };
4977
4978 let resp = router.ready().await.unwrap().call(req).await.unwrap();
4979
4980 match resp.inner {
4981 Ok(McpResponse::ListTools(result)) => {
4982 assert_eq!(result.tools.len(), 1);
4983 assert_eq!(result.tools[0].name, "add");
4984 }
4985 _ => panic!("Expected ListTools response"),
4986 }
4987 }
4988
4989 #[tokio::test]
4990 async fn test_router_call_tool() {
4991 let add_tool = ToolBuilder::new("add")
4992 .description("Add two numbers")
4993 .handler(|input: AddInput| async move {
4994 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
4995 })
4996 .build();
4997
4998 let mut router = McpRouter::new().tool(add_tool);
4999
5000 init_router(&mut router).await;
5002
5003 let req = RouterRequest {
5004 id: RequestId::Number(1),
5005 inner: McpRequest::CallTool(CallToolParams {
5006 input_responses: None,
5007 request_state: None,
5008 name: "add".to_string(),
5009 arguments: serde_json::json!({"a": 2, "b": 3}),
5010 meta: None,
5011 task: None,
5012 }),
5013 extensions: Extensions::new(),
5014 };
5015
5016 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5017
5018 match resp.inner {
5019 Ok(McpResponse::CallTool(result)) => {
5020 assert!(!result.is_error);
5021 match &result.content[0] {
5023 Content::Text { text, .. } => assert_eq!(text, "5"),
5024 _ => panic!("Expected text content"),
5025 }
5026 }
5027 _ => panic!("Expected CallTool response"),
5028 }
5029 }
5030
5031 async fn init_jsonrpc_service(service: &mut JsonRpcService<McpRouter>, router: &McpRouter) {
5033 let init_req = JsonRpcRequest::new(0, "initialize").with_params(serde_json::json!({
5034 "protocolVersion": "2025-11-25",
5035 "capabilities": {},
5036 "clientInfo": { "name": "test", "version": "1.0" }
5037 }));
5038 let _ = service.call_single(init_req).await.unwrap();
5039 router.handle_notification(McpNotification::Initialized);
5040 }
5041
5042 #[tokio::test]
5043 async fn test_jsonrpc_service() {
5044 let add_tool = ToolBuilder::new("add")
5045 .description("Add two numbers")
5046 .handler(|input: AddInput| async move {
5047 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5048 })
5049 .build();
5050
5051 let router = McpRouter::new().tool(add_tool);
5052 let mut service = JsonRpcService::new(router.clone());
5053
5054 init_jsonrpc_service(&mut service, &router).await;
5056
5057 let req = JsonRpcRequest::new(1, "tools/list");
5058
5059 let resp = service.call_single(req).await.unwrap();
5060
5061 match resp {
5062 JsonRpcResponse::Result(r) => {
5063 assert_eq!(r.id, RequestId::Number(1));
5064 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5065 assert_eq!(tools.len(), 1);
5066 }
5067 JsonRpcResponse::Error(_) => panic!("Expected success response"),
5068 _ => panic!("unexpected response variant"),
5069 }
5070 }
5071
5072 #[tokio::test]
5073 async fn test_batch_request() {
5074 let add_tool = ToolBuilder::new("add")
5075 .description("Add two numbers")
5076 .handler(|input: AddInput| async move {
5077 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5078 })
5079 .build();
5080
5081 let router = McpRouter::new().tool(add_tool);
5082 let mut service = JsonRpcService::new(router.clone())
5083 .protocol_versions(["2025-03-26"])
5084 .unwrap();
5085
5086 init_jsonrpc_service(&mut service, &router).await;
5088
5089 let requests = vec![
5091 JsonRpcRequest::new(1, "tools/list"),
5092 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5093 "name": "add",
5094 "arguments": {"a": 10, "b": 20}
5095 })),
5096 JsonRpcRequest::new(3, "ping"),
5097 ];
5098
5099 let responses = service.call_batch(requests).await.unwrap();
5100
5101 assert_eq!(responses.len(), 3);
5102
5103 match &responses[0] {
5105 JsonRpcResponse::Result(r) => {
5106 assert_eq!(r.id, RequestId::Number(1));
5107 let tools = r.result.get("tools").unwrap().as_array().unwrap();
5108 assert_eq!(tools.len(), 1);
5109 }
5110 JsonRpcResponse::Error(_) => panic!("Expected success for tools/list"),
5111 _ => panic!("unexpected response variant"),
5112 }
5113
5114 match &responses[1] {
5116 JsonRpcResponse::Result(r) => {
5117 assert_eq!(r.id, RequestId::Number(2));
5118 let content = r.result.get("content").unwrap().as_array().unwrap();
5119 let text = content[0].get("text").unwrap().as_str().unwrap();
5120 assert_eq!(text, "30");
5121 }
5122 JsonRpcResponse::Error(_) => panic!("Expected success for tools/call"),
5123 _ => panic!("unexpected response variant"),
5124 }
5125
5126 match &responses[2] {
5128 JsonRpcResponse::Result(r) => {
5129 assert_eq!(r.id, RequestId::Number(3));
5130 }
5131 JsonRpcResponse::Error(_) => panic!("Expected success for ping"),
5132 _ => panic!("unexpected response variant"),
5133 }
5134 }
5135
5136 #[tokio::test]
5137 async fn test_empty_batch_error() {
5138 let router = McpRouter::new();
5139 let mut service = JsonRpcService::new(router);
5140
5141 let result = service.call_batch(vec![]).await;
5142 assert!(result.is_err());
5143 }
5144
5145 #[tokio::test]
5150 async fn test_progress_token_extraction() {
5151 use crate::context::{ServerNotification, notification_channel};
5152 use crate::protocol::ProgressToken;
5153 use std::sync::Arc;
5154 use std::sync::atomic::{AtomicBool, Ordering};
5155
5156 let progress_reported = Arc::new(AtomicBool::new(false));
5158 let progress_ref = progress_reported.clone();
5159
5160 let tool = ToolBuilder::new("progress_tool")
5162 .description("Tool that reports progress")
5163 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5164 let reported = progress_ref.clone();
5165 async move {
5166 ctx.report_progress(50.0, Some(100.0), Some("Halfway"))
5168 .await;
5169 reported.store(true, Ordering::SeqCst);
5170 Ok(CallToolResult::text("done"))
5171 }
5172 })
5173 .build();
5174
5175 let (tx, mut rx) = notification_channel(10);
5177 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5178 let mut service = JsonRpcService::new(router.clone());
5179
5180 init_jsonrpc_service(&mut service, &router).await;
5182
5183 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5185 "name": "progress_tool",
5186 "arguments": {"a": 1, "b": 2},
5187 "_meta": {
5188 "progressToken": "test-token-123"
5189 }
5190 }));
5191
5192 let resp = service.call_single(req).await.unwrap();
5193
5194 match resp {
5196 JsonRpcResponse::Result(_) => {}
5197 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
5198 _ => panic!("unexpected response variant"),
5199 }
5200
5201 assert!(progress_reported.load(Ordering::SeqCst));
5203
5204 let notification = rx.try_recv().expect("Expected progress notification");
5206 match notification {
5207 ServerNotification::Progress(params) => {
5208 assert_eq!(
5209 params.progress_token,
5210 ProgressToken::String("test-token-123".to_string())
5211 );
5212 assert_eq!(params.progress, 50.0);
5213 assert_eq!(params.total, Some(100.0));
5214 assert_eq!(params.message.as_deref(), Some("Halfway"));
5215 }
5216 _ => panic!("Expected Progress notification"),
5217 }
5218 }
5219
5220 #[tokio::test]
5221 async fn test_tool_call_without_progress_token() {
5222 use crate::context::notification_channel;
5223 use std::sync::Arc;
5224 use std::sync::atomic::{AtomicBool, Ordering};
5225
5226 let progress_attempted = Arc::new(AtomicBool::new(false));
5227 let progress_ref = progress_attempted.clone();
5228
5229 let tool = ToolBuilder::new("no_token_tool")
5230 .description("Tool that tries to report progress without token")
5231 .extractor_handler((), move |ctx: Context, Json(_input): Json<AddInput>| {
5232 let attempted = progress_ref.clone();
5233 async move {
5234 ctx.report_progress(50.0, Some(100.0), None).await;
5236 attempted.store(true, Ordering::SeqCst);
5237 Ok(CallToolResult::text("done"))
5238 }
5239 })
5240 .build();
5241
5242 let (tx, mut rx) = notification_channel(10);
5243 let router = McpRouter::new().with_notification_sender(tx).tool(tool);
5244 let mut service = JsonRpcService::new(router.clone());
5245
5246 init_jsonrpc_service(&mut service, &router).await;
5247
5248 let req = JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5250 "name": "no_token_tool",
5251 "arguments": {"a": 1, "b": 2}
5252 }));
5253
5254 let resp = service.call_single(req).await.unwrap();
5255 assert!(matches!(resp, JsonRpcResponse::Result(_)));
5256
5257 assert!(progress_attempted.load(Ordering::SeqCst));
5259
5260 assert!(rx.try_recv().is_err());
5262 }
5263
5264 #[tokio::test]
5265 async fn test_batch_errors_returned_not_dropped() {
5266 let add_tool = ToolBuilder::new("add")
5267 .description("Add two numbers")
5268 .handler(|input: AddInput| async move {
5269 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5270 })
5271 .build();
5272
5273 let router = McpRouter::new().tool(add_tool);
5274 let mut service = JsonRpcService::new(router.clone())
5275 .protocol_versions(["2025-03-26"])
5276 .unwrap();
5277
5278 init_jsonrpc_service(&mut service, &router).await;
5279
5280 let requests = vec![
5282 JsonRpcRequest::new(1, "tools/call").with_params(serde_json::json!({
5284 "name": "add",
5285 "arguments": {"a": 10, "b": 20}
5286 })),
5287 JsonRpcRequest::new(2, "tools/call").with_params(serde_json::json!({
5289 "name": "nonexistent_tool",
5290 "arguments": {}
5291 })),
5292 JsonRpcRequest::new(3, "ping"),
5294 ];
5295
5296 let responses = service.call_batch(requests).await.unwrap();
5297
5298 assert_eq!(responses.len(), 3);
5300
5301 match &responses[0] {
5303 JsonRpcResponse::Result(r) => {
5304 assert_eq!(r.id, RequestId::Number(1));
5305 }
5306 JsonRpcResponse::Error(_) => panic!("Expected success for first request"),
5307 _ => panic!("unexpected response variant"),
5308 }
5309
5310 match &responses[1] {
5312 JsonRpcResponse::Error(e) => {
5313 assert_eq!(e.id, Some(RequestId::Number(2)));
5314 assert!(e.error.message.contains("not found") || e.error.code == -32601);
5316 }
5317 JsonRpcResponse::Result(_) => panic!("Expected error for second request"),
5318 _ => panic!("unexpected response variant"),
5319 }
5320
5321 match &responses[2] {
5323 JsonRpcResponse::Result(r) => {
5324 assert_eq!(r.id, RequestId::Number(3));
5325 }
5326 JsonRpcResponse::Error(_) => panic!("Expected success for third request"),
5327 _ => panic!("unexpected response variant"),
5328 }
5329 }
5330
5331 #[tokio::test]
5336 async fn test_list_resource_templates() {
5337 use crate::resource::ResourceTemplateBuilder;
5338 use std::collections::HashMap;
5339
5340 let template = ResourceTemplateBuilder::new("file:///{path}")
5341 .name("Project Files")
5342 .description("Access project files")
5343 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5344 Ok(ReadResourceResult {
5345 contents: vec![ResourceContent {
5346 uri,
5347 mime_type: None,
5348 text: None,
5349 blob: None,
5350 meta: None,
5351 }],
5352 meta: None,
5353 ..Default::default()
5354 })
5355 });
5356
5357 let mut router = McpRouter::new().resource_template(template);
5358
5359 init_router(&mut router).await;
5361
5362 let req = RouterRequest {
5363 id: RequestId::Number(1),
5364 inner: McpRequest::ListResourceTemplates(ListResourceTemplatesParams::default()),
5365 extensions: Extensions::new(),
5366 };
5367
5368 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5369
5370 match resp.inner {
5371 Ok(McpResponse::ListResourceTemplates(result)) => {
5372 assert_eq!(result.resource_templates.len(), 1);
5373 assert_eq!(result.resource_templates[0].uri_template, "file:///{path}");
5374 assert_eq!(result.resource_templates[0].name, "Project Files");
5375 }
5376 _ => panic!("Expected ListResourceTemplates response"),
5377 }
5378 }
5379
5380 #[tokio::test]
5381 async fn test_read_resource_via_template() {
5382 use crate::resource::ResourceTemplateBuilder;
5383 use std::collections::HashMap;
5384
5385 let template = ResourceTemplateBuilder::new("db://users/{id}")
5386 .name("User Records")
5387 .handler(|uri: String, vars: HashMap<String, String>| async move {
5388 let id = vars.get("id").unwrap().clone();
5389 Ok(ReadResourceResult {
5390 contents: vec![ResourceContent {
5391 uri,
5392 mime_type: Some("application/json".to_string()),
5393 text: Some(format!(r#"{{"id": "{}"}}"#, id)),
5394 blob: None,
5395 meta: None,
5396 }],
5397 meta: None,
5398 ..Default::default()
5399 })
5400 });
5401
5402 let mut router = McpRouter::new().resource_template(template);
5403
5404 init_router(&mut router).await;
5406
5407 let req = RouterRequest {
5409 id: RequestId::Number(1),
5410 inner: McpRequest::ReadResource(ReadResourceParams {
5411 input_responses: None,
5412 request_state: None,
5413 uri: "db://users/123".to_string(),
5414 meta: None,
5415 }),
5416 extensions: Extensions::new(),
5417 };
5418
5419 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5420
5421 match resp.inner {
5422 Ok(McpResponse::ReadResource(result)) => {
5423 assert_eq!(result.contents.len(), 1);
5424 assert_eq!(result.contents[0].uri, "db://users/123");
5425 assert!(result.contents[0].text.as_ref().unwrap().contains("123"));
5426 }
5427 _ => panic!("Expected ReadResource response"),
5428 }
5429 }
5430
5431 #[tokio::test]
5432 async fn test_static_resource_takes_precedence_over_template() {
5433 use crate::resource::{ResourceBuilder, ResourceTemplateBuilder};
5434 use std::collections::HashMap;
5435
5436 let template = ResourceTemplateBuilder::new("file:///{path}")
5438 .name("Files Template")
5439 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5440 Ok(ReadResourceResult {
5441 contents: vec![ResourceContent {
5442 uri,
5443 mime_type: None,
5444 text: Some("from template".to_string()),
5445 blob: None,
5446 meta: None,
5447 }],
5448 meta: None,
5449 ..Default::default()
5450 })
5451 });
5452
5453 let static_resource = ResourceBuilder::new("file:///README.md")
5455 .name("README")
5456 .text("from static resource");
5457
5458 let mut router = McpRouter::new()
5459 .resource_template(template)
5460 .resource(static_resource);
5461
5462 init_router(&mut router).await;
5464
5465 let req = RouterRequest {
5467 id: RequestId::Number(1),
5468 inner: McpRequest::ReadResource(ReadResourceParams {
5469 input_responses: None,
5470 request_state: None,
5471 uri: "file:///README.md".to_string(),
5472 meta: None,
5473 }),
5474 extensions: Extensions::new(),
5475 };
5476
5477 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5478
5479 match resp.inner {
5480 Ok(McpResponse::ReadResource(result)) => {
5481 assert_eq!(
5483 result.contents[0].text.as_deref(),
5484 Some("from static resource")
5485 );
5486 }
5487 _ => panic!("Expected ReadResource response"),
5488 }
5489 }
5490
5491 #[tokio::test]
5492 async fn test_resource_not_found_when_no_match() {
5493 use crate::resource::ResourceTemplateBuilder;
5494 use std::collections::HashMap;
5495
5496 let template = ResourceTemplateBuilder::new("db://users/{id}")
5497 .name("Users")
5498 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5499 Ok(ReadResourceResult {
5500 contents: vec![ResourceContent {
5501 uri,
5502 mime_type: None,
5503 text: None,
5504 blob: None,
5505 meta: None,
5506 }],
5507 meta: None,
5508 ..Default::default()
5509 })
5510 });
5511
5512 let mut router = McpRouter::new().resource_template(template);
5513
5514 init_router(&mut router).await;
5516
5517 let req = RouterRequest {
5519 id: RequestId::Number(1),
5520 inner: McpRequest::ReadResource(ReadResourceParams {
5521 input_responses: None,
5522 request_state: None,
5523 uri: "db://posts/123".to_string(),
5524 meta: None,
5525 }),
5526 extensions: Extensions::new(),
5527 };
5528
5529 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5530
5531 match resp.inner {
5532 Err(err) => {
5533 assert!(err.message.contains("not found"));
5534 }
5535 Ok(_) => panic!("Expected error for non-matching URI"),
5536 }
5537 }
5538
5539 #[tokio::test]
5540 async fn test_capabilities_include_resources_with_only_templates() {
5541 use crate::resource::ResourceTemplateBuilder;
5542 use std::collections::HashMap;
5543
5544 let template = ResourceTemplateBuilder::new("file:///{path}")
5545 .name("Files")
5546 .handler(|uri: String, _vars: HashMap<String, String>| async move {
5547 Ok(ReadResourceResult {
5548 contents: vec![ResourceContent {
5549 uri,
5550 mime_type: None,
5551 text: None,
5552 blob: None,
5553 meta: None,
5554 }],
5555 meta: None,
5556 ..Default::default()
5557 })
5558 });
5559
5560 let mut router = McpRouter::new().resource_template(template);
5561
5562 let init_req = RouterRequest {
5564 id: RequestId::Number(0),
5565 inner: McpRequest::Initialize(InitializeParams {
5566 protocol_version: "2025-11-25".to_string(),
5567 capabilities: ClientCapabilities {
5568 roots: None,
5569 sampling: None,
5570 elicitation: None,
5571 tasks: None,
5572 experimental: None,
5573 extensions: None,
5574 },
5575 client_info: Implementation {
5576 name: "test".to_string(),
5577 version: "1.0".to_string(),
5578 ..Default::default()
5579 },
5580 meta: None,
5581 }),
5582 extensions: Extensions::new(),
5583 };
5584 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5585
5586 match resp.inner {
5587 Ok(McpResponse::Initialize(result)) => {
5588 assert!(result.capabilities.resources.is_some());
5590 }
5591 _ => panic!("Expected Initialize response"),
5592 }
5593 }
5594
5595 #[tokio::test]
5600 async fn test_log_sends_notification() {
5601 use crate::context::notification_channel;
5602
5603 let (tx, mut rx) = notification_channel(10);
5604 let router = McpRouter::new().with_notification_sender(tx);
5605
5606 let sent = router.log_info("Test message");
5608 assert!(sent);
5609
5610 let notification = rx.try_recv().unwrap();
5612 match notification {
5613 ServerNotification::LogMessage(params) => {
5614 assert_eq!(params.level, LogLevel::Info);
5615 let data = params.data;
5616 assert_eq!(
5617 data.get("message").unwrap().as_str().unwrap(),
5618 "Test message"
5619 );
5620 }
5621 _ => panic!("Expected LogMessage notification"),
5622 }
5623 }
5624
5625 #[tokio::test]
5626 async fn test_log_with_custom_params() {
5627 use crate::context::notification_channel;
5628
5629 let (tx, mut rx) = notification_channel(10);
5630 let router = McpRouter::new().with_notification_sender(tx);
5631
5632 let params = LoggingMessageParams::new(
5634 LogLevel::Error,
5635 serde_json::json!({
5636 "error": "Connection failed",
5637 "host": "localhost"
5638 }),
5639 )
5640 .with_logger("database");
5641
5642 let sent = router.log(params);
5643 assert!(sent);
5644
5645 let notification = rx.try_recv().unwrap();
5646 match notification {
5647 ServerNotification::LogMessage(params) => {
5648 assert_eq!(params.level, LogLevel::Error);
5649 assert_eq!(params.logger.as_deref(), Some("database"));
5650 let data = params.data;
5651 assert_eq!(
5652 data.get("error").unwrap().as_str().unwrap(),
5653 "Connection failed"
5654 );
5655 }
5656 _ => panic!("Expected LogMessage notification"),
5657 }
5658 }
5659
5660 #[tokio::test]
5661 async fn test_log_without_channel_returns_false() {
5662 let router = McpRouter::new();
5664
5665 assert!(!router.log_info("Test"));
5667 assert!(!router.log_warning("Test"));
5668 assert!(!router.log_error("Test"));
5669 assert!(!router.log_debug("Test"));
5670 }
5671
5672 #[tokio::test]
5673 async fn test_logging_capability_with_channel() {
5674 use crate::context::notification_channel;
5675
5676 let (tx, _rx) = notification_channel(10);
5677 let mut router = McpRouter::new().with_notification_sender(tx);
5678
5679 let init_req = RouterRequest {
5681 id: RequestId::Number(0),
5682 inner: McpRequest::Initialize(InitializeParams {
5683 protocol_version: "2025-11-25".to_string(),
5684 capabilities: ClientCapabilities {
5685 roots: None,
5686 sampling: None,
5687 elicitation: None,
5688 tasks: None,
5689 experimental: None,
5690 extensions: None,
5691 },
5692 client_info: Implementation {
5693 name: "test".to_string(),
5694 version: "1.0".to_string(),
5695 ..Default::default()
5696 },
5697 meta: None,
5698 }),
5699 extensions: Extensions::new(),
5700 };
5701 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5702
5703 match resp.inner {
5704 Ok(McpResponse::Initialize(result)) => {
5705 assert!(result.capabilities.logging.is_some());
5707 }
5708 _ => panic!("Expected Initialize response"),
5709 }
5710 }
5711
5712 #[tokio::test]
5713 async fn test_no_logging_capability_without_channel() {
5714 let mut router = McpRouter::new();
5715
5716 let init_req = RouterRequest {
5718 id: RequestId::Number(0),
5719 inner: McpRequest::Initialize(InitializeParams {
5720 protocol_version: "2025-11-25".to_string(),
5721 capabilities: ClientCapabilities {
5722 roots: None,
5723 sampling: None,
5724 elicitation: None,
5725 tasks: None,
5726 experimental: None,
5727 extensions: None,
5728 },
5729 client_info: Implementation {
5730 name: "test".to_string(),
5731 version: "1.0".to_string(),
5732 ..Default::default()
5733 },
5734 meta: None,
5735 }),
5736 extensions: Extensions::new(),
5737 };
5738 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
5739
5740 match resp.inner {
5741 Ok(McpResponse::Initialize(result)) => {
5742 assert!(result.capabilities.logging.is_none());
5744 }
5745 _ => panic!("Expected Initialize response"),
5746 }
5747 }
5748
5749 #[tokio::test]
5754 async fn test_create_task_via_call_tool() {
5755 let add_tool = ToolBuilder::new("add")
5756 .description("Add two numbers")
5757 .task_support(TaskSupportMode::Optional)
5758 .handler(|input: AddInput| async move {
5759 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5760 })
5761 .build();
5762
5763 let mut router = McpRouter::new().tool(add_tool);
5764 init_router(&mut router).await;
5765
5766 let req = RouterRequest {
5767 id: RequestId::Number(1),
5768 inner: McpRequest::CallTool(CallToolParams {
5769 input_responses: None,
5770 request_state: None,
5771 name: "add".to_string(),
5772 arguments: serde_json::json!({"a": 5, "b": 10}),
5773 meta: None,
5774 task: Some(TaskRequestParams { ttl: None }),
5775 }),
5776 extensions: Extensions::new(),
5777 };
5778
5779 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5780
5781 match resp.inner {
5782 Ok(McpResponse::CreateTask(result)) => {
5783 assert!(!result.task.task_id.is_empty());
5784 assert_eq!(result.task.status, TaskStatus::Working);
5785 }
5786 _ => panic!("Expected CreateTask response"),
5787 }
5788 }
5789
5790 struct CountingTaskStore {
5793 inner: MemoryTaskStore,
5794 creates: std::sync::atomic::AtomicUsize,
5795 gets: std::sync::atomic::AtomicUsize,
5796 completes: std::sync::atomic::AtomicUsize,
5797 }
5798
5799 impl CountingTaskStore {
5800 fn new() -> Self {
5801 Self {
5802 inner: MemoryTaskStore::new(),
5803 creates: std::sync::atomic::AtomicUsize::new(0),
5804 gets: std::sync::atomic::AtomicUsize::new(0),
5805 completes: std::sync::atomic::AtomicUsize::new(0),
5806 }
5807 }
5808 }
5809
5810 #[async_trait::async_trait]
5811 impl TaskStore for CountingTaskStore {
5812 async fn create_task(
5813 &self,
5814 tool_name: &str,
5815 arguments: serde_json::Value,
5816 ttl: Option<u64>,
5817 owner: crate::async_task::TaskOwner,
5818 ) -> crate::async_task::Result<(String, crate::async_task::CancellationToken)> {
5819 self.creates
5820 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5821 self.inner
5822 .create_task(tool_name, arguments, ttl, owner)
5823 .await
5824 }
5825
5826 async fn get_task(&self, task_id: &str) -> crate::async_task::Result<Option<TaskObject>> {
5827 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5828 self.inner.get_task(task_id).await
5829 }
5830
5831 async fn task_owner(
5832 &self,
5833 task_id: &str,
5834 ) -> crate::async_task::Result<Option<crate::async_task::TaskOwner>> {
5835 self.inner.task_owner(task_id).await
5836 }
5837
5838 async fn get_task_result(
5839 &self,
5840 task_id: &str,
5841 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
5842 self.gets.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5845 self.inner.get_task_result(task_id).await
5846 }
5847
5848 async fn wait_for_completion(
5849 &self,
5850 task_id: &str,
5851 ) -> crate::async_task::Result<Option<crate::async_task::TaskSnapshot>> {
5852 self.inner.wait_for_completion(task_id).await
5853 }
5854
5855 async fn list_tasks(
5856 &self,
5857 status_filter: Option<TaskStatus>,
5858 ) -> crate::async_task::Result<Vec<TaskObject>> {
5859 self.inner.list_tasks(status_filter).await
5860 }
5861
5862 async fn require_input(
5863 &self,
5864 task_id: &str,
5865 requests: crate::protocol::InputRequests,
5866 message: Option<&str>,
5867 ) -> crate::async_task::Result<bool> {
5868 self.inner.require_input(task_id, requests, message).await
5869 }
5870
5871 async fn outstanding_input_requests(
5872 &self,
5873 task_id: &str,
5874 ) -> crate::async_task::Result<Option<crate::protocol::InputRequests>> {
5875 self.inner.outstanding_input_requests(task_id).await
5876 }
5877
5878 async fn apply_input_responses(
5879 &self,
5880 task_id: &str,
5881 responses: crate::protocol::InputResponses,
5882 ) -> crate::async_task::Result<Option<crate::async_task::AppliedInputResponses>> {
5883 self.inner.apply_input_responses(task_id, responses).await
5884 }
5885
5886 async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> crate::async_task::Result<bool> {
5887 self.inner.set_ttl(task_id, ttl_ms).await
5888 }
5889
5890 async fn complete_task(
5891 &self,
5892 task_id: &str,
5893 result: CallToolResult,
5894 ) -> crate::async_task::Result<bool> {
5895 self.completes
5896 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5897 self.inner.complete_task(task_id, result).await
5898 }
5899
5900 async fn fail_task(
5901 &self,
5902 task_id: &str,
5903 error: JsonRpcError,
5904 ) -> crate::async_task::Result<bool> {
5905 self.inner.fail_task(task_id, error).await
5906 }
5907
5908 async fn cancel_task(
5909 &self,
5910 task_id: &str,
5911 reason: Option<&str>,
5912 ) -> crate::async_task::Result<Option<TaskObject>> {
5913 self.inner.cancel_task(task_id, reason).await
5914 }
5915 }
5916
5917 #[tokio::test]
5918 async fn test_injected_task_store_used_by_dispatch() {
5919 let store = Arc::new(CountingTaskStore::new());
5920
5921 let add_tool = ToolBuilder::new("add")
5922 .description("Add two numbers")
5923 .task_support(TaskSupportMode::Optional)
5924 .handler(|input: AddInput| async move {
5925 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
5926 })
5927 .build();
5928
5929 let mut router = McpRouter::new()
5930 .tool(add_tool)
5931 .task_store(store.clone() as Arc<dyn TaskStore>);
5932 init_router(&mut router).await;
5933
5934 let req = RouterRequest {
5936 id: RequestId::Number(1),
5937 inner: McpRequest::CallTool(CallToolParams {
5938 input_responses: None,
5939 request_state: None,
5940 name: "add".to_string(),
5941 arguments: serde_json::json!({"a": 2, "b": 3}),
5942 meta: None,
5943 task: Some(TaskRequestParams { ttl: None }),
5944 }),
5945 extensions: Extensions::new(),
5946 };
5947 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5948 let task_id = match resp.inner {
5949 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
5950 other => panic!("Expected CreateTask response, got {other:?}"),
5951 };
5952
5953 assert_eq!(
5954 store.creates.load(std::sync::atomic::Ordering::Relaxed),
5955 1,
5956 "create_task must go through the injected store"
5957 );
5958
5959 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
5961 assert_eq!(
5962 store.completes.load(std::sync::atomic::Ordering::Relaxed),
5963 1,
5964 "complete_task must go through the injected store"
5965 );
5966
5967 let gets_before = store.gets.load(std::sync::atomic::Ordering::Relaxed);
5969 let req = RouterRequest {
5970 id: RequestId::Number(2),
5971 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
5972 task_id: task_id.clone(),
5973 meta: None,
5974 }),
5975 extensions: Extensions::new(),
5976 };
5977 let resp = router.ready().await.unwrap().call(req).await.unwrap();
5978 match resp.inner {
5979 Ok(McpResponse::GetTaskInfo(info)) => {
5980 assert_eq!(info.task_id, task_id);
5981 assert_eq!(info.status, TaskStatus::Completed);
5982 }
5983 other => panic!("Expected GetTaskInfo response, got {other:?}"),
5984 }
5985 assert!(
5986 store.gets.load(std::sync::atomic::Ordering::Relaxed) > gets_before,
5987 "tasks/get must go through the injected store"
5988 );
5989 }
5990
5991 #[tokio::test]
5992 async fn test_removed_tasks_methods_get_method_not_found() {
5993 let mut router = McpRouter::new();
5997 init_router(&mut router).await;
5998
5999 for method in ["tasks/list", "tasks/result"] {
6000 let req = RouterRequest {
6001 id: RequestId::Number(1),
6002 inner: McpRequest::Unknown {
6003 method: method.to_string(),
6004 params: None,
6005 },
6006 extensions: Extensions::new(),
6007 };
6008
6009 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6010
6011 match resp.inner {
6012 Err(err) => {
6013 assert_eq!(err.code, -32601, "{method} must be MethodNotFound");
6014 }
6015 other => panic!("Expected MethodNotFound error for {method}, got {other:?}"),
6016 }
6017 }
6018 }
6019
6020 #[tokio::test]
6021 async fn test_task_lifecycle_complete() {
6022 let add_tool = ToolBuilder::new("add")
6023 .description("Add two numbers")
6024 .task_support(TaskSupportMode::Optional)
6025 .handler(|input: AddInput| async move {
6026 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6027 })
6028 .build();
6029
6030 let mut router = McpRouter::new().tool(add_tool);
6031 init_router(&mut router).await;
6032
6033 let req = RouterRequest {
6035 id: RequestId::Number(1),
6036 inner: McpRequest::CallTool(CallToolParams {
6037 input_responses: None,
6038 request_state: None,
6039 name: "add".to_string(),
6040 arguments: serde_json::json!({"a": 7, "b": 8}),
6041 meta: None,
6042 task: Some(TaskRequestParams { ttl: None }),
6043 }),
6044 extensions: Extensions::new(),
6045 };
6046
6047 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6048 let task_id = match resp.inner {
6049 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6050 _ => panic!("Expected CreateTask response"),
6051 };
6052
6053 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
6055
6056 let req = RouterRequest {
6060 id: RequestId::Number(2),
6061 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6062 task_id: task_id.clone(),
6063 meta: None,
6064 }),
6065 extensions: Extensions::new(),
6066 };
6067
6068 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6069
6070 match resp.inner {
6071 Ok(McpResponse::GetTaskInfo(info)) => {
6072 assert_eq!(info.task_id, task_id);
6073 assert_eq!(info.status, TaskStatus::Completed);
6074 }
6075 _ => panic!("Expected GetTaskInfo response"),
6076 }
6077 }
6078
6079 #[tokio::test]
6080 async fn test_task_cancellation() {
6081 let slow_tool = ToolBuilder::new("slow")
6083 .description("Slow tool")
6084 .task_support(TaskSupportMode::Optional)
6085 .handler(|_input: serde_json::Value| async move {
6086 tokio::time::sleep(tokio::time::Duration::from_secs(60)).await;
6087 Ok(CallToolResult::text("done"))
6088 })
6089 .build();
6090
6091 let mut router = McpRouter::new().tool(slow_tool);
6092 init_router(&mut router).await;
6093
6094 let req = RouterRequest {
6096 id: RequestId::Number(1),
6097 inner: McpRequest::CallTool(CallToolParams {
6098 input_responses: None,
6099 request_state: None,
6100 name: "slow".to_string(),
6101 arguments: serde_json::json!({}),
6102 meta: None,
6103 task: Some(TaskRequestParams { ttl: None }),
6104 }),
6105 extensions: Extensions::new(),
6106 };
6107
6108 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6109 let task_id = match resp.inner {
6110 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6111 _ => panic!("Expected CreateTask response"),
6112 };
6113
6114 let req = RouterRequest {
6116 id: RequestId::Number(2),
6117 inner: McpRequest::CancelTask(CancelTaskParams {
6118 task_id: task_id.clone(),
6119 reason: Some("Test cancellation".to_string()),
6120 meta: None,
6121 }),
6122 extensions: Extensions::new(),
6123 };
6124
6125 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6126
6127 match resp.inner {
6129 Ok(McpResponse::CancelTask(EmptyResult {})) => {}
6130 other => panic!("Expected empty CancelTask ack, got {other:?}"),
6131 }
6132
6133 let req = RouterRequest {
6135 id: RequestId::Number(3),
6136 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6137 task_id: task_id.clone(),
6138 meta: None,
6139 }),
6140 extensions: Extensions::new(),
6141 };
6142 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6143 match resp.inner {
6144 Ok(McpResponse::GetTaskInfo(info)) => {
6145 assert_eq!(info.status, TaskStatus::Cancelled);
6146 }
6147 _ => panic!("Expected GetTaskInfo response"),
6148 }
6149 }
6150
6151 #[tokio::test]
6152 async fn test_get_task_info() {
6153 let add_tool = ToolBuilder::new("add")
6154 .description("Add two numbers")
6155 .task_support(TaskSupportMode::Optional)
6156 .handler(|input: AddInput| async move {
6157 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6158 })
6159 .build();
6160
6161 let mut router = McpRouter::new().tool(add_tool);
6162 init_router(&mut router).await;
6163
6164 let req = RouterRequest {
6166 id: RequestId::Number(1),
6167 inner: McpRequest::CallTool(CallToolParams {
6168 input_responses: None,
6169 request_state: None,
6170 name: "add".to_string(),
6171 arguments: serde_json::json!({"a": 1, "b": 2}),
6172 meta: None,
6173 task: Some(TaskRequestParams { ttl: Some(600_000) }),
6174 }),
6175 extensions: Extensions::new(),
6176 };
6177
6178 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6179 let task_id = match resp.inner {
6180 Ok(McpResponse::CreateTask(result)) => result.task.task_id,
6181 _ => panic!("Expected CreateTask response"),
6182 };
6183
6184 let req = RouterRequest {
6186 id: RequestId::Number(2),
6187 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6188 task_id: task_id.clone(),
6189 meta: None,
6190 }),
6191 extensions: Extensions::new(),
6192 };
6193
6194 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6195
6196 match resp.inner {
6197 Ok(McpResponse::GetTaskInfo(info)) => {
6198 assert_eq!(info.task_id, task_id);
6199 assert!(info.created_at.contains('T')); assert_eq!(info.ttl, Some(600_000));
6201 }
6202 _ => panic!("Expected GetTaskInfo response"),
6203 }
6204 }
6205
6206 #[tokio::test]
6207 async fn test_task_forbidden_tool_rejects_task_params() {
6208 let tool = ToolBuilder::new("sync_only")
6209 .description("Sync only tool")
6210 .handler(|_input: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
6211 .build();
6212
6213 let mut router = McpRouter::new().tool(tool);
6214 init_router(&mut router).await;
6215
6216 let req = RouterRequest {
6218 id: RequestId::Number(1),
6219 inner: McpRequest::CallTool(CallToolParams {
6220 input_responses: None,
6221 request_state: None,
6222 name: "sync_only".to_string(),
6223 arguments: serde_json::json!({}),
6224 meta: None,
6225 task: Some(TaskRequestParams { ttl: None }),
6226 }),
6227 extensions: Extensions::new(),
6228 };
6229
6230 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6231
6232 match resp.inner {
6233 Err(e) => {
6234 assert!(e.message.contains("does not support async tasks"));
6235 }
6236 _ => panic!("Expected error response"),
6237 }
6238 }
6239
6240 #[tokio::test]
6241 async fn test_get_nonexistent_task() {
6242 let mut router = McpRouter::new();
6243 init_router(&mut router).await;
6244
6245 let req = RouterRequest {
6246 id: RequestId::Number(1),
6247 inner: McpRequest::GetTaskInfo(GetTaskInfoParams {
6248 task_id: "task-999".to_string(),
6249 meta: None,
6250 }),
6251 extensions: Extensions::new(),
6252 };
6253
6254 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6255
6256 match resp.inner {
6257 Err(e) => {
6258 assert!(e.message.contains("not found"));
6259 }
6260 _ => panic!("Expected error response"),
6261 }
6262 }
6263
6264 #[tokio::test]
6269 async fn test_subscribe_to_resource() {
6270 use crate::resource::ResourceBuilder;
6271
6272 let resource = ResourceBuilder::new("file:///test.txt")
6273 .name("Test File")
6274 .text("Hello");
6275
6276 let mut router = McpRouter::new().resource(resource);
6277 init_router(&mut router).await;
6278
6279 let req = RouterRequest {
6281 id: RequestId::Number(1),
6282 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6283 uri: "file:///test.txt".to_string(),
6284 meta: None,
6285 }),
6286 extensions: Extensions::new(),
6287 };
6288
6289 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6290
6291 match resp.inner {
6292 Ok(McpResponse::SubscribeResource(_)) => {
6293 assert!(router.is_subscribed("file:///test.txt"));
6295 }
6296 _ => panic!("Expected SubscribeResource response"),
6297 }
6298 }
6299
6300 #[tokio::test]
6301 async fn test_unsubscribe_from_resource() {
6302 use crate::resource::ResourceBuilder;
6303
6304 let resource = ResourceBuilder::new("file:///test.txt")
6305 .name("Test File")
6306 .text("Hello");
6307
6308 let mut router = McpRouter::new().resource(resource);
6309 init_router(&mut router).await;
6310
6311 let req = RouterRequest {
6313 id: RequestId::Number(1),
6314 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6315 uri: "file:///test.txt".to_string(),
6316 meta: None,
6317 }),
6318 extensions: Extensions::new(),
6319 };
6320 let _ = router.ready().await.unwrap().call(req).await.unwrap();
6321 assert!(router.is_subscribed("file:///test.txt"));
6322
6323 let req = RouterRequest {
6325 id: RequestId::Number(2),
6326 inner: McpRequest::UnsubscribeResource(UnsubscribeResourceParams {
6327 uri: "file:///test.txt".to_string(),
6328 meta: None,
6329 }),
6330 extensions: Extensions::new(),
6331 };
6332
6333 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6334
6335 match resp.inner {
6336 Ok(McpResponse::UnsubscribeResource(_)) => {
6337 assert!(!router.is_subscribed("file:///test.txt"));
6339 }
6340 _ => panic!("Expected UnsubscribeResource response"),
6341 }
6342 }
6343
6344 #[tokio::test]
6345 async fn test_subscribe_nonexistent_resource() {
6346 let mut router = McpRouter::new();
6347 init_router(&mut router).await;
6348
6349 let req = RouterRequest {
6350 id: RequestId::Number(1),
6351 inner: McpRequest::SubscribeResource(SubscribeResourceParams {
6352 uri: "file:///nonexistent.txt".to_string(),
6353 meta: None,
6354 }),
6355 extensions: Extensions::new(),
6356 };
6357
6358 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6359
6360 match resp.inner {
6361 Err(e) => {
6362 assert!(e.message.contains("not found"));
6363 }
6364 _ => panic!("Expected error response"),
6365 }
6366 }
6367
6368 #[tokio::test]
6369 async fn test_notify_resource_updated() {
6370 use crate::context::notification_channel;
6371 use crate::resource::ResourceBuilder;
6372
6373 let (tx, mut rx) = notification_channel(10);
6374
6375 let resource = ResourceBuilder::new("file:///test.txt")
6376 .name("Test File")
6377 .text("Hello");
6378
6379 let router = McpRouter::new()
6380 .resource(resource)
6381 .with_notification_sender(tx);
6382
6383 router.subscribe("file:///test.txt");
6385
6386 let sent = router.notify_resource_updated("file:///test.txt");
6388 assert!(sent);
6389
6390 let notification = rx.try_recv().unwrap();
6392 match notification {
6393 ServerNotification::ResourceUpdated { uri } => {
6394 assert_eq!(uri, "file:///test.txt");
6395 }
6396 _ => panic!("Expected ResourceUpdated notification"),
6397 }
6398 }
6399
6400 #[tokio::test]
6401 async fn test_notify_resource_updated_not_subscribed() {
6402 use crate::context::notification_channel;
6403 use crate::resource::ResourceBuilder;
6404
6405 let (tx, mut rx) = notification_channel(10);
6406
6407 let resource = ResourceBuilder::new("file:///test.txt")
6408 .name("Test File")
6409 .text("Hello");
6410
6411 let router = McpRouter::new()
6412 .resource(resource)
6413 .with_notification_sender(tx);
6414
6415 let sent = router.notify_resource_updated("file:///test.txt");
6417 assert!(!sent); assert!(rx.try_recv().is_err());
6421 }
6422
6423 #[tokio::test]
6424 async fn test_notify_resources_list_changed() {
6425 use crate::context::notification_channel;
6426
6427 let (tx, mut rx) = notification_channel(10);
6428 let router = McpRouter::new().with_notification_sender(tx);
6429
6430 let sent = router.notify_resources_list_changed();
6431 assert!(sent);
6432
6433 let notification = rx.try_recv().unwrap();
6434 match notification {
6435 ServerNotification::ResourcesListChanged => {}
6436 _ => panic!("Expected ResourcesListChanged notification"),
6437 }
6438 }
6439
6440 #[tokio::test]
6441 async fn test_subscribed_uris() {
6442 use crate::resource::ResourceBuilder;
6443
6444 let resource1 = ResourceBuilder::new("file:///a.txt").name("A").text("A");
6445
6446 let resource2 = ResourceBuilder::new("file:///b.txt").name("B").text("B");
6447
6448 let router = McpRouter::new().resource(resource1).resource(resource2);
6449
6450 router.subscribe("file:///a.txt");
6452 router.subscribe("file:///b.txt");
6453
6454 let uris = router.subscribed_uris();
6455 assert_eq!(uris.len(), 2);
6456 assert!(uris.contains(&"file:///a.txt".to_string()));
6457 assert!(uris.contains(&"file:///b.txt".to_string()));
6458 }
6459
6460 #[tokio::test]
6461 async fn test_subscription_capability_advertised() {
6462 use crate::resource::ResourceBuilder;
6463
6464 let resource = ResourceBuilder::new("file:///test.txt")
6465 .name("Test")
6466 .text("Hello");
6467
6468 let mut router = McpRouter::new().resource(resource);
6469
6470 let init_req = RouterRequest {
6472 id: RequestId::Number(0),
6473 inner: McpRequest::Initialize(InitializeParams {
6474 protocol_version: "2025-11-25".to_string(),
6475 capabilities: ClientCapabilities {
6476 roots: None,
6477 sampling: None,
6478 elicitation: None,
6479 tasks: None,
6480 experimental: None,
6481 extensions: None,
6482 },
6483 client_info: Implementation {
6484 name: "test".to_string(),
6485 version: "1.0".to_string(),
6486 ..Default::default()
6487 },
6488 meta: None,
6489 }),
6490 extensions: Extensions::new(),
6491 };
6492 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
6493
6494 match resp.inner {
6495 Ok(McpResponse::Initialize(result)) => {
6496 let resources_cap = result.capabilities.resources.unwrap();
6498 assert!(resources_cap.subscribe);
6499 }
6500 _ => panic!("Expected Initialize response"),
6501 }
6502 }
6503
6504 #[tokio::test]
6505 async fn test_completion_handler() {
6506 let router = McpRouter::new()
6507 .server_info("test", "1.0")
6508 .completion_handler(|params: CompleteParams| async move {
6509 let prefix = ¶ms.argument.value;
6511 let suggestions: Vec<String> = vec!["alpha", "beta", "gamma"]
6512 .into_iter()
6513 .filter(|s| s.starts_with(prefix))
6514 .map(String::from)
6515 .collect();
6516 Ok(CompleteResult::new(suggestions))
6517 });
6518
6519 let init_req = RouterRequest {
6521 id: RequestId::Number(0),
6522 inner: McpRequest::Initialize(InitializeParams {
6523 protocol_version: "2025-11-25".to_string(),
6524 capabilities: ClientCapabilities::default(),
6525 client_info: Implementation {
6526 name: "test".to_string(),
6527 version: "1.0".to_string(),
6528 ..Default::default()
6529 },
6530 meta: None,
6531 }),
6532 extensions: Extensions::new(),
6533 };
6534 let resp = router
6535 .clone()
6536 .ready()
6537 .await
6538 .unwrap()
6539 .call(init_req)
6540 .await
6541 .unwrap();
6542
6543 match resp.inner {
6545 Ok(McpResponse::Initialize(result)) => {
6546 assert!(result.capabilities.completions.is_some());
6547 }
6548 _ => panic!("Expected Initialize response"),
6549 }
6550
6551 router.handle_notification(McpNotification::Initialized);
6553
6554 let complete_req = RouterRequest {
6556 id: RequestId::Number(1),
6557 inner: McpRequest::Complete(CompleteParams {
6558 reference: CompletionReference::prompt("test-prompt"),
6559 argument: CompletionArgument::new("query", "al"),
6560 context: None,
6561 meta: None,
6562 }),
6563 extensions: Extensions::new(),
6564 };
6565 let resp = router
6566 .clone()
6567 .ready()
6568 .await
6569 .unwrap()
6570 .call(complete_req)
6571 .await
6572 .unwrap();
6573
6574 match resp.inner {
6575 Ok(McpResponse::Complete(result)) => {
6576 assert_eq!(result.completion.values, vec!["alpha"]);
6577 }
6578 _ => panic!("Expected Complete response"),
6579 }
6580 }
6581
6582 #[tokio::test]
6583 async fn test_completion_without_handler_returns_empty() {
6584 let router = McpRouter::new().server_info("test", "1.0");
6585
6586 let init_req = RouterRequest {
6588 id: RequestId::Number(0),
6589 inner: McpRequest::Initialize(InitializeParams {
6590 protocol_version: "2025-11-25".to_string(),
6591 capabilities: ClientCapabilities::default(),
6592 client_info: Implementation {
6593 name: "test".to_string(),
6594 version: "1.0".to_string(),
6595 ..Default::default()
6596 },
6597 meta: None,
6598 }),
6599 extensions: Extensions::new(),
6600 };
6601 let resp = router
6602 .clone()
6603 .ready()
6604 .await
6605 .unwrap()
6606 .call(init_req)
6607 .await
6608 .unwrap();
6609
6610 match resp.inner {
6612 Ok(McpResponse::Initialize(result)) => {
6613 assert!(result.capabilities.completions.is_none());
6614 }
6615 _ => panic!("Expected Initialize response"),
6616 }
6617
6618 router.handle_notification(McpNotification::Initialized);
6620
6621 let complete_req = RouterRequest {
6623 id: RequestId::Number(1),
6624 inner: McpRequest::Complete(CompleteParams {
6625 reference: CompletionReference::prompt("test-prompt"),
6626 argument: CompletionArgument::new("query", "al"),
6627 context: None,
6628 meta: None,
6629 }),
6630 extensions: Extensions::new(),
6631 };
6632 let resp = router
6633 .clone()
6634 .ready()
6635 .await
6636 .unwrap()
6637 .call(complete_req)
6638 .await
6639 .unwrap();
6640
6641 match resp.inner {
6642 Ok(McpResponse::Complete(result)) => {
6643 assert!(result.completion.values.is_empty());
6644 }
6645 _ => panic!("Expected Complete response"),
6646 }
6647 }
6648
6649 #[tokio::test]
6650 async fn test_tool_filter_list() {
6651 use crate::filter::CapabilityFilter;
6652 use crate::tool::Tool;
6653
6654 let public_tool = ToolBuilder::new("public")
6655 .description("Public tool")
6656 .handler(|_: AddInput| async move { Ok(CallToolResult::text("public")) })
6657 .build();
6658
6659 let admin_tool = ToolBuilder::new("admin")
6660 .description("Admin tool")
6661 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6662 .build();
6663
6664 let mut router = McpRouter::new()
6665 .tool(public_tool)
6666 .tool(admin_tool)
6667 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| tool.name != "admin"));
6668
6669 init_router(&mut router).await;
6671
6672 let req = RouterRequest {
6673 id: RequestId::Number(1),
6674 inner: McpRequest::ListTools(ListToolsParams::default()),
6675 extensions: Extensions::new(),
6676 };
6677
6678 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6679
6680 match resp.inner {
6681 Ok(McpResponse::ListTools(result)) => {
6682 assert_eq!(result.tools.len(), 1);
6684 assert_eq!(result.tools[0].name, "public");
6685 }
6686 _ => panic!("Expected ListTools response"),
6687 }
6688 }
6689
6690 #[tokio::test]
6691 async fn test_tool_filter_call_denied() {
6692 use crate::filter::CapabilityFilter;
6693 use crate::tool::Tool;
6694
6695 let admin_tool = ToolBuilder::new("admin")
6696 .description("Admin tool")
6697 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6698 .build();
6699
6700 let mut router = McpRouter::new()
6701 .tool(admin_tool)
6702 .tool_filter(CapabilityFilter::new(|_, _: &Tool| false)); init_router(&mut router).await;
6706
6707 let req = RouterRequest {
6708 id: RequestId::Number(1),
6709 inner: McpRequest::CallTool(CallToolParams {
6710 input_responses: None,
6711 request_state: None,
6712 name: "admin".to_string(),
6713 arguments: serde_json::json!({"a": 1, "b": 2}),
6714 meta: None,
6715 task: None,
6716 }),
6717 extensions: Extensions::new(),
6718 };
6719
6720 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6721
6722 match resp.inner {
6724 Err(e) => {
6725 assert_eq!(e.code, -32601); }
6727 _ => panic!("Expected JsonRpc error"),
6728 }
6729 }
6730
6731 #[tokio::test]
6732 async fn test_tool_filter_call_allowed() {
6733 use crate::filter::CapabilityFilter;
6734 use crate::tool::Tool;
6735
6736 let public_tool = ToolBuilder::new("public")
6737 .description("Public tool")
6738 .handler(|input: AddInput| async move {
6739 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
6740 })
6741 .build();
6742
6743 let mut router = McpRouter::new()
6744 .tool(public_tool)
6745 .tool_filter(CapabilityFilter::new(|_, _: &Tool| true)); init_router(&mut router).await;
6749
6750 let req = RouterRequest {
6751 id: RequestId::Number(1),
6752 inner: McpRequest::CallTool(CallToolParams {
6753 input_responses: None,
6754 request_state: None,
6755 name: "public".to_string(),
6756 arguments: serde_json::json!({"a": 1, "b": 2}),
6757 meta: None,
6758 task: None,
6759 }),
6760 extensions: Extensions::new(),
6761 };
6762
6763 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6764
6765 match resp.inner {
6766 Ok(McpResponse::CallTool(result)) => {
6767 assert!(!result.is_error);
6768 }
6769 _ => panic!("Expected CallTool response"),
6770 }
6771 }
6772
6773 #[tokio::test]
6774 async fn test_tool_filter_custom_denial() {
6775 use crate::filter::{CapabilityFilter, DenialBehavior};
6776 use crate::tool::Tool;
6777
6778 let admin_tool = ToolBuilder::new("admin")
6779 .description("Admin tool")
6780 .handler(|_: AddInput| async move { Ok(CallToolResult::text("admin")) })
6781 .build();
6782
6783 let mut router = McpRouter::new().tool(admin_tool).tool_filter(
6784 CapabilityFilter::new(|_, _: &Tool| false)
6785 .denial_behavior(DenialBehavior::Unauthorized),
6786 );
6787
6788 init_router(&mut router).await;
6790
6791 let req = RouterRequest {
6792 id: RequestId::Number(1),
6793 inner: McpRequest::CallTool(CallToolParams {
6794 input_responses: None,
6795 request_state: None,
6796 name: "admin".to_string(),
6797 arguments: serde_json::json!({"a": 1, "b": 2}),
6798 meta: None,
6799 task: None,
6800 }),
6801 extensions: Extensions::new(),
6802 };
6803
6804 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6805
6806 match resp.inner {
6808 Err(e) => {
6809 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
6811 }
6812 _ => panic!("Expected JsonRpc error"),
6813 }
6814 }
6815
6816 #[tokio::test]
6817 async fn test_resource_filter_list() {
6818 use crate::filter::CapabilityFilter;
6819 use crate::resource::{Resource, ResourceBuilder};
6820
6821 let public_resource = ResourceBuilder::new("file:///public.txt")
6822 .name("Public File")
6823 .text("public content");
6824
6825 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6826 .name("Secret File")
6827 .text("secret content");
6828
6829 let mut router = McpRouter::new()
6830 .resource(public_resource)
6831 .resource(secret_resource)
6832 .resource_filter(CapabilityFilter::new(|_, r: &Resource| {
6833 !r.name.contains("Secret")
6834 }));
6835
6836 init_router(&mut router).await;
6838
6839 let req = RouterRequest {
6840 id: RequestId::Number(1),
6841 inner: McpRequest::ListResources(ListResourcesParams::default()),
6842 extensions: Extensions::new(),
6843 };
6844
6845 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6846
6847 match resp.inner {
6848 Ok(McpResponse::ListResources(result)) => {
6849 assert_eq!(result.resources.len(), 1);
6851 assert_eq!(result.resources[0].name, "Public File");
6852 }
6853 _ => panic!("Expected ListResources response"),
6854 }
6855 }
6856
6857 #[tokio::test]
6858 async fn test_resource_filter_read_denied() {
6859 use crate::filter::CapabilityFilter;
6860 use crate::resource::{Resource, ResourceBuilder};
6861
6862 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6863 .name("Secret File")
6864 .text("secret content");
6865
6866 let mut router = McpRouter::new()
6867 .resource(secret_resource)
6868 .resource_filter(CapabilityFilter::new(|_, _: &Resource| false)); init_router(&mut router).await;
6872
6873 let req = RouterRequest {
6874 id: RequestId::Number(1),
6875 inner: McpRequest::ReadResource(ReadResourceParams {
6876 input_responses: None,
6877 request_state: None,
6878 uri: "file:///secret.txt".to_string(),
6879 meta: None,
6880 }),
6881 extensions: Extensions::new(),
6882 };
6883
6884 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6885
6886 match resp.inner {
6888 Err(e) => {
6889 assert_eq!(e.code, -32601); }
6891 _ => panic!("Expected JsonRpc error"),
6892 }
6893 }
6894
6895 #[tokio::test]
6896 async fn test_resource_filter_read_allowed() {
6897 use crate::filter::CapabilityFilter;
6898 use crate::resource::{Resource, ResourceBuilder};
6899
6900 let public_resource = ResourceBuilder::new("file:///public.txt")
6901 .name("Public File")
6902 .text("public content");
6903
6904 let mut router = McpRouter::new()
6905 .resource(public_resource)
6906 .resource_filter(CapabilityFilter::new(|_, _: &Resource| true)); init_router(&mut router).await;
6910
6911 let req = RouterRequest {
6912 id: RequestId::Number(1),
6913 inner: McpRequest::ReadResource(ReadResourceParams {
6914 input_responses: None,
6915 request_state: None,
6916 uri: "file:///public.txt".to_string(),
6917 meta: None,
6918 }),
6919 extensions: Extensions::new(),
6920 };
6921
6922 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6923
6924 match resp.inner {
6925 Ok(McpResponse::ReadResource(result)) => {
6926 assert_eq!(result.contents.len(), 1);
6927 assert_eq!(result.contents[0].text.as_deref(), Some("public content"));
6928 }
6929 _ => panic!("Expected ReadResource response"),
6930 }
6931 }
6932
6933 #[tokio::test]
6934 async fn test_resource_filter_custom_denial() {
6935 use crate::filter::{CapabilityFilter, DenialBehavior};
6936 use crate::resource::{Resource, ResourceBuilder};
6937
6938 let secret_resource = ResourceBuilder::new("file:///secret.txt")
6939 .name("Secret File")
6940 .text("secret content");
6941
6942 let mut router = McpRouter::new().resource(secret_resource).resource_filter(
6943 CapabilityFilter::new(|_, _: &Resource| false)
6944 .denial_behavior(DenialBehavior::Unauthorized),
6945 );
6946
6947 init_router(&mut router).await;
6949
6950 let req = RouterRequest {
6951 id: RequestId::Number(1),
6952 inner: McpRequest::ReadResource(ReadResourceParams {
6953 input_responses: None,
6954 request_state: None,
6955 uri: "file:///secret.txt".to_string(),
6956 meta: None,
6957 }),
6958 extensions: Extensions::new(),
6959 };
6960
6961 let resp = router.ready().await.unwrap().call(req).await.unwrap();
6962
6963 match resp.inner {
6965 Err(e) => {
6966 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
6968 }
6969 _ => panic!("Expected JsonRpc error"),
6970 }
6971 }
6972
6973 #[tokio::test]
6974 async fn test_prompt_filter_list() {
6975 use crate::filter::CapabilityFilter;
6976 use crate::prompt::{Prompt, PromptBuilder};
6977
6978 let public_prompt = PromptBuilder::new("greeting")
6979 .description("A greeting")
6980 .user_message("Hello!");
6981
6982 let admin_prompt = PromptBuilder::new("system_debug")
6983 .description("Admin prompt")
6984 .user_message("Debug");
6985
6986 let mut router = McpRouter::new()
6987 .prompt(public_prompt)
6988 .prompt(admin_prompt)
6989 .prompt_filter(CapabilityFilter::new(|_, p: &Prompt| {
6990 !p.name.contains("system")
6991 }));
6992
6993 init_router(&mut router).await;
6995
6996 let req = RouterRequest {
6997 id: RequestId::Number(1),
6998 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
6999 extensions: Extensions::new(),
7000 };
7001
7002 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7003
7004 match resp.inner {
7005 Ok(McpResponse::ListPrompts(result)) => {
7006 assert_eq!(result.prompts.len(), 1);
7008 assert_eq!(result.prompts[0].name, "greeting");
7009 }
7010 _ => panic!("Expected ListPrompts response"),
7011 }
7012 }
7013
7014 #[tokio::test]
7015 async fn test_prompt_filter_get_denied() {
7016 use crate::filter::CapabilityFilter;
7017 use crate::prompt::{Prompt, PromptBuilder};
7018 use std::collections::HashMap;
7019
7020 let admin_prompt = PromptBuilder::new("system_debug")
7021 .description("Admin prompt")
7022 .user_message("Debug");
7023
7024 let mut router = McpRouter::new()
7025 .prompt(admin_prompt)
7026 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| false)); init_router(&mut router).await;
7030
7031 let req = RouterRequest {
7032 id: RequestId::Number(1),
7033 inner: McpRequest::GetPrompt(GetPromptParams {
7034 input_responses: None,
7035 request_state: None,
7036 name: "system_debug".to_string(),
7037 arguments: HashMap::new(),
7038 meta: None,
7039 }),
7040 extensions: Extensions::new(),
7041 };
7042
7043 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7044
7045 match resp.inner {
7047 Err(e) => {
7048 assert_eq!(e.code, -32601); }
7050 _ => panic!("Expected JsonRpc error"),
7051 }
7052 }
7053
7054 #[tokio::test]
7055 async fn test_prompt_filter_get_allowed() {
7056 use crate::filter::CapabilityFilter;
7057 use crate::prompt::{Prompt, PromptBuilder};
7058 use std::collections::HashMap;
7059
7060 let public_prompt = PromptBuilder::new("greeting")
7061 .description("A greeting")
7062 .user_message("Hello!");
7063
7064 let mut router = McpRouter::new()
7065 .prompt(public_prompt)
7066 .prompt_filter(CapabilityFilter::new(|_, _: &Prompt| true)); init_router(&mut router).await;
7070
7071 let req = RouterRequest {
7072 id: RequestId::Number(1),
7073 inner: McpRequest::GetPrompt(GetPromptParams {
7074 input_responses: None,
7075 request_state: None,
7076 name: "greeting".to_string(),
7077 arguments: HashMap::new(),
7078 meta: None,
7079 }),
7080 extensions: Extensions::new(),
7081 };
7082
7083 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7084
7085 match resp.inner {
7086 Ok(McpResponse::GetPrompt(result)) => {
7087 assert_eq!(result.messages.len(), 1);
7088 }
7089 _ => panic!("Expected GetPrompt response"),
7090 }
7091 }
7092
7093 #[tokio::test]
7094 async fn test_prompt_filter_custom_denial() {
7095 use crate::filter::{CapabilityFilter, DenialBehavior};
7096 use crate::prompt::{Prompt, PromptBuilder};
7097 use std::collections::HashMap;
7098
7099 let admin_prompt = PromptBuilder::new("system_debug")
7100 .description("Admin prompt")
7101 .user_message("Debug");
7102
7103 let mut router = McpRouter::new().prompt(admin_prompt).prompt_filter(
7104 CapabilityFilter::new(|_, _: &Prompt| false)
7105 .denial_behavior(DenialBehavior::Unauthorized),
7106 );
7107
7108 init_router(&mut router).await;
7110
7111 let req = RouterRequest {
7112 id: RequestId::Number(1),
7113 inner: McpRequest::GetPrompt(GetPromptParams {
7114 input_responses: None,
7115 request_state: None,
7116 name: "system_debug".to_string(),
7117 arguments: HashMap::new(),
7118 meta: None,
7119 }),
7120 extensions: Extensions::new(),
7121 };
7122
7123 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7124
7125 match resp.inner {
7127 Err(e) => {
7128 assert_eq!(e.code, -32007); assert!(e.message.contains("Unauthorized"));
7130 }
7131 _ => panic!("Expected JsonRpc error"),
7132 }
7133 }
7134
7135 #[derive(Debug, Deserialize, JsonSchema)]
7140 struct StringInput {
7141 value: String,
7142 }
7143
7144 #[tokio::test]
7145 async fn test_router_merge_tools() {
7146 let tool_a = ToolBuilder::new("tool_a")
7148 .description("Tool A")
7149 .handler(|_: StringInput| async move { Ok(CallToolResult::text("A")) })
7150 .build();
7151
7152 let router_a = McpRouter::new().tool(tool_a);
7153
7154 let tool_b = ToolBuilder::new("tool_b")
7156 .description("Tool B")
7157 .handler(|_: StringInput| async move { Ok(CallToolResult::text("B")) })
7158 .build();
7159 let tool_c = ToolBuilder::new("tool_c")
7160 .description("Tool C")
7161 .handler(|_: StringInput| async move { Ok(CallToolResult::text("C")) })
7162 .build();
7163
7164 let router_b = McpRouter::new().tool(tool_b).tool(tool_c);
7165
7166 let mut merged = McpRouter::new()
7168 .server_info("merged", "1.0")
7169 .merge(router_a)
7170 .merge(router_b);
7171
7172 init_router(&mut merged).await;
7173
7174 let req = RouterRequest {
7176 id: RequestId::Number(1),
7177 inner: McpRequest::ListTools(ListToolsParams::default()),
7178 extensions: Extensions::new(),
7179 };
7180
7181 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7182
7183 match resp.inner {
7184 Ok(McpResponse::ListTools(result)) => {
7185 assert_eq!(result.tools.len(), 3);
7186 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7187 assert!(names.contains(&"tool_a"));
7188 assert!(names.contains(&"tool_b"));
7189 assert!(names.contains(&"tool_c"));
7190 }
7191 _ => panic!("Expected ListTools response"),
7192 }
7193 }
7194
7195 #[tokio::test]
7196 async fn test_router_merge_overwrites_duplicates() {
7197 let tool_v1 = ToolBuilder::new("shared")
7199 .description("Version 1")
7200 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v1")) })
7201 .build();
7202
7203 let router_a = McpRouter::new().tool(tool_v1);
7204
7205 let tool_v2 = ToolBuilder::new("shared")
7207 .description("Version 2")
7208 .handler(|_: StringInput| async move { Ok(CallToolResult::text("v2")) })
7209 .build();
7210
7211 let router_b = McpRouter::new().tool(tool_v2);
7212
7213 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7215
7216 init_router(&mut merged).await;
7217
7218 let req = RouterRequest {
7219 id: RequestId::Number(1),
7220 inner: McpRequest::ListTools(ListToolsParams::default()),
7221 extensions: Extensions::new(),
7222 };
7223
7224 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7225
7226 match resp.inner {
7227 Ok(McpResponse::ListTools(result)) => {
7228 assert_eq!(result.tools.len(), 1);
7229 assert_eq!(result.tools[0].name, "shared");
7230 assert_eq!(result.tools[0].description.as_deref(), Some("Version 2"));
7231 }
7232 _ => panic!("Expected ListTools response"),
7233 }
7234 }
7235
7236 #[tokio::test]
7237 async fn test_router_merge_resources() {
7238 use crate::resource::ResourceBuilder;
7239
7240 let router_a = McpRouter::new().resource(
7242 ResourceBuilder::new("file:///a.txt")
7243 .name("File A")
7244 .text("content a"),
7245 );
7246
7247 let router_b = McpRouter::new().resource(
7248 ResourceBuilder::new("file:///b.txt")
7249 .name("File B")
7250 .text("content b"),
7251 );
7252
7253 let mut merged = McpRouter::new().merge(router_a).merge(router_b);
7254
7255 init_router(&mut merged).await;
7256
7257 let req = RouterRequest {
7258 id: RequestId::Number(1),
7259 inner: McpRequest::ListResources(ListResourcesParams::default()),
7260 extensions: Extensions::new(),
7261 };
7262
7263 let resp = merged.ready().await.unwrap().call(req).await.unwrap();
7264
7265 match resp.inner {
7266 Ok(McpResponse::ListResources(result)) => {
7267 assert_eq!(result.resources.len(), 2);
7268 let uris: Vec<&str> = result.resources.iter().map(|r| r.uri.as_str()).collect();
7269 assert!(uris.contains(&"file:///a.txt"));
7270 assert!(uris.contains(&"file:///b.txt"));
7271 }
7272 _ => panic!("Expected ListResources response"),
7273 }
7274 }
7275
7276 #[tokio::test]
7277 async fn test_router_merge_prompts() {
7278 use crate::prompt::PromptBuilder;
7279
7280 let router_a =
7281 McpRouter::new().prompt(PromptBuilder::new("prompt_a").user_message("Hello A"));
7282
7283 let router_b =
7284 McpRouter::new().prompt(PromptBuilder::new("prompt_b").user_message("Hello B"));
7285
7286 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::ListPrompts(ListPromptsParams::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::ListPrompts(result)) => {
7300 assert_eq!(result.prompts.len(), 2);
7301 let names: Vec<&str> = result.prompts.iter().map(|p| p.name.as_str()).collect();
7302 assert!(names.contains(&"prompt_a"));
7303 assert!(names.contains(&"prompt_b"));
7304 }
7305 _ => panic!("Expected ListPrompts response"),
7306 }
7307 }
7308
7309 #[tokio::test]
7310 async fn test_router_nest_prefixes_tools() {
7311 let tool_query = ToolBuilder::new("query")
7313 .description("Query the database")
7314 .handler(|_: StringInput| async move { Ok(CallToolResult::text("query result")) })
7315 .build();
7316 let tool_insert = ToolBuilder::new("insert")
7317 .description("Insert into database")
7318 .handler(|_: StringInput| async move { Ok(CallToolResult::text("insert result")) })
7319 .build();
7320
7321 let db_router = McpRouter::new().tool(tool_query).tool(tool_insert);
7322
7323 let mut router = McpRouter::new()
7325 .server_info("nested", "1.0")
7326 .nest("db", db_router);
7327
7328 init_router(&mut router).await;
7329
7330 let req = RouterRequest {
7331 id: RequestId::Number(1),
7332 inner: McpRequest::ListTools(ListToolsParams::default()),
7333 extensions: Extensions::new(),
7334 };
7335
7336 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7337
7338 match resp.inner {
7339 Ok(McpResponse::ListTools(result)) => {
7340 assert_eq!(result.tools.len(), 2);
7341 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7342 assert!(names.contains(&"db.query"));
7343 assert!(names.contains(&"db.insert"));
7344 }
7345 _ => panic!("Expected ListTools response"),
7346 }
7347 }
7348
7349 #[tokio::test]
7350 async fn test_router_nest_call_prefixed_tool() {
7351 let tool = ToolBuilder::new("echo")
7352 .description("Echo input")
7353 .handler(|input: StringInput| async move { Ok(CallToolResult::text(&input.value)) })
7354 .build();
7355
7356 let nested_router = McpRouter::new().tool(tool);
7357
7358 let mut router = McpRouter::new().nest("api", nested_router);
7359
7360 init_router(&mut router).await;
7361
7362 let req = RouterRequest {
7364 id: RequestId::Number(1),
7365 inner: McpRequest::CallTool(CallToolParams {
7366 input_responses: None,
7367 request_state: None,
7368 name: "api.echo".to_string(),
7369 arguments: serde_json::json!({"value": "hello world"}),
7370 meta: None,
7371 task: None,
7372 }),
7373 extensions: Extensions::new(),
7374 };
7375
7376 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7377
7378 match resp.inner {
7379 Ok(McpResponse::CallTool(result)) => {
7380 assert!(!result.is_error);
7381 match &result.content[0] {
7382 Content::Text { text, .. } => assert_eq!(text, "hello world"),
7383 _ => panic!("Expected text content"),
7384 }
7385 }
7386 _ => panic!("Expected CallTool response"),
7387 }
7388 }
7389
7390 #[tokio::test]
7391 async fn test_router_multiple_nests() {
7392 let db_tool = ToolBuilder::new("query")
7393 .description("Database query")
7394 .handler(|_: StringInput| async move { Ok(CallToolResult::text("db")) })
7395 .build();
7396
7397 let api_tool = ToolBuilder::new("fetch")
7398 .description("API fetch")
7399 .handler(|_: StringInput| async move { Ok(CallToolResult::text("api")) })
7400 .build();
7401
7402 let db_router = McpRouter::new().tool(db_tool);
7403 let api_router = McpRouter::new().tool(api_tool);
7404
7405 let mut router = McpRouter::new()
7406 .nest("db", db_router)
7407 .nest("api", api_router);
7408
7409 init_router(&mut router).await;
7410
7411 let req = RouterRequest {
7412 id: RequestId::Number(1),
7413 inner: McpRequest::ListTools(ListToolsParams::default()),
7414 extensions: Extensions::new(),
7415 };
7416
7417 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7418
7419 match resp.inner {
7420 Ok(McpResponse::ListTools(result)) => {
7421 assert_eq!(result.tools.len(), 2);
7422 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7423 assert!(names.contains(&"db.query"));
7424 assert!(names.contains(&"api.fetch"));
7425 }
7426 _ => panic!("Expected ListTools response"),
7427 }
7428 }
7429
7430 #[tokio::test]
7431 async fn test_router_merge_and_nest_combined() {
7432 let tool_a = ToolBuilder::new("local")
7434 .description("Local tool")
7435 .handler(|_: StringInput| async move { Ok(CallToolResult::text("local")) })
7436 .build();
7437
7438 let nested_tool = ToolBuilder::new("remote")
7439 .description("Remote tool")
7440 .handler(|_: StringInput| async move { Ok(CallToolResult::text("remote")) })
7441 .build();
7442
7443 let nested_router = McpRouter::new().tool(nested_tool);
7444
7445 let mut router = McpRouter::new()
7446 .tool(tool_a)
7447 .nest("external", nested_router);
7448
7449 init_router(&mut router).await;
7450
7451 let req = RouterRequest {
7452 id: RequestId::Number(1),
7453 inner: McpRequest::ListTools(ListToolsParams::default()),
7454 extensions: Extensions::new(),
7455 };
7456
7457 let resp = router.ready().await.unwrap().call(req).await.unwrap();
7458
7459 match resp.inner {
7460 Ok(McpResponse::ListTools(result)) => {
7461 assert_eq!(result.tools.len(), 2);
7462 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
7463 assert!(names.contains(&"local"));
7464 assert!(names.contains(&"external.remote"));
7465 }
7466 _ => panic!("Expected ListTools response"),
7467 }
7468 }
7469
7470 #[tokio::test]
7471 async fn test_router_merge_preserves_server_info() {
7472 let child_router = McpRouter::new()
7473 .server_info("child", "2.0")
7474 .instructions("Child instructions");
7475
7476 let mut router = McpRouter::new()
7477 .server_info("parent", "1.0")
7478 .instructions("Parent instructions")
7479 .merge(child_router);
7480
7481 init_router(&mut router).await;
7482
7483 let init_req = RouterRequest {
7485 id: RequestId::Number(99),
7486 inner: McpRequest::Initialize(InitializeParams {
7487 protocol_version: "2025-11-25".to_string(),
7488 capabilities: ClientCapabilities::default(),
7489 client_info: Implementation {
7490 name: "test".to_string(),
7491 version: "1.0".to_string(),
7492 ..Default::default()
7493 },
7494 meta: None,
7495 }),
7496 extensions: Extensions::new(),
7497 };
7498
7499 let child_router2 = McpRouter::new().server_info("child", "2.0");
7501 let mut fresh_router = McpRouter::new()
7502 .server_info("parent", "1.0")
7503 .merge(child_router2);
7504
7505 let resp = fresh_router
7506 .ready()
7507 .await
7508 .unwrap()
7509 .call(init_req)
7510 .await
7511 .unwrap();
7512
7513 match resp.inner {
7514 Ok(McpResponse::Initialize(result)) => {
7515 assert_eq!(result.server_info.name, "parent");
7516 assert_eq!(result.server_info.version, "1.0");
7517 }
7518 _ => panic!("Expected Initialize response"),
7519 }
7520 }
7521
7522 #[tokio::test]
7527 async fn test_auto_instructions_tools_only() {
7528 let tool_a = ToolBuilder::new("alpha")
7529 .description("Alpha tool")
7530 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7531 .build();
7532 let tool_b = ToolBuilder::new("beta")
7533 .description("Beta tool")
7534 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7535 .build();
7536
7537 let mut router = McpRouter::new()
7538 .auto_instructions()
7539 .tool(tool_a)
7540 .tool(tool_b);
7541
7542 let resp = send_initialize(&mut router).await;
7543 let instructions = resp.instructions.expect("should have instructions");
7544
7545 assert!(instructions.contains("## Tools"));
7546 assert!(instructions.contains("- **alpha**: Alpha tool"));
7547 assert!(instructions.contains("- **beta**: Beta tool"));
7548 assert!(!instructions.contains("## Resources"));
7550 assert!(!instructions.contains("## Prompts"));
7551 }
7552
7553 #[tokio::test]
7554 async fn test_auto_instructions_with_annotations() {
7555 let read_only_tool = ToolBuilder::new("query")
7556 .description("Run a query")
7557 .read_only()
7558 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7559 .build();
7560 let destructive_tool = ToolBuilder::new("delete")
7561 .description("Delete a record")
7562 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7563 .build();
7564 let idempotent_tool = ToolBuilder::new("upsert")
7565 .description("Upsert a record")
7566 .non_destructive()
7567 .idempotent()
7568 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7569 .build();
7570
7571 let mut router = McpRouter::new()
7572 .auto_instructions()
7573 .tool(read_only_tool)
7574 .tool(destructive_tool)
7575 .tool(idempotent_tool);
7576
7577 let resp = send_initialize(&mut router).await;
7578 let instructions = resp.instructions.unwrap();
7579
7580 assert!(instructions.contains("- **query**: Run a query [read-only]"));
7581 assert!(instructions.contains("- **delete**: Delete a record\n"));
7583 assert!(instructions.contains("- **upsert**: Upsert a record [idempotent]"));
7584 }
7585
7586 #[tokio::test]
7587 async fn test_auto_instructions_with_resources() {
7588 use crate::resource::ResourceBuilder;
7589
7590 let resource = ResourceBuilder::new("file:///schema.sql")
7591 .name("Schema")
7592 .description("Database schema")
7593 .text("CREATE TABLE ...");
7594
7595 let mut router = McpRouter::new().auto_instructions().resource(resource);
7596
7597 let resp = send_initialize(&mut router).await;
7598 let instructions = resp.instructions.unwrap();
7599
7600 assert!(instructions.contains("## Resources"));
7601 assert!(instructions.contains("- **file:///schema.sql**: Database schema"));
7602 assert!(!instructions.contains("## Tools"));
7603 }
7604
7605 #[tokio::test]
7606 async fn test_auto_instructions_with_resource_templates() {
7607 use crate::resource::ResourceTemplateBuilder;
7608
7609 let template = ResourceTemplateBuilder::new("file:///{path}")
7610 .name("File")
7611 .description("Read a file by path")
7612 .handler(
7613 |_uri: String, _vars: std::collections::HashMap<String, String>| async move {
7614 Ok(crate::ReadResourceResult::text("content", "text/plain"))
7615 },
7616 );
7617
7618 let mut router = McpRouter::new()
7619 .auto_instructions()
7620 .resource_template(template);
7621
7622 let resp = send_initialize(&mut router).await;
7623 let instructions = resp.instructions.unwrap();
7624
7625 assert!(instructions.contains("## Resources"));
7626 assert!(instructions.contains("- **file:///{path}**: Read a file by path"));
7627 }
7628
7629 #[tokio::test]
7630 async fn test_auto_instructions_with_prompts() {
7631 use crate::prompt::PromptBuilder;
7632
7633 let prompt = PromptBuilder::new("write_query")
7634 .description("Help write a SQL query")
7635 .user_message("Write a query for: {task}");
7636
7637 let mut router = McpRouter::new().auto_instructions().prompt(prompt);
7638
7639 let resp = send_initialize(&mut router).await;
7640 let instructions = resp.instructions.unwrap();
7641
7642 assert!(instructions.contains("## Prompts"));
7643 assert!(instructions.contains("- **write_query**: Help write a SQL query"));
7644 assert!(!instructions.contains("## Tools"));
7645 }
7646
7647 #[tokio::test]
7648 async fn test_auto_instructions_all_sections() {
7649 use crate::prompt::PromptBuilder;
7650 use crate::resource::ResourceBuilder;
7651
7652 let tool = ToolBuilder::new("query")
7653 .description("Execute SQL")
7654 .read_only()
7655 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7656 .build();
7657 let resource = ResourceBuilder::new("db://schema")
7658 .name("Schema")
7659 .description("Full database schema")
7660 .text("schema");
7661 let prompt = PromptBuilder::new("write_query")
7662 .description("Help write a SQL query")
7663 .user_message("Write a query");
7664
7665 let mut router = McpRouter::new()
7666 .auto_instructions()
7667 .tool(tool)
7668 .resource(resource)
7669 .prompt(prompt);
7670
7671 let resp = send_initialize(&mut router).await;
7672 let instructions = resp.instructions.unwrap();
7673
7674 assert!(instructions.contains("## Tools"));
7676 assert!(instructions.contains("## Resources"));
7677 assert!(instructions.contains("## Prompts"));
7678
7679 let tools_pos = instructions.find("## Tools").unwrap();
7681 let resources_pos = instructions.find("## Resources").unwrap();
7682 let prompts_pos = instructions.find("## Prompts").unwrap();
7683 assert!(tools_pos < resources_pos);
7684 assert!(resources_pos < prompts_pos);
7685 }
7686
7687 #[tokio::test]
7688 async fn test_auto_instructions_with_prefix_and_suffix() {
7689 let tool = ToolBuilder::new("echo")
7690 .description("Echo input")
7691 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7692 .build();
7693
7694 let mut router = McpRouter::new()
7695 .auto_instructions_with(
7696 Some("This server provides echo capabilities."),
7697 Some("Contact admin@example.com for support."),
7698 )
7699 .tool(tool);
7700
7701 let resp = send_initialize(&mut router).await;
7702 let instructions = resp.instructions.unwrap();
7703
7704 assert!(instructions.starts_with("This server provides echo capabilities."));
7705 assert!(instructions.ends_with("Contact admin@example.com for support."));
7706 assert!(instructions.contains("## Tools"));
7707 assert!(instructions.contains("- **echo**: Echo input"));
7708 }
7709
7710 #[tokio::test]
7711 async fn test_auto_instructions_prefix_only() {
7712 let tool = ToolBuilder::new("echo")
7713 .description("Echo input")
7714 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7715 .build();
7716
7717 let mut router = McpRouter::new()
7718 .auto_instructions_with(Some("My server intro."), None::<String>)
7719 .tool(tool);
7720
7721 let resp = send_initialize(&mut router).await;
7722 let instructions = resp.instructions.unwrap();
7723
7724 assert!(instructions.starts_with("My server intro."));
7725 assert!(instructions.contains("- **echo**: Echo input"));
7726 }
7727
7728 #[tokio::test]
7729 async fn test_auto_instructions_empty_router() {
7730 let mut router = McpRouter::new().auto_instructions();
7731
7732 let resp = send_initialize(&mut router).await;
7733 let instructions = resp.instructions.expect("should have instructions");
7734
7735 assert!(!instructions.contains("## Tools"));
7737 assert!(!instructions.contains("## Resources"));
7738 assert!(!instructions.contains("## Prompts"));
7739 assert!(instructions.is_empty());
7740 }
7741
7742 #[tokio::test]
7743 async fn test_auto_instructions_overrides_manual() {
7744 let tool = ToolBuilder::new("echo")
7745 .description("Echo input")
7746 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7747 .build();
7748
7749 let mut router = McpRouter::new()
7750 .instructions("This will be overridden")
7751 .auto_instructions()
7752 .tool(tool);
7753
7754 let resp = send_initialize(&mut router).await;
7755 let instructions = resp.instructions.unwrap();
7756
7757 assert!(!instructions.contains("This will be overridden"));
7758 assert!(instructions.contains("- **echo**: Echo input"));
7759 }
7760
7761 #[tokio::test]
7762 async fn test_no_auto_instructions_returns_manual() {
7763 let tool = ToolBuilder::new("echo")
7764 .description("Echo input")
7765 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7766 .build();
7767
7768 let mut router = McpRouter::new()
7769 .instructions("Manual instructions here")
7770 .tool(tool);
7771
7772 let resp = send_initialize(&mut router).await;
7773 let instructions = resp.instructions.unwrap();
7774
7775 assert_eq!(instructions, "Manual instructions here");
7776 }
7777
7778 #[tokio::test]
7779 async fn test_auto_instructions_no_description_fallback() {
7780 let tool = ToolBuilder::new("mystery")
7781 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7782 .build();
7783
7784 let mut router = McpRouter::new().auto_instructions().tool(tool);
7785
7786 let resp = send_initialize(&mut router).await;
7787 let instructions = resp.instructions.unwrap();
7788
7789 assert!(instructions.contains("- **mystery**: No description"));
7790 }
7791
7792 #[tokio::test]
7793 async fn test_auto_instructions_sorted_alphabetically() {
7794 let tool_z = ToolBuilder::new("zebra")
7795 .description("Z tool")
7796 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7797 .build();
7798 let tool_a = ToolBuilder::new("alpha")
7799 .description("A tool")
7800 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7801 .build();
7802 let tool_m = ToolBuilder::new("middle")
7803 .description("M tool")
7804 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7805 .build();
7806
7807 let mut router = McpRouter::new()
7808 .auto_instructions()
7809 .tool(tool_z)
7810 .tool(tool_a)
7811 .tool(tool_m);
7812
7813 let resp = send_initialize(&mut router).await;
7814 let instructions = resp.instructions.unwrap();
7815
7816 let alpha_pos = instructions.find("**alpha**").unwrap();
7817 let middle_pos = instructions.find("**middle**").unwrap();
7818 let zebra_pos = instructions.find("**zebra**").unwrap();
7819 assert!(alpha_pos < middle_pos);
7820 assert!(middle_pos < zebra_pos);
7821 }
7822
7823 #[tokio::test]
7824 async fn test_auto_instructions_read_only_and_idempotent_tags() {
7825 let tool = ToolBuilder::new("safe_update")
7826 .description("Safe update operation")
7827 .idempotent()
7828 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7829 .build();
7830
7831 let mut router = McpRouter::new().auto_instructions().tool(tool);
7832
7833 let resp = send_initialize(&mut router).await;
7834 let instructions = resp.instructions.unwrap();
7835
7836 assert!(
7837 instructions.contains("[idempotent]"),
7838 "got: {}",
7839 instructions
7840 );
7841 }
7842
7843 #[tokio::test]
7844 async fn test_auto_instructions_lazy_generation() {
7845 let mut router = McpRouter::new().auto_instructions();
7848
7849 let tool = ToolBuilder::new("late_tool")
7850 .description("Added after auto_instructions")
7851 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7852 .build();
7853
7854 router = router.tool(tool);
7855
7856 let resp = send_initialize(&mut router).await;
7857 let instructions = resp.instructions.unwrap();
7858
7859 assert!(instructions.contains("- **late_tool**: Added after auto_instructions"));
7860 }
7861
7862 #[tokio::test]
7863 async fn test_auto_instructions_multiple_annotation_tags() {
7864 let tool = ToolBuilder::new("update")
7865 .description("Update a record")
7866 .annotations(ToolAnnotations {
7867 read_only_hint: true,
7868 idempotent_hint: true,
7869 ..Default::default()
7870 })
7871 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7872 .build();
7873
7874 let mut router = McpRouter::new().auto_instructions().tool(tool);
7875
7876 let resp = send_initialize(&mut router).await;
7877 let instructions = resp.instructions.unwrap();
7878
7879 assert!(
7880 instructions.contains("[read-only, idempotent]"),
7881 "got: {}",
7882 instructions
7883 );
7884 }
7885
7886 #[tokio::test]
7887 async fn test_auto_instructions_no_annotations_no_tags() {
7888 let tool = ToolBuilder::new("fetch")
7890 .description("Fetch data")
7891 .handler(|_: AddInput| async move { Ok(CallToolResult::text("ok")) })
7892 .build();
7893
7894 let mut router = McpRouter::new().auto_instructions().tool(tool);
7895
7896 let resp = send_initialize(&mut router).await;
7897 let instructions = resp.instructions.unwrap();
7898
7899 assert!(
7901 !instructions.contains('['),
7902 "should have no tags, got: {}",
7903 instructions
7904 );
7905 assert!(instructions.contains("- **fetch**: Fetch data"));
7906 }
7907
7908 async fn send_initialize(router: &mut McpRouter) -> InitializeResult {
7910 let init_req = RouterRequest {
7911 id: RequestId::Number(0),
7912 inner: McpRequest::Initialize(InitializeParams {
7913 protocol_version: "2025-11-25".to_string(),
7914 capabilities: ClientCapabilities {
7915 roots: None,
7916 sampling: None,
7917 elicitation: None,
7918 tasks: None,
7919 experimental: None,
7920 extensions: None,
7921 },
7922 client_info: Implementation {
7923 name: "test".to_string(),
7924 version: "1.0".to_string(),
7925 ..Default::default()
7926 },
7927 meta: None,
7928 }),
7929 extensions: Extensions::new(),
7930 };
7931 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
7932 match resp.inner {
7933 Ok(McpResponse::Initialize(result)) => result,
7934 other => panic!("Expected Initialize response, got {:?}", other),
7935 }
7936 }
7937
7938 #[tokio::test]
7939 async fn test_notify_tools_list_changed() {
7940 let (tx, mut rx) = crate::context::notification_channel(16);
7941
7942 let router = McpRouter::new()
7943 .server_info("test", "1.0")
7944 .with_notification_sender(tx);
7945
7946 assert!(router.notify_tools_list_changed());
7947
7948 let notification = rx.recv().await.unwrap();
7949 assert!(matches!(notification, ServerNotification::ToolsListChanged));
7950 }
7951
7952 #[tokio::test]
7953 async fn test_notify_prompts_list_changed() {
7954 let (tx, mut rx) = crate::context::notification_channel(16);
7955
7956 let router = McpRouter::new()
7957 .server_info("test", "1.0")
7958 .with_notification_sender(tx);
7959
7960 assert!(router.notify_prompts_list_changed());
7961
7962 let notification = rx.recv().await.unwrap();
7963 assert!(matches!(
7964 notification,
7965 ServerNotification::PromptsListChanged
7966 ));
7967 }
7968
7969 #[tokio::test]
7970 async fn test_notify_without_sender_returns_false() {
7971 let router = McpRouter::new().server_info("test", "1.0");
7972
7973 assert!(!router.notify_tools_list_changed());
7974 assert!(!router.notify_prompts_list_changed());
7975 assert!(!router.notify_resources_list_changed());
7976 }
7977
7978 #[tokio::test]
7979 async fn test_list_changed_capabilities_with_notification_sender() {
7980 let (tx, _rx) = crate::context::notification_channel(16);
7981 let tool = ToolBuilder::new("test")
7982 .description("test")
7983 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
7984 .build();
7985
7986 let mut router = McpRouter::new()
7987 .server_info("test", "1.0")
7988 .tool(tool)
7989 .with_notification_sender(tx);
7990
7991 init_router(&mut router).await;
7992
7993 let caps = router.capabilities();
7994 let tools_cap = caps.tools.expect("tools capability should be present");
7995 assert!(
7996 tools_cap.list_changed,
7997 "tools.listChanged should be true when notification sender is configured"
7998 );
7999 }
8000
8001 #[tokio::test]
8002 async fn test_list_changed_capabilities_without_notification_sender() {
8003 let tool = ToolBuilder::new("test")
8004 .description("test")
8005 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8006 .build();
8007
8008 let mut router = McpRouter::new().server_info("test", "1.0").tool(tool);
8009
8010 init_router(&mut router).await;
8011
8012 let caps = router.capabilities();
8013 let tools_cap = caps.tools.expect("tools capability should be present");
8014 assert!(
8015 !tools_cap.list_changed,
8016 "tools.listChanged should be false without notification sender"
8017 );
8018 }
8019
8020 #[tokio::test]
8021 async fn test_set_logging_level_filters_messages() {
8022 let (tx, mut rx) = crate::context::notification_channel(16);
8023
8024 let mut router = McpRouter::new()
8025 .server_info("test", "1.0")
8026 .with_notification_sender(tx);
8027
8028 init_router(&mut router).await;
8029
8030 let set_level_req = RouterRequest {
8032 id: RequestId::Number(99),
8033 inner: McpRequest::SetLoggingLevel(SetLogLevelParams {
8034 level: LogLevel::Warning,
8035 meta: None,
8036 }),
8037 extensions: crate::context::Extensions::new(),
8038 };
8039 let resp = router
8040 .ready()
8041 .await
8042 .unwrap()
8043 .call(set_level_req)
8044 .await
8045 .unwrap();
8046 assert!(matches!(resp.inner, Ok(McpResponse::SetLoggingLevel(_))));
8047
8048 let ctx = router.create_context(RequestId::Number(100), None);
8050
8051 ctx.send_log(LoggingMessageParams::new(
8053 LogLevel::Error,
8054 serde_json::Value::Null,
8055 ));
8056 assert!(
8057 rx.try_recv().is_ok(),
8058 "Error should pass through Warning filter"
8059 );
8060
8061 ctx.send_log(LoggingMessageParams::new(
8063 LogLevel::Info,
8064 serde_json::Value::Null,
8065 ));
8066 assert!(
8067 rx.try_recv().is_err(),
8068 "Info should be filtered at Warning level"
8069 );
8070 }
8071
8072 #[test]
8073 fn test_paginate_no_page_size() {
8074 let items = vec![1, 2, 3, 4, 5];
8075 let (page, cursor) = paginate(items.clone(), None, None).unwrap();
8076 assert_eq!(page, items);
8077 assert!(cursor.is_none());
8078 }
8079
8080 #[test]
8081 fn test_paginate_first_page() {
8082 let items = vec![1, 2, 3, 4, 5];
8083 let (page, cursor) = paginate(items, None, Some(2)).unwrap();
8084 assert_eq!(page, vec![1, 2]);
8085 assert!(cursor.is_some());
8086 }
8087
8088 #[test]
8089 fn test_paginate_middle_page() {
8090 let items = vec![1, 2, 3, 4, 5];
8091 let (page1, cursor1) = paginate(items.clone(), None, Some(2)).unwrap();
8092 assert_eq!(page1, vec![1, 2]);
8093
8094 let (page2, cursor2) = paginate(items, cursor1.as_deref(), Some(2)).unwrap();
8095 assert_eq!(page2, vec![3, 4]);
8096 assert!(cursor2.is_some());
8097 }
8098
8099 #[test]
8100 fn test_paginate_last_page() {
8101 let items = vec![1, 2, 3, 4, 5];
8102 let cursor = encode_cursor(4);
8104 let (page, next) = paginate(items, Some(&cursor), Some(2)).unwrap();
8105 assert_eq!(page, vec![5]);
8106 assert!(next.is_none());
8107 }
8108
8109 #[test]
8110 fn test_paginate_exact_boundary() {
8111 let items = vec![1, 2, 3, 4];
8112 let (page, cursor) = paginate(items, None, Some(4)).unwrap();
8113 assert_eq!(page, vec![1, 2, 3, 4]);
8114 assert!(cursor.is_none());
8115 }
8116
8117 #[test]
8118 fn test_paginate_invalid_cursor() {
8119 let items = vec![1, 2, 3];
8120 let result = paginate(items, Some("not-valid-base64!@#$"), Some(2));
8121 assert!(result.is_err());
8122 }
8123
8124 #[test]
8125 fn test_cursor_round_trip() {
8126 let offset = 42;
8127 let encoded = encode_cursor(offset);
8128 let decoded = decode_cursor(&encoded).unwrap();
8129 assert_eq!(decoded, offset);
8130 }
8131
8132 #[tokio::test]
8133 async fn test_list_tools_pagination() {
8134 let tool_a = ToolBuilder::new("alpha")
8135 .description("a")
8136 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8137 .build();
8138 let tool_b = ToolBuilder::new("beta")
8139 .description("b")
8140 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8141 .build();
8142 let tool_c = ToolBuilder::new("gamma")
8143 .description("c")
8144 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8145 .build();
8146
8147 let mut router = McpRouter::new()
8148 .server_info("test", "1.0")
8149 .page_size(2)
8150 .tool(tool_a)
8151 .tool(tool_b)
8152 .tool(tool_c);
8153
8154 init_router(&mut router).await;
8155
8156 let req = RouterRequest {
8158 id: RequestId::Number(1),
8159 inner: McpRequest::ListTools(ListToolsParams {
8160 cursor: None,
8161 meta: None,
8162 }),
8163 extensions: Extensions::new(),
8164 };
8165 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8166 let (tools, next_cursor) = match resp.inner {
8167 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8168 other => panic!("Expected ListTools, got {:?}", other),
8169 };
8170 assert_eq!(tools.len(), 2);
8171 assert_eq!(tools[0].name, "alpha");
8172 assert_eq!(tools[1].name, "beta");
8173 assert!(next_cursor.is_some());
8174
8175 let req = RouterRequest {
8177 id: RequestId::Number(2),
8178 inner: McpRequest::ListTools(ListToolsParams {
8179 cursor: next_cursor,
8180 meta: None,
8181 }),
8182 extensions: Extensions::new(),
8183 };
8184 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8185 let (tools, next_cursor) = match resp.inner {
8186 Ok(McpResponse::ListTools(result)) => (result.tools, result.next_cursor),
8187 other => panic!("Expected ListTools, got {:?}", other),
8188 };
8189 assert_eq!(tools.len(), 1);
8190 assert_eq!(tools[0].name, "gamma");
8191 assert!(next_cursor.is_none());
8192 }
8193
8194 #[tokio::test]
8195 async fn test_list_tools_no_pagination_by_default() {
8196 let tool_a = ToolBuilder::new("alpha")
8197 .description("a")
8198 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8199 .build();
8200 let tool_b = ToolBuilder::new("beta")
8201 .description("b")
8202 .handler(|_input: AddInput| async { Ok(CallToolResult::text("ok")) })
8203 .build();
8204
8205 let mut router = McpRouter::new()
8206 .server_info("test", "1.0")
8207 .tool(tool_a)
8208 .tool(tool_b);
8209
8210 init_router(&mut router).await;
8211
8212 let req = RouterRequest {
8213 id: RequestId::Number(1),
8214 inner: McpRequest::ListTools(ListToolsParams {
8215 cursor: None,
8216 meta: None,
8217 }),
8218 extensions: Extensions::new(),
8219 };
8220 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8221 match resp.inner {
8222 Ok(McpResponse::ListTools(result)) => {
8223 assert_eq!(result.tools.len(), 2);
8224 assert!(result.next_cursor.is_none());
8225 }
8226 other => panic!("Expected ListTools, got {:?}", other),
8227 }
8228 }
8229
8230 #[cfg(feature = "dynamic-tools")]
8235 mod dynamic_tools_tests {
8236 use super::*;
8237
8238 #[tokio::test]
8239 async fn test_dynamic_tools_register_and_list() {
8240 let (router, registry) = McpRouter::new()
8241 .server_info("test", "1.0")
8242 .with_dynamic_tools();
8243
8244 let tool = ToolBuilder::new("dynamic_echo")
8245 .description("Dynamic echo")
8246 .handler(|input: AddInput| async move {
8247 Ok(CallToolResult::text(format!("{}", input.a)))
8248 })
8249 .build();
8250
8251 registry.register(tool);
8252
8253 let mut router = router;
8254 init_router(&mut router).await;
8255
8256 let req = RouterRequest {
8257 id: RequestId::Number(1),
8258 inner: McpRequest::ListTools(ListToolsParams::default()),
8259 extensions: Extensions::new(),
8260 };
8261
8262 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8263 match resp.inner {
8264 Ok(McpResponse::ListTools(result)) => {
8265 assert_eq!(result.tools.len(), 1);
8266 assert_eq!(result.tools[0].name, "dynamic_echo");
8267 }
8268 _ => panic!("Expected ListTools response"),
8269 }
8270 }
8271
8272 #[tokio::test]
8273 async fn test_dynamic_tools_unregister() {
8274 let (router, registry) = McpRouter::new()
8275 .server_info("test", "1.0")
8276 .with_dynamic_tools();
8277
8278 let tool = ToolBuilder::new("temp")
8279 .description("Temporary")
8280 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8281 .build();
8282
8283 registry.register(tool);
8284 assert!(registry.contains("temp"));
8285
8286 let removed = registry.unregister("temp");
8287 assert!(removed);
8288 assert!(!registry.contains("temp"));
8289
8290 assert!(!registry.unregister("temp"));
8292
8293 let mut router = router;
8294 init_router(&mut router).await;
8295
8296 let req = RouterRequest {
8297 id: RequestId::Number(1),
8298 inner: McpRequest::ListTools(ListToolsParams::default()),
8299 extensions: Extensions::new(),
8300 };
8301
8302 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8303 match resp.inner {
8304 Ok(McpResponse::ListTools(result)) => {
8305 assert_eq!(result.tools.len(), 0);
8306 }
8307 _ => panic!("Expected ListTools response"),
8308 }
8309 }
8310
8311 #[tokio::test]
8312 async fn test_dynamic_tools_merged_with_static() {
8313 let static_tool = ToolBuilder::new("static_tool")
8314 .description("Static")
8315 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8316 .build();
8317
8318 let (router, registry) = McpRouter::new()
8319 .server_info("test", "1.0")
8320 .tool(static_tool)
8321 .with_dynamic_tools();
8322
8323 let dynamic_tool = ToolBuilder::new("dynamic_tool")
8324 .description("Dynamic")
8325 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8326 .build();
8327
8328 registry.register(dynamic_tool);
8329
8330 let mut router = router;
8331 init_router(&mut router).await;
8332
8333 let req = RouterRequest {
8334 id: RequestId::Number(1),
8335 inner: McpRequest::ListTools(ListToolsParams::default()),
8336 extensions: Extensions::new(),
8337 };
8338
8339 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8340 match resp.inner {
8341 Ok(McpResponse::ListTools(result)) => {
8342 assert_eq!(result.tools.len(), 2);
8343 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8344 assert!(names.contains(&"static_tool"));
8345 assert!(names.contains(&"dynamic_tool"));
8346 }
8347 _ => panic!("Expected ListTools response"),
8348 }
8349 }
8350
8351 #[tokio::test]
8352 async fn test_static_tools_shadow_dynamic() {
8353 let static_tool = ToolBuilder::new("shared")
8354 .description("Static version")
8355 .handler(|_: AddInput| async { Ok(CallToolResult::text("static")) })
8356 .build();
8357
8358 let (router, registry) = McpRouter::new()
8359 .server_info("test", "1.0")
8360 .tool(static_tool)
8361 .with_dynamic_tools();
8362
8363 let dynamic_tool = ToolBuilder::new("shared")
8364 .description("Dynamic version")
8365 .handler(|_: AddInput| async { Ok(CallToolResult::text("dynamic")) })
8366 .build();
8367
8368 registry.register(dynamic_tool);
8369
8370 let mut router = router;
8371 init_router(&mut router).await;
8372
8373 let req = RouterRequest {
8375 id: RequestId::Number(1),
8376 inner: McpRequest::ListTools(ListToolsParams::default()),
8377 extensions: Extensions::new(),
8378 };
8379
8380 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8381 match resp.inner {
8382 Ok(McpResponse::ListTools(result)) => {
8383 assert_eq!(result.tools.len(), 1);
8384 assert_eq!(result.tools[0].name, "shared");
8385 assert_eq!(
8386 result.tools[0].description.as_deref(),
8387 Some("Static version")
8388 );
8389 }
8390 _ => panic!("Expected ListTools response"),
8391 }
8392
8393 let req = RouterRequest {
8395 id: RequestId::Number(2),
8396 inner: McpRequest::CallTool(CallToolParams {
8397 input_responses: None,
8398 request_state: None,
8399 name: "shared".to_string(),
8400 arguments: serde_json::json!({"a": 1, "b": 2}),
8401 meta: None,
8402 task: None,
8403 }),
8404 extensions: Extensions::new(),
8405 };
8406
8407 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8408 match resp.inner {
8409 Ok(McpResponse::CallTool(result)) => {
8410 assert!(!result.is_error);
8411 match &result.content[0] {
8412 Content::Text { text, .. } => assert_eq!(text, "static"),
8413 _ => panic!("Expected text content"),
8414 }
8415 }
8416 _ => panic!("Expected CallTool response"),
8417 }
8418 }
8419
8420 #[tokio::test]
8421 async fn test_dynamic_tools_call() {
8422 let (router, registry) = McpRouter::new()
8423 .server_info("test", "1.0")
8424 .with_dynamic_tools();
8425
8426 let tool = ToolBuilder::new("add")
8427 .description("Add two numbers")
8428 .handler(|input: AddInput| async move {
8429 Ok(CallToolResult::text(format!("{}", input.a + input.b)))
8430 })
8431 .build();
8432
8433 registry.register(tool);
8434
8435 let mut router = router;
8436 init_router(&mut router).await;
8437
8438 let req = RouterRequest {
8439 id: RequestId::Number(1),
8440 inner: McpRequest::CallTool(CallToolParams {
8441 input_responses: None,
8442 request_state: None,
8443 name: "add".to_string(),
8444 arguments: serde_json::json!({"a": 3, "b": 4}),
8445 meta: None,
8446 task: None,
8447 }),
8448 extensions: Extensions::new(),
8449 };
8450
8451 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8452 match resp.inner {
8453 Ok(McpResponse::CallTool(result)) => {
8454 assert!(!result.is_error);
8455 match &result.content[0] {
8456 Content::Text { text, .. } => assert_eq!(text, "7"),
8457 _ => panic!("Expected text content"),
8458 }
8459 }
8460 _ => panic!("Expected CallTool response"),
8461 }
8462 }
8463
8464 #[tokio::test]
8465 async fn test_dynamic_tools_notification_on_register() {
8466 let (tx, mut rx) = crate::context::notification_channel(16);
8467 let (router, registry) = McpRouter::new()
8468 .server_info("test", "1.0")
8469 .with_dynamic_tools();
8470 let _router = router.with_notification_sender(tx);
8471
8472 let tool = ToolBuilder::new("notified")
8473 .description("Test")
8474 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8475 .build();
8476
8477 registry.register(tool);
8478
8479 let notification = rx.recv().await.unwrap();
8480 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8481 }
8482
8483 #[tokio::test]
8484 async fn test_dynamic_tools_notification_on_unregister() {
8485 let (tx, mut rx) = crate::context::notification_channel(16);
8486 let (router, registry) = McpRouter::new()
8487 .server_info("test", "1.0")
8488 .with_dynamic_tools();
8489 let _router = router.with_notification_sender(tx);
8490
8491 let tool = ToolBuilder::new("notified")
8492 .description("Test")
8493 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8494 .build();
8495
8496 registry.register(tool);
8497 let _ = rx.recv().await.unwrap();
8499
8500 registry.unregister("notified");
8501 let notification = rx.recv().await.unwrap();
8502 assert!(matches!(notification, ServerNotification::ToolsListChanged));
8503 }
8504
8505 #[tokio::test]
8506 async fn test_dynamic_tools_no_notification_on_empty_unregister() {
8507 let (tx, mut rx) = crate::context::notification_channel(16);
8508 let (router, registry) = McpRouter::new()
8509 .server_info("test", "1.0")
8510 .with_dynamic_tools();
8511 let _router = router.with_notification_sender(tx);
8512
8513 assert!(!registry.unregister("nonexistent"));
8515
8516 assert!(rx.try_recv().is_err());
8518 }
8519
8520 #[tokio::test]
8521 async fn test_dynamic_tools_filter_applies() {
8522 use crate::filter::CapabilityFilter;
8523
8524 let (router, registry) = McpRouter::new()
8525 .server_info("test", "1.0")
8526 .tool_filter(CapabilityFilter::new(|_, tool: &Tool| {
8527 tool.name != "hidden"
8528 }))
8529 .with_dynamic_tools();
8530
8531 let visible = ToolBuilder::new("visible")
8532 .description("Visible")
8533 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8534 .build();
8535
8536 let hidden = ToolBuilder::new("hidden")
8537 .description("Hidden")
8538 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8539 .build();
8540
8541 registry.register(visible);
8542 registry.register(hidden);
8543
8544 let mut router = router;
8545 init_router(&mut router).await;
8546
8547 let req = RouterRequest {
8549 id: RequestId::Number(1),
8550 inner: McpRequest::ListTools(ListToolsParams::default()),
8551 extensions: Extensions::new(),
8552 };
8553
8554 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8555 match resp.inner {
8556 Ok(McpResponse::ListTools(result)) => {
8557 assert_eq!(result.tools.len(), 1);
8558 assert_eq!(result.tools[0].name, "visible");
8559 }
8560 _ => panic!("Expected ListTools response"),
8561 }
8562
8563 let req = RouterRequest {
8565 id: RequestId::Number(2),
8566 inner: McpRequest::CallTool(CallToolParams {
8567 input_responses: None,
8568 request_state: None,
8569 name: "hidden".to_string(),
8570 arguments: serde_json::json!({"a": 1, "b": 2}),
8571 meta: None,
8572 task: None,
8573 }),
8574 extensions: Extensions::new(),
8575 };
8576
8577 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8578 match resp.inner {
8579 Err(e) => {
8580 assert_eq!(e.code, -32601); }
8582 _ => panic!("Expected JsonRpc error"),
8583 }
8584 }
8585
8586 #[tokio::test]
8587 async fn test_dynamic_tools_capabilities_advertised() {
8588 let (mut router, _registry) = McpRouter::new()
8590 .server_info("test", "1.0")
8591 .with_dynamic_tools();
8592
8593 let init_req = RouterRequest {
8594 id: RequestId::Number(1),
8595 inner: McpRequest::Initialize(InitializeParams {
8596 protocol_version: "2025-11-25".to_string(),
8597 capabilities: ClientCapabilities::default(),
8598 client_info: Implementation {
8599 name: "test".to_string(),
8600 version: "1.0".to_string(),
8601 ..Default::default()
8602 },
8603 meta: None,
8604 }),
8605 extensions: Extensions::new(),
8606 };
8607
8608 let resp = router.ready().await.unwrap().call(init_req).await.unwrap();
8609 match resp.inner {
8610 Ok(McpResponse::Initialize(result)) => {
8611 assert!(result.capabilities.tools.is_some());
8612 }
8613 _ => panic!("Expected Initialize response"),
8614 }
8615 }
8616
8617 #[tokio::test]
8618 async fn test_dynamic_tools_multi_session_notification() {
8619 let (tx1, mut rx1) = crate::context::notification_channel(16);
8620 let (tx2, mut rx2) = crate::context::notification_channel(16);
8621
8622 let (router, registry) = McpRouter::new()
8623 .server_info("test", "1.0")
8624 .with_dynamic_tools();
8625
8626 let _session1 = router.clone().with_notification_sender(tx1);
8628 let _session2 = router.clone().with_notification_sender(tx2);
8629
8630 let tool = ToolBuilder::new("broadcast")
8631 .description("Test")
8632 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8633 .build();
8634
8635 registry.register(tool);
8636
8637 let n1 = rx1.recv().await.unwrap();
8639 let n2 = rx2.recv().await.unwrap();
8640 assert!(matches!(n1, ServerNotification::ToolsListChanged));
8641 assert!(matches!(n2, ServerNotification::ToolsListChanged));
8642 }
8643
8644 #[tokio::test]
8645 async fn test_dynamic_tools_call_not_found() {
8646 let (router, _registry) = McpRouter::new()
8647 .server_info("test", "1.0")
8648 .with_dynamic_tools();
8649
8650 let mut router = router;
8651 init_router(&mut router).await;
8652
8653 let req = RouterRequest {
8654 id: RequestId::Number(1),
8655 inner: McpRequest::CallTool(CallToolParams {
8656 input_responses: None,
8657 request_state: None,
8658 name: "nonexistent".to_string(),
8659 arguments: serde_json::json!({}),
8660 meta: None,
8661 task: None,
8662 }),
8663 extensions: Extensions::new(),
8664 };
8665
8666 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8667 match resp.inner {
8668 Err(e) => {
8669 assert_eq!(e.code, -32601);
8670 }
8671 _ => panic!("Expected method not found error"),
8672 }
8673 }
8674
8675 #[tokio::test]
8676 async fn test_dynamic_tools_registry_list() {
8677 let (_, registry) = McpRouter::new()
8678 .server_info("test", "1.0")
8679 .with_dynamic_tools();
8680
8681 assert!(registry.list().is_empty());
8682
8683 let tool = ToolBuilder::new("tool_a")
8684 .description("A")
8685 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8686 .build();
8687 registry.register(tool);
8688
8689 let tool = ToolBuilder::new("tool_b")
8690 .description("B")
8691 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8692 .build();
8693 registry.register(tool);
8694
8695 let tools = registry.list();
8696 assert_eq!(tools.len(), 2);
8697 let names: Vec<&str> = tools.iter().map(|t| t.name.as_str()).collect();
8698 assert!(names.contains(&"tool_a"));
8699 assert!(names.contains(&"tool_b"));
8700 }
8701 } #[tokio::test]
8704 async fn test_tool_if_true_registers() {
8705 let tool = ToolBuilder::new("conditional")
8706 .description("Conditional tool")
8707 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8708 .build();
8709
8710 let mut router = McpRouter::new().tool_if(true, tool);
8711 init_router(&mut router).await;
8712
8713 let req = RouterRequest {
8714 id: RequestId::Number(1),
8715 inner: McpRequest::ListTools(ListToolsParams::default()),
8716 extensions: Extensions::new(),
8717 };
8718 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8719 match resp.inner {
8720 Ok(McpResponse::ListTools(result)) => {
8721 assert_eq!(result.tools.len(), 1);
8722 assert_eq!(result.tools[0].name, "conditional");
8723 }
8724 _ => panic!("Expected ListTools response"),
8725 }
8726 }
8727
8728 #[tokio::test]
8729 async fn test_tool_if_false_skips() {
8730 let tool = ToolBuilder::new("conditional")
8731 .description("Conditional tool")
8732 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8733 .build();
8734
8735 let mut router = McpRouter::new().tool_if(false, tool);
8736 init_router(&mut router).await;
8737
8738 let req = RouterRequest {
8739 id: RequestId::Number(1),
8740 inner: McpRequest::ListTools(ListToolsParams::default()),
8741 extensions: Extensions::new(),
8742 };
8743 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8744 match resp.inner {
8745 Ok(McpResponse::ListTools(result)) => {
8746 assert_eq!(result.tools.len(), 0);
8747 }
8748 _ => panic!("Expected ListTools response"),
8749 }
8750 }
8751
8752 #[tokio::test]
8753 async fn test_tools_if_batch_conditional() {
8754 let tools = vec![
8755 ToolBuilder::new("a")
8756 .description("Tool A")
8757 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8758 .build(),
8759 ToolBuilder::new("b")
8760 .description("Tool B")
8761 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8762 .build(),
8763 ];
8764
8765 let mut router = McpRouter::new().tools_if(false, tools);
8766 init_router(&mut router).await;
8767
8768 let req = RouterRequest {
8769 id: RequestId::Number(1),
8770 inner: McpRequest::ListTools(ListToolsParams::default()),
8771 extensions: Extensions::new(),
8772 };
8773 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8774 match resp.inner {
8775 Ok(McpResponse::ListTools(result)) => {
8776 assert_eq!(result.tools.len(), 0);
8777 }
8778 _ => panic!("Expected ListTools response"),
8779 }
8780 }
8781
8782 #[test]
8783 fn test_resource_if_true_registers() {
8784 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8785 .name("test")
8786 .text("hello");
8787
8788 let router = McpRouter::new().resource_if(true, resource);
8789 assert_eq!(router.inner.resources.len(), 1);
8790 }
8791
8792 #[test]
8793 fn test_resource_if_false_skips() {
8794 let resource = crate::resource::ResourceBuilder::new("file:///test.txt")
8795 .name("test")
8796 .text("hello");
8797
8798 let router = McpRouter::new().resource_if(false, resource);
8799 assert_eq!(router.inner.resources.len(), 0);
8800 }
8801
8802 #[test]
8803 fn test_prompt_if_true_registers() {
8804 let prompt = crate::prompt::PromptBuilder::new("greet")
8805 .description("Greeting")
8806 .user_message("Hello!");
8807
8808 let router = McpRouter::new().prompt_if(true, prompt);
8809 assert_eq!(router.inner.prompts.len(), 1);
8810 }
8811
8812 #[test]
8813 fn test_prompt_if_false_skips() {
8814 let prompt = crate::prompt::PromptBuilder::new("greet")
8815 .description("Greeting")
8816 .user_message("Hello!");
8817
8818 let router = McpRouter::new().prompt_if(false, prompt);
8819 assert_eq!(router.inner.prompts.len(), 0);
8820 }
8821
8822 #[tokio::test]
8823 async fn test_disable_tool_hides_from_list() {
8824 let safe = ToolBuilder::new("safe")
8825 .description("Safe tool")
8826 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8827 .build();
8828 let dangerous = ToolBuilder::new("dangerous")
8829 .description("Dangerous tool")
8830 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8831 .build();
8832 let mut router = McpRouter::new().tool(safe).tool(dangerous);
8833 init_router(&mut router).await;
8834
8835 router.disable_tool("dangerous");
8836 assert!(router.is_tool_enabled("safe"));
8837 assert!(!router.is_tool_enabled("dangerous"));
8838
8839 let req = RouterRequest {
8840 id: RequestId::Number(1),
8841 inner: McpRequest::ListTools(ListToolsParams::default()),
8842 extensions: Extensions::new(),
8843 };
8844 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8845 match resp.inner {
8846 Ok(McpResponse::ListTools(result)) => {
8847 let names: Vec<&str> = result.tools.iter().map(|t| t.name.as_str()).collect();
8848 assert_eq!(names, vec!["safe"]);
8849 }
8850 _ => panic!("Expected ListTools response"),
8851 }
8852 }
8853
8854 #[tokio::test]
8855 async fn test_disable_tool_blocks_call() {
8856 let dangerous = ToolBuilder::new("dangerous")
8857 .description("Dangerous tool")
8858 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
8859 .build();
8860 let mut router = McpRouter::new().tool(dangerous);
8861 init_router(&mut router).await;
8862
8863 router.disable_tool("dangerous");
8864
8865 let req = RouterRequest {
8866 id: RequestId::Number(2),
8867 inner: McpRequest::CallTool(CallToolParams {
8868 input_responses: None,
8869 request_state: None,
8870 name: "dangerous".to_string(),
8871 arguments: serde_json::json!({"a": 1, "b": 2}),
8872 meta: None,
8873 task: None,
8874 }),
8875 extensions: Extensions::new(),
8876 };
8877 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8878 let err = resp.inner.expect_err("disabled tool should error");
8879 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
8880 }
8881
8882 #[tokio::test]
8883 async fn test_enable_tool_restores_visibility() {
8884 let tool = ToolBuilder::new("flippy")
8885 .description("Toggleable tool")
8886 .handler(|_: AddInput| async { Ok(CallToolResult::text("ran")) })
8887 .build();
8888 let mut router = McpRouter::new().tool(tool);
8889 init_router(&mut router).await;
8890
8891 router.disable_tool("flippy");
8892 router.enable_tool("flippy");
8893 assert!(router.is_tool_enabled("flippy"));
8894
8895 let req = RouterRequest {
8896 id: RequestId::Number(3),
8897 inner: McpRequest::CallTool(CallToolParams {
8898 input_responses: None,
8899 request_state: None,
8900 name: "flippy".to_string(),
8901 arguments: serde_json::json!({"a": 1, "b": 2}),
8902 meta: None,
8903 task: None,
8904 }),
8905 extensions: Extensions::new(),
8906 };
8907 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8908 match resp.inner {
8909 Ok(McpResponse::CallTool(result)) => {
8910 assert_eq!(result.first_text(), Some("ran"));
8911 }
8912 _ => panic!("Expected CallTool response"),
8913 }
8914 }
8915
8916 #[tokio::test]
8917 async fn test_disable_propagates_through_fresh_session() {
8918 let tool = ToolBuilder::new("shared")
8919 .description("Shared across sessions")
8920 .handler(|_: AddInput| async { Ok(CallToolResult::text("ok")) })
8921 .build();
8922 let router = McpRouter::new().tool(tool);
8923
8924 router.disable_tool("shared");
8926 let mut child = router.with_fresh_session();
8927 init_router(&mut child).await;
8928 assert!(!child.is_tool_enabled("shared"));
8929
8930 let req = RouterRequest {
8931 id: RequestId::Number(4),
8932 inner: McpRequest::ListTools(ListToolsParams::default()),
8933 extensions: Extensions::new(),
8934 };
8935 let resp = child.ready().await.unwrap().call(req).await.unwrap();
8936 match resp.inner {
8937 Ok(McpResponse::ListTools(result)) => {
8938 assert!(result.tools.is_empty());
8939 }
8940 _ => panic!("Expected ListTools response"),
8941 }
8942 }
8943
8944 #[tokio::test]
8945 async fn test_disable_resource_and_prompt() {
8946 let resource = crate::resource::ResourceBuilder::new("file:///hidden.txt")
8947 .name("hidden")
8948 .text("secret");
8949 let prompt = crate::prompt::PromptBuilder::new("hidden_prompt")
8950 .description("hidden")
8951 .user_message("hello");
8952
8953 let mut router = McpRouter::new().resource(resource).prompt(prompt);
8954 init_router(&mut router).await;
8955
8956 router.disable_resource("file:///hidden.txt");
8957 router.disable_prompt("hidden_prompt");
8958 assert!(!router.is_resource_enabled("file:///hidden.txt"));
8959 assert!(!router.is_prompt_enabled("hidden_prompt"));
8960
8961 let req = RouterRequest {
8963 id: RequestId::Number(5),
8964 inner: McpRequest::ListResources(ListResourcesParams::default()),
8965 extensions: Extensions::new(),
8966 };
8967 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8968 match resp.inner {
8969 Ok(McpResponse::ListResources(result)) => {
8970 assert!(result.resources.is_empty());
8971 }
8972 _ => panic!("Expected ListResources response"),
8973 }
8974
8975 let req = RouterRequest {
8977 id: RequestId::Number(6),
8978 inner: McpRequest::ReadResource(ReadResourceParams {
8979 input_responses: None,
8980 request_state: None,
8981 uri: "file:///hidden.txt".to_string(),
8982 meta: None,
8983 }),
8984 extensions: Extensions::new(),
8985 };
8986 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8987 let err = resp.inner.expect_err("disabled resource should error");
8988 assert_eq!(err.code, -32602); let req = RouterRequest {
8992 id: RequestId::Number(7),
8993 inner: McpRequest::ListPrompts(ListPromptsParams::default()),
8994 extensions: Extensions::new(),
8995 };
8996 let resp = router.ready().await.unwrap().call(req).await.unwrap();
8997 match resp.inner {
8998 Ok(McpResponse::ListPrompts(result)) => {
8999 assert!(result.prompts.is_empty());
9000 }
9001 _ => panic!("Expected ListPrompts response"),
9002 }
9003
9004 let req = RouterRequest {
9006 id: RequestId::Number(8),
9007 inner: McpRequest::GetPrompt(GetPromptParams {
9008 input_responses: None,
9009 request_state: None,
9010 name: "hidden_prompt".to_string(),
9011 arguments: Default::default(),
9012 meta: None,
9013 }),
9014 extensions: Extensions::new(),
9015 };
9016 let resp = router.ready().await.unwrap().call(req).await.unwrap();
9017 let err = resp.inner.expect_err("disabled prompt should error");
9018 assert_eq!(err.code, crate::error::ErrorCode::MethodNotFound as i32);
9019 }
9020
9021 #[test]
9022 fn test_router_request_new() {
9023 let req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9024 assert_eq!(req.id, RequestId::Number(1));
9025 assert!(req.extensions.is_empty());
9026 }
9027
9028 #[test]
9029 fn test_with_inner_preserves_extensions() {
9030 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9031 req.extensions.insert(42u32);
9032
9033 let rewritten = req.with_inner(McpRequest::ListTools(Default::default()));
9034 assert!(matches!(rewritten.inner, McpRequest::ListTools(_)));
9035 assert_eq!(rewritten.id, RequestId::Number(1));
9036 assert_eq!(rewritten.extensions.get::<u32>(), Some(&42));
9037 }
9038
9039 #[test]
9040 fn test_with_id_and_inner_preserves_extensions() {
9041 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9042 req.extensions.insert(String::from("token-abc"));
9043
9044 let rewritten = req.with_id_and_inner(
9045 RequestId::Number(99),
9046 McpRequest::ListResources(Default::default()),
9047 );
9048 assert_eq!(rewritten.id, RequestId::Number(99));
9049 assert!(matches!(rewritten.inner, McpRequest::ListResources(_)));
9050 assert_eq!(
9051 rewritten.extensions.get::<String>(),
9052 Some(&String::from("token-abc"))
9053 );
9054 }
9055
9056 #[test]
9057 fn test_clone_with_inner_preserves_extensions() {
9058 let mut req = RouterRequest::new(RequestId::Number(1), McpRequest::Ping);
9059 req.extensions.insert(true);
9060
9061 let cloned = req.clone_with_inner(McpRequest::ListTools(Default::default()));
9062
9063 assert!(matches!(req.inner, McpRequest::Ping));
9065 assert_eq!(req.extensions.get::<bool>(), Some(&true));
9066
9067 assert!(matches!(cloned.inner, McpRequest::ListTools(_)));
9069 assert_eq!(cloned.extensions.get::<bool>(), Some(&true));
9070 }
9071
9072 #[test]
9073 fn test_router_response_is_error() {
9074 let ok_resp = RouterResponse {
9075 id: RequestId::Number(1),
9076 inner: Ok(McpResponse::Pong(Default::default())),
9077 };
9078 assert!(!ok_resp.is_error());
9079
9080 let err_resp = RouterResponse {
9081 id: RequestId::Number(2),
9082 inner: Err(JsonRpcError::internal_error("boom")),
9083 };
9084 assert!(err_resp.is_error());
9085 }
9086
9087 #[test]
9088 fn test_extensions_len_and_is_empty() {
9089 let mut ext = Extensions::new();
9090 assert!(ext.is_empty());
9091 assert_eq!(ext.len(), 0);
9092
9093 ext.insert(42u32);
9094 assert!(!ext.is_empty());
9095 assert_eq!(ext.len(), 1);
9096
9097 ext.insert(String::from("hello"));
9098 assert_eq!(ext.len(), 2);
9099 }
9100
9101 #[test]
9102 fn test_router_response_serde_roundtrip() {
9103 let response = RouterResponse {
9105 id: RequestId::Number(1),
9106 inner: Ok(McpResponse::Empty(EmptyResult {})),
9107 };
9108 let json = serde_json::to_string(&response).unwrap();
9109 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9110 assert_eq!(deserialized.id, RequestId::Number(1));
9111 assert!(!deserialized.is_error());
9112
9113 let response = RouterResponse {
9115 id: RequestId::String("req-2".into()),
9116 inner: Err(JsonRpcError::method_not_found("unknown")),
9117 };
9118 let json = serde_json::to_string(&response).unwrap();
9119 let deserialized: RouterResponse = serde_json::from_str(&json).unwrap();
9120 assert_eq!(deserialized.id, RequestId::String("req-2".into()));
9121 assert!(deserialized.is_error());
9122 }
9123
9124 #[tokio::test]
9131 async fn test_discover_dispatch_via_jsonrpc_service() {
9132 let router = McpRouter::new().server_info("unit-test-server", "4.2.0");
9135 let mut service = JsonRpcService::new(router);
9136
9137 let req = JsonRpcRequest::new(1, "server/discover");
9138 let resp = service.call_single(req).await.unwrap();
9139
9140 match resp {
9141 JsonRpcResponse::Result(r) => {
9142 let versions = r
9144 .result
9145 .get("supportedVersions")
9146 .and_then(|v| v.as_array())
9147 .expect("result.supportedVersions must be an array");
9148 assert!(!versions.is_empty(), "supportedVersions must not be empty");
9149
9150 assert_eq!(
9152 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"],
9153 "unit-test-server",
9154 "serverInfo.name must match configured value"
9155 );
9156 assert_eq!(
9157 r.result["_meta"]["io.modelcontextprotocol/serverInfo"]["version"], "4.2.0",
9158 "serverInfo.version must match configured value"
9159 );
9160
9161 assert!(
9164 r.result.get("protocolVersion").is_none(),
9165 "server/discover must NOT include protocolVersion: {:?}",
9166 r.result
9167 );
9168 }
9169 JsonRpcResponse::Error(e) => panic!("Expected success, got error: {:?}", e),
9170 _ => panic!("unexpected response variant"),
9171 }
9172 }
9173
9174 #[tokio::test]
9175 async fn test_discover_does_not_require_initialization() {
9176 let router = McpRouter::new().server_info("fresh-router", "1.0.0");
9179 let mut service = JsonRpcService::new(router);
9180
9181 let req = JsonRpcRequest::new(2, "server/discover");
9182 let resp = service.call_single(req).await.unwrap();
9183
9184 assert!(
9186 !matches!(resp, JsonRpcResponse::Error(_)),
9187 "server/discover must not require initialization: {:?}",
9188 resp
9189 );
9190 }
9191}
9192
9193#[cfg(test)]
9194mod cursor_property_tests {
9195 use super::{decode_cursor, encode_cursor};
9196 use proptest::prelude::*;
9197
9198 fn arb_cursor_text() -> BoxedStrategy<String> {
9199 prop_oneof![
9200 8 => prop::collection::vec(any::<char>(), 0..512)
9201 .prop_map(|chars| chars.into_iter().collect()),
9202 1 => Just("\0\r\n\t\u{001b}\u{007f}".repeat(64)),
9203 1 => Just("A".repeat(16 * 1024)),
9204 ]
9205 .boxed()
9206 }
9207
9208 proptest! {
9209 #![proptest_config(ProptestConfig::with_cases(512))]
9210
9211 #[test]
9213 fn cursor_round_trips(offset in any::<usize>()) {
9214 prop_assert_eq!(decode_cursor(&encode_cursor(offset)).unwrap(), offset);
9215 }
9216
9217 #[test]
9219 fn decode_cursor_never_panics(s in arb_cursor_text()) {
9220 let _ = decode_cursor(&s);
9221 }
9222 }
9223}