#![allow(clippy::cast_possible_wrap)]
use serde::{Deserialize, Serialize};
use std::time::Duration;
use wm_core::{CoreError, Result};
pub trait Embedder: Send + Sync {
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>>;
fn embed(&self, text: &str) -> Result<Vec<f32>> {
self.embed_batch(&[text])?
.into_iter()
.next()
.ok_or_else(|| CoreError::Memory("embedder returned empty result".into()))
}
fn embed_query(&self, query: &str) -> Result<Vec<f32>> {
self.embed(query)
}
fn dimension(&self) -> usize;
fn is_available(&self) -> bool;
fn backend_name(&self) -> &'static str;
fn cache_namespace(&self) -> String {
self.backend_name().to_string()
}
fn preferred_max_batch_texts(&self) -> usize {
usize::MAX
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbedderConfig {
pub endpoint: String,
pub model: String,
pub dimension: usize,
pub timeout: Duration,
}
impl EmbedderConfig {
#[must_use]
pub fn from_env() -> Option<Self> {
let endpoint = std::env::var("WM_EMBEDDER_ENDPOINT").ok()?;
if !is_endpoint_safe(&endpoint) {
tracing::warn!(
"embedder endpoint rejected by SSRF validation: {}",
endpoint
);
return None;
}
let model = std::env::var("WM_EMBEDDER_MODEL").unwrap_or_else(|_| "local".into());
let dimension = std::env::var("WM_EMBEDDER_DIM")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(384);
let timeout_ms = std::env::var("WM_EMBEDDER_TIMEOUT_MS")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(30_000);
Some(Self {
endpoint,
model,
dimension,
timeout: Duration::from_millis(timeout_ms),
})
}
}
#[must_use]
pub fn is_endpoint_safe(endpoint: &str) -> bool {
if !endpoint.starts_with("http://") && !endpoint.starts_with("https://") {
return false;
}
let without_scheme = endpoint
.strip_prefix("http://")
.or_else(|| endpoint.strip_prefix("https://"))
.unwrap_or(endpoint);
let host_end = without_scheme
.find(['/', '?', '#'])
.unwrap_or(without_scheme.len());
let host_port = &without_scheme[..host_end];
let host = if host_port.starts_with('[') {
if let Some(end) = host_port.find(']') {
&host_port[1..end]
} else {
return false; }
} else {
host_port.rsplit_once(':').map_or(host_port, |(h, _)| h)
};
if host.is_empty() {
return false;
}
let lower = host.to_ascii_lowercase();
if matches!(
lower.as_str(),
"metadata.google.internal"
| "metadata.aws.internal"
| "metadata"
| "169.254.169.254"
| "169.254.170.2"
) {
return false;
}
true
}
pub struct HttpEmbedder {
config: EmbedderConfig,
agent: ureq::Agent,
available: bool,
concurrency: usize,
}
const HTTP_EMBED_CONCURRENCY_DEFAULT: usize = 4;
fn http_concurrency_from_env() -> usize {
std::env::var("WM_EMBEDDER_HTTP_CONCURRENCY")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|n| *n >= 1)
.unwrap_or(HTTP_EMBED_CONCURRENCY_DEFAULT)
}
impl HttpEmbedder {
#[must_use]
pub fn new(config: EmbedderConfig) -> Self {
let agent = ureq::config::Config::builder()
.timeout_global(Some(config.timeout))
.build()
.new_agent();
Self {
config,
agent,
available: true,
concurrency: http_concurrency_from_env(),
}
}
#[must_use]
pub fn with_concurrency(mut self, concurrency: usize) -> Self {
self.concurrency = concurrency.max(1);
self
}
#[must_use]
pub fn from_env() -> Option<Self> {
EmbedderConfig::from_env().map(Self::new)
}
fn embeddings_url(&self) -> String {
if self.config.endpoint.ends_with("/v1/embeddings") {
self.config.endpoint.clone()
} else if self.config.endpoint.ends_with('/') {
format!("{}v1/embeddings", self.config.endpoint)
} else {
format!("{}/v1/embeddings", self.config.endpoint)
}
}
fn embed_chunk(&self, url: &str, prepared: &[&str]) -> Result<Vec<Vec<f32>>> {
let request = EmbeddingsRequest {
model: &self.config.model,
input: prepared,
};
let response = self
.agent
.post(url)
.header("Content-Type", "application/json")
.send_json(&request)
.map_err(|e| CoreError::Memory(format!("Embedder HTTP error: {e}")))?;
let embed_resp: EmbeddingsResponse = response
.into_body()
.read_json()
.map_err(|e| CoreError::Memory(format!("Embedder response parse error: {e}")))?;
let vectors: Vec<Vec<f32>> = embed_resp.data.into_iter().map(|d| d.embedding).collect();
if vectors.len() != prepared.len() {
return Err(CoreError::Memory(format!(
"Embedder returned {} vectors for {} inputs",
vectors.len(),
prepared.len()
)));
}
Ok(vectors)
}
}
const HTTP_EMBED_MAX_CHARS: usize = 1024;
fn truncate_for_embedding(text: &str) -> &str {
if text.len() <= HTTP_EMBED_MAX_CHARS {
return text;
}
let mut end = HTTP_EMBED_MAX_CHARS;
while !text.is_char_boundary(end) {
end -= 1;
}
&text[..end]
}
impl Embedder for HttpEmbedder {
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let url = self.embeddings_url();
let prepared: Vec<&str> = texts.iter().map(|t| truncate_for_embedding(t)).collect();
let truncated = texts
.iter()
.filter(|t| t.len() > HTTP_EMBED_MAX_CHARS)
.count();
if truncated > 0 {
tracing::debug!(
truncated,
budget_chars = HTTP_EMBED_MAX_CHARS,
"http embedder truncated oversized input(s) to the model window"
);
}
let concurrency = self.concurrency.min(prepared.len());
if concurrency <= 1 {
let vectors = self.embed_chunk(&url, &prepared)?;
if vectors.len() != texts.len() {
return Err(CoreError::Memory(format!(
"Embedder returned {} vectors for {} inputs",
vectors.len(),
texts.len()
)));
}
return Ok(vectors);
}
let chunk_size = prepared.len().div_ceil(concurrency);
let mut results: Vec<Result<Vec<Vec<f32>>>> = Vec::with_capacity(concurrency);
std::thread::scope(|scope| {
let handles: Vec<_> = prepared
.chunks(chunk_size)
.map(|chunk| {
let url = &url;
scope.spawn(move || self.embed_chunk(url, chunk))
})
.collect();
for handle in handles {
results.push(handle.join().unwrap_or_else(|_| {
Err(CoreError::Memory("embedder fan-out thread panicked".into()))
}));
}
});
let mut vectors = Vec::with_capacity(texts.len());
for result in results {
vectors.extend(result?);
}
if vectors.len() != texts.len() {
return Err(CoreError::Memory(format!(
"Embedder returned {} vectors for {} inputs",
vectors.len(),
texts.len()
)));
}
Ok(vectors)
}
fn dimension(&self) -> usize {
self.config.dimension
}
fn is_available(&self) -> bool {
self.available
}
fn backend_name(&self) -> &'static str {
"http"
}
fn cache_namespace(&self) -> String {
format!(
"http:{}:{}:{}",
self.config.endpoint, self.config.model, self.config.dimension
)
}
}
#[derive(Debug, Serialize)]
struct EmbeddingsRequest<'a> {
model: &'a str,
input: &'a [&'a str],
}
#[derive(Debug, Deserialize)]
struct EmbeddingsResponse {
data: Vec<EmbeddingData>,
}
#[derive(Debug, Deserialize)]
struct EmbeddingData {
embedding: Vec<f32>,
}
pub struct StubEmbedder {
dimension: usize,
}
impl StubEmbedder {
#[must_use]
pub const fn new(dimension: usize) -> Self {
Self { dimension }
}
}
impl Default for StubEmbedder {
fn default() -> Self {
Self::new(384)
}
}
impl Embedder for StubEmbedder {
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
use sha2::{Digest, Sha256};
let mut results = Vec::with_capacity(texts.len());
for text in texts {
let mut hasher = Sha256::new();
hasher.update(text.as_bytes());
let hash = hasher.finalize();
let mut embedding = Vec::with_capacity(self.dimension);
for i in 0..self.dimension {
let byte = f32::from(hash[i % hash.len()]);
embedding.push(byte.mul_add(2.0 / 255.0, -1.0)); }
results.push(embedding);
}
Ok(results)
}
fn dimension(&self) -> usize {
self.dimension
}
fn is_available(&self) -> bool {
true
}
fn backend_name(&self) -> &'static str {
"stub"
}
}
#[cfg(feature = "onnx")]
pub struct OrtEmbedder {
shards: Vec<std::sync::Mutex<Option<fastembed::TextEmbedding>>>,
model_name: String,
cache_dir: Option<std::path::PathBuf>,
threads: usize,
dimension: usize,
available: std::sync::atomic::AtomicBool,
batch_size: usize,
}
#[cfg(feature = "onnx")]
fn default_intra_threads() -> usize {
std::thread::available_parallelism()
.map(std::num::NonZero::get)
.unwrap_or(4)
.min(4)
}
#[cfg(feature = "onnx")]
impl OrtEmbedder {
#[must_use]
pub fn new(
model_name: &str,
cache_dir: Option<std::path::PathBuf>,
threads: usize,
dimension: usize,
) -> Self {
let threads = threads.max(1);
let shard_count = (threads / 2).clamp(1, 4);
let shards = (0..shard_count)
.map(|_| std::sync::Mutex::new(None))
.collect();
Self {
shards,
model_name: model_name.to_string(),
cache_dir,
threads,
dimension,
available: std::sync::atomic::AtomicBool::new(false),
batch_size: 32,
}
}
#[must_use]
pub fn from_env() -> Option<Self> {
let model_name = std::env::var("WM_EMBEDDER_ORT_MODEL")
.unwrap_or_else(|_| "BAAI/bge-small-en-v1.5".into());
let cache_dir = std::env::var("WM_EMBEDDER_CACHE_DIR")
.ok()
.map(std::path::PathBuf::from);
let threads = std::env::var("WM_EMBEDDER_ORT_THREADS")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or_else(default_intra_threads);
let dimension = std::env::var("WM_EMBEDDER_DIM")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(384);
Some(Self::new(&model_name, cache_dir, threads, dimension))
}
#[must_use]
pub fn with_shard_override(mut self, shards: usize) -> Self {
self.shards = (0..shards.max(1))
.map(|_| std::sync::Mutex::new(None))
.collect();
self
}
fn resolve_model(&self) -> Option<fastembed::EmbeddingModel> {
match self.model_name.as_str() {
"BAAI/bge-small-en-v1.5" | "bge-small-en-v1.5" | "bge-small" => {
Some(fastembed::EmbeddingModel::BGESmallENV15)
}
"BAAI/bge-small-en-v1.5-q" | "bge-small-en-v1.5-q" | "bge-small-q" => {
Some(fastembed::EmbeddingModel::BGESmallENV15Q)
}
"BAAI/bge-base-en-v1.5" | "bge-base-en-v1.5" | "bge-base" => {
Some(fastembed::EmbeddingModel::BGEBaseENV15)
}
"BAAI/bge-base-en-v1.5-q" | "bge-base-en-v1.5-q" | "bge-base-q" => {
Some(fastembed::EmbeddingModel::BGEBaseENV15Q)
}
"BAAI/bge-large-en-v1.5" | "bge-large-en-v1.5" | "bge-large" => {
Some(fastembed::EmbeddingModel::BGELargeENV15)
}
"sentence-transformers/all-MiniLM-L6-v2" | "all-MiniLM-L6-v2" | "minilm" => {
Some(fastembed::EmbeddingModel::AllMiniLML6V2)
}
"sentence-transformers/all-MiniLM-L6-v2-q" | "all-MiniLM-L6-v2-q" | "minilm-q" => {
Some(fastembed::EmbeddingModel::AllMiniLML6V2Q)
}
"sentence-transformers/all-MiniLM-L12-v2" | "all-MiniLM-L12-v2" => {
Some(fastembed::EmbeddingModel::AllMiniLML12V2)
}
"nomic-ai/nomic-embed-text-v1.5" | "nomic-embed-text-v1.5" | "nomic" => {
Some(fastembed::EmbeddingModel::NomicEmbedTextV15)
}
_ => {
tracing::warn!(
"unknown embedder model '{}', falling back to bge-small-en-v1.5",
self.model_name
);
Some(fastembed::EmbeddingModel::BGESmallENV15)
}
}
}
#[cfg(feature = "onnx")]
#[must_use]
pub const fn threads(&self) -> usize {
self.threads
}
#[cfg(feature = "onnx")]
#[must_use]
pub fn shard_count(&self) -> usize {
self.shards.len()
}
#[cfg(feature = "onnx")]
#[must_use]
fn intra_per_shard(&self) -> usize {
(self.threads / self.shards.len().max(1)).max(1)
}
fn ensure_loaded(&self) -> bool {
if self.available.load(std::sync::atomic::Ordering::Relaxed) {
return true;
}
let intra = self.intra_per_shard();
let mut loaded = Vec::with_capacity(self.shards.len());
for _ in 0..self.shards.len() {
let Some(embedding_model) = self.resolve_model() else {
self.available
.store(false, std::sync::atomic::Ordering::Relaxed);
return false;
};
let mut options = fastembed::TextInitOptions::new(embedding_model)
.with_show_download_progress(false)
.with_intra_threads(intra);
if let Some(ref cache_dir) = self.cache_dir {
options = options.with_cache_dir(cache_dir.clone());
}
match fastembed::TextEmbedding::try_new(options) {
Ok(model) => loaded.push(Some(model)),
Err(e) => {
tracing::warn!("failed to load ONNX embedding model: {e}");
self.available
.store(false, std::sync::atomic::Ordering::Relaxed);
return false;
}
}
}
{
for (shard, model) in self.shards.iter().zip(loaded) {
let mut guard = shard.lock().expect("shard mutex poisoned");
*guard = model;
}
}
self.available
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::info!(
"ONNX embedding model loaded: {} ({} session shards × {} intra threads, dim={})",
self.model_name,
self.shards.len(),
intra,
self.dimension
);
true
}
}
#[cfg(feature = "onnx")]
impl Embedder for OrtEmbedder {
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
if !self.ensure_loaded() {
return Err(CoreError::Memory(
"ONNX embedder model not available".into(),
));
}
let owned_texts: Vec<String> = texts.iter().map(|t| (*t).to_string()).collect();
if owned_texts.len() == 1 {
let mut guard = self.shards[0].lock().expect("shard mutex poisoned");
let Some(ref mut model) = guard.as_mut() else {
return Err(CoreError::Memory("ONNX embedder model not loaded".into()));
};
let embeddings = model
.embed(owned_texts, Some(self.batch_size))
.map_err(|e| CoreError::Memory(format!("ONNX embedding error: {e}")))?;
return Ok(embeddings);
}
let shard_count = self.shards.len();
let mut split: Vec<Vec<String>> = vec![Vec::new(); shard_count];
for (i, text) in owned_texts.into_iter().enumerate() {
split[i % shard_count].push(text);
}
let results: Vec<Result<Vec<Vec<f32>>>> = std::thread::scope(|scope| {
let handles: Vec<_> = split
.into_iter()
.zip(&self.shards)
.map(|(subset, shard)| {
scope.spawn(move || {
let mut guard = shard.lock().expect("shard mutex poisoned");
let Some(ref mut model) = guard.as_mut() else {
return Err(CoreError::Memory("ONNX embedder model not loaded".into()));
};
model
.embed(subset, Some(self.batch_size))
.map_err(|e| CoreError::Memory(format!("ONNX embedding error: {e}")))
})
})
.collect();
handles
.into_iter()
.map(|h| h.join().expect("shard worker panicked"))
.collect()
});
let mut per_shard: Vec<std::collections::VecDeque<Vec<f32>>> = results
.into_iter()
.map(|r| {
r.map(|vectors| {
vectors
.into_iter()
.collect::<std::collections::VecDeque<_>>()
})
})
.collect::<Result<Vec<_>>>()?;
let mut out: Vec<Vec<f32>> = Vec::with_capacity(texts.len());
for i in 0..texts.len() {
let shard_idx = i % shard_count;
let Some(vector) = per_shard[shard_idx].pop_front() else {
return Err(CoreError::Memory(format!(
"embed_batch shard {shard_idx} returned too few vectors"
)));
};
out.push(vector);
}
if out.len() != texts.len() {
return Err(CoreError::Memory(format!(
"embed_batch returned {} vectors for {} inputs",
out.len(),
texts.len()
)));
}
Ok(out)
}
fn dimension(&self) -> usize {
self.dimension
}
fn is_available(&self) -> bool {
self.ensure_loaded()
}
fn backend_name(&self) -> &'static str {
"onnx"
}
fn cache_namespace(&self) -> String {
format!("onnx:{}:{}", self.model_name, self.dimension)
}
fn preferred_max_batch_texts(&self) -> usize {
128
}
}
#[must_use]
pub fn create_embedder() -> Box<dyn Embedder> {
#[cfg(feature = "onnx")]
{
let prefer_ort = std::env::var("WM_EMBEDDER_BACKEND")
.map(|v| v == "onnx" || v == "ort")
.unwrap_or(false);
if prefer_ort {
if let Some(mut ort) = OrtEmbedder::from_env() {
if let Ok(shards) = std::env::var("WM_EMBEDDER_ORT_SHARDS") {
if let Ok(n) = shards.parse::<usize>() {
ort = ort.with_shard_override(n);
}
}
tracing::info!(
"onnx embedder configured (dim={}, shards={})",
ort.dimension(),
ort.shard_count()
);
return Box::new(ort);
}
}
}
if let Some(http) = HttpEmbedder::from_env() {
tracing::info!("http embedder configured (dim={})", http.dimension());
return Box::new(http);
}
#[cfg(feature = "onnx")]
{
if let Some(ort) = OrtEmbedder::from_env() {
tracing::info!(
"onnx embedder configured as default (dim={})",
ort.dimension()
);
return Box::new(ort);
}
}
tracing::info!("no embedder endpoint configured, using stub embedder");
Box::new(StubEmbedder::default())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stub_embedder_dimension() {
let embedder = StubEmbedder::new(128);
assert_eq!(embedder.dimension(), 128);
}
#[test]
fn stub_embedder_single() {
let embedder = StubEmbedder::new(64);
let vec = embedder.embed("hello world").unwrap();
assert_eq!(vec.len(), 64);
for v in &vec {
assert!(*v >= -1.0 && *v <= 1.0);
}
}
#[test]
fn stub_embedder_batch() {
let embedder = StubEmbedder::new(32);
let texts = ["hello", "world", "test"];
let vectors = embedder.embed_batch(&texts).unwrap();
assert_eq!(vectors.len(), 3);
for v in &vectors {
assert_eq!(v.len(), 32);
}
}
#[test]
fn stub_embedder_deterministic() {
let embedder = StubEmbedder::new(64);
let v1 = embedder.embed("same text").unwrap();
let v2 = embedder.embed("same text").unwrap();
assert_eq!(v1, v2);
}
#[test]
fn stub_embedder_different_texts_differ() {
let embedder = StubEmbedder::new(64);
let v1 = embedder.embed("hello").unwrap();
let v2 = embedder.embed("world").unwrap();
assert_ne!(v1, v2);
}
#[test]
fn stub_embedder_empty_batch() {
let embedder = StubEmbedder::new(64);
let vectors = embedder.embed_batch(&[]).unwrap();
assert!(vectors.is_empty());
}
#[test]
fn stub_embedder_is_available() {
let embedder = StubEmbedder::new(64);
assert!(embedder.is_available());
}
#[test]
fn stub_embedder_backend_name() {
let embedder = StubEmbedder::new(64);
assert_eq!(embedder.backend_name(), "stub");
}
#[test]
fn stub_embedder_default_dimension() {
let embedder = StubEmbedder::default();
assert_eq!(embedder.dimension(), 384);
}
#[test]
fn http_embedder_config_from_env_absent() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080".into(),
model: "local".into(),
dimension: 384,
timeout: Duration::from_secs(30),
};
assert_eq!(config.endpoint, "http://localhost:8080");
assert_eq!(config.model, "local");
assert_eq!(config.dimension, 384);
}
#[test]
fn http_embedder_embeddings_url() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080".into(),
model: "local".into(),
dimension: 384,
timeout: Duration::from_secs(30),
};
let embedder = HttpEmbedder::new(config);
assert_eq!(
embedder.embeddings_url(),
"http://localhost:8080/v1/embeddings"
);
}
#[test]
fn http_embedder_embeddings_url_trailing_slash() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080/".into(),
model: "local".into(),
dimension: 384,
timeout: Duration::from_secs(30),
};
let embedder = HttpEmbedder::new(config);
assert_eq!(
embedder.embeddings_url(),
"http://localhost:8080/v1/embeddings"
);
}
#[test]
fn http_embedder_embeddings_url_full_path() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080/v1/embeddings".into(),
model: "local".into(),
dimension: 384,
timeout: Duration::from_secs(30),
};
let embedder = HttpEmbedder::new(config);
assert_eq!(
embedder.embeddings_url(),
"http://localhost:8080/v1/embeddings"
);
}
#[test]
fn http_embedder_dimension() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080".into(),
model: "local".into(),
dimension: 768,
timeout: Duration::from_secs(10),
};
let embedder = HttpEmbedder::new(config);
assert_eq!(embedder.dimension(), 768);
}
#[test]
fn http_embedder_backend_name() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080".into(),
model: "local".into(),
dimension: 384,
timeout: Duration::from_secs(10),
};
let embedder = HttpEmbedder::new(config);
assert_eq!(embedder.backend_name(), "http");
}
#[test]
fn http_embedder_is_available() {
let config = EmbedderConfig {
endpoint: "http://localhost:8080".into(),
model: "local".into(),
dimension: 384,
timeout: Duration::from_secs(10),
};
let embedder = HttpEmbedder::new(config);
assert!(embedder.is_available());
}
#[test]
fn http_embedder_fanout_preserves_order() {
use std::io::{Read, Write};
use std::net::TcpListener;
let texts = ["alpha", "beta", "gamma", "delta", "epsilon", "zeta", "eta"];
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let mut served = 0usize;
let deadline = std::time::Instant::now() + Duration::from_secs(10);
while served < 3 && std::time::Instant::now() < deadline {
listener.set_nonblocking(true).unwrap();
let Ok((mut stream, _)) = listener.accept() else {
std::thread::sleep(Duration::from_millis(5));
continue;
};
stream.set_nonblocking(false).unwrap();
stream
.set_read_timeout(Some(Duration::from_secs(5)))
.unwrap();
let mut buf = Vec::new();
let mut tmp = [0u8; 2048];
let header_end = loop {
let n = stream.read(&mut tmp).unwrap();
buf.extend_from_slice(&tmp[..n]);
if let Some(pos) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
break pos + 4;
}
};
let headers = String::from_utf8_lossy(&buf[..header_end]).to_ascii_lowercase();
let content_length: usize = headers
.lines()
.find_map(|l| l.strip_prefix("content-length:"))
.and_then(|v| v.trim().parse().ok())
.unwrap_or(0);
while buf.len() < header_end + content_length {
let n = stream.read(&mut tmp).unwrap();
buf.extend_from_slice(&tmp[..n]);
}
let req: serde_json::Value =
serde_json::from_slice(&buf[header_end..header_end + content_length]).unwrap();
let inputs = req["input"].as_array().unwrap();
let data: Vec<_> = inputs
.iter()
.map(|v| {
let s = v.as_str().unwrap();
let idx = texts.iter().position(|t| *t == s).unwrap();
serde_json::json!({"embedding": [idx as f32]})
})
.collect();
let body = serde_json::json!({"data": data}).to_string();
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
);
stream.write_all(response.as_bytes()).unwrap();
served += 1;
}
served
});
let embedder = HttpEmbedder::new(EmbedderConfig {
endpoint: format!("http://{addr}"),
model: "mock".into(),
dimension: 1,
timeout: Duration::from_secs(10),
})
.with_concurrency(3);
let vectors = embedder.embed_batch(&texts).unwrap();
assert_eq!(vectors.len(), texts.len());
for (i, v) in vectors.iter().enumerate() {
assert_eq!(v, &vec![i as f32], "fan-out must preserve input order");
}
assert_eq!(
server.join().unwrap(),
3,
"7 inputs at concurrency 3 must fan out into 3 requests"
);
}
#[test]
fn http_embedder_cache_namespace_separates_models_and_dims() {
let make = |model: &str, dim: usize| {
HttpEmbedder::new(EmbedderConfig {
endpoint: "http://localhost:8080".into(),
model: model.into(),
dimension: dim,
timeout: Duration::from_secs(10),
})
};
let base = make("bge-small", 384).cache_namespace();
assert_eq!(base, "http:http://localhost:8080:bge-small:384");
assert_ne!(base, make("nomic", 384).cache_namespace());
assert_ne!(base, make("bge-small", 768).cache_namespace());
}
#[test]
fn truncate_for_embedding_respects_budget_and_utf8_boundaries() {
let short = "short text";
assert_eq!(truncate_for_embedding(short), short);
let ascii = "a".repeat(HTTP_EMBED_MAX_CHARS + 500);
assert_eq!(truncate_for_embedding(&ascii).len(), HTTP_EMBED_MAX_CHARS);
let multibyte = "🦀".repeat(HTTP_EMBED_MAX_CHARS);
let truncated = truncate_for_embedding(&multibyte);
assert!(truncated.len() <= HTTP_EMBED_MAX_CHARS);
assert!(truncated.chars().all(|c| c == '🦀'));
assert!(multibyte.starts_with(truncated));
}
#[test]
fn create_embedder_falls_back_to_stub() {
let embedder = create_embedder();
let name = embedder.backend_name();
assert!(name == "stub" || name == "http" || name == "onnx");
}
#[test]
fn embedder_trait_object() {
let embedder: Box<dyn Embedder> = Box::new(StubEmbedder::new(128));
assert_eq!(embedder.dimension(), 128);
let vec = embedder.embed("test").unwrap();
assert_eq!(vec.len(), 128);
}
#[test]
fn embedder_embed_query() {
let embedder = StubEmbedder::new(64);
let v1 = embedder.embed("hello").unwrap();
let v2 = embedder.embed_query("hello").unwrap();
assert_eq!(v1, v2);
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_backend_name() {
let embedder = OrtEmbedder::new("bge-small", None, 2, 384);
assert_eq!(embedder.backend_name(), "onnx");
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_resolves_quantized_models() {
let q = OrtEmbedder::new("bge-small-q", None, 2, 384);
assert!(matches!(
q.resolve_model(),
Some(fastembed::EmbeddingModel::BGESmallENV15Q)
));
let q_full = OrtEmbedder::new("BAAI/bge-small-en-v1.5-q", None, 2, 384);
assert!(matches!(
q_full.resolve_model(),
Some(fastembed::EmbeddingModel::BGESmallENV15Q)
));
let minilm_q = OrtEmbedder::new("minilm-q", None, 2, 384);
assert!(matches!(
minilm_q.resolve_model(),
Some(fastembed::EmbeddingModel::AllMiniLML6V2Q)
));
let fp32 = OrtEmbedder::new("bge-small", None, 2, 384);
assert!(matches!(
fp32.resolve_model(),
Some(fastembed::EmbeddingModel::BGESmallENV15)
));
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_default_threads_are_capped() {
assert!(
default_intra_threads() <= 4,
"default threads must be capped at 4, got {}",
default_intra_threads()
);
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_dimension() {
let embedder = OrtEmbedder::new("bge-small", None, 2, 384);
assert_eq!(embedder.dimension(), 384);
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_dimension_custom() {
let embedder = OrtEmbedder::new("minilm", None, 1, 256);
assert_eq!(embedder.dimension(), 256);
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_pool_shards_the_thread_budget() {
for threads in [1usize, 2, 4, 8] {
let embedder = OrtEmbedder::new("bge-small", None, threads, 384);
assert_eq!(
embedder.threads(),
threads,
"total budget must be preserved"
);
let shards = embedder.shard_count();
assert!(shards >= 1, "at least one session");
assert!(
shards <= threads,
"shards ({shards}) must not exceed the total budget ({threads})"
);
}
let single = OrtEmbedder::new("bge-small", None, 1, 384);
assert_eq!(single.shard_count(), 1, "threads=1 is the legacy shape");
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_namespace_carries_model_and_dimension() {
let small = OrtEmbedder::new("bge-small", None, 2, 384);
let large = OrtEmbedder::new("bge-large", None, 2, 1024);
assert_ne!(
small.cache_namespace(),
large.cache_namespace(),
"different models must not share cache entries"
);
assert!(small.cache_namespace().contains("bge-small"));
}
#[cfg(feature = "onnx")]
#[test]
#[ignore = "loads the real ONNX model; run explicitly for pool tuning"]
fn ort_pool_shape_microbench() {
let texts: Vec<String> = (0..128)
.map(|i| {
format!(
"Session {i} of the deployment retrospective covered the rollout \
schedule for the telemetry agent, the budget review outcomes, and \
the follow-up decisions about the quarterly report timeline {i}."
)
})
.collect();
let refs: Vec<&str> = texts.iter().map(String::as_str).collect();
let configs: Vec<(usize, Option<usize>)> = vec![
(1, Some(1)),
(2, Some(1)),
(4, Some(1)),
(4, None),
(4, Some(2)),
(8, None),
(8, Some(4)),
];
for (threads, shards) in configs {
let mut embedder = OrtEmbedder::new("bge-small-q", None, threads, 384);
if let Some(n) = shards {
embedder = embedder.with_shard_override(n);
}
let warm = embedder.embed_batch(&refs[..1]).unwrap();
assert_eq!(warm.len(), 1);
let t0 = std::time::Instant::now();
let out = embedder.embed_batch(&refs).unwrap();
let dt = t0.elapsed();
assert_eq!(out.len(), refs.len());
println!(
"pool shape: threads={threads} shards={} intra={} → {} texts in {:.2?} ({:.1} texts/s)",
embedder.shard_count(),
embedder.intra_per_shard(),
refs.len(),
dt,
refs.len() as f64 / dt.as_secs_f64()
);
let t1 = std::time::Instant::now();
let mut total = 0usize;
for chunk in refs.chunks(30) {
total += embedder.embed_batch(chunk).unwrap().len();
}
let dt1 = t1.elapsed();
assert_eq!(total, refs.len());
println!(
" ingest shape (30-text chunks): {:.2?} ({:.1} texts/s)",
dt1,
refs.len() as f64 / dt1.as_secs_f64()
);
}
}
#[cfg(feature = "onnx")]
#[test]
#[ignore = "loads the real ONNX model; run explicitly as a pool-identity gate"]
fn ort_pool_shapes_produce_identical_vectors() {
let text = "determinism probe across session pool shapes";
let single = OrtEmbedder::new("bge-small-q", None, 1, 384).with_shard_override(1);
let pooled = OrtEmbedder::new("bge-small-q", None, 4, 384);
let v1 = single.embed(text).unwrap();
let v2 = pooled.embed(text).unwrap();
assert_eq!(v1.len(), v2.len());
let max_diff = v1
.iter()
.zip(v2.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_diff < 1e-6,
"pool shape changed the vector math (max diff {max_diff})"
);
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_empty_batch() {
let embedder = OrtEmbedder::new("bge-small", None, 2, 384);
let result = embedder.embed_batch(&[]).unwrap();
assert!(result.is_empty());
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_trait_object() {
let embedder: Box<dyn Embedder> = Box::new(OrtEmbedder::new("bge-small", None, 2, 384));
assert_eq!(embedder.dimension(), 384);
assert_eq!(embedder.backend_name(), "onnx");
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_resolve_model_known() {
let embedder = OrtEmbedder::new("bge-small-en-v1.5", None, 2, 384);
assert!(embedder.resolve_model().is_some());
}
#[cfg(feature = "onnx")]
#[test]
fn ort_embedder_resolve_model_unknown_falls_back() {
let embedder = OrtEmbedder::new("some-unknown-model", None, 2, 384);
assert!(embedder.resolve_model().is_some());
}
#[test]
fn endpoint_safe_allows_localhost() {
assert!(is_endpoint_safe("http://localhost:8080"));
assert!(is_endpoint_safe("http://127.0.0.1:8080"));
assert!(is_endpoint_safe("http://localhost:11434/v1/embeddings"));
}
#[test]
fn endpoint_safe_allows_private_ip() {
assert!(is_endpoint_safe("http://10.0.0.2:8080"));
assert!(is_endpoint_safe("http://192.168.1.100:8080"));
assert!(is_endpoint_safe("http://172.16.0.5:8080"));
}
#[test]
fn endpoint_safe_blocks_non_http_schemes() {
assert!(!is_endpoint_safe("file:///etc/passwd"));
assert!(!is_endpoint_safe("gopher://localhost:8080"));
assert!(!is_endpoint_safe("ftp://example.com"));
assert!(!is_endpoint_safe("javascript:alert(1)"));
assert!(!is_endpoint_safe("data:text/plain,hello"));
}
#[test]
fn endpoint_safe_blocks_metadata_endpoints() {
assert!(!is_endpoint_safe("http://169.254.169.254/latest/meta-data"));
assert!(!is_endpoint_safe("http://169.254.170.2/v2/metadata"));
assert!(!is_endpoint_safe(
"http://metadata.google.internal/computeMetadata"
));
assert!(!is_endpoint_safe("http://metadata.aws.internal"));
assert!(!is_endpoint_safe("http://metadata"));
}
#[test]
fn endpoint_safe_blocks_empty_host() {
assert!(!is_endpoint_safe("http://"));
assert!(!is_endpoint_safe("http:///path"));
}
#[test]
fn endpoint_safe_blocks_malformed_ipv6() {
assert!(!is_endpoint_safe("http://[::1:8080"));
}
#[test]
fn endpoint_safe_allows_https() {
assert!(is_endpoint_safe("https://localhost:8080"));
assert!(is_endpoint_safe("https://example.com/api"));
}
#[test]
fn endpoint_safe_allows_ipv6_loopback() {
assert!(is_endpoint_safe("http://[::1]:8080"));
}
}