1use serde::{Deserialize, Deserializer, Serialize, Serializer};
2use serde_json::{Map, Value};
3
4pub const TOKEN_USAGE_SCHEMA_VERSION: &str = "vv-agent.token-usage.v1";
5pub const TASK_TOKEN_USAGE_SCHEMA_VERSION: &str = "vv-agent.task-token-usage.v2";
6
7#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
8#[serde(rename_all = "snake_case")]
9pub enum UsageSource {
10 ProviderReported,
11 Estimated,
12 #[default]
13 AccountingMissing,
14}
15
16#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
17#[serde(rename_all = "snake_case")]
18pub enum CacheUsageStatus {
19 ProviderReported,
20 #[default]
21 AccountingMissing,
22 Unsupported,
23}
24
25fn required_option<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
26where
27 D: Deserializer<'de>,
28 T: Deserialize<'de>,
29{
30 Option::<T>::deserialize(deserializer)
31}
32
33#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
34#[serde(deny_unknown_fields)]
35pub struct CacheUsage {
36 pub status: CacheUsageStatus,
37 #[serde(deserialize_with = "required_option")]
38 pub read_input_tokens: Option<u64>,
39 #[serde(deserialize_with = "required_option")]
40 pub write_input_tokens: Option<u64>,
41 #[serde(deserialize_with = "required_option")]
42 pub uncached_input_tokens: Option<u64>,
43 #[serde(deserialize_with = "required_option")]
44 pub source: Option<String>,
45}
46
47impl Default for CacheUsage {
48 fn default() -> Self {
49 Self {
50 status: CacheUsageStatus::AccountingMissing,
51 read_input_tokens: None,
52 write_input_tokens: None,
53 uncached_input_tokens: None,
54 source: None,
55 }
56 }
57}
58
59#[derive(Debug, Clone, PartialEq, Eq)]
60pub struct TokenUsage {
61 pub input_tokens: Option<u64>,
62 pub output_tokens: Option<u64>,
63 pub total_tokens: Option<u64>,
64 pub reasoning_tokens: Option<u64>,
65 pub usage_source: UsageSource,
66 pub cache_usage: CacheUsage,
67 pub provider_usage: Map<String, Value>,
68}
69
70impl Default for TokenUsage {
71 fn default() -> Self {
72 Self {
73 input_tokens: None,
74 output_tokens: None,
75 total_tokens: None,
76 reasoning_tokens: None,
77 usage_source: UsageSource::AccountingMissing,
78 cache_usage: CacheUsage::default(),
79 provider_usage: Map::new(),
80 }
81 }
82}
83
84impl TokenUsage {
85 pub fn has_usage(&self) -> bool {
86 self.input_tokens.is_some()
87 || self.output_tokens.is_some()
88 || self.total_tokens.is_some()
89 || self.reasoning_tokens.is_some()
90 || self.usage_source != UsageSource::AccountingMissing
91 || self.cache_usage.status != CacheUsageStatus::AccountingMissing
92 }
93}
94
95#[derive(Serialize)]
96struct TokenUsageWireRef<'a> {
97 schema_version: &'static str,
98 input_tokens: Option<u64>,
99 output_tokens: Option<u64>,
100 total_tokens: Option<u64>,
101 reasoning_tokens: Option<u64>,
102 usage_source: UsageSource,
103 cache_usage: &'a CacheUsage,
104 provider_usage: &'a Map<String, Value>,
105}
106
107#[derive(Deserialize)]
108#[serde(deny_unknown_fields)]
109struct TokenUsageWire {
110 schema_version: String,
111 #[serde(deserialize_with = "required_option")]
112 input_tokens: Option<u64>,
113 #[serde(deserialize_with = "required_option")]
114 output_tokens: Option<u64>,
115 #[serde(deserialize_with = "required_option")]
116 total_tokens: Option<u64>,
117 #[serde(deserialize_with = "required_option")]
118 reasoning_tokens: Option<u64>,
119 usage_source: UsageSource,
120 cache_usage: CacheUsage,
121 provider_usage: Map<String, Value>,
122}
123
124impl Serialize for TokenUsage {
125 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
126 where
127 S: Serializer,
128 {
129 TokenUsageWireRef {
130 schema_version: TOKEN_USAGE_SCHEMA_VERSION,
131 input_tokens: self.input_tokens,
132 output_tokens: self.output_tokens,
133 total_tokens: self.total_tokens,
134 reasoning_tokens: self.reasoning_tokens,
135 usage_source: self.usage_source,
136 cache_usage: &self.cache_usage,
137 provider_usage: &self.provider_usage,
138 }
139 .serialize(serializer)
140 }
141}
142
143impl<'de> Deserialize<'de> for TokenUsage {
144 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
145 where
146 D: Deserializer<'de>,
147 {
148 let wire = TokenUsageWire::deserialize(deserializer)?;
149 if wire.schema_version != TOKEN_USAGE_SCHEMA_VERSION {
150 return Err(serde::de::Error::custom(format!(
151 "unsupported TokenUsage schema: {:?}",
152 wire.schema_version
153 )));
154 }
155 Ok(Self {
156 input_tokens: wire.input_tokens,
157 output_tokens: wire.output_tokens,
158 total_tokens: wire.total_tokens,
159 reasoning_tokens: wire.reasoning_tokens,
160 usage_source: wire.usage_source,
161 cache_usage: wire.cache_usage,
162 provider_usage: wire.provider_usage,
163 })
164 }
165}
166
167#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
168#[serde(rename_all = "snake_case")]
169pub enum ModelCallOperation {
170 AgentCycle,
171 SessionMemory,
172 MemoryCompaction,
173}
174
175#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
176#[serde(rename_all = "snake_case")]
177pub enum ModelCallStatus {
178 Completed,
179 Failed,
180 Ambiguous,
181}
182
183#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
184pub struct ModelCallRecord {
185 pub call_id: String,
186 pub operation_id: String,
187 pub attempt: u32,
188 pub operation: ModelCallOperation,
189 pub cycle_index: u32,
190 pub backend: String,
191 pub model: String,
192 pub status: ModelCallStatus,
193 pub usage: TokenUsage,
194 pub error_code: Option<String>,
195}
196
197#[derive(Deserialize)]
198#[serde(deny_unknown_fields)]
199struct ModelCallRecordWire {
200 call_id: String,
201 operation_id: String,
202 attempt: u32,
203 operation: ModelCallOperation,
204 cycle_index: u32,
205 backend: String,
206 model: String,
207 status: ModelCallStatus,
208 usage: TokenUsage,
209 #[serde(deserialize_with = "required_option")]
210 error_code: Option<String>,
211}
212
213impl TryFrom<ModelCallRecordWire> for ModelCallRecord {
214 type Error = String;
215
216 fn try_from(wire: ModelCallRecordWire) -> Result<Self, Self::Error> {
217 for (name, value) in [
218 ("call_id", wire.call_id.as_str()),
219 ("operation_id", wire.operation_id.as_str()),
220 ("backend", wire.backend.as_str()),
221 ("model", wire.model.as_str()),
222 ] {
223 if value.trim().is_empty() {
224 return Err(format!("{name} must be a non-empty string"));
225 }
226 }
227 if wire.attempt == 0 {
228 return Err("attempt must be positive".to_string());
229 }
230 if wire.cycle_index == 0 {
231 return Err("cycle_index must be positive".to_string());
232 }
233 match wire.status {
234 ModelCallStatus::Completed if wire.error_code.is_some() => {
235 return Err("completed model calls require error_code=null".to_string())
236 }
237 ModelCallStatus::Failed | ModelCallStatus::Ambiguous
238 if wire
239 .error_code
240 .as_deref()
241 .is_none_or(|value| value.trim().is_empty()) =>
242 {
243 return Err("failed or ambiguous model calls require an error_code".to_string())
244 }
245 _ => {}
246 }
247 Ok(Self {
248 call_id: wire.call_id,
249 operation_id: wire.operation_id,
250 attempt: wire.attempt,
251 operation: wire.operation,
252 cycle_index: wire.cycle_index,
253 backend: wire.backend,
254 model: wire.model,
255 status: wire.status,
256 usage: wire.usage,
257 error_code: wire.error_code,
258 })
259 }
260}
261
262impl<'de> Deserialize<'de> for ModelCallRecord {
263 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
264 where
265 D: Deserializer<'de>,
266 {
267 ModelCallRecordWire::deserialize(deserializer)?
268 .try_into()
269 .map_err(serde::de::Error::custom)
270 }
271}
272
273#[derive(Debug, Clone, PartialEq, Eq)]
274pub struct TaskTokenUsage {
275 pub input_tokens: Option<u64>,
276 pub output_tokens: Option<u64>,
277 pub total_tokens: Option<u64>,
278 pub reasoning_tokens: Option<u64>,
279 pub cache_usage: CacheUsage,
280 pub model_calls: Vec<ModelCallRecord>,
281}
282
283impl Default for TaskTokenUsage {
284 fn default() -> Self {
285 Self {
286 input_tokens: Some(0),
287 output_tokens: Some(0),
288 total_tokens: Some(0),
289 reasoning_tokens: Some(0),
290 cache_usage: CacheUsage {
291 source: Some("aggregate".to_string()),
292 ..CacheUsage::default()
293 },
294 model_calls: Vec::new(),
295 }
296 }
297}
298
299impl TaskTokenUsage {
300 pub fn add_model_call(&mut self, model_call: ModelCallRecord) -> Result<(), String> {
301 if self
302 .model_calls
303 .iter()
304 .any(|existing| existing.call_id == model_call.call_id)
305 {
306 return Err("model_call_id_duplicate".to_string());
307 }
308 self.model_calls.push(model_call);
309 self.input_tokens = complete_sum(&self.model_calls, |usage| usage.input_tokens);
310 self.output_tokens = complete_sum(&self.model_calls, |usage| usage.output_tokens);
311 self.total_tokens = complete_sum(&self.model_calls, |usage| usage.total_tokens);
312 self.reasoning_tokens = complete_sum(&self.model_calls, |usage| usage.reasoning_tokens);
313 self.cache_usage = aggregate_cache_usage(&self.model_calls);
314 Ok(())
315 }
316}
317
318fn complete_sum(
319 model_calls: &[ModelCallRecord],
320 read: fn(&TokenUsage) -> Option<u64>,
321) -> Option<u64> {
322 if model_calls.is_empty() {
323 return Some(0);
324 }
325 model_calls.iter().try_fold(0_u64, |total, model_call| {
326 total.checked_add(read(&model_call.usage)?)
327 })
328}
329
330fn aggregate_cache_usage(model_calls: &[ModelCallRecord]) -> CacheUsage {
331 if model_calls.is_empty() {
332 return CacheUsage {
333 source: Some("aggregate".to_string()),
334 ..CacheUsage::default()
335 };
336 }
337 let status = if model_calls
338 .iter()
339 .all(|model_call| model_call.usage.cache_usage.status == CacheUsageStatus::ProviderReported)
340 {
341 CacheUsageStatus::ProviderReported
342 } else if model_calls
343 .iter()
344 .all(|model_call| model_call.usage.cache_usage.status == CacheUsageStatus::Unsupported)
345 {
346 CacheUsageStatus::Unsupported
347 } else {
348 CacheUsageStatus::AccountingMissing
349 };
350
351 let complete_cache_sum = |read: fn(&CacheUsage) -> Option<u64>| -> Option<u64> {
352 if status != CacheUsageStatus::ProviderReported {
353 return None;
354 }
355 model_calls.iter().try_fold(0_u64, |total, model_call| {
356 total.checked_add(read(&model_call.usage.cache_usage)?)
357 })
358 };
359
360 CacheUsage {
361 status,
362 read_input_tokens: complete_cache_sum(|usage| usage.read_input_tokens),
363 write_input_tokens: complete_cache_sum(|usage| usage.write_input_tokens),
364 uncached_input_tokens: complete_cache_sum(|usage| usage.uncached_input_tokens),
365 source: Some("aggregate".to_string()),
366 }
367}
368
369#[derive(Serialize)]
370struct TaskTokenUsageWireRef<'a> {
371 schema_version: &'static str,
372 input_tokens: Option<u64>,
373 output_tokens: Option<u64>,
374 total_tokens: Option<u64>,
375 reasoning_tokens: Option<u64>,
376 cache_usage: &'a CacheUsage,
377 model_calls: &'a [ModelCallRecord],
378}
379
380#[derive(Deserialize)]
381#[serde(deny_unknown_fields)]
382struct TaskTokenUsageWire {
383 schema_version: String,
384 #[serde(deserialize_with = "required_option")]
385 input_tokens: Option<u64>,
386 #[serde(deserialize_with = "required_option")]
387 output_tokens: Option<u64>,
388 #[serde(deserialize_with = "required_option")]
389 total_tokens: Option<u64>,
390 #[serde(deserialize_with = "required_option")]
391 reasoning_tokens: Option<u64>,
392 cache_usage: CacheUsage,
393 model_calls: Vec<ModelCallRecord>,
394}
395
396impl Serialize for TaskTokenUsage {
397 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
398 where
399 S: Serializer,
400 {
401 TaskTokenUsageWireRef {
402 schema_version: TASK_TOKEN_USAGE_SCHEMA_VERSION,
403 input_tokens: self.input_tokens,
404 output_tokens: self.output_tokens,
405 total_tokens: self.total_tokens,
406 reasoning_tokens: self.reasoning_tokens,
407 cache_usage: &self.cache_usage,
408 model_calls: &self.model_calls,
409 }
410 .serialize(serializer)
411 }
412}
413
414impl<'de> Deserialize<'de> for TaskTokenUsage {
415 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
416 where
417 D: Deserializer<'de>,
418 {
419 let wire = TaskTokenUsageWire::deserialize(deserializer)?;
420 if wire.schema_version != TASK_TOKEN_USAGE_SCHEMA_VERSION {
421 return Err(serde::de::Error::custom(format!(
422 "unsupported TaskTokenUsage schema: {:?}",
423 wire.schema_version
424 )));
425 }
426 let mut expected = Self::default();
427 for model_call in wire.model_calls {
428 expected
429 .add_model_call(model_call)
430 .map_err(serde::de::Error::custom)?;
431 }
432 if wire.input_tokens != expected.input_tokens
433 || wire.output_tokens != expected.output_tokens
434 || wire.total_tokens != expected.total_tokens
435 || wire.reasoning_tokens != expected.reasoning_tokens
436 || wire.cache_usage != expected.cache_usage
437 {
438 return Err(serde::de::Error::custom(
439 "TaskTokenUsage aggregate does not match model_calls",
440 ));
441 }
442 Ok(expected)
443 }
444}