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
}
}
impl Drop for RuntimeRegistrationLockLease {
fn drop(&mut self) {
if Arc::strong_count(&self.lock) != 1 {
return;
}
let this_lock = Arc::downgrade(&self.lock);
let mut locks = self
.locks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if locks
.get(&self.session_id)
.is_some_and(|registered: &Weak<Mutex<()>>| registered.ptr_eq(&this_lock))
{
locks.remove(&self.session_id);
}
}
}
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);
}
}
}