use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use aws_lc_rs::digest::{SHA256, digest as sha256};
use graph_storage_sdk::models::EmbeddingSpaceId;
use graph_storage_sdk::plugin_api::{
EmbedRequest, EmbedResponse, EmbeddingProviderError, EmbeddingProviderV1,
};
use ort::session::Session;
use ort::session::builder::GraphOptimizationLevel;
use ort::value::Tensor;
use thiserror::Error;
use tokenizers::Tokenizer;
use tokio::sync::Mutex;
use tracing::warn;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum Pooling {
#[default]
Mean,
Cls,
}
impl Pooling {
fn as_str(self) -> &'static str {
match self {
Self::Mean => "mean",
Self::Cls => "cls",
}
}
}
#[derive(Clone, Debug)]
pub struct OnnxProviderConfig {
pub model_path: PathBuf,
pub tokenizer_path: PathBuf,
pub dimension: u32,
pub pooling: Pooling,
pub normalize: bool,
pub max_tokens: usize,
pub intra_op_threads: Option<usize>,
}
impl OnnxProviderConfig {
#[must_use]
pub fn new(model_path: impl Into<PathBuf>, tokenizer_path: impl Into<PathBuf>) -> Self {
Self {
model_path: model_path.into(),
tokenizer_path: tokenizer_path.into(),
dimension: 384,
pooling: Pooling::Mean,
normalize: true,
max_tokens: 256,
intra_op_threads: None,
}
}
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum OnnxLoadError {
#[error("cannot read {what} at {path}: {source}")]
Artifact {
what: &'static str,
path: PathBuf,
source: std::io::Error,
},
#[error("tokenizer at {path} is not loadable: {reason}")]
Tokenizer { path: PathBuf, reason: String },
#[error("ONNX session could not be created: {0}")]
Session(String),
#[error(
"ONNX Runtime did not load within {seconds}s. `ort` 2.0.0-rc.12 hangs \
instead of erroring on an unloadable library, so the thread that \
tried is abandoned rather than killed: check ORT_DYLIB_PATH and \
restart the process"
)]
RuntimeHung { seconds: u64 },
}
const LOAD_TIMEOUT: Duration = Duration::from_secs(30);
const PROBE_TIMEOUT: Duration = Duration::from_secs(10);
pub struct OnnxEmbeddingProvider {
session: Arc<Mutex<Session>>,
tokenizer: Tokenizer,
space: EmbeddingSpaceId,
config: OnnxProviderConfig,
observed: std::sync::Mutex<Option<String>>,
}
impl OnnxEmbeddingProvider {
pub async fn load(config: OnnxProviderConfig) -> Result<Self, OnnxLoadError> {
let model_digest = file_digest("model", &config.model_path)?;
let tokenizer_digest = file_digest("tokenizer", &config.tokenizer_path)?;
let tokenizer = Tokenizer::from_file(&config.tokenizer_path).map_err(|error| {
OnnxLoadError::Tokenizer {
path: config.tokenizer_path.clone(),
reason: error.to_string(),
}
})?;
let session = open_session(&config).await?;
let space = EmbeddingSpaceId::new(
artifact_name(&config.model_path, &model_digest),
artifact_name(&config.tokenizer_path, &tokenizer_digest),
serde_json::json!({
"tokenizer": "file",
"max_tokens": config.max_tokens,
"truncation": "longest_first",
}),
serde_json::json!({ "strategy": config.pooling.as_str() }),
serde_json::json!({ "l2": config.normalize }),
config.dimension,
);
let provider = Arc::new(Self {
session: Arc::new(Mutex::new(session)),
tokenizer,
space,
config,
observed: std::sync::Mutex::new(None),
});
Self::probe(Arc::clone(&provider)).await?;
Ok(Arc::try_unwrap(provider).unwrap_or_else(|_| {
unreachable!("the probe task is the only other holder and it has finished")
}))
}
}
fn artifact_name(path: &Path, digest: &str) -> String {
let name = path.file_name().map_or_else(
|| "unnamed".to_owned(),
|n| n.to_string_lossy().into_owned(),
);
format!("{name}@sha256:{digest}")
}
fn file_digest(what: &'static str, path: &Path) -> Result<String, OnnxLoadError> {
let bytes = std::fs::read(path).map_err(|source| OnnxLoadError::Artifact {
what,
path: path.to_path_buf(),
source,
})?;
Ok(hex::encode(sha256(&SHA256, &bytes)))
}
async fn open_session(config: &OnnxProviderConfig) -> Result<Session, OnnxLoadError> {
let (tx, rx) = tokio::sync::oneshot::channel();
let model_path = config.model_path.clone();
let threads = config.intra_op_threads;
std::thread::Builder::new()
.name("graph-storage-onnx-init".to_owned())
.spawn(move || {
let result = build_session(&model_path, threads);
drop(tx.send(result));
})
.map_err(|error| OnnxLoadError::Session(error.to_string()))?;
match tokio::time::timeout(LOAD_TIMEOUT, rx).await {
Ok(Ok(result)) => result,
Ok(Err(_)) => Err(OnnxLoadError::Session(
"the ONNX init thread ended without a result".to_owned(),
)),
Err(_) => {
warn!(
path = %config.model_path.display(),
"ONNX Runtime did not load in time; leaking the init thread deliberately"
);
Err(OnnxLoadError::RuntimeHung {
seconds: LOAD_TIMEOUT.as_secs(),
})
}
}
}
fn build_session(
model_path: &Path,
intra_op_threads: Option<usize>,
) -> Result<Session, OnnxLoadError> {
let mut builder = Session::builder()
.map_err(|error| OnnxLoadError::Session(error.to_string()))?
.with_optimization_level(GraphOptimizationLevel::Level3)
.map_err(|error| OnnxLoadError::Session(error.to_string()))?;
if let Some(threads) = intra_op_threads {
builder = builder
.with_intra_threads(threads)
.map_err(|error| OnnxLoadError::Session(error.to_string()))?;
}
builder
.commit_from_file(model_path)
.map_err(|error| OnnxLoadError::Session(error.to_string()))
}
#[async_trait]
impl EmbeddingProviderV1 for OnnxEmbeddingProvider {
fn embedding_space(&self) -> &EmbeddingSpaceId {
&self.space
}
fn dimension(&self) -> u32 {
self.config.dimension
}
async fn embed(&self, req: EmbedRequest) -> Result<EmbedResponse, EmbeddingProviderError> {
if req.cancel.is_cancelled() {
return Err(EmbeddingProviderError::Cancelled);
}
if req.budget.is_exhausted() {
return Err(EmbeddingProviderError::Deadline);
}
if req.inputs.is_empty() {
return Ok(EmbedResponse {
vectors: Vec::new(),
space: self.space.clone(),
});
}
let encoded = self.encode(&req.inputs)?;
let mut session = self.session.lock().await;
if req.budget.is_exhausted() {
return Err(EmbeddingProviderError::Deadline);
}
if req.cancel.is_cancelled() {
return Err(EmbeddingProviderError::Cancelled);
}
let outcome = if tokio::runtime::Handle::try_current().is_ok_and(|handle| {
handle.runtime_flavor() == tokio::runtime::RuntimeFlavor::MultiThread
}) {
tokio::task::block_in_place(|| self.run(&mut session, &encoded))
} else {
self.run(&mut session, &encoded)
};
self.record(match &outcome {
Ok(_) => None,
Err(error) => Some(error.to_string()),
});
let vectors = outcome?;
Ok(EmbedResponse {
vectors,
space: self.space.clone(),
})
}
async fn health(&self) -> Result<(), EmbeddingProviderError> {
drop(self.session.lock().await);
match self.failure() {
Some(reason) => Err(EmbeddingProviderError::Unavailable { reason }),
None => Ok(()),
}
}
}
impl OnnxEmbeddingProvider {
async fn probe(provider: Arc<Self>) -> Result<(), OnnxLoadError> {
let inference = tokio::task::spawn_blocking(move || {
let encoded = provider
.encode(std::slice::from_ref(&PROBE_INPUT.to_owned()))
.map_err(|error| OnnxLoadError::Session(format!("probe tokenization: {error}")))?;
let mut session = provider.session.blocking_lock();
provider
.run(&mut session, &encoded)
.map_err(|error| OnnxLoadError::Session(format!("probe inference: {error}")))?;
Ok::<(), OnnxLoadError>(())
});
match tokio::time::timeout(PROBE_TIMEOUT, inference).await {
Ok(Ok(result)) => result,
Ok(Err(join)) => Err(OnnxLoadError::Session(format!(
"the ONNX probe thread ended without a result: {join}"
))),
Err(_) => {
warn!(
seconds = PROBE_TIMEOUT.as_secs(),
"the ONNX session loaded but did not answer one inference in time"
);
Err(OnnxLoadError::RuntimeHung {
seconds: PROBE_TIMEOUT.as_secs(),
})
}
}
}
fn failure(&self) -> Option<String> {
self.observed.lock().map_or(
Some("the observation lock is poisoned".to_owned()),
|seen| seen.clone(),
)
}
fn record(&self, failure: Option<String>) {
if let Ok(mut seen) = self.observed.lock() {
*seen = failure;
}
}
}
const PROBE_INPUT: &str = "graph storage readiness probe";
struct Encoded {
ids: Vec<i64>,
mask: Vec<i64>,
type_ids: Vec<i64>,
rows: usize,
columns: usize,
}
impl OnnxEmbeddingProvider {
fn encode(&self, inputs: &[String]) -> Result<Encoded, EmbeddingProviderError> {
let encodings = self
.tokenizer
.encode_batch(inputs.to_vec(), true)
.map_err(|error| EmbeddingProviderError::Internal(error.to_string()))?;
let columns = encodings
.iter()
.map(|e| e.get_ids().len().min(self.config.max_tokens))
.max()
.unwrap_or(1)
.max(1);
let rows = encodings.len();
let mut ids = vec![0_i64; rows * columns];
let mut mask = vec![0_i64; rows * columns];
let type_ids = vec![0_i64; rows * columns];
for (row, encoding) in encodings.iter().enumerate() {
let take = encoding.get_ids().len().min(columns);
for column in 0..take {
ids[row * columns + column] = i64::from(encoding.get_ids()[column]);
mask[row * columns + column] = i64::from(encoding.get_attention_mask()[column]);
}
}
Ok(Encoded {
ids,
mask,
type_ids,
rows,
columns,
})
}
fn run(
&self,
session: &mut Session,
encoded: &Encoded,
) -> Result<Vec<Vec<f32>>, EmbeddingProviderError> {
let internal = |what: String| EmbeddingProviderError::Internal(what);
let shape = [encoded.rows, encoded.columns];
let tensor = |data: &[i64]| {
Tensor::from_array((shape, data.to_vec().into_boxed_slice()))
.map_err(|error| internal(error.to_string()))
};
let outputs = session
.run(ort::inputs![
"input_ids" => tensor(&encoded.ids)?,
"attention_mask" => tensor(&encoded.mask)?,
"token_type_ids" => tensor(&encoded.type_ids)?,
])
.map_err(|error| EmbeddingProviderError::Unavailable {
reason: error.to_string(),
})?;
let (_, first) = outputs
.iter()
.next()
.ok_or_else(|| internal("the model produced no output".to_owned()))?;
let (out_shape, values) = first
.try_extract_tensor::<f32>()
.map_err(|error| internal(error.to_string()))?;
if out_shape.len() != 3 {
return Err(internal(format!(
"expected a [batch, tokens, hidden] output, got {out_shape:?}"
)));
}
let dim = |axis: usize| {
out_shape
.get(axis)
.copied()
.and_then(|value| usize::try_from(value).ok())
.ok_or_else(|| internal(format!("the model reported a shape of {out_shape:?}")))
};
let (batch, tokens, hidden) = (dim(0)?, dim(1)?, dim(2)?);
if batch != encoded.rows || tokens != encoded.columns {
return Err(internal(format!(
"the model answered a [{batch}, {tokens}, _] batch for a [{}, {}, _] input",
encoded.rows, encoded.columns
)));
}
if hidden != self.config.dimension as usize {
warn!(
got = hidden,
want = self.config.dimension,
model = %self.config.model_path.display(),
"the model's output width is not the configured embedding dimension"
);
return Err(EmbeddingProviderError::SpaceMismatch);
}
Ok(self.pool(values, encoded, hidden))
}
fn pool(&self, values: &[f32], encoded: &Encoded, hidden: usize) -> Vec<Vec<f32>> {
let mut out = Vec::with_capacity(encoded.rows);
for row in 0..encoded.rows {
let base = row * encoded.columns * hidden;
let mut pooled = vec![0.0_f64; hidden];
match self.config.pooling {
Pooling::Cls => {
for (lane, slot) in pooled.iter_mut().enumerate() {
*slot = f64::from(values[base + lane]);
}
}
Pooling::Mean => {
let mut kept = 0.0_f64;
for column in 0..encoded.columns {
if encoded.mask[row * encoded.columns + column] == 0 {
continue;
}
kept += 1.0;
let offset = base + column * hidden;
for (lane, slot) in pooled.iter_mut().enumerate() {
*slot += f64::from(values[offset + lane]);
}
}
if kept > 0.0 {
for slot in &mut pooled {
*slot /= kept;
}
}
}
}
if self.config.normalize {
let norm = pooled.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm > 0.0 {
for slot in &mut pooled {
*slot /= norm;
}
}
}
out.push(narrow(&pooled));
}
out
}
}
#[expect(
clippy::cast_possible_truncation,
reason = "see the function's own documentation"
)]
fn narrow(values: &[f64]) -> Vec<f32> {
values.iter().map(|v| *v as f32).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pooling_is_part_of_the_identity() {
let one = EmbeddingSpaceId::new(
"m@sha256:a",
"t@sha256:b",
serde_json::json!({}),
serde_json::json!({ "strategy": Pooling::Mean.as_str() }),
serde_json::json!({ "l2": true }),
384,
);
let other = EmbeddingSpaceId::new(
"m@sha256:a",
"t@sha256:b",
serde_json::json!({}),
serde_json::json!({ "strategy": Pooling::Cls.as_str() }),
serde_json::json!({ "l2": true }),
384,
);
assert_ne!(one.identity_hash, other.identity_hash);
}
#[test]
fn a_missing_artifact_is_named_in_the_error() {
let error = file_digest("model", Path::new("/nonexistent/model.onnx"))
.expect_err("a missing file cannot be digested");
let rendered = error.to_string();
assert!(rendered.contains("model"), "{rendered}");
assert!(rendered.contains("/nonexistent/model.onnx"), "{rendered}");
}
#[test]
fn the_artifact_name_carries_the_digest_rather_than_the_path() {
assert_eq!(
artifact_name(Path::new("/opt/models/model.onnx"), "abc"),
artifact_name(Path::new("/srv/other/model.onnx"), "abc")
);
}
}