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 if should_shutdown {
184 shutdown.store(true, Ordering::Release);
185 }
186 write_response(&mut stream, &response)
187}
188
189fn dispatch(
190 request: Request,
191 store: &Arc<Mutex<Store>>,
192 database_file: &Path,
193 write_lock: &Mutex<()>,
194 settings: &Settings,
195 commands: &CommandCatalog,
196) -> Response {
197 let result: Result<Response> = (|| match request {
198 Request::Ping => Ok(Response::Pong {
199 version: VERSION.to_owned(),
200 }),
201 Request::Shutdown => Ok(Response::ShuttingDown),
202 Request::Record {
203 command,
204 cwd,
205 exit_code,
206 observed_at_ms,
207 session_id,
208 } => {
209 if cwd.len() > MAX_PATH_BYTES {
210 bail!("working directory is too long");
211 }
212 if session_id.len() > MAX_SESSION_ID_BYTES {
213 bail!("session ID is too long");
214 }
215 if observed_at_ms.abs_diff(now_ms()) > 5 * 60 * 1_000 {
216 bail!("command timestamp is outside the allowed five-minute window");
217 }
218 let _write_guard = write_lock.lock().expect("write lock poisoned");
219 store.lock().expect("store lock poisoned").record(
220 &command,
221 &cwd,
222 exit_code,
223 observed_at_ms,
224 &session_id,
225 settings.history.ignore_leading_space,
226 )?;
227 Ok(Response::Recorded)
228 }
229 Request::Complete {
230 buffer,
231 cursor_byte,
232 cwd,
233 limit,
234 } => {
235 if cwd.len() > MAX_PATH_BYTES {
236 bail!("working directory is too long");
237 }
238 let mut completion = {
239 let store = store.lock().expect("store lock poisoned");
240 engine::complete(
241 &store,
242 commands,
243 &buffer,
244 cursor_byte,
245 &cwd,
246 limit,
247 settings,
248 )?
249 };
250 let limit = limit
251 .unwrap_or(settings.completion.max_candidates)
252 .min(settings.completion.max_candidates);
253 let paths =
254 engine::filesystem_candidates(&buffer, cursor_byte, &cwd, limit.saturating_add(1))?;
255 engine::merge_filesystem_candidates(&mut completion, paths, limit);
256 Ok(Response::Completion(completion))
257 }
258 Request::Fuzzy { query, cwd, limit } => {
259 if cwd.len() > MAX_PATH_BYTES {
260 bail!("working directory is too long");
261 }
262 let completion = engine::fuzzy(
263 &store.lock().expect("store lock poisoned"),
264 commands,
265 &query,
266 &cwd,
267 limit,
268 settings,
269 )?;
270 Ok(Response::Completion(completion))
271 }
272 Request::ImportHistory { path } => {
273 if path.len() > MAX_PATH_BYTES {
274 bail!("history path is too long");
275 }
276 let _write_guard = write_lock.lock().expect("write lock poisoned");
277 let mut import_store = Store::open(database_file)?;
278 let result = import_store
279 .import_zsh_history(Path::new(&path), settings.history.ignore_leading_space)?;
280 Ok(Response::Imported {
281 imported: result.imported,
282 skipped: result.skipped,
283 })
284 }
285 })();
286
287 result.unwrap_or_else(|error| Response::Error {
288 message: format!("{error:#}"),
289 })
290}
291
292fn now_ms() -> i64 {
293 SystemTime::now()
294 .duration_since(UNIX_EPOCH)
295 .unwrap_or_default()
296 .as_millis()
297 .min(i64::MAX as u128) as i64
298}
299
300fn write_response(stream: &mut UnixStream, response: &Response) -> Result<()> {
301 let mut payload = serde_json::to_vec(response)?;
302 if payload.len() > MAX_RESPONSE_BYTES {
303 payload = serde_json::to_vec(&Response::Error {
304 message: "response exceeds 1 MiB".to_owned(),
305 })?;
306 }
307 stream.write_all(&payload)?;
308 stream.write_all(b"\n")?;
309 stream.flush()?;
310 Ok(())
311}
312
313struct SocketGuard {
314 path: PathBuf,
315 device: u64,
316 inode: u64,
317}
318
319impl SocketGuard {
320 fn new(path: PathBuf) -> Result<Self> {
321 let metadata = fs::symlink_metadata(&path)?;
322 Ok(Self {
323 path,
324 device: metadata.dev(),
325 inode: metadata.ino(),
326 })
327 }
328}
329
330impl Drop for SocketGuard {
331 fn drop(&mut self) {
332 let Ok(metadata) = fs::symlink_metadata(&self.path) else {
333 return;
334 };
335 if metadata.file_type().is_socket()
336 && metadata.dev() == self.device
337 && metadata.ino() == self.inode
338 {
339 let _ = fs::remove_file(&self.path);
340 }
341 }
342}
343
344#[cfg(test)]
345mod tests {
346 use super::*;
347 use tempfile::tempdir;
348
349 #[test]
350 fn refuses_to_remove_a_non_socket_path() {
351 let directory = tempdir().unwrap();
352 let path = directory.path().join("aster.sock");
353 fs::write(&path, "keep me").unwrap();
354
355 assert!(prepare_socket(&path).is_err());
356 assert_eq!(fs::read_to_string(path).unwrap(), "keep me");
357 }
358
359 #[test]
360 fn daemon_lock_is_exclusive() {
361 let directory = tempdir().unwrap();
362 let path = directory.path().join("daemon.lock");
363 let first = acquire_daemon_lock(&path).unwrap();
364 assert!(acquire_daemon_lock(&path).is_err());
365 drop(first);
366 assert!(acquire_daemon_lock(&path).is_ok());
367 }
368}