use super::types::{DEFAULT_EMBED_ONNX_BATCH, 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(())
}