1use serde::{Deserialize, Serialize};
2use serde_json::{Map, Value};
3
4#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
5#[serde(rename_all = "snake_case")]
6pub enum UsageSource {
7 ProviderReported,
8 Estimated,
9 #[default]
10 AccountingMissing,
11}
12
13#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
14#[serde(rename_all = "snake_case")]
15pub enum CacheUsageStatus {
16 ProviderReported,
17 #[default]
18 AccountingMissing,
19 Unsupported,
20}
21
22#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
23#[serde(default)]
24pub struct CacheUsage {
25 pub status: CacheUsageStatus,
26 pub read_tokens: Option<u64>,
27 pub write_tokens: Option<u64>,
28 pub uncached_input_tokens: Option<u64>,
29 pub source: Option<String>,
30}
31
32#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
33#[serde(default)]
34pub struct TokenUsage {
35 pub prompt_tokens: u64,
36 pub completion_tokens: u64,
37 pub total_tokens: u64,
38 pub cached_tokens: u64,
39 pub reasoning_tokens: u64,
40 pub input_tokens: u64,
41 pub output_tokens: u64,
42 pub cache_creation_tokens: u64,
43 pub usage_source: UsageSource,
44 pub cache_usage: CacheUsage,
45 pub raw: Value,
46}
47
48impl Default for TokenUsage {
49 fn default() -> Self {
50 Self {
51 prompt_tokens: 0,
52 completion_tokens: 0,
53 total_tokens: 0,
54 cached_tokens: 0,
55 reasoning_tokens: 0,
56 input_tokens: 0,
57 output_tokens: 0,
58 cache_creation_tokens: 0,
59 usage_source: UsageSource::AccountingMissing,
60 cache_usage: CacheUsage::default(),
61 raw: Value::Object(Map::new()),
62 }
63 }
64}
65
66impl TokenUsage {
67 pub fn has_usage(&self) -> bool {
68 self.prompt_tokens > 0
69 || self.completion_tokens > 0
70 || self.total_tokens > 0
71 || self.cached_tokens > 0
72 || self.reasoning_tokens > 0
73 || self.input_tokens > 0
74 || self.output_tokens > 0
75 || self.cache_creation_tokens > 0
76 || self.usage_source != UsageSource::AccountingMissing
77 || self.cache_usage.status != CacheUsageStatus::AccountingMissing
78 }
79}
80
81#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
82pub struct CycleTokenUsage {
83 pub cycle_index: u32,
84 pub usage: TokenUsage,
85}
86
87#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
88#[serde(default)]
89pub struct TaskTokenUsage {
90 pub prompt_tokens: u64,
91 pub completion_tokens: u64,
92 pub total_tokens: u64,
93 pub cached_tokens: u64,
94 pub reasoning_tokens: u64,
95 pub input_tokens: u64,
96 pub output_tokens: u64,
97 pub cache_creation_tokens: u64,
98 pub cache_usage: CacheUsage,
99 pub cycles: Vec<CycleTokenUsage>,
100}
101
102impl TaskTokenUsage {
103 pub fn add_cycle(&mut self, cycle_index: u32, usage: TokenUsage) {
104 if !usage.has_usage() {
105 return;
106 }
107 self.prompt_tokens += usage.prompt_tokens;
108 self.completion_tokens += usage.completion_tokens;
109 self.total_tokens += usage.total_tokens;
110 self.cached_tokens += usage.cached_tokens;
111 self.reasoning_tokens += usage.reasoning_tokens;
112 self.input_tokens += usage.input_tokens;
113 self.output_tokens += usage.output_tokens;
114 self.cache_creation_tokens += usage.cache_creation_tokens;
115 self.cycles.push(CycleTokenUsage { cycle_index, usage });
116 self.cache_usage = aggregate_cache_usage(&self.cycles);
117 }
118}
119
120fn aggregate_cache_usage(cycles: &[CycleTokenUsage]) -> CacheUsage {
121 if cycles.is_empty() {
122 return CacheUsage::default();
123 }
124 let status = if cycles
125 .iter()
126 .all(|cycle| cycle.usage.cache_usage.status == CacheUsageStatus::ProviderReported)
127 {
128 CacheUsageStatus::ProviderReported
129 } else if cycles
130 .iter()
131 .all(|cycle| cycle.usage.cache_usage.status == CacheUsageStatus::Unsupported)
132 {
133 CacheUsageStatus::Unsupported
134 } else {
135 CacheUsageStatus::AccountingMissing
136 };
137
138 let complete_sum = |read: fn(&CacheUsage) -> Option<u64>| -> Option<u64> {
139 if status != CacheUsageStatus::ProviderReported {
140 return None;
141 }
142 cycles.iter().try_fold(0_u64, |total, cycle| {
143 total.checked_add(read(&cycle.usage.cache_usage)?)
144 })
145 };
146
147 CacheUsage {
148 status,
149 read_tokens: complete_sum(|usage| usage.read_tokens),
150 write_tokens: complete_sum(|usage| usage.write_tokens),
151 uncached_input_tokens: complete_sum(|usage| usage.uncached_input_tokens),
152 source: Some("aggregate".to_string()),
153 }
154}