Skip to main content

aster/
daemon.rs

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}