use ineru::{Embedder, Embedding, HashEmbedder};
use std::sync::Arc;
pub fn build_embedder(model_dir: Option<&str>) -> Arc<dyn Embedder> {
if let Ok(raw) = std::env::var("ORT_DYLIB_PATH") {
if let Some(plain) = strip_extended_length_prefix(&raw) {
log::info!("normalized extended-length ORT_DYLIB_PATH to {plain}");
std::env::set_var("ORT_DYLIB_PATH", plain);
}
}
#[cfg(feature = "neural-embeddings")]
if let Some(dir) = model_dir {
let path = std::path::Path::new(dir);
let mut last_err = String::new();
for attempt in 1..=5u32 {
match ineru::NeuralEmbedder::from_path(path) {
Ok(e) => {
log::info!("Using neural embedder from {dir} (attempt {attempt})");
return Arc::new(e);
}
Err(e) => {
last_err = e.to_string();
log::warn!("neural embedder load attempt {attempt}/5 failed: {last_err}");
if attempt < 5 {
std::thread::sleep(std::time::Duration::from_millis(400 * attempt as u64));
}
}
}
}
log::warn!(
"Failed to load neural embedder from {dir} after 5 attempts: {last_err}. Using hash embedder."
);
}
#[cfg(not(feature = "neural-embeddings"))]
if model_dir.is_some() {
log::warn!(
"--embed-model was set but cortex was built without the `neural-embeddings` \
feature; using the hash embedder."
);
}
Arc::new(HashEmbedder::new())
}
pub fn strip_extended_length_prefix(path: &str) -> Option<String> {
if let Some(unc) = path.strip_prefix(r"\\?\UNC\") {
return Some(format!(r"\\{unc}"));
}
path.strip_prefix(r"\\?\").map(str::to_owned)
}
pub fn read_dims(dir: &std::path::Path) -> Option<usize> {
let raw = std::fs::read_to_string(dir.join("embedder.dims")).ok()?;
raw.trim().parse::<usize>().ok()
}
pub fn write_dims(dir: &std::path::Path, dims: usize) {
if let Err(e) = std::fs::write(dir.join("embedder.dims"), dims.to_string()) {
log::warn!("Failed to write embedder.dims sidecar: {e}");
}
}
pub fn read_identity(dir: &std::path::Path) -> Option<String> {
let raw = std::fs::read_to_string(dir.join("embedder.id")).ok()?;
let s = raw.trim();
if s.is_empty() {
None
} else {
Some(s.to_string())
}
}
pub fn write_identity(dir: &std::path::Path, identity: &str) {
if identity.starts_with("pending-") {
log::warn!("refusing to persist embedder identity while pending: {identity}");
return;
}
let final_path = dir.join("embedder.id");
let tmp_path = dir.join("embedder.id.tmp");
if let Err(e) = std::fs::write(&tmp_path, identity) {
log::warn!("Failed to write embedder.id sidecar: {e}");
return;
}
if let Err(e) = std::fs::rename(&tmp_path, &final_path) {
log::warn!("Failed to finalize embedder.id sidecar: {e}");
let _ = std::fs::remove_file(&tmp_path);
}
}
pub fn clear_source_registry(graph: &aingle_graph::GraphDB) -> Result<usize, aingle_graph::Error> {
use aingle_graph::{Predicate, TriplePattern};
let pattern = TriplePattern::any()
.with_predicate(Predicate::named(crate::service::ingest::PRED_SOURCE_HASH));
let ids: Vec<_> = graph.find(pattern)?.into_iter().map(|t| t.id()).collect();
let mut removed = 0;
for id in &ids {
match graph.delete(id) {
Ok(true) => removed += 1,
Ok(false) => {} Err(e) => log::warn!("clear_source_registry: delete failed for {id:?}: {e}"),
}
}
Ok(removed)
}
pub struct SwappableEmbedder {
inner: std::sync::RwLock<Arc<dyn Embedder>>,
dims: usize,
}
struct PendingEmbedder {
dims: usize,
}
impl Embedder for PendingEmbedder {
fn embed_passage(&self, _text: &str) -> Embedding {
Embedding::new(vec![0.0; self.dims])
}
fn embed_query(&self, _text: &str) -> Embedding {
Embedding::new(vec![0.0; self.dims])
}
fn dimensions(&self) -> usize {
self.dims
}
fn identity(&self) -> String {
format!("pending-{}", self.dims)
}
}
impl SwappableEmbedder {
pub fn new_pending(dims: usize) -> Self {
Self {
inner: std::sync::RwLock::new(Arc::new(PendingEmbedder { dims })),
dims,
}
}
pub fn install(&self, delegate: Arc<dyn Embedder>) {
if delegate.dimensions() != self.dims {
log::warn!(
"SwappableEmbedder.install rejected: delegate dims {} != fixed {}",
delegate.dimensions(),
self.dims
);
return;
}
*self.inner.write().expect("swappable embedder poisoned") = delegate;
}
}
impl Embedder for SwappableEmbedder {
fn embed_passage(&self, text: &str) -> Embedding {
let inner = self
.inner
.read()
.expect("swappable embedder poisoned")
.clone();
inner.embed_passage(text)
}
fn embed_query(&self, text: &str) -> Embedding {
let inner = self
.inner
.read()
.expect("swappable embedder poisoned")
.clone();
inner.embed_query(text)
}
fn embed_passages(&self, texts: &[String]) -> Vec<Embedding> {
let inner = self
.inner
.read()
.expect("swappable embedder poisoned")
.clone();
inner.embed_passages(texts)
}
fn dimensions(&self) -> usize {
self.dims
}
fn relevance_thresholds(&self) -> (f32, f32) {
let inner = self
.inner
.read()
.expect("swappable embedder poisoned")
.clone();
inner.relevance_thresholds()
}
fn identity(&self) -> String {
let inner = self
.inner
.read()
.expect("swappable embedder poisoned")
.clone();
inner.identity()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn build_embedder_without_model_is_hash_64d() {
let e = build_embedder(None);
assert_eq!(e.dimensions(), 64);
}
#[test]
fn strips_windows_extended_length_prefix() {
assert_eq!(
strip_extended_length_prefix(r"\\?\C:\app\onnxruntime.dll").as_deref(),
Some(r"C:\app\onnxruntime.dll")
);
assert_eq!(
strip_extended_length_prefix(r"\\?\UNC\srv\share\ort.dll").as_deref(),
Some(r"\\srv\share\ort.dll")
);
assert_eq!(
strip_extended_length_prefix(r"C:\app\onnxruntime.dll"),
None
);
assert_eq!(
strip_extended_length_prefix("/usr/lib/libonnxruntime.so"),
None
);
}
#[test]
fn build_embedder_missing_dir_falls_back_to_hash() {
let e = build_embedder(Some("/nonexistent/model/dir"));
assert_eq!(e.dimensions(), 64);
}
#[test]
fn dims_sidecar_round_trips() {
let dir = tempfile::tempdir().unwrap();
write_dims(dir.path(), 384);
assert_eq!(read_dims(dir.path()), Some(384));
}
#[test]
fn read_dims_absent_is_none() {
let dir = tempfile::tempdir().unwrap();
assert_eq!(read_dims(dir.path()), None);
}
#[test]
fn clear_source_registry_on_empty_graph_is_zero() {
let graph = aingle_graph::GraphDB::memory().unwrap();
assert_eq!(clear_source_registry(&graph).unwrap(), 0);
}
#[test]
fn swappable_reports_fixed_dims_before_and_after_install() {
let s = SwappableEmbedder::new_pending(384);
assert_eq!(s.dimensions(), 384);
let q = s.embed_query("hola");
assert_eq!(q.0.len(), 384);
assert!(q.0.iter().all(|x| *x == 0.0));
s.install(std::sync::Arc::new(Fake384));
assert_eq!(s.dimensions(), 384);
let q2 = s.embed_query("hola");
assert_eq!(q2.0.len(), 384);
assert!(q2.0.iter().any(|x| *x != 0.0));
}
#[test]
fn swappable_rejects_mismatched_dims_install() {
let s = SwappableEmbedder::new_pending(384);
s.install(std::sync::Arc::new(ineru::HashEmbedder::new())); let q = s.embed_query("x");
assert_eq!(q.0.len(), 384);
assert!(q.0.iter().all(|x| *x == 0.0));
}
#[test]
fn swappable_delegates_relevance_thresholds_after_install() {
let s = SwappableEmbedder::new_pending(384);
s.install(std::sync::Arc::new(Fake384));
assert_eq!(s.relevance_thresholds(), (0.80, 0.77));
}
struct Fake384;
impl ineru::Embedder for Fake384 {
fn embed_passage(&self, _t: &str) -> ineru::Embedding {
ineru::Embedding::new(vec![0.5; 384])
}
fn embed_query(&self, _t: &str) -> ineru::Embedding {
ineru::Embedding::new(vec![0.5; 384])
}
fn dimensions(&self) -> usize {
384
}
fn relevance_thresholds(&self) -> (f32, f32) {
(0.80, 0.77)
}
}
}