use rmcp::ErrorData as McpError;
use rmcp::model::{CallToolRequestParams, CallToolResponse, CreateTaskResult};
use rmcp::service::{RequestContext, RoleServer};
use rmcp::task_manager::{TaskExit, TaskOptions};
use super::BasemindServer;
pub(super) const SLOW_CALLS: &[&str] = &[
"admin:rescan",
#[cfg(feature = "documents")]
"memory:documents",
#[cfg(feature = "crawl")]
"web:scrape",
#[cfg(feature = "crawl")]
"web:crawl",
#[cfg(feature = "crawl")]
"web:map",
];
pub(super) fn is_slow_tool(name: &str, arguments: Option<&serde_json::Map<String, serde_json::Value>>) -> bool {
if SLOW_CALLS.contains(&name) {
return true;
}
let Some(mode) = arguments
.and_then(|args| args.get("mode"))
.and_then(serde_json::Value::as_str)
else {
return false;
};
SLOW_CALLS.iter().any(|slow| {
slow.split_once(':')
.is_some_and(|(tool, slow_mode)| tool == name && slow_mode == mode)
})
}
pub(super) fn spawn_slow_tool(
server: &BasemindServer,
request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> CreateTaskResult {
let server = server.clone();
let manager = server.tasks.clone();
let task = manager.spawn(TaskOptions::new(), move |ctx| {
Box::pin(async move {
let mut work = tokio::spawn(async move {
let tcc = rmcp::handler::server::tool::ToolCallContext::new(&server, request, context);
server.tool_router.call(tcc).await
});
let outcome = tokio::select! {
biased;
() = ctx.cancelled() => return Err(TaskExit::Cancelled),
joined = &mut work => joined,
};
match outcome {
Ok(Ok(CallToolResponse::Complete(result))) => Ok(result),
Ok(Ok(_)) => Err(TaskExit::Error(McpError::internal_error(
"tool returned a non-terminal response inside a task",
None,
))),
Ok(Err(error)) => Err(TaskExit::Error(error)),
Err(join_error) => Err(TaskExit::Error(McpError::internal_error(
format!("slow tool task failed to complete: {join_error}"),
None,
))),
}
})
});
CreateTaskResult::new(task)
}
#[cfg(test)]
mod tests {
use super::*;
fn args(json: serde_json::Value) -> serde_json::Map<String, serde_json::Value> {
json.as_object().expect("object").clone()
}
#[test]
fn should_offload_the_slow_admin_mode_and_not_its_fast_siblings() {
assert!(is_slow_tool(
"admin",
Some(&args(serde_json::json!({ "mode": "rescan", "paths": ["src"] })))
));
assert!(!is_slow_tool(
"admin",
Some(&args(serde_json::json!({ "mode": "status" })))
));
assert!(!is_slow_tool("admin", None));
}
#[test]
fn should_not_offload_a_tool_that_is_not_listed() {
assert!(!is_slow_tool("outline", None));
}
#[cfg(feature = "crawl")]
#[test]
fn should_offload_a_consolidated_domain_only_for_the_modes_that_are_slow() {
assert!(is_slow_tool("web", Some(&args(serde_json::json!({ "mode": "crawl" })))));
assert!(is_slow_tool(
"web",
Some(&args(serde_json::json!({ "mode": "scrape" })))
));
}
#[cfg(feature = "crawl")]
#[test]
fn should_not_match_a_mode_across_domains() {
assert!(!is_slow_tool(
"memory",
Some(&args(serde_json::json!({ "mode": "map" })))
));
}
#[cfg(feature = "crawl")]
#[test]
fn should_not_offload_a_domain_call_with_an_absent_or_unknown_mode() {
assert!(!is_slow_tool("web", None));
assert!(!is_slow_tool(
"web",
Some(&args(serde_json::json!({ "mode": "sniff" })))
));
assert!(!is_slow_tool("web", Some(&args(serde_json::json!({ "mode": 7 })))));
}
}