Skip to main content

vv_agent/types/
token_usage.rs

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}