use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use crate::phase::ProcessPhase;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
pub enum ProcessSignal {
Sighup,
Sigterm,
Sigkill,
Sigusr1,
Sigusr2,
Sigstop,
Sigcont,
}
impl ProcessSignal {
pub const ALL: [Self; 7] = [
Self::Sighup,
Self::Sigterm,
Self::Sigkill,
Self::Sigusr1,
Self::Sigusr2,
Self::Sigstop,
Self::Sigcont,
];
pub const fn as_str(self) -> &'static str {
match self {
Self::Sighup => "SIGHUP",
Self::Sigterm => "SIGTERM",
Self::Sigkill => "SIGKILL",
Self::Sigusr1 => "SIGUSR1",
Self::Sigusr2 => "SIGUSR2",
Self::Sigstop => "SIGSTOP",
Self::Sigcont => "SIGCONT",
}
}
pub const fn short_str(self) -> &'static str {
match self {
Self::Sighup => "HUP",
Self::Sigterm => "TERM",
Self::Sigkill => "KILL",
Self::Sigusr1 => "USR1",
Self::Sigusr2 => "USR2",
Self::Sigstop => "STOP",
Self::Sigcont => "CONT",
}
}
}
impl std::fmt::Display for ProcessSignal {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
impl std::str::FromStr for ProcessSignal {
type Err = UnknownSignal;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let upper = s.to_ascii_uppercase();
for sig in Self::ALL {
if upper == sig.as_str() || upper == sig.short_str() {
return Ok(sig);
}
}
Err(UnknownSignal(upper))
}
}
#[derive(Debug, thiserror::Error)]
#[error("unknown signal: {0}")]
pub struct UnknownSignal(pub String);
#[derive(
Clone,
Copy,
Debug,
PartialEq,
Eq,
Hash,
Serialize,
Deserialize,
JsonSchema,
Default,
tatara_lisp::DeriveClosedSet,
)]
#[serde(rename_all = "PascalCase")]
#[closed_set(via = "as_str", generate_unknown, display)]
pub enum SighupStrategy {
#[default]
Reconverge,
Restart,
Noop,
}
impl SighupStrategy {
pub const ALL: [Self; 3] = [Self::Reconverge, Self::Restart, Self::Noop];
pub const fn as_str(self) -> &'static str {
match self {
Self::Reconverge => "Reconverge",
Self::Restart => "Restart",
Self::Noop => "Noop",
}
}
pub const fn sighup_target(self) -> Option<ProcessPhase> {
match self {
Self::Reconverge => Some(ProcessPhase::Reconverging),
Self::Restart => Some(ProcessPhase::Exiting),
Self::Noop => None,
}
}
}
#[cfg(test)]
mod tests {
use super::{ProcessSignal, SighupStrategy, UnknownSighupStrategy};
use crate::phase::ProcessPhase;
use std::str::FromStr;
#[test]
fn all_signals_roundtrip_canonical() {
for sig in ProcessSignal::ALL {
assert_eq!(ProcessSignal::from_str(sig.as_str()).unwrap(), sig);
}
}
#[test]
fn all_signals_roundtrip_short() {
for sig in ProcessSignal::ALL {
assert_eq!(ProcessSignal::from_str(sig.short_str()).unwrap(), sig);
}
}
#[test]
fn short_str_strips_sig_prefix() {
for sig in ProcessSignal::ALL {
assert_eq!(sig.as_str().strip_prefix("SIG"), Some(sig.short_str()));
}
}
#[test]
fn all_is_unique_and_complete() {
let mut seen = std::collections::HashSet::new();
for sig in ProcessSignal::ALL {
assert!(seen.insert(sig), "duplicate variant in ALL: {sig:?}");
}
assert_eq!(seen.len(), ProcessSignal::ALL.len());
}
#[test]
fn lowercase_short_form_accepted() {
assert_eq!(
ProcessSignal::from_str("hup").unwrap(),
ProcessSignal::Sighup
);
assert_eq!(
ProcessSignal::from_str("term").unwrap(),
ProcessSignal::Sigterm
);
}
#[test]
fn unknown_errors() {
let err = ProcessSignal::from_str("sigfoo").unwrap_err();
assert_eq!(err.0, "SIGFOO");
}
#[test]
fn sighup_strategy_is_well_formed_closed_set() {
tatara_lisp::assert_closed_set_well_formed::<SighupStrategy>();
}
#[test]
fn sighup_strategy_as_str_matches_serde() {
for strat in SighupStrategy::ALL {
let serialized = serde_json::to_string(&strat)
.expect("SighupStrategy serializes")
.trim_matches('"')
.to_string();
assert_eq!(
strat.as_str(),
serialized,
"as_str() must match serde output for {strat:?}",
);
}
}
#[test]
fn sighup_strategy_display_matches_as_str() {
for strat in SighupStrategy::ALL {
assert_eq!(strat.to_string(), strat.as_str());
}
}
#[test]
fn unknown_sighup_strategy_errors() {
for bad in ["reconverge", "RESTART", "Suspend", "noop "] {
let err = SighupStrategy::from_str(bad).unwrap_err();
let UnknownSighupStrategy(payload) = &err;
assert_eq!(payload, bad, "error payload should echo input verbatim");
}
}
#[test]
fn sighup_target_truth_table() {
assert_eq!(
SighupStrategy::Reconverge.sighup_target(),
Some(ProcessPhase::Reconverging)
);
assert_eq!(
SighupStrategy::Restart.sighup_target(),
Some(ProcessPhase::Exiting)
);
assert_eq!(SighupStrategy::Noop.sighup_target(), None);
}
#[test]
fn sighup_target_projects_only_to_legal_sighup_transitions() {
for strat in SighupStrategy::ALL {
if let Some(target) = strat.sighup_target() {
let reachable_from_running = ProcessPhase::Running.can_transition_to(target);
let reachable_from_attested = ProcessPhase::Attested.can_transition_to(target);
assert!(
reachable_from_running || reachable_from_attested,
"{strat:?}.sighup_target() = {target:?} must be reachable from \
Running or Attested via can_transition_to",
);
}
}
}
#[test]
fn sighup_target_projection_is_injective() {
let mut seen = std::collections::HashSet::new();
for strat in SighupStrategy::ALL {
if let Some(target) = strat.sighup_target() {
assert!(
seen.insert(target),
"two variants project to the same ProcessPhase: {target:?}",
);
}
}
}
}