Skip to main content

aether_cli/acp/
server.rs

1use super::state::{AcpState, ClientGuard};
2use super::{AcpArgs, AcpRunError, IdleHook, SessionHooks, create_acp_state};
3use crate::output::OutputFormat;
4use crate::prompt::prompt_or_stdin;
5use acp_utils::websocket::WebSocketTransport;
6use std::fs::canonicalize;
7use std::future::Future;
8use std::io;
9use std::net::SocketAddr;
10use std::path::PathBuf;
11use std::sync::Arc;
12use std::time::Duration;
13use thiserror::Error;
14use tokio::net::{TcpListener, TcpStream};
15#[cfg(unix)]
16use tokio::signal::unix::{SignalKind, signal};
17use tokio::task::JoinSet;
18use tokio_tungstenite::tungstenite::{
19    handshake::server::{ErrorResponse, Request},
20    http::StatusCode,
21};
22use tokio_util::sync::CancellationToken;
23use tracing::{info, warn};
24
25#[derive(clap::Args, Debug)]
26pub struct ServerArgs {
27    /// Address to listen on. The raw server is unauthenticated; use a private network or authenticating proxy.
28    #[arg(long, default_value = "127.0.0.1:8765")]
29    pub listen: SocketAddr,
30
31    /// Server workspace used to resolve settings and as the default session directory.
32    #[arg(short = 'C', long, default_value = ".")]
33    pub cwd: PathBuf,
34
35    /// Initial prompt to run.
36    #[arg(long)]
37    pub prompt: Option<String>,
38
39    /// Format for agent events written to stdout/stderr.
40    #[arg(long, default_value = "text")]
41    pub output: OutputFormat,
42
43    /// How many seconds to wait after the agent becomes idle before running the --on-idle command.
44    #[arg(long, requires = "on_idle")]
45    pub idle_after: Option<u64>,
46
47    /// Shell command to spawn when the agent becomes idle.
48    #[arg(long, requires = "idle_after")]
49    pub on_idle: Option<String>,
50
51    #[command(flatten)]
52    pub acp: AcpArgs,
53}
54
55#[derive(Debug, Error)]
56pub enum ServerRunError {
57    #[error("Invalid server workspace {path}: {source}")]
58    Workspace { path: PathBuf, source: io::Error },
59    #[error(transparent)]
60    Initialization(#[from] AcpRunError),
61    #[error("Failed to bind ACP listener at {address}: {source}")]
62    Bind { address: SocketAddr, source: io::Error },
63    #[error("ACP listener failed: {0}")]
64    Accept(#[source] io::Error),
65    #[error("Failed to listen for shutdown signals: {0}")]
66    Signal(#[source] io::Error),
67    #[error("Failed to read the initial prompt from stdin: {0}")]
68    PromptStdin(#[source] io::Error),
69    #[error("Failed to start the initial session: {0}")]
70    InitialSession(#[source] agent_client_protocol::Error),
71}
72
73pub async fn run_server(args: ServerArgs) -> Result<(), ServerRunError> {
74    let cwd = canonicalize(&args.cwd)
75        .and_then(|cwd| {
76            if cwd.is_dir() {
77                Ok(cwd)
78            } else {
79                Err(io::Error::new(io::ErrorKind::NotADirectory, "workspace must be a directory"))
80            }
81        })
82        .map_err(|source| ServerRunError::Workspace { path: args.cwd, source })?;
83
84    let prompt = prompt_or_stdin(args.prompt).map_err(ServerRunError::PromptStdin)?;
85    let hooks = SessionHooks {
86        echo: Some(args.output),
87        idle: args
88            .idle_after
89            .zip(args.on_idle)
90            .map(|(seconds, command)| IdleHook { after: Duration::from_secs(seconds), command }),
91    };
92    let state = Arc::new(create_acp_state(args.acp, &cwd, hooks)?);
93    let server = AcpServer::bind(args.listen, state.clone()).await?;
94    info!(address = %args.listen, cwd = %cwd.display(), "Starting Aether ACP WebSocket server");
95
96    let result = async {
97        if let Some(prompt) = prompt {
98            let session_id = state.start_session(prompt).await.map_err(ServerRunError::InitialSession)?;
99            info!(session_id = %session_id.0, "Started initial server session");
100        }
101        server.run_until(async { shutdown_signal().await.map_err(ServerRunError::Signal) }).await
102    }
103    .await;
104    state.shutdown_all().await;
105    result
106}
107
108/// Owns networking, not session lifetime. Drop closes the listener and aborts
109/// connection tasks; explicit shutdown also waits for output detachment.
110pub(crate) struct AcpServer {
111    listener: TcpListener,
112    state: Arc<AcpState>,
113    connections: JoinSet<()>,
114    stop: CancellationToken,
115}
116
117impl AcpServer {
118    pub(crate) async fn bind(address: SocketAddr, state: Arc<AcpState>) -> Result<Self, ServerRunError> {
119        let listener = TcpListener::bind(address).await.map_err(|source| ServerRunError::Bind { address, source })?;
120        Ok(Self { listener, stop: state.stop_token().child_token(), state, connections: JoinSet::new() })
121    }
122
123    #[cfg(any(test, feature = "testing"))]
124    pub(crate) fn local_addr(&self) -> io::Result<SocketAddr> {
125        self.listener.local_addr()
126    }
127
128    pub(crate) async fn run_until(
129        mut self,
130        stop: impl Future<Output = Result<(), ServerRunError>>,
131    ) -> Result<(), ServerRunError> {
132        let result = tokio::select! {
133            result = self.run() => result,
134            result = stop => result,
135        };
136        self.shutdown().await;
137        result
138    }
139
140    async fn run(&mut self) -> Result<(), ServerRunError> {
141        loop {
142            tokio::select! {
143                biased;
144                () = self.stop.cancelled() => return Ok(()),
145                result = self.connections.join_next(), if !self.connections.is_empty() => {
146                    if let Some(Err(error)) = result {
147                        warn!(%error, "ACP connection task failed");
148                    }
149                }
150                accepted = self.listener.accept() => {
151                    let (socket, peer) = accepted.map_err(ServerRunError::Accept)?;
152                    self.accept(socket, peer);
153                }
154            }
155        }
156    }
157
158    /// Close admission and join all networking, leaving the host running.
159    pub(crate) async fn shutdown(self) {
160        self.stop.cancel();
161        drop(self.listener);
162        let mut connections = self.connections;
163        while let Some(result) = connections.join_next().await {
164            if let Err(error) = result
165                && !error.is_cancelled()
166            {
167                warn!(%error, "ACP connection task failed during shutdown");
168            }
169        }
170    }
171
172    fn accept(&mut self, socket: TcpStream, peer: SocketAddr) {
173        let stop = self.stop.clone();
174        let guard = self.state.try_attach();
175        self.connections.spawn(serve_connection(socket, peer, guard, stop));
176    }
177}
178
179#[expect(clippy::result_large_err, reason = "tungstenite's handshake callback requires an HTTP error response")]
180async fn serve_connection(socket: TcpStream, peer: SocketAddr, guard: Option<ClientGuard>, stop: CancellationToken) {
181    let handshake = tokio::select! {
182        biased;
183        () = stop.cancelled() => {
184            if let Some(guard) = guard {
185                guard.release().await;
186            }
187            return;
188        },
189        result = tokio_tungstenite::accept_hdr_async(socket, |_request: &Request, response| {
190            if guard.is_some() {
191                Ok(response)
192            } else {
193                let mut error = ErrorResponse::new(Some("client already attached".to_string()));
194                *error.status_mut() = StatusCode::CONFLICT;
195                Err(error)
196            }
197        }) => result,
198    };
199    let socket = match handshake {
200        Ok(socket) => socket,
201        Err(error) => {
202            warn!(%peer, %error, "ACP WebSocket handshake failed");
203            if let Some(guard) = guard {
204                guard.release().await;
205            }
206            return;
207        }
208    };
209    if let Err(error) =
210        guard.expect("successful handshake owns the client").serve(WebSocketTransport::new(socket), stop).await
211    {
212        warn!(%peer, %error, "ACP connection failed");
213    }
214}
215
216async fn shutdown_signal() -> io::Result<()> {
217    #[cfg(unix)]
218    {
219        let mut terminate = signal(SignalKind::terminate())?;
220        tokio::select! {
221            result = tokio::signal::ctrl_c() => result,
222            _ = terminate.recv() => Ok(()),
223        }
224    }
225    #[cfg(not(unix))]
226    tokio::signal::ctrl_c().await
227}