use std::sync::LazyLock;
use serde::{Deserialize, Serialize};
#[cfg(feature = "sparse-embeddings")]
pub mod engine;
#[cfg(feature = "sparse-embeddings")]
use std::sync::{Arc, RwLock};
#[cfg(feature = "sparse-embeddings")]
use ahash::AHashMap;
#[cfg(feature = "sparse-embeddings")]
use engine::SparseEmbeddingEngine;
#[cfg(feature = "sparse-embeddings")]
const DEFAULT_MODEL_FILE: &str = "onnx/model.onnx";
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
pub struct SparseEmbedding {
pub indices: Vec<u32>,
pub values: Vec<f32>,
}
#[cfg(all(test, feature = "sparse-embeddings"))]
mod engine_cache_key_tests {
use super::*;
use crate::core::config::acceleration::{AccelerationConfig, ExecutionProviderType};
fn key(
additional_files: &[String],
max_length: usize,
acceleration: &AccelerationConfig,
) -> SparseEmbeddingEngineCacheKey {
SparseEmbeddingEngineCacheKey::new(
"owner/model",
"model.onnx",
additional_files,
"revision",
max_length,
"cache-root".to_string(),
crate::onnx::OnnxAccelerationCacheKey::from_resolved(acceleration.provider.clone(), acceleration.device_id),
)
}
#[test]
fn engine_cache_reuses_equal_configs_and_isolates_distinct_configs() {
let cpu = AccelerationConfig {
provider: ExecutionProviderType::Cpu,
device_id: 0,
};
let cuda = AccelerationConfig {
provider: ExecutionProviderType::Cuda,
device_id: 1,
};
let files = vec!["config.json".to_string(), "weights.onnx.data".to_string()];
let reversed_files = vec!["weights.onnx.data".to_string(), "config.json".to_string()];
let original = key(&files, 512, &cpu);
let equal = key(&files, 512, &cpu);
let mut cache = AHashMap::new();
cache.insert(original, 7_u8);
assert_eq!(cache.get(&equal), Some(&7));
assert_eq!(cache.get(&key(&files, 1024, &cpu)), None);
assert_eq!(cache.get(&key(&files, 512, &cuda)), None);
assert_eq!(cache.get(&key(&[], 512, &cpu)), None);
assert_eq!(cache.get(&key(&reversed_files, 512, &cpu)), None);
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[cfg_attr(feature = "api", derive(utoipa::ToSchema))]
pub struct SparseEmbeddingPreset {
pub name: String,
pub model_repo: String,
pub model_file: String,
pub additional_files: Vec<String>,
pub max_length: usize,
pub description: String,
}
#[cfg(any(feature = "sparse-embeddings", test))]
pub(crate) const SPARSE_EMBEDDING_SHA256_MANIFEST: &str = include_str!("presets.sha256sum");
#[cfg(feature = "sparse-embeddings")]
const SPARSE_EMBEDDING_REVISION: &str = "5b5bccdae54b24d5a117a9e621775ebf97596d20";
pub static SPARSE_EMBEDDING_PRESETS: LazyLock<Vec<SparseEmbeddingPreset>> = LazyLock::new(|| {
vec![
SparseEmbeddingPreset {
name: "splade".to_string(),
model_repo: "xberg-io/sparse-embeddings".to_string(),
model_file: "splade/model.onnx".to_string(),
additional_files: Vec::new(),
max_length: 256,
description: "SPLADE++ EN v1 — English learned sparse retrieval (Apache-2.0).".to_string(),
},
SparseEmbeddingPreset {
name: "opensearch-v3-distill".to_string(),
model_repo: "xberg-io/sparse-embeddings".to_string(),
model_file: "opensearch-v3-distill/model.onnx".to_string(),
additional_files: vec!["opensearch-v3-distill/model.onnx.data".to_string()],
max_length: 512,
description: "OpenSearch neural-sparse v3 distill (DistilBERT MLM, 2026-gen, Apache-2.0). \
Exported with its MLM head; 30522-dim SPLADE sparse vectors, 512 max-len."
.to_string(),
},
]
});
#[cfg(any(feature = "sparse-embedding-presets", feature = "sparse-embeddings"))]
#[cfg_attr(alef, alef(skip))]
pub fn get_preset(name: &str) -> Option<SparseEmbeddingPreset> {
SPARSE_EMBEDDING_PRESETS.iter().find(|p| p.name == name).cloned()
}
#[cfg(any(feature = "sparse-embedding-presets", feature = "sparse-embeddings"))]
#[cfg_attr(alef, alef(skip))]
pub fn list_presets() -> Vec<String> {
SPARSE_EMBEDDING_PRESETS.iter().map(|p| p.name.clone()).collect()
}
#[cfg(feature = "sparse-embeddings")]
type CachedEngine = Arc<SparseEmbeddingEngine>;
#[cfg(feature = "sparse-embeddings")]
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct SparseEmbeddingEngineCacheKey {
repo_name: String,
model_file: String,
additional_files: Vec<String>,
revision: String,
max_length: usize,
cache_root: String,
acceleration: crate::onnx::OnnxAccelerationCacheKey,
}
#[cfg(feature = "sparse-embeddings")]
impl SparseEmbeddingEngineCacheKey {
fn new(
repo_name: &str,
model_file: &str,
additional_files: &[String],
revision: &str,
max_length: usize,
cache_root: String,
acceleration: crate::onnx::OnnxAccelerationCacheKey,
) -> Self {
Self {
repo_name: repo_name.to_string(),
model_file: model_file.to_string(),
additional_files: additional_files.to_vec(),
revision: revision.to_string(),
max_length,
cache_root,
acceleration,
}
}
}
#[cfg(feature = "sparse-embeddings")]
static ENGINE_CACHE: LazyLock<RwLock<AHashMap<SparseEmbeddingEngineCacheKey, CachedEngine>>> =
LazyLock::new(|| RwLock::new(AHashMap::new()));
#[cfg(all(feature = "sparse-embeddings", feature = "tokio-runtime"))]
static SPARSE_SEMAPHORE: LazyLock<Arc<tokio::sync::Semaphore>> = LazyLock::new(|| {
let permits = crate::core::config::concurrency::resolve_batch_concurrency(None, true).max(1);
Arc::new(tokio::sync::Semaphore::new(permits))
});
#[cfg(feature = "sparse-embeddings")]
fn sparse_err(msg: String) -> crate::XbergError {
crate::XbergError::embedding(msg)
}
#[cfg(feature = "sparse-embeddings")]
fn resolve_model_info(
model_type: &crate::core::config::SparseEmbeddingModelType,
config_max_length: usize,
) -> crate::Result<(String, String, Vec<String>, usize)> {
use crate::core::config::SparseEmbeddingModelType as M;
match model_type {
M::Preset { name } => {
let preset =
get_preset(name).ok_or_else(|| sparse_err(format!("Unknown sparse-embedding preset: {name}")))?;
Ok((
preset.model_repo,
preset.model_file,
preset.additional_files,
preset.max_length,
))
}
M::Custom {
model_id,
model_file,
additional_files,
max_length,
} => {
let file = model_file.clone().unwrap_or_else(|| DEFAULT_MODEL_FILE.to_string());
let max_len = match max_length {
Some(v) if *v > 0 => *v as usize,
_ => config_max_length,
};
Ok((model_id.clone(), file, additional_files.clone(), max_len))
}
M::Plugin { .. } => Err(sparse_err(
"Plugin sparse-embedding backends are not yet supported; use Preset or Custom".to_string(),
)),
}
}
#[cfg(feature = "sparse-embeddings")]
fn get_or_init_engine(
repo_name: &str,
model_file: &str,
additional_files: &[String],
max_length: usize,
cache_dir: Option<std::path::PathBuf>,
progress: crate::core::config::DownloadProgress,
accel: Option<crate::core::config::acceleration::AccelerationConfig>,
) -> crate::Result<Arc<SparseEmbeddingEngine>> {
let revision = (repo_name == "xberg-io/sparse-embeddings").then_some(SPARSE_EMBEDDING_REVISION);
let engine_key = SparseEmbeddingEngineCacheKey::new(
repo_name,
model_file,
additional_files,
revision.unwrap_or("main"),
max_length,
crate::model_download::hf_cache_key(cache_dir.as_deref()),
crate::onnx::OnnxAccelerationCacheKey::new(accel.as_ref()),
);
match ENGINE_CACHE.read() {
Ok(cache) => {
if let Some(cached) = cache.get(&engine_key) {
return Ok(Arc::clone(cached));
}
}
Err(poison) => {
if let Some(cached) = poison.get_ref().get(&engine_key) {
return Ok(Arc::clone(cached));
}
}
}
let mut cache = match ENGINE_CACHE.write() {
Ok(guard) => guard,
Err(poison) => poison.into_inner(),
};
if let Some(cached) = cache.get(&engine_key) {
return Ok(Arc::clone(cached));
}
crate::ort_discovery::ensure_ort_available();
let files = crate::onnx::download_model_files(
repo_name,
model_file,
additional_files,
revision,
cache_dir.as_deref(),
progress,
Some(SPARSE_EMBEDDING_SHA256_MANIFEST),
sparse_err,
)?;
let tokenizer = crate::onnx::load_tokenizer(&files, max_length, sparse_err)?;
let session = crate::onnx::build_session(&files.model, accel.as_ref(), sparse_err)?;
let engine = Arc::new(SparseEmbeddingEngine::new(tokenizer, session));
cache.insert(engine_key, Arc::clone(&engine));
Ok(engine)
}
#[cfg(feature = "sparse-embeddings")]
fn map_engine_err(e: engine::SparseEmbedError) -> crate::XbergError {
use engine::SparseEmbedError as E;
match e {
E::Tokenizer(m) => sparse_err(format!("Tokenization failed: {m}")),
E::Ort(err) => {
let msg = err.to_string();
if crate::onnx::looks_like_ort_error(&msg) {
crate::XbergError::MissingDependency(format!(
"ONNX Runtime - {}",
crate::onnx::onnx_runtime_install_message()
))
} else {
sparse_err(format!("Sparse-embedding inference failed: {err}"))
}
}
E::Shape(m) => sparse_err(format!("Unexpected model output shape: {m}")),
E::NoOutput => sparse_err("Sparse-embedding model produced no output".to_string()),
}
}
#[cfg_attr(alef, alef(skip))]
#[cfg(feature = "sparse-embeddings")]
pub fn embed_sparse<T: AsRef<str>>(
texts: &[T],
config: &crate::core::config::SparseEmbeddingConfig,
) -> crate::Result<Vec<SparseEmbedding>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let (repo, model_file, additional, max_len) = resolve_model_info(&config.model, config.max_length)?;
let engine = get_or_init_engine(
&repo,
&model_file,
&additional,
max_len,
config.cache_dir.clone(),
config.into(),
config.acceleration.clone(),
)?;
engine.embed(texts, config.batch_size).map_err(map_engine_err)
}
#[cfg(all(feature = "sparse-embeddings", feature = "tokio-runtime"))]
#[cfg_attr(alef, alef(skip))]
pub async fn embed_sparse_async(
texts: Vec<String>,
config: &crate::core::config::SparseEmbeddingConfig,
) -> crate::Result<Vec<SparseEmbedding>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let config = config.clone();
let permit = SPARSE_SEMAPHORE
.clone()
.acquire_owned()
.await
.map_err(|e| sparse_err(format!("Sparse-embedding semaphore closed: {e}")))?;
tokio::task::spawn_blocking(move || {
let _permit = permit;
embed_sparse(&texts, &config)
})
.await
.map_err(|e| sparse_err(format!("Sparse-embedding task failed: {e}")))?
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn every_preset_file_is_pinned_in_manifest() {
let manifest = crate::model_download::parse_sha256_manifest(SPARSE_EMBEDDING_SHA256_MANIFEST).unwrap();
let pinned: std::collections::HashSet<&str> = manifest.iter().map(|(p, _)| p.as_str()).collect();
for preset in SPARSE_EMBEDDING_PRESETS.iter() {
assert!(
pinned.contains(preset.model_file.as_str()),
"preset {} model_file {} is not pinned in presets.sha256sum",
preset.name,
preset.model_file
);
for sibling in &preset.additional_files {
assert!(
pinned.contains(sibling.as_str()),
"preset {} additional file {} is not pinned in presets.sha256sum",
preset.name,
sibling
);
}
let model_dir = std::path::Path::new(&preset.model_file)
.parent()
.and_then(|p| p.to_str())
.filter(|s| !s.is_empty());
let companion_path = |name: &str| match model_dir {
Some(dir) => format!("{dir}/{name}"),
None => name.to_string(),
};
for required in ["tokenizer.json", "config.json"] {
let path = companion_path(required);
assert!(
pinned.contains(path.as_str()),
"preset {} companion {} is not pinned in presets.sha256sum",
preset.name,
path
);
}
}
}
#[test]
fn preset_catalog_is_nonempty_and_lookup_works() {
assert!(!SPARSE_EMBEDDING_PRESETS.is_empty());
assert!(list_presets().contains(&"splade".to_string()));
let p = get_preset("splade").expect("splade preset present");
assert_eq!(p.model_repo, "xberg-io/sparse-embeddings");
assert_eq!(p.model_file, "splade/model.onnx");
assert!(get_preset("does-not-exist").is_none());
}
#[test]
fn sparse_embedding_serde_roundtrip() {
let se = SparseEmbedding {
indices: vec![3, 17, 200],
values: vec![0.5, 0.25, 0.1],
};
let json = serde_json::to_string(&se).unwrap();
let back: SparseEmbedding = serde_json::from_str(&json).unwrap();
assert_eq!(se, back);
}
}