use async_trait::async_trait;
use lc_core::token_counter::{CharRatioCounter, TokenCounter};
use lc_schema::Message;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use super::{AgentHook, CompletionAction, CompletionContext, CompletionResult, HookError};
pub struct TokenBudgetHook {
budget: usize,
max_calls: Option<usize>,
tokens_used: AtomicUsize,
calls: AtomicUsize,
counter: Option<Arc<dyn TokenCounter>>,
}
impl TokenBudgetHook {
pub fn new(budget: usize) -> Self {
Self {
budget,
max_calls: None,
tokens_used: AtomicUsize::new(0),
calls: AtomicUsize::new(0),
counter: None,
}
}
pub fn with_max_calls(mut self, max_calls: usize) -> Self {
self.max_calls = Some(max_calls);
self
}
pub fn with_counter(mut self, counter: Arc<dyn TokenCounter>) -> Self {
self.counter = Some(counter);
self
}
pub fn budget(&self) -> usize {
self.budget
}
pub fn tokens_used(&self) -> usize {
self.tokens_used.load(Ordering::SeqCst)
}
pub fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
pub fn remaining(&self) -> usize {
self.budget.saturating_sub(self.tokens_used())
}
fn estimate_messages(&self, messages: &[Message]) -> usize {
match &self.counter {
Some(c) => c.count_messages(messages) as usize,
None => CharRatioCounter::new(4).count_messages(messages) as usize,
}
}
}
#[async_trait]
impl AgentHook for TokenBudgetHook {
fn on_before_completion(&self, ctx: &mut CompletionContext) -> CompletionAction {
if let Some(max) = self.max_calls {
if self.calls.load(Ordering::SeqCst) >= max {
return CompletionAction::Reject {
reason: format!("LLM call quota exceeded: max_calls={max}"),
};
}
}
let used = self.tokens_used.load(Ordering::SeqCst);
let estimate = self.estimate_messages(&ctx.messages);
if used.saturating_add(estimate) > self.budget {
return CompletionAction::Reject {
reason: format!(
"token budget exceeded: budget={}, used={used}, estimate={estimate}",
self.budget
),
};
}
self.calls.fetch_add(1, Ordering::SeqCst);
CompletionAction::Continue
}
fn on_after_completion(&self, ctx: &mut CompletionResult) -> Result<(), HookError> {
if let Some(usage) = &ctx.tokens_used {
self.tokens_used
.fetch_add(usage.total_tokens, Ordering::SeqCst);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use lc_core::language_models::TokenUsage;
fn completion_ctx(text: &str) -> CompletionContext {
CompletionContext {
messages: vec![Message::human(text.to_string())],
model: "mock".to_string(),
metadata: std::collections::HashMap::new(),
}
}
#[test]
fn test_allows_within_budget() {
let hook = TokenBudgetHook::new(1_000);
assert!(matches!(
hook.on_before_completion(&mut completion_ctx("short")),
CompletionAction::Continue
));
assert_eq!(hook.calls(), 1);
}
#[test]
fn test_rejects_when_estimated_over_budget() {
let hook = TokenBudgetHook::new(0);
assert!(matches!(
hook.on_before_completion(&mut completion_ctx("x")),
CompletionAction::Reject { .. }
));
}
#[test]
fn test_accumulates_real_usage_after_completion() {
let hook = TokenBudgetHook::new(1_000);
let mut result = CompletionResult {
message: Message::ai("hi"),
tokens_used: Some(TokenUsage {
prompt_tokens: 10,
completion_tokens: 20,
total_tokens: 30,
}),
};
hook.on_after_completion(&mut result).unwrap();
assert_eq!(hook.tokens_used(), 30);
assert_eq!(hook.remaining(), 970);
}
#[test]
fn test_max_calls_quota() {
let hook = TokenBudgetHook::new(1_000).with_max_calls(2);
assert!(matches!(
hook.on_before_completion(&mut completion_ctx("a")),
CompletionAction::Continue
));
assert!(matches!(
hook.on_before_completion(&mut completion_ctx("b")),
CompletionAction::Continue
));
let action = hook.on_before_completion(&mut completion_ctx("c"));
match action {
CompletionAction::Reject { reason } => assert!(reason.contains("quota"), "{reason}"),
other => panic!("expected Reject, got {:?}", other),
}
assert_eq!(hook.calls(), 2);
}
#[test]
fn test_rejects_after_real_usage_exceeds_budget() {
let hook = TokenBudgetHook::new(100);
let mut result = CompletionResult {
message: Message::ai("hi"),
tokens_used: Some(TokenUsage {
prompt_tokens: 0,
completion_tokens: 90,
total_tokens: 90,
}),
};
hook.on_after_completion(&mut result).unwrap();
let action = hook.on_before_completion(&mut completion_ctx("a long enough message"));
match action {
CompletionAction::Reject { reason } => {
assert!(reason.contains("budget"), "{reason}")
}
other => panic!("expected Reject, got {:?}", other),
}
}
}