Skip to main content

pitchfork_cli/
procs.rs

1use crate::Result;
2#[cfg(unix)]
3use crate::settings::settings;
4use miette::IntoDiagnostic;
5use once_cell::sync::Lazy;
6use std::collections::HashMap;
7use std::sync::Mutex;
8use sysinfo::ProcessesToUpdate;
9
10/// Map from parent PID to its child PIDs.
11type ParentToChildren = HashMap<u32, Vec<u32>>;
12
13/// Map from PID to process name and optional executable path.
14type ProcessNames = HashMap<u32, (String, Option<String>)>;
15
16pub struct Procs {
17    system: Mutex<sysinfo::System>,
18}
19
20pub static PROCS: Lazy<Procs> = Lazy::new(Procs::new);
21
22impl Default for Procs {
23    fn default() -> Self {
24        Self::new()
25    }
26}
27
28impl Procs {
29    pub fn new() -> Self {
30        // IMPORTANT: Do NOT call refresh_processes() or System::new_all() here.
31        //
32        // Both refresh the state of every process in the system, which takes
33        // ~500ms on a typical machine. Since PROCS is a Lazy static, the first
34        // access triggers this constructor — and `pitchfork cd` (which only
35        // needs to check if the supervisor PID is alive) would block for that
36        // duration on every directory change.
37        //
38        // See https://github.com/jdx/pitchfork/discussions/439
39        //
40        // Callers that need process info must call refresh_pids() (for specific
41        // PIDs) or refresh_processes() (for full-system stats) explicitly.
42        Self {
43            system: Mutex::new(sysinfo::System::new()),
44        }
45    }
46
47    fn lock_system(&self) -> std::sync::MutexGuard<'_, sysinfo::System> {
48        self.system.lock().unwrap_or_else(|poisoned| {
49            warn!("System mutex was poisoned, recovering");
50            poisoned.into_inner()
51        })
52    }
53
54    pub fn title(&self, pid: u32) -> Option<String> {
55        self.lock_system()
56            .process(sysinfo::Pid::from_u32(pid))
57            .map(|p| p.name().to_string_lossy().to_string())
58    }
59
60    pub fn is_running(&self, pid: u32) -> bool {
61        // Use kill(pid, 0) on Unix for an O(1) liveness check that does not
62        // depend on the process cache being populated. This avoids the need
63        // for a full process refresh just to check a single PID.
64        // ESRCH = process does not exist; EPERM = process exists but owned
65        // by another user (still "running" from our perspective).
66        #[cfg(unix)]
67        {
68            unsafe {
69                if libc::kill(pid as i32, 0) == 0 {
70                    return true;
71                }
72                std::io::Error::last_os_error().raw_os_error() != Some(libc::ESRCH)
73            }
74        }
75        #[cfg(not(unix))]
76        {
77            self.refresh_pids(&[pid]);
78            self.lock_system()
79                .process(sysinfo::Pid::from_u32(pid))
80                .is_some()
81        }
82    }
83
84    /// Walk the /proc tree to find all descendant PIDs.
85    /// Kept for diagnostics/status display; no longer used in the kill path.
86    #[allow(dead_code)]
87    pub fn all_children(&self, pid: u32) -> Vec<u32> {
88        let system = self.lock_system();
89        let all = system.processes();
90        let mut children = vec![];
91        for (child_pid, process) in all {
92            let mut process = process;
93            while let Some(parent) = process.parent() {
94                if parent == sysinfo::Pid::from_u32(pid) {
95                    children.push(child_pid.as_u32());
96                    break;
97                }
98                match system.process(parent) {
99                    Some(p) => process = p,
100                    None => break,
101                }
102            }
103        }
104        children
105    }
106    /// Collect minimal process tree information in a single lock.
107    ///
108    /// Returns a map of parent PID → child PIDs and a map of PID → (name, exe).
109    /// This avoids repeated mutex locking when traversing deep trees.
110    pub fn collect_process_tree_info(&self) -> (ParentToChildren, ProcessNames) {
111        let system = self.lock_system();
112        let all = system.processes();
113        let mut parent_to_children: ParentToChildren = HashMap::new();
114        let mut process_info: ProcessNames = HashMap::new();
115
116        for (pid, proc) in all {
117            let pid_u32 = pid.as_u32();
118            process_info.insert(
119                pid_u32,
120                (
121                    proc.name().to_string_lossy().to_string(),
122                    proc.exe().map(|e| e.to_string_lossy().to_string()),
123                ),
124            );
125
126            if let Some(ppid) = proc.parent() {
127                parent_to_children
128                    .entry(ppid.as_u32())
129                    .or_default()
130                    .push(pid_u32);
131            }
132        }
133
134        (parent_to_children, process_info)
135    }
136    pub async fn kill_process_group_async(
137        &self,
138        pid: u32,
139        stop_signal: i32,
140        stop_timeout: Option<std::time::Duration>,
141    ) -> Result<bool> {
142        tokio::task::spawn_blocking(move || {
143            PROCS.kill_process_group(pid, stop_signal, stop_timeout)
144        })
145        .await
146        .into_diagnostic()?
147    }
148
149    /// Kill an entire process group with graceful shutdown strategy:
150    /// 1. Send the configured stop signal to the process group (-pgid) and wait up to ~3s
151    /// 2. If any processes remain, send SIGKILL to the group
152    ///
153    /// Since daemons are spawned with setsid(), the daemon PID == PGID,
154    /// so this atomically signals all descendant processes.
155    ///
156    /// Returns `Err` if the signal could not be sent (e.g. permission denied).
157    #[cfg(unix)]
158    fn kill_process_group(
159        &self,
160        pid: u32,
161        stop_signal: i32,
162        stop_timeout: Option<std::time::Duration>,
163    ) -> Result<bool> {
164        let pgid = pid as i32;
165        let signal_name = signal_name(stop_signal);
166
167        debug!("killing process group {pgid} with {signal_name}");
168
169        // Send the stop signal to the entire process group.
170        // killpg sends to all processes in the group atomically.
171        // We intentionally skip the zombie check here because the leader may be
172        // a zombie while children in the group are still running.
173        let ret = unsafe { libc::killpg(pgid, stop_signal) };
174        if ret == -1 {
175            let err = std::io::Error::last_os_error();
176            if err.raw_os_error() == Some(libc::ESRCH) {
177                debug!("process group {pgid} no longer exists");
178                return Ok(false);
179            }
180            if err.raw_os_error() == Some(libc::EPERM) {
181                return Err(miette::miette!(
182                    "failed to send {signal_name} to process group {pgid}: permission denied"
183                ));
184            }
185            warn!("failed to send {signal_name} to process group {pgid}: {err}");
186        }
187
188        // Wait for graceful shutdown: fast initial check then slower polling.
189        // Per-daemon timeout overrides the global setting.
190        let stop_timeout = stop_timeout.unwrap_or_else(|| settings().supervisor_stop_timeout());
191        let fast_ms = 10u64;
192        let slow_ms = 50u64;
193        let total_ms = stop_timeout.as_millis().max(1) as u64;
194        let fast_count = ((total_ms / fast_ms) as usize).min(10);
195        let fast_total_ms = fast_ms * fast_count as u64;
196        let remaining_ms = total_ms.saturating_sub(fast_total_ms);
197        let slow_count = (remaining_ms / slow_ms) as usize;
198
199        let fast_checks =
200            std::iter::repeat_n(std::time::Duration::from_millis(fast_ms), fast_count);
201        let slow_checks =
202            std::iter::repeat_n(std::time::Duration::from_millis(slow_ms), slow_count);
203        let mut elapsed_ms = 0u64;
204
205        for sleep_duration in fast_checks.chain(slow_checks) {
206            std::thread::sleep(sleep_duration);
207            self.refresh_pids(&[pid]);
208            elapsed_ms += sleep_duration.as_millis() as u64;
209            if self.is_terminated_or_zombie(sysinfo::Pid::from_u32(pid)) {
210                debug!("process group {pgid} terminated after {signal_name} ({elapsed_ms} ms)",);
211                return Ok(true);
212            }
213        }
214
215        // SIGKILL the entire process group as last resort
216        warn!(
217            "process group {pgid} did not respond to {signal_name} after {}ms, sending SIGKILL",
218            stop_timeout.as_millis()
219        );
220        let ret = unsafe { libc::killpg(pgid, libc::SIGKILL) };
221        if ret == -1 {
222            let err = std::io::Error::last_os_error();
223            if err.raw_os_error() != Some(libc::ESRCH) {
224                warn!("failed to send SIGKILL to process group {pgid}: {err}");
225            }
226        }
227
228        // Brief wait for SIGKILL to take effect
229        std::thread::sleep(std::time::Duration::from_millis(100));
230        Ok(true)
231    }
232
233    #[cfg(not(unix))]
234    fn kill_process_group(
235        &self,
236        pid: u32,
237        _stop_signal: i32,
238        _stop_timeout: Option<std::time::Duration>,
239    ) -> Result<bool> {
240        self.kill(pid, 0, None)
241    }
242
243    pub async fn kill_async(
244        &self,
245        pid: u32,
246        stop_signal: i32,
247        stop_timeout: Option<std::time::Duration>,
248    ) -> Result<bool> {
249        tokio::task::spawn_blocking(move || PROCS.kill(pid, stop_signal, stop_timeout))
250            .await
251            .into_diagnostic()?
252    }
253
254    /// Kill a process with graceful shutdown strategy:
255    /// 1. Send the configured stop signal and wait up to ~3s (10ms intervals for first 100ms, then 50ms intervals)
256    /// 2. If still running, send SIGKILL to force termination
257    ///
258    /// This ensures fast-exiting processes don't wait unnecessarily,
259    /// while stubborn processes eventually get forcefully terminated.
260    ///
261    /// Returns `Err` if the signal could not be sent (e.g. permission denied
262    /// when targeting a process owned by another user/root).
263    fn kill(
264        &self,
265        pid: u32,
266        stop_signal: i32,
267        stop_timeout: Option<std::time::Duration>,
268    ) -> Result<bool> {
269        let sysinfo_pid = sysinfo::Pid::from_u32(pid);
270
271        debug!("killing process {pid}");
272
273        #[cfg(windows)]
274        {
275            let _ = (stop_signal, stop_timeout);
276            self.refresh_pids(&[pid]);
277            if let Some(process) = self.lock_system().process(sysinfo_pid) {
278                process.kill();
279                process.wait();
280            }
281            Ok(true)
282        }
283
284        #[cfg(unix)]
285        {
286            let signal_name = signal_name(stop_signal);
287            // Send stop signal for graceful shutdown using libc::kill directly
288            // so we can distinguish EPERM (permission denied) from ESRCH
289            // (process already gone — possible in a narrow race window).
290            debug!("sending {signal_name} to process {pid}");
291            let ret = unsafe { libc::kill(pid as i32, stop_signal) };
292            if ret == -1 {
293                let err = std::io::Error::last_os_error();
294                if err.raw_os_error() == Some(libc::ESRCH) {
295                    debug!("process {pid} no longer exists");
296                    return Ok(false);
297                }
298                if err.raw_os_error() == Some(libc::EPERM) {
299                    return Err(miette::miette!(
300                        "failed to send {signal_name} to process {pid}: permission denied"
301                    ));
302                }
303                return Err(miette::miette!(
304                    "failed to send {signal_name} to process {pid}: {err}"
305                ));
306            }
307
308            // Fast check: 10ms intervals, then slower 50ms polling for stop_timeout.
309            // Per-daemon timeout overrides the global setting.
310            let stop_timeout = stop_timeout.unwrap_or_else(|| settings().supervisor_stop_timeout());
311            let fast_ms = 10u64;
312            let slow_ms = 50u64;
313            let total_ms = stop_timeout.as_millis().max(1) as u64;
314            let fast_count = ((total_ms / fast_ms) as usize).min(10);
315            let fast_total_ms = fast_ms * fast_count as u64;
316            let remaining_ms = total_ms.saturating_sub(fast_total_ms);
317            let slow_count = (remaining_ms / slow_ms) as usize;
318
319            for i in 0..fast_count {
320                std::thread::sleep(std::time::Duration::from_millis(fast_ms));
321                self.refresh_pids(&[pid]);
322                if self.is_terminated_or_zombie(sysinfo_pid) {
323                    debug!(
324                        "process {pid} terminated after {signal_name} ({} ms)",
325                        (i + 1) * fast_ms as usize
326                    );
327                    return Ok(true);
328                }
329            }
330
331            // Slower check: 50ms intervals for the remainder of stop_timeout
332            for i in 0..slow_count {
333                std::thread::sleep(std::time::Duration::from_millis(slow_ms));
334                self.refresh_pids(&[pid]);
335                if self.is_terminated_or_zombie(sysinfo_pid) {
336                    debug!(
337                        "process {pid} terminated after {signal_name} ({} ms)",
338                        fast_total_ms + (i + 1) as u64 * slow_ms
339                    );
340                    return Ok(true);
341                }
342            }
343
344            // SIGKILL as last resort after stop_timeout
345            warn!(
346                "process {pid} did not respond to {signal_name} after {}ms, sending SIGKILL",
347                stop_timeout.as_millis()
348            );
349            let ret = unsafe { libc::kill(pid as i32, libc::SIGKILL) };
350            if ret == -1 {
351                let err = std::io::Error::last_os_error();
352                if err.raw_os_error() != Some(libc::ESRCH) {
353                    warn!("failed to send SIGKILL to process {pid}: {err}");
354                }
355            }
356
357            // Brief wait for SIGKILL to take effect
358            std::thread::sleep(std::time::Duration::from_millis(100));
359            Ok(true)
360        }
361    }
362
363    /// Check if a process is terminated or is a zombie.
364    /// On Linux, zombie processes still have /proc/[pid] entries but are effectively dead.
365    /// This prevents unnecessary signal escalation for processes that have already exited.
366    fn is_terminated_or_zombie(&self, sysinfo_pid: sysinfo::Pid) -> bool {
367        let system = self.lock_system();
368        match system.process(sysinfo_pid) {
369            None => true,
370            Some(process) => {
371                #[cfg(unix)]
372                {
373                    matches!(process.status(), sysinfo::ProcessStatus::Zombie)
374                }
375                #[cfg(not(unix))]
376                {
377                    let _ = process;
378                    false
379                }
380            }
381        }
382    }
383
384    pub(crate) fn refresh_processes(&self) {
385        self.lock_system()
386            .refresh_processes(ProcessesToUpdate::All, true);
387    }
388
389    /// Refresh only specific PIDs instead of all processes.
390    /// More efficient when you only need to check a small set of known PIDs.
391    pub(crate) fn refresh_pids(&self, pids: &[u32]) {
392        let sysinfo_pids: Vec<sysinfo::Pid> =
393            pids.iter().map(|p| sysinfo::Pid::from_u32(*p)).collect();
394        self.lock_system()
395            .refresh_processes(ProcessesToUpdate::Some(&sysinfo_pids), true);
396    }
397
398    /// Get aggregated stats for multiple process trees in a single pass.
399    ///
400    /// Builds the parent→children map once (O(N)) and then BFS-es from each
401    /// root PID (O(D_i) per daemon). Total cost is O(N + ΣD_i) instead of
402    /// O(D × N) when collecting stats for each daemon separately.
403    pub fn get_batch_group_stats(&self, pids: &[u32]) -> Vec<(u32, Option<ProcessStats>)> {
404        if pids.is_empty() {
405            return Vec::new();
406        }
407
408        let system = self.lock_system();
409        let processes = system.processes();
410
411        let now = std::time::SystemTime::now()
412            .duration_since(std::time::UNIX_EPOCH)
413            .map(|d| d.as_secs())
414            .unwrap_or(0);
415
416        // Build parent → children map once for all daemons
417        let mut children_map: std::collections::HashMap<sysinfo::Pid, Vec<sysinfo::Pid>> =
418            std::collections::HashMap::new();
419        for (child_pid, child) in processes {
420            // Skip Linux userland threads: they report the same memory as their parent process,
421            // so including them would cause massive double-counting.
422            if child.thread_kind().is_some() {
423                continue;
424            }
425            if let Some(ppid) = child.parent() {
426                children_map.entry(ppid).or_default().push(*child_pid);
427            }
428        }
429
430        pids.iter()
431            .map(|&pid| {
432                let root_pid = sysinfo::Pid::from_u32(pid);
433                let Some(root) = processes.get(&root_pid) else {
434                    return (pid, None);
435                };
436
437                let root_disk = root.disk_usage();
438                let mut stats = ProcessStats {
439                    cpu_percent: root.cpu_usage(),
440                    memory_bytes: root.memory(),
441                    uptime_secs: now.saturating_sub(root.start_time()),
442                    disk_read_bytes: root_disk.read_bytes,
443                    disk_write_bytes: root_disk.written_bytes,
444                };
445
446                // BFS from root_pid to find all descendants
447                let mut queue = std::collections::VecDeque::new();
448                if let Some(direct_children) = children_map.get(&root_pid) {
449                    queue.extend(direct_children);
450                }
451                while let Some(child_pid) = queue.pop_front() {
452                    if let Some(child) = processes.get(&child_pid) {
453                        let disk = child.disk_usage();
454                        stats.cpu_percent += child.cpu_usage();
455                        stats.memory_bytes += child.memory();
456                        stats.disk_read_bytes += disk.read_bytes;
457                        stats.disk_write_bytes += disk.written_bytes;
458                    }
459                    if let Some(grandchildren) = children_map.get(&child_pid) {
460                        queue.extend(grandchildren);
461                    }
462                }
463
464                (pid, Some(stats))
465            })
466            .collect()
467    }
468    /// Refresh the process tree, then call [`Self::get_batch_group_stats`].
469    ///
470    /// Convenience wrapper that guarantees `get_batch_group_stats` sees a
471    /// fresh snapshot of /proc (or its equivalent).  Returns a PID →
472    /// [`ProcessStats`] map so callers do not have to repeat the same
473    /// `filter_map`/`collect` boilerplate.
474    pub fn refresh_and_get_batch_stats(&self, pids: &[u32]) -> HashMap<u32, ProcessStats> {
475        self.refresh_processes();
476        self.get_batch_group_stats(pids)
477            .into_iter()
478            .filter_map(|(pid, stats)| stats.map(|s| (pid, s)))
479            .collect()
480    }
481
482    /// Get process-tree stats for multiple root PIDs, omitting roots that no longer exist.
483    pub fn get_batch_tree_stats_map(&self, pids: &[u32]) -> HashMap<u32, ProcessStats> {
484        self.get_batch_group_stats(pids)
485            .into_iter()
486            .filter_map(|(pid, stats)| stats.map(|stats| (pid, stats)))
487            .collect()
488    }
489
490    /// Get process-tree stats (cpu%, memory bytes, uptime secs, disk I/O) for a given root PID.
491    pub fn get_stats(&self, pid: u32) -> Option<ProcessStats> {
492        self.get_batch_group_stats(&[pid])
493            .into_iter()
494            .next()
495            .and_then(|(_, stats)| stats)
496    }
497
498    /// Get extended process information for a given PID
499    pub fn get_extended_stats(&self, pid: u32) -> Option<ExtendedProcessStats> {
500        let system = self.lock_system();
501        let processes = system.processes();
502        let root_pid = sysinfo::Pid::from_u32(pid);
503        let p = processes.get(&root_pid)?;
504
505        let now = std::time::SystemTime::now()
506            .duration_since(std::time::UNIX_EPOCH)
507            .map(|d| d.as_secs())
508            .unwrap_or(0);
509
510        let root_disk = p.disk_usage();
511        let mut aggregate_stats = ProcessStats {
512            cpu_percent: p.cpu_usage(),
513            memory_bytes: p.memory(),
514            uptime_secs: now.saturating_sub(p.start_time()),
515            disk_read_bytes: root_disk.read_bytes,
516            disk_write_bytes: root_disk.written_bytes,
517        };
518
519        let mut children_map: HashMap<sysinfo::Pid, Vec<sysinfo::Pid>> = HashMap::new();
520        for (child_pid, child) in processes {
521            if let Some(ppid) = child.parent() {
522                children_map.entry(ppid).or_default().push(*child_pid);
523            }
524        }
525
526        let mut queue = std::collections::VecDeque::new();
527        if let Some(direct_children) = children_map.get(&root_pid) {
528            queue.extend(direct_children);
529        }
530        while let Some(child_pid) = queue.pop_front() {
531            if let Some(child) = processes.get(&child_pid) {
532                let disk = child.disk_usage();
533                aggregate_stats.cpu_percent += child.cpu_usage();
534                aggregate_stats.memory_bytes += child.memory();
535                aggregate_stats.disk_read_bytes += disk.read_bytes;
536                aggregate_stats.disk_write_bytes += disk.written_bytes;
537            }
538            if let Some(grandchildren) = children_map.get(&child_pid) {
539                queue.extend(grandchildren);
540            }
541        }
542
543        Some(ExtendedProcessStats {
544            name: p.name().to_string_lossy().to_string(),
545            status: format!("{:?}", p.status()),
546            cpu_percent: aggregate_stats.cpu_percent,
547            memory_bytes: aggregate_stats.memory_bytes,
548            virtual_memory_bytes: p.virtual_memory(),
549            uptime_secs: aggregate_stats.uptime_secs,
550            thread_count: p.tasks().map(|t| t.len()).unwrap_or(0),
551        })
552    }
553}
554
555#[derive(Debug, Clone, Copy)]
556pub struct ProcessStats {
557    pub cpu_percent: f32,
558    pub memory_bytes: u64,
559    pub uptime_secs: u64,
560    pub disk_read_bytes: u64,
561    pub disk_write_bytes: u64,
562}
563
564impl ProcessStats {
565    pub fn memory_display(&self) -> String {
566        format_bytes(self.memory_bytes)
567    }
568
569    pub fn cpu_display(&self) -> String {
570        format!("{:.1}%", self.cpu_percent)
571    }
572
573    pub fn uptime_display(&self) -> String {
574        format_duration(self.uptime_secs)
575    }
576
577    pub fn disk_read_display(&self) -> String {
578        format_bytes_per_sec(self.disk_read_bytes)
579    }
580
581    pub fn disk_write_display(&self) -> String {
582        format_bytes_per_sec(self.disk_write_bytes)
583    }
584}
585
586#[derive(Debug, Clone)]
587pub struct ExtendedProcessStats {
588    pub name: String,
589    pub status: String,
590    pub cpu_percent: f32,
591    pub memory_bytes: u64,
592    pub virtual_memory_bytes: u64,
593    pub uptime_secs: u64,
594    pub thread_count: usize,
595}
596
597fn format_bytes(bytes: u64) -> String {
598    humanbyte::to_string(bytes, humanbyte::Format::IEC)
599}
600
601fn format_duration(secs: u64) -> String {
602    if secs < 60 {
603        format!("{secs}s")
604    } else if secs < 3600 {
605        format!("{}m {}s", secs / 60, secs % 60)
606    } else if secs < 86400 {
607        let hours = secs / 3600;
608        let mins = (secs % 3600) / 60;
609        format!("{hours}h {mins}m")
610    } else {
611        let days = secs / 86400;
612        let hours = (secs % 86400) / 3600;
613        format!("{days}d {hours}h")
614    }
615}
616
617fn format_bytes_per_sec(bytes: u64) -> String {
618    format!("{}/s", humanbyte::to_string(bytes, humanbyte::Format::IEC))
619}
620
621#[cfg(unix)]
622fn signal_name(sig: i32) -> &'static str {
623    match sig {
624        libc::SIGHUP => "SIGHUP",
625        libc::SIGINT => "SIGINT",
626        libc::SIGQUIT => "SIGQUIT",
627        libc::SIGTERM => "SIGTERM",
628        libc::SIGUSR1 => "SIGUSR1",
629        libc::SIGUSR2 => "SIGUSR2",
630        libc::SIGKILL => "SIGKILL",
631        _ => "UNKNOWN",
632    }
633}
634
635#[cfg(test)]
636mod format_tests {
637    use super::*;
638
639    #[test]
640    fn test_format_bytes() {
641        assert_eq!(format_bytes(512), "512 B");
642        assert_eq!(format_bytes(1024), "1.0 KiB");
643        assert_eq!(format_bytes(1536), "1.5 KiB");
644        assert_eq!(format_bytes(50 * 1024 * 1024), "50.0 MiB");
645        assert_eq!(format_bytes(3 * 1024 * 1024 * 1024), "3.0 GiB");
646        // rolls over past GiB instead of showing e.g. "1100.0GB"
647        assert_eq!(format_bytes(1100 * 1024 * 1024 * 1024), "1.1 TiB");
648    }
649
650    #[test]
651    fn test_format_bytes_per_sec() {
652        assert_eq!(format_bytes_per_sec(512), "512 B/s");
653        assert_eq!(format_bytes_per_sec(1536), "1.5 KiB/s");
654        assert_eq!(format_bytes_per_sec(2 * 1024 * 1024), "2.0 MiB/s");
655    }
656}
657
658#[cfg(all(test, unix))]
659mod tests {
660    use super::*;
661    use std::os::unix::process::CommandExt;
662    use std::process::{Child, Command, Stdio};
663    use std::time::{Duration, Instant};
664
665    struct ChildGuard(Child);
666
667    impl Drop for ChildGuard {
668        fn drop(&mut self) {
669            let pid = self.0.id() as i32;
670            // The test process is started in its own session, so PID == PGID.
671            let _ = unsafe { libc::killpg(pid, libc::SIGKILL) };
672            let _ = self.0.wait();
673        }
674    }
675
676    #[test]
677    fn get_stats_includes_descendant_rss() {
678        let mut command = Command::new("sh");
679        command
680            .args(["-c", "sleep 30 & wait"])
681            .stdin(Stdio::null())
682            .stdout(Stdio::null())
683            .stderr(Stdio::null());
684        unsafe {
685            command.pre_exec(|| {
686                if libc::setsid() == -1 {
687                    return Err(std::io::Error::last_os_error());
688                }
689                Ok(())
690            });
691        }
692
693        let parent = command.spawn().expect("failed to spawn process tree");
694        let parent_pid = parent.id();
695        let _parent = ChildGuard(parent);
696
697        let procs = Procs::new();
698        let deadline = Instant::now() + Duration::from_secs(5);
699        let mut child_pids = Vec::new();
700        while Instant::now() < deadline {
701            procs.refresh_processes();
702            child_pids = procs.all_children(parent_pid);
703            if !child_pids.is_empty() {
704                break;
705            }
706            std::thread::sleep(Duration::from_millis(50));
707        }
708        assert!(
709            !child_pids.is_empty(),
710            "test process tree did not appear under parent pid {parent_pid}"
711        );
712
713        procs.refresh_processes();
714        child_pids = procs.all_children(parent_pid);
715        assert!(
716            !child_pids.is_empty(),
717            "test process tree disappeared under parent pid {parent_pid}"
718        );
719        let root_pid = sysinfo::Pid::from_u32(parent_pid);
720        let direct_memory = {
721            let system = procs.lock_system();
722            system
723                .process(root_pid)
724                .expect("parent process should exist")
725                .memory()
726        };
727        let descendant_memory = {
728            let system = procs.lock_system();
729            child_pids
730                .iter()
731                .filter_map(|pid| system.process(sysinfo::Pid::from_u32(*pid)))
732                .map(|process| process.memory())
733                .sum::<u64>()
734        };
735        assert!(
736            descendant_memory > 0,
737            "descendants {child_pids:?} should have nonzero RSS"
738        );
739
740        let stats = procs
741            .get_stats(parent_pid)
742            .expect("parent process should have aggregate stats");
743
744        assert_eq!(
745            stats.memory_bytes,
746            direct_memory + descendant_memory,
747            "get_stats should include descendant RSS for parent pid {parent_pid}; \
748             descendants: {child_pids:?}, direct RSS: {direct_memory}, \
749             descendant RSS: {descendant_memory}, reported RSS: {}",
750            stats.memory_bytes
751        );
752    }
753}