use mj_core::state::RecoveryObservation;
use std::collections::{BTreeMap, BTreeSet};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use tokio::sync::{mpsc, watch};
pub(crate) fn worker_target_mutex(session_id: &str) -> Arc<Mutex<()>> {
static LOCKS: std::sync::OnceLock<Mutex<BTreeMap<String, std::sync::Weak<Mutex<()>>>>> =
std::sync::OnceLock::new();
let mut locks = LOCKS
.get_or_init(Mutex::default)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
locks.retain(|_, lock| lock.strong_count() > 0);
let slot = locks.entry(session_id.to_owned()).or_default();
if let Some(lock) = slot.upgrade() {
return lock;
}
let lock = Arc::new(Mutex::new(()));
*slot = Arc::downgrade(&lock);
lock
}
#[derive(Clone)]
pub struct RecoveryObserver {
pub observations: mpsc::UnboundedSender<RecoveryObservation>,
pub gate: Arc<RecoveryGate>,
}
pub struct RecoveryReservation {
session_id: String,
gate: Arc<RecoveryGate>,
}
impl Drop for RecoveryReservation {
fn drop(&mut self) {
self.gate.release(&self.session_id);
}
}
pub struct RecoveryGate {
state: Mutex<RecoveryGateState>,
busy: watch::Sender<BTreeSet<String>>,
}
impl Default for RecoveryGate {
fn default() -> Self {
Self {
state: Mutex::default(),
busy: watch::channel(BTreeSet::new()).0,
}
}
}
#[derive(Default)]
struct RecoveryGateState {
busy: BTreeMap<String, Arc<AtomicBool>>,
reservations: BTreeMap<String, usize>,
}
impl RecoveryGate {
pub fn reserve(self: &Arc<Self>, session_id: &str) -> RecoveryReservation {
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
*state.reservations.entry(session_id.to_owned()).or_default() += 1;
RecoveryReservation {
session_id: session_id.to_owned(),
gate: self.clone(),
}
}
fn release(&self, session_id: &str) {
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
let Some(count) = state.reservations.get_mut(session_id) else {
return;
};
*count -= 1;
if *count == 0 {
state.reservations.remove(session_id);
}
}
pub fn try_start(&self, session_id: &str) -> Option<Arc<AtomicBool>> {
let cancelled = {
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
if state.busy.contains_key(session_id) || state.reservations.contains_key(session_id) {
return None;
}
let cancelled = Arc::new(AtomicBool::new(false));
state.busy.insert(session_id.to_owned(), cancelled.clone());
cancelled
};
self.publish_busy();
Some(cancelled)
}
pub fn finish(&self, session_id: &str) {
self.state
.lock()
.unwrap_or_else(|error| error.into_inner())
.busy
.remove(session_id);
self.publish_busy();
}
fn publish_busy(&self) {
let busy = self.busy_sessions();
self.busy.send_replace(busy);
}
pub fn subscribe(&self) -> watch::Receiver<BTreeSet<String>> {
self.busy.subscribe()
}
pub fn is_busy(&self, session_id: &str) -> bool {
self.state
.lock()
.unwrap_or_else(|error| error.into_inner())
.busy
.contains_key(session_id)
}
pub fn cancel_busy(&self, session_id: &str) {
if let Some(cancelled) = self
.state
.lock()
.unwrap_or_else(|error| error.into_inner())
.busy
.get(session_id)
{
cancelled.store(true, Ordering::Release);
}
}
pub fn cancel_all(&self) {
for cancelled in self
.state
.lock()
.unwrap_or_else(|error| error.into_inner())
.busy
.values()
{
cancelled.store(true, Ordering::Release);
}
}
pub fn busy_sessions(&self) -> BTreeSet<String> {
self.state
.lock()
.unwrap_or_else(|error| error.into_inner())
.busy
.keys()
.cloned()
.collect()
}
}
impl RecoveryObserver {
pub fn observe(&self, observation: RecoveryObservation) {
let session_id = observation.session.id.clone();
if let Err(error) = self.observations.send(observation) {
tracing::debug!(
%session_id,
%error,
"recovery observation dropped because the coordinator stopped"
);
}
}
pub fn is_busy(&self, session_id: &str) -> bool {
self.gate.is_busy(session_id)
}
pub fn reserve(&self, session_id: &str) -> RecoveryReservation {
self.gate.reserve(session_id)
}
pub fn cancel_busy(&self, session_id: &str) {
self.gate.cancel_busy(session_id);
}
pub async fn wait_idle(&self, session_id: &str) {
let mut busy = self.gate.subscribe();
while self.is_busy(session_id) {
if busy.changed().await.is_err() {
break;
}
}
}
}