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 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");
}
}