use super::*;
use crate::model::EmbeddingModel;
#[test]
fn test_max_batch_size_constant() {
assert_eq!(DEFAULT_MAX_BATCH_SIZE, 1000);
}
#[test]
fn test_max_text_bytes_constant() {
assert_eq!(MAX_TEXT_BYTES, 32768);
}
#[test]
#[allow(deprecated)]
fn test_max_text_chars_deprecated_alias_matches() {
assert_eq!(MAX_TEXT_CHARS, MAX_TEXT_BYTES);
}
#[cfg(feature = "native")]
mod native_tests {
use super::*;
use crate::{EmbeddingModel, ModelConfig};
#[test]
fn test_native_service_supports_only_loaded_model() {
let service = NativeEmbeddingService::default();
assert!(service.supports_model(EmbeddingModel::BgeSmallEnV15));
assert!(!service.supports_model(EmbeddingModel::BgeBaseEnV15));
assert!(!service.supports_model(EmbeddingModel::BgeLargeEnV15));
assert!(!service.supports_model(EmbeddingModel::MultilingualE5Small));
assert!(!service.supports_model(EmbeddingModel::MultilingualE5Base));
assert!(!service.supports_model(EmbeddingModel::Qwen3Embedding0_6B));
assert!(!service.supports_model(EmbeddingModel::Qwen3Embedding4B));
assert!(!service.supports_model(EmbeddingModel::TextEmbedding3Small));
}
#[test]
fn test_native_service_with_model_supports_only_that_model() {
let service = NativeEmbeddingService::with_model(EmbeddingModel::MultilingualE5Small);
assert!(service.supports_model(EmbeddingModel::MultilingualE5Small));
assert!(!service.supports_model(EmbeddingModel::BgeSmallEnV15));
}
#[test]
fn test_native_service_name() {
let service = NativeEmbeddingService::default();
assert_eq!(service.name(), "native-bert");
}
#[test]
fn test_native_service_with_model_config_qwen_default() {
use crate::model::ModelConfig;
let cfg = ModelConfig::new(EmbeddingModel::Qwen3Embedding4B);
let service = NativeEmbeddingService::with_model_config(cfg).unwrap();
assert!(service.supports_model(EmbeddingModel::Qwen3Embedding4B));
assert!(!service.supports_model(EmbeddingModel::Qwen3Embedding0_6B));
}
#[test]
fn test_native_service_with_model_config_invalid_dim_rejected() {
use crate::model::ModelConfig;
let cfg = ModelConfig {
model: EmbeddingModel::BgeSmallEnV15,
output_dim: Some(128),
};
assert!(NativeEmbeddingService::with_model_config(cfg).is_err());
}
#[test]
fn test_native_service_model_config_returns_configured_dim() {
let cfg = ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(1024)).unwrap();
let service = NativeEmbeddingService::with_model_config(cfg).unwrap();
let returned = service.model_config(EmbeddingModel::Qwen3Embedding4B);
assert_eq!(returned.output_dim, Some(1024));
assert_eq!(returned.dimensions(), 1024);
}
#[test]
fn test_native_service_model_config_unknown_model_returns_native() {
let service = NativeEmbeddingService::default(); let returned = service.model_config(EmbeddingModel::BgeBaseEnV15);
assert_eq!(returned.output_dim, None);
assert_eq!(returned.dimensions(), 768);
}
static ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn test_with_model_from_env_absent_returns_native_dim() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
let svc =
NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding4B).unwrap();
assert_eq!(
svc.model_config(EmbeddingModel::Qwen3Embedding4B)
.output_dim,
None
);
assert_eq!(
svc.model_config(EmbeddingModel::Qwen3Embedding4B)
.dimensions(),
2560
);
}
#[test]
fn test_with_model_from_env_dim_1024() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::set_var("LATTICE_EMBED_DIM", "1024") };
let result = NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding4B);
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
let svc = result.unwrap();
let cfg = svc.model_config(EmbeddingModel::Qwen3Embedding4B);
assert_eq!(cfg.output_dim, Some(1024));
assert_eq!(cfg.dimensions(), 1024);
}
#[test]
fn test_with_model_from_env_qwen_06b_dim_512() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::set_var("LATTICE_EMBED_DIM", "512") };
let result =
NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding0_6B);
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
let svc = result.unwrap();
let cfg = svc.model_config(EmbeddingModel::Qwen3Embedding0_6B);
assert_eq!(cfg.output_dim, Some(512));
assert_eq!(cfg.dimensions(), 512);
}
#[test]
fn test_with_model_from_env_invalid_value_returns_error() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::set_var("LATTICE_EMBED_DIM", "not_a_number") };
let result = NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding4B);
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
assert!(result.is_err(), "non-numeric LATTICE_EMBED_DIM must fail");
}
#[test]
fn test_with_model_from_env_dim_below_minimum_returns_error() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::set_var("LATTICE_EMBED_DIM", "16") };
let result = NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding4B);
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
assert!(result.is_err(), "dim < 32 must be rejected");
}
#[test]
fn test_with_model_from_env_dim_above_native_returns_error() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::set_var("LATTICE_EMBED_DIM", "9999") };
let result = NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding4B);
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
assert!(result.is_err(), "dim > native must be rejected");
}
#[test]
fn test_with_model_from_env_empty_string_treated_as_absent() {
let _g = ENV_MUTEX.lock().unwrap();
unsafe { std::env::set_var("LATTICE_EMBED_DIM", "") };
let result = NativeEmbeddingService::with_model_from_env(EmbeddingModel::Qwen3Embedding4B);
unsafe { std::env::remove_var("LATTICE_EMBED_DIM") };
let svc = result.unwrap();
assert_eq!(
svc.model_config(EmbeddingModel::Qwen3Embedding4B)
.output_dim,
None
);
}
#[tokio::test]
async fn test_native_service_embed_wrong_model_returns_error() {
let service = NativeEmbeddingService::default(); let texts = vec!["hello".to_string()];
let result = service.embed(&texts, EmbeddingModel::BgeBaseEnV15).await;
assert!(result.is_err(), "expected error for wrong model");
let err = result.unwrap_err().to_string();
assert!(
err.contains("BgeBaseEnV15") || err.contains("requested model"),
"error should mention the model mismatch, got: {err}"
);
}
#[tokio::test]
async fn test_native_service_ascii_text_too_long_reports_bytes() {
let text = "a".repeat(36_556);
assert_eq!(text.len(), 36_556);
assert_eq!(text.chars().count(), 36_556);
let service = NativeEmbeddingService::default();
let texts = vec![text];
let err = service
.embed(&texts, EmbeddingModel::BgeSmallEnV15)
.await
.expect_err("text over MAX_TEXT_BYTES must be rejected");
let msg = err.to_string();
assert!(msg.contains("36556"), "got: {msg}");
assert!(msg.contains("bytes"), "got: {msg}");
}
#[tokio::test]
async fn test_native_service_multibyte_text_too_long_uses_byte_count() {
let text = format!("{}{}", "é".repeat(78), "a".repeat(32_690));
assert_eq!(text.chars().count(), 32_768);
assert_eq!(text.len(), 32_846);
let service = NativeEmbeddingService::default();
let texts = vec![text];
let err = service
.embed(&texts, EmbeddingModel::BgeSmallEnV15)
.await
.expect_err("32846-byte text must be rejected under byte semantics");
let msg = err.to_string();
assert!(msg.contains("32846"), "got: {msg}");
assert!(
!msg.contains("32768 bytes exceeds") && !msg.contains("32768 chars exceeds"),
"reported length must be the byte count (32846), not the char count: {msg}"
);
}
#[tokio::test]
async fn test_cached_service_multibyte_text_too_long_uses_byte_count() {
use crate::service::CachedEmbeddingService;
use std::sync::Arc;
let text = format!("{}{}", "é".repeat(78), "a".repeat(32_690));
assert_eq!(text.chars().count(), 32_768);
assert_eq!(text.len(), 32_846);
let service =
CachedEmbeddingService::with_default_cache(Arc::new(NativeEmbeddingService::default()));
let texts = vec![text];
let err = service
.embed(&texts, EmbeddingModel::BgeSmallEnV15)
.await
.expect_err("32846-byte text must be rejected under byte semantics");
let msg = err.to_string();
assert!(msg.contains("32846"), "got: {msg}");
}
#[tokio::test]
async fn test_role_path_validates_caller_text_not_prepared_text() {
let model = EmbeddingModel::BgeSmallEnV15;
assert!(
model.max_instruction_bytes() > 0,
"test needs a model that prepends a query instruction"
);
let text = "a".repeat(MAX_TEXT_BYTES);
assert_eq!(text.len(), MAX_TEXT_BYTES);
let service = NativeEmbeddingService::default();
if let Err(e) = service.embed_query(&[text], model).await {
assert!(
!matches!(e, crate::error::EmbedError::TextTooLong { .. }),
"caller text at exactly the cap must not be rejected for length: {e}"
);
}
}
#[tokio::test]
async fn test_role_path_still_rejects_caller_text_over_the_cap() {
let text = "a".repeat(MAX_TEXT_BYTES + 1);
let service = NativeEmbeddingService::default();
let err = service
.embed_query(&[text], EmbeddingModel::BgeSmallEnV15)
.await
.expect_err("caller text over the cap must be rejected");
assert!(
matches!(err, crate::error::EmbedError::TextTooLong { max, .. } if max == MAX_TEXT_BYTES),
"must report the published cap, got: {err}"
);
}
#[tokio::test]
async fn test_cached_role_path_validates_caller_text() {
use crate::service::CachedEmbeddingService;
use std::sync::Arc;
let text = "a".repeat(MAX_TEXT_BYTES);
let service =
CachedEmbeddingService::with_default_cache(Arc::new(NativeEmbeddingService::default()));
if let Err(e) = service
.embed_query(&[text], EmbeddingModel::BgeSmallEnV15)
.await
{
assert!(
!matches!(e, crate::error::EmbedError::TextTooLong { .. }),
"caller text at exactly the cap must not be rejected for length: {e}"
);
}
}
}
#[test]
fn test_max_instruction_bytes_matches_the_longest_instruction() {
for model in [
EmbeddingModel::BgeSmallEnV15,
EmbeddingModel::MultilingualE5Small,
EmbeddingModel::Qwen3Embedding0_6B,
EmbeddingModel::AllMiniLmL6V2,
] {
let q = model.query_instruction().map_or(0, str::len);
let d = model.document_instruction().map_or(0, str::len);
assert_eq!(
model.max_instruction_bytes(),
q.max(d),
"{model:?} reports the wrong instruction bound"
);
}
assert_eq!(EmbeddingModel::AllMiniLmL6V2.max_instruction_bytes(), 0);
}
mod external_impl_contract {
use crate::error::{EmbedError, Result};
use crate::model::EmbeddingModel;
use crate::service::{
EmbeddingRole, EmbeddingService, MAX_TEXT_BYTES, apply_prefix, validate_texts,
};
use async_trait::async_trait;
fn exact_cap_embed(texts: &[String]) -> Result<Vec<Vec<f32>>> {
for t in texts {
if t.len() > MAX_TEXT_BYTES {
return Err(EmbedError::TextTooLong {
length: t.len(),
max: MAX_TEXT_BYTES,
});
}
}
Ok(texts.iter().map(|_| vec![0.0]).collect())
}
struct ExactCapNoOverride;
#[async_trait]
impl EmbeddingService for ExactCapNoOverride {
async fn embed(&self, texts: &[String], _model: EmbeddingModel) -> Result<Vec<Vec<f32>>> {
exact_cap_embed(texts)
}
fn supports_model(&self, _model: EmbeddingModel) -> bool {
true
}
fn name(&self) -> &'static str {
"exact-cap-no-override"
}
}
struct ExactCapWithOverride;
#[async_trait]
impl EmbeddingService for ExactCapWithOverride {
async fn embed(&self, texts: &[String], _model: EmbeddingModel) -> Result<Vec<Vec<f32>>> {
exact_cap_embed(texts)
}
async fn embed_with_role(
&self,
texts: &[String],
model: EmbeddingModel,
role: EmbeddingRole,
) -> Result<Vec<Vec<f32>>> {
validate_texts(texts)?;
let prepared = apply_prefix(texts, role.instruction(model));
let backstop = MAX_TEXT_BYTES + model.max_instruction_bytes();
for t in &prepared {
if t.len() > backstop {
return Err(EmbedError::TextTooLong {
length: t.len(),
max: MAX_TEXT_BYTES,
});
}
}
Ok(prepared.iter().map(|_| vec![0.0]).collect())
}
fn supports_model(&self, _model: EmbeddingModel) -> bool {
true
}
fn name(&self) -> &'static str {
"exact-cap-with-override"
}
}
#[tokio::test]
async fn default_role_forwards_prepared_text_to_embed() {
let text = "a".repeat(MAX_TEXT_BYTES);
assert_eq!(text.len(), MAX_TEXT_BYTES);
let err = ExactCapNoOverride
.embed_query(&[text], EmbeddingModel::BgeSmallEnV15)
.await
.expect_err("default forwards the lengthened string through exact-cap embed");
assert!(
matches!(err, EmbedError::TextTooLong { .. }),
"expected the inherited exact-cap embed to reject prepared text, got: {err}"
);
}
#[tokio::test]
async fn overriding_the_role_method_preserves_the_caller_cap() {
let text = "a".repeat(MAX_TEXT_BYTES);
let out = ExactCapWithOverride
.embed_query(&[text], EmbeddingModel::BgeSmallEnV15)
.await
.expect("overriding embed_with_role admits cap-sized caller text");
assert_eq!(out.len(), 1, "one embedding per input");
}
#[cfg(feature = "native")]
#[tokio::test]
async fn cached_wrapper_preserves_external_role_override() {
use crate::service::CachedEmbeddingService;
use std::sync::Arc;
let text = "a".repeat(MAX_TEXT_BYTES);
let service = CachedEmbeddingService::new(Arc::new(ExactCapWithOverride), 128);
let out = service
.embed_query(&[text], EmbeddingModel::BgeSmallEnV15)
.await
.expect("cached miss must preserve the inner role override");
assert_eq!(out.len(), 1, "one embedding per input");
}
#[tokio::test]
async fn over_cap_caller_text_reports_published_cap_on_both_shapes() {
let text = "a".repeat(MAX_TEXT_BYTES + 1);
let one = std::slice::from_ref(&text);
for res in [
ExactCapNoOverride
.embed_query(one, EmbeddingModel::BgeSmallEnV15)
.await,
ExactCapWithOverride
.embed_query(one, EmbeddingModel::BgeSmallEnV15)
.await,
] {
let err = res.expect_err("over-cap caller text must be rejected");
assert!(
matches!(err, EmbedError::TextTooLong { max, .. } if max == MAX_TEXT_BYTES),
"must report the published cap, got: {err}"
);
}
}
}
#[test]
fn test_e5_query_instruction() {
assert_eq!(
EmbeddingModel::MultilingualE5Small.query_instruction(),
Some("query: "),
"E5 small must return 'query: ' prefix"
);
assert_eq!(
EmbeddingModel::MultilingualE5Base.query_instruction(),
Some("query: "),
"E5 base must return 'query: ' prefix"
);
}
#[test]
fn test_e5_document_instruction() {
assert_eq!(
EmbeddingModel::MultilingualE5Small.document_instruction(),
Some("passage: "),
"E5 small must return 'passage: ' document prefix"
);
assert_eq!(
EmbeddingModel::MultilingualE5Base.document_instruction(),
Some("passage: "),
"E5 base must return 'passage: ' document prefix"
);
}
#[test]
fn test_bge_query_instruction() {
assert_eq!(
EmbeddingModel::BgeSmallEnV15.query_instruction(),
Some("Represent this sentence for searching relevant passages: "),
"BGE small must return the retrieval query instruction"
);
assert_eq!(
EmbeddingModel::BgeBaseEnV15.query_instruction(),
Some("Represent this sentence for searching relevant passages: "),
"BGE base must return the retrieval query instruction"
);
assert_eq!(
EmbeddingModel::BgeLargeEnV15.query_instruction(),
Some("Represent this sentence for searching relevant passages: "),
"BGE large must return the retrieval query instruction"
);
assert_eq!(
EmbeddingModel::BgeSmallEnV15.document_instruction(),
None,
"BGE passages must stay unprefixed"
);
}
#[test]
fn test_bge_minilm_no_document_instruction() {
assert_eq!(
EmbeddingModel::BgeSmallEnV15.document_instruction(),
None,
"BGE small must not have document prefix"
);
assert_eq!(EmbeddingModel::BgeBaseEnV15.document_instruction(), None);
assert_eq!(EmbeddingModel::BgeLargeEnV15.document_instruction(), None);
assert_eq!(EmbeddingModel::AllMiniLmL6V2.document_instruction(), None);
assert_eq!(
EmbeddingModel::ParaphraseMultilingualMiniLmL12V2.document_instruction(),
None
);
}
#[test]
fn test_qwen_no_document_instruction() {
assert_eq!(
EmbeddingModel::Qwen3Embedding0_6B.document_instruction(),
None,
"Qwen document side uses raw text"
);
}
#[test]
fn test_apply_prefix_some() {
let texts = vec!["hello".to_string(), "world".to_string()];
let result = apply_prefix(&texts, Some("query: "));
assert_eq!(result, vec!["query: hello", "query: world"]);
}
#[test]
fn test_apply_prefix_none() {
let texts = vec!["hello".to_string()];
let result = apply_prefix(&texts, None);
assert_eq!(result, texts);
}
#[test]
fn test_embedding_role_cache_tags_distinct() {
let q = EmbeddingRole::Query.cache_tag();
let p = EmbeddingRole::Passage.cache_tag();
let g = EmbeddingRole::Generic.cache_tag();
assert_ne!(q, p);
assert_ne!(q, g);
assert_ne!(p, g);
}