use std::collections::HashMap;
use crate::llmtrim::gate::{GateKind, PlanEntry, Scope, Transform};
use crate::llmtrim::ir::Request;
use crate::llmtrim::provider::Provider;
use crate::llmtrim::quality_gate::{self, COVERAGE_THRESHOLD};
use crate::llmtrim::tokenizer::{TokenCounter, Tokens};
fn tools_unchanged(req: &Request, snapshot: &Request) -> bool {
req.raw().get("tools") == snapshot.raw().get("tools")
}
fn quality_gated_by_name(name: &str) -> bool {
matches!(name, "retrieve" | "toolout")
}
#[derive(Debug, Clone)]
pub struct StageReport {
pub name: String,
pub applied: bool,
pub tokens_before: Tokens,
pub tokens_after: Tokens,
pub note: Option<String>,
}
#[derive(Debug, Clone, Default)]
pub struct PipelineOutcome {
pub stages: Vec<StageReport>,
pub plan: Vec<PlanEntry>,
pub input_tokens_before: Tokens,
pub input_tokens_after: Tokens,
pub frozen_input_tokens: Tokens,
}
fn count_content(req: &Request, provider: &dyn Provider, counter: &dyn TokenCounter) -> usize {
provider
.content_text_pointers(req)
.iter()
.filter_map(|p| req.get_str(p))
.map(|s| counter.count(s))
.sum()
}
fn count_content_cached(
req: &Request,
provider: &dyn Provider,
counter: &dyn TokenCounter,
cache: &mut HashMap<String, usize>,
) -> usize {
provider
.content_text_pointers(req)
.iter()
.filter_map(|p| req.get_str(p))
.map(|s| match cache.get(s) {
Some(&c) => c,
None => {
let c = counter.count(s);
cache.insert(s.to_string(), c);
c
}
})
.sum()
}
fn count_tools(req: &Request, counter: &dyn TokenCounter) -> usize {
req.raw()
.get("tools")
.map_or(0, |tools| counter.count(&tools.to_string()))
}
pub fn content_tokens(req: &Request, provider: &dyn Provider, counter: &dyn TokenCounter) -> usize {
count_content(req, provider, counter) + count_tools(req, counter)
}
pub fn run(
req: &mut Request,
provider: &dyn Provider,
counter: &dyn TokenCounter,
stages: &[Box<dyn Transform>],
) -> PipelineOutcome {
run_gated(req, provider, counter, stages, true)
}
pub fn run_gated(
req: &mut Request,
provider: &dyn Provider,
counter: &dyn TokenCounter,
stages: &[Box<dyn Transform>],
quality_gate: bool,
) -> PipelineOutcome {
let mut plan: Vec<PlanEntry> = Vec::new();
let mut reports = Vec::with_capacity(stages.len());
let query = if quality_gate {
quality_gate::query_terms(req, provider)
} else {
Vec::new()
};
let profile = std::env::var_os("LLMTRIM_PROFILE").is_some();
let mut seg_cache: HashMap<String, usize> = HashMap::new();
let mut content = count_content_cached(req, provider, counter, &mut seg_cache);
let mut tools = count_tools(req, counter);
let input_tokens_before = content + tools;
let frozen_input_tokens: usize = crate::llmtrim::cache_zone::frozen_pointers(req, provider)
.iter()
.filter_map(|p| req.get_str(p))
.map(|s| {
seg_cache
.get(s)
.copied()
.unwrap_or_else(|| counter.count(s))
})
.sum();
for stage in stages {
let scope = stage.scope();
let before = content + tools;
let snapshot = req.clone();
let plan_mark = plan.len();
let timer = profile.then(std::time::Instant::now);
let (applied, after, note) = match stage.apply(req, provider, &mut plan) {
Err(e) => {
*req = snapshot;
plan.truncate(plan_mark);
(false, before, Some(format!("error: {e}")))
}
Ok(()) => {
let (new_content, new_tools) = if req.raw() == snapshot.raw() {
(content, tools)
} else {
let new_content = match scope {
Scope::Tools => content,
Scope::Content | Scope::Both => {
count_content_cached(req, provider, counter, &mut seg_cache)
}
};
let new_tools = if tools_unchanged(req, &snapshot) {
tools
} else {
count_tools(req, counter)
};
(new_content, new_tools)
};
let after = new_content + new_tools;
if stage.gate_kind() == GateKind::InputTokens && after >= before {
*req = snapshot;
plan.truncate(plan_mark);
(false, before, Some("no token reduction".to_string()))
} else if quality_gate
&& !query.is_empty()
&& stage.gate_kind() == GateKind::InputTokens
&& scope != Scope::Tools
&& (stage.quality_gated() || quality_gated_by_name(stage.name()))
&& {
let source = quality_gate::context_text(&snapshot, provider);
let compressed = quality_gate::context_text(req, provider);
quality_gate::coverage(&source, &compressed, &query) < COVERAGE_THRESHOLD
}
{
*req = snapshot;
plan.truncate(plan_mark);
(
false,
before,
Some("quality-gate reverted: coverage below threshold".to_string()),
)
} else {
content = new_content;
tools = new_tools;
(true, after, None)
}
}
};
if let Some(timer) = timer {
eprintln!(
"llmtrim profile: {:>14} {:>7.2} ms {:>6} -> {:<6} tok{}",
stage.name(),
timer.elapsed().as_secs_f64() * 1000.0,
before,
after,
if applied { "" } else { " (reverted)" },
);
}
reports.push(StageReport {
name: stage.name().to_string(),
applied,
tokens_before: Tokens(before),
tokens_after: Tokens(after),
note,
});
}
let input_tokens_after = content + tools;
PipelineOutcome {
stages: reports,
plan,
input_tokens_before: Tokens(input_tokens_before),
input_tokens_after: Tokens(input_tokens_after),
frozen_input_tokens: Tokens(frozen_input_tokens),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llmtrim::gate::GateKind;
use crate::llmtrim::ir::ProviderKind;
use crate::llmtrim::provider::OpenAiProvider;
use crate::llmtrim::tokenizer::counter_for;
use serde_json::Value;
struct SetContent {
name: String,
text: String,
}
impl Transform for SetContent {
fn name(&self) -> &str {
&self.name
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn apply(
&self,
req: &mut Request,
_provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> anyhow::Result<()> {
req.set("/messages/0/content", Value::String(self.text.clone()));
Ok(())
}
}
struct AddSystem;
impl Transform for AddSystem {
fn name(&self) -> &str {
"add-system"
}
fn gate_kind(&self) -> GateKind {
GateKind::OutputShaping
}
fn apply(
&self,
req: &mut Request,
provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> anyhow::Result<()> {
provider.add_system_instruction(req, "no preamble, no restating the question");
Ok(())
}
}
struct Boom;
impl Transform for Boom {
fn name(&self) -> &str {
"boom"
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn apply(
&self,
req: &mut Request,
_provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> anyhow::Result<()> {
req.set("/messages/0/content", Value::String("damaged".to_string()));
anyhow::bail!("intentional failure");
}
}
fn fresh() -> (Request, Box<dyn TokenCounter>) {
let req = Request::parse(
ProviderKind::OpenAi,
r#"{"messages":[{"role":"user","content":"this is a fairly long original message about widgets and gadgets"}]}"#,
)
.unwrap();
(
req,
counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap(),
)
}
#[test]
fn input_stage_applies_when_it_shrinks() {
let (mut req, counter) = fresh();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(SetContent {
name: "shrink".into(),
text: "hi".into(),
})];
let out = run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(out.stages[0].applied);
assert_eq!(req.get_str("/messages/0/content"), Some("hi"));
assert!(out.input_tokens_after < out.input_tokens_before);
}
#[test]
fn input_stage_reverts_when_it_bloats() {
let (mut req, counter) = fresh();
let original = req.get_str("/messages/0/content").unwrap().to_string();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(SetContent {
name: "bloat".into(),
text: "word ".repeat(80),
})];
let out = run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(!out.stages[0].applied);
assert_eq!(out.stages[0].note.as_deref(), Some("no token reduction"));
assert_eq!(req.get_str("/messages/0/content"), Some(original.as_str()));
}
#[test]
fn output_shaping_stage_is_never_reverted_on_tokens() {
let (mut req, counter) = fresh();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(AddSystem)];
let out = run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(
out.stages[0].applied,
"output-shaping must apply despite adding tokens"
);
assert_eq!(
req.raw()
.pointer("/messages/0/role")
.and_then(Value::as_str),
Some("system")
);
}
#[test]
fn erroring_stage_is_reverted_and_does_not_block() {
let (mut req, counter) = fresh();
let original = req.get_str("/messages/0/content").unwrap().to_string();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(Boom)];
let out = run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(!out.stages[0].applied);
assert!(out.stages[0].note.as_deref().unwrap().contains("error"));
assert_eq!(req.get_str("/messages/0/content"), Some(original.as_str()));
}
struct ShrinkTools;
impl Transform for ShrinkTools {
fn name(&self) -> &str {
"shrink-tools"
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn scope(&self) -> Scope {
Scope::Tools
}
fn apply(
&self,
req: &mut Request,
_provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> anyhow::Result<()> {
req.set("/tools", serde_json::json!([{"name": "f"}]));
Ok(())
}
}
fn fresh_with_tools() -> (Request, Box<dyn TokenCounter>) {
let req = Request::parse(
ProviderKind::OpenAi,
r#"{"messages":[{"role":"user","content":"this is a fairly long original message about widgets and gadgets"}],"tools":[{"type":"function","function":{"name":"search_documents","description":"search a large corpus for relevant passages","parameters":{"q":"string"}}}]}"#,
)
.unwrap();
(
req,
counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap(),
)
}
#[test]
fn content_stage_preserves_tools_token_count() {
let (mut req, counter) = fresh_with_tools();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(SetContent {
name: "shrink".into(),
text: "hi".into(),
})];
let out = run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(out.stages[0].applied);
let independent = content_tokens(&req, &OpenAiProvider, counter.as_ref());
assert_eq!(out.input_tokens_after.0, independent);
assert!(out.input_tokens_after.0 > count_content(&req, &OpenAiProvider, counter.as_ref()));
}
#[test]
fn tools_stage_recounts_when_tools_change() {
let (mut req, counter) = fresh_with_tools();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(ShrinkTools)];
let out = run(&mut req, &OpenAiProvider, counter.as_ref(), &stages);
assert!(out.stages[0].applied);
assert!(out.input_tokens_after < out.input_tokens_before);
assert_eq!(
out.input_tokens_after.0,
content_tokens(&req, &OpenAiProvider, counter.as_ref())
);
}
}