use std::collections::{HashMap, HashSet};
use std::sync::RwLock;
use anyhow::Result;
use once_cell::sync::Lazy;
use serde_json::Value;
use unicode_segmentation::UnicodeSegmentation;
use crate::llmtrim::gate::{GateKind, PlanEntry, Transform};
use crate::llmtrim::ir::Request;
use crate::llmtrim::provider::Provider;
pub struct ToolStage {
pub select: bool,
pub trim_desc: bool,
pub minify_schema: bool,
pub max_desc_chars: usize,
}
impl Transform for ToolStage {
fn name(&self) -> &str {
if self.trim_desc { "tool_trim" } else { "tools" }
}
fn gate_kind(&self) -> GateKind {
GateKind::InputTokens
}
fn scope(&self) -> crate::llmtrim::gate::Scope {
crate::llmtrim::gate::Scope::Tools }
fn apply(
&self,
req: &mut Request,
provider: &dyn Provider,
_plan: &mut Vec<PlanEntry>,
) -> Result<()> {
if self.select {
select_tools(req, provider);
}
if self.minify_schema {
minify_tool_schemas(req, self.max_desc_chars);
}
if self.trim_desc {
provider.truncate_tool_descriptions(req, self.max_desc_chars);
}
Ok(())
}
}
const SCHEMA_KEYS: [&str; 2] = ["input_schema", "parameters"];
fn minify_tool_schemas(req: &mut Request, max_desc_chars: usize) {
let Some(Value::Array(tools)) = req.raw_mut().get_mut("tools") else {
return;
};
for tool in tools.iter_mut() {
if let Some(decls) = tool
.get_mut("functionDeclarations")
.and_then(Value::as_array_mut)
{
for d in decls.iter_mut() {
minify_schema_at(d, max_desc_chars);
}
continue;
}
let scope = match tool.get_mut("function").filter(|f| f.is_object()) {
Some(f) => f,
None => tool,
};
minify_schema_at(scope, max_desc_chars);
}
}
fn minify_schema_at(scope: &mut Value, max_desc_chars: usize) {
let Some(obj) = scope.as_object_mut() else {
return;
};
for key in SCHEMA_KEYS {
if let Some(schema) = obj.get_mut(key) {
crate::llmtrim::stages::tool_schema::minify_schema(schema, max_desc_chars);
}
}
}
pub(crate) fn detect_lang(sample: &str) -> Option<whatlang::Lang> {
whatlang::detect(sample)
.filter(|info| info.is_reliable())
.map(|info| info.lang())
}
pub(crate) fn stopword_set(sample: &str) -> &'static HashSet<&'static str> {
use stop_words::LANGUAGE as L;
use whatlang::Lang;
let head = if sample.len() > LANG_DETECT_MAX_BYTES {
let mut end = LANG_DETECT_MAX_BYTES;
while !sample.is_char_boundary(end) {
end -= 1;
}
&sample[..end]
} else {
sample
};
let language = match detect_lang(head) {
Some(Lang::Fra) => L::French,
Some(Lang::Spa) => L::Spanish,
Some(Lang::Deu) => L::German,
Some(Lang::Ita) => L::Italian,
Some(Lang::Por) => L::Portuguese,
Some(Lang::Nld) => L::Dutch,
Some(Lang::Rus) => L::Russian,
Some(Lang::Jpn) => L::Japanese,
Some(Lang::Kor) => L::Korean,
Some(Lang::Cmn) => L::Chinese,
Some(Lang::Ara) => L::Arabic,
Some(Lang::Tur) => L::Turkish,
Some(Lang::Pol) => L::Polish,
Some(Lang::Swe) => L::Swedish,
Some(Lang::Dan) => L::Danish,
Some(Lang::Fin) => L::Finnish,
Some(Lang::Ell) => L::Greek,
Some(Lang::Hun) => L::Hungarian,
Some(Lang::Ron) => L::Romanian,
Some(Lang::Ces) => L::Czech,
Some(Lang::Ukr) => L::Ukrainian,
Some(Lang::Vie) => L::Vietnamese,
Some(Lang::Ind) => L::Indonesian,
Some(Lang::Hin) => L::Hindi,
_ => L::English,
};
intern_stopwords(stop_words::get(language))
}
fn intern_stopwords(words: &'static [&'static str]) -> &'static HashSet<&'static str> {
static CACHE: Lazy<RwLock<HashMap<usize, &'static HashSet<&'static str>>>> =
Lazy::new(|| RwLock::new(HashMap::new()));
let key = words.as_ptr() as usize;
if let Some(&set) = CACHE
.read()
.expect("stopword cache lock poisoned")
.get(&key)
{
return set;
}
CACHE
.write()
.expect("stopword cache lock poisoned")
.entry(key)
.or_insert_with(|| Box::leak(Box::new(words.iter().copied().collect())))
}
pub(crate) fn lex_words(s: &str) -> Vec<String> {
s.unicode_words().map(str::to_lowercase).collect()
}
pub(crate) fn fnv1a(bytes: impl IntoIterator<Item = u8>) -> u64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for b in bytes {
h ^= b as u64;
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
h
}
fn content_words<'a>(lower: &'a str, stop: &HashSet<&str>) -> HashSet<&'a str> {
lower
.unicode_words()
.flat_map(|w| w.split('_'))
.filter(|w| w.len() >= 2 && !stop.contains(w))
.collect()
}
const LANG_SAMPLE_BYTES: usize = 2048;
const LANG_DETECT_MAX_BYTES: usize = 8 * 1024;
const TOOL_QUERY_MAX_BYTES: usize = 16 * 1024;
const FIELD_W_NAME: f64 = 4.0;
const FIELD_W_PARAMS: f64 = 2.0;
const FIELD_W_DESC: f64 = 1.0;
const BM25_K1: f64 = 1.2;
const BM25_B: f64 = 0.75;
fn select_tools(req: &mut Request, provider: &dyn Provider) {
if !is_first_turn(req) {
return;
}
let descriptors = provider.tool_descriptors(req);
if descriptors.len() < 2 {
return; }
let pointers = provider.content_text_pointers(req);
let mut lower = String::new();
for p in pointers.iter().rev() {
if let Some(s) = req.get_str(p) {
lower.push_str(&s.to_lowercase());
lower.push(' ');
if lower.len() >= TOOL_QUERY_MAX_BYTES {
break;
}
}
}
let sample_end = lower.len().min(LANG_SAMPLE_BYTES);
let stop = stopword_set(lower.get(..sample_end).unwrap_or(&lower));
let query = content_words(&lower, stop);
if query.is_empty() {
return;
}
let param_fields = tool_param_words(req, stop);
let docs: Vec<ToolDoc> = descriptors
.iter()
.enumerate()
.map(|(i, (name, desc))| ToolDoc {
name: bag(&content_words(&name.to_lowercase(), stop)),
params: param_fields.get(i).cloned().unwrap_or_default(),
desc: bag(&content_words(&desc.to_lowercase(), stop)),
})
.collect();
let scores = bm25f_scores(&docs, &query);
let mut mentioned: HashSet<&str> = HashSet::new();
for p in pointers.iter() {
if mentioned.len() == descriptors.len() {
break;
}
if let Some(s) = req.get_str(p) {
for (name, _) in descriptors.iter() {
if !mentioned.contains(name.as_str()) && contains_standalone(s, name) {
mentioned.insert(name.as_str());
}
}
}
}
let keep: Vec<bool> = descriptors
.iter()
.zip(&scores)
.map(|((name, _), &s)| mentioned.contains(name.as_str()) || s > 0.0)
.collect();
if keep.iter().any(|&k| k) {
provider.retain_tools(req, &keep);
}
}
fn contains_standalone(text: &str, name: &str) -> bool {
if name.is_empty() {
return false;
}
let is_ident = |c: char| c.is_alphanumeric() || c == '_';
let mut from = 0;
while let Some(i) = text[from..].find(name) {
let start = from + i;
let end = start + name.len();
let before_ok = text[..start]
.chars()
.next_back()
.is_none_or(|c| !is_ident(c));
let after_ok = text[end..].chars().next().is_none_or(|c| !is_ident(c));
if before_ok && after_ok {
return true;
}
from = start + name.chars().next().map_or(1, char::len_utf8);
}
false
}
#[derive(Default)]
struct ToolDoc {
name: Vec<(String, u32)>,
params: Vec<(String, u32)>,
desc: Vec<(String, u32)>,
}
fn bag(words: &HashSet<&str>) -> Vec<(String, u32)> {
words.iter().map(|w| (w.to_string(), 1)).collect()
}
fn tool_param_words(req: &Request, stop: &HashSet<&str>) -> Vec<Vec<(String, u32)>> {
let Some(tools) = req.raw().get("tools").and_then(Value::as_array) else {
return Vec::new();
};
let mut out = Vec::new();
for tool in tools {
if let Some(decls) = tool.get("functionDeclarations").and_then(Value::as_array) {
for d in decls {
out.push(prop_name_bag(d, stop));
}
continue;
}
let scope = tool
.get("function")
.filter(|f| f.is_object())
.unwrap_or(tool);
out.push(prop_name_bag(scope, stop));
}
out
}
fn prop_name_bag(scope: &Value, stop: &HashSet<&str>) -> Vec<(String, u32)> {
let props = SCHEMA_KEYS
.iter()
.find_map(|k| scope.pointer(&format!("/{k}/properties")))
.and_then(Value::as_object);
let Some(props) = props else {
return Vec::new();
};
let joined = props
.keys()
.cloned()
.collect::<Vec<_>>()
.join(" ")
.to_lowercase();
bag(&content_words(&joined, stop))
}
fn bm25f_scores(docs: &[ToolDoc], query: &HashSet<&str>) -> Vec<f64> {
let n = docs.len();
if n == 0 {
return Vec::new();
}
let field_len = |d: &ToolDoc, f: usize| -> f64 {
let v = [&d.name, &d.params, &d.desc][f];
v.iter().map(|(_, c)| *c as u64).sum::<u64>() as f64
};
let weights = [FIELD_W_NAME, FIELD_W_PARAMS, FIELD_W_DESC];
let mut avg = [0.0f64; 3];
for (f, a) in avg.iter_mut().enumerate() {
let total: f64 = docs.iter().map(|d| field_len(d, f)).sum();
*a = (total / n as f64).max(1.0); }
let term_df = |term: &str| -> usize {
docs.iter()
.filter(|d| {
[&d.name, &d.params, &d.desc]
.iter()
.any(|fld| fld.iter().any(|(w, _)| w == term))
})
.count()
};
let idf = |term: &str| -> f64 {
let df = term_df(term) as f64;
(((n as f64 - df + 0.5) / (df + 0.5)) + 1.0).ln().max(0.0)
};
let idfs: Vec<(&str, f64)> = query.iter().map(|t| (*t, idf(t))).collect();
docs.iter()
.map(|d| {
let fields = [&d.name, &d.params, &d.desc];
let lens = [field_len(d, 0), field_len(d, 1), field_len(d, 2)];
idfs.iter()
.map(|(term, w_idf)| {
let mut tf = 0.0f64;
for f in 0..3 {
let raw = fields[f]
.iter()
.find(|(w, _)| w == term)
.map_or(0u32, |(_, c)| *c) as f64;
if raw > 0.0 {
let norm = 1.0 - BM25_B + BM25_B * lens[f] / avg[f];
tf += weights[f] * raw / norm;
}
}
if tf > 0.0 {
w_idf * tf / (BM25_K1 + tf)
} else {
0.0
}
})
.sum()
})
.collect()
}
pub(crate) fn is_first_turn(req: &Request) -> bool {
if !tools_used_in_history(req).is_empty() {
return false;
}
let raw = req.raw();
let non_system_turns = ["messages", "input", "contents"]
.iter()
.filter_map(|k| raw.get(*k).and_then(Value::as_array))
.map(|a| {
a.iter()
.filter(|m| m.get("role").and_then(Value::as_str) != Some("system"))
.count()
})
.max()
.unwrap_or(0);
non_system_turns <= 1
}
fn tools_used_in_history(req: &Request) -> HashSet<String> {
let mut used = HashSet::new();
let raw = req.raw();
let turns = raw
.get("messages")
.or_else(|| raw.get("contents"))
.or_else(|| raw.get("input"))
.and_then(Value::as_array);
let Some(turns) = turns else {
return used;
};
for m in turns {
if m.get("type").and_then(Value::as_str) == Some("function_call")
&& let Some(n) = m.get("name").and_then(Value::as_str)
{
used.insert(n.to_string());
}
if let Some(calls) = m.get("tool_calls").and_then(Value::as_array) {
for c in calls {
if let Some(n) = c.pointer("/function/name").and_then(Value::as_str) {
used.insert(n.to_string());
}
}
}
if let Some(blocks) = m.get("content").and_then(Value::as_array) {
for b in blocks {
if b.get("type").and_then(Value::as_str) == Some("tool_use")
&& let Some(n) = b.get("name").and_then(Value::as_str)
{
used.insert(n.to_string());
}
}
}
if let Some(parts) = m.get("parts").and_then(Value::as_array) {
for p in parts {
if let Some(n) = p.pointer("/functionCall/name").and_then(Value::as_str) {
used.insert(n.to_string());
}
}
}
}
used
}
pub(crate) fn is_structured_segment(text: &str) -> bool {
let t = text.trim();
if t.is_empty() {
return false;
}
if t.starts_with(['{', '[']) && serde_json::from_str::<Value>(t).is_ok() {
return true;
}
if t.contains("},{") || t.contains("}, {") {
return true;
}
let lines: Vec<&str> = t.lines().map(str::trim).filter(|l| !l.is_empty()).collect();
if lines.len() >= 3 {
for delim in [',', '\t', '|', ';'] {
let counts: Vec<usize> = lines.iter().map(|l| l.matches(delim).count()).collect();
let cols = counts.iter().copied().max().unwrap_or(0);
if cols >= 1 && counts.iter().filter(|&&c| c == cols).count() * 4 >= lines.len() * 3 {
return true; }
}
}
if lines.len() >= 3 && lines.iter().filter(|l| is_kv_line(l)).count() * 4 >= lines.len() * 3 {
return true;
}
let mut symbols = 0usize;
let mut nonspace = 0usize;
for c in t.chars() {
if c.is_whitespace() {
continue;
}
nonspace += 1;
if !c.is_alphanumeric() {
symbols += 1;
}
}
nonspace >= 40 && symbols * 100 >= nonspace * 22
}
fn is_kv_line(line: &str) -> bool {
match line.find([':', '=']) {
Some(i) if i > 0 && i + 1 < line.len() => {
let key = line[..i].trim();
!key.is_empty() && key.chars().count() <= 40 && !key.contains(['.', '!', '?'])
}
_ => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llmtrim::ir::ProviderKind;
use crate::llmtrim::pipeline;
use crate::llmtrim::provider::{AnthropicProvider, OpenAiProvider};
use crate::llmtrim::tokenizer::counter_for;
use serde_json::{Value, json};
#[test]
fn structured_detects_json_and_record_arrays() {
assert!(is_structured_segment("[{\"a\":1},{\"a\":2}]"));
assert!(is_structured_segment("{\"k\": \"v\"}"));
assert!(is_structured_segment(
"[{\"occupation\":\"Sales\"},{\"occupation\":\"Tech\"}] then a question"
));
}
#[test]
fn structured_detects_csv_tsv_and_markdown_tables() {
assert!(is_structured_segment(
"name,age,city\nJohn,30,NYC\nJane,25,LA\nBob,40,SF"
));
assert!(is_structured_segment(
"| col | val |\n|-----|-----|\n| a | 1 |\n| b | 2 |"
));
}
#[test]
fn structured_detects_key_value_config() {
assert!(is_structured_segment(
"host: localhost\nport: 8080\ndebug: true\nname: app"
));
assert!(is_structured_segment("KEY=val\nFOO=bar\nBAZ=qux"));
}
#[test]
fn structured_detects_code_by_symbol_density() {
assert!(is_structured_segment(
"for (let i = 0; i < n; i++) { out[i] = (a[i] + b[i]) * w - bias / 2; }"
));
}
#[test]
fn prose_is_not_structured_in_any_script() {
assert!(!is_structured_segment(
"The quick brown fox jumps over the lazy dog. It was a calm, bright morning, \
and nothing at all seemed out of the ordinary on that particular day."
));
assert!(!is_structured_segment(
"这是一段用于测试的中文散文文本,它包含足够多的汉字以超过长度阈值,\
但是标点符号很少,因此不应该被误判成结构化数据或者表格。"
));
assert!(!is_structured_segment(
"Note: this is an ordinary sentence that merely happens to contain a colon."
));
}
fn openai_tools() -> Value {
json!([
{"type":"function","function":{"name":"get_weather","description":"Get the weather forecast for a city","parameters":{}}},
{"type":"function","function":{"name":"send_email","description":"Send an email to a recipient","parameters":{}}},
{"type":"function","function":{"name":"run_sql","description":"Execute a SQL query against the database","parameters":{}}}
])
}
fn select_stage() -> Box<dyn Transform> {
Box::new(ToolStage {
select: true,
trim_desc: false,
minify_schema: false,
max_desc_chars: 200,
})
}
#[test]
fn stage_name_reflects_lossy_trim() {
let trimming = ToolStage {
select: false,
trim_desc: true,
minify_schema: false,
max_desc_chars: 200,
};
assert_eq!(trimming.name(), "tool_trim");
let lossless = ToolStage {
select: true,
trim_desc: false,
minify_schema: true,
max_desc_chars: 200,
};
assert_eq!(lossless.name(), "tools");
}
#[test]
fn openai_selection_keeps_relevant_tool() {
let body = json!({
"model":"gpt-4o",
"messages":[{"role":"user","content":"what is the weather forecast in Paris today?"}],
"tools": openai_tools()
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let out = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
assert!(
out.stages[0].applied,
"dropping irrelevant tools reduces tokens"
);
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert_eq!(names, vec!["get_weather"], "only the weather tool is kept");
}
#[test]
fn selection_skipped_once_a_tool_was_invoked() {
let body = json!({
"model":"gpt-4o",
"messages":[
{"role":"assistant","tool_calls":[{"function":{"name":"run_sql"}}]},
{"role":"user","content":"now what is the weather forecast in Paris?"}
],
"tools": openai_tools()
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let out = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
assert!(!out.stages[0].applied, "selection skipped mid-loop");
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert_eq!(
names,
vec!["get_weather", "send_email", "run_sql"],
"the full toolset ships unchanged mid-loop: {names:?}"
);
}
#[test]
fn keeps_all_when_nothing_matches() {
let body = json!({
"model":"gpt-4o",
"messages":[{"role":"user","content":"hello there friend"}],
"tools": openai_tools()
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let _ = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
assert_eq!(
req.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.len(),
3,
"weak query keeps the whole toolset (safety)"
);
}
#[test]
fn french_query_uses_french_stopwords() {
let body = json!({
"model":"gpt-4o",
"messages":[{"role":"user","content":"quelle est la météo pour la ville de Paris aujourd'hui"}],
"tools":[
{"type":"function","function":{"name":"meteo","description":"Obtenir les prévisions météo pour une ville","parameters":{}}},
{"type":"function","function":{"name":"envoyer_email","description":"Envoyer un courriel à un destinataire","parameters":{}}}
]
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let _ = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert_eq!(names, vec!["meteo"], "only the weather tool kept (French)");
}
#[test]
fn anthropic_selection_and_trim() {
let long_desc = "x".repeat(400);
let body = json!({
"max_tokens":100,
"messages":[{"role":"user","content":"run a sql query on the orders table please"}],
"tools":[
{"name":"run_sql","description": long_desc,"input_schema":{}},
{"name":"get_weather","description":"weather forecast","input_schema":{}}
]
});
let mut req = Request::from_value(ProviderKind::Anthropic, body);
let counter = counter_for(ProviderKind::Anthropic, None).unwrap();
let stages: Vec<Box<dyn Transform>> = vec![Box::new(ToolStage {
select: true,
trim_desc: true,
minify_schema: false,
max_desc_chars: 50,
})];
pipeline::run(&mut req, &AnthropicProvider, counter.as_ref(), &stages);
let tools = req.raw().get("tools").and_then(Value::as_array).unwrap();
assert_eq!(tools.len(), 1, "only run_sql kept");
let desc = tools[0].get("description").and_then(Value::as_str).unwrap();
assert!(
desc.chars().count() <= 51,
"description trimmed to max+ellipsis"
);
}
#[test]
fn google_function_call_not_dropped() {
let body = json!({
"contents": [
{"role": "user", "parts": [{"text": "what is the weather?"}]},
{"role": "model", "parts": [{"functionCall": {"name": "get_weather", "args": {"city": "Paris"}}}]},
{"role": "user", "parts": [{"functionResponse": {"name": "get_weather", "response": {"temp": "15°C"}}}]}
],
"tools": [{"functionDeclarations": [
{"name": "get_weather", "description": "weather forecast"},
{"name": "run_sql", "description": "database query"}
]}]
});
let req = Request::from_value(ProviderKind::Google, body);
let used = tools_used_in_history(&req);
assert!(
used.contains("get_weather"),
"Google functionCall must be tracked"
);
}
#[test]
fn bm25f_name_match_outranks_description_match() {
let mk = |words: &[&str]| -> Vec<(String, u32)> {
words.iter().map(|w| (w.to_string(), 1)).collect()
};
let docs = vec![
ToolDoc {
name: mk(&["weather"]),
params: mk(&["city"]),
desc: mk(&["forecast", "data"]),
},
ToolDoc {
name: mk(&["search"]),
params: Vec::new(),
desc: mk(&["look", "up", "the", "weather", "and", "more"]),
},
];
let query: HashSet<&str> = ["weather"].into_iter().collect();
let scores = bm25f_scores(&docs, &query);
assert!(
scores[0] > scores[1],
"name-field match must score above description-only: {scores:?}"
);
}
#[test]
fn bm25f_selection_prefers_name_field_end_to_end() {
let body = json!({
"model":"gpt-4o",
"messages":[{"role":"user","content":"please run_sql on the users table"}],
"tools":[
{"type":"function","function":{"name":"run_sql","description":"Execute a query","parameters":{"type":"object","properties":{"query":{"type":"string"}}}}},
{"type":"function","function":{"name":"notes","description":"Keep notes about sql and other topics you discuss","parameters":{}}}
]
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let _ = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert!(
names.contains(&"run_sql"),
"name-matched tool kept: {names:?}"
);
}
#[test]
fn bm25f_already_invoked_survives_at_rank_bottom() {
let body = json!({
"model":"gpt-4o",
"messages":[
{"role":"assistant","tool_calls":[{"function":{"name":"send_email"}}]},
{"role":"user","content":"now run a sql query on the orders table"}
],
"tools": openai_tools()
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let _ = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert!(names.contains(&"run_sql"), "relevant tool kept");
assert!(
names.contains(&"send_email"),
"already-invoked tool kept despite zero relevance now: {names:?}"
);
}
#[test]
fn bm25f_is_deterministic() {
let mk = |words: &[&str]| -> Vec<(String, u32)> {
words.iter().map(|w| (w.to_string(), 1)).collect()
};
let docs = vec![
ToolDoc {
name: mk(&["weather", "city"]),
params: mk(&["city"]),
desc: mk(&["forecast"]),
},
ToolDoc {
name: mk(&["sql", "run"]),
params: mk(&["query"]),
desc: mk(&["database"]),
},
ToolDoc {
name: mk(&["email"]),
params: mk(&["to"]),
desc: mk(&["send", "message"]),
},
];
let query: HashSet<&str> = ["weather", "city", "run"].into_iter().collect();
let a = bm25f_scores(&docs, &query);
let b = bm25f_scores(&docs, &query);
assert_eq!(a, b, "BM25F is deterministic");
}
#[test]
fn bm25f_param_property_names_are_scored() {
let body = json!({
"model":"gpt-4o",
"messages":[{"role":"user","content":"set the timezone please"}],
"tools":[
{"type":"function","function":{"name":"configure","description":"adjust settings","parameters":{"type":"object","properties":{"timezone":{"type":"string"}}}}},
{"type":"function","function":{"name":"unrelated","description":"does other things","parameters":{}}}
]
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let _ = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert!(
names.contains(&"configure"),
"tool selected via its parameter name: {names:?}"
);
}
#[test]
fn is_first_turn_across_wire_shapes() {
use ProviderKind::{Anthropic, Google, OpenAi};
let ft = |kind, body: Value| is_first_turn(&Request::from_value(kind, body));
assert!(ft(
OpenAi,
json!({"messages":[{"role":"system","content":"s"},{"role":"user","content":"hi"}]})
));
assert!(!ft(
OpenAi,
json!({"messages":[
{"role":"system","content":"s"},{"role":"user","content":"hi"},
{"role":"assistant","tool_calls":[{"function":{"name":"f"}}]},
{"role":"tool","tool_call_id":"1","content":"x"},
{"role":"user","content":"again"}]})
));
assert!(ft(
Anthropic,
json!({"system":"s","messages":[{"role":"user","content":"hi"}]})
));
assert!(!ft(
Anthropic,
json!({"system":"s","messages":[
{"role":"user","content":"hi"},
{"role":"assistant","content":[{"type":"tool_use","name":"f","input":{}}]}]})
));
assert!(ft(
Google,
json!({"contents":[{"role":"user","parts":[{"text":"hi"}]}]})
));
assert!(!ft(
Google,
json!({"contents":[
{"role":"user","parts":[{"text":"hi"}]},
{"role":"model","parts":[{"functionCall":{"name":"f"}}]}]})
));
assert!(ft(
OpenAi,
json!({"instructions":"s","input":[{"role":"user","content":"hi"}]})
));
assert!(!ft(
OpenAi,
json!({"instructions":"s","input":[
{"type":"function_call","name":"f","arguments":"{}","call_id":"1"}]})
));
}
#[test]
fn multi_turn_conversation_is_not_pruned() {
let body = json!({
"model":"gpt-4o",
"messages":[
{"role":"user","content":"Use ToolSearch with query select:<name> to load tool schemas before calling them."},
{"role":"user","content":"what is the weather forecast in Paris today?"}
],
"tools":[
{"type":"function","function":{"name":"get_weather","description":"Get the weather forecast for a city","parameters":{}}},
{"type":"function","function":{"name":"ToolSearch","description":"Load deferred tool schemas","parameters":{}}},
{"type":"function","function":{"name":"send_email","description":"Send an email to a recipient","parameters":{}}}
]
});
let mut req = Request::from_value(ProviderKind::OpenAi, body);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let out = pipeline::run(
&mut req,
&OpenAiProvider,
counter.as_ref(),
&[select_stage()],
);
assert!(
!out.stages[0].applied,
"selection is skipped past the first turn (no churn of the cached tool block)"
);
let names: Vec<&str> = req
.raw()
.get("tools")
.and_then(Value::as_array)
.unwrap()
.iter()
.filter_map(|t| t.pointer("/function/name").and_then(Value::as_str))
.collect();
assert_eq!(
names,
vec!["get_weather", "ToolSearch", "send_email"],
"the full toolset ships unchanged mid-conversation: {names:?}"
);
}
#[test]
fn contains_standalone_word_boundaries() {
assert!(contains_standalone("use ToolSearch now", "ToolSearch"));
assert!(contains_standalone(
"query select:ToolSearch.",
"ToolSearch"
)); assert!(contains_standalone("ToolSearch", "ToolSearch")); assert!(!contains_standalone("MyToolSearcher", "ToolSearch")); assert!(!contains_standalone("ToolSearch_v2", "ToolSearch")); assert!(!contains_standalone("use toolsearch now", "ToolSearch")); assert!(contains_standalone(
"call mcp__server__thing here",
"mcp__server__thing"
));
assert!(!contains_standalone(
"mcp__server__thing2",
"mcp__server__thing"
));
assert!(!contains_standalone("préToolSearch", "ToolSearch"));
}
#[test]
fn stage_minifies_openai_and_anthropic_schemas() {
let verbose = json!({
"$schema": "https://json-schema.org/draft/2020-12/schema",
"title": "Args",
"type": "object",
"additionalProperties": false,
"properties": {"q": {"type": ["string"], "title": "Q", "description": "text"}},
"required": ["q"]
});
let oa = json!({
"model":"gpt-4o",
"messages":[{"role":"user","content":"hi"}],
"tools":[{"type":"function","function":{"name":"search","description":"d","parameters": verbose.clone()}}]
});
let mut req = Request::from_value(ProviderKind::OpenAi, oa);
let counter = counter_for(ProviderKind::OpenAi, Some("gpt-4o")).unwrap();
let stage: Vec<Box<dyn Transform>> = vec![Box::new(ToolStage {
select: false,
trim_desc: false,
minify_schema: true,
max_desc_chars: 300,
})];
pipeline::run(&mut req, &OpenAiProvider, counter.as_ref(), &stage);
let schema = req.raw().pointer("/tools/0/function/parameters").unwrap();
assert!(schema.get("$schema").is_none(), "$schema dropped");
assert!(schema.get("title").is_none(), "root title dropped");
assert_eq!(schema.get("additionalProperties"), Some(&json!(false)));
assert_eq!(schema.get("required"), Some(&json!(["q"])));
assert_eq!(
schema.pointer("/properties/q/type").and_then(Value::as_str),
Some("string"),
"single-type array collapsed"
);
let an = json!({
"max_tokens": 100,
"messages":[{"role":"user","content":"hi"}],
"tools":[{"name":"search","description":"d","input_schema": verbose}]
});
let mut req = Request::from_value(ProviderKind::Anthropic, an);
let counter = counter_for(ProviderKind::Anthropic, None).unwrap();
let stage: Vec<Box<dyn Transform>> = vec![Box::new(ToolStage {
select: false,
trim_desc: false,
minify_schema: true,
max_desc_chars: 300,
})];
pipeline::run(&mut req, &AnthropicProvider, counter.as_ref(), &stage);
let schema = req.raw().pointer("/tools/0/input_schema").unwrap();
assert!(schema.get("$schema").is_none(), "anthropic $schema dropped");
assert_eq!(schema.get("additionalProperties"), Some(&json!(false)));
assert_eq!(
schema.pointer("/properties/q/type").and_then(Value::as_str),
Some("string")
);
}
}