use std::sync::Arc;
use std::time::Instant;
use crate::dense_cache::{DenseCache, Embeddable};
use crate::embedding::EmbedderError;
use crate::fusion::{RETRIEVE_DEPTH, RRF_K, rrf_fuse};
use crate::indexing::searchable_text;
use crate::method::SearchMethod;
use crate::search::bm25_search;
use crate::tool::Tool;
use crate::trace::{
ChurnKind, NoopSink, Origin, SearchHitTrace, SearchStage, TraceEvent, TraceSink,
};
pub struct SearchHit {
pub tool_id: String,
pub score: f32,
}
impl Embeddable for Tool {
fn embed_id(&self) -> &str {
&self.id
}
fn embed_text(&self) -> String {
searchable_text(self)
}
}
pub struct ToolRegistry {
tools: Vec<Tool>,
sink: Arc<dyn TraceSink>,
dense: DenseCache,
}
impl Default for ToolRegistry {
fn default() -> Self {
Self::new()
}
}
impl ToolRegistry {
pub fn new() -> Self {
Self {
tools: Vec::new(),
sink: Arc::new(NoopSink),
dense: DenseCache::new(),
}
}
pub fn with_trace_sink(sink: Arc<dyn TraceSink>) -> Self {
Self {
tools: Vec::new(),
sink,
dense: DenseCache::new(),
}
}
pub fn set_trace_sink(&mut self, sink: Arc<dyn TraceSink>) {
self.sink = sink;
}
pub fn record_event(&self, event: TraceEvent) {
self.sink.record(event);
}
pub fn register(&mut self, tool: Tool) {
let tool_id = tool.id.clone();
self.tools.push(tool);
self.sink.record(TraceEvent::IndexChurn {
kind: ChurnKind::Add,
tool_id,
});
}
pub fn search(&self, query: &str, top_k: usize) -> Vec<SearchHit> {
self.search_with_origin(query, top_k, Origin::Direct)
}
pub fn search_with_origin(&self, query: &str, top_k: usize, origin: Origin) -> Vec<SearchHit> {
self.bm25_search_traced(query, top_k, origin)
}
pub fn search_with_method(
&self,
query: &str,
top_k: usize,
origin: Origin,
method: SearchMethod,
) -> Result<Vec<SearchHit>, EmbedderError> {
match method {
SearchMethod::Bm25 => Ok(self.bm25_search_traced(query, top_k, origin)),
SearchMethod::Semantic => self.semantic_search_traced(query, top_k, origin),
SearchMethod::Hybrid => self.hybrid_search_traced(query, top_k, origin),
}
}
pub fn build_embeddings(&self) -> Result<(), EmbedderError> {
self.dense.extend(&self.tools, self.sink.as_ref())
}
fn bm25_search_traced(&self, query: &str, top_k: usize, origin: Origin) -> Vec<SearchHit> {
let started = Instant::now();
let hits: Vec<SearchHit> = bm25_search(
self.tools
.iter()
.map(|t| (t.id.clone(), searchable_text(t))),
query,
top_k,
)
.into_iter()
.map(|(tool_id, score)| SearchHit { tool_id, score })
.collect();
let took_ms = started.elapsed().as_millis() as u64;
let top_score = hits.first().map(|h| h.score as f64);
self.record_search(
query,
origin,
top_k,
&hits,
vec![SearchStage {
name: "bm25".into(),
took_ms,
top_score,
}],
took_ms,
);
hits
}
fn semantic_search_traced(
&self,
query: &str,
top_k: usize,
origin: Origin,
) -> Result<Vec<SearchHit>, EmbedderError> {
let started = Instant::now();
if self.tools.is_empty() || top_k == 0 {
self.record_search(query, origin, top_k, &[], Vec::new(), 0);
return Ok(Vec::new());
}
self.dense.require_built(self.tools.len())?;
let query_vec = self.dense.embed_query(query, self.sink.as_ref())?;
let t = Instant::now();
let ranked = self.dense.ranked(&self.tools, &query_vec, top_k);
let stage_ms = t.elapsed().as_millis() as u64;
let hits: Vec<SearchHit> = ranked
.into_iter()
.map(|(tool_id, score)| SearchHit { tool_id, score })
.collect();
let took_ms = started.elapsed().as_millis() as u64;
let top_score = hits.first().map(|h| h.score as f64);
self.record_search(
query,
origin,
top_k,
&hits,
vec![SearchStage {
name: "dense".into(),
took_ms: stage_ms,
top_score,
}],
took_ms,
);
Ok(hits)
}
fn hybrid_search_traced(
&self,
query: &str,
top_k: usize,
origin: Origin,
) -> Result<Vec<SearchHit>, EmbedderError> {
let started = Instant::now();
if self.tools.is_empty() || top_k == 0 {
self.record_search(query, origin, top_k, &[], Vec::new(), 0);
return Ok(Vec::new());
}
let depth = RETRIEVE_DEPTH.max(top_k);
let t = Instant::now();
let bm25_ranked = bm25_search(
self.tools
.iter()
.map(|t| (t.id.clone(), searchable_text(t))),
query,
depth,
);
let bm25_stage = SearchStage {
name: "bm25".into(),
took_ms: t.elapsed().as_millis() as u64,
top_score: bm25_ranked.first().map(|(_, s)| *s as f64),
};
self.dense.require_built(self.tools.len())?;
let t = Instant::now();
let query_vec = self.dense.embed_query(query, self.sink.as_ref())?;
let dense_ranked = self.dense.ranked(&self.tools, &query_vec, depth);
let dense_stage = SearchStage {
name: "dense".into(),
took_ms: t.elapsed().as_millis() as u64,
top_score: dense_ranked.first().map(|(_, s)| *s as f64),
};
let t = Instant::now();
let bm25_ids: Vec<String> = bm25_ranked.into_iter().map(|(id, _)| id).collect();
let dense_ids: Vec<String> = dense_ranked.into_iter().map(|(id, _)| id).collect();
let mut fused = rrf_fuse(&[&bm25_ids, &dense_ids], RRF_K);
fused.truncate(top_k);
let rrf_stage = SearchStage {
name: "rrf".into(),
took_ms: t.elapsed().as_millis() as u64,
top_score: fused.first().map(|(_, s)| *s as f64),
};
let hits: Vec<SearchHit> = fused
.into_iter()
.map(|(tool_id, score)| SearchHit { tool_id, score })
.collect();
let took_ms = started.elapsed().as_millis() as u64;
self.record_search(
query,
origin,
top_k,
&hits,
vec![bm25_stage, dense_stage, rrf_stage],
took_ms,
);
Ok(hits)
}
#[allow(clippy::too_many_arguments)]
fn record_search(
&self,
query: &str,
origin: Origin,
top_k: usize,
hits: &[SearchHit],
stages: Vec<SearchStage>,
took_ms: u64,
) {
self.sink.record(TraceEvent::Search {
query: query.to_string(),
origin,
top_k: top_k as u32,
hits: hits
.iter()
.map(|h| SearchHitTrace {
tool_id: h.tool_id.clone(),
score: h.score as f64,
})
.collect(),
stages,
took_ms,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embedding::Embedder;
use crate::trace::MemorySink;
struct StubEmbedder;
impl StubEmbedder {
fn vec_for(text: &str) -> Vec<f32> {
let t = text.to_lowercase();
if t.contains("read") {
vec![1.0, 0.0, 0.0]
} else if t.contains("delete") || t.contains("remove") {
vec![0.0, 1.0, 0.0]
} else {
vec![0.0, 0.0, 1.0]
}
}
}
impl Embedder for StubEmbedder {
fn embed_doc(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
Ok(StubEmbedder::vec_for(text))
}
fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
Ok(StubEmbedder::vec_for(text))
}
}
struct FailingEmbedder;
impl Embedder for FailingEmbedder {
fn embed_doc(&self, _: &str) -> Result<Vec<f32>, EmbedderError> {
Err(EmbedderError::Inference {
source: "stub failure".into(),
})
}
fn embed_query(&self, _: &str) -> Result<Vec<f32>, EmbedderError> {
Err(EmbedderError::Inference {
source: "stub failure".into(),
})
}
}
struct CountingEmbedder {
doc_calls: std::sync::atomic::AtomicUsize,
}
impl CountingEmbedder {
fn new() -> Self {
Self {
doc_calls: std::sync::atomic::AtomicUsize::new(0),
}
}
fn doc_calls(&self) -> usize {
self.doc_calls.load(std::sync::atomic::Ordering::SeqCst)
}
}
impl Embedder for CountingEmbedder {
fn embed_doc(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
self.doc_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(StubEmbedder::vec_for(text))
}
fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
Ok(StubEmbedder::vec_for(text))
}
}
fn with_embedder(embedder: Arc<dyn Embedder>) -> ToolRegistry {
ToolRegistry {
tools: Vec::new(),
sink: Arc::new(NoopSink),
dense: DenseCache::with_embedder(embedder),
}
}
fn tool(id: &str, description: &str) -> Tool {
Tool {
id: id.into(),
name: id.into(),
description: description.into(),
input_schema: serde_json::json!({}),
output_schema: serde_json::json!({}),
}
}
fn catalog(embedder: Arc<dyn Embedder>) -> ToolRegistry {
let mut reg = with_embedder(embedder);
reg.register(tool("read_file", "read a file"));
reg.register(tool("delete_file", "delete a file"));
reg
}
#[test]
fn default_search_is_bm25_and_infallible() {
let mut reg = ToolRegistry::new();
reg.register(tool("read_file", "read the contents of a file"));
reg.register(tool("delete_file", "delete a file"));
let hits = reg.search("read a file", 5);
assert_eq!(hits.first().map(|h| h.tool_id.as_str()), Some("read_file"));
}
#[test]
fn bm25_never_loads_the_model() {
let reg = catalog(Arc::new(FailingEmbedder));
let hits = reg
.search_with_method("read", 5, Origin::Direct, SearchMethod::Bm25)
.expect("bm25 is infallible");
assert_eq!(hits.first().map(|h| h.tool_id.as_str()), Some("read_file"));
}
#[test]
fn semantic_ranks_via_injected_embedder() {
let reg = catalog(Arc::new(StubEmbedder));
reg.build_embeddings().unwrap();
let hits = reg
.search_with_method("read something", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.first().map(|h| h.tool_id.as_str()), Some("read_file"));
}
#[test]
fn semantic_without_embeddings_errors() {
let reg = catalog(Arc::new(StubEmbedder));
assert!(matches!(
reg.search_with_method("read", 5, Origin::Direct, SearchMethod::Semantic),
Err(EmbedderError::EmbeddingsNotBuilt)
));
}
#[test]
fn build_embeddings_surfaces_embedder_error_instead_of_panicking() {
let reg = catalog(Arc::new(FailingEmbedder));
assert!(matches!(
reg.build_embeddings(),
Err(EmbedderError::Inference { .. })
));
}
#[test]
fn hybrid_fuses_bm25_and_dense() {
let reg = catalog(Arc::new(StubEmbedder));
reg.build_embeddings().unwrap();
let hits = reg
.search_with_method("read a file", 5, Origin::Direct, SearchMethod::Hybrid)
.unwrap();
assert_eq!(hits.first().map(|h| h.tool_id.as_str()), Some("read_file"));
}
#[test]
fn hybrid_recalls_a_tool_bm25_alone_misses() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(tool("records_mgr", "manage old records archive")); reg.register(tool("deleter", "delete entries")); reg.build_embeddings().unwrap();
let q = "remove old records";
let bm25 = reg
.search_with_method(q, 5, Origin::Direct, SearchMethod::Bm25)
.unwrap();
let semantic = reg
.search_with_method(q, 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
let hybrid = reg
.search_with_method(q, 5, Origin::Direct, SearchMethod::Hybrid)
.unwrap();
assert_eq!(
bm25.first().map(|h| h.tool_id.as_str()),
Some("records_mgr")
);
assert_eq!(
semantic.first().map(|h| h.tool_id.as_str()),
Some("deleter")
);
assert!(!bm25.iter().any(|h| h.tool_id == "deleter"));
let ids: Vec<&str> = hybrid.iter().map(|h| h.tool_id.as_str()).collect();
assert!(
ids.contains(&"records_mgr") && ids.contains(&"deleter"),
"hybrid should fuse both arms, got {ids:?}"
);
}
#[test]
fn semantic_stage_is_named_dense() {
let sink = Arc::new(MemorySink::new("s"));
let mut reg = catalog(Arc::new(StubEmbedder));
reg.set_trace_sink(sink.clone());
reg.build_embeddings().unwrap();
reg.search_with_method("read", 5, Origin::Agent, SearchMethod::Semantic)
.unwrap();
let events = sink.drain();
assert!(events.iter().any(|e| matches!(
&e.event,
TraceEvent::Search { stages, .. } if stages.iter().any(|s| s.name == "dense")
)));
}
#[test]
fn hybrid_emits_three_stages() {
let sink = Arc::new(MemorySink::new("s"));
let mut reg = catalog(Arc::new(StubEmbedder));
reg.set_trace_sink(sink.clone());
reg.build_embeddings().unwrap();
reg.search_with_method("read", 5, Origin::Agent, SearchMethod::Hybrid)
.unwrap();
let events = sink.drain();
assert!(events.iter().any(|e| matches!(
&e.event,
TraceEvent::Search { stages, .. }
if stages.iter().any(|s| s.name == "bm25")
&& stages.iter().any(|s| s.name == "dense")
&& stages.iter().any(|s| s.name == "rrf")
)));
}
#[test]
fn build_embeddings_after_register_embeds_only_the_new_tool() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(tool("read_file", "read a file"));
reg.register(tool("delete_file", "delete a file"));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 2);
reg.register(tool("reader_v2", "read a file too"));
reg.build_embeddings().unwrap();
assert_eq!(
counter.doc_calls(),
3,
"only the newly-registered tool should be embedded"
);
let hits = reg
.search_with_method("read", 10, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert!(hits.iter().any(|h| h.tool_id == "reader_v2"));
}
#[test]
fn build_embeddings_precomputes_so_search_embeds_no_docs() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(tool("read_file", "read a file"));
reg.register(tool("delete_file", "delete a file"));
reg.build_embeddings().unwrap();
assert_eq!(
counter.doc_calls(),
2,
"build_embeddings embeds the corpus up front"
);
reg.search_with_method("read", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(
counter.doc_calls(),
2,
"a search after build_embeddings embeds only the query, no documents"
);
}
#[test]
fn build_embeddings_is_idempotent() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(tool("read_file", "read a file"));
reg.build_embeddings().unwrap();
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 1);
}
#[test]
fn re_register_updates_the_ranked_vector() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(tool("t", "read a file")); reg.build_embeddings().unwrap();
reg.register(tool("t", "delete a file")); reg.build_embeddings().unwrap();
let hits = reg
.search_with_method("delete", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.first().map(|h| h.tool_id.as_str()), Some("t"));
assert!(hits[0].score > 0.9, "ranks with the re-registered vector");
}
#[test]
fn empty_registry_semantic_returns_no_hits_without_loading() {
let reg = with_embedder(Arc::new(FailingEmbedder));
let hits = reg
.search_with_method("anything", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert!(hits.is_empty());
}
#[test]
fn register_and_search_emit_trace_events() {
let sink = Arc::new(MemorySink::new("test-session"));
let mut reg = ToolRegistry::with_trace_sink(sink.clone());
reg.register(tool("read_file", "read a file"));
reg.search_with_origin("read", 5, Origin::Agent);
let events = sink.drain();
assert!(events.iter().any(|e| matches!(
e.event,
TraceEvent::IndexChurn {
kind: ChurnKind::Add,
..
}
)));
assert!(events.iter().any(|e| matches!(
&e.event,
TraceEvent::Search { origin: Origin::Agent, hits, .. } if !hits.is_empty()
)));
}
}