use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, Instant};
use candle_core::{DType, Device, IndexOp, Tensor};
use candle_nn::VarBuilder;
use candle_transformers::models::bert::{BertModel, Config};
use hf_hub::api::sync::{ApiBuilder, ApiRepo};
use hf_hub::{Repo, RepoType};
use serde::Deserialize;
use tokenizers::{Tokenizer, TruncationDirection, TruncationParams, TruncationStrategy};
use crate::embedding_config::{
DEFAULT_REPO, DEFAULT_REVISION, EmbeddingModel, OLLAMA_DEFAULT_URL, Pooling,
endpoint_fingerprint, fingerprint_suffix, huggingface_fingerprint, local_fingerprint,
};
use crate::trace::{EmbedderLoadStatus, TraceEvent, TraceSink};
const ENDPOINT_TIMEOUT_SECS: u64 = 30;
const ENDPOINT_BATCH_SIZE: usize = 64;
const ENDPOINT_RESPONSE_LIMIT_BYTES: u64 = 64 * 1024 * 1024;
const DEFAULT_SLOW_LOAD_MS: u64 = 5_000;
const SLOW_LOAD_REASON: &str = "embedding model load was slow — this machine may be underpowered \
for in-process CPU inference; expect slow embedding builds and queries";
#[derive(Debug, Clone)]
pub enum EmbedderError {
Download {
model: String,
source: String,
},
CacheUnwritable {
source: String,
},
Load {
model: String,
source: String,
},
Inference {
source: String,
},
EmbeddingsNotBuilt,
DimensionMismatch {
expected: usize,
got: usize,
model: String,
},
ModelMismatch {
built: String,
active: String,
},
Config {
message: String,
},
NotCached {
model: String,
revision: Option<String>,
},
}
impl EmbedderError {
fn hint(&self) -> &'static str {
match self {
EmbedderError::Download { .. } => {
"check source availability, connectivity, and the configured model identifier"
}
EmbedderError::CacheUnwritable { .. } => {
"check ~/.cache/huggingface permissions and free disk space (or set HF_HOME)"
}
EmbedderError::Load { .. } => {
"check that model files and configuration are present, readable, and compatible"
}
EmbedderError::Inference { .. } => {
"check model input/configuration or the endpoint response and retry"
}
EmbedderError::EmbeddingsNotBuilt => {
"embed the corpus before running a semantic/hybrid search"
}
EmbedderError::DimensionMismatch { .. } => {
"the configured model changed; re-embed the corpus with the new model"
}
EmbedderError::ModelMismatch { .. } => {
"the configured model changed; re-embed the corpus with the new model"
}
EmbedderError::Config { .. } => {
"give exactly one embedding source (a model id / path / url) with its required fields"
}
EmbedderError::NotCached { .. } => {
"Ratel auto-downloads only the default model; pre-download this one, pass download=true, or use a local path / endpoint"
}
}
}
}
impl std::fmt::Display for EmbedderError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let hint = self.hint();
match self {
EmbedderError::Download { model, source } => write!(
f,
"failed to download embedding model {model}: {source} (hint: {hint})"
),
EmbedderError::CacheUnwritable { source } => write!(
f,
"embedding model cache is not writable: {source} (hint: {hint})"
),
EmbedderError::Load { model, source } => write!(
f,
"failed to load embedding model {model}: {source} (hint: {hint})"
),
EmbedderError::Inference { source } => {
write!(f, "embedding failed: {source} (hint: {hint})")
}
EmbedderError::DimensionMismatch {
expected,
got,
model,
} => write!(
f,
"embedding dimension mismatch for {model}: expected {expected}, got {got} (hint: {hint})"
),
EmbedderError::ModelMismatch { built, active } => write!(
f,
"embedding model mismatch: cache was built with {built}, active model is {active} (hint: {hint})"
),
EmbedderError::Config { message } => write!(f, "{message} (hint: {hint})"),
EmbedderError::NotCached { model, revision } => {
let rev = revision
.as_deref()
.map(|r| format!(" --revision {r}"))
.unwrap_or_default();
write!(
f,
"embedding model {model} is not in the local HuggingFace cache — download it \
first: `huggingface-cli download {model}{rev}` (hint: {hint})"
)
}
EmbedderError::EmbeddingsNotBuilt => {
write!(
f,
"embeddings are not computed for semantic search (hint: {hint})"
)
}
}
}
}
impl std::error::Error for EmbedderError {}
pub(crate) struct Embedded<T> {
pub(crate) value: T,
pub(crate) fingerprint: String,
}
pub(crate) trait Embedder: Send + Sync {
fn embed_doc(&self, text: &str) -> Result<Vec<f32>, EmbedderError>;
fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbedderError>;
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbedderError> {
texts.iter().map(|t| self.embed_doc(t)).collect()
}
fn embed_query_with_identity(&self, text: &str) -> Result<Embedded<Vec<f32>>, EmbedderError> {
Ok(Embedded {
value: self.embed_query(text)?,
fingerprint: self.fingerprint(),
})
}
fn embed_batch_with_identity(
&self,
texts: &[String],
) -> Result<Embedded<Vec<Vec<f32>>>, EmbedderError> {
Ok(Embedded {
value: self.embed_batch(texts)?,
fingerprint: self.fingerprint(),
})
}
fn fingerprint(&self) -> String {
"unknown".to_string()
}
}
struct DownloadNotice {
model: String,
bytes: u64,
}
#[derive(Default)]
struct LoadNotices {
download: Option<DownloadNotice>,
pooling_assumed: Option<String>,
}
type ResolvedEmbedder = (Arc<dyn Embedder>, Option<u64>, LoadNotices);
fn embedder_for(model: &EmbeddingModel) -> Result<ResolvedEmbedder, EmbedderError> {
static CELL: OnceLock<Mutex<HashMap<String, LoadSlot<dyn Embedder>>>> = OnceLock::new();
let cache = CELL.get_or_init(|| Mutex::new(HashMap::new()));
let mut notices = LoadNotices::default();
let (emb, load_ms) = get_or_load_keyed(cache, &model.embedder_cache_key(), || {
let (emb, n) = build_embedder(model)?;
notices = n;
Ok(emb)
})?;
Ok((emb, load_ms, notices))
}
fn build_embedder(
model: &EmbeddingModel,
) -> Result<(Arc<dyn Embedder>, LoadNotices), EmbedderError> {
let query_prefix = model.query_prefix();
let doc_prefix = model.doc_prefix();
let pooling = model.pooling_override();
match model {
EmbeddingModel::Default => {
let (e, n) = CandleEmbedder::load_hf(
DEFAULT_REPO,
DEFAULT_REVISION,
query_prefix,
doc_prefix,
pooling,
true,
)?;
Ok((Arc::new(e), n))
}
EmbeddingModel::HuggingFace {
repo,
revision,
download,
..
} => {
let (e, n) = CandleEmbedder::load_hf(
repo,
revision.as_deref().unwrap_or("main"),
query_prefix,
doc_prefix,
pooling,
*download,
)?;
Ok((Arc::new(e), n))
}
EmbeddingModel::Local { path, .. } => {
let (e, n) = CandleEmbedder::load_path(path, query_prefix, doc_prefix, pooling)?;
Ok((Arc::new(e), n))
}
EmbeddingModel::Endpoint {
url,
model,
api_key_env,
..
} => {
let e = EndpointEmbedder::new(
url.clone(),
model.clone(),
api_key_env.clone(),
query_prefix.into(),
doc_prefix.into(),
)?;
Ok((Arc::new(e), LoadNotices::default()))
}
}
}
type LoadSlot<T> = Arc<Mutex<Option<Arc<T>>>>;
fn get_or_load_keyed<T: ?Sized>(
cache: &Mutex<HashMap<String, LoadSlot<T>>>,
key: &str,
load: impl FnOnce() -> Result<Arc<T>, EmbedderError>,
) -> Result<(Arc<T>, Option<u64>), EmbedderError> {
let slot = {
let mut guard = cache.lock().expect("embedder cache mutex poisoned");
Arc::clone(guard.entry(key.to_string()).or_default())
};
let mut slot = slot.lock().expect("embedder cache slot mutex poisoned");
if let Some(existing) = slot.as_ref() {
return Ok((existing.clone(), None));
}
let started = Instant::now();
let loaded = load()?; let took_ms = started.elapsed().as_millis() as u64;
*slot = Some(loaded.clone());
Ok((loaded, Some(took_ms)))
}
pub(crate) fn embedder_with_telemetry(
model: &EmbeddingModel,
sink: &dyn TraceSink,
) -> Result<Arc<dyn Embedder>, EmbedderError> {
let display = model.display_name();
let (result, load_ms, notices) = match embedder_for(model) {
Ok((emb, ms, notices)) => (Ok(emb), ms, notices),
Err(e) => (Err(e), None, LoadNotices::default()),
};
if let Some(DownloadNotice { model, bytes }) = notices.download {
let mb = bytes as f64 / 1_048_576.0;
eprintln!("ratel: downloaded embedding model {model} ({mb:.0} MB, one-time)");
sink.record(TraceEvent::EmbedderDownload { model, bytes });
}
if let Some(model) = notices.pooling_assumed {
eprintln!(
"ratel: pooling not detected for {model}, assuming mean; \
set pooling=\"cls\"|\"mean\" to override"
);
sink.record(TraceEvent::EmbedderPoolingAssumed {
model,
pooling: "mean".to_string(),
});
}
if let Some(event) = embedder_load_event(&display, load_ms, result.as_ref().err()) {
if let TraceEvent::EmbedderLoad {
status,
took_ms,
reason,
..
} = &event
&& !matches!(status, EmbedderLoadStatus::Ok)
{
eprintln!(
"ratel: embedding model load {status:?} ({took_ms}ms): {}",
reason.as_deref().unwrap_or("")
);
}
sink.record(event);
}
result
}
fn embedder_load_event(
model: &str,
load_ms: Option<u64>,
error: Option<&EmbedderError>,
) -> Option<TraceEvent> {
match (load_ms, error) {
(_, Some(err)) => Some(TraceEvent::EmbedderLoad {
model: model.to_string(),
status: EmbedderLoadStatus::Failed,
took_ms: load_ms.unwrap_or(0),
reason: Some(err.to_string()),
}),
(Some(ms), None) => {
let slow = ms > slow_load_ms();
Some(TraceEvent::EmbedderLoad {
model: model.to_string(),
status: if slow {
EmbedderLoadStatus::Slow
} else {
EmbedderLoadStatus::Ok
},
took_ms: ms,
reason: slow.then(|| SLOW_LOAD_REASON.to_string()),
})
}
(None, None) => None,
}
}
fn slow_load_ms() -> u64 {
std::env::var("RATEL_EMBED_SLOW_MS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(DEFAULT_SLOW_LOAD_MS)
}
const EMBED_BATCH_CHUNK: usize = 32;
pub(crate) struct CandleEmbedder {
model: BertModel,
tokenizer: Tokenizer,
device: Device,
pooling: Pooling,
query_prefix: String,
doc_prefix: String,
fingerprint: String,
}
struct Loaded {
config: PathBuf,
tokenizer: PathBuf,
weights: PathBuf,
pooling: Pooling,
}
impl CandleEmbedder {
fn load_hf(
repo_id: &str,
revision: &str,
query_prefix: &str,
doc_prefix: &str,
pooling_override: Option<Pooling>,
allow_download: bool,
) -> Result<(Self, LoadNotices), EmbedderError> {
let device = Device::Cpu;
let repo_spec =
Repo::with_revision(repo_id.to_string(), RepoType::Model, revision.to_string());
let cache_repo = hf_hub::Cache::from_env().repo(repo_spec.clone());
let (config_path, tokenizer_path, weights_path, pooling_file, download) = if allow_download
{
let was_cached = cache_repo.get("model.safetensors").is_some()
|| cache_repo.get("pytorch_model.bin").is_some();
if !was_cached {
eprintln!(
"ratel: downloading embedding model {repo_id} (one-time; this may take a moment)…"
);
}
let api = ApiBuilder::from_env()
.build()
.map_err(|e| EmbedderError::Download {
model: repo_id.to_string(),
source: e.to_string(),
})?;
let repo = api.repo(repo_spec);
let config = fetch_cached(&repo, "config.json", repo_id)?;
let tokenizer = fetch_cached(&repo, "tokenizer.json", repo_id)?;
let weights = match fetch_cached(&repo, "model.safetensors", repo_id) {
Ok(p) => p,
Err(EmbedderError::Download { source, .. }) if is_not_found(&source) => {
fetch_cached(&repo, "pytorch_model.bin", repo_id)?
}
Err(e) => return Err(e),
};
let pooling_file = fetch_optional(&repo, "1_Pooling/config.json");
let notice = (!was_cached).then(|| DownloadNotice {
model: repo_id.to_string(),
bytes: [&config, &tokenizer, &weights]
.iter()
.filter_map(|p| std::fs::metadata(p).ok().map(|m| m.len()))
.sum(),
});
(config, tokenizer, weights, pooling_file, notice)
} else {
let not_cached = || EmbedderError::NotCached {
model: repo_id.to_string(),
revision: (revision != "main").then(|| revision.to_string()),
};
let config = cache_repo.get("config.json").ok_or_else(not_cached)?;
let tokenizer = cache_repo.get("tokenizer.json").ok_or_else(not_cached)?;
let weights = cache_repo
.get("model.safetensors")
.or_else(|| cache_repo.get("pytorch_model.bin"))
.ok_or_else(not_cached)?;
let pooling_file = cache_repo.get("1_Pooling/config.json");
(config, tokenizer, weights, pooling_file, None)
};
let detected =
pooling_override.or_else(|| pooling_file.and_then(|p| detect_pooling_file(&p)));
let (pooling, pooling_assumed) = resolve_pooling(detected);
let notices = LoadNotices {
download,
pooling_assumed: pooling_assumed.then(|| repo_id.to_string()),
};
let sha = snapshot_sha(&weights_path).unwrap_or_else(|| revision.to_string());
let loaded = Loaded {
config: config_path,
tokenizer: tokenizer_path,
weights: weights_path,
pooling,
};
let embedder = Self::build(
device,
&loaded,
query_prefix,
doc_prefix,
huggingface_fingerprint(repo_id, &sha),
repo_id,
)?;
Ok((embedder, notices))
}
fn load_path(
dir: &Path,
query_prefix: &str,
doc_prefix: &str,
pooling_override: Option<Pooling>,
) -> Result<(Self, LoadNotices), EmbedderError> {
let device = Device::Cpu;
let name = dir.display().to_string();
let config_path = dir.join("config.json");
let tokenizer_path = dir.join("tokenizer.json");
let weights_path = [dir.join("model.safetensors"), dir.join("pytorch_model.bin")]
.into_iter()
.find(|p| p.exists())
.ok_or_else(|| EmbedderError::Load {
model: name.clone(),
source: format!("missing model.safetensors / pytorch_model.bin in {name}"),
})?;
for (p, f) in [
(&config_path, "config.json"),
(&tokenizer_path, "tokenizer.json"),
] {
if !p.exists() {
return Err(EmbedderError::Load {
model: name.clone(),
source: format!(
"missing {f} in {name} — a fast tokenizer.json is required; run \
tokenizer.save_pretrained() upstream, or serve the model via an endpoint"
),
});
}
}
let detected =
pooling_override.or_else(|| detect_pooling_file(&dir.join("1_Pooling/config.json")));
let (pooling, pooling_assumed) = resolve_pooling(detected);
let notices = LoadNotices {
download: None,
pooling_assumed: pooling_assumed.then(|| name.clone()),
};
let loaded = Loaded {
config: config_path,
tokenizer: tokenizer_path,
weights: weights_path,
pooling,
};
let embedder = Self::build(
device,
&loaded,
query_prefix,
doc_prefix,
local_fingerprint(&name),
&name,
)?;
Ok((embedder, notices))
}
fn build(
device: Device,
loaded: &Loaded,
query_prefix: &str,
doc_prefix: &str,
base_fingerprint: String,
model_name: &str,
) -> Result<Self, EmbedderError> {
let load_err = |source: String| EmbedderError::Load {
model: model_name.to_string(),
source,
};
let config_bytes = std::fs::read(&loaded.config).map_err(|e| load_err(e.to_string()))?;
let config: Config =
serde_json::from_slice(&config_bytes).map_err(|e| load_err(e.to_string()))?;
let mut tokenizer =
Tokenizer::from_file(&loaded.tokenizer).map_err(|e| load_err(e.to_string()))?;
tokenizer
.with_truncation(Some(TruncationParams {
max_length: config.max_position_embeddings,
strategy: TruncationStrategy::LongestFirst,
direction: TruncationDirection::Right,
stride: 0,
}))
.map_err(|e| load_err(e.to_string()))?;
let is_safetensors =
loaded.weights.extension().and_then(|e| e.to_str()) == Some("safetensors");
let vb = if is_safetensors {
unsafe {
VarBuilder::from_mmaped_safetensors(&[&loaded.weights], DType::F32, &device)
.map_err(|e| load_err(e.to_string()))?
}
} else {
VarBuilder::from_pth(&loaded.weights, DType::F32, &device)
.map_err(|e| load_err(e.to_string()))?
};
let model = BertModel::load(vb, &config).map_err(|e| EmbedderError::Load {
model: model_name.to_string(),
source: format!(
"{e} — if this is not a BERT-family model it can't run in-process; \
serve it in a local model server and use {{\"ollama\": \"…\"}} or \
{{\"url\", \"model\"}} (e.g. Ollama at {OLLAMA_DEFAULT_URL})"
),
})?;
let fingerprint = format!(
"{base_fingerprint}{}",
fingerprint_suffix(Some(loaded.pooling), query_prefix, doc_prefix)
);
Ok(Self {
model,
tokenizer,
device,
pooling: loaded.pooling,
query_prefix: query_prefix.to_string(),
doc_prefix: doc_prefix.to_string(),
fingerprint,
})
}
fn embed(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
self.embed_inner(text)
.map_err(|e| EmbedderError::Inference {
source: e.to_string(),
})
}
fn embed_inner(&self, text: &str) -> candle_core::Result<Vec<f32>> {
let encoding = self
.tokenizer
.encode(text, true)
.map_err(|e| candle_core::Error::Msg(e.to_string()))?;
let ids = encoding.get_ids();
let input_ids = Tensor::new(ids, &self.device)?.unsqueeze(0)?; let token_type_ids = input_ids.zeros_like()?;
let mask: Vec<u32> = encoding.get_attention_mask().to_vec();
let attention_mask = Tensor::new(mask.as_slice(), &self.device)?.unsqueeze(0)?;
let sequence_output =
self.model
.forward(&input_ids, &token_type_ids, Some(&attention_mask))?;
let pooled = match self.pooling {
Pooling::Cls => sequence_output.i((0, 0))?, Pooling::Mean => mean_pool(&sequence_output, &attention_mask)?,
};
let vec = pooled.to_vec1::<f32>()?;
Ok(l2_normalize(vec))
}
fn embed_batch_inner(&self, texts: &[String]) -> candle_core::Result<Vec<Vec<f32>>> {
let mut out = Vec::with_capacity(texts.len());
for chunk in texts.chunks(EMBED_BATCH_CHUNK) {
let inputs: Vec<String> = if self.doc_prefix.is_empty() {
chunk.to_vec()
} else {
chunk
.iter()
.map(|t| format!("{}{}", self.doc_prefix, t))
.collect()
};
let encodings = self
.tokenizer
.encode_batch(inputs, true)
.map_err(|e| candle_core::Error::Msg(e.to_string()))?;
let n = encodings.len();
let max_len = encodings
.iter()
.map(|e| e.get_ids().len())
.max()
.unwrap_or(0);
let mut ids = vec![0u32; n * max_len];
let mut mask = vec![0u32; n * max_len];
for (row, enc) in encodings.iter().enumerate() {
let e_ids = enc.get_ids();
let e_mask = enc.get_attention_mask();
let base = row * max_len;
ids[base..base + e_ids.len()].copy_from_slice(e_ids);
mask[base..base + e_mask.len()].copy_from_slice(e_mask);
}
let input_ids = Tensor::from_vec(ids, (n, max_len), &self.device)?;
let attention_mask = Tensor::from_vec(mask, (n, max_len), &self.device)?;
let token_type_ids = input_ids.zeros_like()?;
let sequence_output =
self.model
.forward(&input_ids, &token_type_ids, Some(&attention_mask))?;
let pooled = match self.pooling {
Pooling::Cls => sequence_output.narrow(1, 0, 1)?.squeeze(1)?, Pooling::Mean => mean_pool_batch(&sequence_output, &attention_mask)?,
};
for row in pooled.to_vec2::<f32>()? {
out.push(l2_normalize(row));
}
}
Ok(out)
}
}
fn mean_pool(sequence_output: &Tensor, attention_mask: &Tensor) -> candle_core::Result<Tensor> {
let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; let summed = sequence_output.broadcast_mul(&mask)?.sum(1)?; let counts = mask.sum(1)?; summed.broadcast_div(&counts)?.i(0) }
fn mean_pool_batch(
sequence_output: &Tensor,
attention_mask: &Tensor,
) -> candle_core::Result<Tensor> {
let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; let summed = sequence_output.broadcast_mul(&mask)?.sum(1)?; let counts = mask.sum(1)?; summed.broadcast_div(&counts) }
impl Embedder for CandleEmbedder {
fn embed_doc(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
if self.doc_prefix.is_empty() {
self.embed(text)
} else {
self.embed(&format!("{}{}", self.doc_prefix, text))
}
}
fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
if self.query_prefix.is_empty() {
self.embed(text)
} else {
self.embed(&format!("{}{}", self.query_prefix, text))
}
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbedderError> {
self.embed_batch_inner(texts)
.map_err(|e| EmbedderError::Inference {
source: e.to_string(),
})
}
fn fingerprint(&self) -> String {
self.fingerprint.clone()
}
}
#[derive(Deserialize)]
struct PoolingConfig {
#[serde(default)]
pooling_mode_cls_token: bool,
#[serde(default)]
pooling_mode_mean_tokens: bool,
}
fn detect_pooling_file(path: &Path) -> Option<Pooling> {
let bytes = std::fs::read(path).ok()?;
parse_pooling_config(&bytes)
}
fn parse_pooling_config(bytes: &[u8]) -> Option<Pooling> {
let c: PoolingConfig = serde_json::from_slice(bytes).ok()?;
if c.pooling_mode_cls_token {
Some(Pooling::Cls)
} else if c.pooling_mode_mean_tokens {
Some(Pooling::Mean)
} else {
None
}
}
fn resolve_pooling(detected: Option<Pooling>) -> (Pooling, bool) {
match detected {
Some(p) => (p, false),
None => (Pooling::Mean, true),
}
}
fn fetch_optional(repo: &ApiRepo, file: &str) -> Option<PathBuf> {
repo.get(file).ok()
}
fn is_not_found(source: &str) -> bool {
let l = source.to_lowercase();
l.contains("404") || l.contains("not found") || l.contains("entry not found")
}
fn snapshot_sha(weights_path: &Path) -> Option<String> {
let name = weights_path.parent()?.file_name()?.to_str()?;
(name.len() == 40 && name.chars().all(|c| c.is_ascii_hexdigit())).then(|| name.to_string())
}
pub(crate) struct EndpointEmbedder {
url: String,
model: String,
api_key_env: Option<String>,
query_prefix: String,
doc_prefix: String,
agent: ureq::Agent,
fingerprint: String,
}
#[derive(Deserialize)]
struct EmbeddingsResponse {
data: Vec<EmbeddingData>,
#[serde(default)]
model: Option<String>,
}
#[derive(Deserialize)]
struct EmbeddingData {
embedding: Vec<f32>,
index: usize,
}
struct ParsedEmbeddings {
vectors: Vec<Vec<f32>>,
model: Option<String>,
}
impl EndpointEmbedder {
fn new(
url: String,
model: String,
api_key_env: Option<String>,
query_prefix: String,
doc_prefix: String,
) -> Result<Self, EmbedderError> {
let agent: ureq::Agent = ureq::Agent::config_builder()
.timeout_global(Some(Duration::from_secs(ENDPOINT_TIMEOUT_SECS)))
.build()
.into();
let fingerprint = format!(
"{}{}",
endpoint_fingerprint(&url, &model),
fingerprint_suffix(None, &query_prefix, &doc_prefix)
);
Ok(Self {
url,
model,
api_key_env,
query_prefix,
doc_prefix,
agent,
fingerprint,
})
}
fn api_key(&self) -> Result<Option<String>, EmbedderError> {
match &self.api_key_env {
None => Ok(None),
Some(var) => std::env::var(var)
.map(Some)
.map_err(|_| EmbedderError::Config {
message: format!(
"api_key_env=\"{var}\" but that environment variable is not set"
),
}),
}
}
fn request_chunk(&self, inputs: &[String]) -> Result<Embedded<Vec<Vec<f32>>>, EmbedderError> {
let key = self.api_key()?;
let body = serde_json::json!({ "model": self.model, "input": inputs });
let mut req = self
.agent
.post(&self.url)
.header("content-type", "application/json");
if let Some(k) = key {
req = req.header("authorization", &format!("Bearer {k}"));
}
let mut resp = req.send_json(&body).map_err(|e| self.classify(e))?;
let parsed: EmbeddingsResponse = resp
.body_mut()
.with_config()
.limit(ENDPOINT_RESPONSE_LIMIT_BYTES)
.read_json()
.map_err(|e| EmbedderError::Inference {
source: format!("malformed or oversized endpoint response: {e}"),
})?;
let parsed = parse_embeddings(parsed, inputs.len())?;
let resolved_model = parsed.model.as_deref().unwrap_or(&self.model);
Ok(Embedded {
value: parsed.vectors,
fingerprint: self.fingerprint_for_model(resolved_model),
})
}
fn request(&self, inputs: &[String]) -> Result<Embedded<Vec<Vec<f32>>>, EmbedderError> {
if inputs.is_empty() {
return Ok(Embedded {
value: Vec::new(),
fingerprint: self.fingerprint.clone(),
});
}
let mut vectors = Vec::with_capacity(inputs.len());
let mut fingerprint: Option<String> = None;
let mut dimension = None;
for chunk in inputs.chunks(ENDPOINT_BATCH_SIZE) {
let embedded = self.request_chunk(chunk)?;
if let Some(first) = &fingerprint {
if first != &embedded.fingerprint {
return Err(EmbedderError::ModelMismatch {
built: first.clone(),
active: embedded.fingerprint,
});
}
} else {
fingerprint = Some(embedded.fingerprint.clone());
}
let chunk_dimension = embedded
.value
.first()
.expect("non-empty request chunk has a non-empty response")
.len();
if let Some(expected) = dimension {
if expected != chunk_dimension {
return Err(EmbedderError::Inference {
source: format!(
"endpoint returned mixed embedding dimensions across chunks: expected {expected}, got {chunk_dimension}"
),
});
}
} else {
dimension = Some(chunk_dimension);
}
vectors.extend(embedded.value);
}
Ok(Embedded {
value: vectors,
fingerprint: fingerprint.expect("non-empty input produced at least one chunk"),
})
}
fn fingerprint_for_model(&self, model: &str) -> String {
format!(
"{}{}",
endpoint_fingerprint(&self.url, model),
fingerprint_suffix(None, &self.query_prefix, &self.doc_prefix)
)
}
fn classify(&self, e: ureq::Error) -> EmbedderError {
let status = match &e {
ureq::Error::StatusCode(code) => Some(*code),
_ => None,
};
let is_local_ollama =
self.url.contains("localhost:11434") || self.url.contains("127.0.0.1:11434");
match status {
Some(401) | Some(403) => EmbedderError::Config {
message: format!(
"endpoint rejected the request ({}); check api_key_env / the key",
status.unwrap()
),
},
Some(404) => {
let hint = if is_local_ollama {
format!(" — run: ollama pull {}", self.model)
} else {
String::new()
};
EmbedderError::Download {
model: self.model.clone(),
source: format!("endpoint returned 404 for model '{}'{hint}", self.model),
}
}
_ if is_local_ollama => EmbedderError::Download {
model: self.model.clone(),
source: format!(
"could not reach Ollama at {} ({e}) — is it running? start it with \
`ollama serve`, then `ollama pull {}`",
self.url, self.model
),
},
_ => EmbedderError::Download {
model: self.model.clone(),
source: e.to_string(),
},
}
}
}
fn parse_embeddings(
resp: EmbeddingsResponse,
expected_len: usize,
) -> Result<ParsedEmbeddings, EmbedderError> {
if resp.data.len() != expected_len {
return Err(EmbedderError::Inference {
source: format!(
"endpoint returned {} embeddings for {expected_len} inputs",
resp.data.len()
),
});
}
if resp
.model
.as_deref()
.is_some_and(|model| model.trim().is_empty())
{
return Err(EmbedderError::Inference {
source: "endpoint returned a blank model identity".into(),
});
}
let mut ordered: Vec<Option<Vec<f32>>> = (0..expected_len).map(|_| None).collect();
let mut dimension = None;
for data in resp.data {
if data.index >= expected_len {
return Err(EmbedderError::Inference {
source: format!(
"endpoint returned out-of-range embedding index {} for {expected_len} inputs",
data.index
),
});
}
if ordered[data.index].is_some() {
return Err(EmbedderError::Inference {
source: format!("endpoint returned duplicate embedding index {}", data.index),
});
}
let vector = normalize_endpoint_vector(data.embedding, &mut dimension)?;
ordered[data.index] = Some(vector);
}
let vectors = ordered
.into_iter()
.enumerate()
.map(|(index, vector)| {
vector.ok_or_else(|| EmbedderError::Inference {
source: format!("endpoint response is missing embedding index {index}"),
})
})
.collect::<Result<_, _>>()?;
Ok(ParsedEmbeddings {
vectors,
model: resp.model,
})
}
fn normalize_endpoint_vector(
mut vector: Vec<f32>,
dimension: &mut Option<usize>,
) -> Result<Vec<f32>, EmbedderError> {
if vector.is_empty() {
return Err(EmbedderError::Inference {
source: "endpoint returned an empty embedding vector".into(),
});
}
if vector.iter().any(|value| !value.is_finite()) {
return Err(EmbedderError::Inference {
source: "endpoint returned a non-finite embedding value".into(),
});
}
match *dimension {
Some(expected) if vector.len() != expected => {
return Err(EmbedderError::Inference {
source: format!(
"endpoint returned mixed embedding dimensions: expected {expected}, got {}",
vector.len()
),
});
}
None => *dimension = Some(vector.len()),
Some(_) => {}
}
let norm = vector
.iter()
.map(|value| f64::from(*value).powi(2))
.sum::<f64>()
.sqrt();
if !norm.is_finite() || norm == 0.0 {
return Err(EmbedderError::Inference {
source: "endpoint returned a zero or non-normalizable embedding vector".into(),
});
}
for value in &mut vector {
*value = (f64::from(*value) / norm) as f32;
}
Ok(vector)
}
impl Embedder for EndpointEmbedder {
fn embed_doc(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
self.embed_batch_with_identity(std::slice::from_ref(&text.to_string()))?
.value
.into_iter()
.next()
.ok_or_else(|| EmbedderError::Inference {
source: "endpoint returned no embedding".into(),
})
}
fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
Ok(self.embed_query_with_identity(text)?.value)
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, EmbedderError> {
Ok(self.embed_batch_with_identity(texts)?.value)
}
fn embed_query_with_identity(&self, text: &str) -> Result<Embedded<Vec<f32>>, EmbedderError> {
let q = if self.query_prefix.is_empty() {
text.to_string()
} else {
format!("{}{}", self.query_prefix, text)
};
let embedded = self.request(&[q])?;
let fingerprint = embedded.fingerprint;
let value = embedded
.value
.into_iter()
.next()
.ok_or_else(|| EmbedderError::Inference {
source: "endpoint returned no embedding".into(),
})?;
Ok(Embedded { value, fingerprint })
}
fn embed_batch_with_identity(
&self,
texts: &[String],
) -> Result<Embedded<Vec<Vec<f32>>>, EmbedderError> {
if self.doc_prefix.is_empty() {
self.request(texts)
} else {
let prefixed: Vec<String> = texts
.iter()
.map(|t| format!("{}{}", self.doc_prefix, t))
.collect();
self.request(&prefixed)
}
}
fn fingerprint(&self) -> String {
self.fingerprint.clone()
}
}
fn fetch_cached(repo: &ApiRepo, file: &str, model: &str) -> Result<PathBuf, EmbedderError> {
const MAX_ATTEMPTS: u32 = 30;
const BACKOFF: Duration = Duration::from_secs(1);
let mut attempt = 1;
loop {
match repo.get(file) {
Ok(path) => return Ok(path),
Err(e) => {
let msg = e.to_string();
if attempt < MAX_ATTEMPTS && is_lock_contention(&msg) {
attempt += 1;
std::thread::sleep(BACKOFF);
continue;
}
return Err(classify_fetch_error(model, &msg));
}
}
}
}
fn is_lock_contention(err: &str) -> bool {
err.contains("Lock acquisition failed")
}
fn classify_fetch_error(model: &str, msg: &str) -> EmbedderError {
let lower = msg.to_lowercase();
let unwritable = lower.contains("permission denied")
|| lower.contains("read-only")
|| lower.contains("no space")
|| lower.contains("os error 13") || lower.contains("os error 28") || lower.contains("os error 30"); if unwritable {
EmbedderError::CacheUnwritable {
source: msg.to_string(),
}
} else {
EmbedderError::Download {
model: model.to_string(),
source: msg.to_string(),
}
}
}
fn l2_normalize(mut v: Vec<f32>) -> Vec<f32> {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut v {
*x /= norm;
}
}
v
}
#[cfg(test)]
mod tests {
use std::io::{Read, Write};
use std::net::TcpListener;
use std::sync::mpsc;
use super::*;
use crate::{Origin, SearchMethod, Tool, ToolRegistry};
#[test]
fn only_lock_contention_is_retried() {
assert!(is_lock_contention(
"Lock acquisition failed: /home/u/.cache/huggingface/hub/models--BAAI--bge-small-en-v1.5/blobs/abc.lock"
));
assert!(!is_lock_contention("request error: connection refused"));
assert!(!is_lock_contention("Http(reqwest::Error { status: 404 })"));
assert!(!is_lock_contention(
"No such file or directory (os error 2)"
));
}
#[test]
fn classifies_cache_permission_and_space_as_unwritable() {
assert!(matches!(
classify_fetch_error("m", "Permission denied (os error 13)"),
EmbedderError::CacheUnwritable { .. }
));
assert!(matches!(
classify_fetch_error("m", "No space left on device (os error 28)"),
EmbedderError::CacheUnwritable { .. }
));
assert!(matches!(
classify_fetch_error("m", "Read-only file system (os error 30)"),
EmbedderError::CacheUnwritable { .. }
));
}
#[test]
fn classifies_network_and_http_as_download() {
assert!(matches!(
classify_fetch_error("m", "error sending request: dns error: failed to lookup"),
EmbedderError::Download { .. }
));
assert!(matches!(
classify_fetch_error("m", "Http status client error (404 Not Found)"),
EmbedderError::Download { .. }
));
}
#[test]
fn error_display_carries_source_and_hint() {
let s = EmbedderError::Download {
model: "embed-v1 @ https://embeddings.example.test".into(),
source: "connection refused".into(),
}
.to_string();
assert!(s.contains("connection refused"), "got: {s}");
assert!(s.contains("hint:"), "got: {s}");
assert!(!s.contains("revision"), "got: {s}");
let load = EmbedderError::Load {
model: "/models/embed".into(),
source: "missing config.json".into(),
}
.to_string();
assert!(!load.contains("re-download"), "got: {load}");
let inference = EmbedderError::Inference {
source: "endpoint returned duplicate index 0".into(),
}
.to_string();
assert!(!inference.contains("underpowered"), "got: {inference}");
}
#[test]
fn get_or_load_keyed_does_not_cache_failure_and_reports_latency_once() {
let cache: Mutex<HashMap<String, LoadSlot<i32>>> = Mutex::new(HashMap::new());
let boom = || {
Err::<Arc<i32>, _>(EmbedderError::Inference {
source: "boom".into(),
})
};
assert!(get_or_load_keyed(&cache, "k", boom).is_err());
let (v, ms) =
get_or_load_keyed(&cache, "k", || Ok::<_, EmbedderError>(Arc::new(7))).unwrap();
assert_eq!(*v, 7);
assert!(ms.is_some(), "the loading call reports latency");
let (v2, ms2) =
get_or_load_keyed(&cache, "k", || Ok::<_, EmbedderError>(Arc::new(999))).unwrap();
assert_eq!(*v2, 7);
assert!(ms2.is_none(), "warm reuse reports no load latency");
}
#[test]
fn distinct_key_loads_run_concurrently() {
use std::sync::Barrier;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::thread;
let cache: Arc<Mutex<HashMap<String, LoadSlot<i32>>>> =
Arc::new(Mutex::new(HashMap::new()));
let barrier = Arc::new(Barrier::new(2));
let inflight = Arc::new(AtomicUsize::new(0));
let (done_tx, done_rx) = mpsc::channel();
for (i, key) in ["a", "b"].into_iter().enumerate() {
let cache = Arc::clone(&cache);
let barrier = Arc::clone(&barrier);
let inflight = Arc::clone(&inflight);
let done_tx = done_tx.clone();
thread::spawn(move || {
let result = get_or_load_keyed(&cache, key, || {
inflight.fetch_add(1, Ordering::SeqCst);
barrier.wait(); Ok::<_, EmbedderError>(Arc::new(i as i32))
});
let _ = done_tx.send(result.map(|(v, _)| *v));
});
}
drop(done_tx);
let mut got = Vec::new();
for _ in 0..2 {
let value = done_rx
.recv_timeout(Duration::from_secs(5))
.expect("distinct-key loads did not run concurrently (map lock held across load)");
got.push(value.unwrap());
}
got.sort_unstable();
assert_eq!(got, vec![0, 1]);
assert_eq!(inflight.load(Ordering::SeqCst), 2);
}
#[test]
fn same_key_load_is_single_flight() {
use std::sync::Barrier;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::thread;
const THREADS: usize = 8;
let cache: Arc<Mutex<HashMap<String, LoadSlot<i32>>>> =
Arc::new(Mutex::new(HashMap::new()));
let loads = Arc::new(AtomicUsize::new(0));
let start = Arc::new(Barrier::new(THREADS));
let (tx, rx) = mpsc::channel();
for _ in 0..THREADS {
let cache = Arc::clone(&cache);
let loads = Arc::clone(&loads);
let start = Arc::clone(&start);
let tx = tx.clone();
thread::spawn(move || {
start.wait(); let result = get_or_load_keyed(&cache, "shared", || {
loads.fetch_add(1, Ordering::SeqCst);
Ok::<_, EmbedderError>(Arc::new(42))
});
let (value, ms) = result.unwrap();
let _ = tx.send((*value, ms.is_some()));
});
}
drop(tx);
let mut reported = 0;
for _ in 0..THREADS {
let (value, loaded) = rx.recv_timeout(Duration::from_secs(5)).unwrap();
assert_eq!(value, 42);
if loaded {
reported += 1;
}
}
assert_eq!(
loads.load(Ordering::SeqCst),
1,
"load must run exactly once per key"
);
assert_eq!(reported, 1, "exactly one caller reports the load latency");
}
#[test]
fn load_event_flags_slow_ok_failed_and_warm() {
assert!(matches!(
embedder_load_event("m", Some(10_000), None),
Some(TraceEvent::EmbedderLoad {
status: EmbedderLoadStatus::Slow,
took_ms: 10_000,
..
})
));
assert!(matches!(
embedder_load_event("m", Some(5), None),
Some(TraceEvent::EmbedderLoad {
status: EmbedderLoadStatus::Ok,
..
})
));
let err = EmbedderError::Inference { source: "x".into() };
assert!(matches!(
embedder_load_event("m", None, Some(&err)),
Some(TraceEvent::EmbedderLoad {
status: EmbedderLoadStatus::Failed,
..
})
));
assert!(embedder_load_event("m", None, None).is_none());
}
#[test]
fn endpoint_embeddings_are_normalized_and_ordered_by_index() {
let resp = EmbeddingsResponse {
model: Some("resolved-model".into()),
data: vec![
EmbeddingData {
embedding: vec![0.0, 3.0], index: 1,
},
EmbeddingData {
embedding: vec![4.0, 0.0], index: 0,
},
],
};
let out = parse_embeddings(resp, 2).expect("parse");
assert_eq!(out.vectors[0], vec![1.0, 0.0], "index 0 first, normalized");
assert_eq!(out.vectors[1], vec![0.0, 1.0], "index 1 second, normalized");
assert_eq!(out.model.as_deref(), Some("resolved-model"));
}
#[test]
fn endpoint_response_count_mismatch_errors() {
let resp = EmbeddingsResponse {
model: None,
data: vec![EmbeddingData {
embedding: vec![1.0],
index: 0,
}],
};
assert!(matches!(
parse_embeddings(resp, 2),
Err(EmbedderError::Inference { .. })
));
}
fn response(vectors: &[(usize, Vec<f32>)]) -> EmbeddingsResponse {
EmbeddingsResponse {
model: None,
data: vectors
.iter()
.cloned()
.map(|(index, embedding)| EmbeddingData { embedding, index })
.collect(),
}
}
#[test]
fn endpoint_response_requires_an_exact_index_permutation() {
assert!(
serde_json::from_value::<EmbeddingsResponse>(serde_json::json!({
"data": [{ "embedding": [1.0] }]
}))
.is_err(),
"index is required"
);
for malformed in [
response(&[(0, vec![1.0]), (0, vec![1.0])]),
response(&[(0, vec![1.0]), (2, vec![1.0])]),
] {
assert!(matches!(
parse_embeddings(malformed, 2),
Err(EmbedderError::Inference { .. })
));
}
}
#[test]
fn endpoint_response_rejects_invalid_vectors() {
for malformed in [
response(&[(0, vec![])]),
response(&[(0, vec![0.0, 0.0])]),
response(&[(0, vec![f32::NAN, 1.0])]),
response(&[(0, vec![f32::INFINITY, 1.0])]),
response(&[(0, vec![1.0, 0.0]), (1, vec![1.0])]),
] {
let expected_len = malformed.data.len();
assert!(matches!(
parse_embeddings(malformed, expected_len),
Err(EmbedderError::Inference { .. })
));
}
}
#[test]
fn endpoint_rejects_a_response_over_64_mib() {
let (url, requests_rx, server) = mock_endpoint(vec![MockReply::Oversized]);
let embedder = EndpointEmbedder::new(
url,
"requested-model".into(),
None,
String::new(),
String::new(),
)
.unwrap();
let err = embedder
.embed_batch(&["one".to_string()])
.expect_err("oversized response must fail");
let requests = requests_rx.recv_timeout(Duration::from_secs(5)).unwrap();
server.join().unwrap();
assert!(err.to_string().contains("oversized"), "got: {err}");
assert_eq!(requests.len(), 1);
}
#[test]
fn endpoint_batches_65_inputs_as_64_plus_1_and_preserves_global_order() {
let (url, requests_rx, server) = mock_endpoint(vec![
MockReply::Embeddings("resolved-model"),
MockReply::Embeddings("resolved-model"),
]);
let embedder = EndpointEmbedder::new(
url,
"requested-model".into(),
None,
String::new(),
String::new(),
)
.unwrap();
let inputs = (0..65).map(|index| index.to_string()).collect::<Vec<_>>();
let embedded = embedder.embed_batch_with_identity(&inputs).unwrap();
let requests = requests_rx.recv_timeout(Duration::from_secs(5)).unwrap();
server.join().unwrap();
assert_eq!(
requests
.iter()
.map(|request| request.inputs.len())
.collect::<Vec<_>>(),
vec![64, 1]
);
assert_eq!(
requests
.into_iter()
.flat_map(|request| request.inputs)
.collect::<Vec<_>>(),
inputs
);
assert_eq!(embedded.value.len(), 65);
assert!(embedded.fingerprint.contains("resolved-model"));
}
#[test]
fn second_chunk_failure_commits_neither_chunk_and_retry_sends_all_inputs() {
let (url, requests_rx, server) = mock_endpoint(vec![
MockReply::Embeddings("resolved-model"),
MockReply::Status(500),
MockReply::Embeddings("resolved-model"),
MockReply::Embeddings("resolved-model"),
]);
let mut registry = ToolRegistry::with_embedding(EmbeddingModel::Endpoint {
url,
model: "requested-model".into(),
api_key_env: None,
query_prefix: None,
doc_prefix: None,
});
for index in 0..65 {
registry.register(tool_for_endpoint(index));
}
assert!(registry.build_embeddings().is_err());
registry.build_embeddings().unwrap();
let requests = requests_rx.recv_timeout(Duration::from_secs(5)).unwrap();
server.join().unwrap();
assert_eq!(
requests
.iter()
.map(|request| request.inputs.len())
.collect::<Vec<_>>(),
vec![64, 1, 64, 1]
);
}
#[test]
fn endpoint_cache_separates_api_key_env_names_and_sends_each_bearer_token() {
const KEY_A: &str = "RATEL_CORE_ENDPOINT_TEST_KEY_A";
const KEY_B: &str = "RATEL_CORE_ENDPOINT_TEST_KEY_B";
unsafe {
std::env::set_var(KEY_A, "alpha-token");
std::env::set_var(KEY_B, "beta-token");
}
let (url, requests_rx, server) = mock_endpoint(vec![
MockReply::Embeddings("resolved-model"),
MockReply::Embeddings("resolved-model"),
]);
for (id, env_name) in [("a", KEY_A), ("b", KEY_B)] {
let mut registry = ToolRegistry::with_embedding(EmbeddingModel::Endpoint {
url: url.clone(),
model: "requested-model".into(),
api_key_env: Some(env_name.into()),
query_prefix: None,
doc_prefix: None,
});
registry.register(Tool {
id: id.into(),
name: id.into(),
description: "endpoint auth test".into(),
input_schema: serde_json::json!({}),
output_schema: serde_json::json!({}),
});
registry.build_embeddings().unwrap();
}
let requests = requests_rx.recv_timeout(Duration::from_secs(5)).unwrap();
server.join().unwrap();
unsafe {
std::env::remove_var(KEY_A);
std::env::remove_var(KEY_B);
}
assert_eq!(
requests
.into_iter()
.map(|request| request.authorization)
.collect::<Vec<_>>(),
vec![
Some("Bearer alpha-token".into()),
Some("Bearer beta-token".into())
]
);
}
#[test]
fn response_model_drift_is_hard_and_rebuild_adopts_the_new_identity() {
let (url, requests_rx, server) = mock_endpoint(vec![
MockReply::Embeddings("model-a"),
MockReply::Embeddings("model-b"),
MockReply::Embeddings("model-b"),
MockReply::Embeddings("model-b"),
]);
let mut registry = ToolRegistry::with_embedding(EmbeddingModel::Endpoint {
url,
model: "requested-model".into(),
api_key_env: None,
query_prefix: None,
doc_prefix: None,
});
registry.register(tool_for_endpoint(0));
registry.build_embeddings().unwrap();
assert!(matches!(
registry.search_with_method("tool", 1, Origin::Direct, SearchMethod::Semantic),
Err(EmbedderError::ModelMismatch { .. })
));
registry.rebuild_embeddings().unwrap();
assert_eq!(
registry
.search_with_method("tool", 1, Origin::Direct, SearchMethod::Semantic)
.unwrap()
.len(),
1
);
let requests = requests_rx.recv_timeout(Duration::from_secs(5)).unwrap();
server.join().unwrap();
assert_eq!(requests.len(), 4);
}
fn tool_for_endpoint(index: usize) -> Tool {
Tool {
id: format!("tool-{index}"),
name: format!("tool-{index}"),
description: format!("endpoint tool {index}"),
input_schema: serde_json::json!({}),
output_schema: serde_json::json!({}),
}
}
#[derive(Clone, Copy)]
enum MockReply {
Embeddings(&'static str),
Status(u16),
Oversized,
}
struct MockRequest {
inputs: Vec<String>,
authorization: Option<String>,
}
type MockServer = (
String,
mpsc::Receiver<Vec<MockRequest>>,
std::thread::JoinHandle<()>,
);
fn mock_endpoint(replies: Vec<MockReply>) -> MockServer {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let url = format!("http://{}/v1/embeddings", listener.local_addr().unwrap());
let (requests_tx, requests_rx) = mpsc::channel();
let server = std::thread::spawn(move || {
let deadline = Instant::now() + Duration::from_secs(5);
let mut replies = std::collections::VecDeque::from(replies);
let mut requests = Vec::new();
while let Some(reply) = replies.front().copied() {
match listener.accept() {
Ok((mut stream, _)) => {
stream.set_nonblocking(false).unwrap();
let (body, authorization) = read_http_request(&mut stream);
let inputs = body["input"]
.as_array()
.expect("input array")
.iter()
.map(|value| value.as_str().expect("string input").to_string())
.collect::<Vec<_>>();
write_mock_response(&mut stream, reply, inputs.len());
requests.push(MockRequest {
inputs,
authorization,
});
replies.pop_front();
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
if Instant::now() >= deadline {
break;
}
std::thread::sleep(Duration::from_millis(5));
}
Err(error) => panic!("accept failed: {error}"),
}
}
requests_tx.send(requests).unwrap();
});
(url, requests_rx, server)
}
fn write_mock_response(stream: &mut std::net::TcpStream, reply: MockReply, input_len: usize) {
let (status, response) = match reply {
MockReply::Embeddings(model) => {
let data = (0..input_len)
.map(|index| {
serde_json::json!({
"index": index,
"embedding": [1.0, 0.0]
})
})
.collect::<Vec<_>>();
(
"200 OK",
serde_json::json!({ "data": data, "model": model }).to_string(),
)
}
MockReply::Status(code) => {
("500 Internal Server Error", format!("{{\"code\":{code}}}"))
}
MockReply::Oversized => {
write_oversized_response(stream);
return;
}
};
write!(
stream,
"HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response.len(),
response
)
.unwrap();
}
fn write_oversized_response(stream: &mut std::net::TcpStream) {
const PREFIX: &[u8] = b"{\"data\":[],\"padding\":\"";
const SUFFIX: &[u8] = b"\"}";
const CHUNK: &[u8] = &[b'x'; 64 * 1024];
let padding_len = ENDPOINT_RESPONSE_LIMIT_BYTES;
let content_len = PREFIX.len() as u64 + padding_len + SUFFIX.len() as u64;
write!(
stream,
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {content_len}\r\nconnection: close\r\n\r\n"
)
.unwrap();
if stream.write_all(PREFIX).is_err() {
return;
}
let mut remaining = padding_len;
while remaining > 0 {
let len = remaining.min(CHUNK.len() as u64) as usize;
if stream.write_all(&CHUNK[..len]).is_err() {
return;
}
remaining -= len as u64;
}
let _ = stream.write_all(SUFFIX);
}
fn read_http_request(stream: &mut std::net::TcpStream) -> (serde_json::Value, Option<String>) {
stream
.set_read_timeout(Some(Duration::from_secs(2)))
.unwrap();
let mut request = Vec::new();
let mut buffer = [0_u8; 4096];
loop {
let read = stream.read(&mut buffer).unwrap();
assert!(read > 0, "connection closed before request body");
request.extend_from_slice(&buffer[..read]);
if let Some(header_end) = request.windows(4).position(|window| window == b"\r\n\r\n") {
let body_start = header_end + 4;
let headers = std::str::from_utf8(&request[..header_end]).unwrap();
let content_len = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().unwrap())
})
.expect("content-length");
if request.len() >= body_start + content_len {
let authorization = headers.lines().find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("authorization")
.then(|| value.trim().to_string())
});
let body =
serde_json::from_slice(&request[body_start..body_start + content_len])
.unwrap();
return (body, authorization);
}
}
}
}
#[test]
fn mean_pool_averages_only_unmasked_tokens() {
let dev = Device::Cpu;
let seq = Tensor::new(&[[[1.0f32, 2.0], [3.0, 4.0]]], &dev).unwrap();
let all = Tensor::new(&[[1u32, 1]], &dev).unwrap();
assert_eq!(
mean_pool(&seq, &all).unwrap().to_vec1::<f32>().unwrap(),
vec![2.0, 3.0]
);
let first = Tensor::new(&[[1u32, 0]], &dev).unwrap();
assert_eq!(
mean_pool(&seq, &first).unwrap().to_vec1::<f32>().unwrap(),
vec![1.0, 2.0]
);
}
#[test]
fn mean_pool_batch_averages_each_row_over_its_own_unmasked_tokens() {
let dev = Device::Cpu;
let seq = Tensor::new(
&[[[1.0f32, 2.0], [3.0, 4.0]], [[5.0, 6.0], [7.0, 8.0]]],
&dev,
)
.unwrap();
let mask = Tensor::new(&[[1u32, 1], [1, 0]], &dev).unwrap();
assert_eq!(
mean_pool_batch(&seq, &mask)
.unwrap()
.to_vec2::<f32>()
.unwrap(),
vec![vec![2.0, 3.0], vec![5.0, 6.0]]
);
}
#[test]
fn parse_pooling_config_maps_cls_mean_and_none() {
assert_eq!(
parse_pooling_config(br#"{"pooling_mode_cls_token": true}"#),
Some(Pooling::Cls)
);
assert_eq!(
parse_pooling_config(br#"{"pooling_mode_mean_tokens": true}"#),
Some(Pooling::Mean)
);
assert_eq!(
parse_pooling_config(br#"{"pooling_mode_max_tokens": true}"#),
None
);
assert_eq!(parse_pooling_config(b"not json"), None);
}
#[test]
fn resolve_pooling_assumes_mean_and_flags_it() {
assert_eq!(resolve_pooling(Some(Pooling::Cls)), (Pooling::Cls, false));
assert_eq!(resolve_pooling(Some(Pooling::Mean)), (Pooling::Mean, false));
assert_eq!(resolve_pooling(None), (Pooling::Mean, true));
}
#[test]
fn is_not_found_distinguishes_missing_file_from_network_error() {
assert!(is_not_found("Http status client error (404 Not Found)"));
assert!(is_not_found("Entry Not Found"));
assert!(!is_not_found("error sending request: connection refused"));
}
#[test]
fn endpoint_missing_api_key_env_is_a_config_error() {
let e = EndpointEmbedder::new(
"http://localhost:11434/v1/embeddings".into(),
"nomic".into(),
Some("RATEL_TEST_DEFINITELY_UNSET_KEY".into()),
String::new(),
String::new(),
)
.unwrap();
let err = e.api_key().unwrap_err();
assert!(matches!(err, EmbedderError::Config { .. }));
assert!(err.to_string().contains("RATEL_TEST_DEFINITELY_UNSET_KEY"));
}
#[test]
#[ignore = "downloads the ~130 MB bge model; run with `cargo test -- --ignored`"]
fn embeds_to_unit_norm_384_vectors_deterministically() {
let e = embedder_for(&EmbeddingModel::Default)
.expect("load embedder")
.0;
let a = e.embed_doc("read a file from disk").expect("embed");
let b = e.embed_doc("read a file from disk").expect("embed");
assert_eq!(a.len(), 384, "bge-small is 384-dim");
assert_eq!(a, b, "same text must embed identically (determinism)");
let norm = a.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-3, "expected unit norm, got {norm}");
}
#[test]
#[ignore = "downloads the ~130 MB bge model; run with `cargo test -- --ignored`"]
fn embed_batch_matches_looping_embed_doc_cls() {
let e = embedder_for(&EmbeddingModel::Default)
.expect("load embedder")
.0;
let docs: Vec<String> = (0..EMBED_BATCH_CHUNK + 5)
.map(|i| format!("{}read a file from disk", "word ".repeat(i % 7)))
.collect();
let batched = e.embed_batch(&docs).expect("embed_batch");
let looped: Vec<Vec<f32>> = docs
.iter()
.map(|d| e.embed_doc(d).expect("embed_doc"))
.collect();
assert_eq!(
batched, looped,
"batched embedding must be bit-for-bit identical to the per-doc path"
);
}
#[test]
#[ignore = "downloads a mean-pooled model (gte-small); run with `cargo test -- --ignored`"]
fn embed_batch_matches_looping_embed_doc_mean() {
let model = EmbeddingModel::HuggingFace {
repo: "thenlper/gte-small".into(),
revision: None,
query_prefix: None,
doc_prefix: None,
pooling: None, download: true, };
let e = embedder_for(&model).expect("load gte-small").0;
let docs: Vec<String> = (0..EMBED_BATCH_CHUNK + 3)
.map(|i| format!("{}deploy the service", "x ".repeat(i % 5)))
.collect();
let batched = e.embed_batch(&docs).expect("embed_batch");
let looped: Vec<Vec<f32>> = docs
.iter()
.map(|d| e.embed_doc(d).expect("embed_doc"))
.collect();
assert_eq!(
batched, looped,
"mean-pooled batched embedding must be bit-for-bit identical to the per-doc path"
);
}
#[test]
#[ignore = "downloads the ~130 MB bge model; run with `cargo test -- --ignored`"]
fn query_prefix_changes_the_embedding() {
let e = embedder_for(&EmbeddingModel::Default)
.expect("load embedder")
.0;
let doc = e.embed_doc("delete a file").expect("embed");
let query = e.embed_query("delete a file").expect("embed");
assert_ne!(doc, query, "query instruction prefix must shift the vector");
}
#[test]
#[ignore = "downloads the ~130 MB bge model; run with `cargo test -- --ignored`"]
fn ranks_synonyms_above_lexically_unrelated_text() {
let e = embedder_for(&EmbeddingModel::Default)
.expect("load embedder")
.0;
let q = e.embed_query("remove a file").expect("embed");
let delete = e
.embed_doc("delete a path from the filesystem")
.expect("embed");
let weather = e
.embed_doc("get the current weather forecast")
.expect("embed");
let dot = |a: &[f32], b: &[f32]| a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
assert!(
dot(&q, &delete) > dot(&q, &weather),
"semantic match should beat an unrelated tool"
);
}
#[test]
#[ignore = "downloads a mean-pooled model (gte-small); run with `cargo test -- --ignored`"]
fn mean_pooled_model_ranks_synonyms_correctly() {
let model = EmbeddingModel::HuggingFace {
repo: "thenlper/gte-small".into(),
revision: None,
query_prefix: None,
doc_prefix: None,
pooling: None, download: true, };
let e = embedder_for(&model).expect("load gte-small").0;
let q = e.embed_query("remove a file").expect("embed");
let delete = e
.embed_doc("delete a path from the filesystem")
.expect("embed");
let weather = e
.embed_doc("get the current weather forecast")
.expect("embed");
let dot = |a: &[f32], b: &[f32]| a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
assert!(
dot(&q, &delete) > dot(&q, &weather),
"mean-pooled semantic match should beat an unrelated tool"
);
}
}