#![allow(clippy::expect_used, clippy::unwrap_used)]
use std::time::Duration;
use graph_storage_sdk::models::RemainingBudget;
use graph_storage_sdk::plugin_api::{EmbedRequest, EmbeddingProviderV1};
use onnx_embedding_plugin::{OnnxEmbeddingProvider, OnnxProviderConfig};
use tokio_util::sync::CancellationToken;
async fn provider() -> Option<OnnxEmbeddingProvider> {
let required = std::env::var("GRAPH_STORAGE_ONNX_REQUIRED").is_ok();
let missing = |what: &str| {
assert!(
!required,
"GRAPH_STORAGE_ONNX_REQUIRED is set but {what} is not available"
);
eprintln!("skipping the ONNX lane: {what} is not set");
None::<OnnxEmbeddingProvider>
};
if std::env::var("ORT_DYLIB_PATH").is_err() {
return missing("ORT_DYLIB_PATH");
}
let Ok(model) = std::env::var("GRAPH_STORAGE_ONNX_MODEL") else {
return missing("GRAPH_STORAGE_ONNX_MODEL");
};
let Ok(tokenizer) = std::env::var("GRAPH_STORAGE_ONNX_TOKENIZER") else {
return missing("GRAPH_STORAGE_ONNX_TOKENIZER");
};
match OnnxEmbeddingProvider::load(OnnxProviderConfig::new(model, tokenizer)).await {
Ok(provider) => Some(provider),
Err(error) => {
assert!(
!required,
"GRAPH_STORAGE_ONNX_REQUIRED is set but the provider did not load: {error}"
);
eprintln!("skipping the ONNX lane: {error}");
None
}
}
}
async fn provider_with(
adjust: impl FnOnce(&mut OnnxProviderConfig),
) -> Option<OnnxEmbeddingProvider> {
let model = std::env::var("GRAPH_STORAGE_ONNX_MODEL").ok()?;
let tokenizer = std::env::var("GRAPH_STORAGE_ONNX_TOKENIZER").ok()?;
let mut config = OnnxProviderConfig::new(model, tokenizer);
adjust(&mut config);
OnnxEmbeddingProvider::load(config).await.ok()
}
async fn embed(provider: &OnnxEmbeddingProvider, texts: &[&str]) -> Vec<Vec<f32>> {
provider
.embed(EmbedRequest {
inputs: texts.iter().map(|t| (*t).to_owned()).collect(),
budget: RemainingBudget::starting_now(Duration::from_mins(1)),
cancel: CancellationToken::new(),
})
.await
.expect("the model embeds")
.vectors
}
fn cosine(one: &[f32], other: &[f32]) -> f64 {
one.iter()
.zip(other)
.map(|(a, b)| f64::from(*a) * f64::from(*b))
.sum()
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn readiness_reports_a_session_that_has_been_proven_to_run() {
let Some(provider) = provider().await else {
return;
};
provider
.health()
.await
.expect("a provider that loaded has run its probe");
embed(&provider, &["something to embed"]).await;
provider
.health()
.await
.expect("a session that just produced a vector is healthy");
}
#[tokio::test]
async fn a_session_whose_width_contradicts_the_configuration_does_not_load() {
if std::env::var("ORT_DYLIB_PATH").is_err() {
return;
}
let (Ok(model), Ok(tokenizer)) = (
std::env::var("GRAPH_STORAGE_ONNX_MODEL"),
std::env::var("GRAPH_STORAGE_ONNX_TOKENIZER"),
) else {
return;
};
let mut config = OnnxProviderConfig::new(model, tokenizer);
config.dimension += 1;
let error = OnnxEmbeddingProvider::load(config)
.await
.err()
.expect("a width the model cannot produce is a load failure");
assert!(
error.to_string().contains("embedding space mismatch"),
"the refusal names the width disagreement: {error}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn the_onnx_provider_honours_the_contract_on_a_threaded_runtime() {
let Some(provider) = provider().await else {
return;
};
graph_storage_sdk::contract::assert_embedding_provider(&provider).await;
let batch = |texts: &[&str]| EmbedRequest {
inputs: texts.iter().map(|t| (*t).to_owned()).collect(),
budget: RemainingBudget::starting_now(Duration::from_mins(1)),
cancel: CancellationToken::new(),
};
let (first, second) = tokio::join!(
provider.embed(batch(&["one sentence", "another"])),
provider.embed(batch(&["a third"])),
);
assert_eq!(first.expect("the first batch embeds").vectors.len(), 2);
assert_eq!(second.expect("the second batch embeds").vectors.len(), 1);
}
#[tokio::test]
async fn the_onnx_provider_honours_the_contract() {
let Some(provider) = provider().await else {
return;
};
graph_storage_sdk::contract::assert_embedding_provider(&provider).await;
}
#[tokio::test]
async fn related_sentences_sit_closer_than_unrelated_ones() {
let Some(provider) = provider().await else {
return;
};
let vectors = embed(
&provider,
&[
"A hardcoded password was committed to the deployment script.",
"Someone checked a plaintext credential into the deploy config.",
"The kitchen renovation is scheduled for next spring.",
],
)
.await;
let related = cosine(&vectors[0], &vectors[1]);
let unrelated = cosine(&vectors[0], &vectors[2]);
assert!(
related > 0.5,
"paraphrases should be plainly similar, scored {related}"
);
assert!(
unrelated < 0.3,
"unrelated sentences scored {unrelated}; a high floor here means the \
pooling is averaging in padding rather than tokens"
);
}
#[tokio::test]
async fn one_text_embeds_to_one_vector() {
let Some(provider) = provider().await else {
return;
};
let text = "Hardcoded credential in deploy script";
let first = embed(&provider, &[text]).await;
let second = embed(&provider, &[text]).await;
assert_eq!(first, second);
}
#[tokio::test]
async fn a_vector_does_not_depend_on_how_much_padding_follows_it() {
let Some(provider) = provider().await else {
return;
};
let Some(tight) = provider_with(|config| config.max_tokens = 16).await else {
return;
};
let text = "A short sentence.";
let roomy = embed(&provider, &[text]).await;
let cramped = embed(&tight, &[text]).await;
let agreement = cosine(&roomy[0], &cramped[0]);
assert!(
(agreement - 1.0).abs() < 1e-4,
"the same text embedded differently under a different amount of \
padding: cosine {agreement}"
);
}