kcode_k1_codex_token_usage/
lib.rs1use kcode_k1_codex_conversations::{Error, ErrorKind};
2use serde_json::Value;
3
4const METHOD: &str = "thread/tokenUsage/updated";
5
6#[derive(Clone, Debug, PartialEq, Eq)]
7pub enum TokenUsageUpdated {
8 Absent,
9 Malformed(Error),
10 Valid(TokenUsage),
11}
12
13#[derive(Clone, Debug, PartialEq, Eq)]
14pub struct TokenUsageScope {
15 pub thread_id: String,
16 pub turn_id: String,
17}
18
19#[derive(Clone, Debug, PartialEq, Eq)]
20pub struct TokenUsage {
21 pub scope: TokenUsageScope,
22 pub total: TokenUsageBreakdown,
23 pub last: TokenUsageBreakdown,
24 pub model_context_window: Option<i64>,
25}
26
27#[derive(Clone, Debug, PartialEq, Eq)]
28pub struct TokenUsageBreakdown {
29 pub input_tokens: i64,
30 pub cached_input_tokens: i64,
31 pub cache_write_input_tokens: i64,
32 pub output_tokens: i64,
33 pub reasoning_output_tokens: i64,
34 pub total_tokens: i64,
35}
36
37pub fn decode_token_usage_updated(message: Value) -> TokenUsageUpdated {
38 if message.get("method").and_then(Value::as_str) != Some(METHOD) {
39 return TokenUsageUpdated::Absent;
40 }
41 match decode(&message) {
42 Ok(usage) => TokenUsageUpdated::Valid(usage),
43 Err(error) => TokenUsageUpdated::Malformed(error),
44 }
45}
46
47fn decode(message: &Value) -> Result<TokenUsage, Error> {
48 let params = object_field(message, "params")?;
49 let scope = TokenUsageScope {
50 thread_id: string_field(params, "threadId")?.to_owned(),
51 turn_id: string_field(params, "turnId")?.to_owned(),
52 };
53 let usage = object_field(params, "tokenUsage")?;
54 Ok(TokenUsage {
55 scope,
56 total: breakdown(object_field(usage, "total")?)?,
57 last: breakdown(object_field(usage, "last")?)?,
58 model_context_window: nullable_count(usage, "modelContextWindow")?,
59 })
60}
61
62fn breakdown(value: &Value) -> Result<TokenUsageBreakdown, Error> {
63 Ok(TokenUsageBreakdown {
64 input_tokens: count(value, "inputTokens")?,
65 cached_input_tokens: count(value, "cachedInputTokens")?,
66 cache_write_input_tokens: optional_count(value, "cacheWriteInputTokens")?,
67 output_tokens: count(value, "outputTokens")?,
68 reasoning_output_tokens: count(value, "reasoningOutputTokens")?,
69 total_tokens: count(value, "totalTokens")?,
70 })
71}
72
73fn object_field<'a>(value: &'a Value, field: &str) -> Result<&'a Value, Error> {
74 value
75 .get(field)
76 .filter(|field| field.is_object())
77 .ok_or_else(|| protocol(format!("token usage notification has invalid {field}")))
78}
79
80fn string_field<'a>(value: &'a Value, field: &str) -> Result<&'a str, Error> {
81 value
82 .get(field)
83 .and_then(Value::as_str)
84 .ok_or_else(|| protocol(format!("token usage notification has invalid {field}")))
85}
86
87fn count(value: &Value, field: &str) -> Result<i64, Error> {
88 value
89 .get(field)
90 .and_then(Value::as_i64)
91 .filter(|count| *count >= 0)
92 .ok_or_else(|| protocol(format!("token usage notification has invalid {field}")))
93}
94
95fn optional_count(value: &Value, field: &str) -> Result<i64, Error> {
96 if value.get(field).is_none() {
97 Ok(0)
98 } else {
99 count(value, field)
100 }
101}
102
103fn nullable_count(value: &Value, field: &str) -> Result<Option<i64>, Error> {
104 match value.get(field) {
105 Some(Value::Null) => Ok(None),
106 Some(_) => count(value, field).map(Some),
107 None => Err(protocol(format!(
108 "token usage notification has invalid {field}"
109 ))),
110 }
111}
112
113fn protocol(message: String) -> Error {
114 Error::new(ErrorKind::Protocol, message)
115}
116
117#[cfg(test)]
118mod tests {
119 use super::*;
120 use serde_json::json;
121
122 fn notification() -> Value {
123 json!({
124 "method": METHOD,
125 "params": {
126 "threadId": "thread-a",
127 "turnId": "turn-b",
128 "tokenUsage": {
129 "total": counts(10),
130 "last": counts(1),
131 "modelContextWindow": 128000
132 }
133 }
134 })
135 }
136
137 fn counts(base: i64) -> Value {
138 json!({
139 "inputTokens": base,
140 "cachedInputTokens": base + 1,
141 "cacheWriteInputTokens": base + 2,
142 "outputTokens": base + 3,
143 "reasoningOutputTokens": base + 4,
144 "totalTokens": base + 5
145 })
146 }
147
148 fn malformed(value: Value) {
149 assert!(matches!(
150 decode_token_usage_updated(value),
151 TokenUsageUpdated::Malformed(Error {
152 kind: ErrorKind::Protocol,
153 ..
154 })
155 ));
156 }
157
158 #[test]
159 fn decodes_full_notification() {
160 assert_eq!(
161 decode_token_usage_updated(notification()),
162 TokenUsageUpdated::Valid(TokenUsage {
163 scope: TokenUsageScope {
164 thread_id: "thread-a".into(),
165 turn_id: "turn-b".into(),
166 },
167 total: TokenUsageBreakdown {
168 input_tokens: 10,
169 cached_input_tokens: 11,
170 cache_write_input_tokens: 12,
171 output_tokens: 13,
172 reasoning_output_tokens: 14,
173 total_tokens: 15,
174 },
175 last: TokenUsageBreakdown {
176 input_tokens: 1,
177 cached_input_tokens: 2,
178 cache_write_input_tokens: 3,
179 output_tokens: 4,
180 reasoning_output_tokens: 5,
181 total_tokens: 6,
182 },
183 model_context_window: Some(128000),
184 })
185 );
186 }
187
188 #[test]
189 fn accepts_null_window_absent_cache_write_and_unknown_fields() {
190 let mut value = notification();
191 let usage = value.pointer_mut("/params/tokenUsage").unwrap();
192 usage["modelContextWindow"] = Value::Null;
193 usage["extra"] = json!(true);
194 for part in ["total", "last"] {
195 usage[part]
196 .as_object_mut()
197 .unwrap()
198 .remove("cacheWriteInputTokens");
199 usage[part]["extra"] = json!("ignored");
200 }
201 assert!(matches!(
202 decode_token_usage_updated(value),
203 TokenUsageUpdated::Valid(TokenUsage {
204 model_context_window: None,
205 total: TokenUsageBreakdown {
206 cache_write_input_tokens: 0,
207 ..
208 },
209 last: TokenUsageBreakdown {
210 cache_write_input_tokens: 0,
211 ..
212 },
213 ..
214 })
215 ));
216 }
217
218 #[test]
219 fn rejects_wrong_method_and_nonmatching_messages_as_absent() {
220 for value in [json!({"method": "thread/other"}), json!({}), Value::Null] {
221 assert_eq!(decode_token_usage_updated(value), TokenUsageUpdated::Absent);
222 }
223 }
224
225 #[test]
226 fn rejects_every_invalid_scope_string_class() {
227 for field in ["threadId", "turnId"] {
228 for bad in [
229 Value::Null,
230 json!(1),
231 json!(-1),
232 json!(1.5),
233 json!(u64::MAX),
234 ] {
235 let mut value = notification();
236 value["params"].as_object_mut().unwrap().remove(field);
237 malformed(value);
238 let mut value = notification();
239 value["params"][field] = bad;
240 malformed(value);
241 }
242 }
243 }
244
245 #[test]
246 fn rejects_every_invalid_required_count_class() {
247 for part in ["total", "last"] {
248 for field in [
249 "inputTokens",
250 "cachedInputTokens",
251 "outputTokens",
252 "reasoningOutputTokens",
253 "totalTokens",
254 ] {
255 for bad in [
256 Value::Null,
257 json!("1"),
258 json!(-1),
259 json!(1.5),
260 json!(u64::MAX),
261 ] {
262 let mut value = notification();
263 value["params"]["tokenUsage"][part]
264 .as_object_mut()
265 .unwrap()
266 .remove(field);
267 malformed(value);
268 let mut value = notification();
269 value["params"]["tokenUsage"][part][field] = bad;
270 malformed(value);
271 }
272 }
273 }
274 }
275
276 #[test]
277 fn rejects_every_invalid_optional_and_window_count_class() {
278 for part in ["total", "last"] {
279 for bad in [
280 Value::Null,
281 json!("1"),
282 json!(-1),
283 json!(1.5),
284 json!(u64::MAX),
285 ] {
286 let mut value = notification();
287 value["params"]["tokenUsage"][part]["cacheWriteInputTokens"] = bad;
288 malformed(value);
289 }
290 }
291 for bad in [
292 Value::Null,
293 json!("1"),
294 json!(-1),
295 json!(1.5),
296 json!(u64::MAX),
297 ] {
298 let mut value = notification();
299 value["params"]["tokenUsage"]["modelContextWindow"] = bad;
300 if value.pointer("/params/tokenUsage/modelContextWindow") == Some(&Value::Null) {
301 continue;
302 }
303 malformed(value);
304 }
305 let mut value = notification();
306 value["params"]["tokenUsage"]
307 .as_object_mut()
308 .unwrap()
309 .remove("modelContextWindow");
310 malformed(value);
311 }
312}