1use std::fmt;
28use std::sync::atomic::{AtomicI64, Ordering};
29use std::sync::{Arc, Mutex as StdMutex};
30
31use hashbrown::HashMap;
32use std::time::Duration;
33
34use serde_json::Value;
35use tokio::io::{AsyncBufRead, AsyncWriteExt, BufReader};
36use tokio::process::{Child, ChildStderr, ChildStdin, ChildStdout};
37use tokio::sync::{mpsc, oneshot};
38use tokio::time::timeout;
39
40use crate::error::{AcpError, AcpResult};
41use vtcode_commons::sanitizer::{PROVIDER_DIAGNOSTIC_MAX_BYTES, sanitize_provider_diagnostic};
42
43const WRITE_CHANNEL_CAPACITY: usize = 64;
46
47const MAX_JSON_RPC_MESSAGE_BYTES: usize = 64 * 1024 * 1024;
51
52type PendingRequestMap = HashMap<String, oneshot::Sender<AcpResult<Value>>>;
53type PendingRequestStore = Arc<StdMutex<PendingRequestMap>>;
54
55type NotificationHandler = Arc<dyn Fn(Value) -> anyhow::Result<()> + Send + Sync>;
60
61#[derive(Debug, Clone, Copy)]
62pub struct StdioTransportOptions {
63 pub include_jsonrpc_version: bool,
64}
65
66impl Default for StdioTransportOptions {
67 fn default() -> Self {
68 Self { include_jsonrpc_version: true }
69 }
70}
71
72pub struct StdioTransport {
84 write_tx: mpsc::Sender<String>,
85 pending: PendingRequestStore,
86 request_counter: AtomicI64,
87 notification_handler: Arc<StdMutex<Option<NotificationHandler>>>,
88 child: StdMutex<Option<Child>>,
89 rpc_timeout: Duration,
90 options: StdioTransportOptions,
91}
92
93impl StdioTransport {
94 pub fn from_child(
99 child: Child,
100 stdin: ChildStdin,
101 stdout: ChildStdout,
102 stderr: ChildStderr,
103 rpc_timeout: Duration,
104 ) -> Self {
105 Self::from_child_with_options(child, stdin, stdout, stderr, rpc_timeout, StdioTransportOptions::default())
106 }
107
108 pub fn from_child_with_options(
109 child: Child,
110 stdin: ChildStdin,
111 stdout: ChildStdout,
112 stderr: ChildStderr,
113 rpc_timeout: Duration,
114 options: StdioTransportOptions,
115 ) -> Self {
116 let (write_tx, write_rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
117 let pending = Arc::new(StdMutex::new(HashMap::new()));
118 let notification_handler = Arc::new(StdMutex::new(None));
119
120 spawn_writer(write_rx, stdin);
121 spawn_stderr_logger(stderr);
122 spawn_reader(stdout, Arc::clone(&pending), Arc::clone(¬ification_handler));
123
124 Self {
125 write_tx,
126 pending,
127 request_counter: AtomicI64::new(1),
128 notification_handler,
129 child: StdMutex::new(Some(child)),
130 rpc_timeout,
131 options,
132 }
133 }
134
135 #[cfg(test)]
140 fn new_for_testing(write_tx: mpsc::Sender<String>, rpc_timeout: Duration) -> Self {
141 Self::new_for_testing_with_options(write_tx, rpc_timeout, StdioTransportOptions::default())
142 }
143
144 #[cfg(test)]
145 fn new_for_testing_with_options(
146 write_tx: mpsc::Sender<String>,
147 rpc_timeout: Duration,
148 options: StdioTransportOptions,
149 ) -> Self {
150 Self {
151 write_tx,
152 pending: Arc::new(StdMutex::new(HashMap::new())),
153 request_counter: AtomicI64::new(1),
154 notification_handler: Arc::new(StdMutex::new(None)),
155 child: StdMutex::new(None),
156 rpc_timeout,
157 options,
158 }
159 }
160
161 pub fn set_notification_handler(&self, handler: NotificationHandler) {
167 if let Ok(mut guard) = self.notification_handler.lock() {
168 *guard = Some(handler);
169 }
170 }
171
172 pub async fn call(&self, method: &str, params: Value) -> AcpResult<Value> {
182 let id = self.request_counter.fetch_add(1, Ordering::Relaxed);
183 let id_value = Value::from(id);
184 let pending_key = response_id_key(&id_value);
185 let (tx, rx) = oneshot::channel();
186 drop(
187 self.pending
188 .lock()
189 .map_err(|_err| AcpError::Internal("stdio transport pending mutex poisoned".into()))?
190 .insert(pending_key.clone(), tx),
191 );
192 let _pending_guard = PendingRequestGuard::new(Arc::clone(&self.pending), pending_key.clone());
193
194 let mut payload = serde_json::json!({
195 "jsonrpc": "2.0",
196 "id": id,
197 "method": method,
198 "params": params,
199 });
200 maybe_strip_jsonrpc_field(&mut payload, self.options);
201 self.send_raw(payload)?;
202
203 timeout(self.rpc_timeout, rx)
204 .await
205 .map_err(|_err| AcpError::Timeout(format!("{method} timed out")))?
206 .map_err(|_err| AcpError::Internal(format!("{method} response channel closed")))
207 .and_then(|r| r)
208 }
209
210 pub fn notify(&self, method: &str, params: Value) -> AcpResult<()> {
216 let mut payload = serde_json::json!({
217 "jsonrpc": "2.0",
218 "method": method,
219 "params": params,
220 });
221 maybe_strip_jsonrpc_field(&mut payload, self.options);
222 self.send_raw(payload)
223 }
224
225 fn respond(&self, id: i64, result: Value) -> AcpResult<()> {
234 self.respond_value(Value::from(id), result)
235 }
236
237 pub fn respond_value(&self, id: Value, result: Value) -> AcpResult<()> {
238 let mut payload = serde_json::json!({
239 "jsonrpc": "2.0",
240 "id": id,
241 "result": result,
242 });
243 maybe_strip_jsonrpc_field(&mut payload, self.options);
244 self.send_raw(payload)
245 }
246
247 fn respond_error(&self, id: i64, code: i32, message: impl Into<String>) -> AcpResult<()> {
253 self.respond_error_value(Value::from(id), code, message)
254 }
255
256 pub fn respond_error_value(&self, id: Value, code: i32, message: impl Into<String>) -> AcpResult<()> {
257 let mut payload = serde_json::json!({
258 "jsonrpc": "2.0",
259 "id": id,
260 "error": {
261 "code": code,
262 "message": message.into(),
263 },
264 });
265 maybe_strip_jsonrpc_field(&mut payload, self.options);
266 self.send_raw(payload)
267 }
268
269 fn send_raw(&self, payload: Value) -> AcpResult<()> {
270 let text = serde_json::to_string(&payload)?;
271 if text.len() > MAX_JSON_RPC_MESSAGE_BYTES {
272 return Err(AcpError::Internal(format!(
273 "stdio transport JSON-RPC frame exceeds {MAX_JSON_RPC_MESSAGE_BYTES} byte limit"
274 )));
275 }
276 self.write_tx.try_send(text).map_err(|e| match e {
277 mpsc::error::TrySendError::Full(_) => {
278 AcpError::Internal("stdio transport write channel full; subprocess may be slow".into())
279 }
280 mpsc::error::TrySendError::Closed(_) => AcpError::Internal("stdio transport writer channel closed".into()),
281 })
282 }
283}
284
285struct PendingRequestGuard {
286 pending: PendingRequestStore,
287 key: String,
288}
289
290impl PendingRequestGuard {
291 fn new(pending: PendingRequestStore, key: String) -> Self {
292 Self { pending, key }
293 }
294}
295
296impl Drop for PendingRequestGuard {
297 fn drop(&mut self) {
298 drop(self.pending.lock().unwrap_or_else(|error| error.into_inner()).remove(&self.key));
299 }
300}
301
302impl fmt::Debug for StdioTransport {
303 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
304 f.debug_struct("StdioTransport")
305 .field("request_counter", &self.request_counter.load(Ordering::Relaxed))
306 .field("rpc_timeout", &self.rpc_timeout)
307 .finish_non_exhaustive()
308 }
309}
310
311impl Drop for StdioTransport {
312 fn drop(&mut self) {
313 if let Ok(mut child) = self.child.lock()
314 && let Some(child) = child.as_mut()
315 {
316 drop(child.start_kill());
317 }
318 }
319}
320
321fn spawn_writer(mut write_rx: mpsc::Receiver<String>, mut stdin: ChildStdin) {
326 drop(tokio::spawn(async move {
327 while let Some(payload) = write_rx.recv().await {
328 if stdin.write_all(payload.as_bytes()).await.is_err()
329 || stdin.write_all(b"\n").await.is_err()
330 || stdin.flush().await.is_err()
331 {
332 tracing::warn!(
333 target: "vtcode.stdio_transport",
334 "stdin write failed; writer task exiting"
335 );
336 break;
337 }
338 }
339 }));
340}
341
342fn spawn_stderr_logger(stderr: ChildStderr) {
343 drop(tokio::spawn(async move {
344 let mut reader = BufReader::new(stderr);
345 loop {
346 match read_bounded_line(&mut reader, PROVIDER_DIAGNOSTIC_MAX_BYTES).await {
347 Ok(Some(line)) => {
348 let diagnostic = sanitize_provider_diagnostic(&line.bytes);
349 tracing::debug!(
350 target: "vtcode.stdio_transport.stderr",
351 truncated = line.truncated,
352 "{diagnostic}"
353 );
354 }
355 Ok(None) => break,
356 Err(error) => {
357 tracing::warn!(
358 target: "vtcode.stdio_transport.stderr",
359 error = %error,
360 "stderr reader failed"
361 );
362 break;
363 }
364 }
365 }
366 }));
367}
368
369async fn read_bounded_line<R: AsyncBufRead + Unpin>(
370 reader: &mut R,
371 max_bytes: usize,
372) -> std::io::Result<Option<BoundedLine>> {
373 let mut line = Vec::with_capacity(max_bytes.min(256));
374 let truncated = vtcode_commons::line_framing::read_bounded_line(
375 reader,
376 &mut line,
377 max_bytes,
378 vtcode_commons::line_framing::LineEnding::ExcludeLf,
379 )
380 .await?;
381 Ok(truncated.map(|truncated| BoundedLine { bytes: line, truncated }))
382}
383
384#[derive(Debug)]
385struct BoundedLine {
386 bytes: Vec<u8>,
387 truncated: bool,
388}
389
390fn fail_pending_calls(pending: &PendingRequestStore, message: &str) {
391 let pending_calls = pending
392 .lock()
393 .unwrap_or_else(|error| error.into_inner())
394 .drain()
395 .map(|(_, sender)| sender)
396 .collect::<Vec<_>>();
397 for sender in pending_calls {
398 drop(sender.send(Err(AcpError::Internal(message.to_string()))));
399 }
400}
401
402fn spawn_reader(
403 stdout: ChildStdout,
404 pending: PendingRequestStore,
405 notification_handler: Arc<StdMutex<Option<NotificationHandler>>>,
406) {
407 drop(tokio::spawn(async move {
408 let mut reader = BufReader::new(stdout);
409 loop {
410 let line = match read_bounded_line(&mut reader, MAX_JSON_RPC_MESSAGE_BYTES).await {
411 Ok(Some(line)) => line,
412 Ok(None) => {
413 fail_pending_calls(&pending, "stdio transport stdout reached EOF");
414 break;
415 }
416 Err(error) => {
417 tracing::warn!("stdio transport: stdout read failed: {error}");
418 fail_pending_calls(&pending, "stdio transport stdout read failed");
419 break;
420 }
421 };
422 if line.truncated {
423 tracing::warn!(
424 target: "vtcode.stdio_transport",
425 "stdout JSON-RPC frame exceeded {MAX_JSON_RPC_MESSAGE_BYTES} byte limit; discarded"
426 );
427 continue;
428 }
429 if line.bytes.iter().all(u8::is_ascii_whitespace) {
430 continue;
431 }
432 let message: Value = match serde_json::from_slice(&line.bytes) {
433 Ok(v) => v,
434 Err(e) => {
435 tracing::warn!("stdio transport: JSON decode failed: {e}");
436 continue;
437 }
438 };
439
440 if let Some(id) = response_id(&message) {
443 let result = extract_rpc_result(&message);
444 let tx = pending.lock().unwrap_or_else(|e| e.into_inner()).remove(&response_id_key(&id));
445 if let Some(tx) = tx {
446 drop(tx.send(result));
447 }
448 continue;
449 }
450
451 if let Some(handler) = notification_handler.lock().unwrap_or_else(|e| e.into_inner()).as_ref().cloned()
454 && let Err(e) = handler(message)
455 {
456 tracing::warn!("stdio transport: notification handler error: {e}");
457 }
458 }
459 }));
460}
461
462fn response_id(message: &Value) -> Option<Value> {
468 if message.get("result").is_some() || message.get("error").is_some() {
469 message.get("id").cloned()
470 } else {
471 None
472 }
473}
474
475fn response_id_key(id: &Value) -> String {
476 serde_json::to_string(id).unwrap_or_else(|_| "null".to_string())
477}
478
479fn maybe_strip_jsonrpc_field(payload: &mut Value, options: StdioTransportOptions) {
480 if options.include_jsonrpc_version {
481 return;
482 }
483
484 if let Some(object) = payload.as_object_mut() {
485 drop(object.remove("jsonrpc"));
486 }
487}
488
489fn extract_rpc_result(message: &Value) -> AcpResult<Value> {
490 if let Some(error) = message.get("error") {
491 let code = error.get("code").and_then(Value::as_i64).unwrap_or_default();
492 let detail = error.get("message").and_then(Value::as_str).unwrap_or("unknown error");
493 Err(AcpError::RemoteError {
494 agent_id: "stdio".into(),
495 message: format!("rpc error {code}: {detail}"),
496 code: Some(i32::try_from(code).unwrap_or(if code < 0 { i32::MIN } else { i32::MAX })),
497 })
498 } else {
499 Ok(message.get("result").cloned().unwrap_or(Value::Null))
500 }
501}
502
503#[cfg(test)]
504mod tests {
505 use super::*;
506
507 #[test]
508 fn response_id_requires_result_or_error() {
509 assert!(
511 response_id(&serde_json::json!({
512 "jsonrpc": "2.0",
513 "method": "some/notification",
514 "params": {}
515 }))
516 .is_none()
517 );
518
519 assert!(
521 response_id(&serde_json::json!({
522 "jsonrpc": "2.0",
523 "id": 7,
524 "method": "permission.request",
525 "params": {}
526 }))
527 .is_none()
528 );
529
530 assert_eq!(
532 response_id(&serde_json::json!({
533 "jsonrpc": "2.0",
534 "id": 3,
535 "result": { "ok": true }
536 })),
537 Some(Value::from(3))
538 );
539
540 assert_eq!(
542 response_id(&serde_json::json!({
543 "jsonrpc": "2.0",
544 "id": 5,
545 "error": { "code": -32601, "message": "method not found" }
546 })),
547 Some(Value::from(5))
548 );
549 }
550
551 #[test]
552 fn extract_rpc_result_propagates_error() {
553 let result = extract_rpc_result(&serde_json::json!({
554 "jsonrpc": "2.0",
555 "id": 1,
556 "error": { "code": -32600, "message": "invalid request" }
557 }));
558 assert!(result.is_err());
559 let err = result.unwrap_err().to_string();
560 assert!(err.contains("invalid request"));
561 }
562
563 #[test]
564 fn extract_rpc_result_returns_result_value() {
565 let result = extract_rpc_result(&serde_json::json!({
566 "jsonrpc": "2.0",
567 "id": 1,
568 "result": { "sessionId": "abc" }
569 }))
570 .unwrap();
571 assert_eq!(result["sessionId"], "abc");
572 }
573
574 #[test]
575 fn notify_serialises_payload_to_write_channel() {
576 let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
577 let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
578
579 transport
580 .notify("session/cancel", serde_json::json!({ "sessionId": "s1" }))
581 .unwrap();
582
583 let raw = rx.try_recv().expect("notification payload");
584 let payload: Value = serde_json::from_str(&raw).unwrap();
585 assert_eq!(payload["method"], "session/cancel");
586 assert_eq!(payload["params"]["sessionId"], "s1");
587 assert!(payload.get("id").is_none(), "notifications must not have id");
588 }
589
590 #[test]
591 fn respond_writes_jsonrpc_result() {
592 let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
593 let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
594
595 transport.respond(42, serde_json::json!({ "ok": true })).unwrap();
596
597 let raw = rx.try_recv().unwrap();
598 let payload: Value = serde_json::from_str(&raw).unwrap();
599 assert_eq!(payload["jsonrpc"], "2.0");
600 assert_eq!(payload["id"], 42);
601 assert_eq!(payload["result"]["ok"], true);
602 }
603
604 #[test]
605 fn respond_error_writes_jsonrpc_error() {
606 let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
607 let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
608
609 transport.respond_error(9, -32601, "method not found").unwrap();
610
611 let raw = rx.try_recv().unwrap();
612 let payload: Value = serde_json::from_str(&raw).unwrap();
613 assert_eq!(payload["id"], 9);
614 assert_eq!(payload["error"]["code"], -32601);
615 assert_eq!(payload["error"]["message"], "method not found");
616 }
617
618 #[test]
619 fn respond_value_supports_string_ids() {
620 let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
621 let transport = StdioTransport::new_for_testing(tx, Duration::from_secs(5));
622
623 transport
624 .respond_value(Value::String("request-1".to_string()), serde_json::json!({ "ok": true }))
625 .unwrap();
626
627 let raw = rx.try_recv().unwrap();
628 let payload: Value = serde_json::from_str(&raw).unwrap();
629 assert_eq!(payload["id"], "request-1");
630 assert_eq!(payload["result"]["ok"], true);
631 }
632
633 #[test]
634 fn can_omit_jsonrpc_field_for_codex_mode() {
635 let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
636 let transport = StdioTransport::new_for_testing_with_options(
637 tx,
638 Duration::from_secs(5),
639 StdioTransportOptions { include_jsonrpc_version: false },
640 );
641
642 transport.notify("initialized", serde_json::json!({})).unwrap();
643
644 let raw = rx.try_recv().unwrap();
645 let payload: Value = serde_json::from_str(&raw).unwrap();
646 assert!(payload.get("jsonrpc").is_none());
647 assert_eq!(payload["method"], "initialized");
648 }
649
650 #[tokio::test]
651 async fn call_timeout_removes_pending_request() {
652 let (tx, mut rx) = mpsc::channel(WRITE_CHANNEL_CAPACITY);
653 let transport = StdioTransport::new_for_testing(tx, Duration::from_millis(1));
654
655 let result = transport.call("slow", Value::Null).await;
656
657 assert!(matches!(result, Err(AcpError::Timeout(_))));
658 drop(rx.recv().await);
659 assert!(transport.pending.lock().unwrap_or_else(|error| error.into_inner()).is_empty());
660 }
661
662 #[tokio::test]
663 async fn bounded_line_cap_excludes_lf_but_retains_cr() -> std::io::Result<()> {
664 let mut reader = BufReader::with_capacity(1, b"abc\nxy\r\n".as_slice());
665 let first = read_bounded_line(&mut reader, 3).await?.expect("first line");
666 assert_eq!(first.bytes, b"abc");
667 assert!(!first.truncated);
668 let second = read_bounded_line(&mut reader, 3).await?.expect("second line");
669 assert_eq!(second.bytes, b"xy\r");
670 assert!(!second.truncated);
671 Ok(())
672 }
673
674 #[tokio::test]
675 async fn bounded_line_reader_drains_oversized_frames() -> std::io::Result<()> {
676 let mut data = vec![b'x'; PROVIDER_DIAGNOSTIC_MAX_BYTES + 32];
677 data.extend_from_slice(b"\nnext\n");
678 let mut reader = BufReader::new(data.as_slice());
679
680 let first = read_bounded_line(&mut reader, PROVIDER_DIAGNOSTIC_MAX_BYTES)
681 .await?
682 .expect("first line");
683 assert!(first.truncated);
684 assert_eq!(first.bytes.len(), PROVIDER_DIAGNOSTIC_MAX_BYTES);
685
686 let second = read_bounded_line(&mut reader, PROVIDER_DIAGNOSTIC_MAX_BYTES)
687 .await?
688 .expect("second line");
689 assert!(!second.truncated);
690 assert_eq!(second.bytes, b"next");
691 Ok(())
692 }
693
694 #[tokio::test]
695 async fn fail_pending_calls_wakes_all_waiters_and_clears_map() {
696 let pending = Arc::new(StdMutex::new(HashMap::new()));
697 let (first_tx, first_rx) = oneshot::channel();
698 let (second_tx, second_rx) = oneshot::channel();
699 drop(pending.lock().unwrap().insert("1".into(), first_tx));
700 drop(pending.lock().unwrap().insert("2".into(), second_tx));
701
702 fail_pending_calls(&pending, "stdout closed");
703
704 assert!(pending.lock().unwrap().is_empty());
705 assert_eq!(first_rx.await.unwrap().unwrap_err().to_string(), "Internal error: stdout closed");
706 assert_eq!(second_rx.await.unwrap().unwrap_err().to_string(), "Internal error: stdout closed");
707 }
708}