Skip to main content

llm/usage/
tokens.rs

1use schemars::JsonSchema;
2use serde::{Deserialize, Serialize};
3use std::fmt;
4use std::ops::{Add, AddAssign};
5
6/// A count of tokens. Sums saturate instead of wrapping.
7#[repr(transparent)]
8#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize, JsonSchema)]
9#[serde(transparent)]
10#[schemars(transparent)]
11pub struct Tokens(u64);
12
13impl Tokens {
14    pub const ZERO: Self = Self(0);
15
16    pub const fn new(value: u64) -> Self {
17        Self(value)
18    }
19
20    pub const fn get(self) -> u64 {
21        self.0
22    }
23
24    pub const fn is_zero(self) -> bool {
25        self.0 == 0
26    }
27
28    pub const fn saturating_add(self, rhs: Self) -> Self {
29        Self(self.0.saturating_add(rhs.0))
30    }
31
32    pub const fn saturating_sub(self, rhs: Self) -> Self {
33        Self(self.0.saturating_sub(rhs.0))
34    }
35}
36
37impl Add for Tokens {
38    type Output = Self;
39    fn add(self, rhs: Self) -> Self::Output {
40        self.saturating_add(rhs)
41    }
42}
43
44impl AddAssign for Tokens {
45    fn add_assign(&mut self, rhs: Self) {
46        *self = self.saturating_add(rhs);
47    }
48}
49
50impl From<u32> for Tokens {
51    fn from(value: u32) -> Self {
52        Self(value.into())
53    }
54}
55
56impl From<Tokens> for u64 {
57    fn from(value: Tokens) -> Self {
58        value.0
59    }
60}
61
62impl TryFrom<Tokens> for i64 {
63    type Error = std::num::TryFromIntError;
64    fn try_from(value: Tokens) -> Result<Self, Self::Error> {
65        value.0.try_into()
66    }
67}
68
69impl From<Tokens> for f64 {
70    #[allow(clippy::cast_precision_loss)]
71    fn from(value: Tokens) -> Self {
72        value.0 as f64
73    }
74}
75
76impl fmt::Display for Tokens {
77    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
78        self.0.fmt(formatter)
79    }
80}
81
82/// Token counts for a single LLM call or an aggregate of calls. Providers fill
83/// in only the dimensions they report; the rest stay `None`.
84#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
85pub struct TokenUsage {
86    pub input_tokens: Tokens,
87    pub output_tokens: Tokens,
88    #[serde(default)]
89    pub cache_read_tokens: Option<Tokens>,
90    #[serde(default)]
91    pub cache_creation_tokens: Option<Tokens>,
92    #[serde(default)]
93    pub input_audio_tokens: Option<Tokens>,
94    #[serde(default)]
95    pub input_video_tokens: Option<Tokens>,
96    #[serde(default)]
97    pub reasoning_tokens: Option<Tokens>,
98    #[serde(default)]
99    pub output_audio_tokens: Option<Tokens>,
100    #[serde(default)]
101    pub accepted_prediction_tokens: Option<Tokens>,
102    #[serde(default)]
103    pub rejected_prediction_tokens: Option<Tokens>,
104}
105
106impl TokenUsage {
107    pub fn new(input_tokens: u64, output_tokens: u64) -> Self {
108        Self { input_tokens: Tokens::new(input_tokens), output_tokens: Tokens::new(output_tokens), ..Self::default() }
109    }
110
111    pub fn is_zero(self) -> bool {
112        let reported = [
113            self.cache_read_tokens,
114            self.cache_creation_tokens,
115            self.input_audio_tokens,
116            self.input_video_tokens,
117            self.reasoning_tokens,
118            self.output_audio_tokens,
119            self.accepted_prediction_tokens,
120            self.rejected_prediction_tokens,
121        ];
122        self.total_tokens().is_zero() && reported.iter().all(|dimension| dimension.unwrap_or_default().is_zero())
123    }
124
125    pub fn total_tokens(self) -> Tokens {
126        self.input_tokens + self.output_tokens
127    }
128}
129
130impl Add for TokenUsage {
131    type Output = Self;
132    fn add(self, rhs: Self) -> Self::Output {
133        Self {
134            input_tokens: self.input_tokens + rhs.input_tokens,
135            output_tokens: self.output_tokens + rhs.output_tokens,
136            cache_read_tokens: add_reported(self.cache_read_tokens, rhs.cache_read_tokens),
137            cache_creation_tokens: add_reported(self.cache_creation_tokens, rhs.cache_creation_tokens),
138            input_audio_tokens: add_reported(self.input_audio_tokens, rhs.input_audio_tokens),
139            input_video_tokens: add_reported(self.input_video_tokens, rhs.input_video_tokens),
140            reasoning_tokens: add_reported(self.reasoning_tokens, rhs.reasoning_tokens),
141            output_audio_tokens: add_reported(self.output_audio_tokens, rhs.output_audio_tokens),
142            accepted_prediction_tokens: add_reported(self.accepted_prediction_tokens, rhs.accepted_prediction_tokens),
143            rejected_prediction_tokens: add_reported(self.rejected_prediction_tokens, rhs.rejected_prediction_tokens),
144        }
145    }
146}
147
148impl AddAssign for TokenUsage {
149    fn add_assign(&mut self, rhs: Self) {
150        *self = *self + rhs;
151    }
152}
153
154fn add_reported(lhs: Option<Tokens>, rhs: Option<Tokens>) -> Option<Tokens> {
155    match (lhs, rhs) {
156        (None, None) => None,
157        _ => Some(lhs.unwrap_or_default() + rhs.unwrap_or_default()),
158    }
159}