acp_utils/agent/
tokio_agent.rs1use agent_client_protocol::util::internal_error;
2use agent_client_protocol::{
3 AcpAgent, AcpAgentConfig, ByteStreams, ConnectTo, Error, INCOMING_TRANSPORT_CLOSED_REASON, Role,
4 is_incoming_transport_closed,
5};
6use std::path::PathBuf;
7use std::process::{ExitStatus, Stdio};
8use std::str::FromStr;
9use tokio::io::{AsyncBufReadExt, BufReader};
10use tokio::process::Command;
11use tokio::sync::oneshot;
12use tokio::time::{Duration, timeout};
13use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt};
14
15pub struct TokioAcpAgent {
16 config: AcpAgentConfig,
17}
18
19impl TokioAcpAgent {
20 pub fn from_command(command: impl Into<PathBuf>, args: Vec<String>) -> Self {
21 Self { config: AcpAgentConfig::new(command).args(args) }
22 }
23
24 pub fn config(&self) -> &AcpAgentConfig {
25 &self.config
26 }
27}
28
29impl<T: Role> ConnectTo<T> for TokioAcpAgent {
30 async fn connect_to(self, client: impl ConnectTo<T::Counterpart>) -> Result<(), Error> {
31 connect_stdio::<T>(self.config, client).await
32 }
33}
34
35impl FromStr for TokioAcpAgent {
36 type Err = Error;
37
38 fn from_str(s: &str) -> Result<Self, Self::Err> {
39 Ok(Self { config: AcpAgent::from_str(s)?.into_config() })
40 }
41}
42
43async fn connect_stdio<T: Role>(config: AcpAgentConfig, client: impl ConnectTo<T::Counterpart>) -> Result<(), Error> {
44 let (stdin, stdout, stderr, mut child) = {
45 let mut cmd = Command::new(config.command());
46 cmd.args(config.arguments());
47 for (name, value) in config.environment() {
48 cmd.env(name, value);
49 }
50
51 let mut child = cmd
52 .stdin(Stdio::piped())
53 .stdout(Stdio::piped())
54 .stderr(Stdio::piped())
55 .kill_on_drop(true)
56 .spawn()
57 .map_err(Error::into_internal_error)?;
58
59 let stdin = child.stdin.take().ok_or_else(|| internal_error("missing child stdin"))?;
60 let stdout = child.stdout.take().ok_or_else(|| internal_error("missing child stdout"))?;
61 let stderr = child.stderr.take().ok_or_else(|| internal_error("missing child stderr"))?;
62 (stdin, stdout, stderr, child)
63 };
64
65 let (stderr_tx, stderr_rx) = oneshot::channel::<String>();
66 tokio::spawn(async move {
67 let mut lines = BufReader::new(stderr).lines();
68 let mut buf = String::new();
69 while let Ok(Some(line)) = lines.next_line().await {
70 if !buf.is_empty() {
71 buf.push('\n');
72 }
73 buf.push_str(&line);
74 }
75 let _ = stderr_tx.send(buf);
76 });
77
78 let child_fut = async move {
79 let status = child.wait().await.map_err(Error::into_internal_error)?;
80 finish_child_exit(status, stderr_rx).await
81 };
82
83 let bytes = ByteStreams::new(stdin.compat_write(), stdout.compat());
84 let protocol_fut = ConnectTo::<T>::connect_to(bytes, client);
85 tokio::pin!(child_fut);
86
87 tokio::select! {
88 result = &mut child_fut => result,
89 result = protocol_fut => match result {
90 Ok(()) => timeout(SHUTDOWN_GRACE_PERIOD, &mut child_fut).await.unwrap_or(Ok(())),
91 Err(protocol_error) if has_incoming_transport_closed(&protocol_error) => {
92 match timeout(SHUTDOWN_GRACE_PERIOD, &mut child_fut).await {
93 Ok(Err(child_error)) => Err(child_error),
94 _ => Err(protocol_error),
95 }
96 }
97 Err(error) => Err(error),
98 },
99 }
100}
101
102fn has_incoming_transport_closed(error: &Error) -> bool {
104 fn data_has_reason(data: &serde_json::Value) -> bool {
105 data.get("reason").and_then(serde_json::Value::as_str) == Some(INCOMING_TRANSPORT_CLOSED_REASON)
106 || data.get("data").is_some_and(data_has_reason)
107 }
108
109 is_incoming_transport_closed(error) || error.data.as_ref().is_some_and(data_has_reason)
110}
111
112async fn finish_child_exit(status: ExitStatus, stderr_rx: oneshot::Receiver<String>) -> Result<(), Error> {
113 if status.success() {
114 return Ok(());
115 }
116
117 let stderr = match timeout(SHUTDOWN_GRACE_PERIOD, stderr_rx).await {
118 Ok(Ok(stderr)) => stderr,
119 _ => String::new(),
120 };
121 let message = if stderr.is_empty() {
122 format!("agent process exited ({status})")
123 } else {
124 format!("agent process exited ({status}): {stderr}")
125 };
126
127 Err(internal_error(message))
128}
129
130const SHUTDOWN_GRACE_PERIOD: Duration = Duration::from_secs(1);