Skip to main content

acp_utils/agent/
tokio_agent.rs

1use 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
102// ACP 2.0.0's SpawnedRun and Task::new wrap prior error data under another `data` field.
103fn 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);