use std::collections::BTreeMap;
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, Instant};
const DRAIN_LAPSE: Duration = Duration::from_secs(10);
const WAIT_REPORT_INTERVAL: Duration = Duration::from_secs(30);
pub(crate) const DESTROY_LABEL: &str = "session destroy";
pub(crate) const DESTROY_WAIT_BOUND: Duration = Duration::from_secs(60);
#[derive(Default)]
pub(crate) struct Gate(Mutex<State>);
#[derive(Default)]
struct State {
closed: bool,
destroying: Vec<String>,
destroy_wait_started: Option<Instant>,
destroy_wait_bound: Option<Duration>,
drain_requested: Option<Instant>,
handoff_wait: Option<(Instant, Instant)>,
active: BTreeMap<&'static str, BTreeMap<u64, Instant>>,
next_hold: u64,
}
impl State {
fn destroy_wait_holds(&mut self) -> bool {
let started = *self.destroy_wait_started.get_or_insert_with(Instant::now);
let bound = self.destroy_wait_bound.unwrap_or(DESTROY_WAIT_BOUND);
if started.elapsed() < bound {
return true;
}
tracing::warn!(
sessions = ?self.destroying,
bound_seconds = bound.as_secs_f64(),
"daemon stops with session destroys unfinished; each session comes back and needs `mj destroy` again"
);
false
}
fn holds_only_destroys_or_nothing(&self) -> bool {
self.active.keys().all(|label| *label == DESTROY_LABEL)
}
fn draining(&self) -> bool {
self.drain_requested
.is_some_and(|requested| requested.elapsed() < DRAIN_LAPSE)
}
fn blockers(&self) -> Vec<Blocker> {
self.active
.iter()
.filter_map(|(label, holds)| {
let (_, oldest) = holds.first_key_value()?;
Some(Blocker {
label,
count: holds.len(),
age: oldest.elapsed(),
})
})
.collect()
}
fn report_refused_handoff(&mut self) {
let now = Instant::now();
let (started, reported) = match self.handoff_wait {
Some(wait) if self.draining() => wait,
_ => {
self.handoff_wait = Some((now, now));
tracing::info!(
blockers = %describe(&self.blockers()),
"daemon upgrade handoff is waiting for daemon-owned work"
);
return;
}
};
if now.duration_since(reported) >= WAIT_REPORT_INTERVAL {
self.handoff_wait = Some((started, now));
tracing::info!(
blockers = %describe(&self.blockers()),
waited_seconds = started.elapsed().as_secs(),
"daemon upgrade handoff is still waiting for daemon-owned work"
);
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Blocker {
pub(crate) label: &'static str,
pub(crate) count: usize,
pub(crate) age: Duration,
}
#[cfg(test)]
impl Blocker {
fn counted_label(&self) -> String {
match self.count {
1 => self.label.to_owned(),
count => format!("{} x{count}", self.label),
}
}
}
impl std::fmt::Display for Blocker {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let seconds = self.age.as_secs();
let age = match seconds {
0..60 => format!("{seconds}s"),
_ => format!("{}m {:02}s", seconds / 60, seconds % 60),
};
match self.count {
1 => write!(formatter, "{} ({age})", self.label),
count => write!(formatter, "{} x{count} (oldest {age})", self.label),
}
}
}
fn describe(blockers: &[Blocker]) -> String {
blockers
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
.join(", ")
}
#[derive(Clone)]
pub(crate) struct Work {
_registration: Arc<Registration>,
}
struct Registration {
gate: Arc<Gate>,
label: &'static str,
hold: u64,
destroying: Option<String>,
}
impl Gate {
pub(crate) fn is_open(&self) -> bool {
!self.lock().closed
}
pub(crate) fn is_draining(&self) -> bool {
let state = self.lock();
state.closed || state.draining()
}
pub(crate) fn enter(self: &Arc<Self>, label: &'static str) -> anyhow::Result<Work> {
let mut state = self.lock();
anyhow::ensure!(!state.closed, "daemon upgrade handoff is underway");
Ok(self.register(&mut state, label))
}
pub(crate) fn enter_unless_draining(
self: &Arc<Self>,
label: &'static str,
) -> anyhow::Result<Work> {
let mut state = self.lock();
anyhow::ensure!(!state.closed, "daemon upgrade handoff is underway");
anyhow::ensure!(!state.draining(), "a daemon upgrade handoff is waiting");
Ok(self.register(&mut state, label))
}
pub(crate) fn enter_destroy(self: &Arc<Self>, session_id: &str) -> anyhow::Result<Work> {
let mut state = self.lock();
anyhow::ensure!(!state.closed, "daemon upgrade handoff is underway");
state.destroying.push(session_id.to_owned());
Ok(self.register_as(&mut state, DESTROY_LABEL, Some(session_id.to_owned())))
}
fn register(self: &Arc<Self>, state: &mut State, label: &'static str) -> Work {
self.register_as(state, label, None)
}
fn register_as(
self: &Arc<Self>,
state: &mut State,
label: &'static str,
destroying: Option<String>,
) -> Work {
let hold = state.next_hold;
state.next_hold += 1;
state
.active
.entry(label)
.or_default()
.insert(hold, Instant::now());
Work {
_registration: Arc::new(Registration {
gate: self.clone(),
label,
hold,
destroying,
}),
}
}
#[cfg(test)]
pub(crate) fn with_destroy_wait_bound(bound: Duration) -> Arc<Self> {
let gate = Arc::new(Self::default());
gate.lock().destroy_wait_bound = Some(bound);
gate
}
pub(crate) fn destroys_hold_the_stop(&self) -> bool {
let mut state = self.lock();
if state.destroying.is_empty() {
return false;
}
state.drain_requested = Some(Instant::now());
state.destroy_wait_holds()
}
fn lock(&self) -> std::sync::MutexGuard<'_, State> {
self.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[cfg(test)]
pub(crate) fn active_labels(&self) -> Vec<String> {
self.blockers().iter().map(Blocker::counted_label).collect()
}
pub(crate) fn blockers(&self) -> Vec<Blocker> {
self.lock().blockers()
}
pub(crate) fn try_close(&self) -> bool {
let mut state = self.lock();
if state.closed {
return true;
}
if !state.holds_only_destroys_or_nothing()
|| (!state.destroying.is_empty() && state.destroy_wait_holds())
{
state.report_refused_handoff();
state.drain_requested = Some(Instant::now());
return false;
}
tracing::info!(
waited_seconds = state
.handoff_wait
.take()
.map_or(0.0, |(started, _)| started.elapsed().as_secs_f64()),
"daemon upgrade handoff closed admission; the daemon is shutting down"
);
state.closed = true;
true
}
}
impl Drop for Registration {
fn drop(&mut self) {
let mut state = self.gate.lock();
if let Some(session_id) = &self.destroying
&& let Some(index) = state.destroying.iter().position(|id| id == session_id)
{
state.destroying.remove(index);
}
if state.destroying.is_empty() {
state.destroy_wait_started = None;
}
let holds = state
.active
.get_mut(self.label)
.expect("registered upgrade work");
holds.remove(&self.hold);
if holds.is_empty() {
state.active.remove(self.label);
}
}
}
pub(crate) fn gate() -> &'static Arc<Gate> {
static GATE: OnceLock<Arc<Gate>> = OnceLock::new();
GATE.get_or_init(Default::default)
}
pub(crate) fn activity(label: &'static str) -> anyhow::Result<Work> {
gate().enter(label)
}
pub(crate) fn activity_unless_draining(label: &'static str) -> anyhow::Result<Work> {
gate().enter_unless_draining(label)
}
pub(crate) fn destroy_activity(session_id: &str) -> anyhow::Result<Work> {
gate().enter_destroy(session_id)
}
pub(crate) fn is_draining() -> bool {
gate().is_draining()
}
#[cfg(test)]
pub(crate) fn active_labels() -> Vec<String> {
gate().active_labels()
}
pub(crate) fn blockers() -> Vec<Blocker> {
gate().blockers()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn handoff_waits_for_every_owner_and_closes_admission_atomically() {
let gate = Arc::new(Gate::default());
let provisioning = gate.enter("provisioning").unwrap();
for _ in 0..1000 {
assert!(!gate.try_close());
}
let steering = gate.enter("steering").unwrap();
drop(provisioning);
assert!(!gate.try_close());
drop(steering);
assert!(gate.try_close());
assert!(gate.enter("new operation").is_err());
assert!(gate.try_close(), "handoff admission is idempotent");
}
#[test]
fn active_labels_name_each_blocker_and_count_repeats() {
let gate = Arc::new(Gate::default());
assert!(gate.active_labels().is_empty());
let first = gate.enter("session lifecycle").unwrap();
let second = gate.enter("session lifecycle").unwrap();
let request = gate.enter("client request").unwrap();
assert_eq!(
gate.active_labels(),
vec![
"client request".to_owned(),
"session lifecycle x2".to_owned()
]
);
drop(second);
assert_eq!(
gate.active_labels(),
vec!["client request".to_owned(), "session lifecycle".to_owned()]
);
drop(first);
drop(request);
assert!(gate.active_labels().is_empty());
}
#[test]
fn blockers_carry_the_age_of_each_labels_oldest_hold() {
let gate = Arc::new(Gate::default());
let oldest = gate.enter("session lifecycle").unwrap();
let newer = gate.enter("session lifecycle").unwrap();
let sync = gate.enter("SessionWiki sync").unwrap();
{
let mut state = gate.lock();
let holds = state.active.get_mut("session lifecycle").unwrap();
let first = holds.values_mut().next().unwrap();
*first -= Duration::from_secs(65);
}
assert_eq!(
gate.blockers()
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>(),
[
"SessionWiki sync (0s)",
"session lifecycle x2 (oldest 1m 05s)"
]
);
drop(oldest);
let lifecycle = gate.blockers()[1].clone();
assert_eq!((lifecycle.label, lifecycle.count), ("session lifecycle", 1));
assert!(lifecycle.age < Duration::from_secs(60), "{lifecycle:?}");
drop((newer, sync));
assert!(gate.blockers().is_empty());
}
#[test]
fn a_waiting_handoff_names_its_blockers_in_the_log_and_reports_when_it_closes() {
let log = crate::test_log::CapturedLog::default();
let _guard = tracing::subscriber::set_default(log.clone());
let gate = Arc::new(Gate::default());
let held = gate.enter("SessionWiki sync").unwrap();
for _ in 0..10 {
assert!(!gate.try_close());
}
let waiting = log.at(tracing::Level::INFO);
assert_eq!(
waiting.len(),
1,
"a wait is named once, not on every attempt: {waiting:?}"
);
assert!(
waiting[0].contains("SessionWiki sync (0s)"),
"the log names the blocker and its age: {waiting:?}"
);
{
let mut state = gate.lock();
let (started, reported) = state.handoff_wait.unwrap();
state.handoff_wait = Some((
started - WAIT_REPORT_INTERVAL,
reported - WAIT_REPORT_INTERVAL,
));
}
assert!(!gate.try_close());
let waiting = log.at(tracing::Level::INFO);
assert_eq!(waiting.len(), 2, "{waiting:?}");
assert!(waiting[1].contains("still waiting"), "{waiting:?}");
drop(held);
assert!(gate.try_close());
let closed = log.at(tracing::Level::INFO);
assert!(
closed[2].contains("closed admission"),
"the log says when the wait ended: {closed:?}"
);
}
#[test]
fn a_waiting_handoff_refuses_deferrable_work_but_admits_control_work() {
let gate = Arc::new(Gate::default());
let running = gate.enter("session lifecycle").unwrap();
let deferrable = gate.enter_unless_draining("worker swap").unwrap();
assert!(!gate.is_draining());
assert!(!gate.try_close());
assert!(gate.is_draining());
assert!(gate.enter_unless_draining("worker swap").is_err());
let control = gate.enter("client request").unwrap();
drop((running, deferrable, control));
assert!(gate.try_close());
assert!(gate.enter_unless_draining("worker swap").is_err());
}
#[test]
fn draining_lapses_when_the_handoff_stops_asking() {
let gate = Arc::new(Gate::default());
let running = gate.enter("session lifecycle").unwrap();
assert!(!gate.try_close());
gate.lock().drain_requested = Some(Instant::now() - DRAIN_LAPSE);
assert!(!gate.is_draining());
assert!(gate.enter_unless_draining("worker swap").is_ok());
drop(running);
}
#[test]
fn a_handoff_waits_for_a_destroy_until_the_bound_and_then_abandons_it_with_a_warning() {
let log = crate::test_log::CapturedLog::default();
let _guard = tracing::subscriber::set_default(log.clone());
let bound = Duration::from_secs(60);
let gate = Gate::with_destroy_wait_bound(bound);
let destroy = gate.enter_destroy("aaaa1111").unwrap();
assert_eq!(gate.active_labels(), vec![DESTROY_LABEL.to_owned()]);
assert!(!gate.try_close(), "a destroy in flight holds the handoff");
assert!(gate.is_draining());
assert!(gate.destroys_hold_the_stop());
assert!(log.at(tracing::Level::WARN).is_empty());
let started = gate.lock().destroy_wait_started.expect("the wait began");
gate.lock().destroy_wait_started = Some(started - bound - Duration::from_secs(1));
assert!(!gate.destroys_hold_the_stop());
assert!(gate.try_close(), "past the bound the gate proceeds");
let warnings = log.at(tracing::Level::WARN);
assert!(
warnings.iter().any(|line| line.contains("aaaa1111")),
"the warning names the abandoned session: {warnings:?}"
);
drop(destroy);
}
#[test]
fn other_work_still_holds_the_handoff_past_the_destroy_bound() {
let gate = Gate::with_destroy_wait_bound(Duration::ZERO);
let destroy = gate.enter_destroy("aaaa1111").unwrap();
let lifecycle = gate.enter("session lifecycle").unwrap();
assert!(!gate.try_close());
drop(lifecycle);
assert!(gate.try_close());
drop(destroy);
}
#[test]
fn a_finished_destroy_lets_the_handoff_close_without_a_warning() {
let log = crate::test_log::CapturedLog::default();
let _guard = tracing::subscriber::set_default(log.clone());
let gate = Gate::with_destroy_wait_bound(Duration::from_secs(60));
let destroy = gate.enter_destroy("aaaa1111").unwrap();
assert!(!gate.try_close());
drop(destroy);
assert!(!gate.destroys_hold_the_stop());
assert!(gate.try_close());
assert!(log.at(tracing::Level::WARN).is_empty());
}
#[test]
fn clones_of_one_operation_count_once_and_hold_until_the_last_drops() {
let gate = Arc::new(Gate::default());
let work = gate.enter("web action").unwrap();
let clone = work.clone();
assert_eq!(gate.active_labels(), vec!["web action".to_owned()]);
drop(work);
assert!(!gate.try_close());
drop(clone);
assert!(gate.try_close());
}
#[test]
fn concurrent_admission_either_owns_work_or_observes_a_committed_handoff() {
for _ in 0..100 {
let gate = Arc::new(Gate::default());
let thread_gate = gate.clone();
let barrier = Arc::new(std::sync::Barrier::new(2));
let thread_barrier = barrier.clone();
let worker = std::thread::spawn(move || {
thread_barrier.wait();
thread_gate.enter("request")
});
barrier.wait();
let closed = gate.try_close();
let admitted = worker.join().unwrap();
assert_ne!(closed, admitted.is_ok());
drop(admitted);
assert!(gate.try_close());
}
}
}