use crate::{Message, ModelProfileWitness, ProviderRequestPressure, Session, ToolDef, ToolName};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ContextBudgetState {
Within,
ForecastExceeded,
Exceeded,
}
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ContextBudgetEstimateProvenance {
#[default]
CanonicalForecast,
ExactProviderTokenCount,
}
impl ContextBudgetState {
pub fn requires_dispatch_refusal(self) -> bool {
matches!(self, Self::Exceeded)
}
}
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", deny_unknown_fields)]
pub struct ContextBudgetFact {
pub state: ContextBudgetState,
pub context_window_tokens: u32,
pub estimated_input_tokens: u64,
pub estimated_tool_tokens: u64,
pub reserved_output_tokens: u32,
pub estimated_total_tokens: u64,
pub remaining_tokens: u64,
pub overage_tokens: u64,
#[serde(default)]
pub estimate_provenance: ContextBudgetEstimateProvenance,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_lowered_encoded_bytes: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_issued_input_tokens: Option<u64>,
}
impl ContextBudgetFact {
pub fn requires_dispatch_refusal(&self) -> bool {
self.state.requires_dispatch_refusal()
&& self.estimate_provenance == ContextBudgetEstimateProvenance::ExactProviderTokenCount
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum ContextBudgetFactError {
#[error("active model profile has no context-window authority")]
ContextWindowUnavailable,
#[error("failed to estimate durable transcript tokens: {detail}")]
InputEstimationFailed {
detail: String,
},
#[error("failed to estimate visible tool '{tool_name}' tokens: {detail}")]
ToolEstimationFailed {
tool_name: ToolName,
detail: String,
},
}
pub fn context_budget_fact_for_session(
session: &Session,
visible_tools: &[Arc<ToolDef>],
reserved_output_tokens: u32,
active_model_profile: &ModelProfileWitness,
) -> Result<ContextBudgetFact, ContextBudgetFactError> {
context_budget_fact_for_messages(
session.messages(),
visible_tools,
reserved_output_tokens,
active_model_profile,
)
}
pub fn context_budget_fact_for_messages(
messages: &[Message],
visible_tools: &[Arc<ToolDef>],
reserved_output_tokens: u32,
active_model_profile: &ModelProfileWitness,
) -> Result<ContextBudgetFact, ContextBudgetFactError> {
context_budget_fact_for_messages_with_pressure(
messages,
visible_tools,
reserved_output_tokens,
active_model_profile,
None,
)
}
pub fn context_budget_fact_for_provider_request(
messages: &[Message],
visible_tools: &[Arc<ToolDef>],
reserved_output_tokens: u32,
active_model_profile: &ModelProfileWitness,
provider_request_pressure: ProviderRequestPressure,
) -> Result<ContextBudgetFact, ContextBudgetFactError> {
context_budget_fact_for_messages_with_pressure(
messages,
visible_tools,
reserved_output_tokens,
active_model_profile,
Some(provider_request_pressure),
)
}
fn context_budget_fact_for_messages_with_pressure(
messages: &[Message],
visible_tools: &[Arc<ToolDef>],
reserved_output_tokens: u32,
active_model_profile: &ModelProfileWitness,
provider_request_pressure: Option<ProviderRequestPressure>,
) -> Result<ContextBudgetFact, ContextBudgetFactError> {
let context_window_tokens = active_model_profile
.context_window()
.ok_or(ContextBudgetFactError::ContextWindowUnavailable)?;
let estimated_input_tokens = crate::agent::compact::estimate_tokens(messages)
.map_err(|error| ContextBudgetFactError::InputEstimationFailed {
detail: error.to_string(),
})?
.saturating_add(messages.len() as u64);
let estimated_tool_tokens = visible_tools.iter().try_fold(
0_u64,
|total, tool| -> Result<u64, ContextBudgetFactError> {
let serialized = serde_json::to_vec(tool.as_ref()).map_err(|error| {
ContextBudgetFactError::ToolEstimationFailed {
tool_name: tool.tool_name(),
detail: error.to_string(),
}
})?;
let tokens = (serialized.len() as u64).div_ceil(4);
Ok(total.saturating_add(tokens))
},
)?;
let canonical_input_and_tools = estimated_input_tokens.saturating_add(estimated_tool_tokens);
let provider_issued_input_tokens =
provider_request_pressure.and_then(|pressure| pressure.provider_issued_input_tokens);
let effective_input_and_tools =
provider_issued_input_tokens.unwrap_or(canonical_input_and_tools);
let estimated_total_tokens =
effective_input_and_tools.saturating_add(u64::from(reserved_output_tokens));
let context_window = u64::from(context_window_tokens);
let state = if estimated_total_tokens > context_window {
if provider_issued_input_tokens.is_some() {
ContextBudgetState::Exceeded
} else {
ContextBudgetState::ForecastExceeded
}
} else {
ContextBudgetState::Within
};
Ok(ContextBudgetFact {
state,
context_window_tokens,
estimated_input_tokens,
estimated_tool_tokens,
reserved_output_tokens,
estimated_total_tokens,
remaining_tokens: context_window.saturating_sub(estimated_total_tokens),
overage_tokens: estimated_total_tokens.saturating_sub(context_window),
estimate_provenance: if provider_issued_input_tokens.is_some() {
ContextBudgetEstimateProvenance::ExactProviderTokenCount
} else {
ContextBudgetEstimateProvenance::CanonicalForecast
},
provider_lowered_encoded_bytes: provider_request_pressure
.map(|pressure| pressure.encoded_bytes),
provider_issued_input_tokens,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::CustomModelConfig;
use crate::model_profile::test_catalog::TEST_CATALOG;
use crate::{Config, ModelRegistry, Provider};
const MODEL: &str = "context-budget-test-model";
fn profile_witness(
context_window: Option<u32>,
) -> Result<ModelProfileWitness, Box<dyn std::error::Error>> {
let mut config = Config::default();
config.models.custom.insert(
MODEL.to_string(),
CustomModelConfig {
provider: Provider::OpenAI,
display_name: None,
context_window,
max_output_tokens: Some(8_192),
vision: None,
web_search: None,
call_timeout_secs: None,
},
);
let registry = ModelRegistry::from_config(&config, *TEST_CATALOG)?;
registry
.profile_witness_for_provider(Provider::OpenAI, MODEL)
.ok_or_else(|| std::io::Error::other("missing exact test profile witness").into())
}
fn session_with_system_bytes(bytes: usize) -> Session {
let mut session = Session::new();
session.append_system_message("a".repeat(bytes));
session
}
#[test]
fn classification_is_deterministic_and_accounts_for_all_components()
-> Result<(), Box<dyn std::error::Error>> {
let session = session_with_system_bytes(1_900);
let profile = profile_witness(Some(1_000))?;
let tools = [Arc::new(ToolDef::new(
"lookup",
"look up an exact record",
serde_json::json!({
"type": "object",
"properties": { "query": { "type": "string" } }
}),
))];
let first = context_budget_fact_for_session(&session, &tools, 100, &profile)?;
let second = context_budget_fact_for_session(&session, &tools, 100, &profile)?;
assert_eq!(first, second);
assert!(first.estimated_input_tokens > 0);
assert!(first.estimated_tool_tokens > 0);
assert_eq!(first.reserved_output_tokens, 100);
assert_eq!(
first.estimated_total_tokens,
first
.estimated_input_tokens
.saturating_add(first.estimated_tool_tokens)
.saturating_add(100)
);
Ok(())
}
#[test]
fn provider_lowered_bytes_are_observable_but_not_token_authority()
-> Result<(), Box<dyn std::error::Error>> {
let session = session_with_system_bytes(20);
let profile = profile_witness(Some(1_000))?;
let pressure = ProviderRequestPressure::new(2_400, Some(10_000));
let fact = context_budget_fact_for_provider_request(
session.messages(),
&[],
100,
&profile,
pressure,
)?;
assert_eq!(
fact.estimate_provenance,
ContextBudgetEstimateProvenance::CanonicalForecast
);
assert_eq!(fact.provider_lowered_encoded_bytes, Some(2_400));
assert!(fact.estimated_total_tokens < 1_000);
assert_eq!(fact.state, ContextBudgetState::Within);
assert!(!fact.requires_dispatch_refusal());
Ok(())
}
#[test]
fn provider_lowered_witness_never_shrinks_the_canonical_forecast()
-> Result<(), Box<dyn std::error::Error>> {
let session = session_with_system_bytes(2_400);
let profile = profile_witness(Some(10_000))?;
let forecast = context_budget_fact_for_session(&session, &[], 100, &profile)?;
let witnessed = context_budget_fact_for_provider_request(
session.messages(),
&[],
100,
&profile,
ProviderRequestPressure::new(12, None),
)?;
assert_eq!(
witnessed.estimated_total_tokens,
forecast.estimated_total_tokens
);
assert_eq!(
witnessed.estimated_input_tokens,
forecast.estimated_input_tokens
);
assert_eq!(
witnessed.estimate_provenance,
ContextBudgetEstimateProvenance::CanonicalForecast
);
assert_eq!(witnessed.provider_lowered_encoded_bytes, Some(12));
Ok(())
}
#[test]
fn exact_provider_token_count_dominates_all_forecasts() -> Result<(), Box<dyn std::error::Error>>
{
let session = session_with_system_bytes(20);
let profile = profile_witness(Some(1_000))?;
let pressure =
ProviderRequestPressure::new(2_400, None).with_provider_issued_input_tokens(1_100);
let fact = context_budget_fact_for_provider_request(
session.messages(),
&[],
100,
&profile,
pressure,
)?;
assert_eq!(
fact.estimate_provenance,
ContextBudgetEstimateProvenance::ExactProviderTokenCount
);
assert_eq!(fact.provider_issued_input_tokens, Some(1_100));
assert_eq!(fact.estimated_total_tokens, 1_200);
assert_eq!(fact.state, ContextBudgetState::Exceeded);
assert!(fact.requires_dispatch_refusal());
Ok(())
}
#[test]
fn state_distinguishes_within_and_forecast_exceeded() -> Result<(), Box<dyn std::error::Error>>
{
let profile = profile_witness(Some(1_000))?;
let within =
context_budget_fact_for_session(&session_with_system_bytes(200), &[], 100, &profile)?;
let exceeded =
context_budget_fact_for_session(&session_with_system_bytes(4_000), &[], 100, &profile)?;
assert_eq!(within.state, ContextBudgetState::Within);
assert_eq!(exceeded.state, ContextBudgetState::ForecastExceeded);
Ok(())
}
#[test]
fn missing_registry_context_window_fails_typed() -> Result<(), Box<dyn std::error::Error>> {
let profile = profile_witness(None)?;
assert!(matches!(
context_budget_fact_for_session(&session_with_system_bytes(20), &[], 0, &profile),
Err(ContextBudgetFactError::ContextWindowUnavailable)
));
Ok(())
}
}