use std::sync::Mutex;
use std::time::Duration;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Containment {
pub max_total_agents: u32,
pub max_concurrent: u32,
pub max_depth: u32,
pub max_total_tokens: u64,
#[serde(default)]
pub max_total_cost: Option<u64>,
#[serde(default)]
pub max_total_duration: Option<Duration>,
}
impl Containment {
pub fn new(
max_total_agents: u32,
max_concurrent: u32,
max_depth: u32,
max_total_tokens: u64,
) -> Self {
Self {
max_total_agents,
max_concurrent,
max_depth,
max_total_tokens,
max_total_cost: None,
max_total_duration: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SpawnRefusal {
AgentCap { max: u32 },
DepthCap { max: u32, requested: u32 },
BudgetExhausted,
}
impl SpawnRefusal {
pub fn cap(&self) -> &'static str {
match self {
SpawnRefusal::AgentCap { .. } => "agents",
SpawnRefusal::DepthCap { .. } => "depth",
SpawnRefusal::BudgetExhausted => "budget",
}
}
}
impl std::fmt::Display for SpawnRefusal {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SpawnRefusal::AgentCap { max } => {
write!(f, "agent cap reached ({max} agents)")
}
SpawnRefusal::DepthCap { max, requested } => {
write!(f, "depth cap reached (max {max}, requested {requested})")
}
SpawnRefusal::BudgetExhausted => write!(f, "the tree's spend ceiling is exhausted"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Draw {
Ok,
Halted,
}
#[derive(Debug)]
pub struct Ledger {
max_total_tokens: u64,
max_total_agents: u32,
max_depth: u32,
state: Mutex<State>,
}
#[derive(Debug)]
struct State {
spent_tokens: u64,
agents: u32,
}
impl Ledger {
pub fn new(c: &Containment) -> Self {
Self {
max_total_tokens: c.max_total_tokens,
max_total_agents: c.max_total_agents,
max_depth: c.max_depth,
state: Mutex::new(State {
spent_tokens: 0,
agents: 1, }),
}
}
pub fn from_state(c: &Containment, spent_tokens: u64, agents: u32) -> Self {
Self {
max_total_tokens: c.max_total_tokens,
max_total_agents: c.max_total_agents,
max_depth: c.max_depth,
state: Mutex::new(State {
spent_tokens,
agents: agents.max(1),
}),
}
}
pub fn remaining_tokens(&self) -> u64 {
let s = self.state.lock().unwrap();
self.max_total_tokens.saturating_sub(s.spent_tokens)
}
pub fn effective_token_budget(&self, contract_max: Option<u64>) -> u64 {
let remaining = self.remaining_tokens();
remaining.min(contract_max.unwrap_or(u64::MAX))
}
pub fn draw_tokens(&self, tokens: u64) -> Draw {
let mut s = self.state.lock().unwrap();
let next = s.spent_tokens.saturating_add(tokens);
if next > self.max_total_tokens {
Draw::Halted
} else {
s.spent_tokens = next;
Draw::Ok
}
}
pub fn spent_tokens(&self) -> u64 {
self.state.lock().unwrap().spent_tokens
}
pub fn register_agent(&self, depth: u32) -> std::result::Result<(), SpawnRefusal> {
if depth > self.max_depth {
return Err(SpawnRefusal::DepthCap {
max: self.max_depth,
requested: depth,
});
}
if self.remaining_tokens() == 0 {
return Err(SpawnRefusal::BudgetExhausted);
}
let mut s = self.state.lock().unwrap();
if s.agents >= self.max_total_agents {
return Err(SpawnRefusal::AgentCap {
max: self.max_total_agents,
});
}
s.agents += 1;
Ok(())
}
pub fn agents(&self) -> u32 {
self.state.lock().unwrap().agents
}
}
#[cfg(test)]
mod tests {
use super::*;
fn containment() -> Containment {
Containment::new(10, 4, 3, 100)
}
#[test]
fn the_ceiling_is_tree_wide_not_per_agent() {
let led = Ledger::new(&containment());
let per_child_contract_budget = 50;
assert!(40 < per_child_contract_budget);
assert_eq!(led.draw_tokens(40), Draw::Ok); assert_eq!(led.draw_tokens(40), Draw::Ok); assert_eq!(led.draw_tokens(40), Draw::Halted); assert_eq!(led.spent_tokens(), 80);
assert!(led.spent_tokens() <= 100);
}
#[test]
fn a_contract_cannot_raise_the_ceiling() {
let led = Ledger::new(&containment()); assert_eq!(led.effective_token_budget(Some(500)), 100);
led.draw_tokens(70);
assert_eq!(led.remaining_tokens(), 30);
assert_eq!(led.effective_token_budget(Some(500)), 30);
assert_eq!(led.effective_token_budget(Some(10)), 10);
assert_eq!(led.effective_token_budget(None), 30);
}
#[test]
fn concurrent_draws_never_overspend_the_ceiling() {
use std::sync::Arc;
use std::thread;
let led = Arc::new(Ledger::new(&Containment::new(10_000, 64, 3, 1_000)));
let mut handles = Vec::new();
for _ in 0..64 {
let l = Arc::clone(&led);
handles.push(thread::spawn(move || {
let mut ok = 0u64;
for _ in 0..100 {
if l.draw_tokens(10) == Draw::Ok {
ok += 10;
}
}
ok
}));
}
let granted: u64 = handles.into_iter().map(|h| h.join().unwrap()).sum();
assert!(led.spent_tokens() <= 1_000, "never exceeds the ceiling");
assert_eq!(granted, led.spent_tokens());
}
#[test]
fn serde_roundtrips_and_is_stable() {
let c = Containment {
max_total_cost: Some(500),
max_total_duration: Some(Duration::from_secs(3600)),
..containment()
};
let json = serde_json::to_string(&c).unwrap();
let back: Containment = serde_json::from_str(&json).unwrap();
assert_eq!(c, back);
}
}