Skip to main content

inferlab_runtime/
process_group.rs

1use crate::operation_bound::{OperationBound, Remaining};
2use serde::{Deserialize, Serialize};
3use std::process::{Child, Output};
4use std::thread;
5use std::time::Duration;
6use thiserror::Error;
7
8#[derive(Debug, Error)]
9pub enum ProcessGroupError {
10    #[error("local process-group identity requires non-zero identifiers")]
11    ZeroIdentity,
12    #[error("local process-group leader must equal its process group")]
13    LeaderMismatch,
14    #[error("process {pid} exited before its identity could be recorded")]
15    ExitedBeforeCapture { pid: u32 },
16    #[error("failed to reap process-group leader: {source}")]
17    Reap {
18        #[source]
19        source: std::io::Error,
20    },
21    #[error("failed to read {path}: {source}")]
22    StatRead {
23        path: String,
24        #[source]
25        source: std::io::Error,
26    },
27    #[error("invalid process stat for pid {pid}")]
28    InvalidStat { pid: u32 },
29    #[error("process stat for pid {pid} has no start time")]
30    MissingStartTime { pid: u32 },
31    #[error("invalid process start time for pid {pid}: {source}")]
32    InvalidStartTime {
33        pid: u32,
34        #[source]
35        source: std::num::ParseIntError,
36    },
37    #[error("process-group query exited with {status}: {stderr}")]
38    QueryExit {
39        status: std::process::ExitStatus,
40        stderr: String,
41    },
42    #[error("process-group cleanup command deadline expired")]
43    Deadline,
44    #[error("process-group cleanup command was interrupted")]
45    Interrupted,
46    #[error("process-group cleanup command failed to launch: {source}")]
47    Launch {
48        #[source]
49        source: std::io::Error,
50    },
51    #[error("process-group cleanup command failed: {source}")]
52    CommandIo {
53        #[source]
54        source: std::io::Error,
55    },
56    #[error("process-group cleanup command wait failed: {source}; cleanup verification: {cleanup}")]
57    WaitCleanup {
58        #[source]
59        source: std::io::Error,
60        cleanup: String,
61    },
62}
63
64#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
65#[serde(rename_all = "snake_case")]
66pub enum TerminationSignal {
67    Term,
68    Kill,
69}
70
71#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
72#[serde(deny_unknown_fields)]
73pub struct SignalEvidence {
74    pub signal: TerminationSignal,
75    pub process_group: u32,
76    pub exit_code: Option<i32>,
77    pub stderr: Option<String>,
78    pub error: Option<String>,
79}
80
81impl SignalEvidence {
82    pub fn succeeded(&self) -> bool {
83        self.error.is_none() && self.exit_code == Some(0)
84    }
85}
86
87#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
88pub struct LocalProcessGroup {
89    pub leader_pid: u32,
90    pub process_group: u32,
91    pub leader_start_time_ticks: u64,
92}
93
94#[derive(Clone, Copy, Debug, Eq, PartialEq)]
95pub enum VerifiedStatus {
96    Alive,
97    Exited,
98    Reused,
99    LeaderMissingWithMembers,
100}
101
102impl LocalProcessGroup {
103    pub const fn unverified(process_group: u32) -> Self {
104        Self {
105            leader_pid: process_group,
106            process_group,
107            leader_start_time_ticks: 0,
108        }
109    }
110
111    pub fn new(
112        leader_pid: u32,
113        process_group: u32,
114        leader_start_time_ticks: u64,
115    ) -> Result<Self, ProcessGroupError> {
116        if leader_pid == 0 || process_group == 0 {
117            return Err(ProcessGroupError::ZeroIdentity);
118        }
119        if leader_pid != process_group {
120            return Err(ProcessGroupError::LeaderMismatch);
121        }
122        Ok(Self {
123            leader_pid,
124            process_group,
125            leader_start_time_ticks,
126        })
127    }
128
129    pub fn capture_child(child: &Child) -> Result<Self, ProcessGroupError> {
130        let leader_pid = child.id();
131        let leader_start_time_ticks = process_start_time(leader_pid)?
132            .ok_or(ProcessGroupError::ExitedBeforeCapture { pid: leader_pid })?;
133        Self::new(leader_pid, leader_pid, leader_start_time_ticks)
134    }
135
136    pub fn identity_matches(&self) -> bool {
137        process_start_time(self.leader_pid)
138            .ok()
139            .flatten()
140            .is_some_and(|current| current == self.leader_start_time_ticks)
141    }
142
143    pub fn verified_status(
144        &self,
145        bound: &OperationBound,
146    ) -> Result<VerifiedStatus, ProcessGroupError> {
147        let leader_start = process_start_time(self.leader_pid)?;
148        match leader_start {
149            Some(actual) if actual != self.leader_start_time_ticks => Ok(VerifiedStatus::Reused),
150            Some(_) => {
151                let alive = self.has_live_members(bound)?;
152                Ok(if alive {
153                    VerifiedStatus::Alive
154                } else {
155                    VerifiedStatus::Exited
156                })
157            }
158            None => {
159                let alive = self.has_live_members(bound)?;
160                Ok(if alive {
161                    VerifiedStatus::LeaderMissingWithMembers
162                } else {
163                    VerifiedStatus::Exited
164                })
165            }
166        }
167    }
168
169    pub fn send_signal(&self, signal: TerminationSignal, bound: &OperationBound) -> SignalEvidence {
170        let signal_argument = match signal {
171            TerminationSignal::Term => "-TERM",
172            TerminationSignal::Kill => "-KILL",
173        };
174        let target = format!("-{}", self.process_group);
175        match cleanup_output(&["kill", signal_argument, "--", &target], bound) {
176            Ok(output) => SignalEvidence {
177                signal,
178                process_group: self.process_group,
179                exit_code: output.status.code(),
180                stderr: Some(String::from_utf8_lossy(&output.stderr).trim().to_owned()),
181                error: None,
182            },
183            Err(error) => SignalEvidence {
184                signal,
185                process_group: self.process_group,
186                exit_code: None,
187                stderr: None,
188                error: Some(error.to_string()),
189            },
190        }
191    }
192
193    pub fn wait_until_stopped(
194        &self,
195        mut child: Option<&mut Child>,
196        bound: &OperationBound,
197        poll_interval: Duration,
198    ) -> Result<bool, ProcessGroupError> {
199        loop {
200            if let Some(child) = child.as_deref_mut() {
201                child
202                    .try_wait()
203                    .map_err(|source| ProcessGroupError::Reap { source })?;
204            }
205            if bound.is_expired() {
206                return Ok(false);
207            }
208            match self.has_live_members(bound) {
209                Ok(false) => return Ok(true),
210                Ok(true) => {}
211                Err(_) if bound.is_expired() => return Ok(false),
212                Err(error) => return Err(error),
213            }
214            match bound.remaining() {
215                Remaining::Finite(remaining) => {
216                    thread::sleep(poll_interval.min(remaining));
217                }
218                Remaining::Expired => return Ok(false),
219                Remaining::Unbounded => thread::sleep(poll_interval),
220            }
221        }
222    }
223
224    pub fn has_live_members(&self, bound: &OperationBound) -> Result<bool, ProcessGroupError> {
225        Ok(!self.live_member_pids(bound)?.is_empty())
226    }
227
228    /// Members fork from the leader at or after its recorded start, so a
229    /// member predating the leader means the pgid was recycled and the
230    /// survivors are not the recorded cohort; the caller must not signal the
231    /// group. Returns the offending (pid, start ticks) pairs; empty means the
232    /// surviving members are consistent with the recorded leader. Defunct
233    /// members are dead already — they need no signal and do not participate.
234    /// A member exiting between enumeration and its stat read is skipped.
235    pub fn cohort_violations(
236        &self,
237        bound: &OperationBound,
238    ) -> Result<Vec<(u32, u64)>, ProcessGroupError> {
239        let mut violations = Vec::new();
240        for pid in self.live_member_pids(bound)? {
241            if let Some(start) = process_start_time(pid)?
242                && start < self.leader_start_time_ticks
243            {
244                violations.push((pid, start));
245            }
246        }
247        Ok(violations)
248    }
249
250    /// Live (non-defunct) member pids of the group.
251    fn live_member_pids(&self, bound: &OperationBound) -> Result<Vec<u32>, ProcessGroupError> {
252        let output = cleanup_output(&["ps", "-eo", "pid=,pgid=,stat="], bound)?;
253        if !output.status.success() {
254            return Err(ProcessGroupError::QueryExit {
255                status: output.status,
256                stderr: String::from_utf8_lossy(&output.stderr).trim().to_owned(),
257            });
258        }
259        let process_group = self.process_group.to_string();
260        Ok(String::from_utf8_lossy(&output.stdout)
261            .lines()
262            .filter_map(|line| {
263                let mut fields = line.split_whitespace();
264                let pid = fields.next()?;
265                let group = fields.next()?;
266                let state = fields.next()?;
267                Some((pid, group, state))
268            })
269            .filter(|(_, group, state)| *group == process_group && !state.starts_with('Z'))
270            .filter_map(|(pid, _, _)| pid.parse().ok())
271            .collect())
272    }
273}
274
275pub fn process_start_time(pid: u32) -> Result<Option<u64>, ProcessGroupError> {
276    let path = format!("/proc/{pid}/stat");
277    let stat = match std::fs::read_to_string(&path) {
278        Ok(stat) => stat,
279        Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
280        Err(source) => return Err(ProcessGroupError::StatRead { path, source }),
281    };
282    let command_end = stat
283        .rfind(')')
284        .ok_or(ProcessGroupError::InvalidStat { pid })?;
285    let start_time = stat[command_end + 1..]
286        .split_whitespace()
287        .nth(19)
288        .ok_or(ProcessGroupError::MissingStartTime { pid })?
289        .parse::<u64>()
290        .map_err(|source| ProcessGroupError::InvalidStartTime { pid, source })?;
291    Ok(Some(start_time))
292}
293
294fn cleanup_output(argv: &[&str], bound: &OperationBound) -> Result<Output, ProcessGroupError> {
295    match crate::container::run_cleanup_with_bound(argv, &[], None, None, bound, None) {
296        Ok(crate::container::BoundedWait::Exited {
297            status,
298            stdout,
299            stderr,
300        }) => Ok(Output {
301            status,
302            stdout,
303            stderr,
304        }),
305        Ok(crate::container::BoundedWait::Expired { .. }) => Err(ProcessGroupError::Deadline),
306        Ok(crate::container::BoundedWait::Interrupted { .. }) => {
307            Err(ProcessGroupError::Interrupted)
308        }
309        Err(crate::container::BoundedError::Launch(source)) => {
310            Err(ProcessGroupError::Launch { source })
311        }
312        Err(
313            crate::container::BoundedError::Stdin(error)
314            | crate::container::BoundedError::Wait(error),
315        ) => Err(ProcessGroupError::CommandIo { source: error }),
316        Err(crate::container::BoundedError::WaitCleanup {
317            source, cleanup, ..
318        }) => Err(ProcessGroupError::WaitCleanup {
319            source,
320            cleanup: cleanup.error.unwrap_or_else(|| {
321                if cleanup.verified {
322                    "verified"
323                } else {
324                    "unverified"
325                }
326                .to_owned()
327            }),
328        }),
329    }
330}