use std::collections::HashMap;
use std::fmt;
use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use serde_json::Value;
use tokio::io::{AsyncBufRead, AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStderr, ChildStdin, ChildStdout};
use tokio::sync::{mpsc, oneshot};
use tokio::time::timeout;
use vtcode_commons::sanitizer::{PROVIDER_DIAGNOSTIC_MAX_BYTES, sanitize_provider_diagnostic};
use super::error::{AcpError, AcpResult};
const WRITE_CHANNEL_CAPACITY: usize = 64;
const STDERR_LINE_MAX_BYTES: usize = PROVIDER_DIAGNOSTIC_MAX_BYTES;
const MAX_JSON_RPC_MESSAGE_BYTES: usize = 64 * 1024 * 1024;
type PendingRequestMap = HashMap<i64, oneshot::Sender<AcpResult<Value>>>;
type PendingRequestStore = Arc<StdMutex<PendingRequestMap>>;
type NotificationHandler = Arc<dyn Fn(Value) -> anyhow::Result<()> + Send + Sync>;
pub struct StdioTransport {
write_tx: mpsc::Sender<String>,
pending: PendingRequestStore,
request_counter: AtomicI64,
notification_handler: Arc<StdMutex<Option<NotificationHandler>>>,
child: StdMutex<Option<Child>>,
rpc_timeout: Duration,
}
impl StdioTransport {
pub fn from_child(
child: Child,
stdin: ChildStdin,
stdout: ChildStdout,
stderr: ChildStderr,
rpc_timeout: Duration,
) -> Self {
let (write_tx, write_rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
let pending = Arc::new(StdMutex::new(HashMap::new()));
let notification_handler = Arc::new(StdMutex::new(None));
spawn_writer(write_rx, stdin);
spawn_stderr_logger(stderr);
spawn_reader(stdout, Arc::clone(&pending), Arc::clone(¬ification_handler));
Self {
write_tx,
pending,
request_counter: AtomicI64::new(1),
notification_handler,
child: StdMutex::new(Some(child)),
rpc_timeout,
}
}
#[cfg(test)]
pub(crate) fn new_for_testing(write_tx: mpsc::Sender<String>, rpc_timeout: Duration) -> Self {
Self {
write_tx,
pending: Arc::new(StdMutex::new(HashMap::new())),
request_counter: AtomicI64::new(1),
notification_handler: Arc::new(StdMutex::new(None)),
child: StdMutex::new(None),
rpc_timeout,
}
}
pub fn set_notification_handler(&self, handler: NotificationHandler) {
if let Ok(mut guard) = self.notification_handler.lock() {
*guard = Some(handler);
}
}
pub async fn call(&self, method: &str, params: Value) -> AcpResult<Value> {
let id = self.request_counter.fetch_add(1, Ordering::Relaxed);
let (tx, rx) = oneshot::channel();
self.pending
.lock()
.map_err(|_e| AcpError::Internal("stdio transport pending mutex poisoned".into()))?
.insert(id, tx);
let _pending_guard = PendingRequestGuard::new(Arc::clone(&self.pending), id);
let payload = serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": method,
"params": params,
});
self.send_raw(payload)?;
timeout(self.rpc_timeout, rx)
.await
.map_err(|_e| AcpError::Timeout(format!("{method} timed out")))?
.map_err(|_e| AcpError::Internal(format!("{method} response channel closed")))
.and_then(|r| r)
}
pub fn notify(&self, method: &str, params: Value) -> AcpResult<()> {
let payload = serde_json::json!({
"jsonrpc": "2.0",
"method": method,
"params": params,
});
self.send_raw(payload)
}
pub fn respond(&self, id: i64, result: Value) -> AcpResult<()> {
let payload = serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"result": result,
});
self.send_raw(payload)
}
pub fn respond_error(&self, id: i64, code: i32, message: impl Into<String>) -> AcpResult<()> {
let payload = serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"error": {
"code": code,
"message": message.into(),
},
});
self.send_raw(payload)
}
fn send_raw(&self, payload: Value) -> AcpResult<()> {
let text = serde_json::to_string(&payload)?;
if text.len() > MAX_JSON_RPC_MESSAGE_BYTES {
return Err(AcpError::Internal(format!(
"stdio transport JSON-RPC frame exceeds {MAX_JSON_RPC_MESSAGE_BYTES} byte limit"
)));
}
self.write_tx.try_send(text).map_err(|e| match e {
mpsc::error::TrySendError::Full(_) => {
AcpError::Internal("stdio transport write channel full; subprocess may be slow".into())
}
mpsc::error::TrySendError::Closed(_) => AcpError::Internal("stdio transport writer channel closed".into()),
})
}
}
struct PendingRequestGuard {
pending: PendingRequestStore,
id: i64,
}
impl PendingRequestGuard {
fn new(pending: PendingRequestStore, id: i64) -> Self {
Self { pending, id }
}
}
impl Drop for PendingRequestGuard {
fn drop(&mut self) {
drop(self.pending.lock().unwrap_or_else(|error| error.into_inner()).remove(&self.id));
}
}
impl fmt::Debug for StdioTransport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StdioTransport")
.field("request_counter", &self.request_counter.load(Ordering::Relaxed))
.field("rpc_timeout", &self.rpc_timeout)
.finish_non_exhaustive()
}
}
impl Drop for StdioTransport {
fn drop(&mut self) {
if let Ok(mut child) = self.child.lock()
&& let Some(child) = child.as_mut()
{
let _ = child.start_kill();
}
}
}
fn spawn_writer(mut write_rx: mpsc::Receiver<String>, mut stdin: ChildStdin) {
tokio::spawn(async move {
while let Some(payload) = write_rx.recv().await {
if stdin.write_all(payload.as_bytes()).await.is_err()
|| stdin.write_all(b"\n").await.is_err()
|| stdin.flush().await.is_err()
{
tracing::warn!(
target: "vtcode.stdio_transport",
"stdin write failed; writer task exiting"
);
break;
}
}
});
}
fn spawn_stderr_logger(stderr: ChildStderr) {
tokio::spawn(async move {
let mut reader = BufReader::new(stderr);
let mut line = Vec::with_capacity(STDERR_LINE_MAX_BYTES);
loop {
match read_bounded_line(&mut reader, &mut line, STDERR_LINE_MAX_BYTES).await {
Ok(None) => break,
Ok(Some(truncated)) => {
let safe_line = sanitize_provider_diagnostic(trim_line_ending(&line));
tracing::debug!(
target: "vtcode.stdio_transport.stderr",
truncated,
"{}",
safe_line
)
}
Err(error) => {
tracing::warn!(
target: "vtcode.stdio_transport.stderr",
error = %error,
"stderr reader failed"
);
break;
}
}
}
});
}
fn spawn_reader(
stdout: ChildStdout,
pending: PendingRequestStore,
notification_handler: Arc<StdMutex<Option<NotificationHandler>>>,
) {
tokio::spawn(async move {
let mut reader = BufReader::new(stdout);
let mut line = Vec::with_capacity(256);
let close_reason = loop {
let truncated = match read_bounded_line(&mut reader, &mut line, MAX_JSON_RPC_MESSAGE_BYTES).await {
Ok(Some(truncated)) => truncated,
Ok(None) => break "stdout stream closed",
Err(error) => {
tracing::warn!("stdio transport reader failed: {error}");
break "stdout reader failed";
}
};
if truncated {
tracing::warn!(
target: "vtcode.stdio_transport",
"stdout JSON-RPC frame exceeded {MAX_JSON_RPC_MESSAGE_BYTES} byte limit; discarded"
);
continue;
}
if line.iter().all(u8::is_ascii_whitespace) {
continue;
}
let message: Value = match serde_json::from_slice(&line) {
Ok(value) => value,
Err(error) => {
tracing::warn!("stdio transport: JSON decode failed: {error}");
continue;
}
};
if let Some(id) = response_id(&message) {
let result = extract_rpc_result(&message);
let tx = pending.lock().unwrap_or_else(|e| e.into_inner()).remove(&id);
if let Some(tx) = tx {
let _ = tx.send(result);
}
continue;
}
if let Some(handler) = notification_handler.lock().unwrap_or_else(|e| e.into_inner()).as_ref().cloned()
&& let Err(e) = handler(message)
{
tracing::warn!("stdio transport: notification handler error: {e}");
}
};
fail_pending_requests(&pending, close_reason);
});
}
pub(super) async fn read_bounded_line<R: AsyncBufRead + Unpin>(
reader: &mut R,
line: &mut Vec<u8>,
max_bytes: usize,
) -> std::io::Result<Option<bool>> {
line.clear();
let mut truncated = false;
loop {
let available = reader.fill_buf().await?;
if available.is_empty() {
return Ok(if line.is_empty() && !truncated {
None
} else {
Some(truncated)
});
}
let newline = available.iter().position(|byte| *byte == b'\n');
let consumed = newline.map_or(available.len(), |position| position + 1);
if line.len() < max_bytes {
let copy_len = (max_bytes - line.len()).min(consumed);
line.extend_from_slice(&available[..copy_len]);
truncated |= copy_len < consumed;
} else {
truncated = true;
}
reader.consume(consumed);
if newline.is_some() {
return Ok(Some(truncated));
}
}
}
pub(super) fn trim_line_ending(mut line: &[u8]) -> &[u8] {
if line.last() == Some(&b'\n') {
line = &line[..line.len().saturating_sub(1)];
}
if line.last() == Some(&b'\r') {
line = &line[..line.len().saturating_sub(1)];
}
line
}
fn fail_pending_requests(pending: &PendingRequestStore, reason: &str) {
let senders: Vec<_> = pending
.lock()
.unwrap_or_else(|error| error.into_inner())
.drain()
.map(|(_, sender)| sender)
.collect();
for sender in senders {
drop(sender.send(Err(AcpError::Internal(reason.to_string()))));
}
}
fn response_id(message: &Value) -> Option<i64> {
if message.get("result").is_some() || message.get("error").is_some() {
message.get("id").and_then(Value::as_i64)
} else {
None
}
}
fn extract_rpc_result(message: &Value) -> AcpResult<Value> {
if let Some(error) = message.get("error") {
let code = error.get("code").and_then(Value::as_i64).unwrap_or_default();
let detail = error.get("message").and_then(Value::as_str).unwrap_or("unknown error");
Err(AcpError::RemoteError {
agent_id: "stdio".into(),
message: format!("rpc error {code}: {detail}"),
code: Some(code as i32),
})
} else {
Ok(message.get("result").cloned().unwrap_or(Value::Null))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn response_id_requires_result_or_error() {
assert!(
response_id(&serde_json::json!({
"jsonrpc": "2.0",
"method": "some/notification",
"params": {}
}))
.is_none()
);
assert!(
response_id(&serde_json::json!({
"jsonrpc": "2.0",
"id": 7,
"method": "permission.request",
"params": {}
}))
.is_none()
);
assert_eq!(
response_id(&serde_json::json!({
"jsonrpc": "2.0",
"id": 3,
"result": { "ok": true }
})),
Some(3)
);
assert_eq!(
response_id(&serde_json::json!({
"jsonrpc": "2.0",
"id": 5,
"error": { "code": -32601, "message": "method not found" }
})),
Some(5)
);
}
#[test]
fn extract_rpc_result_propagates_error() {
let result = extract_rpc_result(&serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"error": { "code": -32600, "message": "invalid request" }
}));
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("invalid request"));
}
#[test]
fn extract_rpc_result_returns_result_value() {
let result = extract_rpc_result(&serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"result": { "sessionId": "abc" }
}))
.unwrap();
assert_eq!(result["sessionId"], "abc");
}
#[test]
fn notify_serialises_payload_to_write_channel() {
let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
transport
.notify("session/cancel", serde_json::json!({ "sessionId": "s1" }))
.unwrap();
let raw = rx.try_recv().expect("notification payload");
let payload: Value = serde_json::from_str(&raw).unwrap();
assert_eq!(payload["method"], "session/cancel");
assert_eq!(payload["params"]["sessionId"], "s1");
assert!(payload.get("id").is_none(), "notifications must not have id");
}
#[test]
fn respond_writes_jsonrpc_result() {
let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
transport.respond(42, serde_json::json!({ "ok": true })).unwrap();
let raw = rx.try_recv().unwrap();
let payload: Value = serde_json::from_str(&raw).unwrap();
assert_eq!(payload["jsonrpc"], "2.0");
assert_eq!(payload["id"], 42);
assert_eq!(payload["result"]["ok"], true);
}
#[test]
fn respond_error_writes_jsonrpc_error() {
let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
transport.respond_error(9, -32601, "method not found").unwrap();
let raw = rx.try_recv().unwrap();
let payload: Value = serde_json::from_str(&raw).unwrap();
assert_eq!(payload["id"], 9);
assert_eq!(payload["error"]["code"], -32601);
assert_eq!(payload["error"]["message"], "method not found");
}
#[tokio::test]
async fn timed_out_call_clears_pending_entry() {
let (tx, _rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
let transport = StdioTransport::new_for_testing(tx, Duration::from_millis(20));
let result = transport.call("session/start", serde_json::json!({})).await;
assert!(matches!(result, Err(AcpError::Timeout(_))));
let pending_len = transport.pending.lock().unwrap().len();
assert_eq!(pending_len, 0, "timed-out call must not leave a pending entry");
}
#[tokio::test]
async fn bounded_stderr_reader_drains_oversized_lines() -> std::io::Result<()> {
let mut data = vec![b'x'; STDERR_LINE_MAX_BYTES + 32];
data.extend_from_slice(b"\nnext\n");
let mut reader = BufReader::new(data.as_slice());
let mut line = Vec::new();
assert_eq!(read_bounded_line(&mut reader, &mut line, STDERR_LINE_MAX_BYTES).await?, Some(true));
assert_eq!(line.len(), STDERR_LINE_MAX_BYTES);
assert_eq!(read_bounded_line(&mut reader, &mut line, STDERR_LINE_MAX_BYTES).await?, Some(false));
assert_eq!(trim_line_ending(&line), b"next");
Ok(())
}
#[tokio::test]
async fn eof_fails_all_pending_calls() {
let pending = Arc::new(StdMutex::new(HashMap::new()));
let (sender, receiver) = oneshot::channel();
pending.lock().unwrap().insert(7, sender);
fail_pending_requests(&pending, "stdout stream closed");
let result = receiver.await.expect("pending call should be notified");
assert!(matches!(result, Err(AcpError::Internal(message)) if message == "stdout stream closed"));
assert!(pending.lock().unwrap().is_empty());
}
}