#![allow(clippy::disallowed_methods)]
use crate::types::{
JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, ToolCallResult, ToolDefinition,
};
use std::collections::HashMap;
use std::sync::mpsc::{self, Sender};
use std::sync::{Arc, Mutex};
pub type NotificationSink = Box<dyn Fn(JsonRpcNotification) + Send + Sync>;
#[derive(Debug)]
pub struct CancelHandle {
pub cancel_tx: Sender<()>,
}
type InFlight = Arc<Mutex<HashMap<serde_json::Value, CancelHandle>>>;
#[derive(Debug)]
pub struct AprMcpServer {
in_flight: InFlight,
#[cfg(feature = "native")]
workers: Vec<std::thread::JoinHandle<()>>,
#[cfg(feature = "native")]
worker_dispatch: crate::tools::DispatchFn,
}
impl Default for AprMcpServer {
fn default() -> Self {
Self {
in_flight: InFlight::default(),
#[cfg(feature = "native")]
workers: Vec::new(),
#[cfg(feature = "native")]
worker_dispatch: dispatch_tool_call_with_sink,
}
}
}
impl AprMcpServer {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn handle_request(&mut self, request: &JsonRpcRequest) -> JsonRpcResponse {
if request.jsonrpc != "2.0" {
return JsonRpcResponse::error(
request.id.clone(),
-32600,
format!(
"Invalid Request: jsonrpc must be \"2.0\", got \"{}\"",
request.jsonrpc
),
);
}
match request.method.as_str() {
"initialize" => self.handle_initialize(request),
"tools/list" => self.handle_tools_list(request),
"tools/call" => self.handle_tools_call_sync(request),
"ping" => JsonRpcResponse::success(request.id.clone(), serde_json::json!({})),
other => JsonRpcResponse::error(
request.id.clone(),
-32601,
format!("Method not found: {other}"),
),
}
}
fn handle_initialize(&self, request: &JsonRpcRequest) -> JsonRpcResponse {
JsonRpcResponse::success(
request.id.clone(),
serde_json::json!({
"protocolVersion": crate::PROTOCOL_VERSION,
"capabilities": {
"tools": { "listChanged": false }
},
"serverInfo": {
"name": crate::SERVER_NAME,
"version": env!("CARGO_PKG_VERSION"),
},
}),
)
}
fn handle_tools_list(&self, request: &JsonRpcRequest) -> JsonRpcResponse {
let tools: Vec<ToolDefinition> = self.tool_definitions();
JsonRpcResponse::success(request.id.clone(), serde_json::json!({ "tools": tools }))
}
fn handle_tools_call_sync(&self, request: &JsonRpcRequest) -> JsonRpcResponse {
let (_tx, rx) = mpsc::channel::<()>();
let result = dispatch_tool_call(&request.params, &rx, None);
JsonRpcResponse::success(
request.id.clone(),
serde_json::to_value(result).unwrap_or_else(|_| serde_json::json!({})),
)
}
#[must_use]
pub fn handle_request_with_sink(
&mut self,
request: &JsonRpcRequest,
sink: &NotificationSink,
) -> Option<JsonRpcResponse> {
if request.jsonrpc != "2.0" {
return Some(JsonRpcResponse::error(
request.id.clone(),
-32600,
format!(
"Invalid Request: jsonrpc must be \"2.0\", got \"{}\"",
request.jsonrpc
),
));
}
if request.method.starts_with("notifications/") {
return None;
}
if request.method != "tools/call" {
return Some(self.handle_request(request));
}
let progress_token = extract_progress_token(&request.params);
let (_tx, rx) = mpsc::channel::<()>();
let sink_for_dispatch = progress_token.as_ref().map(|_| sink);
let result =
dispatch_tool_call_with_sink(&request.params, &rx, sink_for_dispatch, progress_token);
Some(JsonRpcResponse::success(
request.id.clone(),
serde_json::to_value(result).unwrap_or_else(|_| serde_json::json!({})),
))
}
#[must_use]
pub fn tool_definitions(&self) -> Vec<ToolDefinition> {
tool_index().definitions().to_vec()
}
#[must_use]
pub fn register_in_flight(in_flight: &InFlight, id: serde_json::Value) -> mpsc::Receiver<()> {
let (tx, rx) = mpsc::channel::<()>();
let mut guard = in_flight
.lock()
.expect("in_flight mutex not poisoned during register");
guard.insert(id, CancelHandle { cancel_tx: tx });
rx
}
pub fn cancel_in_flight(in_flight: &InFlight, id: &serde_json::Value) -> bool {
let mut guard = in_flight
.lock()
.expect("in_flight mutex not poisoned during cancel");
if let Some(handle) = guard.remove(id) {
let _ = handle.cancel_tx.send(());
true
} else {
false
}
}
fn deregister_in_flight(in_flight: &InFlight, id: &serde_json::Value) {
if let Ok(mut guard) = in_flight.lock() {
guard.remove(id);
}
}
#[cfg(feature = "native")]
pub fn run_stdio(&mut self) -> anyhow::Result<()> {
let stdin = std::io::stdin();
let reader = stdin.lock();
self.serve_stream(reader, Arc::new(Mutex::new(std::io::stdout())))
}
#[cfg(feature = "native")]
pub fn serve_stream<R, W>(&mut self, reader: R, out: Arc<Mutex<W>>) -> anyhow::Result<()>
where
R: std::io::BufRead,
W: std::io::Write + Send + 'static,
{
let outcome = self.read_loop(reader, &out);
self.join_workers();
outcome
}
#[cfg(feature = "native")]
fn read_loop<R, W>(&mut self, mut reader: R, out: &Arc<Mutex<W>>) -> anyhow::Result<()>
where
R: std::io::BufRead,
W: std::io::Write + Send + 'static,
{
let mut buf: Vec<u8> = Vec::new();
loop {
buf.clear();
if reader.read_until(b'\n', &mut buf)? == 0 {
break; }
while matches!(buf.last(), Some(b'\n' | b'\r')) {
buf.pop();
}
let Ok(line) = std::str::from_utf8(&buf) else {
let resp =
JsonRpcResponse::error(None, -32700, "Parse error: message is not valid UTF-8");
write_response(out, &resp)?;
continue;
};
if line.trim().is_empty() {
continue;
}
match parse_incoming(line) {
Ok(req) => self.route_stdio_message(req, out)?,
Err(resp) => write_response(out, &resp)?,
}
self.reap_finished_workers();
}
Ok(())
}
#[cfg(feature = "native")]
fn reap_finished_workers(&mut self) {
self.workers.retain(|h| !h.is_finished());
}
#[cfg(feature = "native")]
fn worker_response<F>(id: &serde_json::Value, dispatch: F) -> JsonRpcResponse
where
F: FnOnce() -> ToolCallResult,
{
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(dispatch)) {
Ok(result) => JsonRpcResponse::success(
Some(id.clone()),
serde_json::to_value(result).unwrap_or_else(|_| serde_json::json!({})),
),
Err(payload) => JsonRpcResponse::error(
Some(id.clone()),
-32603,
format!(
"Internal error: tool panicked: {}",
panic_message(payload.as_ref())
),
),
}
}
#[cfg(feature = "native")]
fn join_workers(&mut self) {
for handle in std::mem::take(&mut self.workers) {
let _ = handle.join();
}
}
#[cfg(feature = "native")]
fn route_stdio_message<W>(
&mut self,
req: JsonRpcRequest,
stdout: &Arc<Mutex<W>>,
) -> anyhow::Result<()>
where
W: std::io::Write + Send + 'static,
{
if req.jsonrpc != "2.0" {
let resp = JsonRpcResponse::error(
req.id.clone(),
-32600,
format!(
"Invalid Request: jsonrpc must be \"2.0\", got \"{}\"",
req.jsonrpc
),
);
return write_response(stdout, &resp);
}
match req.method.as_str() {
"notifications/cancelled" => {
if let Some(request_id) = req.params.get("requestId").cloned() {
let _ = Self::cancel_in_flight(&self.in_flight, &request_id);
}
Ok(())
}
"notifications/initialized" => {
Ok(())
}
"tools/call" => self.spawn_tools_call_worker(req, stdout),
_ => {
if req.id.is_none() {
return Ok(());
}
let resp = self.handle_request(&req);
write_response(stdout, &resp)
}
}
}
#[cfg(feature = "native")]
fn spawn_tools_call_worker<W>(
&mut self,
req: JsonRpcRequest,
stdout: &Arc<Mutex<W>>,
) -> anyhow::Result<()>
where
W: std::io::Write + Send + 'static,
{
let Some(id) = req.id.clone() else {
let resp =
JsonRpcResponse::error(None, -32600, "Invalid Request: tools/call requires an id");
return write_response(stdout, &resp);
};
let cancel_rx = Self::register_in_flight(&self.in_flight, id.clone());
let stdout_clone = Arc::clone(stdout);
let in_flight_clone = Arc::clone(&self.in_flight);
let params = req.params.clone();
let id_for_worker = id.clone();
let progress_token = extract_progress_token(¶ms);
let sink_stdout = Arc::clone(stdout);
let sink: NotificationSink = Box::new(move |notif| {
let _ = write_notification(&sink_stdout, ¬if);
});
let builder = std::thread::Builder::new().name(format!("apr-mcp-call-{id}"));
let dispatch = self.worker_dispatch;
let spawn_result = builder.spawn(move || {
let resp = Self::worker_response(&id_for_worker, || {
let sink_ref = progress_token.as_ref().map(|_| &sink);
dispatch(¶ms, &cancel_rx, sink_ref, progress_token)
});
let _ = write_response(&stdout_clone, &resp);
Self::deregister_in_flight(&in_flight_clone, &id_for_worker);
});
match spawn_result {
Ok(handle) => {
self.workers.push(handle);
Ok(())
}
Err(e) => {
Self::deregister_in_flight(&self.in_flight, &id);
let resp = JsonRpcResponse::error(
Some(id),
-32603,
format!("Internal error: failed to spawn worker thread: {e}"),
);
write_response(stdout, &resp)
}
}
}
#[must_use]
pub fn in_flight_handle(&self) -> InFlight {
Arc::clone(&self.in_flight)
}
}
fn dispatch_tool_call(
params: &serde_json::Value,
cancel_rx: &mpsc::Receiver<()>,
sink: Option<&NotificationSink>,
) -> ToolCallResult {
dispatch_tool_call_with_sink(params, cancel_rx, sink, None)
}
fn dispatch_tool_call_with_sink(
params: &serde_json::Value,
cancel_rx: &mpsc::Receiver<()>,
sink: Option<&NotificationSink>,
progress_token: Option<serde_json::Value>,
) -> ToolCallResult {
let name = params.get("name").and_then(|v| v.as_str());
let arguments = params
.get("arguments")
.cloned()
.unwrap_or_else(|| serde_json::json!({}));
let Some(name) = name else {
return ToolCallResult::error("Missing tool name");
};
match tool_index().dispatch_for(name) {
Some(dispatch_fn) => dispatch_fn(&arguments, cancel_rx, sink, progress_token),
None => ToolCallResult::error(format!("Unknown tool: {name}")),
}
}
fn tool_index() -> &'static crate::tools::ToolIndex {
static INDEX: std::sync::OnceLock<crate::tools::ToolIndex> = std::sync::OnceLock::new();
INDEX.get_or_init(crate::tools::ToolIndex::from_inventory)
}
fn extract_progress_token(params: &serde_json::Value) -> Option<serde_json::Value> {
params
.get("_meta")
.and_then(|m| m.get("progressToken"))
.cloned()
}
#[cfg(feature = "native")]
fn json_type_name(value: &serde_json::Value) -> &'static str {
crate::tools::args::json_type_name(value)
}
fn parse_incoming(line: &str) -> Result<JsonRpcRequest, Box<JsonRpcResponse>> {
let value: serde_json::Value = serde_json::from_str(line).map_err(|e| {
Box::new(JsonRpcResponse::error(
None,
-32700,
format!("Parse error: {e}"),
))
})?;
if value.is_array() {
return Err(Box::new(JsonRpcResponse::error(
None,
-32600,
"Invalid Request: JSON-RPC batch arrays are not supported; \
send one request per line",
)));
}
let Some(obj) = value.as_object() else {
return Err(Box::new(JsonRpcResponse::error(
None,
-32600,
format!(
"Invalid Request: a request must be a JSON object, got {}",
json_type_name(&value)
),
)));
};
let id = obj.get("id").filter(|v| !v.is_null()).cloned();
let jsonrpc = match obj.get("jsonrpc") {
Some(serde_json::Value::String(s)) => s.clone(),
Some(other) => {
return Err(Box::new(JsonRpcResponse::error(
id,
-32600,
format!(
"Invalid Request: \"jsonrpc\" must be the string \"2.0\", got {}",
json_type_name(other)
),
)));
}
None => {
return Err(Box::new(JsonRpcResponse::error(
id,
-32600,
"Invalid Request: missing required field \"jsonrpc\"",
)));
}
};
let method = match obj.get("method") {
Some(serde_json::Value::String(s)) => s.clone(),
Some(other) => {
return Err(Box::new(JsonRpcResponse::error(
id,
-32600,
format!(
"Invalid Request: \"method\" must be a string, got {}",
json_type_name(other)
),
)));
}
None => {
return Err(Box::new(JsonRpcResponse::error(
id,
-32600,
"Invalid Request: missing required field \"method\"",
)));
}
};
Ok(JsonRpcRequest {
jsonrpc,
id,
method,
params: obj
.get("params")
.cloned()
.unwrap_or(serde_json::Value::Null),
})
}
#[cfg(feature = "native")]
fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
if let Some(s) = payload.downcast_ref::<&'static str>() {
(*s).to_string()
} else if let Some(s) = payload.downcast_ref::<String>() {
s.clone()
} else {
"panic payload is not a string".to_string()
}
}
fn write_response<W: std::io::Write>(
stdout: &Arc<Mutex<W>>,
resp: &JsonRpcResponse,
) -> anyhow::Result<()> {
let json = serde_json::to_string(resp)?;
let mut guard = stdout
.lock()
.map_err(|e| anyhow::anyhow!("stdout mutex poisoned: {e}"))?;
writeln!(&mut *guard, "{json}")?;
guard.flush()?;
Ok(())
}
#[cfg(feature = "native")]
fn write_notification<W: std::io::Write>(
stdout: &Arc<Mutex<W>>,
notif: &JsonRpcNotification,
) -> anyhow::Result<()> {
let json = notif.to_json_line()?;
let mut guard = stdout
.lock()
.map_err(|e| anyhow::anyhow!("stdout mutex poisoned: {e}"))?;
writeln!(&mut *guard, "{json}")?;
guard.flush()?;
Ok(())
}
#[cfg(test)]
#[allow(clippy::disallowed_methods)] mod tests {
use super::*;
fn make_request(method: &str, params: serde_json::Value) -> JsonRpcRequest {
JsonRpcRequest {
jsonrpc: "2.0".to_string(),
id: Some(serde_json::json!(1)),
method: method.to_string(),
params,
}
}
#[cfg(feature = "native")]
fn drive(input: &[u8]) -> Vec<serde_json::Value> {
let out = Arc::new(Mutex::new(Vec::<u8>::new()));
let mut server = AprMcpServer::new();
server
.serve_stream(std::io::Cursor::new(input.to_vec()), Arc::clone(&out))
.expect("serve_stream must not propagate an error out of the session");
parse_written_lines(&out)
}
#[cfg(feature = "native")]
fn parse_written_lines(out: &Arc<Mutex<Vec<u8>>>) -> Vec<serde_json::Value> {
let guard = out.lock().expect("output mutex not poisoned");
String::from_utf8_lossy(&guard)
.lines()
.filter(|l| !l.trim().is_empty())
.map(|l| {
serde_json::from_str::<serde_json::Value>(l)
.unwrap_or_else(|e| panic!("non-JSON output line {l:?}: {e}"))
})
.collect()
}
#[cfg(feature = "native")]
fn find_id(responses: &[serde_json::Value], id: i64) -> Option<&serde_json::Value> {
responses.iter().find(|r| r["id"] == serde_json::json!(id))
}
#[cfg(feature = "native")]
const INIT_LINE: &str = r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2024-11-05"}}"#;
#[cfg(feature = "native")]
const CALL_LINE: &str = r#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"apr.version","arguments":{}}}"#;
#[cfg(feature = "native")]
#[test]
fn serve_stream_answers_tools_call_before_returning_on_eof() {
let responses = drive(format!("{INIT_LINE}\n{CALL_LINE}\n").as_bytes());
let call = find_id(&responses, 2).unwrap_or_else(|| {
panic!(
"tools/call response (id=2) was DROPPED at EOF; got {} response(s): {responses:?}",
responses.len()
)
});
assert!(
call.get("error").is_none(),
"tools/call must succeed, got {call:?}"
);
let text = call["result"]["content"][0]["text"]
.as_str()
.unwrap_or_else(|| panic!("missing content text in {call:?}"));
let payload: serde_json::Value =
serde_json::from_str(text).expect("apr.version payload is JSON");
assert_eq!(
payload["server"], "aprender-mcp",
"must be the real apr.version result, not an empty envelope"
);
assert!(
find_id(&responses, 1).is_some(),
"initialize still answered"
);
}
#[cfg(feature = "native")]
#[test]
fn serve_stream_answers_every_pipelined_tools_call() {
let mut input = format!("{INIT_LINE}\n");
for id in 2..=6 {
input.push_str(&format!(
r#"{{"jsonrpc":"2.0","id":{id},"method":"tools/call","params":{{"name":"apr.version","arguments":{{}}}}}}"#
));
input.push('\n');
}
let responses = drive(input.as_bytes());
for id in 1..=6 {
assert!(
find_id(&responses, id).is_some(),
"id={id} unanswered; got {} of 6: {responses:?}",
responses.len()
);
}
}
#[cfg(feature = "native")]
struct FailsAfterInput {
bytes: Vec<u8>,
pos: usize,
}
#[cfg(feature = "native")]
impl std::io::Read for FailsAfterInput {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if self.pos == self.bytes.len() {
return Err(std::io::Error::other("simulated stdin failure"));
}
let n = std::cmp::min(buf.len(), self.bytes.len() - self.pos);
buf[..n].copy_from_slice(&self.bytes[self.pos..self.pos + n]);
self.pos += n;
Ok(n)
}
}
#[cfg(feature = "native")]
struct StallsOnId2 {
sink: Arc<Mutex<Vec<u8>>>,
}
#[cfg(feature = "native")]
const SLOW_WRITE: std::time::Duration = std::time::Duration::from_secs(2);
#[cfg(feature = "native")]
impl std::io::Write for StallsOnId2 {
fn write(&mut self, data: &[u8]) -> std::io::Result<usize> {
if String::from_utf8_lossy(data).contains(r#""id":2"#) {
std::thread::sleep(SLOW_WRITE);
}
self.sink
.lock()
.expect("sink mutex not poisoned")
.extend_from_slice(data);
Ok(data.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
#[cfg(feature = "native")]
#[test]
fn serve_stream_drains_in_flight_workers_when_the_read_loop_aborts() {
let input = format!("{INIT_LINE}\n{CALL_LINE}\n");
let reader = std::io::BufReader::new(FailsAfterInput {
bytes: input.into_bytes(),
pos: 0,
});
let sink = Arc::new(Mutex::new(Vec::<u8>::new()));
let out = Arc::new(Mutex::new(StallsOnId2 {
sink: Arc::clone(&sink),
}));
let mut server = AprMcpServer::new();
let outcome = server.serve_stream(reader, Arc::clone(&out));
assert!(
outcome.is_err(),
"the stdin failure must still be reported after the drain"
);
let written = {
let guard = sink.lock().expect("sink mutex not poisoned");
String::from_utf8_lossy(&guard).into_owned()
};
assert!(
written.contains(r#""id":2"#),
"the in-flight tools/call was answered by NEITHER a result NOR an error \
when the read loop aborted; stdout was: {written:?}"
);
assert!(
written.contains("aprender-mcp"),
"the answer must be the real apr.version payload, not an empty envelope: {written:?}"
);
assert!(
server.workers.is_empty(),
"{} worker(s) were abandoned instead of joined",
server.workers.len()
);
}
#[cfg(feature = "native")]
#[test]
fn serve_stream_survives_invalid_utf8_line() {
let mut input: Vec<u8> = Vec::new();
input.extend_from_slice(br#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#);
input.push(b'\n');
input.push(0xFF); input.push(b'\n');
input.extend_from_slice(br#"{"jsonrpc":"2.0","id":2,"method":"ping"}"#);
input.push(b'\n');
let responses = drive(&input);
assert!(
find_id(&responses, 1).is_some(),
"request before the bad byte must be answered: {responses:?}"
);
let after = find_id(&responses, 2)
.unwrap_or_else(|| panic!("request AFTER the bad byte was lost: {responses:?}"));
assert!(
after.get("error").is_none(),
"request after the bad byte must be served normally, got {after:?}"
);
let parse_err = responses
.iter()
.find(|r| r["error"]["code"] == serde_json::json!(-32700))
.unwrap_or_else(|| panic!("the bad line itself must be reported: {responses:?}"));
assert!(
parse_err["error"]["message"]
.as_str()
.unwrap_or_default()
.contains("UTF-8"),
"the -32700 must name the cause, got {parse_err:?}"
);
}
#[cfg(feature = "native")]
#[test]
fn serve_stream_survives_leading_invalid_utf8() {
let mut input: Vec<u8> = vec![0x80, b'\n'];
input.extend_from_slice(br#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#);
input.push(b'\n');
let responses = drive(&input);
assert!(
find_id(&responses, 1).is_some(),
"the request after a leading bad byte must be answered: {responses:?}"
);
}
#[cfg(feature = "native")]
#[test]
fn serve_stream_protocol_surface_matches_jsonrpc_and_mcp() {
let input = concat!(
r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18"}}"#,
"\n",
r#"{"jsonrpc":"2.0","id":2,"method":"ping"}"#,
"\n",
r#"{"id":3,"method":"tools/list"}"#,
"\n",
r#"{"jsonrpc":"2.0","id":4}"#,
"\n",
r#"[{"jsonrpc":"2.0","id":5,"method":"tools/list"}]"#,
"\n",
r#"{not json"#,
"\n",
r#"{"jsonrpc":"2.0","id":7,"method":"tools/list"}"#,
"\n",
);
let responses = drive(input.as_bytes());
let init = find_id(&responses, 1).expect("initialize answered");
assert!(
init.get("error").is_none(),
"a newer protocolVersion must not abort the handshake: {init:?}"
);
assert_eq!(init["result"]["protocolVersion"], crate::PROTOCOL_VERSION);
let pong = find_id(&responses, 2).expect("ping answered");
assert_eq!(pong["result"], serde_json::json!({}), "ping must pong");
for id in [3, 4] {
let resp = find_id(&responses, id).unwrap_or_else(|| {
panic!("id={id} must be echoed on an Invalid Request: {responses:?}")
});
assert_eq!(
resp["error"]["code"],
serde_json::json!(-32600),
"id={id} must be Invalid Request, not Parse error: {resp:?}"
);
}
let batch = responses
.iter()
.find(|r| {
r["error"]["message"]
.as_str()
.is_some_and(|m| m.contains("batch"))
})
.unwrap_or_else(|| panic!("batch array must be diagnosed as such: {responses:?}"));
assert_eq!(batch["error"]["code"], serde_json::json!(-32600));
assert!(
responses
.iter()
.any(|r| r["error"]["code"] == serde_json::json!(-32700)),
"`{{not json` must remain a Parse error: {responses:?}"
);
assert!(
find_id(&responses, 7).is_some(),
"the loop must keep serving after every malformed line: {responses:?}"
);
}
#[cfg(feature = "native")]
#[test]
fn serve_stream_never_answers_a_notification() {
let responses = drive(
concat!(
r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#,
"\n",
r#"{"jsonrpc":"2.0","method":"tools/list"}"#,
"\n",
r#"{"jsonrpc":"2.0","id":9,"method":"ping"}"#,
"\n",
)
.as_bytes(),
);
assert_eq!(
responses.len(),
1,
"only the id-bearing request may be answered: {responses:?}"
);
assert!(find_id(&responses, 9).is_some());
}
#[test]
fn initialize_returns_protocol_version() {
let mut server = AprMcpServer::new();
let req = make_request("initialize", serde_json::json!({}));
let resp = server.handle_request(&req);
assert!(resp.error.is_none());
let result = resp.result.expect("result present");
assert_eq!(result["protocolVersion"], "2024-11-05");
assert_eq!(result["serverInfo"]["name"], "aprender-mcp");
assert!(result["capabilities"]["tools"].is_object());
}
#[test]
fn tools_list_returns_registered_tools() {
let mut server = AprMcpServer::new();
let req = make_request("tools/list", serde_json::json!({}));
let resp = server.handle_request(&req);
let result = resp.result.expect("result present");
let tools = result["tools"].as_array().expect("tools array");
let names: Vec<&str> = tools.iter().filter_map(|t| t["name"].as_str()).collect();
for expected in [
"apr.version",
"apr.validate",
"apr.tensors",
"apr.bench",
"apr.qa",
"apr.trace",
"apr.run",
"apr.serve",
"apr.finetune",
] {
assert!(names.contains(&expected), "{expected} registered");
}
for tool in tools {
assert_eq!(tool["inputSchema"]["type"], "object");
}
}
#[test]
fn tools_call_version_returns_metadata() {
let mut server = AprMcpServer::new();
let req = make_request(
"tools/call",
serde_json::json!({ "name": "apr.version", "arguments": {} }),
);
let resp = server.handle_request(&req);
let result = resp.result.expect("result present");
let text = result["content"][0]["text"].as_str().expect("text");
let parsed: serde_json::Value = serde_json::from_str(text).expect("json");
assert_eq!(parsed["server"], "aprender-mcp");
assert_eq!(parsed["protocol_version"], "2024-11-05");
}
#[test]
fn unknown_method_returns_method_not_found() {
let mut server = AprMcpServer::new();
let req = make_request("tools/explode", serde_json::json!({}));
let resp = server.handle_request(&req);
assert!(resp.result.is_none());
let err = resp.error.expect("error present");
assert_eq!(err.code, -32601);
}
#[test]
fn tools_call_validate_missing_model_path_is_error() {
let mut server = AprMcpServer::new();
let req = make_request(
"tools/call",
serde_json::json!({ "name": "apr.validate", "arguments": {} }),
);
let resp = server.handle_request(&req);
let result = resp.result.expect("result present");
assert_eq!(result["isError"], true);
let text = result["content"][0]["text"].as_str().expect("text");
assert!(text.contains("model_path"));
}
#[test]
fn tools_call_unknown_tool_returns_is_error() {
let mut server = AprMcpServer::new();
let req = make_request(
"tools/call",
serde_json::json!({ "name": "apr.nonexistent" }),
);
let resp = server.handle_request(&req);
let result = resp.result.expect("result present");
assert_eq!(result["isError"], true);
}
#[test]
fn tools_call_missing_name_returns_is_error() {
let mut server = AprMcpServer::new();
let req = make_request("tools/call", serde_json::json!({}));
let resp = server.handle_request(&req);
let result = resp.result.expect("result present");
assert_eq!(result["isError"], true);
}
#[test]
fn id_is_echoed_back() {
let mut server = AprMcpServer::new();
let req = JsonRpcRequest {
jsonrpc: "2.0".to_string(),
id: Some(serde_json::json!("req-42")),
method: "initialize".to_string(),
params: serde_json::json!({}),
};
let resp = server.handle_request(&req);
assert_eq!(resp.id, Some(serde_json::json!("req-42")));
}
#[test]
fn cancel_in_flight_signals_and_deregisters() {
let server = AprMcpServer::new();
let id = serde_json::json!(99);
let rx = AprMcpServer::register_in_flight(&server.in_flight, id.clone());
let signalled = AprMcpServer::cancel_in_flight(&server.in_flight, &id);
assert!(signalled, "live id should signal");
let received = rx.try_recv();
assert!(received.is_ok(), "cancel signal must be deliverable");
let signalled_again = AprMcpServer::cancel_in_flight(&server.in_flight, &id);
assert!(
!signalled_again,
"cancelling an already-removed id is a no-op"
);
}
#[test]
fn initialize_negotiates_down_instead_of_erroring() {
for proposed in ["2025-06-18", "2025-03-26", "2024-10-07", "latest", ""] {
let mut server = AprMcpServer::new();
let req = make_request(
"initialize",
serde_json::json!({ "protocolVersion": proposed }),
);
let resp = server.handle_request(&req);
assert!(
resp.error.is_none(),
"proposing {proposed:?} must not abort the handshake, got {:?}",
resp.error
);
let result = resp.result.expect("result present");
assert_eq!(
result["protocolVersion"],
crate::PROTOCOL_VERSION,
"server must answer with the version it actually speaks"
);
}
}
#[test]
fn initialize_ignores_non_string_protocol_version() {
let mut server = AprMcpServer::new();
let req = make_request("initialize", serde_json::json!({ "protocolVersion": 2025 }));
let resp = server.handle_request(&req);
assert!(resp.error.is_none());
let result = resp.result.expect("result present");
assert_eq!(result["protocolVersion"], crate::PROTOCOL_VERSION);
}
#[test]
fn ping_returns_empty_result() {
let mut server = AprMcpServer::new();
let req = make_request("ping", serde_json::json!({}));
let resp = server.handle_request(&req);
assert!(
resp.error.is_none(),
"ping must not error: {:?}",
resp.error
);
assert_eq!(resp.result, Some(serde_json::json!({})));
assert_eq!(resp.id, Some(serde_json::json!(1)), "id echoed");
}
#[cfg(feature = "native")]
#[test]
fn worker_response_turns_a_tool_panic_into_32603_for_the_same_id() {
let resp = AprMcpServer::worker_response(&serde_json::json!(7), || {
panic!("tool exploded while dispatching")
});
assert_eq!(
resp.id,
Some(serde_json::json!(7)),
"the id the client must correlate on has to survive the panic: {resp:?}"
);
assert!(
resp.result.is_none(),
"a panic must not also produce a result: {resp:?}"
);
let err = resp
.error
.expect("a panicking tool must produce an ERROR, never silence");
assert_eq!(
err.code, -32603,
"a tool panic is an Internal error: {err:?}"
);
assert!(
err.message.contains("panicked"),
"the -32603 must name the cause: {err:?}"
);
assert!(
err.message.contains("tool exploded while dispatching"),
"the panic message itself must reach the client: {err:?}"
);
}
#[cfg(feature = "native")]
#[test]
fn worker_response_passes_a_normal_tool_result_through_unchanged() {
let resp = AprMcpServer::worker_response(&serde_json::json!("abc"), || {
ToolCallResult::success("payload".to_string())
});
assert_eq!(resp.id, Some(serde_json::json!("abc")), "id echoed");
assert!(resp.error.is_none(), "no error on the happy path: {resp:?}");
let result = resp.result.expect("result present");
assert_eq!(
result["content"][0]["text"], "payload",
"the tool's own payload must reach the client verbatim: {result:?}"
);
}
#[cfg(feature = "native")]
fn panicking_dispatch(
_args: &serde_json::Value,
_cancel_rx: &mpsc::Receiver<()>,
_sink: Option<&NotificationSink>,
_progress_token: Option<serde_json::Value>,
) -> ToolCallResult {
panic!("PANIC-PROBE-2608 exploded inside tool dispatch")
}
#[cfg(feature = "native")]
#[test]
fn a_panicking_tool_is_answered_through_the_real_stdio_path() {
let out = Arc::new(Mutex::new(Vec::<u8>::new()));
let mut server = AprMcpServer::new();
server.worker_dispatch = panicking_dispatch;
let outcome = server.serve_stream(
std::io::Cursor::new(format!("{INIT_LINE}\n{CALL_LINE}\n").into_bytes()),
Arc::clone(&out),
);
assert!(
outcome.is_ok(),
"a panicking TOOL must not take the SESSION down: {outcome:?}"
);
let responses = parse_written_lines(&out);
let call = find_id(&responses, 2).unwrap_or_else(|| {
panic!(
"the tools/call whose tool panicked was answered by NEITHER a result NOR an \
error — the worker never routed through the panic guard; got {} response(s): \
{responses:?}",
responses.len()
)
});
assert!(
call.get("result").is_none(),
"a panic must not also produce a result: {call:?}"
);
assert_eq!(
call["error"]["code"],
serde_json::json!(-32603),
"a tool panic is an Internal error for the SAME id: {call:?}"
);
let message = call["error"]["message"]
.as_str()
.unwrap_or_else(|| panic!("error message must be a string: {call:?}"));
assert!(
message.contains("PANIC-PROBE-2608 exploded inside tool dispatch"),
"the panic's own message must reach the client, not a generic envelope: {message:?}"
);
assert!(
find_id(&responses, 1).is_some(),
"initialize must still be answered: {responses:?}"
);
assert!(
server.workers.is_empty(),
"{} worker(s) were abandoned instead of joined",
server.workers.len()
);
}
#[cfg(feature = "native")]
#[test]
fn default_worker_dispatch_is_the_real_tool_dispatcher() {
let server = AprMcpServer::new();
let (_cancel_tx, cancel_rx) = mpsc::channel();
let result = (server.worker_dispatch)(
&serde_json::json!({ "name": "apr.version", "arguments": {} }),
&cancel_rx,
None,
None,
);
assert!(
result.is_error.is_none(),
"the real dispatcher answers apr.version: {result:?}"
);
let payload: serde_json::Value = serde_json::from_str(&result.content[0].text)
.unwrap_or_else(|e| panic!("apr.version payload is JSON: {e}"));
assert_eq!(
payload["server"], "aprender-mcp",
"the default must be the real dispatcher, not a stub: {payload:?}"
);
}
#[test]
fn missing_required_field_is_invalid_request_with_id_echoed() {
for line in [
r#"{"id":1,"method":"tools/list"}"#,
r#"{"jsonrpc":"2.0","id":1}"#,
r#"{"jsonrpc":2.0,"id":1,"method":"tools/list"}"#,
r#"{"jsonrpc":"2.0","id":1,"method":42}"#,
] {
let resp = parse_incoming(line).expect_err("must be rejected");
let err = resp.error.as_ref().expect("error present");
assert_eq!(
err.code, -32600,
"{line} must be Invalid Request, got {err:?}"
);
assert_eq!(
resp.id,
Some(serde_json::json!(1)),
"{line} must echo the client's id so it can correlate"
);
}
}
#[test]
fn malformed_json_is_still_parse_error() {
for line in [
r#"{not json"#,
r#"{"jsonrpc":"2.0","id":1,"method":"tools/li"#,
] {
let resp = parse_incoming(line).expect_err("must be rejected");
let err = resp.error.as_ref().expect("error present");
assert_eq!(err.code, -32700, "{line} must remain a Parse error");
}
}
#[test]
fn batch_array_is_diagnosed_as_unsupported_batching() {
let line = r#"[{"jsonrpc":"2.0","id":1,"method":"tools/list"},{"jsonrpc":"2.0","id":2,"method":"tools/list"}]"#;
let resp = parse_incoming(line).expect_err("batch must be rejected");
let err = resp.error.as_ref().expect("error present");
assert_eq!(err.code, -32600);
assert!(
err.message.contains("batch"),
"message must name batching, got: {}",
err.message
);
assert!(
!err.message.contains("expected a string"),
"must not leak serde's field-level error, got: {}",
err.message
);
}
#[test]
fn non_object_request_is_invalid_request() {
for line in ["42", r#""hello""#, "null", "true"] {
let resp = parse_incoming(line).expect_err("must be rejected");
let err = resp.error.as_ref().expect("error present");
assert_eq!(err.code, -32600, "{line} must be Invalid Request");
}
}
#[test]
fn well_formed_request_parses_unchanged() {
let req = parse_incoming(r#"{"jsonrpc":"2.0","id":"abc","method":"tools/list"}"#)
.expect("well-formed request must parse");
assert_eq!(req.jsonrpc, "2.0");
assert_eq!(req.method, "tools/list");
assert_eq!(req.id, Some(serde_json::json!("abc")));
assert_eq!(req.params, serde_json::Value::Null);
let with_params = parse_incoming(
r#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"apr.version"}}"#,
)
.expect("params must round-trip");
assert_eq!(with_params.params["name"], "apr.version");
}
#[test]
fn null_and_absent_id_both_parse_as_notification() {
let absent = parse_incoming(r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#)
.expect("parse");
assert!(absent.id.is_none());
let null_id =
parse_incoming(r#"{"jsonrpc":"2.0","id":null,"method":"tools/list"}"#).expect("parse");
assert!(null_id.id.is_none(), "a null id is not an id");
}
#[test]
fn cancel_unknown_id_is_noop() {
let server = AprMcpServer::new();
let id = serde_json::json!("never-registered");
let signalled = AprMcpServer::cancel_in_flight(&server.in_flight, &id);
assert!(!signalled);
}
}