use super::trait_::{BudgetLimit, BudgetSnapshot, Governor, GovernorError, Host, Permit};
use crate::types::{BudgetScope, ProviderId};
use async_trait::async_trait;
use dashmap::DashMap;
use klieo_core::KvStore;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Semaphore;
const DEFAULT_EGRESS_RPS: u32 = 10;
const DEFAULT_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(30);
pub struct TokenBucketGovernorBuilder {
kv: Arc<dyn KvStore>,
llm_limits: Vec<(ProviderId, u32)>,
egress_default_rps: u32,
}
impl TokenBucketGovernorBuilder {
#[must_use]
pub fn llm_limit(mut self, provider: ProviderId, rps: u32) -> Self {
self.llm_limits.push((provider, rps));
self
}
#[must_use]
pub fn egress_default_rps(mut self, rps: u32) -> Self {
self.egress_default_rps = rps;
self
}
#[must_use]
pub fn build(self) -> TokenBucketGovernor {
let llm: DashMap<ProviderId, Arc<Bucket>> = DashMap::new();
for (p, rps) in &self.llm_limits {
llm.insert(p.clone(), Arc::new(Bucket::new(*rps)));
}
let g = TokenBucketGovernor {
kv: self.kv,
llm_buckets: llm,
egress_buckets: DashMap::new(),
egress_default_rps: self.egress_default_rps,
};
g.spawn_refiller();
g
}
}
struct Bucket {
sem: Arc<Semaphore>,
rps: u32,
fenced: std::sync::atomic::AtomicBool,
}
impl Bucket {
fn new(rps: u32) -> Self {
Self {
sem: Arc::new(Semaphore::new(usize::try_from(rps).unwrap_or(usize::MAX))),
rps,
fenced: std::sync::atomic::AtomicBool::new(false),
}
}
fn snapshot(&self) -> BudgetSnapshot {
BudgetSnapshot {
remaining: self.sem.available_permits() as i64,
limit: i64::from(self.rps),
}
}
}
pub struct TokenBucketGovernor {
#[allow(dead_code)]
kv: Arc<dyn KvStore>,
llm_buckets: DashMap<ProviderId, Arc<Bucket>>,
egress_buckets: DashMap<Host, Arc<Bucket>>,
egress_default_rps: u32,
}
impl TokenBucketGovernor {
#[must_use]
pub fn builder(kv: Arc<dyn KvStore>) -> TokenBucketGovernorBuilder {
TokenBucketGovernorBuilder {
kv,
llm_limits: Vec::new(),
egress_default_rps: DEFAULT_EGRESS_RPS,
}
}
fn spawn_refiller(&self) {
let llm = self.llm_buckets.clone();
let egress = self.egress_buckets.clone();
tokio::spawn(async move {
let mut tick = tokio::time::interval(Duration::from_secs(1));
loop {
tick.tick().await;
refill_map(&llm);
refill_map(&egress);
}
});
}
async fn acquire_from(
bucket: Arc<Bucket>,
scope_label: String,
) -> Result<Permit, GovernorError> {
if bucket.fenced.load(std::sync::atomic::Ordering::Acquire) {
return Err(GovernorError::Saturated { scope: scope_label });
}
let sem = Arc::clone(&bucket.sem);
let timeout = tokio::time::timeout(DEFAULT_ACQUIRE_TIMEOUT, sem.acquire_owned()).await;
match timeout {
Ok(Ok(p)) => {
p.forget();
let sem_ref = Arc::clone(&bucket.sem);
Ok(Permit::new(move || {
sem_ref.add_permits(1);
}))
}
Ok(Err(_)) => Err(GovernorError::Unavailable(format!(
"semaphore closed for {scope_label}"
))),
Err(_) => Err(GovernorError::TimedOut {
millis: DEFAULT_ACQUIRE_TIMEOUT.as_millis() as u64,
}),
}
}
}
fn refill_map<K>(map: &DashMap<K, Arc<Bucket>>)
where
K: std::hash::Hash + Eq,
{
for entry in map.iter() {
let b = entry.value();
if b.fenced.load(std::sync::atomic::Ordering::Acquire) {
continue;
}
let want = usize::try_from(b.rps).unwrap_or(usize::MAX);
let have = b.sem.available_permits();
if have < want {
b.sem.add_permits(want - have);
}
}
}
#[async_trait]
impl Governor for TokenBucketGovernor {
async fn acquire_llm(
&self,
provider: ProviderId,
_est_tokens: u32,
) -> Result<Permit, GovernorError> {
let bucket = self
.llm_buckets
.get(&provider)
.map(|e| e.value().clone())
.ok_or_else(|| GovernorError::Saturated {
scope: format!("provider:{provider}"),
})?;
Self::acquire_from(bucket, format!("llm:{provider}")).await
}
async fn acquire_egress(&self, host: Host) -> Result<Permit, GovernorError> {
let bucket = self
.egress_buckets
.entry(host.clone())
.or_insert_with(|| Arc::new(Bucket::new(self.egress_default_rps)))
.value()
.clone();
Self::acquire_from(bucket, format!("egress:{host}")).await
}
async fn budget(&self, scope: BudgetScope) -> BudgetSnapshot {
match scope {
BudgetScope::Provider(p) => self
.llm_buckets
.get(&p)
.map(|e| e.value().snapshot())
.unwrap_or(BudgetSnapshot {
remaining: 0,
limit: 0,
}),
_ => BudgetSnapshot {
remaining: 0,
limit: 0,
},
}
}
async fn fence(&self, scope: BudgetScope, _limit: BudgetLimit) -> Result<(), GovernorError> {
if let BudgetScope::Provider(p) = scope {
if let Some(b) = self.llm_buckets.get(&p) {
b.fenced.store(true, std::sync::atomic::Ordering::Release);
}
}
Ok(())
}
}