magi-code 0.63.1

Repository-aware CLI coding agent for terminal work
Documentation
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 {
    // 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.
    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 {
    project_provider_request_input_tokens("local-ai", request).tokens
}

pub(crate) fn project_provider_request_input_tokens(
    provider_id: &str,
    request: &ProviderRequest,
) -> ContextTokenCount {
    let mut projection = project_provider_conversation_item_tokens(
        provider_id,
        &request.model,
        request.conversation_items_iter(),
    );
    if let Some(tool_definitions) = request.tool_definitions_json_if_enabled() {
        let tool_tokens =
            if let Some(bpe) = tokenizer_for_provider_model(provider_id, &request.model) {
                count_json_value_tokens(&BpeCounter { bpe }, &tool_definitions)
            } else {
                count_json_value_tokens(&FallbackCounter, &tool_definitions)
            };
        projection.tokens = projection.tokens.saturating_add(tool_tokens);
    }
    projection
}

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)
}