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 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
222fn tokenizer_count_text_tokens(bpe: &CoreBPE, text: &str) -> usize {
223 bpe.encode_ordinary(text).len()
224}
225
226trait TokenCounter {
227 fn count_text(&self, text: &str) -> usize;
228
229 fn counts_message_role(&self) -> bool {
230 false
231 }
232}
233
234struct BpeCounter<'a> {
235 bpe: &'a CoreBPE,
236}
237
238impl TokenCounter for BpeCounter<'_> {
239 fn count_text(&self, text: &str) -> usize {
240 tokenizer_count_text_tokens(self.bpe, text)
241 }
242
243 fn counts_message_role(&self) -> bool {
244 true
245 }
246}
247
248struct FallbackCounter;
249
250impl TokenCounter for FallbackCounter {
251 fn count_text(&self, text: &str) -> usize {
252 estimate_text_tokens(text)
253 }
254}
255
256fn count_json_value_tokens(counter: &dyn TokenCounter, value: &serde_json::Value) -> usize {
257 match value {
258 serde_json::Value::String(text) => counter.count_text(text),
259 serde_json::Value::Array(items) => items
260 .iter()
261 .map(|item| count_json_value_tokens(counter, item))
262 .sum::<usize>()
263 .max(1),
264 serde_json::Value::Object(fields) => fields
265 .iter()
266 .map(|(key, value)| counter.count_text(key) + count_json_value_tokens(counter, value))
267 .sum::<usize>()
268 .max(1),
269 serde_json::Value::Null => 1,
270 other => counter.count_text(&other.to_string()),
271 }
272}
273
274fn count_response_item_tokens(counter: &dyn TokenCounter, item: &serde_json::Value) -> usize {
275 match item.get("type").and_then(serde_json::Value::as_str) {
276 Some("function_call") => {
277 counter.count_text("function_call")
278 + item
279 .get("call_id")
280 .and_then(serde_json::Value::as_str)
281 .map(|text| counter.count_text(text))
282 .unwrap_or(0)
283 + item
284 .get("name")
285 .and_then(serde_json::Value::as_str)
286 .map(|text| counter.count_text(text))
287 .unwrap_or(0)
288 + item
289 .get("arguments")
290 .map(|value| count_json_value_tokens(counter, value))
291 .unwrap_or(0)
292 }
293 Some("function_call_output") => {
294 counter.count_text("function_call_output")
295 + item
296 .get("call_id")
297 .and_then(serde_json::Value::as_str)
298 .map(|text| counter.count_text(text))
299 .unwrap_or(0)
300 + item
301 .get("output")
302 .map(|value| count_json_value_tokens(counter, value))
303 .unwrap_or(0)
304 }
305 Some("reasoning") => {
306 counter.count_text("reasoning")
307 + item
308 .get("summary")
309 .map(|value| count_json_value_tokens(counter, value))
310 .unwrap_or(0)
311 + item
312 .get("content")
313 .map(|value| count_json_value_tokens(counter, value))
314 .unwrap_or(0)
315 }
316 _ => {
317 if let Some(role) = item.get("role").and_then(serde_json::Value::as_str) {
318 counter.count_text(role)
319 + item
320 .get("content")
321 .map(|value| count_json_value_tokens(counter, value))
322 .unwrap_or(0)
323 + item
324 .get("tool_calls")
325 .map(|value| count_json_value_tokens(counter, value))
326 .unwrap_or(0)
327 + item
328 .get("tool_call_id")
329 .and_then(serde_json::Value::as_str)
330 .map(|text| counter.count_text(text))
331 .unwrap_or(0)
332 } else {
333 counter.count_text(&item.to_string())
334 }
335 }
336 }
337}
338
339fn count_conversation_item_tokens(
340 counter: &dyn TokenCounter,
341 item: &ProviderConversationItem,
342) -> usize {
343 match item {
344 ProviderConversationItem::Message(message) => {
345 let role_tokens = if counter.counts_message_role() {
346 counter.count_text(message.role.as_api_str())
347 } else {
348 0
349 };
350 role_tokens + counter.count_text(&message.content) + 4
351 }
352 ProviderConversationItem::ResponseItem(item) => {
353 count_response_item_tokens(counter, item) + 4
354 }
355 ProviderConversationItem::ToolResult(result) => {
356 counter.count_text(&result.call_id)
357 + counter.count_text(&result.tool_name)
358 + counter.count_text(&result.output)
359 + 4
360 }
361 ProviderConversationItem::LegacyReplayNote {
362 event_type,
363 content,
364 } => counter.count_text(event_type) + counter.count_text(content) + 4,
365 }
366}
367
368fn tokenizer_count_conversation_item_tokens(
369 bpe: &CoreBPE,
370 item: &ProviderConversationItem,
371) -> usize {
372 count_conversation_item_tokens(&BpeCounter { bpe }, item)
373}
374
375fn fallback_estimate_conversation_item_tokens(item: &ProviderConversationItem) -> usize {
376 count_conversation_item_tokens(&FallbackCounter, item)
377}