use std::sync::{Arc, Mutex, PoisonError};
#[cfg_attr(
not(any(feature = "image-vision", feature = "audio-transcribe")),
allow(dead_code)
)]
pub(crate) struct EngineSlot<T> {
state: Mutex<SlotState<T>>,
}
enum SlotState<T> {
Uninit,
Absent,
Ready(Arc<T>),
}
#[cfg_attr(
not(any(feature = "image-vision", feature = "audio-transcribe")),
allow(dead_code)
)]
impl<T> EngineSlot<T> {
pub(crate) const fn new() -> Self {
Self {
state: Mutex::new(SlotState::Uninit),
}
}
pub(crate) fn get_or_init(&self, init: impl FnOnce() -> Option<T>) -> Option<Arc<T>> {
let mut state = self.state.lock().unwrap_or_else(PoisonError::into_inner);
match &*state {
SlotState::Ready(engine) => return Some(Arc::clone(engine)),
SlotState::Absent => return None,
SlotState::Uninit => {}
}
let Some(engine) = init() else {
*state = SlotState::Absent;
return None;
};
let engine = Arc::new(engine);
*state = SlotState::Ready(Arc::clone(&engine));
Some(engine)
}
pub(crate) fn release(&self) -> bool {
let previous = std::mem::replace(
&mut *self.state.lock().unwrap_or_else(PoisonError::into_inner),
SlotState::Uninit,
);
matches!(previous, SlotState::Ready(_))
}
}
#[cfg(test)]
mod tests {
use super::EngineSlot;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
struct Payload(Arc<AtomicUsize>);
impl Drop for Payload {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
fn slot() -> (EngineSlot<Payload>, Arc<AtomicUsize>) {
(EngineSlot::new(), Arc::new(AtomicUsize::new(0)))
}
#[test]
fn initialises_once_and_reuses() {
let (slot, drops) = slot();
let inits = AtomicUsize::new(0);
let build = || {
inits.fetch_add(1, Ordering::SeqCst);
Some(Payload(Arc::clone(&drops)))
};
assert!(slot.get_or_init(build).is_some());
assert!(slot.get_or_init(build).is_some());
assert!(slot.get_or_init(build).is_some());
assert_eq!(inits.load(Ordering::SeqCst), 1, "engine built once per run");
assert_eq!(drops.load(Ordering::SeqCst), 0, "engine still resident");
}
#[test]
fn release_drops_the_engine_exactly_once_and_is_idempotent() {
let (slot, drops) = slot();
assert!(
slot.get_or_init(|| Some(Payload(Arc::clone(&drops))))
.is_some()
);
assert!(slot.release(), "first release reports the live engine");
assert_eq!(drops.load(Ordering::SeqCst), 1, "engine dropped on release");
assert!(!slot.release(), "second release has nothing to drop");
assert!(!slot.release());
assert_eq!(drops.load(Ordering::SeqCst), 1, "no double free");
}
#[test]
fn release_is_safe_when_never_initialised() {
let (slot, drops) = slot();
assert!(!slot.release(), "an untouched slot releases nothing");
assert_eq!(drops.load(Ordering::SeqCst), 0);
}
#[test]
fn unavailable_engine_is_remembered_and_releases_nothing() {
let slot: EngineSlot<Payload> = EngineSlot::new();
let inits = AtomicUsize::new(0);
let build = || {
inits.fetch_add(1, Ordering::SeqCst);
None
};
assert!(slot.get_or_init(build).is_none());
assert!(slot.get_or_init(build).is_none());
assert_eq!(inits.load(Ordering::SeqCst), 1, "missing model probed once");
assert!(!slot.release(), "nothing was resident to release");
}
#[test]
fn release_resets_the_slot_for_a_later_run() {
let (slot, drops) = slot();
assert!(
slot.get_or_init(|| Some(Payload(Arc::clone(&drops))))
.is_some()
);
assert!(slot.release());
assert!(
slot.get_or_init(|| Some(Payload(Arc::clone(&drops))))
.is_some(),
"a released slot re-initialises on demand"
);
assert!(slot.release());
assert_eq!(drops.load(Ordering::SeqCst), 2, "each engine dropped once");
}
#[test]
fn an_outstanding_handle_keeps_the_engine_alive_past_release() {
let (slot, drops) = slot();
let held = slot
.get_or_init(|| Some(Payload(Arc::clone(&drops))))
.expect("engine built");
assert!(slot.release());
assert_eq!(
drops.load(Ordering::SeqCst),
0,
"an in-flight caller is never freed underneath"
);
drop(held);
assert_eq!(
drops.load(Ordering::SeqCst),
1,
"freed with the last handle"
);
}
#[test]
fn concurrent_callers_build_one_engine() {
let (slot, drops) = slot();
let inits = AtomicUsize::new(0);
std::thread::scope(|scope| {
for _ in 0..8 {
scope.spawn(|| {
let engine = slot.get_or_init(|| {
inits.fetch_add(1, Ordering::SeqCst);
Some(Payload(Arc::clone(&drops)))
});
assert!(engine.is_some());
});
}
});
assert_eq!(
inits.load(Ordering::SeqCst),
1,
"one engine for all threads"
);
assert!(slot.release());
assert_eq!(drops.load(Ordering::SeqCst), 1);
}
}