use std::collections::HashMap;
use std::sync::{Arc, Mutex as StdMutex, Weak};
use meerkat_core::InputId;
use meerkat_core::types::SessionId;
use meerkat_session::RuntimeContextAdmissionGuard;
use tokio::sync::Mutex;
use crate::StagedSessionRegistry;
pub type ActiveCapacityGuard = RuntimeContextAdmissionGuard;
pub type StagedCapacityAdmissions = Arc<StdMutex<HashMap<SessionId, ActiveCapacityGuard>>>;
pub fn restore_staged_capacity_admission(
admissions: &StagedCapacityAdmissions,
session_id: SessionId,
admission: ActiveCapacityGuard,
) {
let mut admissions = admissions
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
admissions.insert(session_id, admission);
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StagedCapacityCollision {
pub session_id: SessionId,
}
pub fn insert_staged_capacity_admission(
admissions: &StagedCapacityAdmissions,
session_id: SessionId,
admission: ActiveCapacityGuard,
) -> Result<(), StagedCapacityCollision> {
let mut guard = admissions
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if guard.contains_key(&session_id) {
return Err(StagedCapacityCollision { session_id });
}
guard.insert(session_id, admission);
Ok(())
}
pub fn take_staged_capacity_admission(
admissions: &StagedCapacityAdmissions,
session_id: &SessionId,
) -> Option<ActiveCapacityGuard> {
admissions
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(session_id)
}
pub fn has_staged_capacity_admission(
admissions: &StagedCapacityAdmissions,
session_id: &SessionId,
) -> bool {
admissions
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.contains_key(session_id)
}
pub fn discard_staged_capacity_admission(
admissions: &StagedCapacityAdmissions,
session_id: &SessionId,
) {
drop(take_staged_capacity_admission(admissions, session_id));
}
#[derive(Clone)]
pub struct StagedAdmissionRestore {
pub admissions: StagedCapacityAdmissions,
pub session_id: SessionId,
}
pub struct RuntimePreAdmission {
pub(crate) admission: Option<ActiveCapacityGuard>,
pub(crate) staged_restore: Option<StagedAdmissionRestore>,
}
impl RuntimePreAdmission {
pub fn fresh(admission: ActiveCapacityGuard) -> Self {
Self {
admission: Some(admission),
staged_restore: None,
}
}
pub fn staged(
session_id: SessionId,
admissions: StagedCapacityAdmissions,
admission: ActiveCapacityGuard,
) -> Self {
Self {
admission: Some(admission),
staged_restore: Some(StagedAdmissionRestore {
admissions,
session_id,
}),
}
}
#[allow(clippy::expect_used)]
pub fn into_admission(mut self) -> ActiveCapacityGuard {
self.staged_restore = None;
self.admission
.take()
.expect("runtime pre-admission should not be consumed twice")
}
}
impl From<ActiveCapacityGuard> for RuntimePreAdmission {
fn from(admission: ActiveCapacityGuard) -> Self {
Self::fresh(admission)
}
}
impl Drop for RuntimePreAdmission {
fn drop(&mut self) {
let Some(admission) = self.admission.take() else {
return;
};
if let Some(restore) = self.staged_restore.take() {
restore_staged_capacity_admission(&restore.admissions, restore.session_id, admission);
} else {
drop(admission);
}
}
}
pub struct RuntimePreAdmissionGuard {
admission: Option<RuntimePreAdmission>,
}
impl RuntimePreAdmissionGuard {
pub fn new(admission: impl Into<RuntimePreAdmission>) -> Self {
Self {
admission: Some(admission.into()),
}
}
pub fn take(&mut self) -> Option<RuntimePreAdmission> {
self.admission.take()
}
}
pub struct RuntimePreAdmissionEntry {
pub input_id: InputId,
pub admission: RuntimePreAdmission,
}
pub struct RuntimeRegistrationLockLease {
pub locks: Arc<StdMutex<HashMap<SessionId, Weak<Mutex<()>>>>>,
pub session_id: SessionId,
pub lock: Arc<Mutex<()>>,
}
impl RuntimeRegistrationLockLease {
pub fn mutex(&self) -> &Mutex<()> {
&self.lock
}
}
fn evict_runtime_registration_lock_if_last(
locks: &StdMutex<HashMap<SessionId, Weak<Mutex<()>>>>,
session_id: &SessionId,
lock: &Arc<Mutex<()>>,
) {
let this_lock = Arc::downgrade(lock);
let mut locks = locks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if Arc::strong_count(lock) == 1
&& locks
.get(session_id)
.is_some_and(|registered| registered.ptr_eq(&this_lock))
{
locks.remove(session_id);
}
}
impl Drop for RuntimeRegistrationLockLease {
fn drop(&mut self) {
evict_runtime_registration_lock_if_last(&self.locks, &self.session_id, &self.lock);
}
}
#[cfg(test)]
mod registration_lock_tests {
use super::*;
#[test]
fn lease_eviction_rechecks_concurrent_weak_upgrade_under_map_lock() {
let locks = Arc::new(StdMutex::new(HashMap::new()));
let session_id = SessionId::new();
let lock = Arc::new(Mutex::new(()));
locks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(session_id.clone(), Arc::downgrade(&lock));
assert_eq!(Arc::strong_count(&lock), 1);
let concurrent_lock = locks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(&session_id)
.and_then(Weak::upgrade)
.expect("concurrent caller upgrades the registered mutex");
evict_runtime_registration_lock_if_last(&locks, &session_id, &lock);
let third_caller_lock = locks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(&session_id)
.and_then(Weak::upgrade)
.expect("live upgraded mutex must remain registered");
assert!(Arc::ptr_eq(&concurrent_lock, &third_caller_lock));
assert!(Arc::ptr_eq(&lock, &third_caller_lock));
let _concurrent_guard = concurrent_lock
.try_lock()
.expect("concurrent caller enters the shared critical section");
assert!(
third_caller_lock.try_lock().is_err(),
"third caller must not enter a split critical section"
);
}
}
pub struct StagedArchiveRollbackGuard {
staged_sessions: Arc<StagedSessionRegistry>,
session_id: SessionId,
restore_on_drop: bool,
}
impl StagedArchiveRollbackGuard {
pub fn new(staged_sessions: Arc<StagedSessionRegistry>, session_id: &SessionId) -> Self {
Self {
staged_sessions,
session_id: session_id.clone(),
restore_on_drop: true,
}
}
pub fn disarm(&mut self) {
self.restore_on_drop = false;
}
}
impl Drop for StagedArchiveRollbackGuard {
fn drop(&mut self) {
if !self.restore_on_drop {
return;
}
let staged_sessions = Arc::clone(&self.staged_sessions);
let session_id = self.session_id.clone();
tokio::spawn(async move {
let _ = staged_sessions.restore_archive(&session_id).await;
});
}
}
pub trait RuntimePreAdmissionRestore: Send + Sync {
fn restore_or_release(&self, session_id: &SessionId, input_id: &InputId);
}
pub struct RuntimePreAdmissionRegistration {
runtime: Arc<dyn RuntimePreAdmissionRestore>,
session_id: SessionId,
input_id: InputId,
release_on_drop: bool,
}
impl RuntimePreAdmissionRegistration {
pub fn new(
runtime: Arc<dyn RuntimePreAdmissionRestore>,
session_id: SessionId,
input_id: InputId,
) -> Self {
Self {
runtime,
session_id,
input_id,
release_on_drop: true,
}
}
pub fn disarm(mut self) {
self.release_on_drop = false;
}
}
impl Drop for RuntimePreAdmissionRegistration {
fn drop(&mut self) {
if self.release_on_drop {
self.runtime
.restore_or_release(&self.session_id, &self.input_id);
}
}
}