1use std::borrow::Cow;
60use std::collections::HashMap;
61use std::sync::{
62 Arc,
63 atomic::{AtomicU64, Ordering},
64};
65use std::time::Duration;
66
67use rmcp::ServiceExt;
68use rmcp::model::{
69 CallToolRequest, CallToolResult, ClientRequest, ContentBlock, ListToolsRequest,
70 PaginatedRequestParams, ResourceContents, ServerResult,
71};
72use rmcp::service::PeerRequestOptions;
73use tokio::sync::{Mutex, RwLock};
74
75use crate::tool::ErasedTool;
76use crate::tool::server::{ManagedToolToken, ToolServerHandle};
77use crate::tool::{ToolContext, ToolExecutionError, ToolOutput, ToolResult};
78use rig_core::OneOrMany;
79use rig_core::message::{ImageMediaType, MimeType, ToolResultContent};
80use rig_core::wasm_compat::WasmBoxedFuture;
81
82pub use rmcp::model::Meta;
85
86pub const DEFAULT_MCP_TOOL_TIMEOUT: Duration = Duration::from_secs(300);
94
95pub const DEFAULT_MCP_REFRESH_TIMEOUT: Duration = Duration::from_secs(30);
100
101const MCP_CANCELLATION_GRACE_PERIOD: Duration = Duration::from_secs(1);
104
105#[derive(Clone)]
107pub(crate) struct McpTool {
108 definition: rmcp::model::Tool,
109 client: rmcp::service::ServerSink,
110 timeout: Option<Duration>,
117}
118
119impl McpTool {
120 pub(crate) fn from_mcp_server(
125 definition: rmcp::model::Tool,
126 client: rmcp::service::ServerSink,
127 ) -> Self {
128 Self {
129 definition,
130 client,
131 timeout: Some(DEFAULT_MCP_TOOL_TIMEOUT),
132 }
133 }
134
135 pub(crate) fn with_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
143 self.timeout = timeout.into();
144 self
145 }
146
147 #[cfg(test)]
149 pub(crate) fn timeout(&self) -> Option<Duration> {
150 self.timeout
151 }
152}
153
154#[derive(Debug, thiserror::Error)]
158enum McpArgumentError {
159 #[error("invalid JSON: {0}")]
161 Json(#[from] serde_json::Error),
162 #[error("expected a JSON object or null, got {0}")]
164 NonObject(&'static str),
165}
166
167fn json_value_kind(value: &serde_json::Value) -> &'static str {
168 match value {
169 serde_json::Value::Null => "null",
170 serde_json::Value::Bool(_) => "boolean",
171 serde_json::Value::Number(_) => "number",
172 serde_json::Value::String(_) => "string",
173 serde_json::Value::Array(_) => "array",
174 serde_json::Value::Object(_) => "object",
175 }
176}
177
178fn parse_mcp_arguments(args: &str) -> Result<Option<rmcp::model::JsonObject>, McpArgumentError> {
183 let trimmed = args.trim();
184 if trimmed.is_empty() {
185 return Ok(None);
186 }
187 let value: serde_json::Value = serde_json::from_str(trimmed)?;
188 match value {
189 serde_json::Value::Null => Ok(None),
190 serde_json::Value::Object(_) => Ok(Some(serde_json::from_value(value)?)),
191 value => Err(McpArgumentError::NonObject(json_value_kind(&value))),
192 }
193}
194
195async fn call_mcp_tool(
196 peer: &rmcp::service::ServerSink,
197 params: rmcp::model::CallToolRequestParams,
198 timeout: Option<Duration>,
199) -> Result<CallToolResult, rmcp::ServiceError> {
200 let deadline = timeout.map(|timeout| (tokio::time::Instant::now() + timeout, timeout));
201 let response = send_mcp_request(
202 peer,
203 ClientRequest::CallToolRequest(CallToolRequest::new(params)),
204 deadline,
205 )
206 .await?;
207
208 match response {
209 ServerResult::CallToolResult(result) => Ok(result),
210 _ => Err(rmcp::ServiceError::UnexpectedResponse),
211 }
212}
213
214async fn send_mcp_request(
215 peer: &rmcp::service::ServerSink,
216 request: ClientRequest,
217 deadline: Option<(tokio::time::Instant, Duration)>,
218) -> Result<ServerResult, rmcp::ServiceError> {
219 let handle = match deadline {
220 Some((deadline, timeout)) => {
221 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
222 if remaining.is_zero() {
223 return Err(rmcp::ServiceError::Timeout { timeout });
224 }
225 rig_core::wasm_compat::timeout(
226 remaining,
227 peer.send_cancellable_request(request, PeerRequestOptions::no_options()),
228 )
229 .await
230 .map_err(|_| rmcp::ServiceError::Timeout { timeout })??
231 }
232 None => {
233 peer.send_cancellable_request(request, PeerRequestOptions::no_options())
234 .await?
235 }
236 };
237
238 let Some((deadline, timeout)) = deadline else {
239 return handle.await_response().await;
240 };
241 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
242 let mut handle = handle;
243 match rig_core::wasm_compat::timeout(remaining, &mut handle.rx).await {
244 Ok(response) => response.map_err(|_| rmcp::ServiceError::TransportClosed)?,
245 Err(_) => {
246 cancel_timed_out_request(handle);
247 Err(rmcp::ServiceError::Timeout { timeout })
248 }
249 }
250}
251
252fn cancel_timed_out_request(handle: rmcp::service::RequestHandle<rmcp::service::RoleClient>) {
258 let cancellation = async move {
259 bounded_best_effort_cancellation(
260 handle.cancel(Some(
261 rmcp::service::RequestHandle::<rmcp::service::RoleClient>::REQUEST_TIMEOUT_REASON
262 .to_owned(),
263 )),
264 MCP_CANCELLATION_GRACE_PERIOD,
265 )
266 .await;
267 };
268
269 tokio::spawn(cancellation);
273}
274
275async fn bounded_best_effort_cancellation(
276 cancellation: impl std::future::Future<Output = Result<(), rmcp::ServiceError>>,
277 grace_period: Duration,
278) {
279 let _ = rig_core::wasm_compat::timeout(grace_period, cancellation).await;
280}
281
282impl McpTool {
283 fn execute_mcp(
291 &self,
292 args: String,
293 meta: Option<rmcp::model::Meta>,
294 ) -> WasmBoxedFuture<'_, Result<CallToolResult, ToolExecutionError>> {
295 let name = self.definition.name.clone();
296
297 Box::pin(async move {
298 let arguments = parse_mcp_arguments(&args).map_err(|error| {
301 ToolExecutionError::invalid_args(format!(
302 "MCP tool '{name}' received invalid arguments: {error}"
303 ))
304 .with_source(error)
305 })?;
306 let mut request = arguments
307 .map(|arguments| {
308 rmcp::model::CallToolRequestParams::new(name.clone()).with_arguments(arguments)
309 })
310 .unwrap_or_else(|| rmcp::model::CallToolRequestParams::new(name));
311 request.meta = meta;
312
313 match call_mcp_tool(&self.client, request, self.timeout).await {
314 Ok(result) => Ok(result),
315 Err(
316 error @ rmcp::ServiceError::Timeout {
317 timeout: elapsed_timeout,
318 },
319 ) => {
320 let timeout = self.timeout.unwrap_or(elapsed_timeout);
321 Err(ToolExecutionError::timeout(format!(
322 "MCP tool '{}' timed out after {timeout:?}",
323 self.definition.name
324 ))
325 .with_source(error))
326 }
327 Err(error) => Err(ToolExecutionError::provider(format!(
329 "MCP tool '{}' request failed: {error}",
330 self.definition.name
331 ))
332 .with_source(error)),
333 }
334 })
335 }
336}
337
338fn mcp_content_block_as_json(
339 content: &ContentBlock,
340) -> Result<ToolResultContent, ToolExecutionError> {
341 serde_json::to_value(content)
342 .map(ToolResultContent::json)
343 .map_err(|error| {
344 ToolExecutionError::provider(format!(
345 "failed to preserve an MCP content block as JSON: {error}"
346 ))
347 .with_source(error)
348 })
349}
350
351fn mcp_content_block_to_tool_content(
352 content: &ContentBlock,
353) -> Result<ToolResultContent, ToolExecutionError> {
354 match content {
355 ContentBlock::Text(text) => Ok(ToolResultContent::text(text.text.clone())),
356 ContentBlock::Image(image) => match ImageMediaType::from_mime_type(&image.mime_type) {
357 Some(media_type) => Ok(ToolResultContent::image_base64(
358 image.data.clone(),
359 Some(media_type),
360 None,
361 )),
362 None => mcp_content_block_as_json(content),
363 },
364 ContentBlock::Resource(resource) => match &resource.resource {
365 ResourceContents::TextResourceContents { .. } => mcp_content_block_as_json(content),
369 ResourceContents::BlobResourceContents {
370 mime_type, blob, ..
371 } => match mime_type
372 .as_deref()
373 .and_then(ImageMediaType::from_mime_type)
374 {
375 Some(media_type) => Ok(ToolResultContent::image_base64(
376 blob.clone(),
377 Some(media_type),
378 None,
379 )),
380 _ => mcp_content_block_as_json(content),
381 },
382 _ => mcp_content_block_as_json(content),
383 },
384 ContentBlock::ResourceLink(_) | ContentBlock::Audio(_) => {
385 mcp_content_block_as_json(content)
386 }
387 _ => mcp_content_block_as_json(content),
390 }
391}
392
393fn mcp_result_output(result: &CallToolResult) -> Result<ToolOutput, ToolExecutionError> {
395 let structured = result.structured_content.as_ref();
396 let canonical_fallback = structured.map(serde_json::Value::to_string);
397 let mut replaced_fallback = false;
398 let mut mapped = Vec::with_capacity(result.content.len());
399
400 for block in &result.content {
401 let fallback_structured = if !replaced_fallback {
402 match (block, canonical_fallback.as_deref(), structured) {
403 (ContentBlock::Text(text), Some(fallback), Some(structured))
404 if text.text == fallback =>
405 {
406 Some(structured)
407 }
408 _ => None,
409 }
410 } else {
411 None
412 };
413 if let Some(structured) = fallback_structured {
414 mapped.push(ToolResultContent::json(structured.clone()));
418 replaced_fallback = true;
419 } else {
420 mapped.push(mcp_content_block_to_tool_content(block)?);
421 }
422 }
423
424 if let Some(structured) = structured
425 && !replaced_fallback
426 {
427 mapped.insert(0, ToolResultContent::json(structured.clone()));
432 }
433
434 let mut mapped = mapped.into_iter();
435 if let Some(first) = mapped.next() {
436 let mut ordered = OneOrMany::one(first);
437 for block in mapped {
438 ordered.push(block);
439 }
440 return Ok(ToolOutput::content(ordered));
441 }
442
443 if result.is_error == Some(true) {
444 Ok(ToolOutput::text("the MCP tool reported an error"))
445 } else {
446 Ok(ToolOutput::text(""))
447 }
448}
449
450fn preserve_mcp_result(context: &mut ToolContext, result: &CallToolResult) {
451 if let Some(structured) = result.structured_content.clone() {
452 context.insert_result(structured);
453 }
454 if let Some(meta) = result.meta.clone() {
455 context.insert_result(meta);
456 }
457 context.insert_result(result.clone());
458}
459
460impl ErasedTool for McpTool {
461 fn name(&self) -> String {
462 self.definition.name.to_string()
463 }
464
465 fn description(&self) -> String {
466 self.definition
467 .description
468 .clone()
469 .unwrap_or(Cow::from(""))
470 .to_string()
471 }
472
473 fn parameters(&self) -> serde_json::Value {
474 self.definition.schema_as_json_value()
475 }
476
477 fn is_live(&self) -> bool {
478 !self.client.is_transport_closed()
479 }
480
481 fn execute<'a>(
482 &'a self,
483 args: String,
484 context: &'a mut ToolContext,
485 ) -> WasmBoxedFuture<'a, ToolResult> {
486 let meta = context.get::<rmcp::model::Meta>().cloned();
487 Box::pin(async move {
488 match self.execute_mcp(args, meta).await {
489 Ok(result) => {
490 let is_error = result.is_error == Some(true);
491 preserve_mcp_result(context, &result);
492 let output = match mcp_result_output(&result) {
493 Ok(output) => output,
494 Err(error) => return ToolResult::failed(error),
495 };
496
497 if is_error {
498 ToolResult::failed(
499 ToolExecutionError::other(format!(
500 "MCP tool '{}' reported an execution error",
501 self.definition.name
502 ))
503 .with_model_output(output),
504 )
505 } else {
506 ToolResult::success(output)
507 }
508 }
509 Err(error) => ToolResult::failed(error),
510 }
511 })
512 }
513}
514
515#[derive(Debug, thiserror::Error)]
517pub enum McpClientError {
518 #[error("MCP connection error: {0}")]
520 ConnectionError(String),
521
522 #[error("Failed to fetch MCP tool list: {0}")]
524 ToolFetchError(#[from] rmcp::ServiceError),
525
526 #[error("Timed out fetching MCP tool list after {0:?}")]
528 ToolFetchTimeout(Duration),
529}
530
531#[derive(Default)]
532struct ManagedToolsState {
533 registrations: HashMap<String, ManagedToolToken>,
534 committed_refresh: u64,
535}
536
537#[derive(Default)]
538struct RefreshActivity {
539 active: usize,
540 dirty: bool,
541}
542
543const MAX_CONCURRENT_REFRESHES: usize = 2;
544
545pub struct McpClientHandler {
569 client_info: rmcp::model::ClientInfo,
570 tool_server_handle: ToolServerHandle,
571 timeout: Option<Duration>,
574 refresh_timeout: Duration,
576 managed_tools: Arc<RwLock<ManagedToolsState>>,
580 refresh_activity: Arc<Mutex<RefreshActivity>>,
582 next_refresh: Arc<AtomicU64>,
584}
585
586impl McpClientHandler {
587 pub fn new(client_info: rmcp::model::ClientInfo, tool_server_handle: ToolServerHandle) -> Self {
593 Self {
594 client_info,
595 tool_server_handle,
596 timeout: Some(DEFAULT_MCP_TOOL_TIMEOUT),
597 refresh_timeout: DEFAULT_MCP_REFRESH_TIMEOUT,
598 managed_tools: Arc::new(RwLock::new(ManagedToolsState::default())),
599 refresh_activity: Arc::new(Mutex::new(RefreshActivity::default())),
600 next_refresh: Arc::new(AtomicU64::new(0)),
601 }
602 }
603
604 pub fn with_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
609 self.timeout = timeout.into();
610 self
611 }
612
613 pub fn with_refresh_timeout(mut self, timeout: Duration) -> Self {
615 self.refresh_timeout = timeout;
616 self
617 }
618
619 fn build_tool(&self, tool: rmcp::model::Tool, client: rmcp::service::ServerSink) -> McpTool {
621 McpTool::from_mcp_server(tool, client).with_timeout(self.timeout)
622 }
623
624 fn begin_refresh(&self) -> u64 {
625 self.next_refresh.fetch_add(1, Ordering::SeqCst) + 1
626 }
627
628 async fn fetch_tools(
629 &self,
630 peer: &rmcp::service::ServerSink,
631 ) -> Result<Vec<Arc<dyn ErasedTool>>, McpClientError> {
632 let deadline = tokio::time::Instant::now() + self.refresh_timeout;
633 let mut tools = Vec::new();
634 let mut cursor = None;
635
636 loop {
637 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
638 if remaining.is_zero() {
639 return Err(McpClientError::ToolFetchTimeout(self.refresh_timeout));
640 }
641 let mut params = PaginatedRequestParams::default();
642 params.cursor = cursor;
643 let response = send_mcp_request(
644 peer,
645 ClientRequest::ListToolsRequest(ListToolsRequest::with_param(params)),
646 Some((deadline, self.refresh_timeout)),
647 )
648 .await
649 .map_err(|error| match error {
650 rmcp::ServiceError::Timeout { .. } => {
651 McpClientError::ToolFetchTimeout(self.refresh_timeout)
652 }
653 error => McpClientError::ToolFetchError(error),
654 })?;
655 let page = match response {
656 ServerResult::ListToolsResult(page) => page,
657 _ => {
658 return Err(McpClientError::ToolFetchError(
659 rmcp::ServiceError::UnexpectedResponse,
660 ));
661 }
662 };
663 tools.extend(page.tools);
664 cursor = page.next_cursor;
665 if cursor.is_none() {
666 break;
667 }
668 }
669
670 Ok(tools
671 .into_iter()
672 .map(|tool| Arc::new(self.build_tool(tool, peer.clone())) as Arc<dyn ErasedTool>)
673 .collect())
674 }
675
676 async fn try_start_refresh(&self) -> bool {
677 let mut activity = self.refresh_activity.lock().await;
678 if activity.active >= MAX_CONCURRENT_REFRESHES {
679 activity.dirty = true;
680 false
681 } else {
682 activity.active += 1;
683 true
684 }
685 }
686
687 async fn finish_or_restart_refresh(&self) -> bool {
688 let mut activity = self.refresh_activity.lock().await;
689 if activity.dirty {
690 activity.dirty = false;
691 true
692 } else {
693 activity.active -= 1;
694 false
695 }
696 }
697
698 async fn commit_initial(&self, refresh: u64, tools: Vec<Arc<dyn ErasedTool>>) {
699 let mut managed = self.managed_tools.write().await;
700 if refresh <= managed.committed_refresh {
701 tracing::debug!(refresh, "discarding stale initial MCP tool list");
702 return;
703 }
704 managed.registrations = self
705 .tool_server_handle
706 .add_managed_erased_tools(tools)
707 .await;
708 managed.committed_refresh = refresh;
709 }
710
711 async fn commit_refresh(&self, refresh: u64, tools: Vec<Arc<dyn ErasedTool>>) -> bool {
712 let mut managed = self.managed_tools.write().await;
713 if refresh <= managed.committed_refresh {
714 tracing::debug!(refresh, "discarding stale MCP tool-list response");
715 return false;
716 }
717 let expected = managed.registrations.clone();
718 managed.registrations = self
719 .tool_server_handle
720 .reconcile_managed_erased_tools(expected, tools)
721 .await;
722 managed.committed_refresh = refresh;
723 true
724 }
725
726 pub async fn connect<T, E, A>(
738 self,
739 transport: T,
740 ) -> Result<rmcp::service::RunningService<rmcp::service::RoleClient, Self>, McpClientError>
741 where
742 T: rmcp::transport::IntoTransport<rmcp::service::RoleClient, E, A>,
743 E: std::error::Error + Send + Sync + 'static,
744 {
745 let service = ServiceExt::serve(self, transport)
746 .await
747 .map_err(|e| McpClientError::ConnectionError(e.to_string()))?;
748
749 let handler = service.service();
750 let refresh = handler.begin_refresh();
751 let tools = handler.fetch_tools(service.peer()).await?;
752 handler.commit_initial(refresh, tools).await;
753
754 Ok(service)
755 }
756}
757
758impl rmcp::handler::client::ClientHandler for McpClientHandler {
759 fn get_info(&self) -> rmcp::model::ClientInfo {
760 self.client_info.clone()
761 }
762
763 async fn on_tool_list_changed(
764 &self,
765 context: rmcp::service::NotificationContext<rmcp::service::RoleClient>,
766 ) {
767 if !self.try_start_refresh().await {
768 return;
769 }
770
771 loop {
772 let refresh = self.begin_refresh();
773 match self.fetch_tools(&context.peer).await {
777 Ok(tools) => {
778 if self.commit_refresh(refresh, tools).await {
779 let tool_count = self.managed_tools.read().await.registrations.len();
780 tracing::info!(tool_count, "MCP tool list refreshed successfully");
781 }
782 }
783 Err(error) => tracing::error!("Failed to re-fetch MCP tool list: {error}"),
784 }
785
786 if !self.finish_or_restart_refresh().await {
787 break;
788 }
789 }
790 }
791}
792
793#[cfg(test)]
794mod tests {
795 use std::{
796 future::pending,
797 sync::{
798 Arc,
799 atomic::{AtomicBool, Ordering},
800 },
801 time::Duration,
802 };
803
804 use rmcp::model::*;
805 use rmcp::service::RequestContext;
806 use rmcp::{RoleServer, ServerHandler, ServiceExt};
807 use serde_json::json;
808 use tokio::{
809 sync::{Notify, RwLock},
810 task::JoinHandle,
811 };
812
813 use super::*;
814 use crate::tool::{
815 ToolErrorKind,
816 server::{ToolServer, ToolServerHandle},
817 };
818 use rig_core::message::ToolResultContent as RigToolResultContent;
819
820 #[derive(Clone)]
821 enum Scenario {
822 Success,
823 StructuredSuccess,
824 StructuredOnly,
825 Hang,
826 ServiceError,
827 ToolReportedError,
828 ImageToolReportedError,
829 }
830
831 #[derive(Clone)]
832 struct ScenarioServer {
833 scenario: Scenario,
834 seen: Arc<RwLock<Option<Meta>>>,
835 cancelled: Arc<Notify>,
836 }
837
838 impl ServerHandler for ScenarioServer {
839 fn get_info(&self) -> ServerInfo {
840 ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
841 .with_protocol_version(ProtocolVersion::LATEST)
842 .with_server_info(Implementation::new("rig-mcp-test", "0.1.0"))
843 }
844
845 async fn call_tool(
846 &self,
847 _request: CallToolRequestParams,
848 context: RequestContext<RoleServer>,
849 ) -> Result<CallToolResult, ErrorData> {
850 *self.seen.write().await = Some(context.meta.clone());
851 match self.scenario {
852 Scenario::Success => Ok(CallToolResult::success(vec![ContentBlock::text("ok")])),
853 Scenario::StructuredSuccess => {
854 let mut response = CallToolResult::success(vec![
855 ContentBlock::text("before"),
856 ContentBlock::image("aGVsbG8=", "image/png"),
857 ContentBlock::text("after"),
858 ]);
859 response.structured_content = Some(json!({
860 "answer": 42,
861 "source": "fixture"
862 }));
863 let mut meta = Meta::new();
864 meta.0.insert("response-id".into(), json!("response-123"));
865 response.meta = Some(meta);
866 Ok(response)
867 }
868 Scenario::StructuredOnly => {
869 let mut response = CallToolResult::structured(json!({"answer": 42}));
870 response.content.clear();
871 Ok(response)
872 }
873 Scenario::Hang => {
874 context.ct.cancelled().await;
875 self.cancelled.notify_one();
876 Err(ErrorData::internal_error("fixture request cancelled", None))
877 }
878 Scenario::ServiceError => {
879 Err(ErrorData::internal_error("fixture service failed", None))
880 }
881 Scenario::ToolReportedError => Ok(CallToolResult::error(vec![ContentBlock::text(
882 "tool reported exact failure",
883 )])),
884 Scenario::ImageToolReportedError => {
885 Ok(CallToolResult::error(vec![ContentBlock::image(
886 "ZXJyb3ItaW1hZ2U=",
887 "image/png",
888 )]))
889 }
890 }
891 }
892 }
893
894 struct Fixture {
895 handle: ToolServerHandle,
896 seen: Arc<RwLock<Option<Meta>>>,
897 cancelled: Arc<Notify>,
898 _client: rmcp::service::RunningService<rmcp::service::RoleClient, ClientInfo>,
899 server_task: JoinHandle<()>,
900 }
901
902 async fn fixture(scenario: Scenario, timeout: Option<Duration>) -> Fixture {
903 let seen = Arc::new(RwLock::new(None));
904 let cancelled = Arc::new(Notify::new());
905 let (client_to_server, server_from_client) = tokio::io::duplex(8192);
906 let (server_to_client, client_from_server) = tokio::io::duplex(8192);
907 let server = ScenarioServer {
908 scenario,
909 seen: seen.clone(),
910 cancelled: cancelled.clone(),
911 };
912 let server_task = tokio::spawn(async move {
913 let running = server
914 .serve((server_from_client, server_to_client))
915 .await
916 .expect("server start");
917 running.waiting().await.expect("server error");
918 });
919 let client = ClientInfo::default()
920 .serve((client_from_server, client_to_server))
921 .await
922 .expect("client connect");
923 let definition = Tool::new(
924 "fixture_tool".to_string(),
925 "fixture".to_string(),
926 Arc::new(serde_json::Map::new()),
927 );
928 let handle = ToolServer::new()
929 .rmcp_tool_with_timeout(definition, client.peer().clone(), timeout)
930 .run();
931 Fixture {
932 handle,
933 seen,
934 cancelled,
935 _client: client,
936 server_task,
937 }
938 }
939
940 async fn execute(fixture: &Fixture, args: &str, context: &mut ToolContext) -> ToolResult {
941 tokio::time::timeout(
942 Duration::from_secs(5),
943 fixture.handle.execute("fixture_tool", args, context),
944 )
945 .await
946 .expect("MCP dispatch exceeded the outer safety timeout")
947 }
948
949 #[tokio::test]
950 async fn best_effort_cancellation_drops_stalled_delivery_after_grace_period() {
951 struct DropProbe(Arc<AtomicBool>);
952
953 impl Drop for DropProbe {
954 fn drop(&mut self) {
955 self.0.store(true, Ordering::SeqCst);
956 }
957 }
958
959 let dropped = Arc::new(AtomicBool::new(false));
960 let drop_probe = DropProbe(dropped.clone());
961 let stalled = async move {
962 let _drop_probe = drop_probe;
963 pending::<Result<(), rmcp::ServiceError>>().await
964 };
965
966 tokio::time::timeout(
967 Duration::from_secs(1),
968 bounded_best_effort_cancellation(stalled, Duration::from_millis(10)),
969 )
970 .await
971 .expect("best-effort cancellation exceeded its grace period");
972
973 assert!(dropped.load(Ordering::SeqCst));
974 }
975
976 #[test]
977 fn model_presentation_preserves_unrepresentable_mcp_blocks_as_json() {
978 let blocks = vec![
979 ContentBlock::resource(ResourceContents::TextResourceContents {
980 uri: "file:///reports/summary.txt".to_string(),
981 mime_type: Some("text/plain".to_string()),
982 text: "full report".to_string(),
983 meta: None,
984 }),
985 ContentBlock::resource(ResourceContents::BlobResourceContents {
986 uri: "file:///reports/raw.bin".to_string(),
987 mime_type: Some("application/octet-stream".to_string()),
988 blob: "AAEC".to_string(),
989 meta: None,
990 }),
991 ContentBlock::audio("UklGRg==", "audio/wav"),
992 ContentBlock::resource_link(
993 Resource::new("file:///reports/linked.txt", "linked.txt")
994 .with_mime_type("text/plain"),
995 ),
996 ContentBlock::image("YXZpZg==", "image/avif"),
997 ContentBlock::resource(ResourceContents::BlobResourceContents {
998 uri: "file:///images/chart.avif".to_string(),
999 mime_type: Some("image/avif".to_string()),
1000 blob: "YmxvYi1hdmlm".to_string(),
1001 meta: None,
1002 }),
1003 ];
1004 let expected = blocks
1005 .iter()
1006 .map(|block| {
1007 RigToolResultContent::json(
1008 serde_json::to_value(block).expect("MCP block is JSON serializable"),
1009 )
1010 })
1011 .collect::<Vec<_>>();
1012
1013 let result = CallToolResult::success(blocks);
1014 let content = mcp_result_output(&result)
1015 .expect("MCP content mapping")
1016 .into_content()
1017 .into_iter()
1018 .collect::<Vec<_>>();
1019
1020 assert_eq!(content, expected);
1021 assert!(matches!(
1022 &content[0],
1023 RigToolResultContent::Json { value }
1024 if value["resource"]["uri"] == "file:///reports/summary.txt"
1025 && value["resource"]["mimeType"] == "text/plain"
1026 && value["resource"]["text"] == "full report"
1027 ));
1028 assert!(matches!(
1029 &content[1],
1030 RigToolResultContent::Json { value }
1031 if value["resource"]["uri"] == "file:///reports/raw.bin"
1032 && value["resource"]["mimeType"] == "application/octet-stream"
1033 && value["resource"]["blob"] == "AAEC"
1034 ));
1035 assert!(matches!(
1036 &content[2],
1037 RigToolResultContent::Json { value }
1038 if value["mimeType"] == "audio/wav" && value["data"] == "UklGRg=="
1039 ));
1040 assert!(matches!(
1041 &content[4],
1042 RigToolResultContent::Json { value }
1043 if value["mimeType"] == "image/avif" && value["data"] == "YXZpZg=="
1044 ));
1045 assert!(matches!(
1046 &content[5],
1047 RigToolResultContent::Json { value }
1048 if value["resource"]["uri"] == "file:///images/chart.avif"
1049 && value["resource"]["mimeType"] == "image/avif"
1050 && value["resource"]["blob"] == "YmxvYi1hdmlm"
1051 ));
1052 }
1053
1054 #[test]
1055 fn image_resource_blob_maps_to_an_image_block() {
1056 let result = CallToolResult::success(vec![ContentBlock::resource(
1057 ResourceContents::BlobResourceContents {
1058 uri: "file:///images/chart.png".to_string(),
1059 mime_type: Some("image/png".to_string()),
1060 blob: "aW1hZ2U=".to_string(),
1061 meta: None,
1062 },
1063 )]);
1064
1065 assert_eq!(
1066 mcp_result_output(&result).expect("MCP content mapping"),
1067 ToolOutput::one(RigToolResultContent::image_base64(
1068 "aW1hZ2U=",
1069 Some(ImageMediaType::PNG),
1070 None,
1071 ))
1072 );
1073 }
1074
1075 #[test]
1076 fn string_valued_structured_content_remains_json() {
1077 let mut result = CallToolResult::structured(json!("forty-two"));
1078 result.content.clear();
1079
1080 assert_eq!(
1081 mcp_result_output(&result).expect("MCP content mapping"),
1082 ToolOutput::json(json!("forty-two"))
1083 );
1084 }
1085
1086 #[test]
1087 fn structured_constructors_replace_their_canonical_text_fallback() {
1088 let value = json!({"answer": 42});
1089 for result in [
1090 CallToolResult::structured(value.clone()),
1091 CallToolResult::structured_error(value.clone()),
1092 ] {
1093 assert_eq!(
1094 mcp_result_output(&result).expect("MCP structured output"),
1095 ToolOutput::json(value.clone())
1096 );
1097 }
1098 }
1099
1100 #[test]
1101 fn structured_content_is_kept_alongside_real_rich_blocks() {
1102 let value = json!({"answer": 42});
1103 let mut result = CallToolResult::structured(value.clone());
1104 result
1105 .content
1106 .push(ContentBlock::image("aW1hZ2U=", "image/png"));
1107 result
1108 .content
1109 .push(ContentBlock::text("human-readable note"));
1110
1111 let mut expected = OneOrMany::one(RigToolResultContent::json(value));
1112 expected.push(RigToolResultContent::image_base64(
1113 "aW1hZ2U=",
1114 Some(ImageMediaType::PNG),
1115 None,
1116 ));
1117 expected.push(RigToolResultContent::text("human-readable note"));
1118 assert_eq!(
1119 mcp_result_output(&result).expect("MCP structured rich output"),
1120 ToolOutput::content(expected)
1121 );
1122 }
1123
1124 #[tokio::test]
1125 async fn canonical_dispatch_forwards_context_meta() {
1126 let fixture = fixture(Scenario::Success, Some(Duration::from_secs(1))).await;
1127 let mut meta = Meta::new();
1128 meta.0.insert("authorization".into(), json!("Bearer test"));
1129 let mut context = ToolContext::new();
1130 context.insert(meta);
1131
1132 let result = execute(&fixture, "{}", &mut context).await;
1133 assert!(result.is_success());
1134 assert_eq!(
1135 fixture
1136 .seen
1137 .read()
1138 .await
1139 .as_ref()
1140 .expect("server observed metadata")
1141 .0
1142 .get("authorization"),
1143 Some(&json!("Bearer test"))
1144 );
1145 fixture.server_task.abort();
1146 }
1147
1148 #[tokio::test]
1149 async fn canonical_dispatch_classifies_timeout() {
1150 let fixture = fixture(Scenario::Hang, Some(Duration::from_millis(25))).await;
1151 let result = execute(&fixture, "{}", &mut ToolContext::new()).await;
1152 assert!(result.is_error_kind(ToolErrorKind::Timeout));
1153 assert_eq!(
1154 result.output().as_text(),
1155 Some("MCP tool 'fixture_tool' timed out after 25ms")
1156 );
1157 tokio::time::timeout(Duration::from_secs(1), fixture.cancelled.notified())
1158 .await
1159 .expect("the timed-out MCP request should be cancelled at the peer");
1160 fixture.server_task.abort();
1161 }
1162
1163 #[tokio::test]
1164 async fn canonical_dispatch_classifies_service_error_and_preserves_source() {
1165 let fixture = fixture(Scenario::ServiceError, Some(Duration::from_secs(1))).await;
1166 let result = execute(&fixture, "{}", &mut ToolContext::new()).await;
1167 let error = result.error().expect("structured MCP service error");
1168 assert_eq!(error.kind(), ToolErrorKind::Provider);
1169 assert!(error.is::<rmcp::ServiceError>());
1170 assert!(error.message().contains("fixture service failed"));
1171 let output = result.output().render();
1172 assert!(output.contains("MCP tool 'fixture_tool' request failed"));
1173 assert!(output.contains("fixture service failed"));
1174 fixture.server_task.abort();
1175 }
1176
1177 #[tokio::test]
1178 async fn canonical_dispatch_preserves_tool_reported_error_message() {
1179 let fixture = fixture(Scenario::ToolReportedError, Some(Duration::from_secs(1))).await;
1180 let result = execute(&fixture, "{}", &mut ToolContext::new()).await;
1181 assert!(result.is_error_kind(ToolErrorKind::Other));
1182 assert_eq!(
1183 result.output(),
1184 &ToolOutput::one(RigToolResultContent::text("tool reported exact failure"))
1185 );
1186 assert_eq!(
1187 result.error().map(ToolExecutionError::message),
1188 Some("MCP tool 'fixture_tool' reported an execution error")
1189 );
1190 fixture.server_task.abort();
1191 }
1192
1193 #[tokio::test]
1194 async fn canonical_dispatch_preserves_non_text_tool_error_content() {
1195 let fixture = fixture(
1196 Scenario::ImageToolReportedError,
1197 Some(Duration::from_secs(1)),
1198 )
1199 .await;
1200 let mut context = ToolContext::new();
1201 let result = execute(&fixture, "{}", &mut context).await;
1202
1203 assert!(result.is_error_kind(ToolErrorKind::Other));
1204 assert_eq!(
1205 result.output(),
1206 &ToolOutput::one(RigToolResultContent::image_base64(
1207 "ZXJyb3ItaW1hZ2U=",
1208 Some(ImageMediaType::PNG),
1209 None,
1210 ))
1211 );
1212 let raw = context
1213 .result::<CallToolResult>()
1214 .expect("raw MCP error result metadata");
1215 assert_eq!(raw.is_error, Some(true));
1216 assert!(matches!(raw.content.as_slice(), [ContentBlock::Image(_)]));
1217 fixture.server_task.abort();
1218 }
1219
1220 #[tokio::test]
1221 async fn canonical_dispatch_preserves_ordered_content_and_response_metadata() {
1222 let fixture = fixture(Scenario::StructuredSuccess, Some(Duration::from_secs(1))).await;
1223 let mut context = ToolContext::new();
1224 let result = execute(&fixture, "{}", &mut context).await;
1225
1226 let mut expected_content = OneOrMany::one(RigToolResultContent::json(json!({
1227 "answer": 42,
1228 "source": "fixture"
1229 })));
1230 expected_content.push(RigToolResultContent::text("before"));
1231 expected_content.push(RigToolResultContent::image_base64(
1232 "aGVsbG8=",
1233 Some(ImageMediaType::PNG),
1234 None,
1235 ));
1236 expected_content.push(RigToolResultContent::text("after"));
1237 assert_eq!(result.output(), &ToolOutput::content(expected_content));
1238
1239 let raw = context
1240 .result::<CallToolResult>()
1241 .expect("raw MCP result metadata");
1242 assert_eq!(raw.content.len(), 3);
1243 assert_eq!(
1244 raw.structured_content,
1245 Some(json!({"answer": 42, "source": "fixture"}))
1246 );
1247 assert_eq!(
1248 context.result::<serde_json::Value>(),
1249 Some(&json!({"answer": 42, "source": "fixture"}))
1250 );
1251 assert_eq!(
1252 context
1253 .result::<Meta>()
1254 .and_then(|meta| meta.0.get("response-id")),
1255 Some(&json!("response-123"))
1256 );
1257 fixture.server_task.abort();
1258 }
1259
1260 #[tokio::test]
1261 async fn canonical_dispatch_uses_structured_content_when_blocks_are_empty() {
1262 let fixture = fixture(Scenario::StructuredOnly, Some(Duration::from_secs(1))).await;
1263 let mut context = ToolContext::new();
1264 let result = execute(&fixture, "{}", &mut context).await;
1265
1266 assert_eq!(result.output(), &ToolOutput::json(json!({"answer": 42})));
1267 assert_eq!(
1268 context.result::<serde_json::Value>(),
1269 Some(&json!({"answer": 42}))
1270 );
1271 fixture.server_task.abort();
1272 }
1273
1274 #[tokio::test]
1275 async fn canonical_dispatch_classifies_invalid_json_and_preserves_source() {
1276 let fixture = fixture(Scenario::Success, Some(Duration::from_secs(1))).await;
1277 let result = execute(&fixture, "{", &mut ToolContext::new()).await;
1278 let error = result.error().expect("structured argument error");
1279 assert_eq!(error.kind(), ToolErrorKind::InvalidArgs);
1280 assert!(matches!(
1281 error.downcast_ref::<McpArgumentError>(),
1282 Some(McpArgumentError::Json(_))
1283 ));
1284 let output = result.output().render();
1285 assert!(output.contains("MCP tool 'fixture_tool' received invalid arguments"));
1286 assert!(output.contains("invalid JSON"));
1287 fixture.server_task.abort();
1288 }
1289
1290 #[tokio::test]
1291 async fn canonical_dispatch_rejects_non_object_arguments() {
1292 let fixture = fixture(Scenario::Success, Some(Duration::from_secs(1))).await;
1293 for args in [r#"[1,2]"#, r#""text""#, "7", "true"] {
1294 let result = execute(&fixture, args, &mut ToolContext::new()).await;
1295 assert!(
1296 result.is_error_kind(ToolErrorKind::InvalidArgs),
1297 "{args} must not be coerced into an argument-less MCP call"
1298 );
1299 }
1300
1301 for args in ["", "null"] {
1303 let result = execute(&fixture, args, &mut ToolContext::new()).await;
1304 assert!(
1305 result.is_success(),
1306 "{args:?} should remain a no-argument call"
1307 );
1308 }
1309 fixture.server_task.abort();
1310 }
1311}
1312
1313#[cfg(test)]
1314mod migrated_tests {
1315 use super::{MAX_CONCURRENT_REFRESHES, McpClientError, McpClientHandler};
1316 use crate::tool::{DynamicTool, ToolOutput, server::ToolServer};
1317 use rmcp::{
1318 RoleServer, ServerHandler, ServiceExt, handler::client::ClientHandler, model::*,
1319 service::RequestContext,
1320 };
1321 use std::{
1322 sync::{
1323 Arc,
1324 atomic::{AtomicUsize, Ordering},
1325 },
1326 time::Duration,
1327 };
1328 use tokio::sync::{Notify, RwLock};
1329
1330 #[derive(Clone)]
1331 struct DynamicToolServer {
1332 tools: Arc<RwLock<Vec<Tool>>>,
1333 }
1334 impl DynamicToolServer {
1335 fn new(tools: Vec<Tool>) -> Self {
1336 Self {
1337 tools: Arc::new(RwLock::new(tools)),
1338 }
1339 }
1340 async fn set_tools(&self, tools: Vec<Tool>) {
1341 *self.tools.write().await = tools;
1342 }
1343 }
1344 impl ServerHandler for DynamicToolServer {
1345 fn get_info(&self) -> ServerInfo {
1346 ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
1347 .with_protocol_version(ProtocolVersion::LATEST)
1348 .with_server_info(Implementation::new("test-dynamic-server", "0.1.0"))
1349 }
1350 async fn list_tools(
1351 &self,
1352 _: Option<PaginatedRequestParams>,
1353 _: RequestContext<RoleServer>,
1354 ) -> Result<ListToolsResult, ErrorData> {
1355 Ok(ListToolsResult::with_all_items(
1356 self.tools.read().await.clone(),
1357 ))
1358 }
1359 async fn call_tool(
1360 &self,
1361 request: CallToolRequestParams,
1362 _: RequestContext<RoleServer>,
1363 ) -> Result<CallToolResult, ErrorData> {
1364 Ok(CallToolResult::success(vec![ContentBlock::text(format!(
1365 "called {}",
1366 request.name
1367 ))]))
1368 }
1369 }
1370
1371 #[derive(Clone)]
1372 struct OrderedRefreshServer {
1373 tools: Arc<RwLock<Vec<Tool>>>,
1374 list_calls: Arc<AtomicUsize>,
1375 first_refresh_started: Arc<Notify>,
1376 release_first_refresh: Arc<Notify>,
1377 first_refresh_returned: Arc<Notify>,
1378 }
1379
1380 impl OrderedRefreshServer {
1381 fn new(tools: Vec<Tool>) -> Self {
1382 Self {
1383 tools: Arc::new(RwLock::new(tools)),
1384 list_calls: Arc::new(AtomicUsize::new(0)),
1385 first_refresh_started: Arc::new(Notify::new()),
1386 release_first_refresh: Arc::new(Notify::new()),
1387 first_refresh_returned: Arc::new(Notify::new()),
1388 }
1389 }
1390
1391 async fn set_tools(&self, tools: Vec<Tool>) {
1392 *self.tools.write().await = tools;
1393 }
1394 }
1395
1396 impl ServerHandler for OrderedRefreshServer {
1397 fn get_info(&self) -> ServerInfo {
1398 ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
1399 .with_protocol_version(ProtocolVersion::LATEST)
1400 .with_server_info(Implementation::new("test-ordered-refresh-server", "0.1.0"))
1401 }
1402
1403 async fn list_tools(
1404 &self,
1405 _: Option<PaginatedRequestParams>,
1406 _: RequestContext<RoleServer>,
1407 ) -> Result<ListToolsResult, ErrorData> {
1408 let call = self.list_calls.fetch_add(1, Ordering::SeqCst);
1409 let tools = self.tools.read().await.clone();
1410
1411 if call == 1 {
1414 self.first_refresh_started.notify_one();
1415 self.release_first_refresh.notified().await;
1416 self.first_refresh_returned.notify_one();
1417 }
1418
1419 Ok(ListToolsResult::with_all_items(tools))
1420 }
1421 }
1422
1423 #[derive(Clone)]
1424 struct HangingListServer;
1425
1426 impl ServerHandler for HangingListServer {
1427 fn get_info(&self) -> ServerInfo {
1428 ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
1429 .with_protocol_version(ProtocolVersion::LATEST)
1430 .with_server_info(Implementation::new("test-hanging-list-server", "0.1.0"))
1431 }
1432
1433 async fn list_tools(
1434 &self,
1435 _: Option<PaginatedRequestParams>,
1436 _: RequestContext<RoleServer>,
1437 ) -> Result<ListToolsResult, ErrorData> {
1438 std::future::pending().await
1439 }
1440 }
1441
1442 fn make_tool(name: &str, description: &str) -> Tool {
1443 Tool::new(
1444 name.to_string(),
1445 description.to_string(),
1446 Arc::new(serde_json::Map::new()),
1447 )
1448 }
1449
1450 fn make_dynamic_tool(name: &str, description: &str) -> DynamicTool {
1451 DynamicTool::new(
1452 name,
1453 description,
1454 serde_json::json!({"type": "object", "properties": {}}),
1455 |_context, _args| Box::pin(async { Ok(ToolOutput::text("local")) }),
1456 )
1457 }
1458
1459 async fn connect<S>(
1460 server: S,
1461 handle: crate::tool::server::ToolServerHandle,
1462 ) -> (
1463 rmcp::service::RunningService<rmcp::RoleClient, McpClientHandler>,
1464 tokio::task::JoinHandle<rmcp::service::RunningService<rmcp::RoleServer, S>>,
1465 )
1466 where
1467 S: ServerHandler,
1468 {
1469 let (c2s, sfc) = tokio::io::duplex(8192);
1470 let (s2c, cfs) = tokio::io::duplex(8192);
1471 let server_task =
1472 tokio::spawn(async move { server.serve((sfc, s2c)).await.expect("server start") });
1473 let service = McpClientHandler::new(ClientInfo::default(), handle)
1474 .connect((cfs, c2s))
1475 .await
1476 .expect("connect");
1477 (service, server_task)
1478 }
1479
1480 #[tokio::test]
1481 async fn client_handler_registers_initial_tools() {
1482 let server = DynamicToolServer::new(vec![
1483 make_tool("tool_a", "First"),
1484 make_tool("tool_b", "Second"),
1485 ]);
1486 let handle = ToolServer::new().run();
1487 let (client, task) = connect(server, handle.clone()).await;
1488 let defs = handle.get_tool_defs(None).await.unwrap();
1489 assert_eq!(
1490 defs.iter().map(|d| d.name.as_str()).collect::<Vec<_>>(),
1491 vec!["tool_a", "tool_b"]
1492 );
1493 client.cancel().await.unwrap();
1494 task.abort();
1495 }
1496
1497 #[tokio::test]
1498 async fn disconnected_handler_tools_are_retired_on_snapshot() {
1499 let server = DynamicToolServer::new(vec![make_tool("tool_a", "First")]);
1500 let handle = ToolServer::new().run();
1501 let (client, task) = connect(server, handle.clone()).await;
1502 assert_eq!(handle.get_tool_defs(None).await.unwrap().len(), 1);
1503
1504 client.cancel().await.unwrap();
1505
1506 let defs = handle.get_tool_defs(None).await.unwrap();
1507 assert!(
1508 defs.is_empty(),
1509 "a disconnected sole owner must not remain provider-visible"
1510 );
1511 task.abort();
1512 }
1513
1514 #[tokio::test]
1515 async fn disconnected_handler_tools_are_retired_on_direct_dispatch() {
1516 let server = DynamicToolServer::new(vec![make_tool("tool_a", "First")]);
1517 let handle = ToolServer::new().run();
1518 let (client, task) = connect(server, handle.clone()).await;
1519 assert_eq!(handle.get_tool_defs(None).await.unwrap().len(), 1);
1520
1521 client.cancel().await.unwrap();
1522
1523 let result = handle
1524 .execute("tool_a", "{}", &mut crate::tool::ToolContext::new())
1525 .await;
1526 assert_eq!(
1527 result.error().expect("disconnected tool must fail").kind(),
1528 crate::tool::ToolErrorKind::NotFound
1529 );
1530 task.abort();
1531 }
1532
1533 #[tokio::test]
1534 async fn initial_tool_fetch_is_bounded_by_the_refresh_timeout() {
1535 let (c2s, sfc) = tokio::io::duplex(8192);
1536 let (s2c, cfs) = tokio::io::duplex(8192);
1537 let server_task = tokio::spawn(async move {
1538 HangingListServer
1539 .serve((sfc, s2c))
1540 .await
1541 .expect("server start")
1542 });
1543 let refresh_timeout = Duration::from_millis(25);
1544 let result = McpClientHandler::new(ClientInfo::default(), ToolServer::new().run())
1545 .with_refresh_timeout(refresh_timeout)
1546 .connect((cfs, c2s))
1547 .await;
1548
1549 assert!(matches!(
1550 result,
1551 Err(McpClientError::ToolFetchTimeout(timeout)) if timeout == refresh_timeout
1552 ));
1553 server_task.abort();
1554 }
1555
1556 #[tokio::test]
1557 async fn refresh_activity_is_bounded_and_coalesces_excess_notifications() {
1558 let handler = McpClientHandler::new(ClientInfo::default(), ToolServer::new().run());
1559
1560 assert!(handler.try_start_refresh().await);
1561 assert!(handler.try_start_refresh().await);
1562 assert!(!handler.try_start_refresh().await);
1563 {
1564 let activity = handler.refresh_activity.lock().await;
1565 assert_eq!(activity.active, MAX_CONCURRENT_REFRESHES);
1566 assert!(activity.dirty);
1567 }
1568
1569 assert!(handler.finish_or_restart_refresh().await);
1570 assert!(!handler.finish_or_restart_refresh().await);
1571 assert!(!handler.finish_or_restart_refresh().await);
1572 let activity = handler.refresh_activity.lock().await;
1573 assert_eq!(activity.active, 0);
1574 assert!(!activity.dirty);
1575 }
1576
1577 #[tokio::test]
1578 async fn client_handler_refreshes_on_tool_list_changed() {
1579 let server = DynamicToolServer::new(vec![make_tool("alpha", "Alpha")]);
1580 let handle = ToolServer::new().run();
1581 let (c2s, sfc) = tokio::io::duplex(8192);
1582 let (s2c, cfs) = tokio::io::duplex(8192);
1583 let copy = server.clone();
1584 let task = tokio::spawn(async move { copy.serve((sfc, s2c)).await.expect("server start") });
1585 let client = McpClientHandler::new(ClientInfo::default(), handle.clone())
1586 .connect((cfs, c2s))
1587 .await
1588 .unwrap();
1589 assert_eq!(handle.get_tool_defs(None).await.unwrap()[0].name, "alpha");
1590 server
1591 .set_tools(vec![make_tool("beta", "Beta"), make_tool("gamma", "Gamma")])
1592 .await;
1593 let running = task.await.unwrap();
1594 running.peer().notify_tool_list_changed().await.unwrap();
1595 tokio::time::timeout(Duration::from_secs(2), async {
1596 loop {
1597 let defs = handle.get_tool_defs(None).await.unwrap();
1598 if defs.len() == 2 {
1599 break;
1600 }
1601 tokio::task::yield_now().await;
1602 }
1603 })
1604 .await
1605 .expect("refresh");
1606 let names = handle
1607 .get_tool_defs(None)
1608 .await
1609 .unwrap()
1610 .into_iter()
1611 .map(|d| d.name)
1612 .collect::<Vec<_>>();
1613 assert_eq!(names, vec!["beta", "gamma"]);
1614 client.cancel().await.unwrap();
1615 }
1616
1617 #[tokio::test]
1618 async fn concurrent_refreshes_cannot_roll_back_a_newer_tool_list() {
1619 let server = OrderedRefreshServer::new(vec![make_tool("stale", "Stale snapshot")]);
1620 let server_control = server.clone();
1621 let handle = ToolServer::new().run();
1622 let (client, server_task) = connect(server, handle.clone()).await;
1623 let running_server = server_task.await.unwrap();
1624
1625 running_server
1626 .peer()
1627 .notify_tool_list_changed()
1628 .await
1629 .unwrap();
1630 tokio::time::timeout(
1631 Duration::from_secs(2),
1632 server_control.first_refresh_started.notified(),
1633 )
1634 .await
1635 .expect("first refresh fetch started");
1636
1637 assert!(
1638 client.service().managed_tools.try_write().is_ok(),
1639 "a hung network fetch must not hold the managed-registry lock"
1640 );
1641
1642 server_control
1643 .set_tools(vec![make_tool("newest", "Newest snapshot")])
1644 .await;
1645 running_server
1646 .peer()
1647 .notify_tool_list_changed()
1648 .await
1649 .unwrap();
1650
1651 tokio::time::timeout(Duration::from_secs(2), async {
1652 loop {
1653 let defs = handle.get_tool_defs(None).await.unwrap();
1654 if defs.len() == 1 && defs[0].name == "newest" {
1655 break;
1656 }
1657 tokio::task::yield_now().await;
1658 }
1659 })
1660 .await
1661 .expect("newest refresh committed while the older fetch remained hung");
1662
1663 server_control.release_first_refresh.notify_one();
1666 tokio::time::timeout(
1667 Duration::from_secs(2),
1668 server_control.first_refresh_returned.notified(),
1669 )
1670 .await
1671 .expect("delayed refresh response returned");
1672 for _ in 0..10 {
1673 tokio::task::yield_now().await;
1674 }
1675 let defs = handle.get_tool_defs(None).await.unwrap();
1676 assert_eq!(defs.len(), 1);
1677 assert_eq!(defs[0].name, "newest");
1678
1679 assert_eq!(server_control.list_calls.load(Ordering::SeqCst), 3);
1680 client.cancel().await.unwrap();
1681 }
1682
1683 #[tokio::test]
1684 async fn refresh_rebuilds_owned_tools_in_latest_server_order() {
1685 let server =
1686 DynamicToolServer::new(vec![make_tool("alpha", "Alpha"), make_tool("beta", "Beta")]);
1687 let server_control = server.clone();
1688 let handle = ToolServer::new().run();
1689 let (client, server_task) = connect(server, handle.clone()).await;
1690 server_control
1691 .set_tools(vec![
1692 make_tool("beta", "Beta refreshed"),
1693 make_tool("gamma", "Gamma"),
1694 make_tool("alpha", "Alpha refreshed"),
1695 ])
1696 .await;
1697 let running_server = server_task.await.unwrap();
1698 running_server
1699 .peer()
1700 .notify_tool_list_changed()
1701 .await
1702 .unwrap();
1703
1704 tokio::time::timeout(Duration::from_secs(2), async {
1705 loop {
1706 let defs = handle.get_tool_defs(None).await.unwrap();
1707 let names = defs
1708 .iter()
1709 .map(|definition| definition.name.as_str())
1710 .collect::<Vec<_>>();
1711 if names == ["beta", "gamma", "alpha"] && defs[0].description == "Beta refreshed" {
1712 break;
1713 }
1714 tokio::task::yield_now().await;
1715 }
1716 })
1717 .await
1718 .expect("latest MCP order committed");
1719 client.cancel().await.unwrap();
1720 }
1721
1722 #[tokio::test]
1723 async fn one_refresh_reclaims_a_name_after_a_peer_owner_disappears() {
1724 let handle = ToolServer::new().run();
1725 let first_server = DynamicToolServer::new(vec![make_tool("shared", "First owner")]);
1726 let first_control = first_server.clone();
1727 let (first_client, first_server_task) = connect(first_server, handle.clone()).await;
1728 let first_running_server = first_server_task.await.unwrap();
1729
1730 let second_server = DynamicToolServer::new(vec![make_tool("shared", "Second owner")]);
1731 let second_control = second_server.clone();
1732 let (second_client, second_server_task) = connect(second_server, handle.clone()).await;
1733 let second_running_server = second_server_task.await.unwrap();
1734 assert_eq!(
1735 handle.get_tool_defs(None).await.unwrap()[0].description,
1736 "Second owner"
1737 );
1738
1739 second_control.set_tools(Vec::new()).await;
1740 second_running_server
1741 .peer()
1742 .notify_tool_list_changed()
1743 .await
1744 .unwrap();
1745 tokio::time::timeout(Duration::from_secs(2), async {
1746 loop {
1747 if handle.get_tool_defs(None).await.unwrap().is_empty() {
1748 break;
1749 }
1750 tokio::task::yield_now().await;
1751 }
1752 })
1753 .await
1754 .expect("second owner removed its registration");
1755
1756 first_control
1760 .set_tools(vec![make_tool("shared", "First owner refreshed")])
1761 .await;
1762 first_running_server
1763 .peer()
1764 .notify_tool_list_changed()
1765 .await
1766 .unwrap();
1767 tokio::time::timeout(Duration::from_secs(2), async {
1768 loop {
1769 let defs = handle.get_tool_defs(None).await.unwrap();
1770 if defs.len() == 1 && defs[0].description == "First owner refreshed" {
1771 break;
1772 }
1773 tokio::task::yield_now().await;
1774 }
1775 })
1776 .await
1777 .expect("one refresh reclaimed the empty slot");
1778
1779 second_client.cancel().await.unwrap();
1780 first_client.cancel().await.unwrap();
1781 }
1782
1783 #[tokio::test]
1784 async fn refresh_does_not_replace_a_newer_local_registration() {
1785 let server = DynamicToolServer::new(vec![make_tool("alpha", "MCP alpha")]);
1786 let server_control = server.clone();
1787 let handle = ToolServer::new().run();
1788 let (client, server_task) = connect(server, handle.clone()).await;
1789
1790 handle
1791 .add_dynamic_tool(make_dynamic_tool("alpha", "Local alpha"))
1792 .await;
1793 server_control
1794 .set_tools(vec![make_tool("refresh_complete", "Refresh sentinel")])
1795 .await;
1796 let running_server = server_task.await.unwrap();
1797 running_server
1798 .peer()
1799 .notify_tool_list_changed()
1800 .await
1801 .unwrap();
1802
1803 tokio::time::timeout(Duration::from_secs(2), async {
1804 loop {
1805 let defs = handle.get_tool_defs(None).await.unwrap();
1806 if defs
1807 .iter()
1808 .any(|definition| definition.name == "refresh_complete")
1809 {
1810 break;
1811 }
1812 tokio::task::yield_now().await;
1813 }
1814 })
1815 .await
1816 .expect("MCP refresh completed");
1817
1818 let defs = handle.get_tool_defs(None).await.unwrap();
1819 let alpha = defs
1820 .iter()
1821 .find(|definition| definition.name == "alpha")
1822 .expect("alpha remains registered");
1823 assert_eq!(alpha.description, "Local alpha");
1824
1825 let result = handle
1826 .execute("alpha", "{}", &mut crate::tool::ToolContext::new())
1827 .await;
1828 assert_eq!(result.output(), &ToolOutput::text("local"));
1829 client.cancel().await.unwrap();
1830 }
1831
1832 #[tokio::test]
1833 async fn one_handler_refresh_protects_live_peer_and_reclaims_after_disconnect() {
1834 let server_a = DynamicToolServer::new(vec![make_tool("alpha", "Handler A")]);
1835 let server_a_control = server_a.clone();
1836 let server_b = DynamicToolServer::new(vec![make_tool("alpha", "Handler B")]);
1837 let handle = ToolServer::new().run();
1838
1839 let (client_a, server_task_a) = connect(server_a, handle.clone()).await;
1840 let (client_b, server_task_b) = connect(server_b, handle.clone()).await;
1841
1842 server_a_control
1843 .set_tools(vec![
1844 make_tool("alpha", "Refreshed handler A"),
1845 make_tool("a_refresh_complete", "Refresh sentinel"),
1846 ])
1847 .await;
1848 let running_server_a = server_task_a.await.unwrap();
1849 let _running_server_b = server_task_b.await.unwrap();
1850 running_server_a
1851 .peer()
1852 .notify_tool_list_changed()
1853 .await
1854 .unwrap();
1855
1856 tokio::time::timeout(Duration::from_secs(2), async {
1857 loop {
1858 let defs = handle.get_tool_defs(None).await.unwrap();
1859 if defs
1860 .iter()
1861 .any(|definition| definition.name == "a_refresh_complete")
1862 {
1863 break;
1864 }
1865 tokio::task::yield_now().await;
1866 }
1867 })
1868 .await
1869 .expect("handler A refresh completed");
1870
1871 let defs = handle.get_tool_defs(None).await.unwrap();
1872 let alpha = defs
1873 .iter()
1874 .find(|definition| definition.name == "alpha")
1875 .expect("alpha remains registered");
1876 assert_eq!(alpha.description, "Handler B");
1877
1878 client_b.cancel().await.unwrap();
1882 server_a_control
1883 .set_tools(vec![make_tool("alpha", "Reclaimed handler A")])
1884 .await;
1885 running_server_a
1886 .peer()
1887 .notify_tool_list_changed()
1888 .await
1889 .unwrap();
1890
1891 tokio::time::timeout(Duration::from_secs(2), async {
1892 loop {
1893 let defs = handle.get_tool_defs(None).await.unwrap();
1894 if defs
1895 .iter()
1896 .any(|definition| definition.description == "Reclaimed handler A")
1897 {
1898 break;
1899 }
1900 tokio::task::yield_now().await;
1901 }
1902 })
1903 .await
1904 .expect("handler A reclaimed the disconnected peer's registration");
1905
1906 let result = handle
1907 .execute("alpha", "{}", &mut crate::tool::ToolContext::new())
1908 .await;
1909 assert!(
1910 result.is_success(),
1911 "reclaimed tool should execute: {result:?}"
1912 );
1913
1914 client_a.cancel().await.unwrap();
1915 }
1916
1917 #[test]
1918 fn client_handler_get_info_delegates() {
1919 let info = ClientInfo::new(
1920 ClientCapabilities::default(),
1921 Implementation::new("test-client", "1.0.0"),
1922 );
1923 let handler = McpClientHandler::new(info, ToolServer::new().run());
1924 let returned = handler.get_info();
1925 assert_eq!(returned.client_info.name, "test-client");
1926 assert_eq!(returned.client_info.version, "1.0.0");
1927 }
1928
1929 #[tokio::test]
1930 async fn mcp_tool_preserves_provider_definition() {
1931 let tool = make_tool("search_docs", "Search the docs");
1932 let server = DynamicToolServer::new(vec![tool.clone()]);
1933 let (c2s, sfc) = tokio::io::duplex(8192);
1934 let (s2c, cfs) = tokio::io::duplex(8192);
1935 let task = tokio::spawn(async move {
1936 let running = server.serve((sfc, s2c)).await.unwrap();
1937 running.waiting().await.unwrap();
1938 });
1939 let client = ClientInfo::default().serve((cfs, c2s)).await.unwrap();
1940 let handle = ToolServer::new()
1941 .rmcp_tool(tool, client.peer().clone())
1942 .run();
1943 let defs = handle.get_tool_defs(None).await.unwrap();
1944 assert_eq!(defs.len(), 1);
1945 assert_eq!(defs[0].name, "search_docs");
1946 assert_eq!(defs[0].description, "Search the docs");
1947 client.cancel().await.unwrap();
1948
1949 let defs = handle.get_tool_defs(None).await.unwrap();
1950 assert!(
1951 defs.is_empty(),
1952 "a disconnected directly registered MCP tool must not remain provider-visible"
1953 );
1954 task.abort();
1955 }
1956
1957 #[tokio::test]
1958 async fn disconnected_directly_registered_mcp_tool_is_retired_on_dispatch() {
1959 let tool = make_tool("search_docs", "Search the docs");
1960 let server = DynamicToolServer::new(vec![tool.clone()]);
1961 let (c2s, sfc) = tokio::io::duplex(8192);
1962 let (s2c, cfs) = tokio::io::duplex(8192);
1963 let task = tokio::spawn(async move {
1964 let running = server.serve((sfc, s2c)).await.unwrap();
1965 running.waiting().await.unwrap();
1966 });
1967 let client = ClientInfo::default().serve((cfs, c2s)).await.unwrap();
1968 let handle = ToolServer::new()
1969 .rmcp_tool(tool, client.peer().clone())
1970 .run();
1971
1972 client.cancel().await.unwrap();
1973
1974 let result = handle
1975 .execute("search_docs", "{}", &mut crate::tool::ToolContext::new())
1976 .await;
1977 assert_eq!(
1978 result.error().expect("disconnected tool must fail").kind(),
1979 crate::tool::ToolErrorKind::NotFound
1980 );
1981 task.abort();
1982 }
1983}