use std::sync::Arc;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use super::{McpServer, McpServerConfig, ServerError, ServerStatus};
use crate::protocol::JsonRpcMessage;
pub struct McpStdioServer {
server: Arc<McpServer>,
}
impl McpStdioServer {
pub async fn run(config: McpServerConfig) -> Result<(), ServerError> {
let (server, mut channels) = McpServer::new(config);
let stdio_server = Self {
server: Arc::clone(&server),
};
let stdout_handle = tokio::spawn(async move {
let mut stdout = tokio::io::stdout();
while let Some(outbound) = channels.outbound_rx.recv().await {
let json = match outbound.to_json() {
Ok(j) => j,
Err(e) => {
eprintln!("Failed to serialize outbound message: {}", e);
continue;
}
};
if let Err(e) = stdout.write_all(json.as_bytes()).await {
eprintln!("Failed to write to stdout: {}", e);
break;
}
if let Err(e) = stdout.write_all(b"\n").await {
eprintln!("Failed to write newline to stdout: {}", e);
break;
}
if let Err(e) = stdout.flush().await {
eprintln!("Failed to flush stdout: {}", e);
break;
}
}
});
let stdin = tokio::io::stdin();
let mut reader = BufReader::new(stdin);
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => {
break;
}
Ok(_) => {
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
match JsonRpcMessage::parse(trimmed) {
Ok(message) => {
let inbound = message.into_client_inbound();
if channels.inbound_tx.send(inbound).await.is_err() {
break;
}
}
Err(e) => {
let error_response = crate::protocol::JsonRpcResponse::error(
crate::protocol::JsonRpcId::Null,
-32700,
format!("Parse error: {}", e),
None,
);
let outbound =
crate::protocol::ServerOutbound::Response(error_response);
if channels.outbound_tx.send(outbound).await.is_err() {
break;
}
}
}
}
Err(e) => {
return Err(ServerError::Io(e));
}
}
if stdio_server.server.status() != ServerStatus::Running {
break;
}
}
server.stop();
let _ = stdout_handle.await;
Ok(())
}
pub fn server(&self) -> &Arc<McpServer> {
&self.server
}
}
#[cfg(test)]
mod tests {
use crate::protocol::{JsonRpcId, ServerOutbound};
use tokio::sync::mpsc;
#[test]
fn test_stdio_server_module_exists() {
}
#[tokio::test]
async fn test_outbound_message_synchronization() {
let (outbound_tx, mut outbound_rx) = mpsc::channel::<ServerOutbound>(256);
let tx1 = outbound_tx.clone();
let tx2 = outbound_tx.clone();
let tx3 = outbound_tx.clone();
let handles = vec![
tokio::spawn(async move {
for i in 0..10 {
let response = crate::protocol::JsonRpcResponse::success(
JsonRpcId::Number(i),
serde_json::json!({"msg": format!("response_{}", i)}),
);
tx1.send(ServerOutbound::Response(response)).await.unwrap();
}
}),
tokio::spawn(async move {
for i in 10..20 {
let response = crate::protocol::JsonRpcResponse::error(
JsonRpcId::Number(i),
-32700,
format!("Parse error {}", i),
None,
);
tx2.send(ServerOutbound::Response(response)).await.unwrap();
}
}),
tokio::spawn(async move {
for i in 20..30 {
let notification =
crate::protocol::JsonRpcNotification::new(format!("notify_{}", i), None);
tx3.send(ServerOutbound::Notification(notification))
.await
.unwrap();
}
}),
];
for handle in handles {
handle.await.unwrap();
}
drop(outbound_tx);
let mut messages = Vec::new();
while let Some(msg) = outbound_rx.recv().await {
let json = msg.to_json().unwrap();
let parsed: serde_json::Value = serde_json::from_str(&json)
.expect("Each message should be valid JSON - no interleaving");
messages.push(parsed);
}
assert_eq!(messages.len(), 30, "All messages should be received");
for msg in &messages {
assert!(
msg.get("jsonrpc").is_some(),
"Each message should have jsonrpc field"
);
}
}
#[tokio::test]
async fn test_single_writer_pattern() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
let (outbound_tx, mut outbound_rx) = mpsc::channel::<ServerOutbound>(256);
let write_count = Arc::new(AtomicUsize::new(0));
let concurrent_writes = Arc::new(AtomicUsize::new(0));
let max_concurrent = Arc::new(AtomicUsize::new(0));
let write_count_clone = Arc::clone(&write_count);
let concurrent_clone = Arc::clone(&concurrent_writes);
let max_clone = Arc::clone(&max_concurrent);
let writer_handle = tokio::spawn(async move {
while let Some(outbound) = outbound_rx.recv().await {
let current = concurrent_clone.fetch_add(1, Ordering::SeqCst) + 1;
let mut max = max_clone.load(Ordering::SeqCst);
while current > max {
match max_clone.compare_exchange(
max,
current,
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => break,
Err(m) => max = m,
}
}
let _json = outbound.to_json().unwrap();
tokio::task::yield_now().await;
write_count_clone.fetch_add(1, Ordering::SeqCst);
concurrent_clone.fetch_sub(1, Ordering::SeqCst);
}
});
let mut send_handles = Vec::new();
for batch in 0..5 {
let tx = outbound_tx.clone();
send_handles.push(tokio::spawn(async move {
for i in 0..10 {
let response = crate::protocol::JsonRpcResponse::success(
JsonRpcId::Number(batch * 10 + i),
serde_json::json!({}),
);
tx.send(ServerOutbound::Response(response)).await.unwrap();
}
}));
}
for handle in send_handles {
handle.await.unwrap();
}
drop(outbound_tx);
writer_handle.await.unwrap();
assert_eq!(write_count.load(Ordering::SeqCst), 50);
assert!(
max_concurrent.load(Ordering::SeqCst) <= 1,
"Single writer should never have more than 1 concurrent write"
);
}
}