use anyhow::{Context, Result};
use mcpmesh_net::errors::synthesized_limited;
use mcpmesh_net::framing::{FrameReader, Inbound, write_frame};
use mcpmesh_net::transport::NdjsonTransport;
use serde_json::Value;
use tokio::io::{AsyncRead, AsyncWrite, BufReader};
pub mod socket;
pub mod spawn;
pub(crate) const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
pub(crate) async fn pump<TR, TW, SR, SW>(
initialize: Value,
transport: &mut NdjsonTransport<TR, TW>,
server_read: SR,
mut server_write: SW,
auditor: crate::audit::RequestAuditor,
rate: crate::limits::RateGate,
) -> Result<()>
where
TR: AsyncRead + Send + Unpin,
TW: AsyncWrite + Send + Unpin,
SR: AsyncRead + Send + Unpin,
SW: AsyncWrite + Send + Unpin,
{
write_frame(&mut server_write, &initialize)
.await
.context("forward initialize to local MCP server")?;
let mut server_write = Some(server_write);
let throttle_writer = transport.writer();
let transport_writer = transport.writer();
let mut server_out = FrameReader::new(BufReader::new(server_read), MAX_FRAME_BYTES);
let to_server = async {
loop {
match transport.recv_value().await {
Ok(Some(frame)) => {
if frame.get("method").is_some()
&& let Err(retry_after_ms) = rate.admit()
{
if let Some(id) = frame.get("id").filter(|v| !v.is_null()).cloned() {
let _ = throttle_writer
.send_value(synthesized_limited(id, retry_after_ms))
.await;
}
continue;
}
auditor.on_request(&frame);
let Some(w) = server_write.as_mut() else {
break;
};
if write_frame(w, &frame).await.is_err() {
break; }
}
Ok(None) => break, Err(_) => break, }
}
if let Some(mut w) = server_write.take() {
use tokio::io::AsyncWriteExt;
let _ = w.shutdown().await;
} std::future::pending::<()>().await
};
let to_transport = async {
loop {
match server_out.next().await {
Ok(Some(Inbound::Frame(frame))) => {
let bytes_out = serde_json::to_vec(&frame)
.map(|v| v.len() as u64)
.unwrap_or(0);
auditor.on_response(&frame, bytes_out);
if transport_writer.send_value(frame).await.is_err() {
break; }
}
Ok(Some(Inbound::Violation(_))) => break,
Ok(None) => break, Err(_) => break, }
}
};
tokio::select! {
() = to_server => {}
() = to_transport => {}
}
let _ = transport.shutdown().await;
Ok(())
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use serde_json::json;
use tokio::io::duplex;
use tokio::time::timeout;
use super::*;
use crate::audit::{AuditSink, RequestAuditor};
use crate::limits::{RateGate, RateLimiter};
#[tokio::test]
async fn a_server_initiated_notification_reaches_the_peer_unmetered() {
timeout(Duration::from_secs(10), async {
let (mut peer_w, tr) = duplex(64 * 1024);
let (tw, peer_r) = duplex(64 * 1024);
let mut transport = NdjsonTransport::new(tr, tw, MAX_FRAME_BYTES);
let (server_write, srv_stdin) = duplex(64 * 1024);
let (srv_stdout, server_read) = duplex(64 * 1024);
let server = tokio::spawn(async move {
let mut srv_w = srv_stdout;
let mut reader = FrameReader::new(srv_stdin, MAX_FRAME_BYTES);
let mut seen = 0usize;
while let Ok(Some(Inbound::Frame(f))) = reader.next().await {
seen += 1;
if f["method"] == "tools/call" {
write_frame(&mut srv_w, &json!({"jsonrpc":"2.0","id":f["id"]}))
.await
.unwrap();
write_frame(
&mut srv_w,
&json!({"jsonrpc":"2.0","method":"notifications/message",
"params":{"level":"info","data":"pushed"}}),
)
.await
.unwrap();
}
}
seen
});
let rate = RateGate::new(
std::sync::Arc::new(crate::limits::RateLimiter::per_minute(1, 1)),
Some(mcpmesh_net::EndpointId::from_bytes([9u8; 32])),
);
let pump_task = tokio::spawn(async move {
pump(
json!({"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}),
&mut transport,
server_read,
server_write,
RequestAuditor::new(AuditSink::disabled(), Some("bob".into()), "echo".into()),
rate,
)
.await
});
write_frame(
&mut peer_w,
&json!({"jsonrpc":"2.0","id":2,"method":"tools/call","params":{}}),
)
.await
.unwrap();
write_frame(
&mut peer_w,
&json!({"jsonrpc":"2.0","id":3,"method":"tools/call","params":{}}),
)
.await
.unwrap();
let mut peer_reader = FrameReader::new(peer_r, MAX_FRAME_BYTES);
let mut saw_limited = false;
let reply = loop {
match peer_reader.next().await.unwrap() {
Some(Inbound::Frame(f)) => {
if f["error"]["code"] == -32053 {
saw_limited = true;
continue;
}
if f["id"] == 2 {
break f;
}
}
other => panic!("expected the id=2 reply, got {other:?}"),
}
};
assert_eq!(reply["id"], 2, "the solicited reply arrives");
assert!(
saw_limited,
"the per-identity budget must be EXHAUSTED by now — without that this test cannot \
distinguish 'Direction B is unmetered' from 'there was budget left'"
);
let pushed = match peer_reader.next().await.unwrap() {
Some(Inbound::Frame(f)) => f,
other => panic!("expected the server-initiated notification, got {other:?}"),
};
assert_eq!(
pushed["method"], "notifications/message",
"an unsolicited server notification must reach the peer — this is what makes push \
possible instead of polling (#91): {pushed}"
);
assert!(
pushed.get("id").is_none_or(|v| v.is_null()),
"and it is a notification, not a request: {pushed}"
);
drop(peer_w);
let _ = server.await;
let _ = pump_task.await;
})
.await
.expect("server-initiated notification test timed out");
}
#[tokio::test]
async fn transport_eof_does_not_drop_replies_still_inside_the_server() {
timeout(Duration::from_secs(10), async {
for _ in 0..25 {
let (mut peer_w, tr) = duplex(64 * 1024);
let (tw, peer_r) = duplex(64 * 1024);
let mut transport = NdjsonTransport::new(tr, tw, MAX_FRAME_BYTES);
let (server_write, srv_stdin) = duplex(64 * 1024);
let (srv_stdout, server_read) = duplex(64 * 1024);
let server = tokio::spawn(async move {
let mut srv_w = srv_stdout;
let mut reader = FrameReader::new(srv_stdin, MAX_FRAME_BYTES);
let mut seen = Vec::new();
while let Ok(Some(Inbound::Frame(f))) = reader.next().await {
seen.push(f);
}
for f in &seen {
write_frame(&mut srv_w, &json!({"jsonrpc": "2.0", "id": f["id"]}))
.await
.unwrap();
}
seen.len()
});
let pump_task = tokio::spawn(async move {
pump(
json!({"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}),
&mut transport,
server_read,
server_write,
RequestAuditor::new(
AuditSink::disabled(),
Some("bob".into()),
"echo".into(),
),
RateGate::new(RateLimiter::unlimited_shared(), None),
)
.await
});
write_frame(
&mut peer_w,
&json!({"jsonrpc":"2.0","id":2,"method":"tools/call","params":{}}),
)
.await
.unwrap();
drop(peer_w);
let mut peer_reader = FrameReader::new(peer_r, MAX_FRAME_BYTES);
for expect_id in [1, 2] {
match peer_reader.next().await.unwrap() {
Some(Inbound::Frame(f)) => assert_eq!(
f["id"], expect_id,
"the reply to request {expect_id} must survive transport EOF: {f}"
),
other => panic!("expected the id={expect_id} reply, got {other:?}"),
}
}
assert!(
peer_reader.next().await.unwrap().is_none(),
"after the server's output EOF the session closes cleanly"
);
assert_eq!(
server.await.unwrap(),
2,
"the server saw initialize + request"
);
pump_task.await.unwrap().expect("pump returns Ok");
}
})
.await
.expect("pump drain test timed out");
}
}