use std::borrow::Cow;
use tracing::warn;
use crate::classify::tiers::llm::LlmClassifier;
use crate::classify::tiers::llm_context::{with_context, CommitContext};
use crate::core::config::{
default_max_input_tokens_for, LlmSource, LLM_DEFAULT_MAX_INPUT_TOKENS,
LLM_MIN_INPUT_ROOM_TOKENS,
};
pub const BYTES_PER_TOKEN: usize = 1;
pub(crate) const USER_PREFIX: &str = "Classify this commit message:\n\n";
const FRAMING_TOKENS: usize = 64;
const MARKER_OPEN: &str = "\n[truncated ";
const MARKER_CLOSE: &str = " bytes]";
const MARKER_MAX: usize = MARKER_OPEN.len() + 20 + MARKER_CLOSE.len();
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PromptBudget {
max_input_tokens: usize,
}
impl Default for PromptBudget {
fn default() -> Self {
Self::new(LLM_DEFAULT_MAX_INPUT_TOKENS)
}
}
impl PromptBudget {
pub fn new(max_input_tokens: usize) -> Self {
Self { max_input_tokens }
}
pub fn for_source(source: &LlmSource) -> Self {
Self::new(default_max_input_tokens_for(source))
}
pub fn max_input_tokens(&self) -> usize {
self.max_input_tokens
}
fn room(self, system: &str) -> usize {
let fixed = system.len() + USER_PREFIX.len() + FRAMING_TOKENS * BYTES_PER_TOKEN;
self.max_input_tokens
.saturating_mul(BYTES_PER_TOKEN)
.saturating_sub(fixed)
}
pub(crate) fn fit<'a>(
self,
system: &str,
message: &'a str,
ctx: Option<&CommitContext>,
) -> Option<Cow<'a, str>> {
let room = self.room(system);
if room < LLM_MIN_INPUT_ROOM_TOKENS * BYTES_PER_TOKEN {
warn!(
max_input_tokens = self.max_input_tokens,
system_bytes = system.len(),
"LLM input budget leaves too little room for the commit; nothing sent (#178)"
);
return None;
}
let full = with_context(message, ctx);
if full.len() <= room {
return Some(full);
}
let block = ctx.map(CommitContext::render_plain).unwrap_or_default();
let block_room = room.saturating_sub(message.len()).max(room / 4);
let block = cut(&block, block_room);
let message = cut(message, room.saturating_sub(block.len()));
let text = format!("{message}{block}");
debug_assert!(text.len() <= room);
warn!(
prompt_bytes = full.len(),
sent_bytes = text.len(),
max_input_tokens = self.max_input_tokens,
"LLM prompt over the input budget; sending it truncated (#178)"
);
Some(Cow::Owned(text))
}
}
fn cut(text: &str, max: usize) -> Cow<'_, str> {
if text.len() <= max {
return Cow::Borrowed(text);
}
let mut end = max.saturating_sub(MARKER_MAX);
while !text.is_char_boundary(end) {
end -= 1;
}
let dropped = text.len() - end;
Cow::Owned(format!(
"{}{MARKER_OPEN}{dropped}{MARKER_CLOSE}",
&text[..end]
))
}
impl LlmClassifier {
pub fn with_max_input_tokens(mut self, max_input_tokens: Option<usize>) -> Self {
self.budget = max_input_tokens.map(PromptBudget::new);
self
}
pub fn prompt_budget(&self) -> PromptBudget {
self.budget.unwrap_or_else(|| {
PromptBudget::for_source(&match self.provider_label() {
"bedrock" => LlmSource::Bedrock,
"anthropic-api" => LlmSource::AnthropicApi,
_ => LlmSource::Openrouter,
})
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::classify::tiers::llm::SYSTEM_PROMPT;
#[test]
fn default_budget_fits_every_default_model() {
assert_eq!(BYTES_PER_TOKEN, 1, "tokens <= bytes is the bound relied on");
let openai = PromptBudget::for_source(&LlmSource::Openrouter).max_input_tokens();
assert_eq!(openai, PromptBudget::default().max_input_tokens());
assert_eq!(openai, 100_000);
assert!(openai + 2_048 + FRAMING_TOKENS <= 128_000, "{openai}");
for claude in [LlmSource::Bedrock, LlmSource::AnthropicApi] {
let tokens = PromptBudget::for_source(&claude).max_input_tokens();
assert_eq!(tokens, 190_000, "{claude:?}");
assert!(tokens + 2_048 + FRAMING_TOKENS <= 200_000, "{tokens}");
}
}
#[tokio::test]
async fn room_under_the_floor_sends_nothing() {
let fixed = SYSTEM_PROMPT.len() + USER_PREFIX.len() + FRAMING_TOKENS;
let at_floor = PromptBudget::new(fixed + LLM_MIN_INPUT_ROOM_TOKENS);
assert!(at_floor.fit(SYSTEM_PROMPT, "fix: a", None).is_some());
let under = PromptBudget::new(fixed + LLM_MIN_INPUT_ROOM_TOKENS - 1);
assert!(under.fit(SYSTEM_PROMPT, "fix: a", None).is_none());
assert!(PromptBudget::new(0)
.fit(SYSTEM_PROMPT, "fix: a", None)
.is_none());
let server = wiremock::MockServer::start().await;
let llm = LlmClassifier::new("m", Some("sk-test".to_string()))
.with_endpoint(format!("{}/v1/chat/completions", server.uri()))
.with_max_input_tokens(Some(under.max_input_tokens()));
let call = llm.classify_detailed("fix: a").await;
assert_eq!(
call.outcome,
crate::classify::tiers::llm_prompt::LlmOutcome::Failed
);
assert!(call.verdict.is_none());
let sent = server.received_requests().await.expect("recording on");
assert!(sent.is_empty(), "{} requests sent", sent.len());
}
#[test]
fn cut_keeps_whole_characters_and_counts_the_dropped_bytes() {
let text = "é€😀".repeat(40);
for max in MARKER_MAX..text.len() {
let out = cut(&text, max);
assert!(out.len() <= max, "max {max}: {} bytes", out.len());
let (kept, note) = out.split_once(MARKER_OPEN).expect("marker");
assert!(text.starts_with(kept));
assert_eq!(note, format!("{}{MARKER_CLOSE}", text.len() - kept.len()));
}
assert!(matches!(cut(&text, text.len()), Cow::Borrowed(_)));
}
#[test]
fn bedrock_request_fits_the_budget() {
let message = "chore: vendored lockfile ".repeat(42_000);
let budget = PromptBudget::for_source(&LlmSource::Bedrock);
let text = budget.fit(SYSTEM_PROMPT, &message, None).expect("room");
let req = crate::classify::tiers::bedrock::converse_request("m", SYSTEM_PROMPT, &text);
let json = serde_json::to_value(&req).expect("serialize");
let sent: usize = json["messages"]
.as_array()
.expect("messages")
.iter()
.map(|m| m["content"].as_str().expect("text").len())
.sum();
assert!(sent <= budget.max_input_tokens(), "{sent} bytes sent");
assert!(text.contains(MARKER_OPEN), "marker present");
let whole = PromptBudget::new(usize::MAX).fit(SYSTEM_PROMPT, &message, None);
assert!(matches!(whole, Some(Cow::Borrowed(m)) if m == message));
}
#[tokio::test]
async fn configured_budget_is_the_one_applied() {
use crate::core::config::LlmConfig;
use crate::core::creds::CredentialSource;
let cfg: LlmConfig =
serde_yaml::from_str("source: openrouter\napi_key_env: K\nmax_input_tokens: 10000\n")
.expect("parse");
let unset: LlmConfig = serde_yaml::from_str("source: openrouter\n").expect("parse");
assert_eq!(unset.max_input_tokens, None);
let creds = CredentialSource::fixed([("K", "sk-test")]); let llm = LlmClassifier::from_llm_config_with_creds(&cfg, "m", &creds)
.await
.expect("build");
assert_eq!(llm.prompt_budget(), PromptBudget::new(10_000));
let unset_cfg = LlmConfig {
api_key_env: "K".to_string(),
..unset
};
let llm = LlmClassifier::from_llm_config_with_creds(&unset_cfg, "m", &creds)
.await
.expect("build");
assert_eq!(llm.prompt_budget(), PromptBudget::default());
let message = "x".repeat(40_000);
let whole = PromptBudget::default().fit("sys", &message, None);
assert_eq!(whole.as_deref(), Some(message.as_str()));
let text = PromptBudget::new(10_000)
.fit("sys", &message, None)
.expect("room");
assert!(
3 + USER_PREFIX.len() + text.len() <= 10_000,
"{}",
text.len()
);
let (kept, note) = text.split_once(MARKER_OPEN).expect("marker");
assert_eq!(note, format!("{}{MARKER_CLOSE}", 40_000 - kept.len()));
}
}