use anda_core::{BoxError, CancellationToken, FunctionDefinition, Json, ToolOutput, Usage};
use rmcp::{
Peer, RoleClient,
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, CancelTaskParams, ContentBlock,
CreateTaskResult, DEFAULT_MRTR_MAX_ROUNDS, GetTaskParams, InputRequiredResult, TaskPayload,
},
};
use serde_json::json;
use std::{
collections::hash_map::DefaultHasher,
hash::{Hash, Hasher},
time::Duration,
};
use super::McpTasksConfig;
use tokio::time::Instant;
const TASK_CANCEL_TIMEOUT: Duration = Duration::from_secs(2);
pub(crate) const MAX_LOCAL_NAME_ATTEMPTS: usize = 8;
pub(crate) const MRTR_STATE_ROUND_DELAY: Duration = Duration::from_millis(200);
pub(crate) const TASK_POLL_INTERVAL: Duration = Duration::from_secs(1);
pub(crate) const TASK_POLL_INTERVAL_MIN: Duration = Duration::from_millis(250);
pub(crate) const TASK_POLL_INTERVAL_MAX: Duration = Duration::from_secs(10);
pub(crate) const DEFAULT_TASK_MAX_WAIT_SECS: u64 = 300;
pub(crate) const MAX_TASK_MAX_WAIT_SECS: u64 = 24 * 60 * 60;
pub(crate) async fn call_tool_rounds(
route: &McpToolRoute,
peer: &Peer<RoleClient>,
mut params: CallToolRequestParams,
tasks: Option<&McpTasksConfig>,
cancellation: &CancellationToken,
request_timeout: Duration,
elicitation: Option<&super::interaction::ElicitationDispatcher>,
) -> Result<CallToolResult, BoxError> {
for _ in 0..DEFAULT_MRTR_MAX_ROUNDS {
let response = tokio::select! {
biased;
_ = cancellation.cancelled() => return Err("MCP tool call cancelled".into()),
response = tokio::time::timeout(request_timeout, peer.call_tool_once(params.clone())) => {
response.map_err(|_| format!("MCP tool {} request timed out", route.name))??
}
};
match response {
CallToolResponse::Complete(result) => return Ok(result),
CallToolResponse::InputRequired(result) => {
if let Some(requests) = result
.input_requests
.as_ref()
.filter(|requests| !requests.is_empty())
{
let Some(dispatcher) = elicitation.filter(|_| {
requests.values().all(|request| {
matches!(request, rmcp::model::InputRequest::Elicitation(_))
})
}) else {
return Ok(input_required_error(route, &result));
};
let mut responses = std::collections::BTreeMap::new();
for (key, request) in requests {
let rmcp::model::InputRequest::Elicitation(request) = request else {
unreachable!()
};
let response = dispatcher
.elicit(request.params.clone(), cancellation)
.await?;
responses.insert(key.clone(), serde_json::to_value(response)?);
}
params.input_responses = Some(responses);
params.request_state = result.request_state;
continue;
}
let Some(request_state) = result.request_state else {
return Err(format!(
"MCP tool {} returned an input_required result with neither \
input requests nor request state",
route.name
)
.into());
};
params.request_state = Some(request_state);
params.input_responses = None;
tokio::time::sleep(MRTR_STATE_ROUND_DELAY).await;
}
CallToolResponse::Task(task) => {
return await_task(route, peer, task, tasks, cancellation).await;
}
other => {
return Err(format!(
"MCP tool {} returned an unsupported response: {other:?}",
route.name
)
.into());
}
}
}
Err(format!(
"MCP tool {} did not complete within {DEFAULT_MRTR_MAX_ROUNDS} input_required rounds",
route.name
)
.into())
}
async fn await_task(
route: &McpToolRoute,
peer: &Peer<RoleClient>,
created: CreateTaskResult,
tasks: Option<&McpTasksConfig>,
cancellation: &CancellationToken,
) -> Result<CallToolResult, BoxError> {
let task_id = created.task.task_id.clone();
let mut cleanup = TaskCleanup {
peer: peer.clone(),
task_id: task_id.clone(),
armed: true,
};
let Some(tasks) = tasks else {
cleanup.armed = false;
cancel_task(peer, &task_id).await;
return Err(format!(
"MCP tool {} returned a task handle, but the tasks extension is not enabled \
for server {}",
route.name, route.server_id
)
.into());
};
let max_wait = tasks.max_wait();
let deadline = Instant::now() + max_wait;
let mut interval = task_poll_interval(created.task.poll_interval_ms);
loop {
let poll = tokio::select! {
biased;
_ = cancellation.cancelled() => {
cleanup.armed = false;
cancel_task(peer, &task_id).await;
return Err("MCP task cancelled".into());
}
poll = tokio::time::timeout_at(deadline, async {
tokio::time::sleep(interval).await;
peer.get_task(GetTaskParams::new(task_id.clone())).await
}) => poll,
};
let task = match poll {
Ok(result) => result?.task,
Err(_) => {
cleanup.armed = false;
cancel_task(peer, &task_id).await;
return Err(format!(
"MCP tool {} task {task_id} did not finish within {}s",
route.name,
max_wait.as_secs()
)
.into());
}
};
interval = task_poll_interval(task.task.poll_interval_ms);
match task.payload {
TaskPayload::Working => continue,
TaskPayload::Completed { result } => {
cleanup.armed = false;
return serde_json::from_value(Json::Object(result)).map_err(|err| {
format!(
"MCP tool {} returned an unreadable task result: {err}",
route.name
)
.into()
});
}
TaskPayload::Failed { error } => {
cleanup.armed = false;
return Err(format!(
"MCP tool {} task {task_id} failed: {}",
route.name,
Json::Object(error)
)
.into());
}
TaskPayload::Cancelled => {
cleanup.armed = false;
return Err(format!("MCP tool {} task {task_id} was cancelled", route.name).into());
}
TaskPayload::InputRequired { input_requests } => {
cleanup.armed = false;
cancel_task(peer, &task_id).await;
return Ok(unsupported_input_error(
route,
input_requests.keys().map(String::as_str),
));
}
_ => {
return Err(format!(
"MCP tool {} task {task_id} reported an unsupported status",
route.name
)
.into());
}
}
}
}
struct TaskCleanup {
peer: Peer<RoleClient>,
task_id: String,
armed: bool,
}
impl Drop for TaskCleanup {
fn drop(&mut self) {
if self.armed
&& let Ok(runtime) = tokio::runtime::Handle::try_current()
{
let peer = self.peer.clone();
let task_id = self.task_id.clone();
runtime.spawn(async move {
cancel_task(&peer, &task_id).await;
});
}
}
}
#[derive(Debug, Clone)]
pub struct McpToolRoute {
pub name: String,
pub server_id: String,
pub remote_name: String,
pub definition: FunctionDefinition,
pub tool: rmcp::model::Tool,
pub server_generation: u64,
pub catalog_revision: u64,
}
pub(super) fn tool_is_model_visible(tool: &rmcp::model::Tool) -> bool {
let visibility = tool
.meta
.as_deref()
.and_then(|meta| meta.get("ui"))
.and_then(|ui| ui.get("visibility"));
match visibility {
None => true,
Some(Json::Array(targets)) => targets
.iter()
.any(|target| target.as_str() == Some("model")),
Some(_) => false,
}
}
pub(super) fn model_schema(schema: &serde_json::Map<String, Json>) -> Json {
let mut schema = schema.clone();
if schema.get("type").and_then(Json::as_str) == Some("object")
&& schema.get("properties").is_none_or(Json::is_null)
{
schema.insert("properties".into(), json!({}));
}
Json::Object(schema)
}
pub(crate) fn task_poll_interval(poll_interval_ms: Option<u64>) -> Duration {
poll_interval_ms
.map(Duration::from_millis)
.unwrap_or(TASK_POLL_INTERVAL)
.clamp(TASK_POLL_INTERVAL_MIN, TASK_POLL_INTERVAL_MAX)
}
async fn cancel_task(peer: &Peer<RoleClient>, task_id: &str) {
match tokio::time::timeout(
TASK_CANCEL_TIMEOUT,
peer.cancel_task(CancelTaskParams::new(task_id)),
)
.await
{
Ok(Ok(_)) => {}
Ok(Err(err)) => log::debug!("MCP task {task_id} could not be cancelled: {err}"),
Err(_) => log::debug!("MCP task {task_id} cancellation acknowledgement timed out"),
}
}
pub(crate) fn input_required_error(
route: &McpToolRoute,
result: &InputRequiredResult,
) -> CallToolResult {
let keys = result
.input_requests
.iter()
.flat_map(|requests| requests.keys().map(String::as_str));
unsupported_input_error(route, keys)
}
pub(crate) fn unsupported_input_error<'a>(
route: &McpToolRoute,
request_keys: impl Iterator<Item = &'a str>,
) -> CallToolResult {
let keys: Vec<&str> = request_keys.collect();
let requested = if keys.is_empty() {
String::new()
} else {
format!(" (requests: {})", keys.join(", "))
};
CallToolResult::error(vec![ContentBlock::text(format!(
"MCP tool {} on server {} requires client-side input{requested}, which this host \
does not provide for this call: sampling and roots are not supported; elicitation requires opt-in. Call the tool \
with complete arguments, or use a different tool.",
route.remote_name, route.server_id
))])
}
pub(crate) fn mcp_result_to_tool_output(
route: &McpToolRoute,
result: CallToolResult,
limits: &super::McpLimits,
) -> ToolOutput<Json> {
let mut output = ToolOutput::new(json!({
"server_id": route.server_id,
"tool": route.remote_name,
"structured_content": result.structured_content,
"content": result.content,
"_meta": result.meta,
}));
output.model_output = Some(super::presentation::present_result(&output.output, limits));
output.is_error = result.is_error;
output.usage = Usage {
requests: 1,
..Usage::default()
};
output
}
pub(crate) fn sanitize_name_part(input: &str) -> String {
let mut out = String::new();
let mut previous_underscore = false;
for c in input.chars() {
let c = c.to_ascii_lowercase();
let valid = matches!(c, 'a'..='z' | '0'..='9');
if valid {
out.push(c);
previous_underscore = false;
} else if !previous_underscore {
out.push('_');
previous_underscore = true;
}
}
let trimmed = out.trim_matches('_').to_string();
let mut normalized = if trimmed.is_empty() {
"x".to_string()
} else {
trimmed
};
if !normalized
.chars()
.next()
.is_some_and(|c: char| c.is_ascii_lowercase())
{
normalized.insert(0, 'x');
}
normalized
}
pub(crate) fn shorten_with_hash(base: &str, key: &str) -> String {
let mut hasher = DefaultHasher::new();
key.hash(&mut hasher);
let suffix = format!("{:08x}", hasher.finish() as u32);
let max_prefix = 64usize.saturating_sub(suffix.len() + 1);
let mut prefix = base.chars().take(max_prefix).collect::<String>();
prefix = prefix.trim_end_matches('_').to_string();
format!("{}_{}", prefix, suffix)
}