use crate::context::ContextQuality;
use iron_providers::TokenUsage;
#[derive(Debug, Clone, Default, PartialEq)]
pub struct TokenUsageTotals {
pub input_tokens: u64,
pub output_tokens: u64,
pub cached_input_tokens: u64,
pub cache_creation_input_tokens: u64,
pub cache_read_input_tokens: u64,
pub reasoning_output_tokens: u64,
}
#[derive(Debug, Clone, Default)]
pub struct SessionTokenTracker {
baseline_input_tokens: Option<usize>,
delta_tokens_after_baseline: usize,
accumulated_input_tokens: u64,
accumulated_output_tokens: u64,
accumulated_cached_input_tokens: u64,
accumulated_cache_creation_input_tokens: u64,
accumulated_cache_read_input_tokens: u64,
accumulated_reasoning_output_tokens: u64,
}
impl SessionTokenTracker {
pub fn record_provider_usage(&mut self, usage: &TokenUsage) {
if let Some(input) = usage.input_tokens {
self.baseline_input_tokens = Some(input as usize);
self.delta_tokens_after_baseline = 0;
self.accumulated_input_tokens = self.accumulated_input_tokens.saturating_add(input);
}
if let Some(output) = usage.output_tokens {
self.accumulated_output_tokens = self.accumulated_output_tokens.saturating_add(output);
}
if let Some(cached) = usage.cached_input_tokens {
self.accumulated_cached_input_tokens =
self.accumulated_cached_input_tokens.saturating_add(cached);
}
if let Some(creation) = usage.cache_creation_input_tokens {
self.accumulated_cache_creation_input_tokens = self
.accumulated_cache_creation_input_tokens
.saturating_add(creation);
}
if let Some(read) = usage.cache_read_input_tokens {
self.accumulated_cache_read_input_tokens = self
.accumulated_cache_read_input_tokens
.saturating_add(read);
}
if let Some(reasoning) = usage.reasoning_output_tokens {
self.accumulated_reasoning_output_tokens = self
.accumulated_reasoning_output_tokens
.saturating_add(reasoning);
}
}
pub fn add_delta(&mut self, tokens: usize) {
self.delta_tokens_after_baseline = self.delta_tokens_after_baseline.saturating_add(tokens);
}
pub fn invalidate_baseline(&mut self) {
self.baseline_input_tokens = None;
self.delta_tokens_after_baseline = 0;
}
pub fn estimate_current_context(&self) -> Option<usize> {
self.baseline_input_tokens
.map(|baseline| baseline.saturating_add(self.delta_tokens_after_baseline))
}
pub fn has_baseline(&self) -> bool {
self.baseline_input_tokens.is_some()
}
pub fn quality(&self) -> ContextQuality {
if self.baseline_input_tokens.is_some() && self.delta_tokens_after_baseline == 0 {
ContextQuality::Exact
} else {
ContextQuality::Estimated
}
}
pub fn accumulated_input_tokens(&self) -> u64 {
self.accumulated_input_tokens
}
pub fn accumulated_output_tokens(&self) -> u64 {
self.accumulated_output_tokens
}
pub fn accumulated_cached_input_tokens(&self) -> u64 {
self.accumulated_cached_input_tokens
}
pub fn accumulated_cache_creation_input_tokens(&self) -> u64 {
self.accumulated_cache_creation_input_tokens
}
pub fn accumulated_cache_read_input_tokens(&self) -> u64 {
self.accumulated_cache_read_input_tokens
}
pub fn accumulated_reasoning_output_tokens(&self) -> u64 {
self.accumulated_reasoning_output_tokens
}
pub fn accumulated_totals(&self) -> TokenUsageTotals {
TokenUsageTotals {
input_tokens: self.accumulated_input_tokens,
output_tokens: self.accumulated_output_tokens,
cached_input_tokens: self.accumulated_cached_input_tokens,
cache_creation_input_tokens: self.accumulated_cache_creation_input_tokens,
cache_read_input_tokens: self.accumulated_cache_read_input_tokens,
reasoning_output_tokens: self.accumulated_reasoning_output_tokens,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tracker_without_baseline_returns_none() {
let tracker = SessionTokenTracker::default();
assert!(!tracker.has_baseline());
assert_eq!(tracker.estimate_current_context(), None);
assert_eq!(tracker.quality(), ContextQuality::Estimated);
}
#[test]
fn tracker_records_baseline_and_accumulates_totals() {
let mut tracker = SessionTokenTracker::default();
let usage = TokenUsage {
input_tokens: Some(100),
output_tokens: Some(50),
total_tokens: None,
cached_input_tokens: Some(10),
cache_creation_input_tokens: Some(5),
cache_read_input_tokens: Some(15),
reasoning_output_tokens: Some(20),
};
tracker.record_provider_usage(&usage);
assert!(tracker.has_baseline());
assert_eq!(tracker.estimate_current_context(), Some(100));
assert_eq!(tracker.quality(), ContextQuality::Exact);
assert_eq!(tracker.accumulated_input_tokens(), 100);
assert_eq!(tracker.accumulated_output_tokens(), 50);
assert_eq!(tracker.accumulated_cached_input_tokens(), 10);
assert_eq!(tracker.accumulated_cache_creation_input_tokens(), 5);
assert_eq!(tracker.accumulated_cache_read_input_tokens(), 15);
assert_eq!(tracker.accumulated_reasoning_output_tokens(), 20);
}
#[test]
fn tracker_adds_delta_after_baseline() {
let mut tracker = SessionTokenTracker::default();
tracker.record_provider_usage(&TokenUsage {
input_tokens: Some(100),
..TokenUsage::default()
});
tracker.add_delta(25);
assert_eq!(tracker.estimate_current_context(), Some(125));
assert_eq!(tracker.quality(), ContextQuality::Estimated);
}
#[test]
fn tracker_resets_on_new_usage() {
let mut tracker = SessionTokenTracker::default();
tracker.record_provider_usage(&TokenUsage {
input_tokens: Some(100),
..TokenUsage::default()
});
tracker.add_delta(25);
tracker.record_provider_usage(&TokenUsage {
input_tokens: Some(200),
..TokenUsage::default()
});
assert_eq!(tracker.estimate_current_context(), Some(200));
assert_eq!(tracker.quality(), ContextQuality::Exact);
}
#[test]
fn tracker_invalidates_baseline() {
let mut tracker = SessionTokenTracker::default();
tracker.record_provider_usage(&TokenUsage {
input_tokens: Some(100),
..TokenUsage::default()
});
tracker.add_delta(25);
tracker.invalidate_baseline();
assert!(!tracker.has_baseline());
assert_eq!(tracker.estimate_current_context(), None);
assert_eq!(tracker.quality(), ContextQuality::Estimated);
}
#[test]
fn tracker_does_not_establish_baseline_without_input_tokens() {
let mut tracker = SessionTokenTracker::default();
tracker.record_provider_usage(&TokenUsage {
input_tokens: None,
output_tokens: Some(50),
..TokenUsage::default()
});
assert!(!tracker.has_baseline());
assert_eq!(tracker.estimate_current_context(), None);
assert_eq!(tracker.accumulated_output_tokens(), 50);
}
}