use anyhow::Result;
use futures_util::StreamExt;
use futures_util::stream::{self, Stream};
use mcp_execution_server::service::GeneratorService;
use rmcp::RoleServer;
use rmcp::ServiceExt;
use rmcp::model::{GetExtensions, JsonRpcMessage};
use rmcp::service::{RxJsonRpcMessage, TxJsonRpcMessage};
use rmcp::transport::async_rw::{JsonRpcMessageCodec, JsonRpcMessageCodecError};
use rmcp::transport::stdio;
use std::collections::VecDeque;
use std::sync::Arc;
use std::task::Poll;
use tokio::io::AsyncRead;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tokio_util::bytes::BytesMut;
use tokio_util::codec::{Decoder, FramedRead, FramedWrite};
use tokio_util::sync::PollSemaphore;
use tracing_subscriber::{EnvFilter, layer::SubscriberExt, util::SubscriberInitExt};
const MAX_REQUEST_LINE_SIZE: usize = 4 * 1024 * 1024;
const MAX_CONCURRENT_REQUESTS: usize = 8;
fn attach_permit(
mut message: RxJsonRpcMessage<RoleServer>,
permit: OwnedSemaphorePermit,
) -> RxJsonRpcMessage<RoleServer> {
if let JsonRpcMessage::Request(ref mut request) = message {
request.request.extensions_mut().insert(Arc::new(permit));
}
message
}
enum DecodedFrame {
Message(Box<RxJsonRpcMessage<RoleServer>>),
Malformed(JsonRpcMessageCodecError),
Skipped,
}
struct RecoveringCodec {
inner: JsonRpcMessageCodec<RxJsonRpcMessage<RoleServer>>,
blank_scan_from: usize,
assume_mid_discard: bool,
}
impl RecoveringCodec {
fn new(max_length: usize) -> Self {
Self {
inner: JsonRpcMessageCodec::new_with_max_length(max_length),
blank_scan_from: 0,
assume_mid_discard: false,
}
}
fn peek_blank_line(scan_from: &mut usize, max_length: usize, buf: &BytesMut) -> Option<bool> {
let bound = std::cmp::min(max_length.saturating_add(1), buf.len());
if *scan_from >= bound {
return None;
}
if let Some(offset) = buf[*scan_from..bound]
.iter()
.position(|&byte| byte == b'\n')
{
let newline_at = *scan_from + offset;
Some(buf[..newline_at].iter().all(u8::is_ascii_whitespace))
} else {
*scan_from = bound;
None
}
}
fn fold(
result: Result<Option<RxJsonRpcMessage<RoleServer>>, JsonRpcMessageCodecError>,
len_before: usize,
len_after: usize,
is_blank: bool,
) -> std::io::Result<Option<DecodedFrame>> {
match result {
Ok(Some(message)) => Ok(Some(DecodedFrame::Message(Box::new(message)))),
Ok(None) if len_after < len_before => Ok(Some(DecodedFrame::Skipped)),
Ok(None) => Ok(None),
Err(JsonRpcMessageCodecError::Io(error)) => Err(error),
Err(JsonRpcMessageCodecError::Serde(_)) if is_blank => Ok(Some(DecodedFrame::Skipped)),
Err(other) => Ok(Some(DecodedFrame::Malformed(other))),
}
}
fn drive(
&mut self,
buf: &mut BytesMut,
decode_step: impl FnOnce(
&mut JsonRpcMessageCodec<RxJsonRpcMessage<RoleServer>>,
&mut BytesMut,
) -> Result<
Option<RxJsonRpcMessage<RoleServer>>,
JsonRpcMessageCodecError,
>,
) -> std::io::Result<Option<DecodedFrame>> {
let max_length = self.inner.max_length();
let is_blank = !self.assume_mid_discard
&& Self::peek_blank_line(&mut self.blank_scan_from, max_length, buf).unwrap_or(false);
let len_before = buf.len();
let result = decode_step(&mut self.inner, buf);
let len_after = buf.len();
if len_after < len_before {
self.blank_scan_from = 0;
}
match &result {
Ok(Some(_)) | Err(JsonRpcMessageCodecError::Serde(_)) => {
self.assume_mid_discard = false;
}
Err(JsonRpcMessageCodecError::MaxLineLengthExceeded) => {
self.assume_mid_discard = true;
}
Ok(None) | Err(_) => {}
}
Self::fold(result, len_before, len_after, is_blank)
}
}
impl Decoder for RecoveringCodec {
type Item = DecodedFrame;
type Error = std::io::Error;
fn decode(&mut self, buf: &mut BytesMut) -> std::io::Result<Option<DecodedFrame>> {
self.drive(buf, JsonRpcMessageCodec::decode)
}
fn decode_eof(&mut self, buf: &mut BytesMut) -> std::io::Result<Option<DecodedFrame>> {
self.drive(buf, JsonRpcMessageCodec::decode_eof)
}
}
fn bounded_request_stream<R>(
reader: R,
max_length: usize,
concurrency_limit: Arc<Semaphore>,
) -> impl Stream<Item = RxJsonRpcMessage<RoleServer>> + Send + Unpin + 'static
where
R: AsyncRead + Send + Unpin + 'static,
{
let mut framed = FramedRead::new(reader, RecoveringCodec::new(max_length));
let mut semaphore = PollSemaphore::new(concurrency_limit);
let mut pending_admission: VecDeque<RxJsonRpcMessage<RoleServer>> = VecDeque::new();
stream::poll_fn(move |cx| {
loop {
if !pending_admission.is_empty() {
match semaphore.poll_acquire(cx) {
Poll::Ready(Some(permit)) => {
let message = pending_admission
.pop_front()
.expect("just checked pending_admission is non-empty");
return Poll::Ready(Some(attach_permit(message, permit)));
}
Poll::Ready(None) => {
tracing::error!("request concurrency semaphore closed unexpectedly");
return Poll::Ready(None);
}
Poll::Pending => {
if pending_admission.len() >= MAX_CONCURRENT_REQUESTS {
return Poll::Pending;
}
}
}
}
return match framed.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(DecodedFrame::Message(message)))) => {
let message = *message;
if matches!(message, JsonRpcMessage::Request(_)) {
pending_admission.push_back(message);
continue;
}
Poll::Ready(Some(message))
}
Poll::Ready(Some(Ok(DecodedFrame::Malformed(reason)))) => {
tracing::warn!(%reason, "dropping oversized or malformed request line");
continue;
}
Poll::Ready(Some(Ok(DecodedFrame::Skipped))) => {
tracing::trace!("inner codec consumed input without producing a message");
continue;
}
Poll::Ready(Some(Err(error))) => {
tracing::error!(%error, "stdin read failed; ending session");
Poll::Ready(None)
}
Poll::Ready(None) if !pending_admission.is_empty() => {
Poll::Pending
}
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
};
}
})
}
#[tokio::main]
async fn main() -> Result<()> {
tracing_subscriber::registry()
.with(
EnvFilter::try_from_default_env()
.unwrap_or_else(|_| EnvFilter::new("info,mcp_execution_server=debug")),
)
.with(
tracing_subscriber::fmt::layer()
.with_writer(std::io::stderr)
.with_target(true),
)
.init();
tracing::info!(
"Starting mcp-execution-server v{}",
env!("CARGO_PKG_VERSION")
);
let (stdin, stdout) = stdio();
let sink = FramedWrite::new(
stdout,
JsonRpcMessageCodec::<TxJsonRpcMessage<RoleServer>>::new(),
);
let stream = bounded_request_stream(
stdin,
MAX_REQUEST_LINE_SIZE,
Arc::new(Semaphore::new(MAX_CONCURRENT_REQUESTS)),
);
let service = GeneratorService::new().serve((sink, stream)).await?;
service.waiting().await?;
tracing::info!("Server shutdown complete");
Ok(())
}
#[cfg(test)]
mod tests {
use super::{
AsyncRead, GetExtensions, JsonRpcMessage, MAX_CONCURRENT_REQUESTS, MAX_REQUEST_LINE_SIZE,
OwnedSemaphorePermit, RoleServer, RxJsonRpcMessage, Semaphore, Stream, StreamExt,
bounded_request_stream,
};
use mcp_execution_server::service::GeneratorService;
use rmcp::ServiceExt;
use rmcp::model::NumberOrString;
use rmcp::service::TxJsonRpcMessage;
use rmcp::transport::async_rw::JsonRpcMessageCodec;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Context, Poll, Waker};
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader, ReadBuf};
use tokio_util::codec::FramedWrite;
const TEST_MAX: usize = 64;
const TEST_CONCURRENCY: usize = 8;
fn semaphore(permits: usize) -> Arc<Semaphore> {
Arc::new(Semaphore::new(permits))
}
async fn next_or_timeout<S>(stream: &mut S) -> Option<S::Item>
where
S: Stream + Unpin,
{
tokio::time::timeout(Duration::from_secs(2), stream.next())
.await
.expect("stream.next() must resolve within 2s instead of hanging")
}
fn request_id(message: &RxJsonRpcMessage<RoleServer>) -> i64 {
match message {
JsonRpcMessage::Request(request) => match &request.id {
NumberOrString::Number(id) => *id,
NumberOrString::String(id) => {
panic!("test fixtures only use numeric request ids, got {id:?}")
}
},
_ => panic!("expected a JsonRpcMessage::Request"),
}
}
fn oversized_line() -> Vec<u8> {
let mut line = vec![b'x'; TEST_MAX * 3];
line.push(b'\n');
line
}
fn valid_notification_line() -> Vec<u8> {
br#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#
.iter()
.copied()
.chain(std::iter::once(b'\n'))
.collect()
}
fn valid_request_line(id: u64) -> Vec<u8> {
format!(r#"{{"jsonrpc":"2.0","id":{id},"method":"ping"}}"#)
.into_bytes()
.into_iter()
.chain(std::iter::once(b'\n'))
.collect()
}
fn take_permit(message: &mut RxJsonRpcMessage<RoleServer>) -> Arc<OwnedSemaphorePermit> {
match message {
JsonRpcMessage::Request(request) => request
.request
.extensions_mut()
.get::<Arc<OwnedSemaphorePermit>>()
.cloned()
.expect("a decoded request must carry a permit inserted by attach_permit"),
_ => panic!("expected a JsonRpcMessage::Request"),
}
}
struct Script {
chunks: Vec<Vec<u8>>,
idx: usize,
}
impl AsyncRead for Script {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if self.idx < self.chunks.len() {
let chunk = self.chunks[self.idx].clone();
self.idx += 1;
debug_assert!(
chunk.len() <= buf.remaining(),
"test fixture chunk exceeds the reader's spare buffer capacity"
);
buf.put_slice(&chunk);
return Poll::Ready(Ok(()));
}
Poll::Ready(Ok(())) }
}
struct ScriptThenIdle {
chunks: Vec<Vec<u8>>,
idx: usize,
}
impl AsyncRead for ScriptThenIdle {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if self.idx < self.chunks.len() {
let chunk = self.chunks[self.idx].clone();
self.idx += 1;
debug_assert!(
chunk.len() <= buf.remaining(),
"test fixture chunk exceeds the reader's spare buffer capacity"
);
buf.put_slice(&chunk);
return Poll::Ready(Ok(()));
}
Poll::Pending
}
}
const MAX_ERR_AFTER_POLLS: usize = 8;
struct ErrAfter {
chunks: Vec<Vec<u8>>,
idx: usize,
error_polls: Arc<AtomicUsize>,
}
impl AsyncRead for ErrAfter {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if self.idx < self.chunks.len() {
let chunk = self.chunks[self.idx].clone();
self.idx += 1;
debug_assert!(
chunk.len() <= buf.remaining(),
"test fixture chunk exceeds the reader's spare buffer capacity"
);
buf.put_slice(&chunk);
return Poll::Ready(Ok(()));
}
if self.error_polls.fetch_add(1, Ordering::SeqCst) >= MAX_ERR_AFTER_POLLS {
return Poll::Ready(Ok(())); }
Poll::Ready(Err(io::Error::other("persistent read failure")))
}
}
#[tokio::test]
async fn recovers_from_oversized_lines_and_keeps_serving() {
let script = Script {
chunks: vec![
oversized_line(),
oversized_line(),
valid_notification_line(),
],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(
stream.next().await.is_some(),
"the trailing valid line must still decode after two oversized lines"
);
assert!(
stream.next().await.is_none(),
"stream ends cleanly at EOF after the valid message"
);
}
#[tokio::test]
async fn oversized_line_and_valid_request_in_one_chunk_do_not_stall() {
let mut one_chunk = oversized_line();
one_chunk.extend(valid_request_line(1));
let script = ScriptThenIdle {
chunks: vec![one_chunk],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
let message = next_or_timeout(&mut stream)
.await
.expect("the valid request sharing a chunk with the oversized line must decode without waiting for more input");
assert_eq!(request_id(&message), 1);
}
#[tokio::test]
async fn malformed_json_and_valid_request_in_one_chunk_do_not_stall() {
let mut one_chunk = b"not valid json\n".to_vec();
one_chunk.extend(valid_request_line(1));
let script = ScriptThenIdle {
chunks: vec![one_chunk],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
let message = next_or_timeout(&mut stream)
.await
.expect("the valid request sharing a chunk with the malformed line must decode without waiting for more input");
assert_eq!(request_id(&message), 1);
}
#[tokio::test]
async fn multiple_malformed_lines_then_valid_request_in_one_chunk_do_not_stall() {
let mut one_chunk = oversized_line();
one_chunk.extend(b"not valid json\n");
one_chunk.extend(oversized_line());
one_chunk.extend(valid_request_line(1));
let script = ScriptThenIdle {
chunks: vec![one_chunk],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
let message = next_or_timeout(&mut stream).await.expect(
"the valid request behind three consecutive malformed lines in one chunk must decode without waiting for more input",
);
assert_eq!(request_id(&message), 1);
}
#[tokio::test]
async fn chunk_with_only_malformed_lines_and_no_valid_request_stays_pending() {
let mut one_chunk = oversized_line();
one_chunk.extend(b"not valid json\n");
let script = ScriptThenIdle {
chunks: vec![one_chunk],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
match Pin::new(&mut stream).poll_next(&mut cx) {
Poll::Pending => {}
Poll::Ready(item) => panic!(
"expected Poll::Pending after discarding only-malformed lines with no valid request behind them, got Some={}",
item.is_some()
),
}
}
struct WarnCounter(Arc<AtomicUsize>);
impl tracing::Subscriber for WarnCounter {
fn enabled(&self, metadata: &tracing::Metadata<'_>) -> bool {
*metadata.level() == tracing::Level::WARN
}
fn new_span(&self, _span: &tracing::span::Attributes<'_>) -> tracing::span::Id {
tracing::span::Id::from_u64(1)
}
fn record(&self, _span: &tracing::span::Id, _values: &tracing::span::Record<'_>) {}
fn record_follows_from(&self, _span: &tracing::span::Id, _follows: &tracing::span::Id) {}
fn event(&self, event: &tracing::Event<'_>) {
if *event.metadata().level() == tracing::Level::WARN {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
fn enter(&self, _span: &tracing::span::Id) {}
fn exit(&self, _span: &tracing::span::Id) {}
}
#[tokio::test]
async fn blank_lines_are_skipped_without_a_warn_log() {
let warn_count = Arc::new(AtomicUsize::new(0));
let _tracing_guard = tracing::subscriber::set_default(WarnCounter(Arc::clone(&warn_count)));
let script = Script {
chunks: vec![
b"\n".to_vec(),
b" \n".to_vec(),
b"\r\n".to_vec(),
valid_notification_line(),
],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(
stream.next().await.is_some(),
"the valid line must decode with no item emitted for the blank lines before it"
);
assert!(stream.next().await.is_none());
assert_eq!(
warn_count.load(Ordering::SeqCst),
0,
"blank lines must not produce a warn!-level log record"
);
}
#[tokio::test]
async fn malformed_non_blank_line_still_warns() {
let warn_count = Arc::new(AtomicUsize::new(0));
let _tracing_guard = tracing::subscriber::set_default(WarnCounter(Arc::clone(&warn_count)));
let mut one_chunk = b"not valid json\n".to_vec();
one_chunk.extend(valid_notification_line());
let script = Script {
chunks: vec![one_chunk],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(stream.next().await.is_some());
assert!(stream.next().await.is_none());
assert_eq!(
warn_count.load(Ordering::SeqCst),
1,
"a genuinely malformed line must still be warned about"
);
}
#[tokio::test]
async fn splits_whitespace_only_line_across_reads_without_panicking() {
let script = Script {
chunks: vec![
b" ".to_vec(), b"\n".to_vec(),
valid_notification_line(),
],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(
stream.next().await.is_some(),
"the valid line must still decode after a blank line split across two reads"
);
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn warns_for_malformed_line_immediately_after_oversized_discard() {
let oversized_whitespace_run = vec![b' '; TEST_MAX * 3]; let mut rest = b"\nnot json at all\n".to_vec();
rest.extend(valid_notification_line());
let script = Script {
chunks: vec![oversized_whitespace_run, rest],
idx: 0,
};
let warn_count = Arc::new(AtomicUsize::new(0));
let _tracing_guard = tracing::subscriber::set_default(WarnCounter(Arc::clone(&warn_count)));
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(
stream.next().await.is_some(),
"the valid notification after the malformed line must still decode"
);
assert!(stream.next().await.is_none());
assert_eq!(
warn_count.load(Ordering::SeqCst),
2,
"one warning for the oversized run and one for the malformed line that follows \
it -- neither is blank, so neither may be suppressed"
);
}
#[tokio::test]
async fn ignored_notification_and_valid_request_in_one_chunk_do_not_stall() {
let mut one_chunk = br#"{"jsonrpc":"1.0","method":"foo"}"#.to_vec();
one_chunk.push(b'\n');
one_chunk.extend(valid_request_line(1));
let script = ScriptThenIdle {
chunks: vec![one_chunk],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
let message = next_or_timeout(&mut stream)
.await
.expect("the valid request sharing a chunk with an ignored non-standard line must decode without waiting for more input");
assert_eq!(request_id(&message), 1);
}
#[tokio::test]
async fn ends_cleanly_when_oversized_line_is_last_before_eof() {
let script = Script {
chunks: vec![oversized_line()],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn ends_cleanly_on_unterminated_oversized_line_then_eof() {
let script = Script {
chunks: vec![vec![b'y'; TEST_MAX * 3]], idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn ends_session_on_persistent_io_error_without_spinning() {
let error_polls = Arc::new(AtomicUsize::new(0));
let reader = ErrAfter {
chunks: vec![valid_notification_line()],
idx: 0,
error_polls: error_polls.clone(),
};
let mut stream =
bounded_request_stream::<ErrAfter>(reader, TEST_MAX, semaphore(TEST_CONCURRENCY));
assert!(
stream.next().await.is_some(),
"the valid line decodes before the reader starts failing"
);
assert!(
stream.next().await.is_none(),
"a persistent I/O error must end the session rather than recover"
);
assert_eq!(
error_polls.load(Ordering::SeqCst),
1,
"the I/O error must be surfaced on the first failing poll, not retried in a hot loop"
);
}
#[tokio::test]
async fn accepts_line_at_exact_cap_and_rejects_one_byte_over() {
let mut content = valid_notification_line();
assert_eq!(content.pop(), Some(b'\n'), "fixture must end in a newline");
let boundary_max = content.len();
let mut at_cap = content.clone();
at_cap.push(b'\n');
let mut stream = bounded_request_stream(
Script {
chunks: vec![at_cap],
idx: 0,
},
boundary_max,
semaphore(TEST_CONCURRENCY),
);
assert!(
stream.next().await.is_some(),
"a line whose content is exactly max_length bytes must be accepted"
);
assert!(stream.next().await.is_none());
let mut one_over = content;
one_over.push(b' ');
one_over.push(b'\n');
let mut stream = bounded_request_stream(
Script {
chunks: vec![one_over],
idx: 0,
},
boundary_max,
semaphore(TEST_CONCURRENCY),
);
assert!(
stream.next().await.is_none(),
"a line one byte over max_length must be dropped, not accepted"
);
}
#[tokio::test]
async fn permit_is_released_after_request_is_dropped() {
let limit = semaphore(1);
let script = Script {
chunks: vec![valid_request_line(1)],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, limit.clone());
let mut message = stream
.next()
.await
.expect("a valid request line must decode");
assert_eq!(
limit.available_permits(),
0,
"the single permit must be held while the request is in flight"
);
let permit = take_permit(&mut message);
drop(message);
assert_eq!(
limit.available_permits(),
0,
"the Arc-wrapped permit clone held by the test must still keep it reserved"
);
drop(permit);
assert_eq!(
limit.available_permits(),
1,
"dropping the last Arc<OwnedSemaphorePermit> must release the permit"
);
}
#[tokio::test]
async fn yields_pending_at_capacity_instead_of_dropping_or_erroring() {
let limit = semaphore(1);
let script = Script {
chunks: vec![valid_request_line(1), valid_request_line(2)],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, limit.clone());
let mut first = stream
.next()
.await
.expect("the first request line must decode and be admitted");
assert_eq!(limit.available_permits(), 0);
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
match Pin::new(&mut stream).poll_next(&mut cx) {
Poll::Pending => {}
Poll::Ready(item) => {
panic!(
"expected Poll::Pending while at capacity, got Some={}",
item.is_some()
)
}
}
let permit = take_permit(&mut first);
drop(first);
drop(permit);
let second = next_or_timeout(&mut stream)
.await
.expect("the second request must be admitted once a permit is available");
assert_eq!(
request_id(&second),
2,
"the admitted request must be the second one (id 2), not a reorder or duplicate of the first"
);
}
#[tokio::test]
async fn oversized_line_consumes_no_permit() {
let limit = semaphore(1);
let script = Script {
chunks: vec![oversized_line(), valid_request_line(1)],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, limit.clone());
let mut message = next_or_timeout(&mut stream)
.await
.expect("the valid request line after the oversized one must still decode");
assert_eq!(
limit.available_permits(),
0,
"only the valid request should have consumed the single permit"
);
let permit = take_permit(&mut message);
drop(message);
drop(permit);
assert_eq!(limit.available_permits(), 1);
}
#[tokio::test]
async fn notification_bypasses_a_request_still_awaiting_a_permit() {
let limit = semaphore(0);
let script = Script {
chunks: vec![valid_request_line(1), valid_notification_line()],
idx: 0,
};
let mut stream = bounded_request_stream(script, TEST_MAX, limit);
let notification = next_or_timeout(&mut stream)
.await
.expect("the notification behind the unadmitted request must still be decoded");
assert!(
matches!(notification, JsonRpcMessage::Notification(_)),
"expected a notification to bypass the request stuck waiting for a permit"
);
}
#[tokio::test]
async fn decode_ahead_queue_stalls_once_it_reaches_capacity() {
let limit = semaphore(0);
let mut chunks: Vec<Vec<u8>> = (1..=MAX_CONCURRENT_REQUESTS as u64)
.map(valid_request_line)
.collect();
chunks.push(valid_notification_line());
let script = Script { chunks, idx: 0 };
let mut stream = bounded_request_stream(script, TEST_MAX, limit);
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
match Pin::new(&mut stream).poll_next(&mut cx) {
Poll::Pending => {}
Poll::Ready(item) => panic!(
"expected Poll::Pending once the decode-ahead queue is full, got Some={}",
item.is_some()
),
}
}
#[tokio::test]
async fn permit_returns_to_capacity_after_real_rmcp_round_trip() {
const CAPACITY: usize = 2;
let limit = semaphore(CAPACITY);
let (client, server) = tokio::io::duplex(64 * 1024);
let (server_read, server_write) = tokio::io::split(server);
let (client_read, mut client_write) = tokio::io::split(client);
let stream = bounded_request_stream(server_read, MAX_REQUEST_LINE_SIZE, limit.clone());
let sink = FramedWrite::new(
server_write,
JsonRpcMessageCodec::<TxJsonRpcMessage<RoleServer>>::new(),
);
let service_task = tokio::spawn(async move {
let service = GeneratorService::new()
.serve((sink, stream))
.await
.expect("initialize handshake must succeed over the bounded stream");
service.waiting().await
});
let mut client_reader = BufReader::new(client_read);
client_write
.write_all(
br#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test-client","version":"0.0.0"}}}
"#,
)
.await
.expect("write initialize request");
let mut init_response = String::new();
tokio::time::timeout(
Duration::from_secs(5),
client_reader.read_line(&mut init_response),
)
.await
.expect("initialize response must arrive within 5s")
.expect("read initialize response");
assert!(
init_response.contains(r#""id":1"#),
"expected an initialize response, got: {init_response}"
);
client_write
.write_all(
br#"{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"list_generated_servers","arguments":{"base_dir":"nonexistent-mcp-permit-test-dir"}}}
"#,
)
.await
.expect("write tools/call request");
let mut tool_response = String::new();
tokio::time::timeout(
Duration::from_secs(5),
client_reader.read_line(&mut tool_response),
)
.await
.expect("tools/call response must arrive within 5s")
.expect("read tools/call response");
assert!(
tool_response.contains(r#""id":2"#),
"expected a tools/call response, got: {tool_response}"
);
assert!(
tool_response.contains(r#""result":"#),
"expected list_generated_servers to succeed, got: {tool_response}"
);
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if limit.available_permits() == CAPACITY {
return;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
})
.await
.unwrap_or_else(|_| {
panic!(
"semaphore did not return to full capacity ({CAPACITY}) within 2s; available: {}",
limit.available_permits()
)
});
drop(client_write);
drop(client_reader);
let _ = tokio::time::timeout(Duration::from_secs(5), service_task).await;
}
}