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    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}