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()
{
match frame.get("id").filter(|v| !v.is_null()).cloned() {
Some(id) => {
let _ = throttle_writer
.send_value(synthesized_limited(id, retry_after_ms))
.await;
}
None => auditor.on_dropped(&frame),
}
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(())
}
pub(crate) fn session_principal(identity: Option<&mcpmesh_net::PeerIdentity>) -> Option<String> {
identity.map(|id| id.endpoint.principal())
}
#[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(flavor = "multi_thread")]
async fn a_rate_limited_notification_is_audited_rather_than_silently_dropped() {
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 _keep = srv_stdout;
let mut reader = FrameReader::new(srv_stdin, MAX_FRAME_BYTES);
let mut seen = 0usize;
while let Ok(Some(Inbound::Frame(_))) = reader.next().await {
seen += 1;
}
seen
});
let rate = RateGate::new(
std::sync::Arc::new(RateLimiter::per_minute(1, 1)),
Some(mcpmesh_net::EndpointId::from_bytes([9u8; 32])),
);
let dir = tempfile::tempdir().unwrap();
let sink = AuditSink::new(crate::audit::log::AuditLog::spawn(dir.path().to_path_buf()));
let mut rx = sink.subscribe().expect("auditing enabled");
let auditor = RequestAuditor::new(sink.clone(), Some("bob".into()), "notes".into());
let pump = tokio::spawn(async move {
let _ = pump(
json!({"jsonrpc":"2.0","id":1,"method":"initialize"}),
&mut transport,
server_read,
server_write,
auditor,
rate,
)
.await;
});
for _ in 0..2 {
write_frame(
&mut peer_w,
&json!({"jsonrpc":"2.0","method":"notifications/progress","params":{}}),
)
.await
.unwrap();
}
drop(peer_w);
let mut statuses = Vec::new();
while let Ok(Ok(rec)) = timeout(Duration::from_secs(3), rx.recv()).await {
if rec.method.as_deref() == Some("notifications/progress") {
statuses.push(rec.status.clone());
if statuses.len() == 2 {
break;
}
}
}
assert!(
statuses.contains(&Some("rate_limited".into())),
"a notification dropped by the limiter must be recorded as rate_limited — \
otherwise the loss is invisible to the sender AND to the operator, which is the \
whole of #76. Saw: {statuses:?}"
);
let _ = pump.await;
let forwarded = server.await.unwrap();
assert!(
forwarded < 3,
"the over-limit notification must NOT have been forwarded: {forwarded}"
);
})
.await
.expect("dropped-notification audit test timed out");
}
#[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");
}
}