use std::collections::HashMap;
use std::sync::{Arc, Mutex, PoisonError};
use tokio::sync::{AcquireError, 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>>>,
}
impl InferencePools {
pub fn new(config: InferencePoolConfig) -> Self {
Self {
config,
semaphores: Mutex::new(HashMap::new()),
}
}
pub async fn acquire(&self, model: &str) -> InferencePermit {
let semaphore = self.semaphore_for(model);
let permit = expect_permit(semaphore.acquire_owned().await);
InferencePermit { _permit: permit }
}
pub fn try_acquire(&self, model: &str) -> Option<InferencePermit> {
let semaphore = self.semaphore_for(model);
match semaphore.try_acquire_owned() {
Ok(permit) => Some(InferencePermit { _permit: permit }),
Err(_) => None,
}
}
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
}
}
fn expect_permit(result: Result<OwnedSemaphorePermit, AcquireError>) -> OwnedSemaphorePermit {
result.expect("inference pool semaphore is never closed")
}
#[derive(Debug)]
pub struct InferencePermit {
_permit: OwnedSemaphorePermit,
}
#[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);
}
}