use crate::prompt_log::RecallQueryShape;
pub(super) const QUERY_TOKEN_BUDGET: usize = 512;
pub const ENV_QUERY_TOKEN_BUDGET: &str = "TRUSTY_MEMORY_PROMPT_QUERY_TOKENS";
const MIN_QUERY_TOKENS: usize = 32;
const MAX_QUERY_TOKENS: usize = 8192;
const ENVELOPE_OPEN: &str = "<task-notification>";
const ENVELOPE_KEEP: [&str; 2] = ["summary", "result"];
const CHARS_PER_SUBTOKEN: usize = 2;
const NON_ASCII_TOKENS_PER_CHAR: usize = 3;
const SPECIAL_TOKEN_OVERHEAD: usize = 2;
pub(super) struct ShapedQuery {
pub(super) text: String,
pub(super) shape: RecallQueryShape,
}
pub(super) fn shape_recall_query(prompt: &str, budget_tokens: usize) -> ShapedQuery {
let original_tokens = estimate_tokens(prompt);
let (body, envelope_stripped) = match strip_notification_envelope(prompt) {
Some(inner) => (inner, true),
None => (prompt.to_string(), false),
};
let (text, units_dropped) = pack_whole_units(&body, budget_tokens);
let sent_tokens = estimate_tokens(&text);
let sent_tokens_max = max_tokens(&text);
ShapedQuery {
shape: RecallQueryShape {
original_tokens,
sent_tokens,
sent_tokens_max,
budget_tokens,
envelope_stripped,
units_dropped,
},
text,
}
}
pub(super) fn warn_if_reshaped(shape: &RecallQueryShape) {
if !shape.reshaped() {
return;
}
tracing::warn!(
original_tokens = shape.original_tokens,
sent_tokens = shape.sent_tokens,
sent_tokens_max = shape.sent_tokens_max,
budget_tokens = shape.budget_tokens,
envelope_stripped = shape.envelope_stripped,
units_dropped = shape.units_dropped,
may_exceed_window = shape.may_exceed_window(),
"prompt-context: recall query reshaped to fit the embedder window (#4972)"
);
}
pub(super) fn configured_query_budget() -> usize {
clamp_query_budget(std::env::var(ENV_QUERY_TOKEN_BUDGET).ok().as_deref())
}
fn clamp_query_budget(raw: Option<&str>) -> usize {
raw.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|n| *n > 0)
.map(|n| n.clamp(MIN_QUERY_TOKENS, MAX_QUERY_TOKENS))
.unwrap_or(QUERY_TOKEN_BUDGET)
}
fn strip_notification_envelope(prompt: &str) -> Option<String> {
let trimmed = prompt.trim_start();
if !trimmed.starts_with(ENVELOPE_OPEN) {
return None;
}
let mut kept: Vec<&str> = ENVELOPE_KEEP
.iter()
.filter_map(|tag| element_text(trimmed, tag))
.map(str::trim)
.filter(|s| !s.is_empty())
.collect();
const ENVELOPE_CLOSE: &str = "</task-notification>";
if let Some(idx) = trimmed.rfind(ENVELOPE_CLOSE) {
let tail = trimmed[idx + ENVELOPE_CLOSE.len()..].trim();
if !tail.is_empty() {
kept.insert(0, tail);
}
}
if kept.is_empty() {
return None;
}
Some(kept.join("\n\n"))
}
fn element_text<'a>(haystack: &'a str, tag: &str) -> Option<&'a str> {
let open = format!("<{tag}>");
let start = haystack.find(&open)? + open.len();
let rest = &haystack[start..];
let close = format!("</{tag}>");
Some(match rest.find(&close) {
Some(end) => &rest[..end],
None => rest,
})
}
fn pack_whole_units(text: &str, budget_tokens: usize) -> (String, usize) {
if estimate_tokens(text) <= budget_tokens {
return (text.to_string(), 0);
}
let lines: Vec<&str> = text.lines().collect();
let (kept, dropped) = pack_units(&lines, "\n", budget_tokens);
if !kept.trim().is_empty() {
return (kept, dropped);
}
let first = lines
.iter()
.find(|l| !l.trim().is_empty())
.copied()
.unwrap_or(text);
let words: Vec<&str> = first.split_whitespace().collect();
let (kept, dropped_words) = pack_units(&words, " ", budget_tokens);
if !kept.trim().is_empty() {
return (kept, dropped_words + lines.len().saturating_sub(1));
}
let kept = pack_chars(first, budget_tokens);
let dropped = text.chars().count().saturating_sub(kept.chars().count());
(kept, dropped)
}
fn pack_chars(text: &str, budget_tokens: usize) -> String {
let mut kept_bytes = 0usize;
let mut committed = SPECIAL_TOKEN_OVERHEAD;
let mut run = Run::default();
for (idx, ch) in text.char_indices() {
let (next_committed, next_run) = if is_cjk(ch) {
(committed + run.cost(Charge::Calibrated) + 1, Run::default())
} else if ch.is_alphanumeric() {
(committed, run.extended(ch))
} else {
(
committed + run.cost(Charge::Calibrated) + usize::from(!ch.is_whitespace()),
Run::default(),
)
};
if next_committed + next_run.cost(Charge::Calibrated) > budget_tokens {
break;
}
committed = next_committed;
run = next_run;
kept_bytes = idx + ch.len_utf8();
}
text[..kept_bytes].to_string()
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Charge {
Calibrated,
Bound,
}
#[derive(Default, Clone, Copy)]
struct Run {
len: usize,
non_ascii: bool,
has_digit: bool,
}
impl Run {
fn cost(self, charge: Charge) -> usize {
if self.len == 0 {
0
} else if self.non_ascii {
self.len * NON_ASCII_TOKENS_PER_CHAR
} else if self.has_digit || charge == Charge::Bound {
self.len
} else {
self.len.div_ceil(CHARS_PER_SUBTOKEN)
}
}
fn extended(self, ch: char) -> Self {
Self {
len: self.len + 1,
non_ascii: self.non_ascii || !ch.is_ascii(),
has_digit: self.has_digit || ch.is_ascii_digit(),
}
}
}
fn pack_units(units: &[&str], sep: &str, budget: usize) -> (String, usize) {
let mut used = SPECIAL_TOKEN_OVERHEAD;
let mut taken = 0usize;
for unit in units {
let cost = piece_tokens(unit, Charge::Calibrated);
if used + cost > budget {
break;
}
used += cost;
taken += 1;
}
(units[..taken].join(sep), units.len() - taken)
}
pub(super) fn estimate_tokens(text: &str) -> usize {
SPECIAL_TOKEN_OVERHEAD + piece_tokens(text, Charge::Calibrated)
}
pub(super) fn max_tokens(text: &str) -> usize {
SPECIAL_TOKEN_OVERHEAD + piece_tokens(text, Charge::Bound)
}
fn piece_tokens(text: &str, charge: Charge) -> usize {
let mut total = 0usize;
let mut run = Run::default();
for ch in text.chars() {
if is_cjk(ch) {
total += run.cost(charge) + 1;
run = Run::default();
continue;
}
if ch.is_alphanumeric() {
run = run.extended(ch);
continue;
}
total += run.cost(charge) + usize::from(!ch.is_whitespace());
run = Run::default();
}
total + run.cost(charge)
}
fn is_cjk(ch: char) -> bool {
matches!(ch as u32,
0x4E00..=0x9FFF
| 0x3400..=0x4DBF
| 0x20000..=0x2A6DF
| 0x2A700..=0x2B73F
| 0x2B740..=0x2B81F
| 0x2B820..=0x2CEAF
| 0xF900..=0xFAFF
| 0x2F800..=0x2FA1F)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn envelope_strip_recovers_the_payload() {
let raw = concat!(
"<task-notification>\n",
"<task-id>a23c46a0439fa7881</task-id>\n",
"<tool-use-id>toolu_01PvoC76SpHX65DVbJi7sPes</tool-use-id>\n",
"<output-file>/private/tmp/claude-502/-Users-masa-projects/tasks/a23.output",
"</output-file>\n",
"<status>completed</status>\n",
"<summary>Agent \"retrieval floor\" finished</summary>\n",
"<note>A task-notification fires each time this agent stops.</note>\n",
"<result>The relevance floor lands at 0.35.</result>\n",
"</task-notification>"
);
let out = strip_notification_envelope(raw).expect("envelope must be recognised");
assert!(
out.contains("retrieval floor") && out.contains("lands at 0.35"),
"summary and result must both survive; got:\n{out}"
);
for framing in [
"task-id",
"toolu_01PvoC76",
"/private/tmp",
"<status>",
"fires each time",
] {
assert!(
!out.contains(framing),
"envelope framing `{framing}` must not reach the embedder; got:\n{out}"
);
}
}
#[test]
fn envelope_strip_tolerates_a_cut_envelope() {
let raw = "<task-notification>\n<task-id>abc</task-id>\n<result>payload survives";
let out = strip_notification_envelope(raw).expect("cut envelope must still strip");
assert_eq!(out, "payload survives");
}
#[test]
fn short_query_passes_through_untouched() {
let prompt = "how does the relevance floor interact with top_k?";
let shaped = shape_recall_query(prompt, QUERY_TOKEN_BUDGET);
assert_eq!(shaped.text, prompt, "a fitting prompt must not be reshaped");
assert!(
!shaped.shape.reshaped(),
"nothing to report on a clean pass"
);
assert_eq!(shaped.shape.units_dropped, 0);
assert!(!shaped.shape.envelope_stripped);
}
#[test]
fn over_window_query_is_reduced_to_whole_units() {
let line = "explain how the retrieval relevance floor interacts with the top_k cap";
let prompt = vec![line; 400].join("\n");
let shaped = shape_recall_query(&prompt, QUERY_TOKEN_BUDGET);
assert!(
shaped.shape.original_tokens > QUERY_TOKEN_BUDGET,
"fixture must exceed the window; got {} tokens",
shaped.shape.original_tokens
);
assert!(
shaped.shape.sent_tokens <= QUERY_TOKEN_BUDGET,
"the sent query must fit the window; got {} tokens",
shaped.shape.sent_tokens
);
assert!(
shaped.shape.units_dropped > 0 && shaped.shape.reshaped(),
"the reduction must be reported, not silent: {:?}",
shaped.shape
);
for kept in shaped.text.lines() {
assert_eq!(kept, line, "a partial unit reached the wire: {kept:?}");
}
assert!(
!shaped.text.is_empty(),
"reduction must not empty the query"
);
}
#[test]
fn single_oversized_line_falls_back_to_whole_words() {
let prompt = vec!["retrieval"; 2000].join(" ");
let shaped = shape_recall_query(&prompt, QUERY_TOKEN_BUDGET);
assert!(shaped.shape.sent_tokens <= QUERY_TOKEN_BUDGET);
assert!(shaped.shape.units_dropped > 0);
assert!(
shaped.text.split_whitespace().all(|w| w == "retrieval"),
"word fallback must not split a word"
);
}
#[test]
fn token_estimate_never_splits_a_unit() {
assert_eq!(estimate_tokens(""), SPECIAL_TOKEN_OVERHEAD);
assert_eq!(piece_tokens("abcdefghijkl", Charge::Calibrated), 6);
assert_eq!(piece_tokens("ab, cd", Charge::Calibrated), 1 + 1 + 1);
assert_eq!(piece_tokens("a_b", Charge::Calibrated), 1 + 1 + 1);
let units = ["alpha beta", "gamma-delta", "epsilon"];
let joined = units.join(" ");
for charge in [Charge::Calibrated, Charge::Bound] {
assert_eq!(
piece_tokens(&joined, charge),
units.iter().map(|u| piece_tokens(u, charge)).sum::<usize>()
);
}
}
#[test]
fn token_estimate_covers_the_measured_scripts() {
let cases: [(&str, String, usize); 9] = [
(
"latin prose",
"The relevance floor lands at 0.35 after measuring the retrieval corpus. "
.repeat(3),
44,
),
(
"chinese han",
"检索相关性下限设定为零点三五这是经过语料库测量后得到的结果".repeat(10),
292,
),
(
"japanese",
"検索の関連性の下限はコーパスを測定した結果です".repeat(12),
278,
),
(
"korean hangul",
"검색 관련성 하한은 코퍼스를 측정한 결과입니다 ".repeat(12),
638,
),
(
"russian cyrillic",
"Нижняя граница релевантности поиска установлена после измерения корпуса "
.repeat(8),
458,
),
(
"greek",
"Το κατώτατο όριο συνάφειας ανάκτησης ορίστηκε μετά τη μέτρηση ".repeat(8),
410,
),
(
"hex digest",
"9f8e7d6c5b4a39281706f5e4d3c2b1a09f8e7d6c ".repeat(24),
890,
),
(
"jwt",
"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0\
.dozjgNryP4J3jVmNHl0w5N_XgL0n3I9PlFUP0THsR8U"
.to_string(),
78,
),
(
"snake_case identifiers",
"filter_drawers_by_relevance_floor_configured ".repeat(40),
442,
),
];
for (label, input, true_tokens) in cases {
let est = estimate_tokens(&input);
assert!(
est >= true_tokens,
"`{label}` underestimates: estimate {est} < true {true_tokens}. \
An underestimate hands the embedder a query it cuts silently while \
RecallQueryShape reports no loss — always err high."
);
}
}
#[test]
fn latin_compounds_are_not_underestimated() {
let cases: [(&str, &str, usize, usize); 4] = [
(
"german compounds",
"Die Rechtsschutzversicherungsgesellschaft veroeffentlichte \
Geschwindigkeitsbegrenzungen und Arbeiterunfallversicherungsgesetze. ",
10,
442,
),
(
"hungarian agglutinative",
"Megszentsegtelenithetetlensegeskedeseitekert \
elkelkaposztastalanitottatok viszontelnezhetetlenseg. ",
10,
372,
),
(
"finnish compounds",
"Lentokonesuihkuturbiinimoottoriapumekaanikkoaliupseerioppilas \
jarjestelmallistyttamattomyydellansakaan. ",
10,
392,
),
(
"dutch compounds",
"Meervoudigepersoonlijkheidsstoornis levensverzekeringsmaatschappij \
aansprakelijkheidsverzekering. ",
10,
362,
),
];
for (label, unit, repeats, true_tokens) in cases {
let est = estimate_tokens(&unit.repeat(repeats));
assert!(
est >= true_tokens,
"`{label}` underestimates: estimate {est} < true {true_tokens}. \
Ordinary Latin prose, not an edge case — the embedder cuts this \
query and RecallQueryShape reports units_dropped: 0."
);
}
let hungarian = "Megszentsegtelenithetetlensegeskedeseitekert \
elkelkaposztastalanitottatok viszontelnezhetetlenseg. "
.repeat(15);
let shaped = shape_recall_query(&hungarian, QUERY_TOKEN_BUDGET);
assert!(
shaped.shape.original_tokens > QUERY_TOKEN_BUDGET,
"a 557-token Hungarian prompt must be seen as over-window; got {} \
against a budget of {QUERY_TOKEN_BUDGET}",
shaped.shape.original_tokens
);
assert!(
shaped.shape.reshaped() && shaped.shape.units_dropped > 0,
"the reduction must reach the metric, not just happen; got {:?}",
shaped.shape
);
}
#[test]
fn shape_flags_a_send_it_cannot_prove_fits() {
let high_entropy = "qzjvxwkfy ".repeat(80);
let shaped = shape_recall_query(&high_entropy, QUERY_TOKEN_BUDGET);
assert!(
!shaped.shape.reshaped(),
"the estimate clears the budget here — that is the premise of the \
test; got {:?}",
shaped.shape
);
assert!(
shaped.shape.may_exceed_window(),
"the true cost is 562 against a 512 window: the shape must not \
report a clean pass it cannot prove; got {:?}",
shaped.shape
);
let ordinary = "how does the relevance floor interact with top_k?";
let shaped = shape_recall_query(ordinary, QUERY_TOKEN_BUDGET);
assert!(
!shaped.shape.may_exceed_window(),
"a short prompt provably fits — a flag that is always on says \
nothing; got {:?}",
shaped.shape
);
}
#[test]
fn max_tokens_bounds_the_real_tokenizer() {
let cases: [(&str, String, usize); 4] = [
("nine random letters", "qzjvxwkfy ".repeat(80), 562),
("repeated letter", "qqqqqqqqq".to_string(), 11),
(
"hungarian agglutinative",
"Megszentsegtelenithetetlensegeskedeseitekert \
elkelkaposztastalanitottatok viszontelnezhetetlenseg. "
.repeat(10),
372,
),
(
"hex digest",
"9f8e7d6c5b4a39281706f5e4d3c2b1a09f8e7d6c ".repeat(24),
890,
),
];
for (label, input, true_tokens) in &cases {
assert!(
max_tokens(input) >= *true_tokens,
"`{label}` breaks the bound: max_tokens {} < true {true_tokens}",
max_tokens(input)
);
}
for (label, input, true_tokens) in &cases[..2] {
assert!(
estimate_tokens(input) < *true_tokens,
"`{label}` was chosen because the calibrated estimate \
underestimates it; if that stopped being true this test no \
longer proves the bound is load-bearing"
);
}
}
#[test]
fn cjk_is_one_token_per_char() {
let han = "检索相关性下限";
assert_eq!(han.chars().count(), 7);
assert_eq!(
piece_tokens(han, Charge::Calibrated),
7,
"each Han codepoint is its own token"
);
assert!(is_cjk('检') && is_cjk('検') && !is_cjk('a') && !is_cjk('я'));
let long_han = "检索相关性下限设定为零点三五".repeat(110);
assert!(long_han.chars().count() > 1_500);
let shaped = shape_recall_query(&long_han, QUERY_TOKEN_BUDGET);
assert!(
shaped.shape.original_tokens > QUERY_TOKEN_BUDGET,
"a 1500-character Chinese prompt must be seen as over-window; got {} tokens",
shaped.shape.original_tokens
);
assert!(
shaped.shape.reshaped() && shaped.shape.units_dropped > 0,
"it must be reduced and the reduction reported; got {:?}",
shaped.shape
);
assert!(shaped.shape.sent_tokens <= QUERY_TOKEN_BUDGET);
}
#[test]
fn unbreakable_oversized_token_still_yields_a_query() {
let blob = "a".repeat(6_000);
for (label, prompt) in [
("bare", blob.clone()),
("one leading newline", format!("\n{blob}")),
("several blank lines", format!("\n\n\n{blob}")),
("space-only first line", format!(" \n{blob}")),
("crlf", format!("\r\n{blob}")),
] {
let shaped = shape_recall_query(&prompt, QUERY_TOKEN_BUDGET);
assert!(
!shaped.text.trim().is_empty(),
"`{label}`: an unbreakable oversized token must not produce an \
empty query — that returns zero drawers, worse than the \
truncation it replaced"
);
assert!(
shaped.text.len() < prompt.len(),
"`{label}`: it must still be reduced"
);
assert!(
blob.starts_with(shaped.text.trim()),
"`{label}`: the query must be a prefix of the payload; got {:?}",
shaped.text
);
assert!(
shaped.shape.sent_tokens <= QUERY_TOKEN_BUDGET,
"`{label}`: got {} tokens",
shaped.shape.sent_tokens
);
assert!(
shaped.shape.reshaped() && shaped.shape.units_dropped > 0,
"`{label}`: the reduction must be reported; got {:?}",
shaped.shape
);
assert!(
shaped.shape.units_dropped < prompt.chars().count(),
"`{label}`: units_dropped {} must count what was dropped, not \
the whole input ({} chars)",
shaped.shape.units_dropped,
prompt.chars().count()
);
}
}
#[test]
fn envelope_trailing_instruction_outranks_the_payload() {
let instruction = "now open the PR and set the milestone";
let bulky = "The agent swept the retrieval floor and reported back. ".repeat(400);
let raw = format!(
"<task-notification>\n\
<task-id>abc123</task-id>\n\
<result>{bulky}</result>\n\
</task-notification>\n\n\
{instruction}"
);
let shaped = shape_recall_query(&raw, QUERY_TOKEN_BUDGET);
assert!(
shaped.shape.units_dropped > 0,
"fixture must actually overflow the budget for this to mean \
anything; got {:?}",
shaped.shape
);
assert!(
shaped.text.contains(instruction),
"the user's instruction outranks the agent report and must survive \
the packer, not just the strip; got:\n{}",
shaped.text
);
}
#[test]
fn pack_chars_matches_the_estimator() {
let text = "abc检索_x9f8e7d6c Нижняя αβγ 0123456789abcdef!! tail".repeat(20);
for budget in [MIN_QUERY_TOKENS, 64, 128, 512] {
let packed = pack_chars(&text, budget);
assert!(
estimate_tokens(&packed) <= budget,
"pack_chars overshot at budget {budget}: {} tokens for {:?}",
estimate_tokens(&packed),
packed
);
}
}
#[test]
fn envelope_strip_keeps_text_after_the_envelope() {
let raw = "<task-notification>\n\
<task-id>abc123</task-id>\n\
<result>the agent finished the sweep</result>\n\
</task-notification>\n\n\
now open the PR and set the milestone";
let out = strip_notification_envelope(raw).expect("envelope must be recognised");
assert!(
out.contains("now open the PR and set the milestone"),
"user text after the envelope must survive; got:\n{out}"
);
assert!(
out.contains("finished the sweep"),
"payload must survive too"
);
assert!(!out.contains("abc123"), "framing must still go");
}
#[test]
fn envelope_with_no_payload_passes_through_whole() {
let raw = "<task-notification>\n\
<task-id>abc123</task-id>\n\
<output-file>/tmp/x.output</output-file>\n\
<status>completed</status>\n\
<summary> </summary>\n\
</task-notification>";
assert_eq!(
strip_notification_envelope(raw),
None,
"a payload-free envelope must pass through, not become an empty query"
);
let shaped = shape_recall_query(raw, QUERY_TOKEN_BUDGET);
assert!(!shaped.text.is_empty(), "the query must not be empty");
assert!(!shaped.shape.envelope_stripped, "nothing was stripped");
}
#[test]
fn configured_query_budget_clamps_to_bounds() {
assert_eq!(clamp_query_budget(None), QUERY_TOKEN_BUDGET);
assert_eq!(clamp_query_budget(Some("256")), 256);
assert_eq!(clamp_query_budget(Some(" 1024 ")), 1024);
assert_eq!(clamp_query_budget(Some("0")), QUERY_TOKEN_BUDGET);
assert_eq!(clamp_query_budget(Some("1")), MIN_QUERY_TOKENS);
assert_eq!(clamp_query_budget(Some("999999")), MAX_QUERY_TOKENS);
assert_eq!(clamp_query_budget(Some("nonsense")), QUERY_TOKEN_BUDGET);
}
}