use std::sync::{Arc, PoisonError, RwLock};
use std::time::Instant;
use indexmap::IndexMap;
use crate::artifact_warm::{ArtifactWarmError, OnArtifactMiss};
use crate::dense_cache::{DenseCache, Embeddable};
use crate::embedding::EmbedderError;
use crate::embedding_artifact::{ArtifactEntryKind, ArtifactError};
use crate::embedding_config::EmbeddingModel;
use crate::fusion::{RETRIEVE_DEPTH, RRF_K, WeightedArm, rrf_fuse_weighted};
use crate::method::SearchMethod;
use crate::search::Bm25Cache;
use crate::skill::Skill;
use crate::skill_indexing::searchable_text;
use crate::tool_registry::AdaptiveRankingStatus;
use crate::trace::{
ChurnKind, NoopSink, Origin, SearchStage, SkillHitTrace, TraceEvent, TraceEventContext,
TraceSink,
};
use crate::usage::{ArmOutcome, Capability, IntentGraph, UsageArm};
pub struct SkillHit {
pub skill_id: String,
pub score: f32,
pub rank: u32,
pub fused: bool,
}
fn to_skill_hits(ranked: Vec<(String, f32)>, fused: bool) -> Vec<SkillHit> {
ranked
.into_iter()
.enumerate()
.map(|(i, (skill_id, score))| SkillHit {
skill_id,
score,
rank: i as u32,
fused,
})
.collect()
}
impl Embeddable for Skill {
fn embed_id(&self) -> &str {
&self.id
}
fn embed_text(&self) -> String {
searchable_text(self)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct ReplaceOutcome {
pub added: usize,
pub removed: usize,
pub updated: usize,
pub unchanged: usize,
}
pub struct SkillRegistry {
skills: IndexMap<String, Skill>,
sink: Arc<dyn TraceSink>,
experimental_catalog_definitions: bool,
bm25: Bm25Cache,
dense: DenseCache,
graph: Option<Arc<RwLock<IntentGraph>>>,
}
impl Default for SkillRegistry {
fn default() -> Self {
Self::new()
}
}
impl SkillRegistry {
pub fn new() -> Self {
Self {
skills: IndexMap::new(),
sink: Arc::new(NoopSink),
experimental_catalog_definitions: false,
bm25: Bm25Cache::new(),
dense: DenseCache::new(),
graph: None,
}
}
pub fn with_trace_sink(sink: Arc<dyn TraceSink>) -> Self {
Self {
skills: IndexMap::new(),
sink,
experimental_catalog_definitions: false,
bm25: Bm25Cache::new(),
dense: DenseCache::new(),
graph: None,
}
}
pub fn with_embedding(model: EmbeddingModel) -> Self {
Self {
skills: IndexMap::new(),
sink: Arc::new(NoopSink),
experimental_catalog_definitions: false,
bm25: Bm25Cache::new(),
dense: DenseCache::with_model(model),
graph: None,
}
}
pub fn set_trace_sink(&mut self, sink: Arc<dyn TraceSink>) {
self.sink = sink;
}
pub fn experimental_enable_catalog_definitions(&mut self) {
self.experimental_catalog_definitions = true;
}
pub fn record_event(&self, event: TraceEvent) {
self.sink.record(event);
}
pub fn record_event_with_context(&self, event: TraceEvent, context: TraceEventContext) {
self.sink.record_with_context(event, context);
}
pub fn set_intent_graph(&mut self, graph: Option<Arc<RwLock<IntentGraph>>>) {
self.graph = graph;
}
pub fn adaptive_ranking_status(&self) -> AdaptiveRankingStatus {
let Some(graph) = self.graph.as_ref() else {
return AdaptiveRankingStatus::Inactive;
};
let Ok(g) = graph.read() else {
return AdaptiveRankingStatus::Unknown;
};
if !g.intents.iter().any(|i| i.centroid.is_some()) {
return AdaptiveRankingStatus::Active;
}
let Some(active_fp) = self.dense.built_fingerprint() else {
return AdaptiveRankingStatus::Active;
};
let active_dim = self.dense.dim().unwrap_or(0);
match g.model_status(&active_fp, active_dim).describe() {
None => AdaptiveRankingStatus::Active,
Some((built, active, dim_mismatch)) => AdaptiveRankingStatus::Paused {
dim_mismatch,
built,
active,
},
}
}
pub fn rebuild_intent_graph(&self) -> Result<(), EmbedderError> {
let Some(graph) = self.graph.as_ref() else {
return Ok(());
};
let members: Vec<(String, Vec<String>)> = {
let g = graph.read().unwrap_or_else(PoisonError::into_inner);
g.intents
.iter()
.map(|i| (i.id.clone(), i.members.clone()))
.collect()
};
let mut per_cluster = Vec::with_capacity(members.len());
let mut fingerprint = None;
for (id, cluster_members) in &members {
let (vectors, fp) = self
.dense
.embed_texts_with_identity(cluster_members, self.sink.as_ref())?;
if !cluster_members.is_empty() {
fingerprint = Some(fp);
}
per_cluster.push((id.clone(), vectors));
}
if let Some(fp) = fingerprint {
let mut g = graph.write().unwrap_or_else(PoisonError::into_inner);
g.rebuild_centroids(per_cluster, fp);
}
Ok(())
}
fn usage_arm(&self, query: &str, query_vec: Option<&[f32]>) -> Option<UsageArm> {
let graph = self.graph.as_ref()?;
let fingerprint = self.dense.built_fingerprint();
let (outcome, mismatch) = {
let guard = graph.read().ok()?;
let mismatch = match (query_vec, &fingerprint) {
(Some(v), Some(fp)) => guard.model_status(fp, v.len()).describe(),
_ => None,
};
if mismatch.is_some() {
(ArmOutcome::NoMatch, mismatch)
} else {
if let (Some(v), Some(fp)) = (query_vec, &fingerprint) {
guard.note_query_vector(query, v, fp);
}
let known = |id: &str| self.skills.contains_key(id);
(guard.arm(query, query_vec, Capability::Skill, &known), None)
}
};
if let Some((built, active, dim_mismatch)) = mismatch {
self.sink.record(TraceEvent::UsageModelMismatch {
built,
active,
dim_mismatch,
});
}
let (intent, similarity, support, promoted, dropped) = outcome.describe();
self.sink.record(TraceEvent::UsageBoost {
intent,
similarity,
support,
promoted,
dropped,
});
outcome.into_arm()
}
fn bm25_docs(&self) -> impl Iterator<Item = (String, String)> + '_ {
self.skills
.values()
.map(|s| (s.id.clone(), searchable_text(s)))
}
fn bm25_index(&self) -> Arc<crate::search::Bm25Index> {
self.bm25.get_or_build(|| self.bm25_docs())
}
fn fuse_arms(arms: &[WeightedArm<'_>], top_k: usize) -> (Vec<SkillHit>, SearchStage) {
let t = Instant::now();
let mut fused = rrf_fuse_weighted(arms, RRF_K);
fused.truncate(top_k);
let stage = SearchStage {
name: "rrf".into(),
took_ms: t.elapsed().as_millis() as u64,
top_score: fused.first().map(|(_, s)| *s as f64),
};
let hits = to_skill_hits(fused, true);
(hits, stage)
}
fn usage_stage(arm: &UsageArm, took_ms: u64) -> SearchStage {
SearchStage {
name: "usage".into(),
took_ms,
top_score: Some(arm.weight() as f64),
}
}
pub fn register(&mut self, skill: Skill) {
let skill_id = skill.id.clone();
let definition = self
.experimental_catalog_definitions
.then(|| TraceEvent::catalog_definition_for_skill(&skill))
.flatten();
let definition_changed = definition.as_ref().is_some_and(|definition| {
self.skills.get(&skill_id).is_none_or(|existing| {
let existing_definition = TraceEvent::catalog_definition_for_skill(existing);
existing_definition
.as_ref()
.and_then(TraceEvent::catalog_definition_hash)
!= definition.catalog_definition_hash()
})
});
self.bm25.invalidate();
if self.skills.insert(skill_id.clone(), skill).is_some() {
self.dense.invalidate(&skill_id);
}
self.sink.record(TraceEvent::SkillChurn {
kind: ChurnKind::Add,
skill_id,
});
if definition_changed && let Some(definition) = definition {
self.sink.record(definition);
}
}
pub fn replace_all(&mut self, skills: Vec<Skill>) -> ReplaceOutcome {
let mut next: IndexMap<String, Skill> = IndexMap::with_capacity(skills.len());
for skill in skills {
next.insert(skill.id.clone(), skill);
}
let mut outcome = ReplaceOutcome::default();
let mut indexed_text_changed = false;
for id in self.skills.keys() {
if !next.contains_key(id) {
self.dense.invalidate(id);
indexed_text_changed = true;
self.sink.record(TraceEvent::SkillChurn {
kind: ChurnKind::Remove,
skill_id: id.clone(),
});
outcome.removed += 1;
}
}
for (id, skill) in &next {
let definition = self
.experimental_catalog_definitions
.then(|| TraceEvent::catalog_definition_for_skill(skill))
.flatten();
let definition_changed = definition.as_ref().is_some_and(|definition| {
self.skills.get(id).is_none_or(|existing| {
let existing_definition = TraceEvent::catalog_definition_for_skill(existing);
existing_definition
.as_ref()
.and_then(TraceEvent::catalog_definition_hash)
!= definition.catalog_definition_hash()
})
});
match self.skills.get(id) {
Some(current) if current == skill => {
outcome.unchanged += 1;
continue;
}
Some(current) => {
if searchable_text(current) != searchable_text(skill) {
self.dense.invalidate(id);
indexed_text_changed = true;
}
outcome.updated += 1;
}
None => {
indexed_text_changed = true;
outcome.added += 1;
}
}
self.sink.record(TraceEvent::SkillChurn {
kind: ChurnKind::Add,
skill_id: id.clone(),
});
if definition_changed && let Some(definition) = definition {
self.sink.record(definition);
}
}
if indexed_text_changed {
self.bm25.invalidate();
}
self.skills = next;
outcome
}
pub fn len(&self) -> usize {
self.skills.len()
}
pub fn is_empty(&self) -> bool {
self.skills.is_empty()
}
pub fn search(&self, query: &str, top_k: usize) -> Vec<SkillHit> {
self.search_with_origin(query, top_k, Origin::Direct)
}
pub fn search_with_origin(&self, query: &str, top_k: usize, origin: Origin) -> Vec<SkillHit> {
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<SkillHit>, EmbedderError> {
self.search_with_method_and_context(
query,
top_k,
origin,
method,
TraceEventContext::default(),
)
}
pub fn search_with_method_and_context(
&self,
query: &str,
top_k: usize,
origin: Origin,
method: SearchMethod,
context: TraceEventContext,
) -> Result<Vec<SkillHit>, EmbedderError> {
match method {
SearchMethod::Bm25 => {
Ok(self.bm25_search_traced_with_context(query, top_k, origin, context))
}
SearchMethod::Semantic => self.semantic_search_traced(query, top_k, origin, context),
SearchMethod::Hybrid => self.hybrid_search_traced(query, top_k, origin, context),
}
}
pub fn build_embeddings(&self) -> Result<(), EmbedderError> {
self.dense.extend(self.skills.values(), self.sink.as_ref())
}
pub fn rebuild_embeddings(&self) -> Result<(), EmbedderError> {
self.dense.rebuild(self.skills.values(), self.sink.as_ref())
}
pub fn warm_embeddings_from_artifact(
&self,
bytes: &[u8],
on_miss: OnArtifactMiss,
) -> Result<(), ArtifactWarmError> {
match on_miss {
OnArtifactMiss::Error => {
let outcome = self.dense.warm_from_artifact(
bytes,
ArtifactEntryKind::Skill,
self.skills.values(),
self.sink.as_ref(),
)?;
if outcome.missing.is_empty() {
Ok(())
} else {
Err(ArtifactWarmError::Incomplete {
missing: outcome.missing,
})
}
}
OnArtifactMiss::Embed => self.dense.with_operation_write(|cache| {
let outcome = cache.warm_from_artifact_locked(
bytes,
ArtifactEntryKind::Skill,
self.skills.values(),
self.sink.as_ref(),
)?;
if outcome.missing.is_empty() {
return Ok(());
}
cache
.extend_locked(self.skills.values(), self.sink.as_ref())
.map_err(ArtifactWarmError::from)
}),
}
}
pub fn build_embedding_artifact(&self) -> Result<Vec<u8>, ArtifactError> {
self.dense.build_artifact(
ArtifactEntryKind::Skill,
self.skills.values(),
self.sink.as_ref(),
)
}
fn bm25_search_traced(&self, query: &str, top_k: usize, origin: Origin) -> Vec<SkillHit> {
self.bm25_search_traced_with_context(query, top_k, origin, TraceEventContext::default())
}
fn bm25_search_traced_with_context(
&self,
query: &str,
top_k: usize,
origin: Origin,
context: TraceEventContext,
) -> Vec<SkillHit> {
let started = Instant::now();
let t = Instant::now();
let arm = self.usage_arm(query, None);
let usage_ms = t.elapsed().as_millis() as u64;
let Some(arm) = arm else {
let hits = to_skill_hits(self.bm25_index().search(query, top_k), false);
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,
context,
);
return hits;
};
let depth = RETRIEVE_DEPTH.max(top_k);
let t = Instant::now();
let bm25_ranked = self.bm25_index().search(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),
};
let bm25_ids: Vec<String> = bm25_ranked.into_iter().map(|(id, _)| id).collect();
let (hits, rrf_stage) =
Self::fuse_arms(&[(&bm25_ids, 1.0), (&arm.ids, arm.weight())], top_k);
let took_ms = started.elapsed().as_millis() as u64;
self.record_search(
query,
origin,
top_k,
&hits,
vec![bm25_stage, Self::usage_stage(&arm, usage_ms), rrf_stage],
took_ms,
context,
);
hits
}
fn semantic_search_traced(
&self,
query: &str,
top_k: usize,
origin: Origin,
context: TraceEventContext,
) -> Result<Vec<SkillHit>, EmbedderError> {
let started = Instant::now();
if self.skills.is_empty() || top_k == 0 {
self.record_search(query, origin, top_k, &[], Vec::new(), 0, context);
return Ok(Vec::new());
}
let depth = if self.graph.is_some() {
RETRIEVE_DEPTH.max(top_k)
} else {
top_k
};
let t = Instant::now();
let (ranked, query_vec) = self.dense.search_returning_query_vec(
self.skills.values(),
query,
depth,
self.sink.as_ref(),
)?;
let stage_ms = t.elapsed().as_millis() as u64;
let t = Instant::now();
let arm = self.usage_arm(query, Some(&query_vec));
let usage_ms = t.elapsed().as_millis() as u64;
let Some(arm) = arm else {
let mut hits = to_skill_hits(ranked, false);
hits.truncate(top_k);
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,
context,
);
return Ok(hits);
};
let dense_stage = SearchStage {
name: "dense".into(),
took_ms: stage_ms,
top_score: ranked.first().map(|(_, s)| *s as f64),
};
let dense_ids: Vec<String> = ranked.into_iter().map(|(id, _)| id).collect();
let (hits, rrf_stage) =
Self::fuse_arms(&[(&dense_ids, 1.0), (&arm.ids, arm.weight())], top_k);
let took_ms = started.elapsed().as_millis() as u64;
self.record_search(
query,
origin,
top_k,
&hits,
vec![dense_stage, Self::usage_stage(&arm, usage_ms), rrf_stage],
took_ms,
context,
);
Ok(hits)
}
fn hybrid_search_traced(
&self,
query: &str,
top_k: usize,
origin: Origin,
context: TraceEventContext,
) -> Result<Vec<SkillHit>, EmbedderError> {
let started = Instant::now();
if self.skills.is_empty() || top_k == 0 {
self.record_search(query, origin, top_k, &[], Vec::new(), 0, context);
return Ok(Vec::new());
}
let depth = RETRIEVE_DEPTH.max(top_k);
let t = Instant::now();
let bm25_ranked = self.bm25_index().search(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),
};
let t = Instant::now();
let (dense_ranked, query_vec) = self.dense.search_returning_query_vec(
self.skills.values(),
query,
depth,
self.sink.as_ref(),
)?;
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 arm = self.usage_arm(query, Some(&query_vec));
let usage_ms = t.elapsed().as_millis() as u64;
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 arms: Vec<WeightedArm<'_>> = vec![(&bm25_ids, 1.0), (&dense_ids, 1.0)];
if let Some(arm) = &arm {
arms.push((&arm.ids, arm.weight()));
}
let (hits, rrf_stage) = Self::fuse_arms(&arms, top_k);
let mut stages = vec![bm25_stage, dense_stage];
if let Some(arm) = &arm {
stages.push(Self::usage_stage(arm, usage_ms));
}
stages.push(rrf_stage);
let took_ms = started.elapsed().as_millis() as u64;
self.record_search(query, origin, top_k, &hits, stages, took_ms, context);
Ok(hits)
}
#[allow(clippy::too_many_arguments)]
fn record_search(
&self,
query: &str,
origin: Origin,
top_k: usize,
hits: &[SkillHit],
stages: Vec<SearchStage>,
took_ms: u64,
context: TraceEventContext,
) {
self.sink.record_with_context(
TraceEvent::SkillSearch {
query: query.to_string(),
origin,
top_k: top_k as u32,
hits: hits
.iter()
.map(|h| SkillHitTrace {
skill_id: h.skill_id.clone(),
score: h.score as f64,
})
.collect(),
stages,
took_ms,
},
context,
);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embedding::Embedder;
use crate::test_support::{
FailOnEmbedStub, FpCountingEmbedder, PanicOnEmbedStub, build_test_artifact, unit,
};
use crate::trace::MemorySink;
struct StubEmbedder;
impl StubEmbedder {
fn vec_for(text: &str) -> Vec<f32> {
let t = text.to_lowercase();
if t.contains("api") || t.contains("rest") {
vec![1.0, 0.0, 0.0]
} else if t.contains("frontend") || t.contains("slides") {
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 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>) -> SkillRegistry {
SkillRegistry {
skills: IndexMap::new(),
sink: Arc::new(NoopSink),
experimental_catalog_definitions: false,
bm25: Bm25Cache::new(),
dense: DenseCache::with_embedder(embedder),
graph: None,
}
}
fn skill(id: &str, name: &str, description: &str, tags: &[&str]) -> Skill {
Skill {
id: id.into(),
name: name.into(),
description: description.into(),
experimental_searchable_description: None,
tags: tags.iter().map(|t| (*t).into()).collect(),
tools: vec![],
metadata: std::collections::HashMap::new(),
body: format!("# {name}\n\nbody"),
}
}
fn catalog() -> SkillRegistry {
let mut reg = SkillRegistry::new();
reg.register(skill(
"frontend-slides",
"frontend-slides",
"Build animation-rich HTML presentations from scratch",
&["frontend", "presentations"],
));
reg.register(skill(
"api-design",
"api-design",
"REST API design patterns: resource naming, status codes, pagination",
&["backend", "api"],
));
reg
}
#[test]
fn mutation_after_a_warmed_search_is_visible_in_the_next_search() {
let mut reg = catalog();
for _ in 0..3 {
let _ = reg.search("REST API design", 5);
}
reg.register(skill(
"migrations",
"migrations",
"Write reversible database migrations",
&["backend"],
));
assert_eq!(
reg.search("reversible database migrations", 5)[0].skill_id,
"migrations",
"a skill registered after searches must rank immediately"
);
let _ = reg.search("presentations", 5); reg.replace_all(vec![
skill(
"api-design",
"api-design",
"GraphQL schema federation",
&["backend"],
),
skill(
"migrations",
"migrations",
"Write reversible database migrations",
&["backend"],
),
]);
assert!(
reg.search("animation-rich HTML presentations", 5)
.is_empty(),
"a skill removed by replace_all must stop matching immediately"
);
assert_eq!(
reg.search("GraphQL schema federation", 5)[0].skill_id,
"api-design",
"content rewritten by replace_all must match immediately"
);
}
fn no_build() -> Vec<(String, String)> {
unreachable!("cache should already be populated by the search path")
}
fn assert_rebuilds(reg: &SkillRegistry) -> Arc<crate::search::Bm25Index> {
let builds = std::cell::Cell::new(0);
let index = reg.bm25.get_or_build(|| {
builds.set(builds.get() + 1);
reg.bm25_docs()
});
assert_eq!(builds.get(), 1, "mutation must drop the cached index");
index
}
#[test]
fn bm25_cache_is_warmed_by_search_and_dropped_by_every_mutator() {
let mut reg = catalog();
let _ = reg.search("REST API design", 5);
let warmed = reg.bm25.get_or_build(no_build);
let _ = reg.search("presentations", 5);
let reused = reg.bm25.get_or_build(no_build);
assert!(
Arc::ptr_eq(&warmed, &reused),
"searches between mutations must reuse one index"
);
reg.register(skill("extra", "extra", "an extra skill", &[]));
let after_register = assert_rebuilds(®);
assert!(!Arc::ptr_eq(&warmed, &after_register));
reg.replace_all(vec![skill("only", "only", "the only skill", &[])]);
let after_replace = assert_rebuilds(®);
assert!(!Arc::ptr_eq(&after_register, &after_replace));
}
#[test]
fn an_unchanged_replace_all_keeps_the_cached_bm25_index() {
let reload = || {
vec![
skill(
"frontend-slides",
"frontend-slides",
"Build animation-rich HTML presentations from scratch",
&["frontend", "presentations"],
),
skill(
"api-design",
"api-design",
"REST API design patterns: resource naming, status codes, pagination",
&["backend", "api"],
),
]
};
let mut reg = catalog();
let _ = reg.search("REST API design", 5);
let warmed = reg.bm25.get_or_build(no_build);
reg.replace_all(reload());
let after_noop = reg.bm25.get_or_build(no_build);
assert!(
Arc::ptr_eq(&warmed, &after_noop),
"an unchanged reload must keep the cached index"
);
let mut body_edit = reload();
body_edit[0].body = "rewritten body — not part of searchable_text".into();
reg.replace_all(body_edit);
let after_body_edit = reg.bm25.get_or_build(no_build);
assert!(
Arc::ptr_eq(&warmed, &after_body_edit),
"a body-only edit is not indexed and must keep the cached index"
);
}
#[test]
fn semantic_search_truncates_to_top_k_when_the_graph_matches_no_cluster() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(skill("api_a", "api_a", "rest api design", &[]));
reg.register(skill("api_b", "api_b", "rest api pagination", &[]));
reg.register(skill("api_c", "api_c", "rest api auth", &[]));
reg.register(skill("api_d", "api_d", "rest api versioning", &[]));
reg.build_embeddings().unwrap();
reg.set_intent_graph(Some(Arc::new(RwLock::new(IntentGraph::empty()))));
let hits = reg
.search_with_method("api", 2, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.len(), 2, "no-match dense path must honor top_k");
}
#[test]
fn skill_hits_carry_rank_and_unfused_scores_without_a_graph() {
let mut reg = SkillRegistry::new();
reg.register(skill(
"design-api",
"design-api",
"design a REST endpoint",
&[],
));
reg.register(skill(
"html-slides",
"html-slides",
"build html slide decks",
&[],
));
let hits = reg.search("design a REST endpoint", 5);
for (i, h) in hits.iter().enumerate() {
assert_eq!(h.rank, i as u32);
assert!(!h.fused, "no graph → not fused");
}
}
#[test]
fn search_ranks_the_relevant_skill_first() {
let reg = catalog();
let hits = reg.search("design a REST endpoint with pagination", 5);
assert_eq!(
hits.first().map(|h| h.skill_id.as_str()),
Some("api-design")
);
}
#[test]
fn experimental_searchable_description_replaces_skill_description_but_keeps_name_and_tags() {
let mut reg = SkillRegistry::new();
let mut overridden = skill(
"billing",
"billing_helper",
"orchestrate zeppelin manifests",
&["finance_ops"],
);
overridden.experimental_searchable_description = Some("reconcile overdue invoices".into());
reg.register(overridden);
assert_eq!(reg.search("overdue invoices", 5)[0].skill_id, "billing");
assert!(reg.search("zeppelin manifests", 5).is_empty());
assert_eq!(reg.search("billing", 5)[0].skill_id, "billing");
assert_eq!(reg.search("finance ops", 5)[0].skill_id, "billing");
}
#[test]
fn search_on_empty_registry_returns_no_hits() {
let reg = SkillRegistry::new();
assert!(reg.search("anything", 5).is_empty());
}
#[test]
fn re_register_replaces_not_appends() {
let mut reg = SkillRegistry::new();
reg.register(skill("s", "s", "REST API design", &["api"]));
reg.register(skill("s", "s", "HTML slides frontend", &["frontend"]));
assert_eq!(reg.len(), 1, "re-register replaces, not appends");
let hits = reg.search("html slides frontend", 5);
assert_eq!(hits.first().map(|h| h.skill_id.as_str()), Some("s"));
assert_eq!(hits.len(), 1, "one id in the corpus yields at most one hit");
}
#[test]
fn re_register_updates_the_ranked_vector() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(skill("s", "s", "REST API design", &["api"])); reg.build_embeddings().unwrap();
reg.register(skill("s", "s", "HTML slides frontend", &["frontend"])); reg.build_embeddings().unwrap();
let hits = reg
.search_with_method("frontend slides", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.first().map(|h| h.skill_id.as_str()), Some("s"));
assert!(
hits[0].score > 0.9,
"ranks with the re-embedded frontend vector"
);
}
#[test]
fn replace_all_drops_ids_absent_from_the_batch() {
let mut reg = catalog();
let outcome = reg.replace_all(vec![skill(
"api-design",
"api-design",
"REST API design patterns: resource naming, status codes, pagination",
&["backend", "api"],
)]);
assert_eq!(reg.len(), 1);
assert_eq!(outcome.removed, 1);
assert!(
reg.search("animation-rich HTML presentations", 5)
.is_empty(),
"a dropped skill must leave the corpus, not linger in the index"
);
}
#[test]
fn replace_all_with_an_empty_batch_clears_the_corpus() {
let mut reg = catalog();
reg.replace_all(Vec::new());
assert!(reg.is_empty());
assert!(reg.search("anything", 5).is_empty());
}
#[test]
fn replace_all_keeps_the_last_of_duplicate_ids() {
let mut reg = SkillRegistry::new();
reg.replace_all(vec![
skill("s", "s", "REST API design", &["api"]),
skill("s", "s", "HTML slides frontend", &["frontend"]),
]);
assert_eq!(reg.len(), 1);
let hits = reg.search("html slides frontend", 5);
assert_eq!(hits.first().map(|h| h.skill_id.as_str()), Some("s"));
}
#[test]
fn replace_all_reports_what_changed() {
let mut reg = SkillRegistry::new();
reg.register(skill("keep", "keep", "REST API design", &["api"]));
reg.register(skill("edit", "edit", "HTML slides", &["frontend"]));
reg.register(skill("drop", "drop", "database migrations", &["data"]));
let outcome = reg.replace_all(vec![
skill("keep", "keep", "REST API design", &["api"]),
skill("edit", "edit", "HTML slides and animations", &["frontend"]),
skill("add", "add", "queue consumers", &["backend"]),
]);
assert_eq!(
outcome,
ReplaceOutcome {
added: 1,
removed: 1,
updated: 1,
unchanged: 1,
}
);
}
#[test]
fn replace_all_embeds_only_what_changed() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(skill("a", "a", "REST API design", &["api"]));
reg.register(skill("b", "b", "HTML slides", &["frontend"]));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 2);
reg.replace_all(vec![
skill("a", "a", "REST API design", &["api"]),
skill("c", "c", "database migrations", &["data"]),
]);
reg.build_embeddings().unwrap();
assert_eq!(
counter.doc_calls(),
3,
"an unchanged id keeps its vector; only the new skill is embedded"
);
}
#[test]
fn replace_all_keeps_the_vector_when_only_the_body_changed() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(skill("a", "a", "REST API design", &["api"]));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 1);
let mut rewritten = skill("a", "a", "REST API design", &["api"]);
rewritten.body = "# a\n\nrewritten instructions".into();
reg.replace_all(vec![rewritten]);
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 1);
}
#[test]
fn replace_all_re_embeds_a_skill_whose_indexed_text_changed() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(skill("s", "s", "REST API design", &["api"])); reg.build_embeddings().unwrap();
reg.replace_all(vec![skill("s", "s", "HTML slides frontend", &["frontend"])]); reg.build_embeddings().unwrap();
let hits = reg
.search_with_method("frontend slides", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.first().map(|h| h.skill_id.as_str()), Some("s"));
assert!(hits[0].score > 0.9, "ranks with the re-embedded vector");
}
#[test]
fn replace_all_re_embeds_only_the_experimental_searchable_description_edit() {
let counter = Arc::new(CountingEmbedder::new());
let sink = Arc::new(MemorySink::new("test-session"));
let mut reg = with_embedder(counter.clone());
reg.set_trace_sink(sink.clone());
reg.register(skill("keep", "keep", "REST API design", &["api"]));
reg.register(skill("edit", "edit", "REST API design", &["api"]));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 2);
sink.drain();
let mut edited = skill("edit", "edit", "REST API design", &["api"]);
edited.experimental_searchable_description = Some("HTML slides frontend".into());
let outcome = reg.replace_all(vec![
skill("keep", "keep", "REST API design", &["api"]),
edited,
]);
assert_eq!(
outcome,
ReplaceOutcome {
added: 0,
removed: 0,
updated: 1,
unchanged: 1,
}
);
let churn: Vec<(ChurnKind, String)> = sink
.drain()
.into_iter()
.filter_map(|envelope| match envelope.event {
TraceEvent::SkillChurn { kind, skill_id } => Some((kind, skill_id)),
_ => None,
})
.collect();
assert_eq!(churn, vec![(ChurnKind::Add, "edit".to_string())]);
reg.build_embeddings().unwrap();
assert_eq!(
counter.doc_calls(),
3,
"only the override-edited skill is re-embedded"
);
}
#[test]
fn replace_all_drops_the_vector_of_a_removed_id() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.register(skill("frontend", "frontend", "HTML slides", &["frontend"]));
reg.build_embeddings().unwrap();
reg.replace_all(vec![
skill("api-design", "api-design", "REST API design", &["api"]),
skill("db", "db", "database migrations", &["data"]),
]);
assert!(
matches!(
reg.search_with_method("database", 5, Origin::Direct, SearchMethod::Semantic),
Err(EmbedderError::EmbeddingsNotBuilt)
),
"a removed id's stale vector must not mask an unembedded new id"
);
reg.build_embeddings().unwrap();
let hits = reg
.search_with_method(
"database migrations",
5,
Origin::Direct,
SearchMethod::Semantic,
)
.unwrap();
assert_eq!(hits.first().map(|h| h.skill_id.as_str()), Some("db"));
}
#[test]
fn replace_all_emits_churn_only_for_real_changes() {
let sink = Arc::new(MemorySink::new("test-session"));
let mut reg = SkillRegistry::with_trace_sink(sink.clone());
reg.register(skill("keep", "keep", "REST API design", &["api"]));
reg.register(skill("drop", "drop", "HTML slides", &["frontend"]));
sink.drain();
reg.replace_all(vec![
skill("keep", "keep", "REST API design", &["api"]),
skill("add", "add", "database migrations", &["data"]),
]);
let churn: Vec<(ChurnKind, String)> = sink
.drain()
.into_iter()
.filter_map(|envelope| match envelope.event {
TraceEvent::SkillChurn { kind, skill_id } => Some((kind, skill_id)),
_ => None,
})
.collect();
assert_eq!(churn.len(), 2, "an unchanged id emits nothing: {churn:?}");
assert!(churn.contains(&(ChurnKind::Remove, "drop".to_string())));
assert!(churn.contains(&(ChurnKind::Add, "add".to_string())));
}
#[test]
fn semantic_ranks_via_injected_embedder() {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.register(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.register(skill(
"frontend-slides",
"frontend-slides",
"HTML slides",
&["frontend"],
));
reg.build_embeddings().unwrap();
let hits = reg
.search_with_method("rest api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(
hits.first().map(|h| h.skill_id.as_str()),
Some("api-design")
);
}
#[test]
fn semantic_uses_experimental_searchable_description_and_keeps_name_and_tags() {
let mut overridden_reg = with_embedder(Arc::new(StubEmbedder));
let mut overridden = skill("target", "catalog", "REST API design", &["general"]);
overridden.experimental_searchable_description = Some("frontend slides".into());
overridden_reg.register(overridden);
overridden_reg.register(skill("decoy", "decoy", "REST API design", &["general"]));
overridden_reg.build_embeddings().unwrap();
let override_hits = overridden_reg
.search_with_method("frontend slides", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(
override_hits.first().map(|h| h.skill_id.as_str()),
Some("target")
);
let mut name_reg = with_embedder(Arc::new(StubEmbedder));
let mut named = skill("named", "api_helper", "unrelated", &["general"]);
named.experimental_searchable_description = Some("frontend slides".into());
name_reg.register(named);
name_reg.register(skill("name-decoy", "decoy", "frontend slides", &[]));
name_reg.build_embeddings().unwrap();
let name_hits = name_reg
.search_with_method("REST API", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(
name_hits.first().map(|h| h.skill_id.as_str()),
Some("named")
);
let mut tag_reg = with_embedder(Arc::new(StubEmbedder));
let mut tagged = skill("tagged", "catalog", "unrelated", &["rest_ops"]);
tagged.experimental_searchable_description = Some("frontend slides".into());
tag_reg.register(tagged);
tag_reg.register(skill("tag-decoy", "decoy", "frontend slides", &[]));
tag_reg.build_embeddings().unwrap();
let tag_hits = tag_reg
.search_with_method("REST API", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(
tag_hits.first().map(|h| h.skill_id.as_str()),
Some("tagged")
);
}
#[test]
fn build_embeddings_after_register_embeds_only_the_new_skill() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.register(skill("frontend", "frontend", "HTML slides", &["frontend"]));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 2);
reg.register(skill("api-v2", "api-v2", "REST API v2", &["api"]));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 3, "only the new skill is embedded");
}
#[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(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.build_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 1);
reg.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(
counter.doc_calls(),
1,
"a search after build_embeddings embeds only the query"
);
}
#[test]
fn rebuild_embeddings_recomputes_the_full_skill_corpus() {
let counter = Arc::new(CountingEmbedder::new());
let mut reg = with_embedder(counter.clone());
reg.register(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.register(skill("frontend", "frontend", "HTML slides", &["frontend"]));
reg.build_embeddings().unwrap();
reg.rebuild_embeddings().unwrap();
assert_eq!(counter.doc_calls(), 4, "rebuild embeds every skill again");
}
#[test]
fn hybrid_emits_three_stages() {
let sink = Arc::new(MemorySink::new("s"));
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.set_trace_sink(sink.clone());
reg.register(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.build_embeddings().unwrap();
reg.search_with_method("api", 5, Origin::Agent, SearchMethod::Hybrid)
.unwrap();
let events = sink.drain();
assert!(events.iter().any(|e| matches!(
&e.event,
TraceEvent::SkillSearch { 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 register_and_search_emit_trace_events() {
let sink = Arc::new(MemorySink::new("test-session"));
let mut reg = SkillRegistry::with_trace_sink(sink.clone());
reg.register(skill(
"api-design",
"api-design",
"REST API design",
&["api"],
));
reg.search_with_origin("api design", 5, Origin::Agent);
let events = sink.drain();
assert!(events.iter().any(|e| matches!(
e.event,
TraceEvent::SkillChurn {
kind: ChurnKind::Add,
..
}
)));
assert!(events.iter().any(|e| matches!(
&e.event,
TraceEvent::SkillSearch { origin: Origin::Agent, hits, .. } if !hits.is_empty()
)));
}
fn graph_with_model(
skill_id: &str,
centroid: Vec<f32>,
model: &str,
) -> Arc<RwLock<IntentGraph>> {
let c: Vec<String> = centroid.iter().map(|x| x.to_string()).collect();
let json = format!(
r#"{{"v":1,"built_from_ts":1,"model":"{model}",
"intents":[{{"id":"i0","label":"l","terms":[],
"members":["rest api design"],"centroid":[{}],
"support":9,"tools":{{}},"skills":{{"{skill_id}":1.0}}}}]}}"#,
c.join(",")
);
Arc::new(RwLock::new(IntentGraph::from_json(&json).expect("valid")))
}
fn model_mismatch_events(sink: &MemorySink) -> Vec<(String, String, bool)> {
sink.drain()
.into_iter()
.filter_map(|e| match e.event {
TraceEvent::UsageModelMismatch {
built,
active,
dim_mismatch,
} => Some((built, active, dim_mismatch)),
_ => None,
})
.collect()
}
fn mismatch_registry(sink: Arc<MemorySink>) -> SkillRegistry {
let mut reg = with_embedder(Arc::new(StubEmbedder));
reg.set_trace_sink(sink);
reg.register(skill("api-design", "api-design", "rest api design", &[]));
reg.register(skill("frontend", "frontend", "frontend slides", &[]));
reg.build_embeddings().unwrap();
reg
}
#[test]
fn a_same_dim_model_mismatch_pauses_the_arm_and_warns() {
let sink = Arc::new(MemorySink::new("s"));
let mut reg = mismatch_registry(sink.clone());
reg.set_intent_graph(Some(graph_with_model(
"frontend",
vec![1.0, 0.0, 0.0],
"a-different-model",
)));
let hits = reg
.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.last().map(|h| h.skill_id.as_str()), Some("frontend"));
assert!(hits.iter().all(|h| !h.fused), "no fusion — the arm paused");
let events = model_mismatch_events(&sink);
assert_eq!(events.len(), 1);
assert!(!events[0].2, "same-dim swap → dim_mismatch false");
}
#[test]
fn a_dim_mismatch_pauses_the_arm_and_warns() {
let sink = Arc::new(MemorySink::new("s"));
let mut reg = mismatch_registry(sink.clone());
reg.set_intent_graph(Some(graph_with_model(
"frontend",
vec![1.0, 0.0, 0.0, 0.0, 0.0],
"some-model",
)));
let hits = reg
.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(hits.last().map(|h| h.skill_id.as_str()), Some("frontend"));
let events = model_mismatch_events(&sink);
assert_eq!(events.len(), 1);
assert!(events[0].2, "different width → dim_mismatch true");
}
#[test]
fn rebuild_intent_graph_restores_the_arm_after_a_model_change() {
let sink = Arc::new(MemorySink::new("s"));
let mut reg = mismatch_registry(sink.clone());
reg.set_intent_graph(Some(graph_with_model(
"api-design",
vec![1.0, 0.0, 0.0],
"a-different-model",
)));
reg.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert_eq!(model_mismatch_events(&sink).len(), 1);
reg.rebuild_intent_graph().unwrap();
let after = reg
.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert!(after.iter().all(|h| h.fused), "arm resumed → fused ranking");
assert!(
model_mismatch_events(&sink).is_empty(),
"no mismatch after rebuild"
);
}
fn poison(graph: &Arc<RwLock<IntentGraph>>) {
let g = graph.clone();
let _ = std::thread::spawn(move || {
let _guard = g.write().expect("first writer takes the lock");
panic!("intentional poison");
})
.join();
assert!(
graph.is_poisoned(),
"lock should be poisoned after the panic"
);
}
#[test]
fn a_dense_graph_on_a_bm25_catalog_reports_active_not_unknown() {
let mut reg = SkillRegistry::new();
reg.register(skill("api-design", "api-design", "rest api design", &[]));
reg.set_intent_graph(Some(graph_with_model(
"api-design",
vec![1.0, 0.0, 0.0],
"some-model",
)));
assert_eq!(reg.adaptive_ranking_status(), AdaptiveRankingStatus::Active);
}
#[test]
fn a_poisoned_graph_lock_reports_unknown_not_a_panic() {
let mut reg = SkillRegistry::new();
reg.register(skill("api-design", "api-design", "rest api design", &[]));
let graph = Arc::new(RwLock::new(IntentGraph::empty()));
poison(&graph);
reg.set_intent_graph(Some(graph));
assert_eq!(
reg.adaptive_ranking_status(),
AdaptiveRankingStatus::Unknown
);
}
#[test]
fn a_poisoned_graph_lock_degrades_the_search_path_not_a_panic() {
let mut reg = mismatch_registry(Arc::new(MemorySink::new("s")));
let graph = graph_with_model("frontend", vec![0.0, 1.0, 0.0], "m");
poison(&graph);
reg.set_intent_graph(Some(graph));
let hits = reg
.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.unwrap();
assert!(hits.iter().all(|h| !h.fused), "poisoned lock → arm paused");
}
#[test]
fn rebuild_recovers_a_poisoned_graph_lock_not_a_panic() {
let mut reg = SkillRegistry::new();
reg.register(skill("api-design", "api-design", "rest api design", &[]));
let graph = Arc::new(RwLock::new(IntentGraph::empty()));
poison(&graph);
reg.set_intent_graph(Some(graph));
assert!(reg.rebuild_intent_graph().is_ok());
}
#[test]
fn warm_embeddings_error_ok_when_artifact_covers_corpus() {
let a = skill("api-design", "api-design", "rest api design", &[]);
let b = skill("frontend", "frontend", "frontend slides", &[]);
let bytes = build_test_artifact(
ArtifactEntryKind::Skill,
[&a, &b],
"fp-warm",
vec![unit([1.0, 0.0, 0.0]), unit([0.0, 1.0, 0.0])],
);
let counter = Arc::new(FpCountingEmbedder::new("fp-warm", StubEmbedder::vec_for));
let mut reg = with_embedder(counter.clone());
reg.register(a);
reg.register(b);
reg.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Error)
.unwrap();
assert_eq!(counter.docs(), 0);
assert!(
reg.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.is_ok()
);
}
#[test]
fn warm_embeddings_error_fails_when_ids_missing() {
let a = skill("api-design", "api-design", "rest api design", &[]);
let b = skill("frontend", "frontend", "frontend slides", &[]);
let bytes = build_test_artifact(
ArtifactEntryKind::Skill,
[&a],
"fp-warm",
vec![unit([1.0, 0.0, 0.0])],
);
let counter = Arc::new(FpCountingEmbedder::new("fp-warm", StubEmbedder::vec_for));
let mut reg = with_embedder(counter.clone());
reg.register(a);
reg.register(b);
assert!(matches!(
reg.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Error),
Err(ArtifactWarmError::Incomplete { missing }) if missing == ["frontend"]
));
assert_eq!(counter.docs(), 0);
assert!(matches!(
reg.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic),
Err(EmbedderError::EmbeddingsNotBuilt)
));
}
#[test]
fn warm_embeddings_embed_completes_only_missing_ids() {
let a = skill("api-design", "api-design", "rest api design", &[]);
let b = skill("frontend", "frontend", "frontend slides", &[]);
let bytes = build_test_artifact(
ArtifactEntryKind::Skill,
[&a],
"fp-warm",
vec![unit([1.0, 0.0, 0.0])],
);
let counter = Arc::new(FpCountingEmbedder::new("fp-warm", StubEmbedder::vec_for));
let mut reg = with_embedder(counter.clone());
reg.register(a);
reg.register(b);
reg.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Embed)
.unwrap();
assert_eq!(counter.docs(), 1, "only the missing skill is embedded");
assert!(
reg.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.is_ok()
);
}
#[test]
fn warm_embeddings_embed_policy_propagates_build_embeddings_failure() {
let a = skill("api-design", "api-design", "rest api design", &[]);
let b = skill("frontend", "frontend", "frontend slides", &[]);
let bytes = build_test_artifact(
ArtifactEntryKind::Skill,
[&a],
"fp-warm",
vec![unit([1.0, 0.0, 0.0])],
);
let mut reg = with_embedder(Arc::new(FailOnEmbedStub::new("fp-warm")));
reg.register(a);
reg.register(b);
assert!(matches!(
reg.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Embed),
Err(ArtifactWarmError::Embedder(EmbedderError::Inference { .. }))
));
}
#[test]
fn warm_embeddings_propagates_warm_error_without_embed() {
let a = skill("api-design", "api-design", "rest api design", &[]);
let bytes = build_test_artifact(
ArtifactEntryKind::Skill,
[&a],
"fp-artifact",
vec![unit([1.0, 0.0, 0.0])],
);
let counter = Arc::new(FpCountingEmbedder::new("fp-active", StubEmbedder::vec_for));
let mut reg = with_embedder(counter.clone());
reg.register(a);
assert!(matches!(
reg.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Embed),
Err(ArtifactWarmError::Warm(
crate::WarmError::ArtifactModelMismatch { .. }
))
));
assert_eq!(counter.docs(), 0);
}
#[test]
fn build_embedding_artifact_round_trips_via_warm() {
let a = skill("api-design", "api-design", "rest api design", &[]);
let b = skill("frontend", "frontend", "frontend slides", &[]);
let builder = Arc::new(FpCountingEmbedder::new("fp-warm", StubEmbedder::vec_for));
let mut reg_a = with_embedder(builder.clone());
reg_a.register(skill("api-design", "api-design", "rest api design", &[]));
reg_a.register(skill("frontend", "frontend", "frontend slides", &[]));
let bytes = reg_a.build_embedding_artifact().unwrap();
assert_eq!(
builder.docs(),
2,
"build embeds each corpus document exactly once"
);
let warmer = Arc::new(FpCountingEmbedder::new("fp-warm", StubEmbedder::vec_for));
let mut reg_b = with_embedder(warmer.clone());
reg_b.register(a);
reg_b.register(b);
reg_b
.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Error)
.unwrap();
assert_eq!(warmer.docs(), 0, "warm must not re-embed covered ids");
assert!(
reg_b
.search_with_method("api", 5, Origin::Direct, SearchMethod::Semantic)
.is_ok()
);
}
#[test]
fn build_embedding_artifact_propagates_embedder_failure() {
let mut reg = with_embedder(Arc::new(FailOnEmbedStub::new("fp-warm")));
reg.register(skill("api-design", "api-design", "rest api design", &[]));
assert!(matches!(
reg.build_embedding_artifact(),
Err(ArtifactError::Embedder(EmbedderError::Inference { .. }))
));
}
#[test]
fn build_embedding_artifact_empty_corpus_is_valid() {
let reg = with_embedder(Arc::new(PanicOnEmbedStub::new("unused")));
let bytes = reg.build_embedding_artifact().unwrap();
reg.warm_embeddings_from_artifact(&bytes, OnArtifactMiss::Error)
.unwrap();
}
}