use crate::rpc::{Incoming, Notification, Request, Response, frame};
use crate::wire::method;
use serde_json::{Value, json};
use std::collections::HashMap;
use std::io::{self, BufReader, Read, Write};
use std::os::unix::net::{UnixListener, UnixStream};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PeerOrigin {
Stdio,
Management,
}
impl PeerOrigin {
pub fn as_str(self) -> &'static str {
match self {
PeerOrigin::Stdio => "stdio",
PeerOrigin::Management => "management",
}
}
}
pub enum ServeStream {
Unix(UnixStream),
#[cfg(feature = "vsock")]
Vsock(vsock::VsockStream),
Http(Box<dyn Write + Send>),
}
impl ServeStream {
pub fn try_clone(&self) -> io::Result<ServeStream> {
match self {
ServeStream::Unix(s) => s.try_clone().map(ServeStream::Unix),
#[cfg(feature = "vsock")]
ServeStream::Vsock(s) => s.try_clone().map(ServeStream::Vsock),
ServeStream::Http(_) => Err(io::Error::new(
io::ErrorKind::Unsupported,
"http SSE sink is not clonable",
)),
}
}
pub fn set_write_timeout(&self, dur: Option<Duration>) -> io::Result<()> {
match self {
ServeStream::Unix(s) => s.set_write_timeout(dur),
#[cfg(feature = "vsock")]
ServeStream::Vsock(s) => s.set_write_timeout(dur),
ServeStream::Http(_) => Ok(()),
}
}
pub fn write_notification(&mut self, note: &Notification) -> io::Result<()> {
if let ServeStream::Http(sink) = self {
let json = serde_json::to_string(note).map_err(io::Error::other)?;
sink.write_all(format!("data: {json}\n\n").as_bytes())?;
return sink.flush();
}
frame::write_line(self, note)
}
pub fn write_response(&mut self, resp: &Response) -> io::Result<()> {
if let ServeStream::Http(sink) = self {
let json = serde_json::to_string(resp).map_err(io::Error::other)?;
sink.write_all(format!("data: {json}\n\n").as_bytes())?;
return sink.flush();
}
frame::write_line(self, resp)
}
}
impl Read for ServeStream {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
ServeStream::Unix(s) => s.read(buf),
#[cfg(feature = "vsock")]
ServeStream::Vsock(s) => s.read(buf),
ServeStream::Http(_) => Ok(0),
}
}
}
impl Write for ServeStream {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
ServeStream::Unix(s) => s.write(buf),
#[cfg(feature = "vsock")]
ServeStream::Vsock(s) => s.write(buf),
ServeStream::Http(s) => s.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
ServeStream::Unix(s) => s.flush(),
#[cfg(feature = "vsock")]
ServeStream::Vsock(s) => s.flush(),
ServeStream::Http(s) => s.flush(),
}
}
}
pub type SharedWriter = Arc<Mutex<ServeStream>>;
pub struct Subscriber {
conn: u64,
writer: SharedWriter,
}
pub type SubRegistry = Arc<Mutex<HashMap<String, Vec<Subscriber>>>>;
pub fn register_subscriber(subs: &SubRegistry, uri: &str, conn: u64, writer: &SharedWriter) {
let mut g = subs.lock().unwrap_or_else(|e| e.into_inner());
let list = g.entry(uri.to_string()).or_default();
if !list.iter().any(|s| s.conn == conn) {
list.push(Subscriber {
conn,
writer: Arc::clone(writer),
});
}
}
pub fn drop_subscription(subs: &SubRegistry, uri: &str, conn: u64) {
let mut g = subs.lock().unwrap_or_else(|e| e.into_inner());
if let Some(list) = g.get_mut(uri) {
list.retain(|s| s.conn != conn);
if list.is_empty() {
g.remove(uri);
}
}
}
pub fn remove_conn_subscriptions(subs: &SubRegistry, conn: u64) {
let mut g = subs.lock().unwrap_or_else(|e| e.into_inner());
g.retain(|_uri, list| {
list.retain(|s| s.conn != conn);
!list.is_empty()
});
}
pub fn notify_resource_updated(subs: &SubRegistry, uri: &str) {
let writers: Vec<SharedWriter> = {
let mut g = subs.lock().unwrap_or_else(|e| e.into_inner());
match g.remove(uri) {
Some(list) => list.into_iter().map(|s| s.writer).collect(),
None => return,
}
};
push_updated(&writers, uri);
}
pub fn notify_resource_updated_keep(subs: &SubRegistry, uri: &str) {
let writers: Vec<SharedWriter> = {
let g = subs.lock().unwrap_or_else(|e| e.into_inner());
match g.get(uri) {
Some(list) => list.iter().map(|s| Arc::clone(&s.writer)).collect(),
None => return,
}
};
push_updated(&writers, uri);
}
fn push_updated(writers: &[SharedWriter], uri: &str) {
let note = Notification::new(
method::NOTIFY_RESOURCES_UPDATED,
Some(json!({ "uri": uri })),
);
for w in writers {
if let Ok(mut wl) = w.lock() {
let _ = wl.write_notification(¬e);
}
}
}
pub fn broadcast_distinct(subs: &SubRegistry, note: &Notification) {
let writers: Vec<SharedWriter> = {
let g = subs.lock().unwrap_or_else(|e| e.into_inner());
let mut seen: Vec<*const Mutex<ServeStream>> = Vec::new();
let mut out: Vec<SharedWriter> = Vec::new();
for list in g.values() {
for s in list {
let ptr = Arc::as_ptr(&s.writer);
if !seen.contains(&ptr) {
seen.push(ptr);
out.push(Arc::clone(&s.writer));
}
}
}
out
};
for w in writers {
if let Ok(mut wl) = w.lock() {
let _ = wl.write_notification(note);
}
}
}
pub trait Handler: Send + Sync + 'static {
fn dispatch(
&self,
req: Request,
origin: PeerOrigin,
writer: &SharedWriter,
conn: u64,
) -> Response;
fn on_connect(&self, _origin: PeerOrigin, _conn: u64) {}
fn streams(&self, _method: &str) -> bool {
false
}
fn on_disconnect(&self, _origin: PeerOrigin, _conn: u64) {}
}
pub fn lifecycle_response(
req: &Request,
server_info: &Value,
capabilities: &Value,
) -> Option<Response> {
match req.method.as_str() {
"initialize" => {
let requested = req
.params
.as_ref()
.and_then(|p| p.get("protocolVersion"))
.and_then(Value::as_str);
let version = match requested {
Some(v) if crate::version::is_supported_version(v) => v,
_ => crate::version::PROTOCOL_VERSION,
};
Some(Response::ok(
req.id.clone(),
json!({
"protocolVersion": version,
"capabilities": capabilities,
"serverInfo": server_info,
}),
))
}
method::SERVER_DISCOVER => Some(Response::ok(
req.id.clone(),
json!({
"resultType": "complete",
"supportedVersions": crate::version::SUPPORTED_PROTOCOL_VERSIONS,
"capabilities": capabilities,
"serverInfo": server_info,
}),
)),
"ping" => Some(Response::ok(req.id.clone(), json!({}))),
_ => None,
}
}
pub fn handle_conn(
stream: ServeStream,
origin: PeerOrigin,
handler: &Arc<dyn Handler>,
subs: &SubRegistry,
conn_counter: &AtomicU64,
write_timeout: Duration,
) {
let writer: SharedWriter = match stream.try_clone() {
Ok(w) => {
let _ = w.set_write_timeout(Some(write_timeout));
Arc::new(Mutex::new(w))
}
Err(_) => return,
};
let conn = conn_counter.fetch_add(1, Ordering::Relaxed);
handler.on_connect(origin, conn);
let mut reader = BufReader::new(stream);
while let Ok(Some(bytes)) = frame::read_line(&mut reader) {
if let Ok(Incoming::Request(req)) = serde_json::from_slice::<Incoming>(&bytes) {
let resp = handler.dispatch(req, origin, &writer, conn);
let wrote = writer
.lock()
.is_ok_and(|mut w| frame::write_line(&mut *w, &resp).is_ok());
if !wrote {
break; }
}
}
remove_conn_subscriptions(subs, conn); handler.on_disconnect(origin, conn);
}
pub fn bind_unix(path: &str) -> io::Result<UnixListener> {
let _ = std::fs::remove_file(path);
UnixListener::bind(path)
}
pub fn spawn_accept_unix(
listener: UnixListener,
handler: Arc<dyn Handler>,
subs: SubRegistry,
conn_counter: Arc<AtomicU64>,
write_timeout: Duration,
) -> io::Result<()> {
thread::Builder::new()
.name("serve-mcp".into())
.spawn(move || {
for stream in listener.incoming().flatten() {
let handler = Arc::clone(&handler);
let subs = Arc::clone(&subs);
let conn_counter = Arc::clone(&conn_counter);
thread::Builder::new()
.name("serve-mcp-conn".into())
.spawn(move || {
handle_conn(
ServeStream::Unix(stream),
PeerOrigin::Management,
&handler,
&subs,
&conn_counter,
write_timeout,
)
})
.ok();
}
})
.map(|_| ())
}
#[cfg(feature = "vsock")]
pub fn bind_vsock(cid: u32, port: u32) -> io::Result<vsock::VsockListener> {
vsock::VsockListener::bind_with_cid_port(cid, port)
}
#[cfg(feature = "vsock")]
pub fn spawn_accept_vsock(
listener: vsock::VsockListener,
handler: Arc<dyn Handler>,
subs: SubRegistry,
conn_counter: Arc<AtomicU64>,
write_timeout: Duration,
) -> io::Result<()> {
thread::Builder::new()
.name("serve-mcp-vsock".into())
.spawn(move || {
for stream in listener.incoming().flatten() {
let handler = Arc::clone(&handler);
let subs = Arc::clone(&subs);
let conn_counter = Arc::clone(&conn_counter);
thread::Builder::new()
.name("serve-mcp-conn".into())
.spawn(move || {
handle_conn(
ServeStream::Vsock(stream),
PeerOrigin::Management,
&handler,
&subs,
&conn_counter,
write_timeout,
)
})
.ok();
}
})
.map(|_| ())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::rpc::Request;
use std::io::BufReader;
use std::os::unix::net::UnixStream;
fn info() -> Value {
json!({"name": "test-server", "version": "9.9.9"})
}
fn caps() -> Value {
json!({"tools": {}, "resources": {"subscribe": true}})
}
#[test]
fn initialize_echoes_a_supported_requested_version() {
let want = crate::version::SUPPORTED_PROTOCOL_VERSIONS[1]; let req = Request::new(1, "initialize", Some(json!({"protocolVersion": want})));
let resp = lifecycle_response(&req, &info(), &caps()).expect("lifecycle handled");
let r = resp.result.expect("ok");
assert_eq!(r["protocolVersion"], want);
assert_eq!(r["serverInfo"]["name"], "test-server");
assert!(
r["capabilities"]["resources"]["subscribe"]
.as_bool()
.unwrap()
);
}
#[test]
fn initialize_falls_back_to_latest_legacy_for_an_unsupported_version() {
let req = Request::new(
1,
"initialize",
Some(json!({"protocolVersion": "1999-01-01"})),
);
let resp = lifecycle_response(&req, &info(), &caps()).expect("handled");
let r = resp.result.expect("ok");
assert_eq!(r["protocolVersion"], crate::version::PROTOCOL_VERSION);
}
#[test]
fn initialize_defaults_when_no_version_is_requested() {
let req = Request::new(1, "initialize", Some(json!({})));
let resp = lifecycle_response(&req, &info(), &caps()).expect("handled");
assert_eq!(
resp.result.expect("ok")["protocolVersion"],
crate::version::PROTOCOL_VERSION
);
}
#[test]
fn server_discover_advertises_every_supported_version() {
let req = Request::new(7, method::SERVER_DISCOVER, None);
let resp = lifecycle_response(&req, &info(), &caps()).expect("handled");
let r = resp.result.expect("ok");
assert_eq!(r["resultType"], "complete");
let listed = r["supportedVersions"].as_array().expect("array");
assert_eq!(
listed.len(),
crate::version::SUPPORTED_PROTOCOL_VERSIONS.len()
);
assert!(
listed
.iter()
.any(|v| v == crate::version::FIRST_MODERN_VERSION)
);
assert!(listed.iter().any(|v| v == crate::version::PROTOCOL_VERSION));
assert_eq!(r["serverInfo"]["name"], "test-server");
}
#[test]
fn ping_is_an_empty_ok_and_domain_methods_fall_through() {
let ping = Request::new(1, "ping", None);
assert_eq!(
lifecycle_response(&ping, &info(), &caps())
.expect("handled")
.result,
Some(json!({}))
);
let dom = Request::new(2, "tools/call", None);
assert!(lifecycle_response(&dom, &info(), &caps()).is_none());
}
fn wired() -> (SharedWriter, BufReader<UnixStream>) {
let (tx, rx) = UnixStream::pair().unwrap();
rx.set_read_timeout(Some(Duration::from_millis(250)))
.unwrap();
(
Arc::new(Mutex::new(ServeStream::Unix(tx))),
BufReader::new(rx),
)
}
fn read_note(rx: &mut BufReader<UnixStream>) -> (String, String) {
let bytes = frame::read_line(rx).expect("read").expect("a frame");
let v: Value = serde_json::from_slice(&bytes).expect("json");
let method = v["method"].as_str().unwrap_or_default().to_string();
let uri = v["params"]["uri"].as_str().unwrap_or_default().to_string();
(method, uri)
}
fn assert_silent(rx: &mut BufReader<UnixStream>) {
assert!(
!matches!(frame::read_line(rx), Ok(Some(_))),
"expected no further push"
);
}
#[test]
fn register_is_idempotent_per_connection() {
let subs: SubRegistry = Arc::new(Mutex::new(HashMap::new()));
let (w, _rx) = wired();
register_subscriber(&subs, "res://a", 1, &w);
register_subscriber(&subs, "res://a", 1, &w); let g = subs.lock().unwrap();
assert_eq!(g.get("res://a").unwrap().len(), 1);
}
#[test]
fn notify_updated_consumes_but_keep_retains() {
let subs: SubRegistry = Arc::new(Mutex::new(HashMap::new()));
let (w, mut rx) = wired();
register_subscriber(&subs, "res://run", 1, &w);
notify_resource_updated_keep(&subs, "res://run");
let (m, uri) = read_note(&mut rx);
assert_eq!(m, method::NOTIFY_RESOURCES_UPDATED);
assert_eq!(uri, "res://run");
assert!(subs.lock().unwrap().contains_key("res://run"));
notify_resource_updated(&subs, "res://run");
let (_m, uri2) = read_note(&mut rx);
assert_eq!(uri2, "res://run");
assert!(!subs.lock().unwrap().contains_key("res://run"));
}
#[test]
fn drop_and_conn_cleanup_remove_subscriptions() {
let subs: SubRegistry = Arc::new(Mutex::new(HashMap::new()));
let (w, _rx) = wired();
register_subscriber(&subs, "res://a", 1, &w);
register_subscriber(&subs, "res://b", 1, &w);
drop_subscription(&subs, "res://a", 1);
assert!(!subs.lock().unwrap().contains_key("res://a"));
assert!(subs.lock().unwrap().contains_key("res://b"));
remove_conn_subscriptions(&subs, 1);
assert!(subs.lock().unwrap().is_empty());
}
#[test]
fn broadcast_distinct_writes_once_per_connection() {
let subs: SubRegistry = Arc::new(Mutex::new(HashMap::new()));
let (w, mut rx) = wired();
register_subscriber(&subs, "res://a", 1, &w);
register_subscriber(&subs, "res://b", 1, &w);
let note = Notification::new(method::NOTIFY_TOOLS_LIST_CHANGED, None);
broadcast_distinct(&subs, ¬e);
let (m, _) = read_note(&mut rx);
assert_eq!(m, method::NOTIFY_TOOLS_LIST_CHANGED);
assert_silent(&mut rx); }
}