use crate::llmtrim::ir::Request;
use crate::llmtrim::provider::{Provider, Role};
use crate::llmtrim::stages::tools::lex_words;
const CONTEXT_MIN_SEGMENT_CHARS: usize = 600;
fn question_turn(req: &Request, provider: &dyn Provider) -> Option<usize> {
provider
.content_text_pointers(req)
.iter()
.filter(|p| provider.role_at(req, p) == Some(Role::User))
.filter_map(|p| crate::llmtrim::provider::turn_index(p))
.max()
}
pub fn context_text(req: &Request, provider: &dyn Provider) -> String {
let q_turn = question_turn(req, provider);
provider
.content_text_pointers(req)
.iter()
.filter(|p| crate::llmtrim::provider::turn_index(p) != q_turn)
.filter_map(|p| req.get_str(p))
.collect::<Vec<_>>()
.join(" ")
}
pub fn query_terms(req: &Request, provider: &dyn Provider) -> Vec<String> {
let q_turn = question_turn(req, provider);
let mut text = String::new();
for p in provider.content_text_pointers(req) {
let Some(s) = req.get_str(&p) else { continue };
if crate::llmtrim::provider::turn_index(&p) == q_turn
&& s.chars().count() < CONTEXT_MIN_SEGMENT_CHARS
{
text.push_str(s);
text.push(' ');
}
}
lex_words(&text)
}
pub const COVERAGE_THRESHOLD: f64 = 0.5;
pub fn coverage(source: &str, compressed: &str, query_terms: &[String]) -> f64 {
let src = lex_words(source);
if src.is_empty() {
return 1.0; }
let comp = lex_words(compressed);
let q: std::collections::HashSet<String> =
query_terms.iter().flat_map(|t| lex_words(t)).collect();
if q.is_empty() {
return bigram_coverage(&src, &comp);
}
let comp_uni: std::collections::HashSet<&str> = comp.iter().map(String::as_str).collect();
let comp_bi = bigram_set(&comp);
let mut rel_uni: std::collections::HashSet<&str> = std::collections::HashSet::new();
for w in &src {
if q.contains(w) {
rel_uni.insert(w.as_str());
}
}
let mut rel_bi: std::collections::HashSet<(&str, &str)> = std::collections::HashSet::new();
for pair in src.windows(2) {
if q.contains(&pair[0]) || q.contains(&pair[1]) {
rel_bi.insert((pair[0].as_str(), pair[1].as_str()));
}
}
let total = rel_uni.len() + rel_bi.len();
if total == 0 {
return 1.0;
}
let kept = rel_uni.iter().filter(|w| comp_uni.contains(*w)).count()
+ rel_bi
.iter()
.filter(|(a, b)| comp_bi.contains(&(*a, *b)))
.count();
kept as f64 / total as f64
}
fn bigram_coverage(src: &[String], comp: &[String]) -> f64 {
if src.len() < 2 {
if src.is_empty() {
return 1.0;
}
let comp_uni: std::collections::HashSet<&str> = comp.iter().map(String::as_str).collect();
let src_uni: std::collections::HashSet<&str> = src.iter().map(String::as_str).collect();
let kept = src_uni.iter().filter(|w| comp_uni.contains(*w)).count();
return kept as f64 / src_uni.len() as f64;
}
let comp_bi = bigram_set(comp);
let src_bi = bigram_set(src);
let kept = src_bi.iter().filter(|b| comp_bi.contains(b)).count();
kept as f64 / src_bi.len() as f64
}
fn bigram_set(words: &[String]) -> std::collections::HashSet<(&str, &str)> {
words
.windows(2)
.map(|p| (p[0].as_str(), p[1].as_str()))
.collect()
}
pub fn density(source: &str, compressed: &str) -> f64 {
let src = lex_words(source);
let comp = lex_words(compressed);
if comp.is_empty() || src.is_empty() {
return 0.0;
}
let mut fragments: Vec<usize> = Vec::new();
let mut i = 0usize; while i < comp.len() {
let mut best = 0usize;
for s in 0..src.len() {
if src[s] != comp[i] {
continue;
}
let mut len = 0usize;
while s + len < src.len() && i + len < comp.len() && src[s + len] == comp[i + len] {
len += 1;
}
if len > best {
best = len;
}
}
if best == 0 {
fragments.push(0);
i += 1;
} else {
fragments.push(best);
i += best;
}
}
let sum: usize = fragments.iter().sum();
sum as f64 / fragments.len() as f64
}
#[derive(Debug, Clone, Copy)]
pub struct CoverageScore {
pub coverage: f64,
pub answer_kept: bool,
}
pub fn calibrate_threshold(scores: &[CoverageScore], target_recall: f64, alpha: f64) -> f64 {
if scores.is_empty() {
return 0.0;
}
let mut bad: Vec<f64> = scores
.iter()
.filter(|s| !s.answer_kept)
.map(|s| s.coverage)
.collect();
if bad.is_empty() {
return 0.0;
}
bad.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = bad.len();
let rank = (((n + 1) as f64) * (1.0 - alpha)).ceil() as usize;
let idx = rank.clamp(1, n) - 1;
let q = bad[idx];
let tau = (q + f64::EPSILON).min(1.0);
let (acc_total, acc_kept) = scores
.iter()
.filter(|s| s.coverage >= tau)
.fold((0usize, 0usize), |(t, k), s| {
(t + 1, k + usize::from(s.answer_kept))
});
let achieved = if acc_total == 0 {
1.0
} else {
acc_kept as f64 / acc_total as f64
};
if achieved + 1e-9 < target_recall {
return bad[n - 1].min(1.0);
}
tau
}
#[cfg(test)]
mod tests {
use super::*;
fn terms(ws: &[&str]) -> Vec<String> {
ws.iter().map(|s| s.to_string()).collect()
}
#[test]
fn coverage_high_when_query_content_survives() {
let source = "The vault access code is 7741 and the door is on the left.";
let compressed = "vault access code is 7741";
let q = terms(&["vault", "code", "7741"]);
let c = coverage(source, compressed, &q);
assert!(c >= 0.75, "answer span retained → high coverage, got {c}");
assert!(c > COVERAGE_THRESHOLD, "stays above the gate, got {c}");
}
#[test]
fn coverage_drops_when_answer_removed() {
let source = "The vault access code is 7741 and the door is on the left.";
let compressed = "the door is on the left";
let q = terms(&["vault", "code", "7741"]);
let c = coverage(source, compressed, &q);
assert!(
c < COVERAGE_THRESHOLD,
"answer dropped → low coverage, got {c}"
);
}
#[test]
fn coverage_empty_source_is_one() {
assert_eq!(coverage("", "anything", &terms(&["x"])), 1.0);
assert_eq!(coverage("", "", &[]), 1.0);
}
#[test]
fn coverage_empty_compressed_is_zero_when_relevant() {
let c = coverage("the code is 7741", "", &terms(&["code", "7741"]));
assert_eq!(c, 0.0, "everything relevant lost");
}
#[test]
fn coverage_no_query_falls_back_to_bigram_coverage() {
let source = "alpha beta gamma delta epsilon";
let half = coverage(source, "alpha beta gamma", &[]);
assert!(
half > 0.0 && half < 1.0,
"partial bigram coverage, got {half}"
);
assert_eq!(coverage(source, source, &[]), 1.0);
assert_eq!(coverage(source, "totally unrelated words here", &[]), 0.0);
}
#[test]
fn coverage_offtopic_query_is_full() {
let c = coverage(
"a calm bright morning by the sea",
"morning",
&terms(&["zebra", "xyzzy"]),
);
assert_eq!(c, 1.0);
}
#[test]
fn coverage_is_unicode_aware_cjk() {
let source = "金库的访问密码是 七七四一 然后门在左边";
let q = terms(&["密码", "七七四一"]);
let keep = coverage(source, "访问密码是 七七四一", &q);
let drop = coverage(source, "然后门在左边", &q);
assert!(
keep > drop,
"keeping the code covers more ({keep} vs {drop})"
);
assert!(keep >= 0.5, "CJK answer retained, got {keep}");
}
#[test]
fn density_matches_hand_computed_fragments() {
let d = density("a b c d e f", "a b c x e f");
assert!((d - 5.0 / 3.0).abs() < 1e-9, "expected 5/3, got {d}");
}
#[test]
fn density_full_copy_equals_length() {
let d = density("one two three four", "one two three four");
assert!(
(d - 4.0).abs() < 1e-9,
"verbatim copy → density = token count, got {d}"
);
}
#[test]
fn density_empty_is_zero() {
assert_eq!(density("", "abc"), 0.0);
assert_eq!(density("abc", ""), 0.0);
}
fn synthetic_cases() -> Vec<CoverageScore> {
let mut v = Vec::new();
for c in [
0.95, 0.90, 0.86, 0.82, 0.78, 0.74, 0.72, 0.70, 0.66, 0.62, 0.60, 0.58,
] {
v.push(CoverageScore {
coverage: c,
answer_kept: true,
});
}
for c in [
0.10, 0.15, 0.20, 0.25, 0.30, 0.36, 0.42, 0.46, 0.50, 0.50, 0.50, 0.50,
] {
v.push(CoverageScore {
coverage: c,
answer_kept: false,
});
}
v
}
#[test]
fn calibrate_returns_threshold_in_the_gap() {
let tau = calibrate_threshold(&synthetic_cases(), 0.9, 0.1);
assert!(
tau > 0.49 && tau < 0.58,
"threshold lands in the gap, got {tau}"
);
}
#[test]
fn calibrate_holds_on_held_out() {
let all = synthetic_cases();
let (cal, holdout): (Vec<_>, Vec<_>) =
all.iter().enumerate().partition(|(i, _)| i % 2 == 0);
let cal: Vec<CoverageScore> = cal.into_iter().map(|(_, s)| *s).collect();
let holdout: Vec<CoverageScore> = holdout.into_iter().map(|(_, s)| *s).collect();
let target = 0.9;
let tau = calibrate_threshold(&cal, target, 0.1);
let accepted: Vec<&CoverageScore> = holdout.iter().filter(|s| s.coverage >= tau).collect();
assert!(
!accepted.is_empty(),
"threshold must admit some held-out cases"
);
let kept = accepted.iter().filter(|s| s.answer_kept).count();
let rate = kept as f64 / accepted.len() as f64;
assert!(
rate >= target - 1e-9,
"held-out acceptance retains the answer ≥ {target}: got {rate} at τ={tau}"
);
}
#[test]
fn calibrate_all_good_accepts_everything() {
let cases = vec![
CoverageScore {
coverage: 0.9,
answer_kept: true,
},
CoverageScore {
coverage: 0.3,
answer_kept: true,
},
];
assert_eq!(calibrate_threshold(&cases, 0.9, 0.1), 0.0);
}
#[test]
fn shipped_threshold_matches_calibration() {
let tau = calibrate_threshold(&synthetic_cases(), 0.9, 0.1);
assert!(
COVERAGE_THRESHOLD <= tau + 1e-9,
"shipped threshold {COVERAGE_THRESHOLD} must be ≤ calibrated {tau} (no weaker guarantee than calibrated)"
);
}
}
#[cfg(test)]
mod pipeline_tests {
use super::*;
use crate::llmtrim::config::DenseConfig;
use crate::llmtrim::gate::{GateKind, PlanEntry, Scope, Transform};
use crate::llmtrim::ir::ProviderKind;
use crate::llmtrim::pipeline;
use crate::llmtrim::provider::{OpenAiProvider, Provider};
use crate::llmtrim::stages::RetrieveStage;
use crate::llmtrim::tokenizer::counter_for;
use serde_json::{Value, json};
const ANSWER: &str =
"The logistics division reported quarterly revenue of 4.2 million dollars this period.";
fn long_context_with_answer() -> String {
let mut paras = vec![ANSWER.to_string()];
for i in 0..10 {
paras.push(format!(
"Note {i}: the cafeteria menu rotates weekly and parking is available out back."
));
}
let ctx = paras.join("\n\n");
assert!(
ctx.chars().count() >= CONTEXT_MIN_SEGMENT_CHARS,
"context is long enough"
);
ctx
}
fn rag_request() -> Value {
json!({"model":"gpt-4o","messages":[
{"role":"user","content": long_context_with_answer()},
{"role":"user","content":"what was the quarterly revenue for the logistics division?"}]})
}
struct DropAnswer;
impl Transform for DropAnswer {
fn name(&self) -> &str {
"drop-answer"
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn scope(&self) -> Scope {
Scope::Content
}
fn quality_gated(&self) -> bool {
true
}
fn apply(
&self,
req: &mut Request,
_provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> anyhow::Result<()> {
req.set(
"/messages/0/content",
Value::String(
"Weather was pleasant. The cafeteria served pasta today.".to_string(),
),
);
Ok(())
}
}
fn counter() -> Box<dyn crate::llmtrim::tokenizer::TokenCounter> {
counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap()
}
#[test]
fn quality_gate_reverts_a_cut_that_deletes_the_answer() {
let body = rag_request();
let original_ctx = long_context_with_answer();
let mut req = Request::from_value(ProviderKind::OpenAi, body.clone());
let c = counter();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(DropAnswer)];
let mut probe = Request::from_value(ProviderKind::OpenAi, body);
let token_only =
pipeline::run_gated(&mut probe, &OpenAiProvider, c.as_ref(), &stages, false);
assert!(
token_only.stages[0].applied,
"token gate alone accepts the token-saving cut"
);
assert!(token_only.input_tokens_after < token_only.input_tokens_before);
let out = pipeline::run(&mut req, &OpenAiProvider, c.as_ref(), &stages);
assert!(
!out.stages[0].applied,
"quality gate reverts the answer-deleting cut"
);
assert!(
out.stages[0]
.note
.as_deref()
.is_some_and(|n| n.contains("quality-gate")),
"report names the quality-gate revert, got {:?}",
out.stages[0].note
);
assert_eq!(
req.get_str("/messages/0/content"),
Some(original_ctx.as_str()),
"context restored intact after revert"
);
}
#[test]
fn quality_gate_reverts_only_the_offending_stage_in_a_chain() {
let mut req = Request::from_value(ProviderKind::OpenAi, rag_request());
let c = counter();
struct DropFiller;
impl Transform for DropFiller {
fn name(&self) -> &str {
"drop-filler"
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn scope(&self) -> Scope {
Scope::Content
}
fn quality_gated(&self) -> bool {
true
}
fn apply(
&self,
req: &mut Request,
_p: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> anyhow::Result<()> {
req.set(
"/messages/0/content",
Value::String(format!("{ANSWER}\n\n[filler omitted]")),
);
Ok(())
}
}
let stages: Vec<Box<dyn Transform>> = vec![Box::new(DropFiller), Box::new(DropAnswer)];
let out = pipeline::run(&mut req, &OpenAiProvider, c.as_ref(), &stages);
assert!(out.stages[0].applied, "answer-preserving cut sticks");
assert!(
!out.stages[1].applied
&& out.stages[1]
.note
.as_deref()
.is_some_and(|n| n.contains("quality-gate")),
"answer-deleting cut is quality-gate reverted, got {:?}",
out.stages[1].note
);
let final_text = req.get_str("/messages/0/content").unwrap();
assert!(final_text.contains("4.2 million"), "answer survived");
assert!(
final_text.contains("omitted"),
"stage 0's filler-drop stuck"
);
}
#[test]
fn quality_gate_keeps_a_prune_that_retains_the_answer() {
let answer =
"The logistics division quarterly revenue was 4.2 million dollars this period.";
let filler: Vec<String> = (0..8)
.map(|i| format!("Paragraph {i}: the cat sat quietly on the warm windowsill at dawn."))
.collect();
let mut paras = vec![answer.to_string()];
paras.extend(filler);
let context = paras.join("\n\n");
let body = json!({"model":"gpt-4o","messages":[
{"role":"user","content":context},
{"role":"user","content":"what was the quarterly revenue for the logistics division?"}]});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let c = counter();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(RetrieveStage {
keep_ratio: 0.3,
min_segment_chars: 200,
reorder: false,
mmr: false,
mmr_lambda: 0.5,
sentence: false,
})];
let out = pipeline::run(&mut req, &OpenAiProvider, c.as_ref(), &stages);
assert!(
out.stages[0].applied,
"answer-preserving prune is kept (coverage stays above the gate)"
);
assert!(
out.input_tokens_after < out.input_tokens_before,
"tokens still cut"
);
let kept = req.get_str("/messages/0/content").unwrap();
assert!(kept.contains("4.2 million"), "answer survived the prune");
assert!(kept.contains("omitted"), "filler was actually pruned");
}
#[test]
fn quality_gate_off_lets_the_cut_through() {
let mut req = Request::from_value(ProviderKind::OpenAi, rag_request());
let stages: Vec<Box<dyn Transform>> = vec![Box::new(DropAnswer)];
let out = pipeline::run_gated(
&mut req,
&OpenAiProvider,
counter().as_ref(),
&stages,
false,
);
assert!(
out.stages[0].applied,
"with the quality gate off, the token-saving cut is kept"
);
assert_eq!(
req.get_str("/messages/0/content"),
Some("Weather was pleasant. The cafeteria served pasta today."),
"the answer-deleting cut took effect (gate off)"
);
}
#[test]
fn quality_gate_skipped_without_a_distinct_question() {
let body = json!({"model":"gpt-4o","messages":[
{"role":"user","content": long_context_with_answer()}]});
let mut req = Request::from_value(ProviderKind::OpenAi, body.clone());
let q = query_terms(
&Request::from_value(ProviderKind::OpenAi, body),
&OpenAiProvider,
);
assert!(q.is_empty(), "monolithic prompt yields no query anchor");
let stages: Vec<Box<dyn Transform>> = vec![Box::new(DropAnswer)];
let out = pipeline::run(&mut req, &OpenAiProvider, counter().as_ref(), &stages);
assert!(
out.stages[0].applied,
"no query → quality gate skipped → token-saving cut kept"
);
}
#[test]
fn quality_gate_default_is_on() {
assert!(DenseConfig::default().quality_gate, "default ON");
for p in [
"safe",
"auto",
"rag",
"agent",
"code",
"aggressive",
"cache",
"reasoning",
] {
assert!(
DenseConfig::preset(p).unwrap().quality_gate,
"preset `{p}` keeps the quality gate on"
);
}
}
}