use super::{CallToolResult, ContentBlock, McpError, McpResult, manager::McpManager};
use crate::{
cancellation::AgentCancellation,
config::{McpHttpServerConfig, McpServerConfig, McpServersSettings},
};
use serde_json::{Value, json};
use std::{
io::{BufRead, BufReader, Read, Write},
net::{TcpListener, TcpStream},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
mpsc::{self, Receiver},
},
thread::{self, JoinHandle},
time::{Duration, Instant},
};
fn wait_until(mut ready: impl FnMut() -> bool) {
let deadline = Instant::now() + Duration::from_secs(5);
while !ready() {
assert!(Instant::now() < deadline, "MCP workflow barrier timed out");
thread::sleep(Duration::from_millis(10));
}
}
#[derive(Default)]
struct HttpSeen {
calls: usize,
get_open: bool,
get_closed: bool,
deletes: usize,
}
struct HttpServer {
seen: Arc<Mutex<HttpSeen>>,
stop: Arc<AtomicBool>,
worker: Option<JoinHandle<()>>,
}
impl HttpServer {
fn start(reply: bool, timeout: u64) -> (Self, McpServerConfig) {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let config = McpServerConfig::Http(McpHttpServerConfig {
url: format!("http://{}/mcp", listener.local_addr().unwrap()),
headers: Default::default(),
oauth: None,
enabled: true,
timeout: Some(timeout),
});
let seen = Arc::new(Mutex::new(HttpSeen::default()));
let stop = Arc::new(AtomicBool::new(false));
let worker_seen = Arc::clone(&seen);
let worker_stop = Arc::clone(&stop);
let worker = thread::spawn(move || {
let mut get: Option<TcpStream> = None;
let mut waiting = Vec::with_capacity(2);
while !worker_stop.load(Ordering::SeqCst) {
let get_closed = get
.as_mut()
.is_some_and(|stream| match stream.read(&mut [0]) {
Ok(0) => true,
Err(error) => matches!(
error.kind(),
std::io::ErrorKind::ConnectionReset
| std::io::ErrorKind::ConnectionAborted
),
_ => false,
});
if get_closed {
worker_seen.lock().unwrap().get_closed = true;
get = None;
}
let (mut stream, _) = match listener.accept() {
Ok(connection) => connection,
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(5));
continue;
}
Err(error) => panic!("accept MCP request: {error}"),
};
stream
.set_read_timeout(Some(Duration::from_millis(250)))
.unwrap();
stream
.set_write_timeout(Some(Duration::from_secs(2)))
.unwrap();
let Some((head, body)) = read_request(&mut stream) else {
continue;
};
if head.starts_with("GET ") {
assert!(get.is_none(), "unexpected second notification stream");
stream.write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: keep-alive\r\n\r\n").unwrap();
stream.set_nonblocking(true).unwrap();
get = Some(stream);
worker_seen.lock().unwrap().get_open = true;
} else if head.starts_with("DELETE ") {
assert!(
head.to_ascii_lowercase()
.contains("mcp-session-id: concurrent-session")
);
worker_seen.lock().unwrap().deletes += 1;
respond(&mut stream, "204 No Content", "");
} else {
let request: Value = serde_json::from_slice(&body).unwrap();
let id = request["id"].clone();
let result = match request["method"].as_str().unwrap() {
"initialize" => json!({
"protocolVersion": "2025-03-26", "capabilities": {"tools": {}},
"serverInfo": {"name": "concurrent-http", "version": "1"}
}),
"notifications/initialized" => {
respond(&mut stream, "204 No Content", "");
continue;
}
"tools/list" => {
json!({"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]})
}
"tools/call" => {
worker_seen.lock().unwrap().calls += 1;
if reply {
waiting.push((stream, request));
if waiting.len() == 2 {
for (mut stream, request) in waiting.drain(..).rev() {
respond(&mut stream, "200 OK", &json!({
"jsonrpc": "2.0", "id": request["id"],
"result": {"content": [{"type": "text", "text": request["params"]["arguments"]["text"]}]}
}).to_string());
}
}
} else {
respond(&mut stream, "202 Accepted", "");
}
continue;
}
method => panic!("unexpected MCP method {method}"),
};
respond(
&mut stream,
"200 OK",
&json!({"jsonrpc": "2.0", "id": id, "result": result}).to_string(),
);
}
}
});
(
Self {
seen,
stop,
worker: Some(worker),
},
config,
)
}
}
impl Drop for HttpServer {
fn drop(&mut self) {
self.stop.store(true, Ordering::SeqCst);
if let Some(worker) = self.worker.take() {
let _ = worker.join();
}
}
}
fn read_request(stream: &mut TcpStream) -> Option<(String, Vec<u8>)> {
let mut reader = BufReader::new(stream);
let mut head = String::new();
loop {
let mut line = String::new();
match reader.read_line(&mut line) {
Ok(0) if head.is_empty() => return None,
Err(error)
if head.is_empty()
&& line.is_empty()
&& matches!(
error.kind(),
std::io::ErrorKind::WouldBlock
| std::io::ErrorKind::TimedOut
| std::io::ErrorKind::ConnectionReset
| std::io::ErrorKind::ConnectionAborted
) =>
{
return None;
}
Ok(0) => panic!("MCP request ended before its headers were complete"),
Err(error) => panic!("read MCP request header: {error}"),
Ok(_) => {}
}
head.push_str(&line);
assert!(head.len() <= 16 * 1024);
if line == "\r\n" {
break;
}
}
let length = head
.lines()
.find_map(|line| {
line.to_ascii_lowercase()
.strip_prefix("content-length:")?
.trim()
.parse::<usize>()
.ok()
})
.unwrap_or(0);
assert!(length <= 16 * 1024);
let mut body = vec![0; length];
reader.read_exact(&mut body).unwrap();
Some((head, body))
}
fn respond(stream: &mut TcpStream, status: &str, body: &str) {
let response = format!(
"HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nMcp-Session-Id: concurrent-session\r\nConnection: close\r\n\r\n{body}",
body.len()
);
let _ = stream.write_all(response.as_bytes());
}
enum Server {
Http(HttpServer),
#[cfg(unix)]
Stdio(tempfile::TempDir),
}
struct Fixture {
manager: McpManager,
server: Server,
}
impl Fixture {
fn new(server: Server, config: McpServerConfig) -> Self {
let settings = McpServersSettings::from([("shared".to_string(), config)]);
let manager = McpManager::from_settings_strict(&settings, None).unwrap();
match &server {
Server::Http(server) => wait_until(|| server.seen.lock().unwrap().get_open),
#[cfg(unix)]
Server::Stdio(_) => {}
}
Self { manager, server }
}
fn http(reply: bool, timeout: u64) -> Self {
let (server, config) = HttpServer::start(reply, timeout);
Self::new(Server::Http(server), config)
}
#[cfg(unix)]
fn stdio(reply: bool, timeout: u64) -> Self {
let temp = tempfile::tempdir().unwrap();
let config = McpServerConfig::Stdio(crate::config::McpStdioServerConfig {
command: "sh".to_string(),
args: vec!["-c".to_string(), STDIO_SERVER.to_string()],
env: [
(
"MCP_TEST_DIR".to_string(),
temp.path().to_str().unwrap().to_string(),
),
(
"MCP_TEST_REPLY".to_string(),
if reply { "yes" } else { "no" }.to_string(),
),
]
.into(),
enabled: true,
timeout: Some(timeout),
});
Self::new(Server::Stdio(temp), config)
}
fn wait_for_calls(&self, count: usize) {
wait_until(|| match &self.server {
Server::Http(server) => server.seen.lock().unwrap().calls == count,
#[cfg(unix)]
Server::Stdio(temp) => std::fs::read_to_string(temp.path().join("calls"))
.is_ok_and(|calls| calls.lines().count() == count),
});
}
fn assert_cleanup(&self) {
match &self.server {
Server::Http(server) => {
wait_until(|| {
let seen = server.seen.lock().unwrap();
seen.get_closed && seen.deletes >= 1
});
assert_eq!(server.seen.lock().unwrap().deletes, 1);
}
#[cfg(unix)]
Server::Stdio(temp) => {
let pid = std::fs::read_to_string(temp.path().join("pid")).unwrap();
wait_until(|| {
!std::process::Command::new("/bin/kill")
.args(["-0", pid.trim()])
.output()
.unwrap()
.status
.success()
});
}
}
}
fn start_call(
&self,
text: &str,
cancellation: AgentCancellation,
) -> (JoinHandle<()>, Receiver<McpResult<CallToolResult>>) {
let call = self.manager.resolve_tool_call("mcp__shared__echo").unwrap();
let arguments = Some(json!({"text": text}));
let (sender, receiver) = mpsc::sync_channel(1);
let worker = thread::spawn(move || {
let _ = sender.send(call.call_tool_cancellable(arguments, &cancellation));
});
(worker, receiver)
}
}
fn finish(
call: (JoinHandle<()>, Receiver<McpResult<CallToolResult>>),
) -> McpResult<CallToolResult> {
let result = call
.1
.recv_timeout(Duration::from_secs(5))
.expect("MCP call did not finish");
call.0.join().unwrap();
result
}
fn overlap(mut fixture: Fixture) {
let first = fixture.start_call("first", AgentCancellation::default());
fixture.wait_for_calls(1);
let canceled = AgentCancellation::new(Arc::new(AtomicBool::new(true)));
let error = fixture
.manager
.resolve_tool_call("mcp__shared__echo")
.unwrap()
.call_tool_cancellable(None, &canceled)
.unwrap_err()
.to_string();
assert!(error.contains("prompt canceled"), "{error}");
let second = fixture.start_call("second", AgentCancellation::default());
for (call, expected) in [(first, "first"), (second, "second")] {
let result = finish(call).unwrap();
assert!(matches!(&result.content[..], [ContentBlock::Text { text }] if text == expected));
}
fixture.manager.shutdown();
fixture.manager.shutdown();
fixture.assert_cleanup();
}
#[derive(Clone, Copy)]
enum Failure {
Cancel,
Timeout,
Shutdown,
}
fn failure(mut fixture: Fixture, failure: Failure) {
let stale_call = fixture
.manager
.resolve_tool_call("mcp__shared__echo")
.unwrap();
let flag = Arc::new(AtomicBool::new(false));
let first = fixture.start_call("first", AgentCancellation::new(Arc::clone(&flag)));
let second = fixture.start_call("second", AgentCancellation::default());
fixture.wait_for_calls(2);
match failure {
Failure::Cancel => flag.store(true, Ordering::SeqCst),
Failure::Timeout => {}
Failure::Shutdown => fixture.manager.shutdown(),
}
let errors = [finish(first).unwrap_err(), finish(second).unwrap_err()];
match failure {
Failure::Cancel => {
assert!(
errors[0].to_string().contains("prompt canceled"),
"{:?}",
errors[0]
);
assert!(errors[1].to_string().contains("closed"), "{:?}", errors[1]);
}
Failure::Timeout => assert!(
errors
.iter()
.any(|error| matches!(error, McpError::Timeout { .. }))
),
Failure::Shutdown => assert!(
errors
.iter()
.all(|error| error.to_string().contains("closed"))
),
}
let error = stale_call
.call_tool_cancellable(None, &AgentCancellation::default())
.unwrap_err()
.to_string();
assert!(error.contains("closed"), "{error}");
fixture.manager.shutdown();
fixture.manager.shutdown();
assert!(
fixture
.manager
.resolve_tool_call("mcp__shared__echo")
.is_err()
);
fixture.assert_cleanup();
}
#[test]
fn same_server_http_calls_overlap() {
overlap(Fixture::http(true, 5));
}
#[test]
fn http_cancellation_releases_siblings_and_deletes_session() {
failure(Fixture::http(false, 30), Failure::Cancel);
}
#[test]
fn http_timeout_closes_shared_connection_and_deletes_session() {
failure(Fixture::http(false, 2), Failure::Timeout);
}
#[test]
fn http_shutdown_releases_siblings_and_deletes_session() {
failure(Fixture::http(false, 30), Failure::Shutdown);
}
#[cfg(unix)]
#[test]
fn same_server_stdio_calls_overlap() {
overlap(Fixture::stdio(true, 5));
}
#[cfg(unix)]
#[test]
fn stdio_cancellation_releases_siblings_and_reaps_child() {
failure(Fixture::stdio(false, 30), Failure::Cancel);
}
#[cfg(unix)]
#[test]
fn stdio_timeout_closes_shared_connection_and_reaps_child() {
failure(Fixture::stdio(false, 2), Failure::Timeout);
}
#[cfg(unix)]
#[test]
fn stdio_shutdown_releases_siblings_and_reaps_child() {
failure(Fixture::stdio(false, 30), Failure::Shutdown);
}
#[cfg(unix)]
const STDIO_SERVER: &str = r#"
printf '%s\n' "$$" > "$MCP_TEST_DIR/pid"
first_id=
while IFS= read -r line; do
method=$(printf '%s' "$line" | sed -n 's/.*"method":"\([^"]*\)".*/\1/p')
id=$(printf '%s' "$line" | sed -n 's/.*"id":\([^,}]*\).*/\1/p')
case "$method" in
initialize)
printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":"2025-03-26","capabilities":{"tools":{}},"serverInfo":{"name":"concurrent-stdio","version":"1"}}}\n' "$id" ;;
tools/list)
printf '{"jsonrpc":"2.0","id":%s,"result":{"tools":[{"name":"echo","inputSchema":{"type":"object"}}]}}\n' "$id" ;;
tools/call)
printf '%s\n' "$id" >> "$MCP_TEST_DIR/calls"
if [ "$MCP_TEST_REPLY" = yes ]; then
text=$(printf '%s' "$line" | sed -n 's/.*"text":"\([^"]*\)".*/\1/p')
if [ -z "$first_id" ]; then
first_id=$id
first_text=$text
else
printf '{"jsonrpc":"2.0","id":%s,"result":{"content":[{"type":"text","text":"%s"}]}}\n' "$id" "$text"
printf '{"jsonrpc":"2.0","id":%s,"result":{"content":[{"type":"text","text":"%s"}]}}\n' "$first_id" "$first_text"
fi
fi ;;
esac
done
"#;