use super::types::{
DEFAULT_EMBED_ONNX_BATCH, EMBED_BATCH_BYTE_BUDGET, EMBED_INPUT_BYTE_CAP,
embed_in_bounded_batches, resolve_embed_onnx_batch,
};
use crate::embedder::test_env::{EnvVarGuard, env_lock};
use anyhow::Result;
use std::cell::RefCell;
fn inputs(n: usize) -> Vec<String> {
(0..n).map(|i| format!("drawer {i}")).collect()
}
fn echo_vectors(chunk: &[String]) -> Vec<Vec<f32>> {
chunk
.iter()
.map(|s| {
let n: f32 = s
.rsplit(' ')
.next()
.and_then(|d| d.parse().ok())
.unwrap_or(-1.0);
vec![n, 0.0]
})
.collect()
}
#[derive(Default)]
struct CallLog {
calls: RefCell<Vec<(usize, usize)>>,
}
impl CallLog {
fn record(&self, chunk: &[String], ceiling: usize) {
self.calls.borrow_mut().push((chunk.len(), ceiling));
}
fn batch_sizes(&self) -> Vec<usize> {
self.calls.borrow().iter().map(|(len, _)| *len).collect()
}
fn ceilings(&self) -> Vec<usize> {
self.calls.borrow().iter().map(|(_, c)| *c).collect()
}
}
#[test]
fn bounded_batches_never_exceed_the_ceiling() {
let log = CallLog::default();
let texts = inputs(600);
let out = embed_in_bounded_batches(&texts, DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
Ok(echo_vectors(chunk))
})
.expect("bounded embed must succeed");
let sizes = log.batch_sizes();
let full = 600 / DEFAULT_EMBED_ONNX_BATCH;
let remainder = 600 % DEFAULT_EMBED_ONNX_BATCH;
let mut expected = vec![DEFAULT_EMBED_ONNX_BATCH; full];
if remainder > 0 {
expected.push(remainder);
}
assert_eq!(
sizes, expected,
"600 inputs must reach ONNX as bounded batches, not one call of 600"
);
assert!(
sizes.iter().all(|n| *n <= DEFAULT_EMBED_ONNX_BATCH),
"no ONNX batch may exceed the ceiling: {sizes:?}"
);
assert_eq!(
sizes.iter().sum::<usize>(),
600,
"chunking must not drop or duplicate an input"
);
assert!(
log.ceilings()
.iter()
.all(|c| *c == DEFAULT_EMBED_ONNX_BATCH),
"every call must carry the ceiling itself, so fastembed's own \
batch_size is pinned and a dynamically-quantised model is never handed \
a batch_size below its input count: {:?}",
log.ceilings()
);
assert_eq!(out.len(), 600, "one vector per input");
}
#[test]
fn bounded_batches_preserve_input_order_and_count() {
let texts = inputs(600);
let out = embed_in_bounded_batches(&texts, DEFAULT_EMBED_ONNX_BATCH, |chunk, _| {
Ok(echo_vectors(chunk))
})
.expect("bounded embed must succeed");
assert_eq!(out.len(), texts.len());
for (i, vector) in out.iter().enumerate() {
assert_eq!(
vector[0], i as f32,
"vector {i} came back out of input order"
);
}
}
#[test]
fn a_short_input_still_makes_one_call() {
let log = CallLog::default();
let texts = inputs(5);
embed_in_bounded_batches(&texts, DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
Ok(echo_vectors(chunk))
})
.expect("bounded embed must succeed");
assert_eq!(log.batch_sizes(), vec![5]);
}
#[test]
fn an_empty_input_makes_no_call() {
let log = CallLog::default();
let out = embed_in_bounded_batches(&[], DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
Ok(echo_vectors(chunk))
})
.expect("an empty batch must succeed");
assert!(log.batch_sizes().is_empty(), "no input, no ONNX call");
assert!(out.is_empty());
}
#[test]
fn a_mid_batch_error_fails_the_whole_call() {
let log = CallLog::default();
let texts = inputs(600);
let err = embed_in_bounded_batches(&texts, DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
if log.batch_sizes().len() == 3 {
anyhow::bail!("ORT session run failed");
}
Ok(echo_vectors(chunk))
})
.map(|v| v.len())
.expect_err("a failing chunk must fail the whole call");
assert!(
format!("{err:#}").contains("ORT session run failed"),
"the underlying error must survive: {err:#}"
);
assert_eq!(
log.batch_sizes().len(),
3,
"the call must stop at the failing chunk, not run the remaining 35"
);
}
#[test]
fn a_short_chunk_result_fails_the_whole_call() {
let log = CallLog::default();
let texts = inputs(600);
let err = embed_in_bounded_batches(&texts, DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
let mut vectors = echo_vectors(chunk);
if log.batch_sizes().len() == 2 {
vectors.pop();
}
Ok(vectors)
})
.map(|v| v.len())
.expect_err("a short chunk result must fail the whole call");
let rendered = format!("{err:#}");
assert!(
rendered.contains("15") && rendered.contains("16"),
"the error must name the mismatched counts: {rendered}"
);
}
#[test]
fn a_zero_ceiling_is_clamped_to_one() {
let log = CallLog::default();
let texts = inputs(3);
embed_in_bounded_batches(&texts, 0, |chunk, ceiling| {
log.record(chunk, ceiling);
Ok(echo_vectors(chunk))
})
.expect("bounded embed must succeed");
assert_eq!(log.batch_sizes(), vec![1, 1, 1]);
assert_eq!(log.ceilings(), vec![1, 1, 1]);
}
#[test]
fn embed_onnx_batch_defaults_when_unset() {
let _g = env_lock();
let _e = EnvVarGuard::apply("TRUSTY_EMBED_ONNX_BATCH", None);
assert_eq!(resolve_embed_onnx_batch(), DEFAULT_EMBED_ONNX_BATCH);
}
#[test]
fn embed_onnx_batch_reads_env() {
let _g = env_lock();
let _e = EnvVarGuard::apply("TRUSTY_EMBED_ONNX_BATCH", Some(" 64 "));
assert_eq!(resolve_embed_onnx_batch(), 64);
}
#[test]
fn embed_onnx_batch_defaults_on_garbage() {
let _g = env_lock();
let _e = EnvVarGuard::apply("TRUSTY_EMBED_ONNX_BATCH", Some("lots"));
assert_eq!(resolve_embed_onnx_batch(), DEFAULT_EMBED_ONNX_BATCH);
}
#[test]
fn embed_onnx_batch_defaults_on_zero() {
let _g = env_lock();
let _e = EnvVarGuard::apply("TRUSTY_EMBED_ONNX_BATCH", Some("0"));
assert_eq!(resolve_embed_onnx_batch(), DEFAULT_EMBED_ONNX_BATCH);
}
#[test]
fn a_resolved_env_ceiling_drives_the_chunking() -> Result<()> {
let _g = env_lock();
let _e = EnvVarGuard::apply("TRUSTY_EMBED_ONNX_BATCH", Some("100"));
let log = CallLog::default();
let texts = inputs(600);
embed_in_bounded_batches(&texts, resolve_embed_onnx_batch(), |chunk, ceiling| {
log.record(chunk, ceiling);
Ok(echo_vectors(chunk))
})?;
assert_eq!(log.batch_sizes(), vec![100; 6]);
Ok(())
}
#[test]
fn long_inputs_split_under_the_byte_budget() {
let long: Vec<String> = (0..16)
.map(|i| format!("{i:04}{}", "x".repeat(1996)))
.collect();
let log = CallLog::default();
let seen = RefCell::new(Vec::new());
let out = embed_in_bounded_batches(&long, DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
seen.borrow_mut().extend(chunk.iter().cloned());
Ok(echo_vectors(chunk))
})
.expect("budgeted embed must succeed");
let per_call = EMBED_BATCH_BYTE_BUDGET / 2000;
assert_eq!(
log.batch_sizes(),
vec![per_call; 16 / per_call],
"2000-byte inputs must split so count × longest stays within the budget"
);
assert_eq!(
*seen.borrow(),
long,
"inputs must reach ONNX unaltered and in order"
);
assert_eq!(out.len(), 16);
assert!(
log.ceilings()
.iter()
.all(|c| *c == DEFAULT_EMBED_ONNX_BATCH),
"the ceiling handed to fastembed is unchanged by the budget"
);
let mut mixed = inputs(15);
mixed.push("y".repeat(EMBED_INPUT_BYTE_CAP * 4));
let log = CallLog::default();
embed_in_bounded_batches(&mixed, DEFAULT_EMBED_ONNX_BATCH, |chunk, ceiling| {
log.record(chunk, ceiling);
Ok(echo_vectors(chunk))
})
.expect("budgeted embed must succeed");
assert_eq!(
log.batch_sizes(),
vec![15, 1],
"one long input must not pad fifteen short ones to its length"
);
}