1use std::fs::{self, File, OpenOptions};
2use std::io::{BufRead, BufReader, Read, Write};
3use std::os::unix::fs::{FileTypeExt, MetadataExt, PermissionsExt};
4use std::os::unix::net::{UnixListener, UnixStream};
5use std::path::{Path, PathBuf};
6use std::sync::atomic::{AtomicBool, Ordering};
7use std::sync::{Arc, Mutex};
8use std::thread;
9use std::time::{Duration, SystemTime, UNIX_EPOCH};
10
11use anyhow::{Context, Result, bail};
12use fs2::FileExt;
13
14use crate::VERSION;
15use crate::commands::CommandCatalog;
16use crate::config::{Paths, Settings};
17use crate::engine;
18use crate::protocol::{PROTOCOL_VERSION, Request, RequestEnvelope, Response};
19use crate::store::Store;
20
21const MAX_REQUEST_BYTES: u64 = 1024 * 1024;
22const MAX_RESPONSE_BYTES: usize = 1024 * 1024;
23const MAX_CONNECTIONS: usize = 64;
24const MAX_PATH_BYTES: usize = 16 * 1024;
25const MAX_SESSION_ID_BYTES: usize = 1024;
26const IO_TIMEOUT: Duration = Duration::from_secs(3);
27
28pub fn serve(paths: Paths, settings: Settings) -> Result<()> {
29 paths.ensure_directories()?;
30 let _daemon_lock = acquire_daemon_lock(&paths.daemon_lock_file)?;
31 prepare_socket(&paths.socket_file)?;
32 let listener = UnixListener::bind(&paths.socket_file)
33 .with_context(|| format!("failed to bind socket {}", paths.socket_file.display()))?;
34 listener.set_nonblocking(true)?;
35 fs::set_permissions(&paths.socket_file, fs::Permissions::from_mode(0o600))?;
36 let _socket_guard = SocketGuard::new(paths.socket_file.clone())?;
37
38 let store = Arc::new(Mutex::new(Store::open(&paths.database_file)?));
39 let database_file = Arc::new(paths.database_file.clone());
40 let write_lock = Arc::new(Mutex::new(()));
41 let settings = Arc::new(settings);
42 let commands = Arc::new(CommandCatalog::discover(
43 paths.command_description_cache.clone(),
44 ));
45 let shutdown = Arc::new(AtomicBool::new(false));
46 let mut workers = Vec::new();
47
48 while !shutdown.load(Ordering::Acquire) {
49 reap_finished_workers(&mut workers);
50 match listener.accept() {
51 Ok((mut stream, _)) if workers.len() >= MAX_CONNECTIONS => {
52 let _ = write_response(
53 &mut stream,
54 &Response::Error {
55 message: "daemon is at its connection limit".to_owned(),
56 },
57 );
58 }
59 Ok((stream, _)) => {
60 let store = Arc::clone(&store);
61 let database_file = Arc::clone(&database_file);
62 let write_lock = Arc::clone(&write_lock);
63 let settings = Arc::clone(&settings);
64 let commands = Arc::clone(&commands);
65 let shutdown = Arc::clone(&shutdown);
66 let worker = thread::Builder::new()
67 .name("aster-client".to_owned())
68 .spawn(move || {
69 if let Err(error) = handle_connection(
70 stream,
71 &store,
72 &database_file,
73 &write_lock,
74 &settings,
75 &commands,
76 &shutdown,
77 ) {
78 eprintln!("aster: request failed: {error:#}");
79 }
80 })?;
81 workers.push(worker);
82 }
83 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
84 thread::sleep(Duration::from_millis(5));
85 }
86 Err(error) => {
87 eprintln!("aster: failed to accept connection: {error}");
88 thread::sleep(Duration::from_millis(50));
89 }
90 }
91 }
92 for worker in workers {
93 let _ = worker.join();
94 }
95 Ok(())
96}
97
98fn reap_finished_workers(workers: &mut Vec<thread::JoinHandle<()>>) {
99 let mut index = 0;
100 while index < workers.len() {
101 if workers[index].is_finished() {
102 let worker = workers.swap_remove(index);
103 let _ = worker.join();
104 } else {
105 index += 1;
106 }
107 }
108}
109
110fn prepare_socket(path: &Path) -> Result<()> {
111 let Ok(metadata) = fs::symlink_metadata(path) else {
112 return Ok(());
113 };
114 if !metadata.file_type().is_socket() {
115 bail!("refusing to replace non-socket path {}", path.display());
116 }
117 if UnixStream::connect(path).is_ok() {
118 bail!("aster daemon is already running at {}", path.display());
119 }
120 fs::remove_file(path)
121 .with_context(|| format!("failed to remove stale socket {}", path.display()))
122}
123
124fn acquire_daemon_lock(path: &Path) -> Result<File> {
125 let file = OpenOptions::new()
126 .create(true)
127 .read(true)
128 .write(true)
129 .truncate(false)
130 .open(path)
131 .with_context(|| format!("failed to open daemon lock {}", path.display()))?;
132 fs::set_permissions(path, fs::Permissions::from_mode(0o600))?;
133 file.try_lock_exclusive()
134 .with_context(|| format!("aster daemon is already running for {}", path.display()))?;
135 Ok(file)
136}
137
138fn handle_connection(
139 mut stream: UnixStream,
140 store: &Arc<Mutex<Store>>,
141 database_file: &Path,
142 write_lock: &Mutex<()>,
143 settings: &Settings,
144 commands: &CommandCatalog,
145 shutdown: &AtomicBool,
146) -> Result<()> {
147 stream.set_read_timeout(Some(IO_TIMEOUT))?;
148 stream.set_write_timeout(Some(IO_TIMEOUT))?;
149
150 let mut payload = String::new();
151 BufReader::new(stream.try_clone()?)
152 .take(MAX_REQUEST_BYTES + 1)
153 .read_line(&mut payload)?;
154 if payload.len() as u64 > MAX_REQUEST_BYTES {
155 return write_response(
156 &mut stream,
157 &Response::Error {
158 message: "request exceeds 1 MiB".to_owned(),
159 },
160 );
161 }
162
163 let response = match serde_json::from_str::<RequestEnvelope>(&payload) {
164 Ok(envelope) if envelope.version == PROTOCOL_VERSION => dispatch(
165 envelope.request,
166 store,
167 database_file,
168 write_lock,
169 settings,
170 commands,
171 ),
172 Ok(envelope) => Response::Error {
173 message: format!(
174 "unsupported protocol version {}; expected {PROTOCOL_VERSION}",
175 envelope.version
176 ),
177 },
178 Err(error) => Response::Error {
179 message: format!("invalid request: {error}"),
180 },
181 };
182 let should_shutdown = matches!(response, Response::ShuttingDown);
183 write_response(&mut stream, &response)?;
184 if should_shutdown {
185 shutdown.store(true, Ordering::Release);
186 }
187 Ok(())
188}
189
190fn dispatch(
191 request: Request,
192 store: &Arc<Mutex<Store>>,
193 database_file: &Path,
194 write_lock: &Mutex<()>,
195 settings: &Settings,
196 commands: &CommandCatalog,
197) -> Response {
198 let result: Result<Response> = (|| match request {
199 Request::Ping => Ok(Response::Pong {
200 version: VERSION.to_owned(),
201 }),
202 Request::Shutdown => Ok(Response::ShuttingDown),
203 Request::Record {
204 command,
205 cwd,
206 exit_code,
207 observed_at_ms,
208 session_id,
209 } => {
210 if cwd.len() > MAX_PATH_BYTES {
211 bail!("working directory is too long");
212 }
213 if session_id.len() > MAX_SESSION_ID_BYTES {
214 bail!("session ID is too long");
215 }
216 if observed_at_ms.abs_diff(now_ms()) > 5 * 60 * 1_000 {
217 bail!("command timestamp is outside the allowed five-minute window");
218 }
219 let _write_guard = write_lock.lock().expect("write lock poisoned");
220 store.lock().expect("store lock poisoned").record(
221 &command,
222 &cwd,
223 exit_code,
224 observed_at_ms,
225 &session_id,
226 settings.history.ignore_leading_space,
227 )?;
228 Ok(Response::Recorded)
229 }
230 Request::Complete {
231 buffer,
232 cursor_byte,
233 cwd,
234 limit,
235 } => {
236 if cwd.len() > MAX_PATH_BYTES {
237 bail!("working directory is too long");
238 }
239 let completion = engine::complete(
240 &store.lock().expect("store lock poisoned"),
241 commands,
242 &buffer,
243 cursor_byte,
244 &cwd,
245 limit,
246 settings,
247 )?;
248 Ok(Response::Completion(completion))
249 }
250 Request::Fuzzy { query, cwd, limit } => {
251 if cwd.len() > MAX_PATH_BYTES {
252 bail!("working directory is too long");
253 }
254 let completion = engine::fuzzy(
255 &store.lock().expect("store lock poisoned"),
256 commands,
257 &query,
258 &cwd,
259 limit,
260 settings,
261 )?;
262 Ok(Response::Completion(completion))
263 }
264 Request::ImportHistory { path } => {
265 if path.len() > MAX_PATH_BYTES {
266 bail!("history path is too long");
267 }
268 let _write_guard = write_lock.lock().expect("write lock poisoned");
269 let mut import_store = Store::open(database_file)?;
270 let result = import_store
271 .import_zsh_history(Path::new(&path), settings.history.ignore_leading_space)?;
272 Ok(Response::Imported {
273 imported: result.imported,
274 skipped: result.skipped,
275 })
276 }
277 })();
278
279 result.unwrap_or_else(|error| Response::Error {
280 message: format!("{error:#}"),
281 })
282}
283
284fn now_ms() -> i64 {
285 SystemTime::now()
286 .duration_since(UNIX_EPOCH)
287 .unwrap_or_default()
288 .as_millis()
289 .min(i64::MAX as u128) as i64
290}
291
292fn write_response(stream: &mut UnixStream, response: &Response) -> Result<()> {
293 let mut payload = serde_json::to_vec(response)?;
294 if payload.len() > MAX_RESPONSE_BYTES {
295 payload = serde_json::to_vec(&Response::Error {
296 message: "response exceeds 1 MiB".to_owned(),
297 })?;
298 }
299 stream.write_all(&payload)?;
300 stream.write_all(b"\n")?;
301 stream.flush()?;
302 Ok(())
303}
304
305struct SocketGuard {
306 path: PathBuf,
307 device: u64,
308 inode: u64,
309}
310
311impl SocketGuard {
312 fn new(path: PathBuf) -> Result<Self> {
313 let metadata = fs::symlink_metadata(&path)?;
314 Ok(Self {
315 path,
316 device: metadata.dev(),
317 inode: metadata.ino(),
318 })
319 }
320}
321
322impl Drop for SocketGuard {
323 fn drop(&mut self) {
324 let Ok(metadata) = fs::symlink_metadata(&self.path) else {
325 return;
326 };
327 if metadata.file_type().is_socket()
328 && metadata.dev() == self.device
329 && metadata.ino() == self.inode
330 {
331 let _ = fs::remove_file(&self.path);
332 }
333 }
334}
335
336#[cfg(test)]
337mod tests {
338 use super::*;
339 use tempfile::tempdir;
340
341 #[test]
342 fn refuses_to_remove_a_non_socket_path() {
343 let directory = tempdir().unwrap();
344 let path = directory.path().join("aster.sock");
345 fs::write(&path, "keep me").unwrap();
346
347 assert!(prepare_socket(&path).is_err());
348 assert_eq!(fs::read_to_string(path).unwrap(), "keep me");
349 }
350
351 #[test]
352 fn daemon_lock_is_exclusive() {
353 let directory = tempdir().unwrap();
354 let path = directory.path().join("daemon.lock");
355 let first = acquire_daemon_lock(&path).unwrap();
356 assert!(acquire_daemon_lock(&path).is_err());
357 drop(first);
358 assert!(acquire_daemon_lock(&path).is_ok());
359 }
360}