Skip to main content

magi_code/context/
tokens.rs

1use std::{collections::BTreeMap, sync::LazyLock};
2
3use crate::{
4    output::ContextUsageSource,
5    providers::{ChatMessage, ProviderConversationItem, ProviderRequest, Usage},
6};
7use schemars::JsonSchema;
8use serde::{Deserialize, Serialize};
9use tiktoken_rs::{CoreBPE, cl100k_base, o200k_base};
10
11static O200K_BASE: LazyLock<CoreBPE> =
12    LazyLock::new(|| o200k_base().expect("embedded o200k_base tokenizer table must load"));
13static CL100K_BASE: LazyLock<CoreBPE> =
14    LazyLock::new(|| cl100k_base().expect("embedded cl100k_base tokenizer table must load"));
15
16#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
17pub struct ContextBudget {
18    #[serde(default = "default_context_enabled")]
19    pub enabled: bool,
20    #[serde(default = "default_max_tokens")]
21    pub max_tokens: usize,
22    #[serde(default = "default_reserve_tokens")]
23    pub reserve_tokens: usize,
24    #[serde(default = "default_keep_recent_tokens")]
25    pub keep_recent_tokens: usize,
26    #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
27    pub model_overrides: BTreeMap<String, ContextBudgetOverride>,
28}
29
30#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
31pub struct ContextBudgetOverride {
32    #[serde(default, skip_serializing_if = "Option::is_none")]
33    pub max_tokens: Option<usize>,
34    #[serde(default, skip_serializing_if = "Option::is_none")]
35    pub reserve_tokens: Option<usize>,
36    #[serde(default, skip_serializing_if = "Option::is_none")]
37    pub keep_recent_tokens: Option<usize>,
38}
39
40impl Default for ContextBudget {
41    fn default() -> Self {
42        Self {
43            enabled: true,
44            max_tokens: default_max_tokens(),
45            reserve_tokens: default_reserve_tokens(),
46            keep_recent_tokens: default_keep_recent_tokens(),
47            model_overrides: BTreeMap::new(),
48        }
49    }
50}
51
52fn default_context_enabled() -> bool {
53    true
54}
55fn default_max_tokens() -> usize {
56    128_000
57}
58fn default_reserve_tokens() -> usize {
59    16_384
60}
61fn default_keep_recent_tokens() -> usize {
62    20_000
63}
64
65impl ContextBudgetOverride {
66    pub(crate) fn is_empty(&self) -> bool {
67        self.max_tokens.is_none()
68            && self.reserve_tokens.is_none()
69            && self.keep_recent_tokens.is_none()
70    }
71}
72
73impl ContextBudget {
74    pub(crate) fn threshold_tokens(&self) -> usize {
75        self.max_tokens.saturating_sub(self.reserve_tokens)
76    }
77
78    pub(crate) fn apply_model_override(&mut self, provider: &str, model: &str) {
79        let key = format!("{provider}/{model}");
80        let Some(model_override) = self.model_overrides.get(&key) else {
81            return;
82        };
83        if let Some(max_tokens) = model_override.max_tokens {
84            self.max_tokens = max_tokens;
85        }
86        if let Some(reserve_tokens) = model_override.reserve_tokens {
87            self.reserve_tokens = reserve_tokens;
88        }
89        if let Some(keep_recent_tokens) = model_override.keep_recent_tokens {
90            self.keep_recent_tokens = keep_recent_tokens;
91        }
92    }
93}
94
95pub(crate) fn estimate_text_tokens(text: &str) -> usize {
96    text.chars().count().div_ceil(4).max(1)
97}
98
99pub fn estimate_messages_tokens(messages: &[ChatMessage]) -> usize {
100    messages
101        .iter()
102        .map(|message| estimate_text_tokens(&message.content) + 4)
103        .sum()
104}
105
106pub fn usage_input_tokens(usage: &Usage) -> usize {
107    // ponytail: usize::MAX saturation assumes 64-bit targets; on 32-bit, provider token counts exceeding u32::MAX would over-trigger compaction. Upgrade to u64 budget math if exact >usize accounting becomes required.
108    usize::try_from(usage.input).unwrap_or(usize::MAX)
109}
110
111#[derive(Debug, Clone, Copy, PartialEq, Eq)]
112pub(crate) struct ContextTokenCount {
113    pub(crate) tokens: usize,
114    pub(crate) source: ContextUsageSource,
115}
116
117#[cfg(test)]
118pub(crate) fn estimate_provider_request_input_tokens(request: &ProviderRequest) -> usize {
119    project_provider_request_input_tokens("local-ai", request).tokens
120}
121
122pub(crate) fn project_provider_request_input_tokens(
123    provider_id: &str,
124    request: &ProviderRequest,
125) -> ContextTokenCount {
126    let mut projection = project_provider_conversation_item_tokens(
127        provider_id,
128        &request.model,
129        request.conversation_items_iter(),
130    );
131    if let Some(tool_definitions) = request.tool_definitions_json_if_enabled() {
132        let tool_tokens =
133            if let Some(bpe) = tokenizer_for_provider_model(provider_id, &request.model) {
134                count_json_value_tokens(&BpeCounter { bpe }, &tool_definitions)
135            } else {
136                count_json_value_tokens(&FallbackCounter, &tool_definitions)
137            };
138        projection.tokens = projection.tokens.saturating_add(tool_tokens);
139    }
140    projection
141}
142
143pub(crate) fn project_provider_conversation_items_tokens(
144    provider_id: &str,
145    model: &str,
146    items: &[ProviderConversationItem],
147) -> ContextTokenCount {
148    project_provider_conversation_item_tokens(provider_id, model, items.iter())
149}
150
151fn project_provider_conversation_item_tokens<'a>(
152    provider_id: &str,
153    model: &str,
154    items: impl Iterator<Item = &'a ProviderConversationItem>,
155) -> ContextTokenCount {
156    let Some(bpe) = tokenizer_for_provider_model(provider_id, model) else {
157        return ContextTokenCount {
158            tokens: items.map(fallback_estimate_conversation_item_tokens).sum(),
159            source: ContextUsageSource::FallbackEstimate,
160        };
161    };
162    ContextTokenCount {
163        tokens: items
164            .map(|item| tokenizer_count_conversation_item_tokens(bpe, item))
165            .sum(),
166        source: ContextUsageSource::TokenizerEstimate,
167    }
168}
169
170pub(crate) fn project_text_tokens(provider_id: &str, model: &str, text: &str) -> ContextTokenCount {
171    if let Some(bpe) = tokenizer_for_provider_model(provider_id, model) {
172        return ContextTokenCount {
173            tokens: tokenizer_count_text_tokens(bpe, text),
174            source: ContextUsageSource::TokenizerProjection,
175        };
176    }
177    ContextTokenCount {
178        tokens: estimate_text_tokens(text),
179        source: ContextUsageSource::FallbackProjection,
180    }
181}
182
183#[derive(Debug, Clone, Copy, PartialEq, Eq)]
184pub(crate) enum TokenEncodingFamily {
185    O200KBase,
186    Cl100KBase,
187}
188
189fn tokenizer_for_provider_model(provider_id: &str, model: &str) -> Option<&'static CoreBPE> {
190    if provider_id != crate::providers::OPENAI_CODEX_PROVIDER {
191        return None;
192    }
193    match token_encoding_family_for_model(model)? {
194        TokenEncodingFamily::O200KBase => Some(&O200K_BASE),
195        TokenEncodingFamily::Cl100KBase => Some(&CL100K_BASE),
196    }
197}
198
199pub(crate) fn token_encoding_family_for_model(model: &str) -> Option<TokenEncodingFamily> {
200    let model = model.to_ascii_lowercase();
201    if model.starts_with("gpt-5")
202        || model.starts_with("gpt-4.1")
203        || model.starts_with("gpt-4o")
204        || model.starts_with("gpt-4.5")
205        || model.starts_with("o1")
206        || model.starts_with("o3")
207        || model.starts_with("o4")
208        || model.starts_with("codex-")
209    {
210        return Some(TokenEncodingFamily::O200KBase);
211    }
212    if model.starts_with("gpt-4")
213        || model.starts_with("gpt-3.5-turbo")
214        || model.starts_with("text-embedding-3")
215        || model == "text-embedding-ada-002"
216    {
217        return Some(TokenEncodingFamily::Cl100KBase);
218    }
219    None
220}
221
222/// Counts one complete provider conversation item with a supported local BPE.
223///
224/// This is intentionally separate from the existing projection functions: offline
225/// measurement must fail closed instead of using their fallback estimate.
226pub(crate) fn count_supported_provider_conversation_item_tokens(
227    provider_id: &str,
228    model: &str,
229    item: &ProviderConversationItem,
230) -> Option<usize> {
231    tokenizer_for_provider_model(provider_id, model)
232        .map(|bpe| tokenizer_count_conversation_item_tokens(bpe, item))
233}
234
235impl TokenEncodingFamily {
236    pub(crate) const fn as_str(self) -> &'static str {
237        match self {
238            Self::O200KBase => "o200k_base",
239            Self::Cl100KBase => "cl100k_base",
240        }
241    }
242}
243
244fn tokenizer_count_text_tokens(bpe: &CoreBPE, text: &str) -> usize {
245    bpe.encode_ordinary(text).len()
246}
247
248trait TokenCounter {
249    fn count_text(&self, text: &str) -> usize;
250
251    fn counts_message_role(&self) -> bool {
252        false
253    }
254}
255
256struct BpeCounter<'a> {
257    bpe: &'a CoreBPE,
258}
259
260impl TokenCounter for BpeCounter<'_> {
261    fn count_text(&self, text: &str) -> usize {
262        tokenizer_count_text_tokens(self.bpe, text)
263    }
264
265    fn counts_message_role(&self) -> bool {
266        true
267    }
268}
269
270struct FallbackCounter;
271
272impl TokenCounter for FallbackCounter {
273    fn count_text(&self, text: &str) -> usize {
274        estimate_text_tokens(text)
275    }
276}
277
278fn count_json_value_tokens(counter: &dyn TokenCounter, value: &serde_json::Value) -> usize {
279    match value {
280        serde_json::Value::String(text) => counter.count_text(text),
281        serde_json::Value::Array(items) => items
282            .iter()
283            .map(|item| count_json_value_tokens(counter, item))
284            .sum::<usize>()
285            .max(1),
286        serde_json::Value::Object(fields) => fields
287            .iter()
288            .map(|(key, value)| counter.count_text(key) + count_json_value_tokens(counter, value))
289            .sum::<usize>()
290            .max(1),
291        serde_json::Value::Null => 1,
292        other => counter.count_text(&other.to_string()),
293    }
294}
295
296fn count_response_item_tokens(counter: &dyn TokenCounter, item: &serde_json::Value) -> usize {
297    match item.get("type").and_then(serde_json::Value::as_str) {
298        Some("function_call") => {
299            counter.count_text("function_call")
300                + item
301                    .get("call_id")
302                    .and_then(serde_json::Value::as_str)
303                    .map(|text| counter.count_text(text))
304                    .unwrap_or(0)
305                + item
306                    .get("name")
307                    .and_then(serde_json::Value::as_str)
308                    .map(|text| counter.count_text(text))
309                    .unwrap_or(0)
310                + item
311                    .get("arguments")
312                    .map(|value| count_json_value_tokens(counter, value))
313                    .unwrap_or(0)
314        }
315        Some("function_call_output") => {
316            counter.count_text("function_call_output")
317                + item
318                    .get("call_id")
319                    .and_then(serde_json::Value::as_str)
320                    .map(|text| counter.count_text(text))
321                    .unwrap_or(0)
322                + item
323                    .get("output")
324                    .map(|value| count_json_value_tokens(counter, value))
325                    .unwrap_or(0)
326        }
327        Some("reasoning") => {
328            counter.count_text("reasoning")
329                + item
330                    .get("summary")
331                    .map(|value| count_json_value_tokens(counter, value))
332                    .unwrap_or(0)
333                + item
334                    .get("content")
335                    .map(|value| count_json_value_tokens(counter, value))
336                    .unwrap_or(0)
337        }
338        _ => {
339            if let Some(role) = item.get("role").and_then(serde_json::Value::as_str) {
340                counter.count_text(role)
341                    + item
342                        .get("content")
343                        .map(|value| count_json_value_tokens(counter, value))
344                        .unwrap_or(0)
345                    + item
346                        .get("tool_calls")
347                        .map(|value| count_json_value_tokens(counter, value))
348                        .unwrap_or(0)
349                    + item
350                        .get("tool_call_id")
351                        .and_then(serde_json::Value::as_str)
352                        .map(|text| counter.count_text(text))
353                        .unwrap_or(0)
354            } else {
355                counter.count_text(&item.to_string())
356            }
357        }
358    }
359}
360
361fn count_conversation_item_tokens(
362    counter: &dyn TokenCounter,
363    item: &ProviderConversationItem,
364) -> usize {
365    match item {
366        ProviderConversationItem::Message(message) => {
367            let role_tokens = if counter.counts_message_role() {
368                counter.count_text(message.role.as_api_str())
369            } else {
370                0
371            };
372            role_tokens + counter.count_text(&message.content) + 4
373        }
374        ProviderConversationItem::ResponseItem(item) => {
375            count_response_item_tokens(counter, item) + 4
376        }
377        ProviderConversationItem::ToolResult(result) => {
378            counter.count_text(&result.call_id)
379                + counter.count_text(&result.tool_name)
380                + counter.count_text(&result.output)
381                + 4
382        }
383        ProviderConversationItem::LegacyReplayNote {
384            event_type,
385            content,
386        } => counter.count_text(event_type) + counter.count_text(content) + 4,
387    }
388}
389
390fn tokenizer_count_conversation_item_tokens(
391    bpe: &CoreBPE,
392    item: &ProviderConversationItem,
393) -> usize {
394    count_conversation_item_tokens(&BpeCounter { bpe }, item)
395}
396
397fn fallback_estimate_conversation_item_tokens(item: &ProviderConversationItem) -> usize {
398    count_conversation_item_tokens(&FallbackCounter, item)
399}