use std::collections::HashSet;
use edgecrab_types::ToolSchema;
use crate::tool_schema_index::{is_hot_tool, partition_schemas};
#[derive(Debug, Clone)]
pub struct CatalogEntry {
pub name: String,
pub description: String,
pub tokens: Vec<String>,
}
fn tokenize(text: &str) -> Vec<String> {
if text.is_empty() {
return Vec::new();
}
text.split(|c: char| !c.is_ascii_alphanumeric())
.filter(|t| !t.is_empty())
.map(|t| t.to_ascii_lowercase())
.collect()
}
fn entry_search_text(schema: &ToolSchema) -> String {
let name_words = schema.name.replace(['_', '.', '-'], " ");
let param_names: String = schema
.parameters
.get("properties")
.and_then(|v| v.as_object())
.map(|props| props.keys().cloned().collect::<Vec<_>>().join(" "))
.unwrap_or_default();
format!("{name_words} {} {param_names}", schema.description)
}
fn catalog_entry(schema: &ToolSchema) -> CatalogEntry {
let text = entry_search_text(schema);
CatalogEntry {
name: schema.name.clone(),
description: schema.description.clone(),
tokens: tokenize(&text),
}
}
pub fn build_deferred_catalog(
schemas: &[ToolSchema],
materialized: &HashSet<String>,
) -> Vec<CatalogEntry> {
let (_, deferred) = partition_schemas(schemas, materialized);
deferred
.iter()
.filter(|s| !is_hot_tool(&s.name))
.map(|schema| catalog_entry(schema))
.collect()
}
pub fn build_registry_catalog(schemas: &[ToolSchema]) -> Vec<CatalogEntry> {
use crate::tool_schema_index::TOOL_SEARCH_NAME;
schemas
.iter()
.filter(|s| s.name != TOOL_SEARCH_NAME)
.map(catalog_entry)
.collect()
}
fn bm25_score(
query_tokens: &[String],
doc_tokens: &[String],
avg_dl: f64,
doc_freq: &std::collections::HashMap<String, usize>,
n_docs: usize,
) -> f64 {
if doc_tokens.is_empty() {
return 0.0;
}
const K1: f64 = 1.5;
const B: f64 = 0.75;
let dl = doc_tokens.len() as f64;
let mut doc_tf: std::collections::HashMap<&str, usize> = std::collections::HashMap::new();
for t in doc_tokens {
*doc_tf.entry(t.as_str()).or_insert(0) += 1;
}
let mut score = 0.0;
for q in query_tokens {
let df = doc_freq.get(q).copied().unwrap_or(0);
if df == 0 {
continue;
}
let idf = ((n_docs as f64 - df as f64 + 0.5) / (df as f64 + 0.5) + 1.0).ln();
let tf = doc_tf.get(q.as_str()).copied().unwrap_or(0) as f64;
if tf == 0.0 {
continue;
}
let norm = tf * (K1 + 1.0) / (tf + K1 * (1.0 - B + B * dl / avg_dl.max(1.0)));
score += idf * norm;
}
score
}
pub fn search_deferred_catalog(catalog: &[CatalogEntry], query: &str, limit: usize) -> Vec<String> {
if catalog.is_empty() || limit == 0 {
return Vec::new();
}
let query_tokens = tokenize(query);
if query_tokens.is_empty() {
return Vec::new();
}
let doc_lengths: Vec<usize> = catalog.iter().map(|e| e.tokens.len()).collect();
let avg_dl = doc_lengths.iter().sum::<usize>() as f64 / doc_lengths.len().max(1) as f64;
let mut doc_freq: std::collections::HashMap<String, usize> = std::collections::HashMap::new();
for entry in catalog {
let seen: HashSet<&str> = entry.tokens.iter().map(String::as_str).collect();
for t in seen {
*doc_freq.entry(t.to_string()).or_insert(0) += 1;
}
}
let n_docs = catalog.len();
let mut scored: Vec<(f64, &CatalogEntry)> = catalog
.iter()
.filter_map(|entry| {
let s = bm25_score(&query_tokens, &entry.tokens, avg_dl, &doc_freq, n_docs);
if s > 0.0 { Some((s, entry)) } else { None }
})
.collect();
if scored.is_empty() {
let ql = query.to_ascii_lowercase();
for entry in catalog {
if entry.name.to_ascii_lowercase().contains(&ql) {
scored.push((0.1, entry));
}
}
}
scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
scored
.into_iter()
.take(limit)
.map(|(_, e)| e.name.clone())
.collect()
}
pub fn looks_like_create_file_intent(user_text: &str) -> bool {
let lower = user_text.to_ascii_lowercase();
let has_create_verb = lower.contains("write ")
|| lower.contains("write a")
|| lower.contains("create ")
|| lower.contains("scaffold")
|| lower.contains("generate ");
let has_path_or_artifact = lower.contains("./")
|| lower.contains("demo/")
|| lower.contains("src/")
|| lower.contains(".html")
|| lower.contains(".js")
|| lower.contains(".css")
|| lower.contains(".rs")
|| lower.contains(".py")
|| lower.contains(".ts")
|| lower.contains(".tsx")
|| lower.contains(".md");
has_create_verb && has_path_or_artifact
}
pub fn prefetch_tools_for_user_message(
user_text: &str,
schemas: &[ToolSchema],
materialized: &HashSet<String>,
max_prefetch: usize,
) -> Vec<String> {
let query = user_text.trim();
if query.is_empty() || max_prefetch == 0 {
return Vec::new();
}
let catalog = build_deferred_catalog(schemas, materialized);
let mut hits = search_deferred_catalog(&catalog, query, max_prefetch);
let mcp_hits = crate::tools::mcp_client::mcp_tool_names_for_user_text(query, schemas);
if !mcp_hits.is_empty() {
let mcp_cap = mcp_hits_cap(schemas);
let deferred: HashSet<&str> = catalog.iter().map(|e| e.name.as_str()).collect();
let mut merged = Vec::new();
for name in &mcp_hits {
if (deferred.contains(name.as_str())
|| name == "mcp_list_tools"
|| name == "mcp_call_tool")
&& !merged.iter().any(|existing: &String| existing == name)
{
merged.push(name.clone());
}
}
for name in hits {
if !merged.iter().any(|existing| existing == &name) {
merged.push(name);
}
}
hits = merged;
let cap = max_prefetch.max(mcp_cap);
if hits.len() > cap {
hits.truncate(cap);
}
}
if looks_like_create_file_intent(query) {
let deferred_names: HashSet<&str> = catalog.iter().map(|e| e.name.as_str()).collect();
hits.retain(|n| n != "skill_manage");
if deferred_names.contains("write_file") && !hits.iter().any(|n| n == "write_file") {
hits.insert(0, "write_file".into());
if hits.len() > max_prefetch {
hits.truncate(max_prefetch);
}
}
}
hits
}
fn mcp_hits_cap(schemas: &[ToolSchema]) -> usize {
let mcp_dynamic = schemas
.iter()
.filter(|s| {
s.name.starts_with("mcp_") && s.name != "mcp_list_tools" && s.name != "mcp_call_tool"
})
.count();
2 + mcp_dynamic.min(8)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn schema(name: &str, desc: &str) -> ToolSchema {
ToolSchema {
name: name.into(),
description: desc.into(),
parameters: json!({
"type": "object",
"properties": { "url": { "type": "string" } }
}),
strict: None,
}
}
#[test]
fn bm25_finds_browser_tools_by_query() {
let schemas = vec![
schema("browser_navigate", "Navigate headless browser to URL"),
schema("browser_snapshot", "Capture accessibility snapshot of page"),
schema("memory_write", "Write persistent memory entry"),
];
let materialized = HashSet::new();
let catalog = build_deferred_catalog(&schemas, &materialized);
let hits = search_deferred_catalog(&catalog, "browser navigate url", 1);
assert!(hits.first().is_some_and(|n| n.contains("browser")));
}
#[test]
fn substring_fallback_matches_partial_name() {
let schemas = vec![schema(
"ha_get_states",
"Fetch Home Assistant entity states",
)];
let catalog = build_deferred_catalog(&schemas, &HashSet::new());
let hits = search_deferred_catalog(&catalog, "ha_get", 5);
assert!(hits.contains(&"ha_get_states".to_string()));
}
#[test]
fn prefetch_returns_deferred_hits() {
let schemas = vec![
schema("browser_navigate", "Navigate headless browser to URL"),
schema("memory_write", "Write persistent memory entry"),
];
let hits = prefetch_tools_for_user_message(
"please navigate the browser to a url",
&schemas,
&HashSet::new(),
3,
);
assert!(hits.iter().any(|n| n.contains("browser")));
}
#[test]
fn mcp_tool_names_for_user_text_matches_server_prefix() {
let schemas = vec![
schema("mcp_list_tools", "List MCP tools"),
schema("mcp_GPS_list_funds", "List funds"),
schema("terminal", "Shell"),
];
let hits =
crate::tools::mcp_client::mcp_tool_names_for_user_text("List Fund in GPS", &schemas);
assert!(hits.is_empty() || hits.contains(&"mcp_list_tools".to_string()));
}
#[test]
fn create_intent_excludes_skill_manage_from_prefetch() {
let schemas = vec![
schema(
"skill_manage",
"Create edit patch delete a skill or write supporting files",
),
schema("web_search", "Search the web for facts"),
schema("browser_navigate", "Navigate headless browser to URL"),
];
let hits = prefetch_tools_for_user_message(
"Write a complete html5 and javascript 3D game in ./demo/game001",
&schemas,
&HashSet::new(),
3,
);
assert!(
!hits.iter().any(|n| n == "skill_manage"),
"create intent must not prefetch skill_manage: {hits:?}"
);
}
#[test]
fn looks_like_create_file_intent_game001() {
assert!(looks_like_create_file_intent(
"Write a complete html5 and javascript 3D game in ./demo/game001"
));
assert!(!looks_like_create_file_intent("what is the weather today"));
}
}