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 #[arg(long, default_value = "127.0.0.1:8765")]
29 pub listen: SocketAddr,
30
31 #[arg(short = 'C', long, default_value = ".")]
33 pub cwd: PathBuf,
34
35 #[arg(long)]
37 pub prompt: Option<String>,
38
39 #[arg(long, default_value = "text")]
41 pub output: OutputFormat,
42
43 #[arg(long, requires = "on_idle")]
45 pub idle_after: Option<u64>,
46
47 #[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
108pub(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 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}