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