Skip to main content

kcode_k1_codex_token_usage/
lib.rs

1use 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}