use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex, MutexGuard, Weak};
use std::time::{Duration, Instant};
use trusty_common::memory_core::PalaceRegistry;
const TICKER_GRACE_INTERVALS: u32 = 3;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize)]
#[serde(rename_all = "snake_case")]
pub enum PalaceLock {
Write,
Commit,
}
#[derive(Debug, Clone)]
struct Stall {
since: Instant,
probing: bool,
mutex: Weak<tokio::sync::Mutex<()>>,
}
fn stamped_against(
stamped: &Weak<tokio::sync::Mutex<()>>,
mutex: &Arc<tokio::sync::Mutex<()>>,
) -> bool {
std::ptr::eq(stamped.as_ptr(), Arc::as_ptr(mutex))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StalledLock {
pub palace: String,
pub lock: PalaceLock,
pub age: Duration,
}
type StallKey = (String, PalaceLock);
#[derive(Debug, Default)]
pub struct LockStallTracker {
stalls: Mutex<HashMap<StallKey, Stall>>,
poisoned: AtomicBool,
last_sweep: Mutex<Option<Instant>>,
ticker: Mutex<Option<(Instant, Duration)>>,
}
pub fn probe_interval(threshold: Duration) -> Duration {
(threshold / 4).clamp(Duration::from_millis(10), Duration::from_secs(30))
}
impl LockStallTracker {
fn guard<'a, T>(&self, m: &'a Mutex<T>) -> MutexGuard<'a, T> {
m.lock().unwrap_or_else(|poisoned| {
self.poisoned.store(true, Ordering::Release);
poisoned.into_inner()
})
}
pub fn sweep_if_due(self: &Arc<Self>, registry: &PalaceRegistry, interval: Duration) {
let now = Instant::now();
{
let mut last = self.guard(&self.last_sweep);
if last.is_some_and(|t| now.saturating_duration_since(t) < interval) {
return;
}
*last = Some(now);
}
self.sweep_at(registry, now);
}
pub(crate) fn sweep_at(self: &Arc<Self>, registry: &PalaceRegistry, now: Instant) {
for id in registry.list() {
let Some(handle) = registry.peek(&id) else {
continue;
};
self.observe(id.as_str(), PalaceLock::Write, &handle.write_mutex, now);
self.observe(id.as_str(), PalaceLock::Commit, &handle.commit_mutex, now);
}
}
pub(crate) fn observe(
self: &Arc<Self>,
palace: &str,
lock: PalaceLock,
mutex: &Arc<tokio::sync::Mutex<()>>,
now: Instant,
) {
let key = (palace.to_string(), lock);
if let Ok(free) = mutex.try_lock() {
drop(free);
let mut stalls = self.guard(&self.stalls);
if stalls
.get(&key)
.is_some_and(|s| !s.probing || !stamped_against(&s.mutex, mutex))
{
stalls.remove(&key);
}
return;
}
let Some(token) = self.claim(key, now, mutex) else {
return;
};
match tokio::runtime::Handle::try_current() {
Ok(rt) => {
rt.spawn(token.wait(Arc::clone(mutex)));
}
Err(_) => drop(token),
}
}
pub(crate) fn claim(
self: &Arc<Self>,
key: StallKey,
now: Instant,
mutex: &Arc<tokio::sync::Mutex<()>>,
) -> Option<ProbeToken> {
let mut stalls = self.guard(&self.stalls);
match stalls.get_mut(&key) {
Some(s) if stamped_against(&s.mutex, mutex) => {
if s.probing {
return None;
}
s.probing = true;
}
_ => {
stalls.insert(
key.clone(),
Stall {
since: now,
probing: true,
mutex: Arc::downgrade(mutex),
},
);
}
}
Some(ProbeToken {
tracker: Arc::clone(self),
key,
mutex: Arc::downgrade(mutex),
acquired: false,
})
}
pub fn oldest_stall_at(&self, now: Instant) -> Option<StalledLock> {
let stalls = self.guard(&self.stalls);
stalls
.iter()
.min_by_key(|(_, s)| s.since)
.map(|((palace, lock), s)| StalledLock {
palace: palace.clone(),
lock: *lock,
age: now.saturating_duration_since(s.since),
})
}
pub(crate) fn beat(&self, now: Instant, interval: Duration) {
*self.guard(&self.ticker) = Some((now, interval));
}
pub fn degraded_at(&self, now: Instant) -> Option<String> {
drop(self.guard(&self.stalls));
let beat = *self.guard(&self.ticker);
if self.poisoned.load(Ordering::Acquire) {
return Some(
"palace lock stall tracking was poisoned by a panic; stall ages may be \
incomplete (#4001)"
.to_string(),
);
}
let (last, interval) = beat?;
let silent = now.saturating_duration_since(last);
(silent > interval * TICKER_GRACE_INTERVALS).then(|| {
format!(
"palace lock stall ticker has not run for {}s (interval {}ms); a held \
lock may go unnoticed between health polls (#4001)",
silent.as_secs(),
interval.as_millis()
)
})
}
}
#[derive(Debug)]
pub(crate) struct ProbeToken {
tracker: Arc<LockStallTracker>,
key: StallKey,
mutex: Weak<tokio::sync::Mutex<()>>,
acquired: bool,
}
impl ProbeToken {
pub(crate) async fn wait(mut self, mutex: Arc<tokio::sync::Mutex<()>>) {
drop(mutex.lock().await);
let mut stalls = self.tracker.guard(&self.tracker.stalls);
if stalls
.get(&self.key)
.is_some_and(|s| Weak::ptr_eq(&s.mutex, &self.mutex))
{
stalls.remove(&self.key);
}
drop(stalls);
self.acquired = true;
}
}
impl Drop for ProbeToken {
fn drop(&mut self) {
if self.acquired {
return;
}
let mut stalls = self.tracker.guard(&self.tracker.stalls);
if let Some(s) = stalls
.get_mut(&self.key)
.filter(|s| Weak::ptr_eq(&s.mutex, &self.mutex))
{
s.probing = false;
}
}
}
pub fn spawn_lock_stall_ticker(
tracker: Arc<LockStallTracker>,
registry: Arc<PalaceRegistry>,
interval: Duration,
) -> tokio::task::JoinHandle<()> {
tracker.beat(Instant::now(), interval);
tokio::spawn(async move {
loop {
tokio::time::sleep(interval).await;
tracker.beat(Instant::now(), interval);
tracker.sweep_if_due(®istry, interval);
}
})
}
#[cfg(test)]
#[path = "lock_stall_tests.rs"]
mod tests;