use serde_json::Value;
use crate::limits::TokenLimits;
const CHARS_PER_TOKEN: usize = 2;
pub const OVERHEAD_TOKENS: u32 = 8192;
pub trait HistoryMessage: Clone {
fn role(&self) -> &str;
fn content(&self) -> &Value;
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct Trimmed<M> {
pub messages: Vec<M>,
pub dropped: usize,
}
fn message_tokens<M: HistoryMessage>(msg: &M) -> u32 {
let len = serde_json::to_string(msg.content())
.map(|s| s.len())
.unwrap_or(0)
+ msg.role().len();
u32::try_from(len / CHARS_PER_TOKEN + 8).unwrap_or(u32::MAX)
}
pub fn estimate_tokens<M: HistoryMessage>(messages: &[M]) -> u32 {
messages
.iter()
.map(message_tokens)
.fold(0u32, u32::saturating_add)
}
fn has_block<M: HistoryMessage>(msg: &M, kind: &str) -> bool {
msg.content().as_array().is_some_and(|blocks| {
blocks
.iter()
.any(|b| b.get("type").and_then(Value::as_str) == Some(kind))
})
}
fn has_tool_use<M: HistoryMessage>(msg: &M) -> bool {
has_block(msg, "tool_use")
}
fn has_tool_result<M: HistoryMessage>(msg: &M) -> bool {
has_block(msg, "tool_result")
}
pub fn trim_history<M: HistoryMessage>(messages: Vec<M>, budget: Option<u32>) -> Trimmed<M> {
let Some(budget) = budget else {
return Trimmed {
messages,
dropped: 0,
};
};
if messages.is_empty() || estimate_tokens(&messages) <= budget {
return Trimmed {
messages,
dropped: 0,
};
}
let first_user = messages
.iter()
.position(|m| m.role() == "user" && !has_tool_result(m));
let reserve = first_user
.map(|i| message_tokens(&messages[i]))
.filter(|t| *t <= budget / 2)
.unwrap_or(0);
let recent_budget = budget - reserve;
let mut keep_from = messages.len();
let mut used = 0u32;
for (i, m) in messages.iter().enumerate().rev() {
let t = message_tokens(m);
if used.saturating_add(t) > recent_budget {
break;
}
used += t;
keep_from = i;
}
while keep_from < messages.len() && has_tool_result(&messages[keep_from]) {
keep_from += 1;
}
let mut keep_to = messages.len();
if keep_from < keep_to && has_tool_use(&messages[keep_to - 1]) {
keep_to -= 1;
}
if keep_from >= keep_to {
let fallback: Vec<M> = messages
.iter()
.rev()
.find(|m| !has_tool_use(*m) && !has_tool_result(*m))
.cloned()
.into_iter()
.collect();
return Trimmed {
dropped: messages.len() - fallback.len(),
messages: fallback,
};
}
let mut kept: Vec<M> = messages[keep_from..keep_to].to_vec();
if let Some(i) = first_user {
if reserve > 0 && i < keep_from {
kept.insert(0, messages[i].clone());
}
}
let dropped = messages.len() - kept.len();
Trimmed {
messages: kept,
dropped,
}
}
pub fn history_budget(limits: Option<&TokenLimits>, max_tokens: u32) -> Option<u32> {
limits?
.input_budget(max_tokens)
.map(|b| b.saturating_sub(OVERHEAD_TOKENS))
}
pub fn retry_budget<M: HistoryMessage>(messages: &[M]) -> u32 {
estimate_tokens(messages) / 2
}
pub fn is_context_overflow(status: u16, body: &str) -> bool {
if status == 413 {
return true;
}
if !matches!(status, 400 | 422) {
return false;
}
let b = body.to_ascii_lowercase();
let about_output_only = (b.contains("max_tokens") || b.contains("max_completion_tokens"))
&& !b.contains("context")
&& !b.contains("prompt")
&& !b.contains("input");
if about_output_only {
return false;
}
const PATTERNS: &[&str] = &[
"context_length_exceeded",
"maximum context length",
"context length",
"context window",
"prompt is too long",
"prompt too long",
"too many tokens",
"exceeds the maximum number of tokens",
"exceeded model token limit",
"input is too long",
"input length",
"input tokens exceed",
"请求内容过长",
"超长",
"输入长度",
"超过最大长度",
"上下文长度",
"超出模型",
];
PATTERNS.iter().any(|p| b.contains(p))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[derive(Debug, Clone, PartialEq)]
struct Msg {
role: String,
content: Value,
}
impl HistoryMessage for Msg {
fn role(&self) -> &str {
&self.role
}
fn content(&self) -> &Value {
&self.content
}
}
fn text(role: &str, s: &str) -> Msg {
Msg {
role: role.into(),
content: json!(s),
}
}
fn tool_use_msg(id: &str) -> Msg {
Msg {
role: "assistant".into(),
content: json!([{ "type": "tool_use", "id": id, "name": "x", "input": {} }]),
}
}
fn tool_result_msg(id: &str, payload: &str) -> Msg {
Msg {
role: "user".into(),
content: json!([{ "type": "tool_result", "tool_use_id": id, "content": payload }]),
}
}
#[test]
fn unknown_budget_never_trims() {
let msgs = vec![text("user", &"x".repeat(100_000))];
let out = trim_history(msgs.clone(), None);
assert_eq!(out.dropped, 0);
assert_eq!(out.messages, msgs);
}
#[test]
fn under_budget_is_untouched() {
let msgs = vec![text("user", "hi"), text("assistant", "hello")];
let out = trim_history(msgs.clone(), Some(100_000));
assert_eq!(out.dropped, 0);
assert_eq!(out.messages, msgs);
}
#[test]
fn keeps_recent_drops_old() {
let msgs = vec![
text("user", &"a".repeat(4000)),
text("assistant", &"b".repeat(4000)),
text("user", &"c".repeat(4000)),
text("assistant", &"d".repeat(40)),
];
let out = trim_history(msgs, Some(2100));
assert!(out.dropped > 0, "该裁却没裁");
let last = out.messages.last().unwrap();
assert_eq!(last.role, "assistant");
assert!(last.content.as_str().unwrap().contains("dddd"));
}
#[test]
fn never_starts_with_orphan_tool_result() {
let msgs = vec![
text("user", &"a".repeat(8000)),
tool_use_msg("t1"),
tool_result_msg("t1", &"r".repeat(8000)),
text("assistant", "done"),
];
let out = trim_history(msgs, Some(4200));
assert!(!has_tool_result(&out.messages[0]), "{:?}", out.messages[0]);
}
#[test]
fn never_ends_with_unanswered_tool_use() {
let msgs = vec![
text("user", "start"),
text("assistant", &"x".repeat(20_000)),
tool_use_msg("t9"),
];
let out = trim_history(msgs, Some(200));
assert!(!out.messages.last().is_some_and(has_tool_use));
}
#[test]
fn preserves_first_user_message_and_counts_correctly() {
let msgs = vec![
text("user", "把这个仓库的测试全部跑一遍"),
text("assistant", &"x".repeat(20_000)),
text("user", "继续"),
];
let total = msgs.len();
let out = trim_history(msgs, Some(600));
assert!(out.dropped > 0);
assert_eq!(out.dropped, total - out.messages.len());
assert!(out.messages[0]
.content
.as_str()
.unwrap()
.contains("把这个仓库"));
}
#[test]
fn oversized_first_user_is_not_forced_back() {
let msgs = vec![
text("user", &"日志".repeat(20_000)),
text("assistant", "看完了"),
text("user", "那第三行什么意思"),
];
let out = trim_history(msgs, Some(200));
assert!(estimate_tokens(&out.messages) <= 200, "{:?}", out.messages);
assert!(!out
.messages
.iter()
.any(|m| m.content.as_str().is_some_and(|c| c.len() > 1000)));
assert_eq!(
out.messages.last().unwrap().content,
json!("那第三行什么意思")
);
}
#[test]
fn extreme_budget_keeps_last_intact_message() {
let msgs = vec![
text("user", &"a".repeat(10_000)),
text("assistant", &"b".repeat(10_000)),
];
let out = trim_history(msgs, Some(1));
assert_eq!(out.messages.len(), 1);
assert_eq!(out.dropped, 1);
assert_eq!(out.messages[0].role, "assistant");
}
#[test]
fn extreme_budget_still_respects_tool_pairing() {
let msgs = vec![text("user", &"a".repeat(10_000)), tool_use_msg("t1")];
let out = trim_history(msgs, Some(1));
assert!(!out.messages.iter().any(has_tool_use));
assert!(!out.messages.iter().any(has_tool_result));
}
#[test]
fn budget_subtracts_output_and_overhead() {
let l = TokenLimits::from_endpoint(Some(128_000), Some(8192));
assert_eq!(
history_budget(Some(&l), 4096),
Some(128_000 - 4096 - OVERHEAD_TOKENS)
);
let tiny = TokenLimits::from_preset(Some(1000), None);
assert_eq!(history_budget(Some(&tiny), 4096), Some(0), "不下溢");
assert_eq!(history_budget(None, 4096), None);
let no_window = TokenLimits::from_preset(None, Some(8192));
assert_eq!(history_budget(Some(&no_window), 4096), None);
}
#[test]
fn retry_budget_halves_until_nothing_left_to_drop() {
let mut cur = vec![text("user", "任务")];
for _ in 0..8 {
cur.push(text("assistant", &"x".repeat(2000)));
cur.push(text("user", &"y".repeat(2000)));
}
let before = estimate_tokens(&cur);
let mut last = u32::MAX;
let mut rounds = 0;
for _ in 0..3 {
let b = retry_budget(&cur);
assert!(b < last, "预算必须单调减小");
last = b;
let next = trim_history(cur.clone(), Some(b));
if next.dropped == 0 {
break;
}
assert_eq!(
next.messages[0].content,
json!("任务"),
"任务描述必须一直在"
);
cur = next.messages;
rounds += 1;
}
assert!(rounds >= 1);
assert!(estimate_tokens(&cur) <= before / 2, "一轮后至少减半");
}
#[test]
fn detects_real_overflow_errors() {
let cases: &[(u16, &str)] = &[
(
400,
r#"{"error":{"message":"This model's maximum context length is 131072 tokens. However, you requested 140000 tokens (135904 in the messages, 4096 in the completion).","type":"invalid_request_error","code":"context_length_exceeded"}}"#,
),
(
400,
r#"{"type":"error","error":{"type":"invalid_request_error","message":"prompt is too long: 215000 tokens > 200000 maximum"}}"#,
),
(
400,
r#"[{"error":{"code":400,"message":"The input token count (1200000) exceeds the maximum number of tokens allowed (1048576).","status":"INVALID_ARGUMENT"}}]"#,
),
(
400,
r#"{"error":{"message":"This endpoint's maximum context length is 200000 tokens. However, you requested about 230000 tokens","code":400}}"#,
),
(
400,
r#"{"error":{"message":"Invalid request: Your request exceeded model token limit: 131072","type":"invalid_request_error"}}"#,
),
(
400,
r#"{"error":{"message":"Range of input length should be [1, 129024]","type":"invalid_request_error","code":"invalid_parameter_error"}}"#,
),
(400, r#"{"error":{"code":"1261","message":"Prompt 超长"}}"#),
(413, "Request Entity Too Large"),
];
for (status, body) in cases {
assert!(is_context_overflow(*status, body), "漏判:{status} {body}");
}
}
#[test]
fn does_not_misfire() {
let cases: &[(u16, &str)] = &[
(
429,
r#"{"error":{"message":"Rate limit reached for gpt-x on tokens per min (TPM): Limit 30000, Used 29000","type":"tokens"}}"#,
),
(
400,
r#"{"error":{"message":"max_tokens is too large: 50000. This model supports at most 8192 completion tokens.","type":"invalid_request_error"}}"#,
),
(
400,
r#"{"type":"error","error":{"type":"invalid_request_error","message":"max_tokens: 64000 > 32000, which is the maximum allowed number of output tokens"}}"#,
),
(401, r#"{"error":{"message":"Incorrect API key provided"}}"#),
(
400,
r#"{"error":{"message":"The model `gpt-9` does not exist"}}"#,
),
(500, "maximum context length"),
];
for (status, body) in cases {
assert!(!is_context_overflow(*status, body), "误判:{status} {body}");
}
}
}