Skip to main content

nucleus/container/
lifecycle.rs

1use crate::container::{ContainerState, ContainerStateManager};
2use crate::error::{NucleusError, Result};
3use nix::sys::signal::{kill, Signal};
4use nix::unistd::Pid;
5use nix::unistd::Uid;
6use std::thread;
7use std::time::Duration;
8use tracing::{info, warn};
9
10/// Container lifecycle operations (stop, kill, delete)
11pub struct ContainerLifecycle;
12
13impl ContainerLifecycle {
14    fn ensure_container_access(state: &ContainerState) -> Result<()> {
15        let current_uid = Uid::effective().as_raw();
16        if current_uid == 0 || current_uid == state.creator_uid {
17            return Ok(());
18        }
19
20        Err(NucleusError::PermissionDenied(format!(
21            "container {} owned by UID {}, caller is UID {}",
22            state.id, state.creator_uid, current_uid
23        )))
24    }
25
26    /// Stop a container gracefully: SIGTERM, wait for timeout, then SIGKILL
27    pub fn stop(state: &ContainerState, timeout_secs: u64) -> Result<()> {
28        Self::ensure_container_access(state)?;
29
30        if !state.is_running() {
31            info!("Container {} is already stopped", state.id);
32            return Ok(());
33        }
34
35        let pid = Pid::from_raw(state.pid as i32);
36
37        // Verify PID is still alive before sending signal.
38        // kill(pid, None) sends signal 0 – a no-op that returns ESRCH if the
39        // PID doesn't exist, protecting against PID recycling TOCTOU races.
40        if let Err(e) = kill(pid, None) {
41            if e == nix::errno::Errno::ESRCH {
42                info!("Process already exited");
43                return Ok(());
44            }
45        }
46
47        // Send SIGTERM
48        info!(
49            "Sending SIGTERM to container {} (PID {})",
50            state.id, state.pid
51        );
52        if let Err(e) = kill(pid, Signal::SIGTERM) {
53            if e == nix::errno::Errno::ESRCH {
54                info!("Process already exited");
55                return Ok(());
56            }
57            return Err(NucleusError::ExecError(format!(
58                "Failed to send SIGTERM: {}",
59                e
60            )));
61        }
62
63        // Wait for process to exit
64        let poll_interval = Duration::from_millis(100);
65        let deadline = Duration::from_secs(timeout_secs);
66        let mut elapsed = Duration::ZERO;
67
68        while elapsed < deadline {
69            if !state.is_running() {
70                info!("Container {} stopped gracefully", state.id);
71                return Ok(());
72            }
73            thread::sleep(poll_interval);
74            elapsed += poll_interval;
75        }
76
77        // Force kill
78        warn!(
79            "Container {} did not stop after {}s, sending SIGKILL",
80            state.id, timeout_secs
81        );
82        if let Err(e) = kill(pid, Signal::SIGKILL) {
83            if e == nix::errno::Errno::ESRCH {
84                return Ok(());
85            }
86            return Err(NucleusError::ExecError(format!(
87                "Failed to send SIGKILL: {}",
88                e
89            )));
90        }
91
92        Ok(())
93    }
94
95    /// Send an arbitrary signal to a container
96    pub fn kill_container(state: &ContainerState, signal: Signal) -> Result<()> {
97        Self::ensure_container_access(state)?;
98
99        if !state.is_running() {
100            return Err(NucleusError::ContainerNotRunning(format!(
101                "Container {} is not running",
102                state.id
103            )));
104        }
105
106        let pid = Pid::from_raw(state.pid as i32);
107        info!(
108            "Sending {:?} to container {} (PID {})",
109            signal, state.id, state.pid
110        );
111
112        kill(pid, signal).map_err(|e| {
113            NucleusError::ExecError(format!("Failed to send signal {:?}: {}", signal, e))
114        })?;
115
116        Ok(())
117    }
118
119    /// Remove a stopped container's state
120    pub fn remove(
121        state_mgr: &ContainerStateManager,
122        state: &ContainerState,
123        force: bool,
124    ) -> Result<()> {
125        Self::ensure_container_access(state)?;
126
127        if state.is_running() {
128            if force {
129                info!("Force removing running container {}", state.id);
130                Self::stop(state, 5)?;
131            } else {
132                return Err(NucleusError::ExecError(format!(
133                    "Container {} is still running. Stop it first or use --force",
134                    state.id
135                )));
136            }
137        }
138
139        // Clean up cgroup directory if present
140        if let Some(ref cgroup_path) = state.cgroup_path {
141            let cgroup = std::path::Path::new(cgroup_path);
142            if cgroup.exists() {
143                if let Err(e) = std::fs::remove_dir_all(cgroup) {
144                    warn!(
145                        "Failed to remove cgroup {}: {} (may still have processes)",
146                        cgroup_path, e
147                    );
148                } else {
149                    info!("Removed cgroup {}", cgroup_path);
150                }
151            }
152        }
153
154        if let Ok(overlay_dir) = state_mgr.rootfs_overlay_dir_path(&state.id) {
155            if overlay_dir.exists() {
156                if let Err(e) = std::fs::remove_dir_all(&overlay_dir) {
157                    warn!(
158                        "Failed to remove rootfs overlay directory {:?}: {}",
159                        overlay_dir, e
160                    );
161                }
162            }
163        }
164
165        state_mgr.delete_state(&state.id)?;
166        info!("Removed container {}", state.id);
167        Ok(())
168    }
169}
170
171/// Parse a signal name or number string into a Signal
172pub fn parse_signal(s: &str) -> Result<Signal> {
173    // Try numeric
174    if let Ok(num) = s.parse::<i32>() {
175        return Signal::try_from(num)
176            .map_err(|_| NucleusError::ConfigError(format!("Invalid signal number: {}", num)));
177    }
178
179    // Normalize: uppercase and strip optional "SIG" prefix
180    let upper = s.to_ascii_uppercase();
181    let normalized = upper.strip_prefix("SIG").unwrap_or(&upper);
182
183    match normalized {
184        "ABRT" | "IOT" => Ok(Signal::SIGABRT),
185        "ALRM" => Ok(Signal::SIGALRM),
186        "BUS" => Ok(Signal::SIGBUS),
187        "CHLD" | "CLD" => Ok(Signal::SIGCHLD),
188        "CONT" => Ok(Signal::SIGCONT),
189        "FPE" => Ok(Signal::SIGFPE),
190        "HUP" => Ok(Signal::SIGHUP),
191        "ILL" => Ok(Signal::SIGILL),
192        "INT" => Ok(Signal::SIGINT),
193        "IO" | "POLL" => Ok(Signal::SIGIO),
194        "KILL" => Ok(Signal::SIGKILL),
195        "PIPE" => Ok(Signal::SIGPIPE),
196        "PROF" => Ok(Signal::SIGPROF),
197        "PWR" => Ok(Signal::SIGPWR),
198        "QUIT" => Ok(Signal::SIGQUIT),
199        "SEGV" => Ok(Signal::SIGSEGV),
200        "STKFLT" => Ok(Signal::SIGSTKFLT),
201        "STOP" => Ok(Signal::SIGSTOP),
202        "SYS" => Ok(Signal::SIGSYS),
203        "TERM" => Ok(Signal::SIGTERM),
204        "TRAP" => Ok(Signal::SIGTRAP),
205        "TSTP" => Ok(Signal::SIGTSTP),
206        "TTIN" => Ok(Signal::SIGTTIN),
207        "TTOU" => Ok(Signal::SIGTTOU),
208        "URG" => Ok(Signal::SIGURG),
209        "USR1" => Ok(Signal::SIGUSR1),
210        "USR2" => Ok(Signal::SIGUSR2),
211        "VTALRM" => Ok(Signal::SIGVTALRM),
212        "WINCH" => Ok(Signal::SIGWINCH),
213        "XCPU" => Ok(Signal::SIGXCPU),
214        "XFSZ" => Ok(Signal::SIGXFSZ),
215        _ => Err(NucleusError::ConfigError(format!("Unknown signal: {}", s))),
216    }
217}
218
219#[cfg(test)]
220mod tests {
221    use super::*;
222    use crate::container::ContainerStateParams;
223
224    #[test]
225    fn test_parse_signal_by_name() {
226        assert_eq!(parse_signal("TERM").unwrap(), Signal::SIGTERM);
227        assert_eq!(parse_signal("SIGTERM").unwrap(), Signal::SIGTERM);
228        assert_eq!(parse_signal("KILL").unwrap(), Signal::SIGKILL);
229        assert_eq!(parse_signal("SIGKILL").unwrap(), Signal::SIGKILL);
230        assert_eq!(parse_signal("INT").unwrap(), Signal::SIGINT);
231        assert_eq!(parse_signal("HUP").unwrap(), Signal::SIGHUP);
232    }
233
234    #[test]
235    fn test_parse_signal_by_number() {
236        assert_eq!(parse_signal("15").unwrap(), Signal::SIGTERM);
237        assert_eq!(parse_signal("9").unwrap(), Signal::SIGKILL);
238        assert_eq!(parse_signal("2").unwrap(), Signal::SIGINT);
239    }
240
241    #[test]
242    fn test_parse_signal_case_insensitive() {
243        assert_eq!(parse_signal("term").unwrap(), Signal::SIGTERM);
244        assert_eq!(parse_signal("sigterm").unwrap(), Signal::SIGTERM);
245        assert_eq!(parse_signal("Term").unwrap(), Signal::SIGTERM);
246    }
247
248    #[test]
249    fn test_parse_signal_all_standard_names() {
250        let cases = vec![
251            ("ABRT", Signal::SIGABRT),
252            ("IOT", Signal::SIGABRT),
253            ("ALRM", Signal::SIGALRM),
254            ("BUS", Signal::SIGBUS),
255            ("CHLD", Signal::SIGCHLD),
256            ("CLD", Signal::SIGCHLD),
257            ("FPE", Signal::SIGFPE),
258            ("ILL", Signal::SIGILL),
259            ("IO", Signal::SIGIO),
260            ("POLL", Signal::SIGIO),
261            ("PIPE", Signal::SIGPIPE),
262            ("PROF", Signal::SIGPROF),
263            ("PWR", Signal::SIGPWR),
264            ("SEGV", Signal::SIGSEGV),
265            ("STKFLT", Signal::SIGSTKFLT),
266            ("SYS", Signal::SIGSYS),
267            ("TRAP", Signal::SIGTRAP),
268            ("TSTP", Signal::SIGTSTP),
269            ("TTIN", Signal::SIGTTIN),
270            ("TTOU", Signal::SIGTTOU),
271            ("URG", Signal::SIGURG),
272            ("VTALRM", Signal::SIGVTALRM),
273            ("WINCH", Signal::SIGWINCH),
274            ("XCPU", Signal::SIGXCPU),
275            ("XFSZ", Signal::SIGXFSZ),
276        ];
277        for (name, expected) in cases {
278            assert_eq!(
279                parse_signal(name).unwrap(),
280                expected,
281                "parse_signal({name}) failed"
282            );
283            // Also with SIG prefix
284            let prefixed = format!("SIG{name}");
285            assert_eq!(
286                parse_signal(&prefixed).unwrap(),
287                expected,
288                "parse_signal({prefixed}) failed"
289            );
290        }
291    }
292
293    #[test]
294    fn test_parse_signal_invalid() {
295        assert!(parse_signal("INVALID").is_err());
296        assert!(parse_signal("999").is_err());
297    }
298
299    #[test]
300    fn test_access_check_owner_allowed() {
301        let uid = Uid::effective().as_raw();
302        let state = ContainerState::new(ContainerStateParams {
303            id: "testid".to_string(),
304            name: "testname".to_string(),
305            pid: 12345,
306            command: vec!["/bin/true".to_string()],
307            memory_limit: None,
308            cpu_limit: None,
309            using_gvisor: false,
310            rootless: true,
311            cgroup_path: None,
312            process_uid: 0,
313            process_gid: 0,
314            additional_gids: Vec::new(),
315        });
316        // Override creator to match current caller for this test.
317        let mut state = state;
318        state.creator_uid = uid;
319        assert!(ContainerLifecycle::ensure_container_access(&state).is_ok());
320    }
321}