use std::sync::Mutex;
use std::sync::atomic::{AtomicU8, AtomicU32, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use serde::{Deserialize, Serialize};
use tokio::sync::{Semaphore, SemaphorePermit};
use turnframe_core::replay::BudgetReport;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct Budget {
pub max_model_calls: Option<u32>,
pub max_chain_depth: Option<u8>,
pub max_parallel: usize,
pub max_prompt_tokens: Option<u64>,
pub max_wall_clock_secs: Option<u64>,
pub per_call_timeout_secs: u64,
}
impl Budget {
#[must_use]
pub const fn understanding() -> Self {
Self {
max_model_calls: Some(32),
max_chain_depth: Some(8),
max_parallel: 6,
max_prompt_tokens: Some(100_000),
max_wall_clock_secs: Some(30),
per_call_timeout_secs: 20,
}
}
#[must_use]
pub const fn narration() -> Self {
Self {
max_model_calls: Some(12),
max_chain_depth: Some(4),
max_parallel: 4,
max_prompt_tokens: Some(40_000),
max_wall_clock_secs: Some(20),
per_call_timeout_secs: 20,
}
}
#[must_use]
pub const fn unbounded() -> Self {
Self {
max_model_calls: None,
max_chain_depth: None,
max_parallel: 6,
max_prompt_tokens: None,
max_wall_clock_secs: None,
per_call_timeout_secs: 60,
}
}
}
impl Default for Budget {
fn default() -> Self {
Self::understanding()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum BudgetBound {
ModelCalls,
ChainDepth,
PromptTokens,
WallClock,
}
impl BudgetBound {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::ModelCalls => "model_calls",
Self::ChainDepth => "chain_depth",
Self::PromptTokens => "prompt_tokens",
Self::WallClock => "wall_clock",
}
}
}
#[derive(Debug)]
pub struct BudgetTracker {
budget: Budget,
started: Instant,
calls: AtomicU32,
tokens: AtomicU64,
max_depth: AtomicU8,
exhausted: Mutex<Option<BudgetBound>>,
permits: Semaphore,
}
impl BudgetTracker {
#[must_use]
pub fn new(budget: Budget) -> Self {
Self {
budget,
started: Instant::now(),
calls: AtomicU32::new(0),
tokens: AtomicU64::new(0),
max_depth: AtomicU8::new(0),
exhausted: Mutex::new(None),
permits: Semaphore::new(budget.max_parallel.max(1)),
}
}
#[must_use]
pub const fn budget(&self) -> &Budget {
&self.budget
}
pub fn reserve(&self, depth: u8) -> Result<(), BudgetBound> {
let refused = self.first_bound(depth);
if let Some(bound) = refused {
self.exhaust(bound);
return Err(bound);
}
let limit = self.budget.max_model_calls.unwrap_or(u32::MAX);
let reserved = self
.calls
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |calls| {
(calls < limit).then_some(calls + 1)
});
if reserved.is_err() {
self.exhaust(BudgetBound::ModelCalls);
return Err(BudgetBound::ModelCalls);
}
self.max_depth.fetch_max(depth, Ordering::SeqCst);
Ok(())
}
fn first_bound(&self, depth: u8) -> Option<BudgetBound> {
if self.remaining_wall_clock() == Some(Duration::ZERO) {
return Some(BudgetBound::WallClock);
}
if self.budget.max_chain_depth.is_some_and(|max| depth > max) {
return Some(BudgetBound::ChainDepth);
}
if self
.budget
.max_prompt_tokens
.is_some_and(|max| self.tokens.load(Ordering::SeqCst) >= max)
{
return Some(BudgetBound::PromptTokens);
}
None
}
fn exhaust(&self, bound: BudgetBound) {
let mut exhausted = self
.exhausted
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
exhausted.get_or_insert(bound);
}
pub async fn permit(&self) -> Option<SemaphorePermit<'_>> {
self.permits.acquire().await.ok()
}
pub fn record_tokens(&self, tokens: u64) {
self.tokens.fetch_add(tokens, Ordering::SeqCst);
}
#[must_use]
pub fn remaining_wall_clock(&self) -> Option<Duration> {
self.budget
.max_wall_clock_secs
.map(|secs| Duration::from_secs(secs).saturating_sub(self.started.elapsed()))
}
#[must_use]
pub fn call_timeout(&self, requested: Option<Duration>) -> Duration {
let per_call = Duration::from_secs(self.budget.per_call_timeout_secs);
let wanted = requested.map_or(per_call, |requested| requested.min(per_call));
match self.remaining_wall_clock() {
Some(left) => wanted.min(left),
None => wanted,
}
}
#[must_use]
pub fn exhausted(&self) -> Option<BudgetBound> {
*self
.exhausted
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[must_use]
pub fn report(&self) -> BudgetReport {
let mut report = BudgetReport::default();
report.model_calls = self.calls.load(Ordering::SeqCst);
report.prompt_tokens = self.tokens.load(Ordering::SeqCst);
report.max_depth = self.max_depth.load(Ordering::SeqCst);
report.exhausted = self.exhausted().map(|bound| bound.as_str().to_owned());
report
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn calls_are_refused_once_the_bound_is_reached_and_the_bound_is_reported() {
let tracker = BudgetTracker::new(Budget {
max_model_calls: Some(2),
..Budget::understanding()
});
assert!(tracker.reserve(1).is_ok());
assert!(tracker.reserve(2).is_ok());
assert_eq!(tracker.reserve(1), Err(BudgetBound::ModelCalls));
let report = tracker.report();
assert_eq!(report.model_calls, 2);
assert_eq!(report.max_depth, 2);
assert_eq!(report.exhausted.as_deref(), Some("model_calls"));
}
#[test]
fn a_chain_too_deep_is_refused_before_a_call_is_counted() {
let tracker = BudgetTracker::new(Budget::understanding());
assert_eq!(tracker.reserve(9), Err(BudgetBound::ChainDepth));
assert_eq!(tracker.report().model_calls, 0);
}
#[test]
fn reported_tokens_stop_the_next_call() {
let tracker = BudgetTracker::new(Budget {
max_prompt_tokens: Some(100),
..Budget::understanding()
});
tracker.record_tokens(100);
assert_eq!(tracker.reserve(1), Err(BudgetBound::PromptTokens));
}
#[test]
fn unbounded_is_a_choice_that_still_keeps_a_deadline() {
let tracker = BudgetTracker::new(Budget::unbounded());
for _ in 0..100 {
assert!(tracker.reserve(200).is_ok());
}
assert_eq!(tracker.call_timeout(None), Duration::from_secs(60));
}
}