use super::*;
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::RwLock;
use rmcp::{
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, ServerCapabilities,
ServerInfo,
},
service::{MaybeSendFuture, RequestContext},
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt,
};
#[derive(Clone)]
struct SlowToolServer {
delay: Duration,
}
impl ServerHandler for SlowToolServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}
fn call_tool(
&self,
_params: CallToolRequestParams,
_ctx: RequestContext<RoleServer>,
) -> impl std::future::Future<Output = Result<CallToolResponse, McpError>> + MaybeSendFuture + '_
{
let delay = self.delay;
async move {
tokio::time::sleep(delay).await;
Ok(CallToolResult::success(vec![ContentBlock::text("ok")]).into())
}
}
}
async fn attach_slow_server(mgr: &mut McpManager, name: &str, delay: Duration) {
let (server_side, client_side) = tokio::io::duplex(8192);
let server = SlowToolServer { delay };
tokio::spawn(async move {
if let Ok(running) = server.serve(server_side).await {
let _ = running.waiting().await;
}
});
let handler = AgentBlockClientHandler::new();
let running = handler
.serve(client_side)
.await
.expect("client handshake should succeed over duplex");
mgr.servers.insert(name.to_string(), running);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn concurrent_call_tool_same_server_does_not_serialize() {
let delay = Duration::from_millis(300);
let mgr = Arc::new(RwLock::new(McpManager::new()));
attach_slow_server(&mut *mgr.write().await, "slow", delay).await;
let start = Instant::now();
let a = {
let mgr = Arc::clone(&mgr);
async move {
mgr.read()
.await
.call_tool("slow", "slow_tool", serde_json::json!({}))
.await
}
};
let b = {
let mgr = Arc::clone(&mgr);
async move {
mgr.read()
.await
.call_tool("slow", "slow_tool", serde_json::json!({}))
.await
}
};
let (r1, r2) = tokio::join!(a, b);
let elapsed = start.elapsed();
r1.expect("first call succeeds");
r2.expect("second call succeeds");
let serialized_budget = delay * 2 - Duration::from_millis(80);
assert!(
elapsed < serialized_budget,
"concurrent call_tool appears serialized: elapsed={:?}, serialized_budget={:?}",
elapsed,
serialized_budget,
);
}
#[tokio::test]
async fn two_reads_coexist_on_rwlock() {
let mgr = Arc::new(RwLock::new(McpManager::new()));
let _g1 = mgr.read().await;
assert!(
mgr.try_read().is_ok(),
"RwLock rejected a concurrent second read guard",
);
}
#[tokio::test]
async fn write_blocks_while_read_held() {
let mgr = Arc::new(RwLock::new(McpManager::new()));
let _g1 = mgr.read().await;
assert!(
mgr.try_write().is_err(),
"write lock acquired while a read guard was held",
);
}
#[derive(Clone)]
struct IsErrorServer;
impl ServerHandler for IsErrorServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
}
async fn call_tool(
&self,
_params: CallToolRequestParams,
_ctx: RequestContext<RoleServer>,
) -> Result<CallToolResponse, McpError> {
Ok(CallToolResult::error(vec![ContentBlock::text("tool blew up")]).into())
}
}
async fn attach_is_error_server(mgr: &mut McpManager, name: &str) {
let (server_side, client_side) = tokio::io::duplex(8192);
tokio::spawn(async move {
if let Ok(running) = IsErrorServer.serve(server_side).await {
let _ = running.waiting().await;
}
});
let handler = AgentBlockClientHandler::new();
let running = handler.serve(client_side).await.expect("handshake");
mgr.servers.insert(name.to_string(), running);
}
#[tokio::test]
async fn is_error_is_passed_through_in_ok_branch() {
let mut mgr = McpManager::new();
attach_is_error_server(&mut mgr, "boom").await;
let val = mgr
.call_tool("boom", "explode", serde_json::json!({}))
.await
.expect("RPC succeeds even when isError=true");
assert_eq!(
val.get("isError").and_then(|v| v.as_bool()),
Some(true),
"isError must be preserved in Ok branch: {val}",
);
let content = val.get("content").and_then(|v| v.as_array()).cloned();
assert!(
content.as_ref().map(|c| !c.is_empty()).unwrap_or(false),
"content blocks must be forwarded alongside isError: {val:?}",
);
}