use async_trait::async_trait;
use futures::Stream;
use futures::StreamExt;
use std::{
io::{Read, Write},
pin::Pin,
sync::{Arc, Mutex},
time::Duration,
};
use tokio::{
io::{AsyncBufReadExt, BufReader as TokioBufReader},
sync::broadcast,
time::timeout,
};
use crate::{
error::{Error, ErrorCode},
transport::{Message, Transport},
};
pub struct StdioTransport {
stdout: Arc<Mutex<std::io::Stdout>>,
receiver: broadcast::Receiver<Result<Message, Error>>,
}
impl StdioTransport {
pub fn new() -> (Self, broadcast::Sender<Result<Message, Error>>) {
let (sender, receiver) = broadcast::channel(100);
let transport = Self {
stdout: Arc::new(Mutex::new(std::io::stdout())),
receiver,
};
let stdin = tokio::io::stdin();
let mut reader = TokioBufReader::new(stdin);
let sender_clone = sender.clone();
tokio::spawn(async move {
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => break, Ok(_) => {
let message = match serde_json::from_str(&line) {
Ok(message) => Ok(message),
Err(err) => Err(Error::Serialization(err.to_string())),
};
if sender_clone.send(message).is_err() {
break;
}
}
Err(err) => {
let _ = sender_clone.send(Err(Error::Io(err.to_string())));
break;
}
}
}
});
(transport, sender)
}
}
#[async_trait]
impl Transport for StdioTransport {
async fn send(&self, message: Message) -> Result<(), Error> {
let mut stdout = self.stdout.lock().map_err(|_e| {
Error::protocol(ErrorCode::InternalError, "Failed to acquire stdout lock")
})?;
let json = serde_json::to_string(&message)?;
writeln!(stdout, "{}", json).map_err(|e| Error::Io(e.to_string()))?;
stdout.flush().map_err(|e| Error::Io(e.to_string()))?;
Ok(())
}
fn receive(&self) -> Pin<Box<dyn Stream<Item = Result<Message, Error>> + Send>> {
let rx = self.receiver.resubscribe();
Box::pin(futures::stream::unfold(rx, |mut rx| async move {
match rx.recv().await {
Ok(msg) => Some((msg, rx)),
Err(_) => None,
}
}))
}
async fn close(&self) -> Result<(), Error> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::{Request, RequestId};
use std::sync::mpsc;
use std::thread;
use std::time::Duration;
use tokio::sync::broadcast;
use tokio::runtime::Runtime;
#[test]
fn test_stdio_transport() {
let (verify_tx, verify_rx) = mpsc::channel();
thread::spawn(move || {
let rt = Runtime::new().unwrap();
let request = Request::new(
"test_method",
Some(serde_json::json!({"key": "value"})),
RequestId::Number(1),
);
let message = Message::Request(request);
let (_, receiver) = broadcast::channel(100);
let transport = StdioTransport {
stdout: Arc::new(Mutex::new(std::io::stdout())),
receiver,
};
let send_result = rt.block_on(transport.send(message.clone()));
verify_tx.send(send_result.is_ok()).unwrap();
});
match verify_rx.recv_timeout(Duration::from_secs(1)) {
Ok(true) => (), Ok(false) => panic!("Failed to send message"),
Err(_) => panic!("Test timed out"),
}
}
}