1use std::sync::Arc;
10use std::time::Duration;
11
12use rmcp::model::{
13 CallToolRequest, CallToolResult, ClientRequest, ContentBlock, ResourceContents, ServerResult,
14};
15use rmcp::service::PeerRequestOptions;
16
17use rig_core::message::{EmptyToolName, ImageMediaType, MimeType, ToolName, ToolResultContent};
18use rig_core::tool::{
19 ContextValue, DynamicTool, ToolContext, ToolContextError, ToolExecutionError, ToolOutput,
20};
21use rig_core::wasm_compat::WasmBoxedFuture;
22
23pub use rmcp::model::Meta;
27
28#[derive(
33 Debug, Clone, Default, PartialEq, rig_core::serde::Serialize, rig_core::serde::Deserialize,
34)]
35#[serde(crate = "rig_core::serde", transparent)]
36pub struct McpMeta(pub Meta);
37
38impl ContextValue for McpMeta {
39 const KEY: &'static str = "rmcp.meta";
40}
41
42#[derive(Debug, Clone, PartialEq, rig_core::serde::Serialize, rig_core::serde::Deserialize)]
45#[serde(crate = "rig_core::serde", transparent)]
46pub struct McpStructuredContent(pub serde_json::Value);
47
48impl ContextValue for McpStructuredContent {
49 const KEY: &'static str = "rmcp.structured_content";
50}
51
52#[derive(
55 Debug, Clone, Default, PartialEq, rig_core::serde::Serialize, rig_core::serde::Deserialize,
56)]
57#[serde(crate = "rig_core::serde", transparent)]
58pub struct McpResponseMeta(pub Meta);
59
60impl ContextValue for McpResponseMeta {
61 const KEY: &'static str = "rmcp.response_meta";
62}
63
64#[derive(Debug, Clone, PartialEq, rig_core::serde::Serialize, rig_core::serde::Deserialize)]
67#[serde(crate = "rig_core::serde", transparent)]
68pub struct McpCallToolResult(pub CallToolResult);
69
70impl ContextValue for McpCallToolResult {
71 const KEY: &'static str = "rmcp.call_tool_result";
72}
73
74pub const DEFAULT_MCP_TOOL_TIMEOUT: Duration = Duration::from_secs(300);
76
77pub const DEFAULT_MCP_REFRESH_TIMEOUT: Duration = Duration::from_secs(30);
82
83const MCP_CANCELLATION_GRACE_PERIOD: Duration = Duration::from_secs(1);
86
87#[derive(Clone)]
93pub struct McpTool {
94 pub(crate) definition: rmcp::model::Tool,
95 pub(crate) client: rmcp::service::ServerSink,
96 pub(crate) timeout: Option<Duration>,
99}
100
101impl McpTool {
102 pub fn from_mcp_server(
106 definition: rmcp::model::Tool,
107 client: rmcp::service::ServerSink,
108 ) -> Self {
109 Self {
110 definition,
111 client,
112 timeout: Some(DEFAULT_MCP_TOOL_TIMEOUT),
113 }
114 }
115
116 #[must_use = "the setting applies to the returned value"]
122 pub fn with_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
123 self.timeout = timeout.into();
124 self
125 }
126
127 pub fn timeout(&self) -> Option<Duration> {
129 self.timeout
130 }
131
132 pub fn definition(&self) -> &rmcp::model::Tool {
134 &self.definition
135 }
136}
137
138#[derive(Debug, thiserror::Error)]
140pub(crate) enum McpArgumentError {
141 #[error("invalid JSON: {0}")]
143 Json(#[from] serde_json::Error),
144 #[error("expected a JSON object or null, got {0}")]
146 NonObject(&'static str),
147}
148
149pub(crate) fn json_value_kind(value: &serde_json::Value) -> &'static str {
150 match value {
151 serde_json::Value::Null => "null",
152 serde_json::Value::Bool(_) => "boolean",
153 serde_json::Value::Number(_) => "number",
154 serde_json::Value::String(_) => "string",
155 serde_json::Value::Array(_) => "array",
156 serde_json::Value::Object(_) => "object",
157 }
158}
159
160pub(crate) fn parse_mcp_arguments(
165 args: &str,
166) -> Result<Option<rmcp::model::JsonObject>, McpArgumentError> {
167 let trimmed = args.trim();
168 if trimmed.is_empty() {
169 return Ok(None);
170 }
171 let value: serde_json::Value = serde_json::from_str(trimmed)?;
172 match value {
173 serde_json::Value::Null => Ok(None),
174 serde_json::Value::Object(_) => Ok(Some(serde_json::from_value(value)?)),
175 value => Err(McpArgumentError::NonObject(json_value_kind(&value))),
176 }
177}
178
179pub(crate) async fn call_mcp_tool(
180 peer: &rmcp::service::ServerSink,
181 params: rmcp::model::CallToolRequestParams,
182 timeout: Option<Duration>,
183) -> Result<CallToolResult, rmcp::ServiceError> {
184 let deadline = timeout.map(|timeout| (tokio::time::Instant::now() + timeout, timeout));
185 let response = send_mcp_request(
186 peer,
187 ClientRequest::CallToolRequest(CallToolRequest::new(params)),
188 deadline,
189 )
190 .await?;
191
192 match response {
193 ServerResult::CallToolResult(result) => Ok(result),
194 _ => Err(rmcp::ServiceError::UnexpectedResponse),
195 }
196}
197
198pub(crate) async fn send_mcp_request(
199 peer: &rmcp::service::ServerSink,
200 request: ClientRequest,
201 deadline: Option<(tokio::time::Instant, Duration)>,
202) -> Result<ServerResult, rmcp::ServiceError> {
203 let handle = match deadline {
204 Some((deadline, timeout)) => {
205 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
206 if remaining.is_zero() {
207 return Err(rmcp::ServiceError::Timeout { timeout });
208 }
209 rig_core::wasm_compat::timeout(
210 remaining,
211 peer.send_cancellable_request(request, PeerRequestOptions::no_options()),
212 )
213 .await
214 .map_err(|_| rmcp::ServiceError::Timeout { timeout })??
215 }
216 None => {
217 peer.send_cancellable_request(request, PeerRequestOptions::no_options())
218 .await?
219 }
220 };
221
222 let Some((deadline, timeout)) = deadline else {
223 return handle.await_response().await;
224 };
225 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
226 let mut handle = handle;
227 match rig_core::wasm_compat::timeout(remaining, &mut handle.rx).await {
228 Ok(response) => response.map_err(|_| rmcp::ServiceError::TransportClosed)?,
229 Err(_) => {
230 cancel_timed_out_request(handle);
231 Err(rmcp::ServiceError::Timeout { timeout })
232 }
233 }
234}
235
236pub(crate) fn cancel_timed_out_request(
240 handle: rmcp::service::RequestHandle<rmcp::service::RoleClient>,
241) {
242 let cancellation = async move {
243 bounded_best_effort_cancellation(
244 handle.cancel(Some(
245 rmcp::service::RequestHandle::<rmcp::service::RoleClient>::REQUEST_TIMEOUT_REASON
246 .to_owned(),
247 )),
248 MCP_CANCELLATION_GRACE_PERIOD,
249 )
250 .await;
251 };
252
253 tokio::spawn(cancellation);
255}
256
257pub(crate) async fn bounded_best_effort_cancellation(
258 cancellation: impl std::future::Future<Output = Result<(), rmcp::ServiceError>>,
259 grace_period: Duration,
260) {
261 let _ = rig_core::wasm_compat::timeout(grace_period, cancellation).await;
262}
263
264impl McpTool {
265 pub fn execute_mcp(
271 &self,
272 args: String,
273 meta: Option<rmcp::model::Meta>,
274 ) -> WasmBoxedFuture<'_, Result<CallToolResult, ToolExecutionError>> {
275 let name = self.definition.name.clone();
276
277 Box::pin(async move {
278 let arguments = parse_mcp_arguments(&args).map_err(|error| {
281 ToolExecutionError::invalid_args(format!(
282 "MCP tool '{name}' received invalid arguments: {error}"
283 ))
284 .with_source(error)
285 })?;
286 let mut request = arguments
287 .map(|arguments| {
288 rmcp::model::CallToolRequestParams::new(name.clone()).with_arguments(arguments)
289 })
290 .unwrap_or_else(|| rmcp::model::CallToolRequestParams::new(name));
291 request.meta = meta;
292
293 match call_mcp_tool(&self.client, request, self.timeout).await {
294 Ok(result) => Ok(result),
295 Err(
296 error @ rmcp::ServiceError::Timeout {
297 timeout: elapsed_timeout,
298 },
299 ) => {
300 let timeout = self.timeout.unwrap_or(elapsed_timeout);
301 Err(ToolExecutionError::timeout(format!(
302 "MCP tool '{}' timed out after {timeout:?}",
303 self.definition.name
304 ))
305 .with_source(error))
306 }
307 Err(error) => Err(ToolExecutionError::provider(format!(
309 "MCP tool '{}' request failed: {error}",
310 self.definition.name
311 ))
312 .with_source(error)),
313 }
314 })
315 }
316}
317
318pub(crate) fn mcp_content_block_as_json(
319 content: &ContentBlock,
320) -> Result<ToolResultContent, ToolExecutionError> {
321 serde_json::to_value(content)
322 .map(ToolResultContent::json)
323 .map_err(|error| {
324 ToolExecutionError::provider(format!(
325 "failed to preserve an MCP content block as JSON: {error}"
326 ))
327 .with_source(error)
328 })
329}
330
331pub(crate) fn mcp_content_block_to_tool_content(
332 content: &ContentBlock,
333) -> Result<ToolResultContent, ToolExecutionError> {
334 match content {
335 ContentBlock::Text(text) => Ok(ToolResultContent::text(text.text.clone())),
336 ContentBlock::Image(image) => match ImageMediaType::from_mime_type(&image.mime_type) {
337 Some(media_type) => Ok(ToolResultContent::image_base64(
338 image.data.clone(),
339 Some(media_type),
340 None,
341 )),
342 None => mcp_content_block_as_json(content),
343 },
344 ContentBlock::Resource(resource) => match &resource.resource {
345 ResourceContents::TextResourceContents { .. } => mcp_content_block_as_json(content),
349 ResourceContents::BlobResourceContents {
350 mime_type, blob, ..
351 } => match mime_type
352 .as_deref()
353 .and_then(ImageMediaType::from_mime_type)
354 {
355 Some(media_type) => Ok(ToolResultContent::image_base64(
356 blob.clone(),
357 Some(media_type),
358 None,
359 )),
360 _ => mcp_content_block_as_json(content),
361 },
362 _ => mcp_content_block_as_json(content),
363 },
364 ContentBlock::ResourceLink(_) | ContentBlock::Audio(_) => {
365 mcp_content_block_as_json(content)
366 }
367 _ => mcp_content_block_as_json(content),
370 }
371}
372
373pub fn mcp_result_output(result: &CallToolResult) -> Result<ToolOutput, ToolExecutionError> {
375 let structured = result.structured_content.as_ref();
376 let canonical_fallback = structured.map(serde_json::Value::to_string);
377 let mut replaced_fallback = false;
378 let mut mapped = Vec::with_capacity(result.content.len());
379
380 for block in &result.content {
381 let fallback_structured = if !replaced_fallback {
382 match (block, canonical_fallback.as_deref(), structured) {
383 (ContentBlock::Text(text), Some(fallback), Some(structured))
384 if text.text == fallback =>
385 {
386 Some(structured)
387 }
388 _ => None,
389 }
390 } else {
391 None
392 };
393 if let Some(structured) = fallback_structured {
394 mapped.push(ToolResultContent::json(structured.clone()));
398 replaced_fallback = true;
399 } else {
400 mapped.push(mcp_content_block_to_tool_content(block)?);
401 }
402 }
403
404 if let Some(structured) = structured
405 && !replaced_fallback
406 {
407 mapped.insert(0, ToolResultContent::json(structured.clone()));
412 }
413
414 if !mapped.is_empty() {
415 return ToolOutput::content(mapped);
416 }
417
418 if result.is_error == Some(true) {
421 Ok(ToolOutput::text("the MCP tool reported an error"))
422 } else {
423 Ok(ToolOutput::text(""))
424 }
425}
426
427#[derive(Debug, thiserror::Error)]
429pub enum McpClientError {
430 #[error("MCP connection error: {0}")]
432 Connection(#[from] rmcp::service::ClientInitializeError),
433
434 #[error("Failed to fetch MCP tool list: {0}")]
436 ToolFetch(#[from] rmcp::ServiceError),
437
438 #[error("Timed out fetching MCP tool list after {0:?}")]
440 ToolFetchTimeout(Duration),
441}
442
443pub fn tools_from_server(
448 tools: impl IntoIterator<Item = rmcp::model::Tool>,
449 client: &rmcp::service::ServerSink,
450) -> Vec<McpTool> {
451 tools
452 .into_iter()
453 .map(|tool| McpTool::from_mcp_server(tool, client.clone()))
454 .collect()
455}
456
457pub fn preserve_mcp_result(
461 context: &mut ToolContext,
462 result: CallToolResult,
463) -> Result<(), ToolContextError> {
464 if let Some(structured) = result.structured_content.clone() {
465 context.insert_result(McpStructuredContent(structured))?;
466 }
467 if let Some(meta) = result.meta.clone() {
468 context.insert_result(McpResponseMeta(meta))?;
469 }
470 context.insert_result(McpCallToolResult(result))?;
471 Ok(())
472}
473
474impl TryFrom<McpTool> for DynamicTool {
483 type Error = EmptyToolName;
484
485 fn try_from(tool: McpTool) -> Result<Self, EmptyToolName> {
486 let name = ToolName::new(tool.definition.name.to_string())?;
487 let description = tool
488 .definition
489 .description
490 .as_deref()
491 .unwrap_or("")
492 .to_string();
493 let parameters = tool.definition.schema_as_json_value();
494 let liveness_client = tool.client.clone();
495 let tool = Arc::new(tool);
496 Ok(DynamicTool::new_with_context(
497 name,
498 description,
499 parameters,
500 move |context: &mut ToolContext, args: serde_json::Value| {
501 let tool = Arc::clone(&tool);
502 let meta = context.get::<McpMeta>();
503 Box::pin(async move {
504 let meta = meta?.map(|meta| meta.0);
505 let result = tool.execute_mcp(args.to_string(), meta).await?;
506 let is_error = result.is_error == Some(true);
507 let output = mcp_result_output(&result);
508 preserve_mcp_result(context, result)?;
509 let output = output?;
510 if is_error {
511 Err(ToolExecutionError::other(format!(
512 "MCP tool '{}' reported an execution error",
513 tool.definition.name
514 ))
515 .with_model_output(output))
516 } else {
517 Ok(output)
518 }
519 })
520 },
521 )
522 .with_liveness(move || !liveness_client.is_transport_closed()))
523 }
524}
525
526const _: fn() = || {
529 fn assert_send_sync_static<T: Send + Sync + 'static>() {}
530 assert_send_sync_static::<McpTool>();
531};