use chrono::{DateTime, Utc};
use turnframe_core::error::{OrchestratorError, PolicyError};
use turnframe_provider::fallback::{AttemptOutcome, ProviderAttempt};
use crate::config::ResourceBudget;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
#[non_exhaustive]
pub enum BudgetLimit {
ModelCalls,
PromptTokens,
WallClock,
}
impl BudgetLimit {
pub const ALL: [Self; 3] = [Self::ModelCalls, Self::PromptTokens, Self::WallClock];
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::ModelCalls => "model_calls",
Self::PromptTokens => "prompt_tokens",
Self::WallClock => "wall_clock",
}
}
#[must_use]
pub fn into_error(self) -> OrchestratorError {
OrchestratorError::Policy(PolicyError::BudgetExhausted {
limit: self.as_str().to_owned(),
})
}
}
impl std::fmt::Display for BudgetLimit {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct BudgetSpend {
pub model_calls: u64,
pub prompt_tokens: u64,
}
impl BudgetSpend {
#[must_use]
pub const fn none() -> Self {
Self {
model_calls: 0,
prompt_tokens: 0,
}
}
#[must_use]
pub const fn with_model_calls(mut self, model_calls: u64) -> Self {
self.model_calls = model_calls;
self
}
#[must_use]
pub const fn with_prompt_tokens(mut self, prompt_tokens: u64) -> Self {
self.prompt_tokens = prompt_tokens;
self
}
#[must_use]
pub fn of(attempts: &[ProviderAttempt]) -> Self {
let mut spend = Self::none();
for attempt in attempts {
if attempt.outcome == AttemptOutcome::Cancelled {
continue;
}
spend.model_calls = spend.model_calls.saturating_add(1);
spend.prompt_tokens = spend
.prompt_tokens
.saturating_add(attempt.input_tokens.unwrap_or(0));
}
spend
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TurnBudget {
budget: ResourceBudget,
started_at: DateTime<Utc>,
}
impl TurnBudget {
#[must_use]
pub const fn new(budget: ResourceBudget, started_at: DateTime<Utc>) -> Self {
Self { budget, started_at }
}
#[must_use]
pub const fn budget(&self) -> &ResourceBudget {
&self.budget
}
#[must_use]
pub const fn started_at(&self) -> DateTime<Utc> {
self.started_at
}
#[must_use]
pub fn exhausted(&self, spent: BudgetSpend, now: DateTime<Utc>) -> Option<BudgetLimit> {
if spent.model_calls >= u64::from(self.budget.max_model_calls) {
return Some(BudgetLimit::ModelCalls);
}
if spent.prompt_tokens >= self.budget.max_prompt_tokens {
return Some(BudgetLimit::PromptTokens);
}
let elapsed = (now - self.started_at).to_std().unwrap_or_default();
if elapsed >= self.budget.max_wall_clock {
return Some(BudgetLimit::WallClock);
}
None
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use turnframe_provider::fallback::FallbackStage;
use turnframe_provider::ids::{AttemptNumber, ModelRef, RequestId};
use turnframe_provider::purpose::ModelPurpose;
use super::*;
fn attempt(outcome: AttemptOutcome, input_tokens: Option<u64>) -> ProviderAttempt {
ProviderAttempt {
attempt: AttemptNumber::FIRST,
request_id: RequestId::nil(),
purpose: ModelPurpose::Extract,
stage: FallbackStage::PreCommit,
model: ModelRef::new("p", "m"),
outcome,
class: None,
latency: Duration::ZERO,
input_tokens,
output_tokens: None,
temperature: None,
finish_reasons: Vec::new(),
}
}
fn started() -> DateTime<Utc> {
DateTime::from_timestamp(1_700_000_000, 0).expect("a valid fixed instant")
}
#[test]
fn every_attempt_is_a_call_and_only_reported_tokens_count() {
let spend = BudgetSpend::of(&[
attempt(
AttemptOutcome::Retried {
code: "timeout".to_owned(),
},
None,
),
attempt(AttemptOutcome::Succeeded, Some(120)),
attempt(AttemptOutcome::Cancelled, Some(999)),
]);
assert_eq!(
spend.model_calls, 2,
"a retry cost a call; a cancel did not"
);
assert_eq!(spend.prompt_tokens, 120);
assert_eq!(BudgetSpend::of(&[]), BudgetSpend::none());
}
#[test]
fn each_bound_reports_itself_and_the_order_is_fixed() {
let budget = ResourceBudget::conservative()
.with_max_model_calls(4)
.with_max_prompt_tokens(100)
.with_max_wall_clock(Duration::from_secs(30));
let turn = TurnBudget::new(budget, started());
assert_eq!(turn.budget(), &budget);
assert_eq!(turn.started_at(), started());
assert_eq!(turn.exhausted(BudgetSpend::none(), started()), None);
assert_eq!(
turn.exhausted(BudgetSpend::none().with_model_calls(4), started()),
Some(BudgetLimit::ModelCalls)
);
assert_eq!(
turn.exhausted(BudgetSpend::none().with_prompt_tokens(100), started()),
Some(BudgetLimit::PromptTokens)
);
assert_eq!(
turn.exhausted(
BudgetSpend::none(),
started() + chrono::Duration::seconds(30)
),
Some(BudgetLimit::WallClock)
);
assert_eq!(
turn.exhausted(
BudgetSpend::none()
.with_model_calls(9)
.with_prompt_tokens(999),
started() + chrono::Duration::seconds(99)
),
Some(BudgetLimit::ModelCalls)
);
assert_eq!(
turn.exhausted(
BudgetSpend::none(),
started() - chrono::Duration::seconds(5)
),
None
);
}
#[test]
fn a_limit_names_itself_in_the_error_it_raises() {
for limit in BudgetLimit::ALL {
let error = limit.into_error();
assert!(
matches!(
&error,
OrchestratorError::Policy(PolicyError::BudgetExhausted { limit: name })
if name == limit.as_str()
),
"{limit} did not name itself: {error}"
);
assert_eq!(limit.to_string(), limit.as_str());
}
}
}