use std::collections::HashMap;
use std::sync::{Arc, Mutex, PoisonError};
use tokio::sync::{AcquireError, Notify, OwnedSemaphorePermit, Semaphore};
const UNBOUNDED_PERMITS: usize = Semaphore::MAX_PERMITS;
#[derive(Debug, Clone, Default)]
pub struct InferencePoolConfig {
per_model: HashMap<String, usize>,
default_limit: Option<usize>,
}
impl InferencePoolConfig {
pub fn new() -> Self {
Self::default()
}
pub fn with_default(mut self, limit: Option<usize>) -> Self {
self.default_limit = limit;
self
}
pub fn set_limit(&mut self, model: impl Into<String>, limit: usize) {
self.per_model.insert(model.into(), limit);
}
pub fn limit_for(&self, model: &str) -> Option<usize> {
self.per_model.get(model).copied().or(self.default_limit)
}
}
#[derive(Debug)]
pub struct InferencePools {
config: InferencePoolConfig,
semaphores: Mutex<HashMap<String, Arc<Semaphore>>>,
wake: Option<Arc<Notify>>,
}
impl InferencePools {
pub fn new(config: InferencePoolConfig) -> Self {
Self {
config,
semaphores: Mutex::new(HashMap::new()),
wake: None,
}
}
pub fn with_wake(mut self, wake: Arc<Notify>) -> Self {
self.wake = Some(wake);
self
}
pub async fn acquire(&self, model: &str) -> InferencePermit {
let semaphore = self.semaphore_for(model);
let permit = expect_permit(semaphore.acquire_owned().await);
self.issue(model, permit)
}
pub fn try_acquire(&self, model: &str) -> Option<InferencePermit> {
let semaphore = self.semaphore_for(model);
match semaphore.try_acquire_owned() {
Ok(permit) => Some(self.issue(model, permit)),
Err(_) => None,
}
}
fn issue(&self, model: &str, permit: OwnedSemaphorePermit) -> InferencePermit {
tracing::trace!(model = %model, "inference slot acquired");
InferencePermit {
permit: Some(permit),
model: model.to_string(),
wake: self.wake.clone(),
}
}
pub fn occupancy(&self) -> Vec<PoolOccupancy> {
let map = self
.semaphores
.lock()
.unwrap_or_else(PoisonError::into_inner);
let mut out: Vec<PoolOccupancy> = map
.iter()
.map(|(model, semaphore)| {
let cap = self.config.limit_for(model);
let free = semaphore.available_permits();
PoolOccupancy {
model: model.clone(),
in_use: cap.unwrap_or(UNBOUNDED_PERMITS).saturating_sub(free),
cap,
}
})
.collect();
out.sort_by(|a, b| a.model.cmp(&b.model)); out
}
fn semaphore_for(&self, model: &str) -> Arc<Semaphore> {
let mut map = self
.semaphores
.lock()
.unwrap_or_else(PoisonError::into_inner);
if let Some(existing) = map.get(model) {
return existing.clone();
}
let permits = self.config.limit_for(model).unwrap_or(UNBOUNDED_PERMITS);
let semaphore = Arc::new(Semaphore::new(permits));
map.insert(model.to_string(), semaphore.clone());
semaphore
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PoolOccupancy {
pub model: String,
pub in_use: usize,
pub cap: Option<usize>,
}
impl PoolOccupancy {
#[must_use]
pub fn is_full(&self) -> bool {
self.cap.is_some_and(|cap| self.in_use >= cap)
}
}
impl std::fmt::Display for PoolOccupancy {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.cap {
Some(cap) => write!(f, "{}={}/{}", self.model, self.in_use, cap),
None => write!(f, "{}={}/unbounded", self.model, self.in_use),
}
}
}
pub(crate) fn expect_permit(
result: Result<OwnedSemaphorePermit, AcquireError>,
) -> OwnedSemaphorePermit {
result.expect("a lane semaphore is never closed")
}
#[derive(Debug)]
pub struct InferencePermit {
permit: Option<OwnedSemaphorePermit>,
model: String,
wake: Option<Arc<Notify>>,
}
impl Drop for InferencePermit {
fn drop(&mut self) {
drop(self.permit.take());
tracing::trace!(model = %self.model, "inference slot released");
if let Some(wake) = &self.wake {
wake.notify_one();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn limit_for_prefers_explicit_over_default() {
let mut cfg = InferencePoolConfig::new().with_default(Some(5));
cfg.set_limit("anthropic:x", 3);
assert_eq!(cfg.limit_for("anthropic:x"), Some(3)); assert_eq!(cfg.limit_for("ollama:gemma"), Some(5)); }
#[test]
fn limit_for_unbounded_when_no_entry_and_no_default() {
let cfg = InferencePoolConfig::new();
assert_eq!(cfg.limit_for("anything"), None);
}
#[test]
fn semaphore_for_is_cached_per_model() {
let pools = InferencePools::new(InferencePoolConfig::new());
let first = pools.semaphore_for("m");
let second = pools.semaphore_for("m"); assert!(Arc::ptr_eq(&first, &second));
let other = pools.semaphore_for("n"); assert!(!Arc::ptr_eq(&first, &other));
}
#[tokio::test]
async fn acquire_bounds_concurrency_and_releases_on_drop() {
let mut cfg = InferencePoolConfig::new();
cfg.set_limit("m", 1);
let pools = Arc::new(InferencePools::new(cfg));
let permit = pools.acquire("m").await;
let pools2 = pools.clone();
let waiting = tokio::spawn(async move { pools2.acquire("m").await });
tokio::task::yield_now().await;
assert!(
!waiting.is_finished(),
"second acquire must wait for a slot"
);
drop(permit); let _second = waiting.await.expect("waiter task should not panic");
}
#[test]
fn try_acquire_returns_none_when_full() {
let mut cfg = InferencePoolConfig::new();
cfg.set_limit("m", 1);
let pools = InferencePools::new(cfg);
let permit = pools.try_acquire("m").expect("first slot is free"); assert!(pools.try_acquire("m").is_none()); drop(permit);
assert!(pools.try_acquire("m").is_some()); }
#[tokio::test]
async fn acquire_unbounded_model_never_blocks() {
let pools = InferencePools::new(InferencePoolConfig::new()); let mut permits = Vec::new();
for _ in 0..64 {
permits.push(pools.acquire("free").await);
}
assert_eq!(permits.len(), 64);
}
#[test]
fn expect_permit_returns_ok_permit() {
let sem = Arc::new(Semaphore::new(1));
let ok = sem.clone().try_acquire_owned().unwrap();
let permit = expect_permit(Ok(ok));
drop(permit);
assert_eq!(sem.available_permits(), 1);
}
#[tokio::test]
#[should_panic(expected = "never closed")]
async fn expect_permit_panics_on_closed_semaphore() {
let sem = Arc::new(Semaphore::new(0));
sem.close();
let _ = expect_permit(sem.acquire_owned().await);
}
#[tokio::test]
async fn dropping_a_permit_frees_the_slot_and_wakes_the_driver() {
leviath_testkit::with_tracing(|| async {
let mut cfg = InferencePoolConfig::new();
cfg.set_limit("m", 1);
let wake = Arc::new(Notify::new());
let pools = InferencePools::new(cfg).with_wake(wake.clone());
let permit = pools.try_acquire("m").expect("the only slot");
assert!(pools.try_acquire("m").is_none(), "pool full");
assert!(
tokio::time::timeout(std::time::Duration::from_millis(20), wake.notified())
.await
.is_err(),
"holding a permit must not wake the driver"
);
drop(permit);
assert!(pools.try_acquire("m").is_some(), "slot freed");
tokio::time::timeout(std::time::Duration::from_millis(20), wake.notified())
.await
.expect("releasing a permit must wake the driver");
})
.await;
}
#[tokio::test]
async fn a_permit_without_a_wake_handle_still_releases() {
let mut cfg = InferencePoolConfig::new();
cfg.set_limit("m", 1);
let pools = InferencePools::new(cfg); let permit = pools.try_acquire("m").expect("the only slot");
assert!(pools.try_acquire("m").is_none());
drop(permit);
assert!(pools.try_acquire("m").is_some());
}
#[tokio::test]
async fn the_awaiting_acquire_path_also_wakes_on_release() {
let wake = Arc::new(Notify::new());
let pools = InferencePools::new(InferencePoolConfig::new()).with_wake(wake.clone());
drop(pools.acquire("m").await);
tokio::time::timeout(std::time::Duration::from_millis(20), wake.notified())
.await
.expect("an awaited permit wakes on release as well");
}
#[tokio::test]
async fn occupancy_reports_in_use_against_each_models_cap() {
let mut cfg = InferencePoolConfig::new().with_default(None); cfg.set_limit("capped", 2);
let pools = InferencePools::new(cfg);
assert!(pools.occupancy().is_empty());
let held = pools.try_acquire("capped").expect("free");
let _unbounded = pools.try_acquire("free").expect("unbounded is always free");
let occ = pools.occupancy();
assert_eq!(
occ,
vec![
PoolOccupancy {
model: "capped".to_string(),
in_use: 1,
cap: Some(2)
},
PoolOccupancy {
model: "free".to_string(),
in_use: 1,
cap: None
},
],
"sorted by model, in-use counted against the cap where there is one"
);
assert!(!occ[0].is_full(), "1 of 2 is not full");
assert!(!occ[1].is_full(), "an unbounded pool is never full");
assert_eq!(occ[0].to_string(), "capped=1/2");
assert_eq!(occ[1].to_string(), "free=1/unbounded");
let _second = pools.try_acquire("capped").expect("second of two");
assert!(pools.occupancy()[0].is_full(), "2 of 2 is full");
drop(held);
assert!(!pools.occupancy()[0].is_full(), "and not full once freed");
}
}