use std::sync::Arc;
use anyhow::{Result, anyhow};
use tokio::sync::OnceCell;
use tokio::sync::mpsc::UnboundedSender;
use crate::app::events::AppEvent;
use crate::shared::api::{EmbedRole, Embedder};
use crate::shared::embed_calibration;
use crate::shared::embed_identity::{CANARY_TEXT, EmbedFingerprint};
use crate::shared::embed_prefix::EmbedConvention;
use crate::shared::i18n::Locale;
use crate::shared::storage::Storage;
use crate::shared::storage::db::ReembedPending;
pub(super) struct EmbedGuard {
inner: Arc<dyn Embedder>,
storage: Arc<Storage>,
model_id: Option<String>,
convention: EmbedConvention,
loc: &'static Locale,
evt_tx: UnboundedSender<AppEvent>,
checked: OnceCell<()>,
}
#[derive(Debug, Default, PartialEq)]
struct Invalidated {
pending: ReembedPending,
stale_profiles: usize,
}
impl Invalidated {
fn total(&self) -> usize {
self.pending.notes + self.pending.attachments + self.pending.rag
}
fn is_empty(&self) -> bool {
*self == Self::default()
}
}
impl EmbedGuard {
pub(super) fn new(
inner: Arc<dyn Embedder>,
storage: Arc<Storage>,
model_id: Option<String>,
convention: EmbedConvention,
loc: &'static Locale,
evt_tx: UnboundedSender<AppEvent>,
) -> Self {
Self {
inner,
storage,
model_id,
convention,
loc,
evt_tx,
checked: OnceCell::new(),
}
}
async fn run_check(&self) -> Result<()> {
let fresh = self
.inner
.embed(vec![CANARY_TEXT.to_string()], EmbedRole::Passage)
.await?
.into_iter()
.next()
.filter(|v| !v.is_empty())
.ok_or_else(|| anyhow!("embedder returned no canary vector"))?;
let stored = self.storage.db().embed_fingerprint().unwrap_or(None);
let current = EmbedFingerprint::new(fresh, self.model_id.clone(), self.convention.id());
match stored {
Some(prev) if prev.matches(¤t) => return Ok(()),
Some(prev) => {
let hit = self.invalidate();
tracing::warn!(
previous = prev.display_id(),
current = current.display_id(),
notes = hit.pending.notes,
attachments = hit.pending.attachments,
rag = hit.pending.rag,
stale_profiles = hit.stale_profiles,
"embedding model changed: vectors from the previous model were invalidated"
);
self.notify(&prev, ¤t, &hit);
}
None => tracing::info!(
model = current.display_id(),
"recorded the embedding model fingerprint"
),
}
if let Err(err) = self.storage.db().set_embed_fingerprint(¤t) {
tracing::warn!(error = %err, "failed to record the embedding fingerprint");
}
self.calibrate().await;
Ok(())
}
async fn calibrate(&self) {
let probes = embed_calibration::probe_texts();
let vectors = match self.inner.embed(probes, EmbedRole::Passage).await {
Ok(v) => v,
Err(err) => {
tracing::warn!(error = %err, "similarity calibration skipped");
return;
}
};
let Some(calibration) = embed_calibration::measure(&vectors) else {
tracing::warn!("similarity calibration probe returned an unusable result");
return;
};
match self.storage.db().set_embed_calibration(&calibration) {
Ok(()) => tracing::info!(
unrelated = calibration.unrelated,
paraphrase = calibration.paraphrase,
"calibrated the similarity scale"
),
Err(err) => tracing::warn!(error = %err, "failed to record the similarity calibration"),
}
}
fn invalidate(&self) -> Invalidated {
let db = self.storage.db();
if let Err(err) = db.bump_embed_generation() {
tracing::error!(error = %err, "failed to retire the previous model's vectors");
}
let pending = db.count_rows_to_reembed().unwrap_or_else(|err| {
tracing::warn!(error = %err, "failed to count vectors awaiting re-embedding");
Default::default()
});
let stale = db.profiles_with_rag_docs().unwrap_or_else(|err| {
tracing::warn!(error = %err, "failed to list profiles with RAG documents");
Vec::new()
});
if !stale.is_empty()
&& let Err(err) = db.set_rag_stale_profiles(&stale)
{
tracing::warn!(error = %err, "failed to record stale knowledge bases");
}
Invalidated {
pending,
stale_profiles: stale.len(),
}
}
fn notify(&self, prev: &EmbedFingerprint, current: &EmbedFingerprint, hit: &Invalidated) {
if hit.is_empty() {
return;
}
let mut msg = self.loc.tf(
"ui.embed.model_changed",
&[("old", prev.display_id()), ("new", current.display_id())],
);
msg.push(' ');
msg.push_str(
&self
.loc
.tf("ui.embed.reindex_hint", &[("n", &hit.total().to_string())]),
);
if hit.stale_profiles > 0 {
msg.push(' ');
msg.push_str(&self.loc.tf(
"ui.embed.rag_stale",
&[("n", &hit.stale_profiles.to_string())],
));
}
if let Some(suggested) = self.suggested_convention() {
msg.push(' ');
msg.push_str(
&self
.loc
.tf("ui.embed.convention_hint", &[("name", suggested.id())]),
);
}
let _ = self.evt_tx.send(AppEvent::Error(msg));
}
fn suggested_convention(&self) -> Option<EmbedConvention> {
let suggested = EmbedConvention::suggested_for(self.model_id.as_deref()?)?;
(suggested != self.convention).then_some(suggested)
}
}
#[async_trait::async_trait]
impl Embedder for EmbedGuard {
async fn embed(&self, texts: Vec<String>, role: EmbedRole) -> Result<Vec<Vec<f32>>> {
let _ = self.checked.get_or_try_init(|| self.run_check()).await;
self.inner.embed(texts, role).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::entities::note::Note;
use crate::shared::api::mock::MockEmbedder;
use crate::shared::i18n::{Lang, locale};
use crate::shared::paths::Paths;
use uuid::Uuid;
struct SaltedEmbedder {
salt: usize,
dim: usize,
}
#[async_trait::async_trait]
impl Embedder for SaltedEmbedder {
async fn embed(&self, texts: Vec<String>, _role: EmbedRole) -> Result<Vec<Vec<f32>>> {
Ok(texts
.iter()
.map(|t| {
let mut v = vec![0.0; self.dim];
for (i, b) in t.bytes().enumerate() {
v[(i + b as usize + self.salt * 7) % self.dim] += 1.0;
}
v[self.salt % self.dim] += 5.0; v
})
.collect())
}
}
struct Fixture {
_dir: tempfile::TempDir,
storage: Arc<Storage>,
rx: tokio::sync::mpsc::UnboundedReceiver<AppEvent>,
tx: UnboundedSender<AppEvent>,
}
fn fixture() -> Fixture {
let dir = tempfile::tempdir().unwrap();
let storage = Arc::new(Storage::open(Paths::with_root(dir.path())).unwrap());
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
Fixture {
_dir: dir,
storage,
rx,
tx,
}
}
fn guard(f: &Fixture, inner: Arc<dyn Embedder>, model: &str) -> EmbedGuard {
guard_with(f, inner, model, EmbedConvention::None)
}
fn guard_with(
f: &Fixture,
inner: Arc<dyn Embedder>,
model: &str,
convention: EmbedConvention,
) -> EmbedGuard {
EmbedGuard::new(
Arc::new(crate::shared::embed_prefix::PrefixedEmbedder::new(
inner, convention,
)),
f.storage.clone(),
Some(model.to_string()),
convention,
locale(Lang::Ru),
f.tx.clone(),
)
}
#[derive(Default)]
struct Recorder {
seen: std::sync::Mutex<Vec<String>>,
}
#[async_trait::async_trait]
impl Embedder for Recorder {
async fn embed(&self, texts: Vec<String>, _role: EmbedRole) -> Result<Vec<Vec<f32>>> {
self.seen.lock().unwrap().extend(texts.iter().cloned());
Ok(texts.iter().map(|_| vec![1.0, 0.0]).collect())
}
}
#[tokio::test]
async fn the_canary_and_the_calibration_go_through_the_prefixer() {
let f = fixture();
let rec = Arc::new(Recorder::default());
guard_with(&f, rec.clone(), "e5", EmbedConvention::E5)
.embed(vec!["настоящий текст".into()], EmbedRole::Passage)
.await
.unwrap();
let seen = rec.seen.lock().unwrap();
assert!(
seen.iter().any(|t| t == &format!("passage: {CANARY_TEXT}")),
"the canary must carry the passage marker: {seen:?}"
);
assert!(
seen.iter().filter(|t| t.starts_with("passage: ")).count() > 30,
"the 32 calibration probes are prefixed too: {}",
seen.len()
);
assert!(
seen.iter().any(|t| t == "passage: настоящий текст"),
"and so is the real call: {seen:?}"
);
}
#[tokio::test]
async fn turning_a_convention_on_reads_as_a_changed_vector_space() {
let mut f = fixture();
let profile = Uuid::new_v4();
let note = Note::new(profile, "заметка", vec![]);
f.storage.db().note_insert(¬e).unwrap();
f.storage
.db()
.note_vector_upsert(note.id, profile, &[1.0, 0.0])
.unwrap();
let embedder = Arc::new(MockEmbedder::new(16));
guard_with(&f, embedder.clone(), "e5", EmbedConvention::None)
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
let gen_before = f.storage.db().embed_generation().unwrap();
while f.rx.try_recv().is_ok() {}
guard_with(&f, embedder, "e5", EmbedConvention::E5)
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
assert!(
f.storage.db().embed_generation().unwrap() > gen_before,
"the generation must be bumped, retiring the old vectors"
);
assert!(
f.rx.try_recv().is_ok(),
"and the user must be told, as for any model change"
);
assert!(f.storage.db().note_get(profile, note.id).unwrap().is_some());
}
#[tokio::test]
async fn the_model_change_notice_hints_at_a_matching_convention() {
let mut f = fixture();
let profile = Uuid::new_v4();
let note = Note::new(profile, "заметка", vec![]);
f.storage.db().note_insert(¬e).unwrap();
f.storage
.db()
.note_vector_upsert(note.id, profile, &[1.0, 0.0])
.unwrap();
guard(&f, Arc::new(SaltedEmbedder { salt: 0, dim: 16 }), "bge-m3")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
while f.rx.try_recv().is_ok() {}
guard(
&f,
Arc::new(SaltedEmbedder { salt: 3, dim: 16 }),
"multilingual-e5-large-instruct-q8_0.gguf",
)
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
let AppEvent::Error(msg) = f.rx.try_recv().expect("a notice") else {
panic!("expected an error event");
};
assert!(msg.contains("e5-instruct"), "{msg}");
}
#[tokio::test]
async fn no_hint_when_the_convention_already_matches() {
let mut f = fixture();
let profile = Uuid::new_v4();
let note = Note::new(profile, "заметка", vec![]);
f.storage.db().note_insert(¬e).unwrap();
f.storage
.db()
.note_vector_upsert(note.id, profile, &[1.0, 0.0])
.unwrap();
guard_with(
&f,
Arc::new(SaltedEmbedder { salt: 0, dim: 16 }),
"bge-m3",
EmbedConvention::E5Instruct,
)
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
while f.rx.try_recv().is_ok() {}
guard_with(
&f,
Arc::new(SaltedEmbedder { salt: 3, dim: 16 }),
"multilingual-e5-large-instruct-q8_0.gguf",
EmbedConvention::E5Instruct,
)
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
let AppEvent::Error(msg) = f.rx.try_recv().expect("a notice") else {
panic!("expected an error event");
};
assert!(
!msg.contains("e5-instruct"),
"nothing to suggest when it is already set: {msg}"
);
}
#[tokio::test]
async fn first_run_records_fingerprint_without_notifying() {
let mut f = fixture();
let g = guard(&f, Arc::new(MockEmbedder::new(16)), "model-a");
g.embed(vec!["hello".into()], EmbedRole::Passage)
.await
.unwrap();
assert!(
f.storage.db().embed_fingerprint().unwrap().is_some(),
"the fingerprint is recorded on the first use"
);
assert!(
f.rx.try_recv().is_err(),
"a first launch must not claim anything changed"
);
}
#[tokio::test]
async fn same_model_is_silent_and_keeps_vectors() {
let mut f = fixture();
let profile = Uuid::new_v4();
let note = Note::new(profile, "a note", vec![]);
f.storage.db().note_insert(¬e).unwrap();
f.storage
.db()
.note_vector_upsert(note.id, profile, &[1.0, 0.0])
.unwrap();
guard(&f, Arc::new(MockEmbedder::new(16)), "model-a")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
let _ = f.rx.try_recv();
guard(&f, Arc::new(MockEmbedder::new(16)), "model-a")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
assert!(f.rx.try_recv().is_err(), "no notice for an unchanged model");
assert!(
f.storage
.db()
.notes_missing_vectors(profile)
.unwrap()
.is_empty(),
"an unchanged model must not invalidate anything"
);
}
#[tokio::test]
async fn same_dimension_model_swap_is_detected_and_invalidates() {
let mut f = fixture();
let profile = Uuid::new_v4();
let note = Note::new(profile, "a note", vec![]);
f.storage.db().note_insert(¬e).unwrap();
f.storage
.db()
.note_vector_upsert(note.id, profile, &[1.0, 0.0])
.unwrap();
guard(&f, Arc::new(SaltedEmbedder { salt: 0, dim: 32 }), "model-a")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
let _ = f.rx.try_recv();
guard(&f, Arc::new(SaltedEmbedder { salt: 3, dim: 32 }), "model-b")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
assert_eq!(
f.storage.db().notes_missing_vectors(profile).unwrap().len(),
1,
"note vectors read as foreign, so the existing backfill re-embeds them"
);
assert_eq!(
f.storage
.db()
.note_list(profile, None, &[], None)
.unwrap()
.len(),
1,
"the note content itself is untouched"
);
match f.rx.try_recv() {
Ok(AppEvent::Error(msg)) => {
assert!(msg.contains("model-a") && msg.contains("model-b"), "{msg}");
}
other => panic!("expected a model-change notice, got {other:?}"),
}
}
#[tokio::test]
async fn rag_documents_are_marked_stale_not_deleted() {
let f = fixture();
let profile = Uuid::new_v4();
f.storage
.db()
.rag_insert(&crate::entities::rag::RagDocument::new(
profile,
"kb.txt",
"chunk",
vec![1.0; 32],
))
.unwrap();
guard(&f, Arc::new(SaltedEmbedder { salt: 0, dim: 32 }), "model-a")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
guard(&f, Arc::new(SaltedEmbedder { salt: 3, dim: 32 }), "model-b")
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap();
assert!(
f.storage.db().rag_is_stale(profile).unwrap(),
"the knowledge base is marked stale for the affected profile"
);
assert_eq!(
f.storage.db().rag_count(profile).unwrap(),
1,
"the user's knowledge base is never deleted — only marked"
);
}
#[tokio::test]
async fn detection_happens_once_per_instance() {
let f = fixture();
let g = guard(&f, Arc::new(MockEmbedder::new(16)), "model-a");
g.embed(vec!["a".into()], EmbedRole::Passage).await.unwrap();
let first = f.storage.db().embed_fingerprint().unwrap();
f.storage
.db()
.set_embed_fingerprint(&EmbedFingerprint::new(
vec![9.0; 4],
Some("junk".into()),
EmbedConvention::None.id(),
))
.unwrap();
g.embed(vec!["b".into()], EmbedRole::Passage).await.unwrap();
assert_ne!(
f.storage.db().embed_fingerprint().unwrap(),
first,
"the check is cached per instance, so the record stays as we left it"
);
}
#[tokio::test]
#[ignore = "requires two live embedding servers (MINDFORK_EMBED_URL, MINDFORK_EMBED_URL_ALT)"]
async fn same_dimension_model_swap_detected_live() {
let (Some(a), Some(b)) = (
crate::shared::api::live_client("MINDFORK_EMBED_URL", "MINDFORK_EMBED_KEY"),
crate::shared::api::live_client("MINDFORK_EMBED_URL_ALT", "MINDFORK_EMBED_KEY_ALT"),
) else {
eprintln!("skip: MINDFORK_EMBED_URL / MINDFORK_EMBED_URL_ALT not set");
return;
};
let model_a: Arc<dyn Embedder> = Arc::new(a);
let model_b: Arc<dyn Embedder> = Arc::new(b);
let mut f = fixture();
let profile = Uuid::new_v4();
let note = Note::new(profile, "the user prefers concise answers", vec![]);
f.storage.db().note_insert(¬e).unwrap();
let vec_a = model_a
.embed(vec![note.content.clone()], EmbedRole::Passage)
.await
.unwrap()
.remove(0);
let dim_a = vec_a.len();
f.storage
.db()
.note_vector_upsert(note.id, profile, &vec_a)
.unwrap();
f.storage
.db()
.rag_insert(&crate::entities::rag::RagDocument::new(
profile, "kb.txt", "a chunk", vec_a,
))
.unwrap();
guard(&f, model_a.clone(), "bge-m3")
.embed(vec!["warm up".into()], EmbedRole::Passage)
.await
.unwrap();
assert!(f.storage.db().embed_fingerprint().unwrap().is_some());
assert!(f.rx.try_recv().is_err(), "a first launch claims nothing");
guard(&f, model_a, "bge-m3")
.embed(vec!["warm up".into()], EmbedRole::Passage)
.await
.unwrap();
assert!(
f.rx.try_recv().is_err(),
"the same live model must not read as a change"
);
assert!(
f.storage
.db()
.notes_missing_vectors(profile)
.unwrap()
.is_empty(),
"the same live model must not invalidate anything"
);
let dim_b = model_b
.embed(vec!["x".into()], EmbedRole::Passage)
.await
.unwrap()[0]
.len();
assert_eq!(
dim_a, dim_b,
"this smoke is only meaningful for two models of the SAME dimensionality \
(that is the case no existing guard can catch)"
);
guard(&f, model_b, "multilingual-e5-large-instruct")
.embed(vec!["warm up".into()], EmbedRole::Passage)
.await
.unwrap();
assert_eq!(
f.storage.db().notes_missing_vectors(profile).unwrap().len(),
1,
"the note's old-model vector must read as foreign so the backfill re-embeds it"
);
assert!(
f.storage.db().rag_is_stale(profile).unwrap(),
"the knowledge base must be marked stale"
);
assert_eq!(
f.storage.db().rag_count(profile).unwrap(),
1,
"the knowledge base itself must be left intact"
);
assert!(
matches!(f.rx.try_recv(), Ok(AppEvent::Error(_))),
"the user must be told"
);
}
#[tokio::test]
#[ignore = "requires two live embedding servers (MINDFORK_EMBED_URL, MINDFORK_EMBED_URL_ALT)"]
async fn similarity_scale_follows_the_model_live() {
use crate::shared::embed_calibration::{REFERENCE_PARAPHRASE, REFERENCE_UNRELATED};
let (Some(a), Some(b)) = (
crate::shared::api::live_client("MINDFORK_EMBED_URL", "MINDFORK_EMBED_KEY"),
crate::shared::api::live_client("MINDFORK_EMBED_URL_ALT", "MINDFORK_EMBED_KEY_ALT"),
) else {
eprintln!("skip: MINDFORK_EMBED_URL / MINDFORK_EMBED_URL_ALT not set");
return;
};
let bge: Arc<dyn Embedder> = Arc::new(a);
let e5: Arc<dyn Embedder> = Arc::new(b);
let f_bge = fixture();
guard(&f_bge, bge.clone(), "bge-m3")
.embed(vec!["warm up".into()], EmbedRole::Passage)
.await
.unwrap();
let c = f_bge
.storage
.db()
.embed_calibration()
.unwrap()
.expect("bge-m3 must calibrate");
eprintln!("bge-m3 calibration: {c:?}");
assert!(
(c.unrelated - REFERENCE_UNRELATED).abs() < 0.05
&& (c.paraphrase - REFERENCE_PARAPHRASE).abs() < 0.05,
"the reference constants must still describe bge-m3: {c:?}"
);
let scale_bge = f_bge.storage.db().similarity_scale();
for t in [0.85, 0.72, 0.62] {
assert!(
(scale_bge.map(t) - t).abs() < 0.05,
"bge-m3 must keep its thresholds: {t} -> {}",
scale_bge.map(t)
);
}
let f_e5 = fixture();
guard(&f_e5, e5.clone(), "e5-large-instruct")
.embed(vec!["warm up".into()], EmbedRole::Passage)
.await
.unwrap();
let scale_e5 = f_e5.storage.db().similarity_scale();
eprintln!(
"e5 calibration: {:?}; 0.72 -> {:.4}, 0.85 -> {:.4}",
f_e5.storage.db().embed_calibration().unwrap(),
scale_e5.map(0.72),
scale_e5.map(0.85)
);
assert!(
scale_e5.map(0.72) > 0.85,
"e5's trait gate must move well above the raw constant: {}",
scale_e5.map(0.72)
);
let unrelated = e5
.embed(
vec![
"the user values brevity in answers".into(),
"the train leaves from platform nine".into(),
],
EmbedRole::Passage,
)
.await
.unwrap();
let s = cosine_for_test(&unrelated[0], &unrelated[1]);
eprintln!("e5 unrelated pair scores {s:.4}");
assert!(
s >= 0.72,
"this smoke is only meaningful while the raw constant misfires on e5 \
(measured 0.75); got {s:.4}"
);
assert!(
s < scale_e5.map(0.72),
"the calibrated trait gate must reject an unrelated pair: {s:.4} vs {:.4}",
scale_e5.map(0.72)
);
}
fn cosine_for_test(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if na == 0.0 || nb == 0.0 {
0.0
} else {
dot / (na * nb)
}
}
#[tokio::test]
async fn unavailable_embedder_does_not_record_anything() {
let f = fixture();
let g = guard(
&f,
Arc::new(crate::shared::api::UnavailableEmbedder),
"model-a",
);
assert!(
g.embed(vec!["x".into()], EmbedRole::Passage).await.is_err(),
"the real error surfaces"
);
assert!(
f.storage.db().embed_fingerprint().unwrap().is_none(),
"a failed check must not be cached as a verdict"
);
}
}