Skip to main content

magi_code/context/
tokens.rs

1use std::{collections::BTreeMap, sync::LazyLock};
2
3use crate::{
4    output::ContextUsageSource,
5    providers::{ProviderConversationItem, ProviderRequest},
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_limits(&mut self, provider: &str, model: &str) {
79        if provider == crate::providers::CLAUDE_SUBSCRIPTION_PROVIDER {
80            // CLI capacity replaces published windows and the generic fallback.
81            self.max_tokens = crate::providers::CLAUDE_SUBSCRIPTION_MAX_PROMPT_TOKENS;
82        }
83        let key = format!("{provider}/{model}");
84        if let Some(model_override) = self.model_overrides.get(&key) {
85            if let Some(max_tokens) = model_override.max_tokens {
86                self.max_tokens = max_tokens;
87            }
88            if let Some(reserve_tokens) = model_override.reserve_tokens {
89                self.reserve_tokens = reserve_tokens;
90            }
91            if let Some(keep_recent_tokens) = model_override.keep_recent_tokens {
92                self.keep_recent_tokens = keep_recent_tokens;
93            }
94        }
95        // Catalog metadata and user overrides cannot raise the CLI transport limit.
96        if provider == crate::providers::CLAUDE_SUBSCRIPTION_PROVIDER {
97            self.enabled = true;
98            // This is an input-only ceiling, not a published input + output window.
99            // Conversation headroom is reserved separately by CompactionSettings.
100            self.reserve_tokens = 0;
101            self.max_tokens = self
102                .max_tokens
103                .min(crate::providers::CLAUDE_SUBSCRIPTION_MAX_PROMPT_TOKENS);
104        }
105    }
106}
107
108pub(crate) fn estimate_text_tokens(text: &str) -> usize {
109    text.chars().count().div_ceil(4).max(1)
110}
111
112#[derive(Debug, Clone, Copy, PartialEq, Eq)]
113pub(crate) struct ContextTokenCount {
114    pub(crate) tokens: usize,
115    pub(crate) source: ContextUsageSource,
116}
117
118pub(crate) fn project_provider_request_input_tokens(
119    provider_id: &str,
120    request: &ProviderRequest,
121) -> ContextTokenCount {
122    let mut projection = project_provider_conversation_item_tokens(
123        provider_id,
124        &request.model,
125        request.conversation_items_iter(),
126    );
127    if let Some(tool_definitions) = request.tool_definitions_json_if_enabled() {
128        let tool_tokens =
129            if let Some(bpe) = tokenizer_for_provider_model(provider_id, &request.model) {
130                count_json_value_tokens(&BpeCounter { bpe }, &tool_definitions)
131            } else {
132                count_json_value_tokens(&FallbackCounter, &tool_definitions)
133            };
134        projection.tokens = projection.tokens.saturating_add(tool_tokens);
135    }
136    projection
137}
138
139pub(crate) fn project_provider_conversation_items_tokens(
140    provider_id: &str,
141    model: &str,
142    items: &[ProviderConversationItem],
143) -> ContextTokenCount {
144    project_provider_conversation_item_tokens(provider_id, model, items.iter())
145}
146
147fn project_provider_conversation_item_tokens<'a>(
148    provider_id: &str,
149    model: &str,
150    items: impl Iterator<Item = &'a ProviderConversationItem>,
151) -> ContextTokenCount {
152    let Some(bpe) = tokenizer_for_provider_model(provider_id, model) else {
153        return ContextTokenCount {
154            tokens: items.map(fallback_estimate_conversation_item_tokens).sum(),
155            source: ContextUsageSource::FallbackEstimate,
156        };
157    };
158    ContextTokenCount {
159        tokens: items
160            .map(|item| tokenizer_count_conversation_item_tokens(bpe, item))
161            .sum(),
162        source: ContextUsageSource::TokenizerEstimate,
163    }
164}
165
166pub(crate) fn project_text_tokens(provider_id: &str, model: &str, text: &str) -> ContextTokenCount {
167    if let Some(bpe) = tokenizer_for_provider_model(provider_id, model) {
168        return ContextTokenCount {
169            tokens: tokenizer_count_text_tokens(bpe, text),
170            source: ContextUsageSource::TokenizerProjection,
171        };
172    }
173    ContextTokenCount {
174        tokens: estimate_text_tokens(text),
175        source: ContextUsageSource::FallbackProjection,
176    }
177}
178
179#[derive(Debug, Clone, Copy, PartialEq, Eq)]
180pub(crate) enum TokenEncodingFamily {
181    O200KBase,
182    Cl100KBase,
183}
184
185fn tokenizer_for_provider_model(provider_id: &str, model: &str) -> Option<&'static CoreBPE> {
186    if provider_id != crate::providers::OPENAI_CODEX_PROVIDER {
187        return None;
188    }
189    match token_encoding_family_for_model(model)? {
190        TokenEncodingFamily::O200KBase => Some(&O200K_BASE),
191        TokenEncodingFamily::Cl100KBase => Some(&CL100K_BASE),
192    }
193}
194
195pub(crate) fn token_encoding_family_for_model(model: &str) -> Option<TokenEncodingFamily> {
196    let model = model.to_ascii_lowercase();
197    if model.starts_with("gpt-5")
198        || model.starts_with("gpt-4.1")
199        || model.starts_with("gpt-4o")
200        || model.starts_with("gpt-4.5")
201        || model.starts_with("o1")
202        || model.starts_with("o3")
203        || model.starts_with("o4")
204        || model.starts_with("codex-")
205    {
206        return Some(TokenEncodingFamily::O200KBase);
207    }
208    if model.starts_with("gpt-4")
209        || model.starts_with("gpt-3.5-turbo")
210        || model.starts_with("text-embedding-3")
211        || model == "text-embedding-ada-002"
212    {
213        return Some(TokenEncodingFamily::Cl100KBase);
214    }
215    None
216}
217
218/// Counts one complete provider conversation item with a supported local BPE.
219///
220/// This is intentionally separate from the existing projection functions: offline
221/// measurement must fail closed instead of using their fallback estimate.
222pub(crate) fn count_supported_provider_conversation_item_tokens(
223    provider_id: &str,
224    model: &str,
225    item: &ProviderConversationItem,
226) -> Option<usize> {
227    tokenizer_for_provider_model(provider_id, model)
228        .map(|bpe| tokenizer_count_conversation_item_tokens(bpe, item))
229}
230
231impl TokenEncodingFamily {
232    pub(crate) const fn as_str(self) -> &'static str {
233        match self {
234            Self::O200KBase => "o200k_base",
235            Self::Cl100KBase => "cl100k_base",
236        }
237    }
238}
239
240fn tokenizer_count_text_tokens(bpe: &CoreBPE, text: &str) -> usize {
241    bpe.encode_ordinary(text).len()
242}
243
244trait TokenCounter {
245    fn count_text(&self, text: &str) -> usize;
246
247    fn counts_message_role(&self) -> bool {
248        false
249    }
250}
251
252struct BpeCounter<'a> {
253    bpe: &'a CoreBPE,
254}
255
256impl TokenCounter for BpeCounter<'_> {
257    fn count_text(&self, text: &str) -> usize {
258        tokenizer_count_text_tokens(self.bpe, text)
259    }
260
261    fn counts_message_role(&self) -> bool {
262        true
263    }
264}
265
266struct FallbackCounter;
267
268impl TokenCounter for FallbackCounter {
269    fn count_text(&self, text: &str) -> usize {
270        estimate_text_tokens(text)
271    }
272}
273
274fn count_json_value_tokens(counter: &dyn TokenCounter, value: &serde_json::Value) -> usize {
275    match value {
276        serde_json::Value::String(text) => counter.count_text(text),
277        serde_json::Value::Array(items) => items
278            .iter()
279            .map(|item| count_json_value_tokens(counter, item))
280            .sum::<usize>()
281            .max(1),
282        serde_json::Value::Object(fields) => fields
283            .iter()
284            .map(|(key, value)| counter.count_text(key) + count_json_value_tokens(counter, value))
285            .sum::<usize>()
286            .max(1),
287        serde_json::Value::Null => 1,
288        other => counter.count_text(&other.to_string()),
289    }
290}
291
292fn count_response_item_tokens(counter: &dyn TokenCounter, item: &serde_json::Value) -> usize {
293    match item.get("type").and_then(serde_json::Value::as_str) {
294        Some("function_call") => {
295            counter.count_text("function_call")
296                + item
297                    .get("call_id")
298                    .and_then(serde_json::Value::as_str)
299                    .map(|text| counter.count_text(text))
300                    .unwrap_or(0)
301                + item
302                    .get("name")
303                    .and_then(serde_json::Value::as_str)
304                    .map(|text| counter.count_text(text))
305                    .unwrap_or(0)
306                + item
307                    .get("arguments")
308                    .map(|value| count_json_value_tokens(counter, value))
309                    .unwrap_or(0)
310        }
311        Some("function_call_output") => {
312            counter.count_text("function_call_output")
313                + item
314                    .get("call_id")
315                    .and_then(serde_json::Value::as_str)
316                    .map(|text| counter.count_text(text))
317                    .unwrap_or(0)
318                + item
319                    .get("output")
320                    .map(|value| count_json_value_tokens(counter, value))
321                    .unwrap_or(0)
322        }
323        Some("reasoning") => {
324            counter.count_text("reasoning")
325                + item
326                    .get("summary")
327                    .map(|value| count_json_value_tokens(counter, value))
328                    .unwrap_or(0)
329                + item
330                    .get("content")
331                    .map(|value| count_json_value_tokens(counter, value))
332                    .unwrap_or(0)
333        }
334        _ => {
335            if let Some(role) = item.get("role").and_then(serde_json::Value::as_str) {
336                counter.count_text(role)
337                    + item
338                        .get("content")
339                        .map(|value| count_json_value_tokens(counter, value))
340                        .unwrap_or(0)
341                    + item
342                        .get("tool_calls")
343                        .map(|value| count_json_value_tokens(counter, value))
344                        .unwrap_or(0)
345                    + item
346                        .get("tool_call_id")
347                        .and_then(serde_json::Value::as_str)
348                        .map(|text| counter.count_text(text))
349                        .unwrap_or(0)
350            } else {
351                counter.count_text(&item.to_string())
352            }
353        }
354    }
355}
356
357fn count_conversation_item_tokens(
358    counter: &dyn TokenCounter,
359    item: &ProviderConversationItem,
360) -> usize {
361    match item {
362        ProviderConversationItem::ReasoningSelection { .. } => 16,
363        ProviderConversationItem::Message(message) => {
364            let role_tokens = if counter.counts_message_role() {
365                counter.count_text(message.role.as_api_str())
366            } else {
367                0
368            };
369            role_tokens + counter.count_text(&message.content) + 4
370        }
371        ProviderConversationItem::ResponseItem(item) => {
372            count_response_item_tokens(counter, item) + 4
373        }
374        ProviderConversationItem::ToolResult(result) => {
375            counter.count_text(&result.call_id)
376                + counter.count_text(&result.tool_name)
377                + counter.count_text(&result.output)
378                + 4
379        }
380        ProviderConversationItem::LegacyReplayNote {
381            event_type,
382            content,
383        } => counter.count_text(event_type) + counter.count_text(content) + 4,
384    }
385}
386
387fn tokenizer_count_conversation_item_tokens(
388    bpe: &CoreBPE,
389    item: &ProviderConversationItem,
390) -> usize {
391    count_conversation_item_tokens(&BpeCounter { bpe }, item)
392}
393
394fn fallback_estimate_conversation_item_tokens(item: &ProviderConversationItem) -> usize {
395    count_conversation_item_tokens(&FallbackCounter, item)
396}