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 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 fn threshold_tokens(&self) -> usize {
75 self.max_tokens.saturating_sub(self.reserve_tokens)
76 }
77
78 pub 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 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 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 request
120 .conversation_items_iter()
121 .map(fallback_estimate_conversation_item_tokens)
122 .sum()
123}
124
125pub(crate) fn project_provider_request_input_tokens(
126 provider_id: &str,
127 request: &ProviderRequest,
128) -> ContextTokenCount {
129 project_provider_conversation_item_tokens(
130 provider_id,
131 &request.model,
132 request.conversation_items_iter(),
133 )
134}
135
136pub(crate) fn project_provider_conversation_items_tokens(
137 provider_id: &str,
138 model: &str,
139 items: &[ProviderConversationItem],
140) -> ContextTokenCount {
141 project_provider_conversation_item_tokens(provider_id, model, items.iter())
142}
143
144fn project_provider_conversation_item_tokens<'a>(
145 provider_id: &str,
146 model: &str,
147 items: impl Iterator<Item = &'a ProviderConversationItem>,
148) -> ContextTokenCount {
149 let Some(bpe) = tokenizer_for_provider_model(provider_id, model) else {
150 return ContextTokenCount {
151 tokens: items.map(fallback_estimate_conversation_item_tokens).sum(),
152 source: ContextUsageSource::FallbackEstimate,
153 };
154 };
155 ContextTokenCount {
156 tokens: items
157 .map(|item| tokenizer_count_conversation_item_tokens(bpe, item))
158 .sum(),
159 source: ContextUsageSource::TokenizerEstimate,
160 }
161}
162
163pub(crate) fn project_text_tokens(provider_id: &str, model: &str, text: &str) -> ContextTokenCount {
164 if let Some(bpe) = tokenizer_for_provider_model(provider_id, model) {
165 return ContextTokenCount {
166 tokens: tokenizer_count_text_tokens(bpe, text),
167 source: ContextUsageSource::TokenizerProjection,
168 };
169 }
170 ContextTokenCount {
171 tokens: estimate_text_tokens(text),
172 source: ContextUsageSource::FallbackProjection,
173 }
174}
175
176#[derive(Debug, Clone, Copy, PartialEq, Eq)]
177pub(crate) enum TokenEncodingFamily {
178 O200KBase,
179 Cl100KBase,
180}
181
182fn tokenizer_for_provider_model(provider_id: &str, model: &str) -> Option<&'static CoreBPE> {
183 if provider_id != crate::providers::OPENAI_CODEX_PROVIDER {
184 return None;
185 }
186 match token_encoding_family_for_model(model)? {
187 TokenEncodingFamily::O200KBase => Some(&O200K_BASE),
188 TokenEncodingFamily::Cl100KBase => Some(&CL100K_BASE),
189 }
190}
191
192pub(crate) fn token_encoding_family_for_model(model: &str) -> Option<TokenEncodingFamily> {
193 let model = model.to_ascii_lowercase();
194 if model.starts_with("gpt-5")
195 || model.starts_with("gpt-4.1")
196 || model.starts_with("gpt-4o")
197 || model.starts_with("gpt-4.5")
198 || model.starts_with("o1")
199 || model.starts_with("o3")
200 || model.starts_with("o4")
201 || model.starts_with("codex-")
202 {
203 return Some(TokenEncodingFamily::O200KBase);
204 }
205 if model.starts_with("gpt-4")
206 || model.starts_with("gpt-3.5-turbo")
207 || model.starts_with("text-embedding-3")
208 || model == "text-embedding-ada-002"
209 {
210 return Some(TokenEncodingFamily::Cl100KBase);
211 }
212 None
213}
214
215fn tokenizer_count_text_tokens(bpe: &CoreBPE, text: &str) -> usize {
216 bpe.encode_ordinary(text).len()
217}
218
219trait TokenCounter {
220 fn count_text(&self, text: &str) -> usize;
221
222 fn counts_message_role(&self) -> bool {
223 false
224 }
225}
226
227struct BpeCounter<'a> {
228 bpe: &'a CoreBPE,
229}
230
231impl TokenCounter for BpeCounter<'_> {
232 fn count_text(&self, text: &str) -> usize {
233 tokenizer_count_text_tokens(self.bpe, text)
234 }
235
236 fn counts_message_role(&self) -> bool {
237 true
238 }
239}
240
241struct FallbackCounter;
242
243impl TokenCounter for FallbackCounter {
244 fn count_text(&self, text: &str) -> usize {
245 estimate_text_tokens(text)
246 }
247}
248
249fn count_json_value_tokens(counter: &dyn TokenCounter, value: &serde_json::Value) -> usize {
250 match value {
251 serde_json::Value::String(text) => counter.count_text(text),
252 serde_json::Value::Array(items) => items
253 .iter()
254 .map(|item| count_json_value_tokens(counter, item))
255 .sum::<usize>()
256 .max(1),
257 serde_json::Value::Object(fields) => fields
258 .iter()
259 .map(|(key, value)| counter.count_text(key) + count_json_value_tokens(counter, value))
260 .sum::<usize>()
261 .max(1),
262 serde_json::Value::Null => 1,
263 other => counter.count_text(&other.to_string()),
264 }
265}
266
267fn count_response_item_tokens(counter: &dyn TokenCounter, item: &serde_json::Value) -> usize {
268 match item.get("type").and_then(serde_json::Value::as_str) {
269 Some("function_call") => {
270 counter.count_text("function_call")
271 + item
272 .get("call_id")
273 .and_then(serde_json::Value::as_str)
274 .map(|text| counter.count_text(text))
275 .unwrap_or(0)
276 + item
277 .get("name")
278 .and_then(serde_json::Value::as_str)
279 .map(|text| counter.count_text(text))
280 .unwrap_or(0)
281 + item
282 .get("arguments")
283 .map(|value| count_json_value_tokens(counter, value))
284 .unwrap_or(0)
285 }
286 Some("function_call_output") => {
287 counter.count_text("function_call_output")
288 + item
289 .get("call_id")
290 .and_then(serde_json::Value::as_str)
291 .map(|text| counter.count_text(text))
292 .unwrap_or(0)
293 + item
294 .get("output")
295 .map(|value| count_json_value_tokens(counter, value))
296 .unwrap_or(0)
297 }
298 Some("reasoning") => {
299 counter.count_text("reasoning")
300 + item
301 .get("summary")
302 .map(|value| count_json_value_tokens(counter, value))
303 .unwrap_or(0)
304 + item
305 .get("content")
306 .map(|value| count_json_value_tokens(counter, value))
307 .unwrap_or(0)
308 }
309 _ => {
310 if let Some(role) = item.get("role").and_then(serde_json::Value::as_str) {
311 counter.count_text(role)
312 + item
313 .get("content")
314 .map(|value| count_json_value_tokens(counter, value))
315 .unwrap_or(0)
316 + item
317 .get("tool_calls")
318 .map(|value| count_json_value_tokens(counter, value))
319 .unwrap_or(0)
320 + item
321 .get("tool_call_id")
322 .and_then(serde_json::Value::as_str)
323 .map(|text| counter.count_text(text))
324 .unwrap_or(0)
325 } else {
326 counter.count_text(&item.to_string())
327 }
328 }
329 }
330}
331
332fn count_conversation_item_tokens(
333 counter: &dyn TokenCounter,
334 item: &ProviderConversationItem,
335) -> usize {
336 match item {
337 ProviderConversationItem::Message(message) => {
338 let role_tokens = if counter.counts_message_role() {
339 counter.count_text(message.role.as_api_str())
340 } else {
341 0
342 };
343 role_tokens + counter.count_text(&message.content) + 4
344 }
345 ProviderConversationItem::ResponseItem(item) => {
346 count_response_item_tokens(counter, item) + 4
347 }
348 ProviderConversationItem::ToolResult(result) => {
349 counter.count_text(&result.call_id)
350 + counter.count_text(&result.tool_name)
351 + counter.count_text(&result.output)
352 + 4
353 }
354 ProviderConversationItem::LegacyReplayNote {
355 event_type,
356 content,
357 } => counter.count_text(event_type) + counter.count_text(content) + 4,
358 }
359}
360
361fn tokenizer_count_conversation_item_tokens(
362 bpe: &CoreBPE,
363 item: &ProviderConversationItem,
364) -> usize {
365 count_conversation_item_tokens(&BpeCounter { bpe }, item)
366}
367
368fn fallback_estimate_conversation_item_tokens(item: &ProviderConversationItem) -> usize {
369 count_conversation_item_tokens(&FallbackCounter, item)
370}