use std::collections::BTreeMap;
use std::io::{BufRead, BufReader, Write};
use serde_json::{Value, json};
use super::protocol::{
self, Request, Response, RpcError, RpcId, decode_request, encode_message, error_response,
invalid_request, method_not_found, parse_error, success_response,
};
pub type MethodHandler = Box<dyn Fn(Option<&Value>) -> Result<Value, RpcError> + Send + Sync>;
pub type Methods = BTreeMap<String, MethodHandler>;
pub const SERVER_PROTOCOL_VERSION: &str = "2024-11-05";
pub const SERVER_NAME: &str = "kb";
pub const SERVER_VERSION: &str = env!("CARGO_PKG_VERSION");
#[must_use]
pub fn builtin_methods() -> Methods {
let mut m: Methods = BTreeMap::new();
m.insert(
"initialize".into(),
Box::new(|_params: Option<&Value>| {
Ok(json!({
"protocolVersion": SERVER_PROTOCOL_VERSION,
"capabilities": { "tools": {} },
"serverInfo": {
"name": SERVER_NAME,
"version": SERVER_VERSION,
}
}))
}),
);
m.insert(
"notifications/initialized".into(),
Box::new(|_params: Option<&Value>| Ok(Value::Null)),
);
m
}
#[must_use]
pub fn handle_line(methods: &Methods, line: &str) -> Option<Response> {
match decode_request(line) {
Err(err) => Some(error_response(
Some(RpcId::Null),
parse_error(&err.to_string()),
)),
Ok(req) => dispatch(methods, &req),
}
}
fn dispatch(methods: &Methods, req: &Request) -> Option<Response> {
if req.jsonrpc != "2.0" {
return Some(error_response(
req.id.clone(),
invalid_request(&format!("expected jsonrpc=\"2.0\", got: {}", req.jsonrpc)),
));
}
let params = req.params.as_ref();
if req.id.is_none() {
if let Some(h) = methods.get(&req.method) {
if let Err(e) = h(params) {
tracing::warn!("notification handler {} failed: {e:?}", req.method);
}
}
return None;
}
match methods.get(&req.method) {
None => Some(error_response(
req.id.clone(),
method_not_found(&req.method),
)),
Some(h) => match h(params) {
Ok(v) => Some(success_response(req.id.clone(), v)),
Err(e) => Some(error_response(req.id.clone(), e)),
},
}
}
pub fn run_mcp_with(methods: &Methods) -> std::io::Result<()> {
let stdin = std::io::stdin();
let mut stdin = BufReader::new(stdin.lock());
let stdout = std::io::stdout();
let mut out = stdout.lock();
let mut buf = String::new();
loop {
buf.clear();
let n = stdin.read_line(&mut buf)?;
if n == 0 {
return Ok(());
}
let line = buf.trim_end_matches(['\n', '\r']);
if line.is_empty() {
continue;
}
if let Some(resp) = handle_line(methods, line) {
let encoded = encode_message(&resp).unwrap_or_else(|e| {
let fallback = error_response(
resp.id.clone(),
protocol::internal_error(&format!("response encode failed: {e}")),
);
encode_message(&fallback).unwrap_or_else(|_| {
r#"{"jsonrpc":"2.0","id":null,"error":{"code":-32603,"message":"Internal error"}}"#.to_string()
})
});
out.write_all(encoded.as_bytes())?;
out.write_all(b"\n")?;
out.flush()?;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mcp_builtin_methods_initialize_returns_protocol_version() {
let m = builtin_methods();
let resp = handle_line(&m, r#"{"jsonrpc":"2.0","id":1,"method":"initialize"}"#).unwrap();
let result = resp.result.unwrap();
assert_eq!(result["protocolVersion"], "2024-11-05");
assert_eq!(result["serverInfo"]["name"], "kb");
assert_eq!(result["serverInfo"]["version"], SERVER_VERSION);
assert!(result["capabilities"]["tools"].is_object());
}
#[test]
fn mcp_server_loop_notifications_produce_no_response() {
let m = builtin_methods();
let r = handle_line(
&m,
r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#,
);
assert!(r.is_none());
}
#[test]
fn mcp_server_loop_unknown_method_returns_method_not_found() {
let m = builtin_methods();
let resp = handle_line(&m, r#"{"jsonrpc":"2.0","id":1,"method":"nope"}"#).unwrap();
let err = resp.error.unwrap();
assert_eq!(err.code, -32601);
}
#[test]
fn mcp_server_loop_unknown_notification_is_silently_dropped() {
let m = builtin_methods();
let r = handle_line(&m, r#"{"jsonrpc":"2.0","method":"nope/unknown"}"#);
assert!(r.is_none());
}
#[test]
fn mcp_server_loop_parse_failure_returns_error_with_null_id() {
let m = builtin_methods();
let resp = handle_line(&m, "not json").unwrap();
let err = resp.error.unwrap();
assert_eq!(err.code, -32700);
assert_eq!(resp.id, Some(RpcId::Null));
}
#[test]
fn mcp_server_loop_wrong_jsonrpc_version_returns_invalid_request() {
let m = builtin_methods();
let resp = handle_line(&m, r#"{"jsonrpc":"1.0","id":1,"method":"initialize"}"#).unwrap();
let err = resp.error.unwrap();
assert_eq!(err.code, -32600);
}
#[test]
fn mcp_builtin_methods_initialize_field_names_and_capabilities() {
let m = builtin_methods();
let resp = handle_line(&m, r#"{"jsonrpc":"2.0","id":1,"method":"initialize"}"#).unwrap();
let r = resp.result.unwrap();
assert!(r.get("protocolVersion").is_some());
assert!(r.get("capabilities").is_some());
assert!(r.get("serverInfo").is_some());
assert_eq!(r["serverInfo"]["name"], "kb");
assert_eq!(r["serverInfo"]["version"], env!("CARGO_PKG_VERSION"));
assert!(r["capabilities"]["tools"].is_object());
}
#[test]
fn mcp_builtin_methods_initialized_notification_is_dropped() {
let m = builtin_methods();
let r = handle_line(
&m,
r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#,
);
assert!(r.is_none());
}
#[test]
fn mcp_builtin_methods_registry_lists_lifecycle_keys() {
let m = builtin_methods();
assert!(m.contains_key("initialize"));
assert!(m.contains_key("notifications/initialized"));
}
}