use std::sync::atomic::{AtomicU8, Ordering};
use tokio::sync::{broadcast, watch};
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum LifecycleState {
Created = 0,
Starting = 1,
Running = 2,
Draining = 3,
Stopped = 4,
Failed = 5,
}
impl LifecycleState {
fn from_u8(v: u8) -> Self {
match v {
0 => Self::Created,
1 => Self::Starting,
2 => Self::Running,
3 => Self::Draining,
4 => Self::Stopped,
5 => Self::Failed,
_ => Self::Failed,
}
}
pub fn is_terminal(self) -> bool {
matches!(self, Self::Stopped | Self::Failed)
}
}
impl std::fmt::Display for LifecycleState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Created => write!(f, "created"),
Self::Starting => write!(f, "starting"),
Self::Running => write!(f, "running"),
Self::Draining => write!(f, "draining"),
Self::Stopped => write!(f, "stopped"),
Self::Failed => write!(f, "failed"),
}
}
}
#[derive(Debug)]
pub(crate) struct Lifecycle {
state: AtomicU8,
ready_tx: watch::Sender<bool>,
terminal_tx: broadcast::Sender<()>,
}
impl Default for Lifecycle {
fn default() -> Self {
Self::new()
}
}
impl Lifecycle {
pub(crate) fn new() -> Self {
let (ready_tx, _) = watch::channel(false);
let (terminal_tx, _) = broadcast::channel(1);
Self {
state: AtomicU8::new(LifecycleState::Created as u8),
ready_tx,
terminal_tx,
}
}
pub(crate) fn state(&self) -> LifecycleState {
LifecycleState::from_u8(self.state.load(Ordering::Acquire))
}
pub(crate) fn start(&self) -> Result<(), crate::server::errors::ServerError> {
let prev = self.state.compare_exchange(
LifecycleState::Created as u8,
LifecycleState::Starting as u8,
Ordering::AcqRel,
Ordering::Acquire,
);
match prev {
Ok(_) => Ok(()),
Err(actual) => {
let state = LifecycleState::from_u8(actual);
if matches!(state, LifecycleState::Running | LifecycleState::Starting) {
Err(crate::server::errors::ServerError::AlreadyStarted)
} else {
Err(crate::server::errors::ServerError::Config(format!(
"cannot start: server is in {} state",
state
)))
}
}
}
}
pub(crate) fn mark_running(&self) -> Result<(), crate::server::errors::ServerError> {
let prev = self.state.compare_exchange(
LifecycleState::Starting as u8,
LifecycleState::Running as u8,
Ordering::AcqRel,
Ordering::Acquire,
);
match prev {
Ok(_) => {
let _ = self.ready_tx.send(true);
Ok(())
}
Err(actual) => Err(crate::server::errors::ServerError::Config(format!(
"cannot mark running: server is in {} state",
LifecycleState::from_u8(actual)
))),
}
}
pub(crate) fn drain(&self) -> Result<(), crate::server::errors::ServerError> {
let prev = self.state.compare_exchange(
LifecycleState::Running as u8,
LifecycleState::Draining as u8,
Ordering::AcqRel,
Ordering::Acquire,
);
match prev {
Ok(_) => {
crate::ops::Logger::global().emit(crate::ops::Event::new(
crate::ops::Severity::Info,
crate::ops::EventKind::DrainingStarted,
"draining in-flight connections",
));
Ok(())
}
Err(actual) => {
let state = LifecycleState::from_u8(actual);
if state == LifecycleState::Created || state == LifecycleState::Starting {
if self
.state
.compare_exchange(
actual,
LifecycleState::Stopped as u8,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
let _ = self.terminal_tx.send(());
Ok(())
} else {
Err(crate::server::errors::ServerError::Config(
"server state changed while shutting down".into(),
))
}
} else if state.is_terminal() {
Ok(())
} else {
Err(crate::server::errors::ServerError::Config(format!(
"cannot drain: server is in {} state",
state
)))
}
}
}
}
pub(crate) fn mark_stopped(&self) -> Result<(), crate::server::errors::ServerError> {
let prev = self.state.compare_exchange(
LifecycleState::Draining as u8,
LifecycleState::Stopped as u8,
Ordering::AcqRel,
Ordering::Acquire,
);
match prev {
Ok(_) => {
let _ = self.terminal_tx.send(());
Ok(())
}
Err(actual) => {
let state = LifecycleState::from_u8(actual);
if state.is_terminal() {
Ok(())
} else {
Err(crate::server::errors::ServerError::Config(format!(
"cannot stop: server is in {} state",
state
)))
}
}
}
}
#[allow(dead_code)]
pub(crate) fn mark_failed(&self) -> Result<(), crate::server::errors::ServerError> {
let current = self.state.load(Ordering::Acquire);
let current_state = LifecycleState::from_u8(current);
if current_state.is_terminal() {
return Ok(());
}
self.state
.store(LifecycleState::Failed as u8, Ordering::Release);
let _ = self.ready_tx.send(true);
let _ = self.terminal_tx.send(());
Ok(())
}
pub(crate) async fn wait_ready(&self) {
let mut rx = self.ready_tx.subscribe();
if *rx.borrow() {
return;
}
let _ = rx.changed().await;
}
pub(crate) fn subscribe_terminal(&self) -> broadcast::Receiver<()> {
self.terminal_tx.subscribe()
}
#[allow(dead_code)]
pub(crate) fn is(&self, expected: LifecycleState) -> bool {
self.state() == expected
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn initial_state_is_created() {
let lc = Lifecycle::new();
assert_eq!(lc.state(), LifecycleState::Created);
assert!(!lc.state().is_terminal());
}
#[test]
fn valid_transitions() {
let lc = Lifecycle::new();
assert!(lc.start().is_ok());
assert_eq!(lc.state(), LifecycleState::Starting);
assert!(lc.mark_running().is_ok());
assert_eq!(lc.state(), LifecycleState::Running);
assert!(lc.drain().is_ok());
assert_eq!(lc.state(), LifecycleState::Draining);
assert!(lc.mark_stopped().is_ok());
assert_eq!(lc.state(), LifecycleState::Stopped);
assert!(lc.state().is_terminal());
}
#[test]
fn double_start_fails() {
let lc = Lifecycle::new();
assert!(lc.start().is_ok());
assert!(lc.mark_running().is_ok());
let err = lc.start().unwrap_err();
assert!(err.to_string().contains("already started"));
}
#[test]
fn shutdown_before_start_stops_lifecycle() {
let lc = Lifecycle::new();
assert!(lc.drain().is_ok());
assert_eq!(lc.state(), LifecycleState::Stopped);
}
#[test]
fn mark_failed_from_any_non_terminal() {
let lc = Lifecycle::new();
assert!(lc.mark_failed().is_ok());
assert_eq!(lc.state(), LifecycleState::Failed);
assert!(lc.state().is_terminal());
}
#[test]
fn mark_stopped_from_already_stopped_is_ok() {
let lc = Lifecycle::new();
assert!(lc.start().is_ok());
assert!(lc.mark_running().is_ok());
assert!(lc.drain().is_ok());
assert!(lc.mark_stopped().is_ok());
assert!(lc.mark_stopped().is_ok());
}
#[test]
fn lifecycle_state_display() {
assert_eq!(LifecycleState::Created.to_string(), "created");
assert_eq!(LifecycleState::Running.to_string(), "running");
assert_eq!(LifecycleState::Draining.to_string(), "draining");
assert_eq!(LifecycleState::Stopped.to_string(), "stopped");
assert_eq!(LifecycleState::Failed.to_string(), "failed");
}
}