Skip to main content

ssh_mcp/background/
spooler.rs

1use std::env;
2use std::path::{Component, Path, PathBuf};
3use std::time::{Duration, SystemTime};
4
5#[cfg(unix)]
6use std::ffi::OsStr;
7#[cfg(unix)]
8use std::os::unix::fs::{MetadataExt, PermissionsExt};
9
10use tokio::fs;
11use tokio::io::{AsyncReadExt, AsyncWriteExt};
12use tracing::warn;
13
14use super::job::{JobState, PersistedJobState};
15use super::{BackgroundError, Result};
16#[cfg(unix)]
17use crate::platform::O_NOFOLLOW_FLAG;
18
19/// Local-only job state and log spooler.
20#[derive(Debug, Clone)]
21pub struct LocalLogSpooler {
22    base_dir: PathBuf,
23}
24
25impl LocalLogSpooler {
26    pub fn new(base_dir: PathBuf) -> Self {
27        Self { base_dir }
28    }
29
30    pub fn new_default() -> Self {
31        #[cfg(unix)]
32        let base_dir = default_spool_dir(
33            env::var_os("XDG_RUNTIME_DIR").as_deref(),
34            &env::temp_dir(),
35            rustix::process::geteuid().as_raw(),
36        );
37        #[cfg(not(unix))]
38        let base_dir = env::temp_dir().join("ssh-mcp");
39
40        Self::new(base_dir)
41    }
42
43    pub fn base_dir(&self) -> &Path {
44        &self.base_dir
45    }
46
47    pub async fn ensure_dir(&self) -> Result<()> {
48        #[cfg(unix)]
49        let expected_uid = Some(rustix::process::geteuid().as_raw());
50        #[cfg(not(unix))]
51        let expected_uid = None;
52
53        self.ensure_dir_inner(expected_uid).await
54    }
55
56    async fn ensure_dir_inner(&self, expected_uid: Option<u32>) -> Result<()> {
57        let created = match create_spool_dir(&self.base_dir).await {
58            Ok(()) => true,
59            Err(e) if e.kind() == std::io::ErrorKind::AlreadyExists => false,
60            Err(e) => return Err(e.into()),
61        };
62
63        let meta = fs::symlink_metadata(&self.base_dir).await?;
64        validate_spool_dir_meta(&meta)?;
65
66        #[cfg(unix)]
67        {
68            let expected_uid = expected_uid.expect("effective UID is available on Unix");
69            validate_spool_dir_owner(&meta, expected_uid)?;
70            if !created && meta.permissions().mode() & 0o022 != 0 {
71                return Err(BackgroundError::InvalidState {
72                    message: "spool directory is group- or world-writable",
73                });
74            }
75            if meta.permissions().mode() & 0o777 != 0o700 {
76                let perms = std::fs::Permissions::from_mode(0o700);
77                fs::set_permissions(&self.base_dir, perms).await?;
78            }
79        }
80        #[cfg(not(unix))]
81        let _ = expected_uid;
82
83        Ok(())
84    }
85
86    pub fn log_path_for(&self, job_id: &str) -> Result<PathBuf> {
87        validate_job_id(job_id)?;
88        Ok(self.base_dir.join(format!("{job_id}.log")))
89    }
90
91    pub fn state_path_for(&self, job_id: &str) -> Result<PathBuf> {
92        validate_job_id(job_id)?;
93        Ok(self.base_dir.join(format!("{job_id}.state")))
94    }
95
96    pub async fn persist_job_state(&self, job: &JobState) -> Result<()> {
97        self.ensure_dir().await?;
98
99        if job.log_path.parent() != Some(self.base_dir()) {
100            return Err(BackgroundError::InvalidState {
101                message: "job log path is outside spool directory",
102            });
103        }
104
105        let path = self.state_path_for(&job.job_id)?;
106        let payload =
107            serde_json::to_vec(&job.to_persisted()).map_err(|_| BackgroundError::InvalidState {
108                message: "failed to serialize persisted job state",
109            })?;
110
111        let mut file = open_spool_write_no_symlink(&path).await?;
112        file.write_all(&payload).await?;
113        file.sync_all().await?;
114        Ok(())
115    }
116
117    pub async fn load_job_state(&self, job_id: &str) -> Result<Option<JobState>> {
118        self.ensure_dir().await?;
119        let path = self.state_path_for(job_id)?;
120
121        let mut file = match open_spool_read_no_symlink(&path).await? {
122            Some(file) => file,
123            None => return Ok(None),
124        };
125
126        let mut payload = Vec::new();
127        file.read_to_end(&mut payload).await?;
128
129        let persisted: PersistedJobState =
130            serde_json::from_slice(&payload).map_err(|_| BackgroundError::InvalidState {
131                message: "failed to parse persisted job state",
132            })?;
133        let job = JobState::from_persisted(persisted)
134            .map_err(|message| BackgroundError::InvalidState { message })?;
135
136        if job.job_id != job_id {
137            return Err(BackgroundError::InvalidState {
138                message: "persisted job id does not match requested job id",
139            });
140        }
141        if job.log_path.parent() != Some(self.base_dir()) {
142            return Err(BackgroundError::InvalidState {
143                message: "persisted log path is outside spool directory",
144            });
145        }
146
147        Ok(Some(job))
148    }
149
150    pub async fn cleanup_old_logs(&self, max_age: Duration) -> Result<usize> {
151        self.ensure_dir().await?;
152
153        let now = SystemTime::now();
154        let mut removed = 0usize;
155
156        let mut entries = match fs::read_dir(&self.base_dir).await {
157            Ok(e) => e,
158            Err(e) => return Err(e.into()),
159        };
160
161        loop {
162            let entry = match entries.next_entry().await {
163                Ok(Some(e)) => e,
164                Ok(None) => break,
165                Err(e) => {
166                    warn!(error = ?e, "failed to read spool directory entry");
167                    continue;
168                }
169            };
170            let path = entry.path();
171
172            let file_name = match entry.file_name().to_str() {
173                Some(s) => s.to_owned(),
174                None => continue,
175            };
176
177            let Some((job_id, ext)) = split_spool_file_name(&file_name) else {
178                continue;
179            };
180            if validate_job_id(job_id).is_err() {
181                continue;
182            }
183            if ext != "log" && ext != "exit" && ext != "state" {
184                continue;
185            }
186
187            let meta = match fs::symlink_metadata(&path).await {
188                Ok(m) => m,
189                Err(e) => {
190                    warn!(path = ?path, error = ?e, "failed to stat spool file");
191                    continue;
192                }
193            };
194
195            let ft = meta.file_type();
196            if ft.is_symlink() || !ft.is_file() {
197                continue;
198            }
199
200            let modified = match meta.modified() {
201                Ok(m) => m,
202                Err(e) => {
203                    warn!(path = ?path, error = ?e, "failed to read mtime");
204                    continue;
205                }
206            };
207
208            let age = match now.duration_since(modified) {
209                Ok(d) => d,
210                Err(e) => {
211                    warn!(path = ?path, error = ?e, "invalid modified time");
212                    continue;
213                }
214            };
215            if age <= max_age {
216                continue;
217            }
218
219            match fs::remove_file(&path).await {
220                Ok(()) => removed += 1,
221                Err(e) => {
222                    warn!(path = ?path, error = ?e, "failed to remove old spool file");
223                }
224            }
225        }
226
227        Ok(removed)
228    }
229}
230
231#[cfg(unix)]
232fn default_spool_dir(runtime_dir: Option<&OsStr>, temp_dir: &Path, effective_uid: u32) -> PathBuf {
233    if let Some(runtime_dir) = runtime_dir
234        .map(PathBuf::from)
235        .filter(|path| path.is_absolute())
236    {
237        return runtime_dir.join("ssh-mcp");
238    }
239
240    let temp_dir = if temp_dir.is_absolute() {
241        temp_dir
242    } else {
243        Path::new("/tmp")
244    };
245    temp_dir.join(format!("ssh-mcp-{effective_uid}"))
246}
247
248#[cfg(unix)]
249async fn create_spool_dir(path: &Path) -> std::io::Result<()> {
250    let mut builder = fs::DirBuilder::new();
251    builder.mode(0o700);
252    builder.create(path).await
253}
254
255#[cfg(not(unix))]
256async fn create_spool_dir(path: &Path) -> std::io::Result<()> {
257    fs::create_dir(path).await
258}
259
260async fn open_spool_write_no_symlink(path: &Path) -> Result<tokio::fs::File> {
261    match fs::symlink_metadata(path).await {
262        Ok(meta) if meta.file_type().is_symlink() => {
263            return Err(BackgroundError::InvalidState {
264                message: "spool metadata path is a symlink",
265            });
266        }
267        Ok(_) => {}
268        Err(e) if e.kind() == std::io::ErrorKind::NotFound => {}
269        Err(e) => return Err(e.into()),
270    }
271
272    let mut opts = tokio::fs::OpenOptions::new();
273    opts.write(true).create(true).truncate(true);
274
275    #[cfg(unix)]
276    {
277        opts.custom_flags(O_NOFOLLOW_FLAG);
278    }
279
280    match opts.open(path).await {
281        Ok(file) => Ok(file),
282        Err(e) => {
283            if let Ok(meta) = fs::symlink_metadata(path).await
284                && meta.file_type().is_symlink()
285            {
286                return Err(BackgroundError::InvalidState {
287                    message: "spool metadata path is a symlink",
288                });
289            }
290            Err(e.into())
291        }
292    }
293}
294
295async fn open_spool_read_no_symlink(path: &Path) -> Result<Option<tokio::fs::File>> {
296    match fs::symlink_metadata(path).await {
297        Ok(meta) => {
298            if meta.file_type().is_symlink() {
299                return Err(BackgroundError::InvalidState {
300                    message: "spool metadata path is a symlink",
301                });
302            }
303            if !meta.is_file() {
304                return Err(BackgroundError::InvalidState {
305                    message: "spool metadata path is not a regular file",
306                });
307            }
308        }
309        Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(None),
310        Err(e) => return Err(e.into()),
311    }
312
313    let mut opts = tokio::fs::OpenOptions::new();
314    opts.read(true);
315
316    #[cfg(unix)]
317    {
318        opts.custom_flags(O_NOFOLLOW_FLAG);
319    }
320
321    match opts.open(path).await {
322        Ok(file) => Ok(Some(file)),
323        Err(e) => {
324            if let Ok(meta) = fs::symlink_metadata(path).await
325                && meta.file_type().is_symlink()
326            {
327                return Err(BackgroundError::InvalidState {
328                    message: "spool metadata path is a symlink",
329                });
330            }
331            Err(e.into())
332        }
333    }
334}
335
336fn validate_spool_dir_meta(meta: &std::fs::Metadata) -> Result<()> {
337    let ft = meta.file_type();
338    if ft.is_symlink() {
339        return Err(BackgroundError::InvalidState {
340            message: "spool directory is a symlink",
341        });
342    }
343    if !ft.is_dir() {
344        return Err(BackgroundError::InvalidState {
345            message: "spool path exists but is not a directory",
346        });
347    }
348    Ok(())
349}
350
351#[cfg(unix)]
352fn validate_spool_dir_owner(meta: &std::fs::Metadata, expected_uid: u32) -> Result<()> {
353    if meta.uid() != expected_uid {
354        return Err(BackgroundError::InvalidState {
355            message: "spool directory is not owned by the effective user",
356        });
357    }
358    Ok(())
359}
360
361fn validate_job_id(job_id: &str) -> Result<()> {
362    if job_id.is_empty() || job_id.len() > 128 {
363        return Err(BackgroundError::InvalidJobId {
364            job_id: job_id.to_owned(),
365        });
366    }
367    if job_id.as_bytes().contains(&0) {
368        return Err(BackgroundError::InvalidJobId {
369            job_id: job_id.to_owned(),
370        });
371    }
372
373    // Require a single, normal path component (reject absolute paths, separators, '.', '..').
374    let p = Path::new(job_id);
375    let mut components = p.components();
376    let Some(Component::Normal(_)) = components.next() else {
377        return Err(BackgroundError::InvalidJobId {
378            job_id: job_id.to_owned(),
379        });
380    };
381    if components.next().is_some() {
382        return Err(BackgroundError::InvalidJobId {
383            job_id: job_id.to_owned(),
384        });
385    }
386
387    // Tighten further: allow only ASCII alnum plus '-' and '_' to keep file naming predictable.
388    if !job_id
389        .bytes()
390        .all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_')
391    {
392        return Err(BackgroundError::InvalidJobId {
393            job_id: job_id.to_owned(),
394        });
395    }
396
397    Ok(())
398}
399
400fn split_spool_file_name(name: &str) -> Option<(&str, &str)> {
401    let (stem, ext) = name.rsplit_once('.')?;
402    if stem.is_empty() || ext.is_empty() {
403        return None;
404    }
405    Some((stem, ext))
406}
407
408#[cfg(test)]
409mod tests {
410    use super::*;
411    use std::time::{Instant, SystemTime};
412
413    #[cfg(unix)]
414    use std::os::unix::fs::{MetadataExt, PermissionsExt, symlink};
415
416    async fn wait_until_older_than(path: &Path, min_age: Duration) {
417        let start = Instant::now();
418        loop {
419            let meta = tokio::fs::metadata(path).await.expect("metadata");
420            let modified = meta.modified().expect("modified time");
421            let age = SystemTime::now()
422                .duration_since(modified)
423                .unwrap_or_else(|_| Duration::from_secs(0));
424
425            if age >= min_age {
426                return;
427            }
428
429            assert!(
430                start.elapsed() < Duration::from_secs(2),
431                "file did not become old enough: {path:?}"
432            );
433            tokio::time::sleep(Duration::from_millis(5)).await;
434        }
435    }
436
437    #[tokio::test]
438    async fn test_ensure_dir_and_log_path_for() {
439        let tmp = tempfile::TempDir::new().expect("tempdir");
440        let base = tmp.path().join("spool");
441        let spooler = LocalLogSpooler::new(base.clone());
442
443        spooler.ensure_dir().await.expect("ensure_dir");
444        let meta = std::fs::metadata(&base).expect("spool dir metadata");
445        assert!(meta.is_dir());
446        #[cfg(unix)]
447        assert_eq!(meta.permissions().mode() & 0o777, 0o700);
448
449        let log = spooler.log_path_for("job_123").expect("log_path_for");
450        assert_eq!(log, base.join("job_123.log"));
451        let state = spooler.state_path_for("job_123").expect("state_path_for");
452        assert_eq!(state, base.join("job_123.state"));
453    }
454
455    #[cfg(unix)]
456    #[test]
457    fn test_default_spool_dir_prefers_xdg_and_isolates_fallback_by_uid() {
458        assert_eq!(
459            default_spool_dir(
460                Some(OsStr::new("/run/user/1000")),
461                Path::new("/var/tmp"),
462                1000,
463            ),
464            PathBuf::from("/run/user/1000/ssh-mcp")
465        );
466        assert_eq!(
467            default_spool_dir(Some(OsStr::new("relative")), Path::new("/var/tmp"), 1000),
468            PathBuf::from("/var/tmp/ssh-mcp-1000")
469        );
470        assert_ne!(
471            default_spool_dir(None, Path::new("/tmp"), 1000),
472            default_spool_dir(None, Path::new("/tmp"), 1001)
473        );
474        assert_eq!(
475            default_spool_dir(None, Path::new("relative"), 1000),
476            PathBuf::from("/tmp/ssh-mcp-1000")
477        );
478    }
479
480    #[cfg(unix)]
481    #[tokio::test]
482    async fn test_ensure_dir_normalizes_owned_permissions() {
483        let tmp = tempfile::TempDir::new().expect("tempdir");
484        let base = tmp.path().join("spool");
485        std::fs::create_dir(&base).expect("create spool dir");
486        std::fs::set_permissions(&base, std::fs::Permissions::from_mode(0o755))
487            .expect("set initial permissions");
488
489        LocalLogSpooler::new(base.clone())
490            .ensure_dir()
491            .await
492            .expect("ensure_dir");
493
494        let mode = std::fs::metadata(base)
495            .expect("spool dir metadata")
496            .permissions()
497            .mode()
498            & 0o777;
499        assert_eq!(mode, 0o700);
500    }
501
502    #[cfg(unix)]
503    #[tokio::test]
504    async fn test_ensure_dir_rejects_symlink() {
505        let tmp = tempfile::TempDir::new().expect("tempdir");
506        let target = tmp.path().join("target");
507        let base = tmp.path().join("spool");
508        std::fs::create_dir(&target).expect("create target dir");
509        symlink(target, &base).expect("create spool symlink");
510
511        let error = LocalLogSpooler::new(base)
512            .ensure_dir()
513            .await
514            .expect_err("symlink must be rejected");
515        assert!(error.to_string().contains("symlink"));
516    }
517
518    #[cfg(unix)]
519    #[tokio::test]
520    async fn test_ensure_dir_rejects_wrong_owner_without_chmod() {
521        let tmp = tempfile::TempDir::new().expect("tempdir");
522        let base = tmp.path().join("spool");
523        std::fs::create_dir(&base).expect("create spool dir");
524        std::fs::set_permissions(&base, std::fs::Permissions::from_mode(0o755))
525            .expect("set initial permissions");
526        let actual_uid = std::fs::metadata(&base).expect("metadata").uid();
527        let spooler = LocalLogSpooler::new(base.clone());
528
529        let error = spooler
530            .ensure_dir_inner(Some(actual_uid ^ 1))
531            .await
532            .expect_err("wrong owner must be rejected");
533
534        assert!(error.to_string().contains("not owned"));
535        let mode = std::fs::metadata(base)
536            .expect("spool dir metadata")
537            .permissions()
538            .mode()
539            & 0o777;
540        assert_eq!(mode, 0o755);
541    }
542
543    #[cfg(unix)]
544    #[tokio::test]
545    async fn test_ensure_dir_rejects_writable_existing_directory() {
546        let tmp = tempfile::TempDir::new().expect("tempdir");
547        let base = tmp.path().join("spool");
548        std::fs::create_dir(&base).expect("create spool dir");
549        std::fs::set_permissions(&base, std::fs::Permissions::from_mode(0o770))
550            .expect("set initial permissions");
551
552        let error = LocalLogSpooler::new(base.clone())
553            .ensure_dir()
554            .await
555            .expect_err("writable spool dir must be rejected");
556
557        assert!(error.to_string().contains("writable"));
558        let mode = std::fs::metadata(base)
559            .expect("spool dir metadata")
560            .permissions()
561            .mode()
562            & 0o777;
563        assert_eq!(mode, 0o770);
564    }
565
566    #[test]
567    fn test_log_path_for_rejects_invalid_job_ids() {
568        let spooler = LocalLogSpooler::new(PathBuf::from("/tmp/ssh-mcp-test"));
569        for job_id in ["", "..", "/abs", "a/b", "a\\b", "job id", "job\n1"] {
570            assert!(spooler.log_path_for(job_id).is_err(), "job_id={job_id}");
571        }
572    }
573
574    #[tokio::test]
575    async fn test_cleanup_old_logs_removes_log_exit_and_state_files_only() {
576        let tmp = tempfile::TempDir::new().expect("tempdir");
577        let base = tmp.path().join("spool");
578        let spooler = LocalLogSpooler::new(base.clone());
579        spooler.ensure_dir().await.expect("ensure_dir");
580
581        let log = base.join("job_1.log");
582        let exit = base.join("job_1.exit");
583        let state = base.join("job_1.state");
584        let keep = base.join("job_1.tmp");
585        tokio::fs::write(&log, "hello\n").await.expect("write log");
586        tokio::fs::write(&exit, "0\n").await.expect("write exit");
587        tokio::fs::write(&state, "{}\n").await.expect("write state");
588        tokio::fs::write(&keep, "x\n").await.expect("write tmp");
589
590        // Avoid flakiness from tight timing windows by waiting until the newest file
591        // is safely older than the max_age threshold.
592        wait_until_older_than(&keep, Duration::from_millis(25)).await;
593        let removed = spooler
594            .cleanup_old_logs(Duration::from_millis(1))
595            .await
596            .expect("cleanup_old_logs");
597
598        assert!(removed >= 3, "expected to remove at least log+exit+state");
599        assert!(!log.exists(), "log should be removed");
600        assert!(!exit.exists(), "exit should be removed");
601        assert!(!state.exists(), "state should be removed");
602        assert!(keep.exists(), "non-log file should be kept");
603    }
604
605    #[tokio::test]
606    async fn test_persist_and_load_job_state_round_trip() {
607        let tmp = tempfile::TempDir::new().expect("tempdir");
608        let base = tmp.path().join("spool");
609        let spooler = LocalLogSpooler::new(base.clone());
610        spooler.ensure_dir().await.expect("ensure_dir");
611
612        let mut job = JobState::new_running(super::super::job::NewRunningJob {
613            job_id: "job_123".to_string(),
614            pid: 4242,
615            log_path: base.join("job_123.log"),
616            command: "wget https://example.test/file".to_string(),
617            connection_id: "test@localhost:22".to_string(),
618        });
619        job.mark_state_lost("stream_error");
620
621        spooler
622            .persist_job_state(&job)
623            .await
624            .expect("persist_job_state");
625
626        let loaded = spooler
627            .load_job_state("job_123")
628            .await
629            .expect("load_job_state")
630            .expect("job should exist");
631
632        assert_eq!(loaded.job_id, job.job_id);
633        assert_eq!(loaded.pid, job.pid);
634        assert_eq!(loaded.status, job.status);
635        assert_eq!(loaded.exit_code, job.exit_code);
636        assert_eq!(loaded.state_reason, job.state_reason);
637        assert_eq!(loaded.command, job.command);
638        assert_eq!(loaded.log_path, job.log_path);
639    }
640}