use std::{collections::BTreeMap, sync::LazyLock};
use crate::{
output::ContextUsageSource,
providers::{ChatMessage, ProviderConversationItem, ProviderRequest, Usage},
};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use tiktoken_rs::{CoreBPE, cl100k_base, o200k_base};
static O200K_BASE: LazyLock<CoreBPE> =
LazyLock::new(|| o200k_base().expect("embedded o200k_base tokenizer table must load"));
static CL100K_BASE: LazyLock<CoreBPE> =
LazyLock::new(|| cl100k_base().expect("embedded cl100k_base tokenizer table must load"));
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub struct ContextBudget {
#[serde(default = "default_context_enabled")]
pub enabled: bool,
#[serde(default = "default_max_tokens")]
pub max_tokens: usize,
#[serde(default = "default_reserve_tokens")]
pub reserve_tokens: usize,
#[serde(default = "default_keep_recent_tokens")]
pub keep_recent_tokens: usize,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub model_overrides: BTreeMap<String, ContextBudgetOverride>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub struct ContextBudgetOverride {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reserve_tokens: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub keep_recent_tokens: Option<usize>,
}
impl Default for ContextBudget {
fn default() -> Self {
Self {
enabled: true,
max_tokens: default_max_tokens(),
reserve_tokens: default_reserve_tokens(),
keep_recent_tokens: default_keep_recent_tokens(),
model_overrides: BTreeMap::new(),
}
}
}
fn default_context_enabled() -> bool {
true
}
fn default_max_tokens() -> usize {
128_000
}
fn default_reserve_tokens() -> usize {
16_384
}
fn default_keep_recent_tokens() -> usize {
20_000
}
impl ContextBudgetOverride {
pub fn is_empty(&self) -> bool {
self.max_tokens.is_none()
&& self.reserve_tokens.is_none()
&& self.keep_recent_tokens.is_none()
}
}
impl ContextBudget {
pub fn threshold_tokens(&self) -> usize {
self.max_tokens.saturating_sub(self.reserve_tokens)
}
pub fn apply_model_override(&mut self, provider: &str, model: &str) {
let key = format!("{provider}/{model}");
let Some(model_override) = self.model_overrides.get(&key) else {
return;
};
if let Some(max_tokens) = model_override.max_tokens {
self.max_tokens = max_tokens;
}
if let Some(reserve_tokens) = model_override.reserve_tokens {
self.reserve_tokens = reserve_tokens;
}
if let Some(keep_recent_tokens) = model_override.keep_recent_tokens {
self.keep_recent_tokens = keep_recent_tokens;
}
}
}
pub fn estimate_text_tokens(text: &str) -> usize {
text.chars().count().div_ceil(4).max(1)
}
pub fn estimate_messages_tokens(messages: &[ChatMessage]) -> usize {
messages
.iter()
.map(|message| estimate_text_tokens(&message.content) + 4)
.sum()
}
pub fn usage_input_tokens(usage: &Usage) -> usize {
usize::try_from(usage.input).unwrap_or(usize::MAX)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct ContextTokenCount {
pub(crate) tokens: usize,
pub(crate) source: ContextUsageSource,
}
#[cfg(test)]
pub(crate) fn estimate_provider_request_input_tokens(request: &ProviderRequest) -> usize {
request
.conversation_items_iter()
.map(fallback_estimate_conversation_item_tokens)
.sum()
}
pub(crate) fn project_provider_request_input_tokens(
provider_id: &str,
request: &ProviderRequest,
) -> ContextTokenCount {
project_provider_conversation_item_tokens(
provider_id,
&request.model,
request.conversation_items_iter(),
)
}
pub(crate) fn project_provider_conversation_items_tokens(
provider_id: &str,
model: &str,
items: &[ProviderConversationItem],
) -> ContextTokenCount {
project_provider_conversation_item_tokens(provider_id, model, items.iter())
}
fn project_provider_conversation_item_tokens<'a>(
provider_id: &str,
model: &str,
items: impl Iterator<Item = &'a ProviderConversationItem>,
) -> ContextTokenCount {
let Some(bpe) = tokenizer_for_provider_model(provider_id, model) else {
return ContextTokenCount {
tokens: items.map(fallback_estimate_conversation_item_tokens).sum(),
source: ContextUsageSource::FallbackEstimate,
};
};
ContextTokenCount {
tokens: items
.map(|item| tokenizer_count_conversation_item_tokens(bpe, item))
.sum(),
source: ContextUsageSource::TokenizerEstimate,
}
}
pub(crate) fn project_text_tokens(provider_id: &str, model: &str, text: &str) -> ContextTokenCount {
if let Some(bpe) = tokenizer_for_provider_model(provider_id, model) {
return ContextTokenCount {
tokens: tokenizer_count_text_tokens(bpe, text),
source: ContextUsageSource::TokenizerProjection,
};
}
ContextTokenCount {
tokens: estimate_text_tokens(text),
source: ContextUsageSource::FallbackProjection,
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum TokenEncodingFamily {
O200KBase,
Cl100KBase,
}
fn tokenizer_for_provider_model(provider_id: &str, model: &str) -> Option<&'static CoreBPE> {
if provider_id != crate::providers::OPENAI_CODEX_PROVIDER {
return None;
}
match token_encoding_family_for_model(model)? {
TokenEncodingFamily::O200KBase => Some(&O200K_BASE),
TokenEncodingFamily::Cl100KBase => Some(&CL100K_BASE),
}
}
pub(crate) fn token_encoding_family_for_model(model: &str) -> Option<TokenEncodingFamily> {
let model = model.to_ascii_lowercase();
if model.starts_with("gpt-5")
|| model.starts_with("gpt-4.1")
|| model.starts_with("gpt-4o")
|| model.starts_with("gpt-4.5")
|| model.starts_with("o1")
|| model.starts_with("o3")
|| model.starts_with("o4")
|| model.starts_with("codex-")
{
return Some(TokenEncodingFamily::O200KBase);
}
if model.starts_with("gpt-4")
|| model.starts_with("gpt-3.5-turbo")
|| model.starts_with("text-embedding-3")
|| model == "text-embedding-ada-002"
{
return Some(TokenEncodingFamily::Cl100KBase);
}
None
}
fn tokenizer_count_text_tokens(bpe: &CoreBPE, text: &str) -> usize {
bpe.encode_ordinary(text).len()
}
trait TokenCounter {
fn count_text(&self, text: &str) -> usize;
fn counts_message_role(&self) -> bool {
false
}
}
struct BpeCounter<'a> {
bpe: &'a CoreBPE,
}
impl TokenCounter for BpeCounter<'_> {
fn count_text(&self, text: &str) -> usize {
tokenizer_count_text_tokens(self.bpe, text)
}
fn counts_message_role(&self) -> bool {
true
}
}
struct FallbackCounter;
impl TokenCounter for FallbackCounter {
fn count_text(&self, text: &str) -> usize {
estimate_text_tokens(text)
}
}
fn count_json_value_tokens(counter: &dyn TokenCounter, value: &serde_json::Value) -> usize {
match value {
serde_json::Value::String(text) => counter.count_text(text),
serde_json::Value::Array(items) => items
.iter()
.map(|item| count_json_value_tokens(counter, item))
.sum::<usize>()
.max(1),
serde_json::Value::Object(fields) => fields
.iter()
.map(|(key, value)| counter.count_text(key) + count_json_value_tokens(counter, value))
.sum::<usize>()
.max(1),
serde_json::Value::Null => 1,
other => counter.count_text(&other.to_string()),
}
}
fn count_response_item_tokens(counter: &dyn TokenCounter, item: &serde_json::Value) -> usize {
match item.get("type").and_then(serde_json::Value::as_str) {
Some("function_call") => {
counter.count_text("function_call")
+ item
.get("call_id")
.and_then(serde_json::Value::as_str)
.map(|text| counter.count_text(text))
.unwrap_or(0)
+ item
.get("name")
.and_then(serde_json::Value::as_str)
.map(|text| counter.count_text(text))
.unwrap_or(0)
+ item
.get("arguments")
.map(|value| count_json_value_tokens(counter, value))
.unwrap_or(0)
}
Some("function_call_output") => {
counter.count_text("function_call_output")
+ item
.get("call_id")
.and_then(serde_json::Value::as_str)
.map(|text| counter.count_text(text))
.unwrap_or(0)
+ item
.get("output")
.map(|value| count_json_value_tokens(counter, value))
.unwrap_or(0)
}
Some("reasoning") => {
counter.count_text("reasoning")
+ item
.get("summary")
.map(|value| count_json_value_tokens(counter, value))
.unwrap_or(0)
+ item
.get("content")
.map(|value| count_json_value_tokens(counter, value))
.unwrap_or(0)
}
_ => {
if let Some(role) = item.get("role").and_then(serde_json::Value::as_str) {
counter.count_text(role)
+ item
.get("content")
.map(|value| count_json_value_tokens(counter, value))
.unwrap_or(0)
+ item
.get("tool_calls")
.map(|value| count_json_value_tokens(counter, value))
.unwrap_or(0)
+ item
.get("tool_call_id")
.and_then(serde_json::Value::as_str)
.map(|text| counter.count_text(text))
.unwrap_or(0)
} else {
counter.count_text(&item.to_string())
}
}
}
}
fn count_conversation_item_tokens(
counter: &dyn TokenCounter,
item: &ProviderConversationItem,
) -> usize {
match item {
ProviderConversationItem::Message(message) => {
let role_tokens = if counter.counts_message_role() {
counter.count_text(message.role.as_api_str())
} else {
0
};
role_tokens + counter.count_text(&message.content) + 4
}
ProviderConversationItem::ResponseItem(item) => {
count_response_item_tokens(counter, item) + 4
}
ProviderConversationItem::ToolResult(result) => {
counter.count_text(&result.call_id)
+ counter.count_text(&result.tool_name)
+ counter.count_text(&result.output)
+ 4
}
ProviderConversationItem::LegacyReplayNote {
event_type,
content,
} => counter.count_text(event_type) + counter.count_text(content) + 4,
}
}
fn tokenizer_count_conversation_item_tokens(
bpe: &CoreBPE,
item: &ProviderConversationItem,
) -> usize {
count_conversation_item_tokens(&BpeCounter { bpe }, item)
}
fn fallback_estimate_conversation_item_tokens(item: &ProviderConversationItem) -> usize {
count_conversation_item_tokens(&FallbackCounter, item)
}