Skip to main content

machi_types/
usage.rs

1//! Token usage accounting.
2
3use std::ops::{Add, AddAssign};
4
5use serde::{Deserialize, Serialize};
6
7/// Prompt-side token details.
8#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
9#[non_exhaustive]
10pub struct PromptTokensDetails {
11    /// Cached prompt tokens.
12    #[serde(default)]
13    pub cached_tokens: u32,
14    /// Audio input tokens.
15    #[serde(default)]
16    pub audio_tokens: u32,
17}
18
19/// Completion-side token details.
20#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
21#[non_exhaustive]
22pub struct CompletionTokensDetails {
23    /// Reasoning tokens.
24    #[serde(default)]
25    pub reasoning_tokens: u32,
26    /// Audio output tokens.
27    #[serde(default)]
28    pub audio_tokens: u32,
29}
30
31/// Aggregated token usage for a sample or turn.
32#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
33#[non_exhaustive]
34pub struct Usage {
35    /// Input / prompt tokens.
36    #[serde(default, alias = "prompt_tokens")]
37    pub input_tokens: u32,
38    /// Output / completion tokens.
39    #[serde(default, alias = "completion_tokens")]
40    pub output_tokens: u32,
41    /// Total tokens when provided by the provider.
42    #[serde(default)]
43    pub total_tokens: u32,
44    /// Cache-read tokens (provider-specific; ledger convenience field).
45    #[serde(default)]
46    pub cache_read_tokens: u32,
47    /// Cache-creation / write tokens.
48    #[serde(default)]
49    pub cache_creation_tokens: u32,
50    /// Reasoning tokens (top-level mirror of completion details when set).
51    #[serde(default)]
52    pub reasoning_tokens: u32,
53    /// Provider API wall time for this sample, when known (milliseconds).
54    #[serde(default)]
55    pub api_duration_ms: u64,
56    /// Prompt details.
57    #[serde(default, alias = "prompt_tokens_details")]
58    pub prompt_details: PromptTokensDetails,
59    /// Completion details.
60    #[serde(default, alias = "completion_tokens_details")]
61    pub completion_details: CompletionTokensDetails,
62}
63
64impl Usage {
65    /// Zero usage.
66    #[must_use]
67    pub const fn zero() -> Self {
68        Self {
69            input_tokens: 0,
70            output_tokens: 0,
71            total_tokens: 0,
72            cache_read_tokens: 0,
73            cache_creation_tokens: 0,
74            reasoning_tokens: 0,
75            api_duration_ms: 0,
76            prompt_details: PromptTokensDetails {
77                cached_tokens: 0,
78                audio_tokens: 0,
79            },
80            completion_details: CompletionTokensDetails {
81                reasoning_tokens: 0,
82                audio_tokens: 0,
83            },
84        }
85    }
86
87    /// Construct from input/output token counts.
88    #[must_use]
89    pub const fn new(input_tokens: u32, output_tokens: u32) -> Self {
90        Self {
91            input_tokens,
92            output_tokens,
93            total_tokens: input_tokens.saturating_add(output_tokens),
94            cache_read_tokens: 0,
95            cache_creation_tokens: 0,
96            reasoning_tokens: 0,
97            api_duration_ms: 0,
98            prompt_details: PromptTokensDetails {
99                cached_tokens: 0,
100                audio_tokens: 0,
101            },
102            completion_details: CompletionTokensDetails {
103                reasoning_tokens: 0,
104                audio_tokens: 0,
105            },
106        }
107    }
108
109    /// Recompute `total_tokens` as input + output when total is zero.
110    #[must_use]
111    pub const fn normalized(mut self) -> Self {
112        if self.total_tokens == 0 {
113            self.total_tokens = self.input_tokens.saturating_add(self.output_tokens);
114        }
115        self
116    }
117}
118
119impl Add for Usage {
120    type Output = Self;
121
122    fn add(self, rhs: Self) -> Self::Output {
123        Self {
124            input_tokens: self.input_tokens.saturating_add(rhs.input_tokens),
125            output_tokens: self.output_tokens.saturating_add(rhs.output_tokens),
126            total_tokens: self.total_tokens.saturating_add(rhs.total_tokens),
127            cache_read_tokens: self.cache_read_tokens.saturating_add(rhs.cache_read_tokens),
128            cache_creation_tokens: self
129                .cache_creation_tokens
130                .saturating_add(rhs.cache_creation_tokens),
131            reasoning_tokens: self.reasoning_tokens.saturating_add(rhs.reasoning_tokens),
132            api_duration_ms: self.api_duration_ms.saturating_add(rhs.api_duration_ms),
133            prompt_details: PromptTokensDetails {
134                cached_tokens: self
135                    .prompt_details
136                    .cached_tokens
137                    .saturating_add(rhs.prompt_details.cached_tokens),
138                audio_tokens: self
139                    .prompt_details
140                    .audio_tokens
141                    .saturating_add(rhs.prompt_details.audio_tokens),
142            },
143            completion_details: CompletionTokensDetails {
144                reasoning_tokens: self
145                    .completion_details
146                    .reasoning_tokens
147                    .saturating_add(rhs.completion_details.reasoning_tokens),
148                audio_tokens: self
149                    .completion_details
150                    .audio_tokens
151                    .saturating_add(rhs.completion_details.audio_tokens),
152            },
153        }
154        .normalized()
155    }
156}
157
158impl AddAssign for Usage {
159    fn add_assign(&mut self, rhs: Self) {
160        *self = *self + rhs;
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    #[test]
169    fn add_normalizes_total() {
170        let a = Usage {
171            input_tokens: 10,
172            output_tokens: 5,
173            ..Usage::zero()
174        };
175        let b = Usage {
176            input_tokens: 1,
177            output_tokens: 1,
178            ..Usage::zero()
179        };
180        let sum = (a + b).normalized();
181        assert_eq!(sum.input_tokens, 11);
182        assert_eq!(sum.output_tokens, 6);
183        assert_eq!(sum.total_tokens, 17);
184    }
185
186    #[test]
187    fn serde_aliases() {
188        let raw = r#"{"prompt_tokens":3,"completion_tokens":4}"#;
189        let u: Usage = serde_json::from_str(raw).expect("parse");
190        assert_eq!(u.input_tokens, 3);
191        assert_eq!(u.output_tokens, 4);
192    }
193}