use anyhow::Result;
use ndarray::Array2;
use next_plaid_onnx::Colbert;
use parking_lot::Mutex;
use std::path::{Path, PathBuf};
use std::sync::OnceLock;
pub type TokenEmbeddings = Array2<f32>;
const DEFAULT_ORT_BATCH: usize = 32;
fn default_index_threads(cores: usize, max_concurrent_builds: usize) -> usize {
if max_concurrent_builds <= 1 {
cores.clamp(2, 8)
} else {
(cores / 2).clamp(2, 8)
}
}
fn index_ort_threads() -> usize {
let cores = std::thread::available_parallelism().map_or(4, std::num::NonZeroUsize::get);
let default = default_index_threads(cores, crate::index::gate::max_concurrent_builds());
crate::config::env_usize("SEMANTEX_INDEX_ORT_THREADS", default)
}
pub struct ColbertEmbedder {
model_dir: PathBuf,
threads: usize,
use_coreml: bool,
encoder: OnceLock<Mutex<Colbert>>,
build_lock: std::sync::Mutex<()>,
#[doc(hidden)]
build_count: std::sync::atomic::AtomicUsize,
}
static GLOBAL_COLBERT: OnceLock<ColbertEmbedder> = OnceLock::new();
static COLBERT_INIT_LOCK: parking_lot::Mutex<()> = parking_lot::Mutex::new(());
impl ColbertEmbedder {
pub fn new(model_dir: &Path) -> Result<Self> {
Self::with_threads(
model_dir,
crate::config::env_usize("SEMANTEX_ORT_THREADS", 4),
)
}
pub fn for_indexing(model_dir: &Path) -> Result<Self> {
Self::with_threads(model_dir, index_ort_threads())
}
fn with_threads(model_dir: &Path, threads: usize) -> Result<Self> {
if !model_dir.exists() {
anyhow::bail!("ColBERT model dir does not exist: {}", model_dir.display());
}
let use_coreml = std::env::var("SEMANTEX_COREML").is_ok_and(|v| v == "1");
Ok(Self {
model_dir: model_dir.to_path_buf(),
threads: threads.max(1),
use_coreml,
encoder: OnceLock::new(),
build_lock: std::sync::Mutex::new(()),
build_count: std::sync::atomic::AtomicUsize::new(0),
})
}
fn build_encoder(&self) -> Result<Colbert> {
self.build_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let batch_size = crate::config::env_usize("SEMANTEX_ORT_BATCH", DEFAULT_ORT_BATCH);
#[allow(unused_mut)]
let mut builder = Colbert::builder(&self.model_dir)
.with_quantized(true)
.with_threads(self.threads)
.with_batch_size(batch_size);
#[cfg(target_os = "macos")]
{
let provider = if self.use_coreml {
next_plaid_onnx::ExecutionProvider::CoreML
} else {
next_plaid_onnx::ExecutionProvider::Cpu
};
builder = builder.with_execution_provider(provider);
}
#[cfg(not(target_os = "macos"))]
{
let _ = self.use_coreml;
builder = builder.with_execution_provider(next_plaid_onnx::ExecutionProvider::Cpu);
}
builder.build()
}
fn encoder(&self) -> Result<&Mutex<Colbert>> {
if let Some(enc) = self.encoder.get() {
return Ok(enc);
}
let _guard = self
.build_lock
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(enc) = self.encoder.get() {
return Ok(enc);
}
let built = self.build_encoder()?;
let _ = self.encoder.set(Mutex::new(built));
Ok(self
.encoder
.get()
.expect("encoder set in previous statement"))
}
pub fn warm_up(&self) -> Result<()> {
let enc = self.encoder()?;
let _ = enc.lock().encode_queries(&["warmup"])?;
Ok(())
}
pub fn global(model_dir: &Path) -> Result<&'static ColbertEmbedder> {
if let Some(embedder) = GLOBAL_COLBERT.get() {
return Ok(embedder);
}
let _guard = COLBERT_INIT_LOCK.lock();
if let Some(embedder) = GLOBAL_COLBERT.get() {
return Ok(embedder);
}
tracing::info!("Initializing global ColBERT encoder singleton (lazy)");
let embedder = Self::new(model_dir)?;
let _ = GLOBAL_COLBERT.set(embedder);
Ok(GLOBAL_COLBERT.get().expect("just set"))
}
pub fn encode_query(&self, text: &str) -> Result<TokenEmbeddings> {
let encoder = self.encoder()?;
let mut embeddings = encoder.lock().encode_queries(&[text])?;
embeddings
.pop()
.ok_or_else(|| anyhow::anyhow!("encode_queries returned empty result"))
}
pub fn encode_documents(&self, texts: &[String]) -> Result<Vec<TokenEmbeddings>> {
let encoder = self.encoder()?;
let refs: Vec<&str> = texts.iter().map(String::as_str).collect();
let embeddings = encoder.lock().encode_documents(&refs, None)?;
Ok(embeddings)
}
#[doc(hidden)]
pub fn is_initialized(&self) -> bool {
self.encoder.get().is_some()
}
#[doc(hidden)]
pub fn build_count(&self) -> usize {
self.build_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_index_threads_uses_all_cores_when_only_one_build_slot_exists() {
assert_eq!(default_index_threads(4, 1), 4);
assert_eq!(default_index_threads(2, 1), 2, "clamped to floor of 2");
assert_eq!(default_index_threads(16, 1), 8, "clamped to ceiling of 8");
assert_eq!(default_index_threads(1, 1), 2, "clamped to floor of 2");
}
#[test]
fn default_index_threads_halves_cores_when_multiple_build_slots_exist() {
assert_eq!(default_index_threads(8, 2), 4);
assert_eq!(default_index_threads(32, 4), 8, "clamped to ceiling of 8");
assert_eq!(default_index_threads(4, 2), 2, "clamped to floor of 2");
}
#[test]
fn new_is_lazy_does_not_build_session() {
let tmp = tempfile::TempDir::new().unwrap();
let embedder = ColbertEmbedder::new(tmp.path())
.expect("constructor should succeed for any existing directory");
assert!(
!embedder.is_initialized(),
"ONNX session must not be materialized at construction time"
);
}
#[test]
fn new_rejects_missing_model_dir() {
let res = ColbertEmbedder::new(Path::new("/nonexistent/path/that/does/not/exist"));
assert!(res.is_err(), "missing model dir must fail at construction");
}
#[test]
fn encoder_init_is_serialized_under_concurrency() {
use std::sync::Arc;
use std::sync::Barrier;
use std::sync::atomic::{AtomicUsize, Ordering};
struct TestEmbedder {
encoder: OnceLock<u32>,
build_lock: std::sync::Mutex<()>,
build_count: AtomicUsize,
}
impl TestEmbedder {
fn new() -> Self {
Self {
encoder: OnceLock::new(),
build_lock: std::sync::Mutex::new(()),
build_count: AtomicUsize::new(0),
}
}
fn get_or_init(&self) -> &u32 {
if let Some(v) = self.encoder.get() {
return v;
}
let _guard = self
.build_lock
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(v) = self.encoder.get() {
return v;
}
self.build_count.fetch_add(1, Ordering::Relaxed);
std::thread::sleep(std::time::Duration::from_millis(50));
let _ = self.encoder.set(42);
self.encoder.get().expect("encoder set above")
}
}
let embedder = Arc::new(TestEmbedder::new());
let n_threads = 8;
let barrier = Arc::new(Barrier::new(n_threads));
let mut handles = Vec::with_capacity(n_threads);
for _ in 0..n_threads {
let emb = Arc::clone(&embedder);
let b = Arc::clone(&barrier);
handles.push(std::thread::spawn(move || {
b.wait();
let v = emb.get_or_init();
assert_eq!(*v, 42);
}));
}
for h in handles {
h.join().unwrap();
}
let count = embedder.build_count.load(Ordering::Relaxed);
assert_eq!(
count, 1,
"build path must be invoked exactly once under {n_threads} concurrent \
first-callers (observed {count}). If this fails, the check-build-set \
pattern in `ColbertEmbedder::encoder` has regressed — see Finding 10."
);
}
#[test]
fn build_count_does_not_grow_on_cached_reads() {
let tmp = tempfile::TempDir::new().unwrap();
let embedder = ColbertEmbedder::new(tmp.path()).unwrap();
assert_eq!(embedder.build_count(), 0);
assert!(!embedder.is_initialized());
}
}