use super::*;
#[test]
fn options_default_equals_new_and_is_cpu_and_gpu() {
assert_eq!(TextEmbedderOptions::default(), TextEmbedderOptions::new());
assert_eq!(TextEmbedderOptions::new().compute(), DEFAULT_TEXT_COMPUTE);
assert_eq!(DEFAULT_TEXT_COMPUTE, ComputeUnits::CpuAndGpu);
}
#[test]
fn options_with_and_set_compute() {
let opts = TextEmbedderOptions::new().with_compute(ComputeUnits::CpuAndNeuralEngine);
assert_eq!(opts.compute(), ComputeUnits::CpuAndNeuralEngine);
let mut opts = TextEmbedderOptions::new();
opts.set_compute(ComputeUnits::CpuOnly);
assert_eq!(opts.compute(), ComputeUnits::CpuOnly);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_roundtrip() {
let opts = TextEmbedderOptions::new().with_compute(ComputeUnits::All);
let json = serde_json::to_string(&opts).unwrap();
assert!(json.contains("all"), "serialized: {json}");
let back: TextEmbedderOptions = serde_json::from_str(&json).unwrap();
assert_eq!(back, opts);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_defaults_missing_compute() {
let back: TextEmbedderOptions = serde_json::from_str("{}").unwrap();
assert_eq!(back, TextEmbedderOptions::new());
}
const T: usize = 64;
#[test]
fn build_window_right_pad_places_prefix_and_pads_suffix() {
let ids = [10u32, 20, 30];
let w = build_window(&ids, 7, PadSide::Right, T).expect("window");
assert_eq!(w.len(), T);
assert_eq!(&w[..3], &[10i32, 20, 30]);
assert!(w[3..].iter().all(|&x| x == 7), "suffix must be pad_id");
}
#[test]
fn build_window_left_pad_places_suffix_and_pads_prefix() {
let ids = [10u32, 20, 30];
let w = build_window(&ids, 7, PadSide::Left, T).expect("window");
assert_eq!(w.len(), T);
assert!(w[..T - 3].iter().all(|&x| x == 7), "prefix must be pad_id");
assert_eq!(&w[T - 3..], &[10i32, 20, 30]);
}
#[test]
fn build_window_full_window_has_no_pad() {
let ids: Vec<u32> = (0..T as u32).collect();
let w_right = build_window(&ids, 7, PadSide::Right, T).expect("full window");
let w_left = build_window(&ids, 7, PadSide::Left, T).expect("full window");
let expected: Vec<i32> = (0..T as i32).collect();
assert_eq!(w_right, expected);
assert_eq!(w_left, expected);
}
#[test]
fn build_window_rejects_overlong_ids_with_typed_error() {
let overlong = vec![1u32; T + 1];
match build_window(&overlong, 0, PadSide::Right, T) {
Err(Error::TokenCount(e)) => {
assert_eq!(e.got(), T + 1);
assert_eq!(e.max(), T);
}
other => panic!("expected TokenCount, got {other:?}"),
}
}
#[test]
fn build_window_rejects_out_of_range_token_id() {
match build_window(&[u32::MAX], 0, PadSide::Right, T) {
Err(Error::TokenIdRange(id)) => assert_eq!(id, u32::MAX),
other => panic!("expected TokenIdRange, got {other:?}"),
}
}
const TINY_TOKENIZER: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": { "type": "Whitespace" },
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": { "<pad>": 0, "a": 1, "b": 2, "c": 3, "d": 4, "e": 5 },
"unk_token": "<pad>"
}
}"#;
#[test]
fn configured_tokenizer_truncates_to_window_and_disables_padding() {
let max_tokens = 4;
let tok = configured_tokenizer_from_bytes(TINY_TOKENIZER.as_bytes(), max_tokens)
.expect("configure tiny tokenizer");
let ids = tok
.encode("a b c d e a b c", false)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(ids.len(), max_tokens, "must truncate to the window");
let short = tok.encode("a b", false).expect("encode").get_ids().to_vec();
assert_eq!(
short,
vec![1u32, 2],
"short input stays unpadded by the tokenizer"
);
let window = build_window(&short, 0, PadSide::Right, max_tokens).expect("window");
assert_eq!(window, vec![1i32, 2, 0, 0]);
}
fn tokenizer_bytes_with_special_overhead(added: usize) -> Vec<u8> {
use tokenizers::processors::template::{SpecialToken, TemplateProcessing};
let mut tokenizer =
Tokenizer::from_bytes(TINY_TOKENIZER.as_bytes()).expect("load the tiny tokenizer");
let special = SpecialToken::new(
"<sp>".to_string(),
vec![0u32; added],
vec!["<pad>".to_string(); added],
)
.expect("ids and tokens are the same length");
let template = TemplateProcessing::builder()
.try_single("<sp> $A")
.expect("single template")
.try_pair("<sp> $A $B")
.expect("pair template")
.special_tokens(vec![special])
.build()
.expect("build the template post-processor");
tokenizer.with_post_processor(Some(template));
tokenizer
.to_string(false)
.expect("serialize the tokenizer")
.into_bytes()
}
#[test]
fn overhead_fixture_installs_the_claimed_special_token_count() {
use tokenizers::PostProcessor;
for added in [1usize, 2, 5] {
let bytes = tokenizer_bytes_with_special_overhead(added);
let tok = Tokenizer::from_bytes(&bytes).expect("reload the fixture");
let post = tok.get_post_processor().expect("the fixture has one");
assert_eq!(post.added_tokens(false), added, "single-sequence overhead");
}
}
#[test]
fn configure_tokenizer_refuses_overhead_over_the_window() {
let bytes = tokenizer_bytes_with_special_overhead(2);
match configured_tokenizer_from_bytes(&bytes, 1) {
Err(Error::SpecialTokenOverhead(overhead)) => {
assert_eq!(overhead.added(), 2);
assert_eq!(overhead.window(), 1);
}
other => panic!("expected SpecialTokenOverhead, got {other:?}"),
}
}
#[test]
fn configure_tokenizer_refuses_overhead_equal_to_the_window() {
let bytes = tokenizer_bytes_with_special_overhead(2);
match configured_tokenizer_from_bytes(&bytes, 2) {
Err(Error::SpecialTokenOverhead(overhead)) => {
assert_eq!(overhead.added(), 2);
assert_eq!(overhead.window(), 2);
}
other => panic!("expected SpecialTokenOverhead, got {other:?}"),
}
let mut raw = Tokenizer::from_bytes(&bytes).expect("load the fixture");
raw
.with_truncation(Some(TruncationParams {
max_length: 2,
strategy: TruncationStrategy::LongestFirst,
stride: 0,
direction: TruncationDirection::Right,
}))
.expect("added == window does not overflow");
let ids = raw
.encode("a b c d e", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(
ids,
vec![0u32, 0],
"at a zero effective window every input is the specials alone"
);
}
#[test]
fn configure_tokenizer_accepts_overhead_below_the_window() {
let bytes = tokenizer_bytes_with_special_overhead(2);
let tok = configured_tokenizer_from_bytes(&bytes, 3).expect("2 specials fit a 3-token window");
let ids = tok
.encode("a b c d e", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(ids.len(), 3, "two specials plus one real token");
assert_eq!(ids[2], 1, "the real token survives (`a`)");
}
#[test]
fn configure_tokenizer_accepts_a_tokenizer_without_a_post_processor() {
for max_tokens in [1usize, 4, 64] {
configured_tokenizer_from_bytes(TINY_TOKENIZER.as_bytes(), max_tokens)
.expect("no post-processor is zero overhead");
}
}
#[test]
fn from_memory_accepts_a_real_tokenizer_past_the_guard() {
let real = TINY_TOKENIZER.as_bytes();
assert!(
ensure_not_placeholder(real).is_ok(),
"a real (non-sentinel) tokenizer must pass the placeholder guard"
);
match TextEmbedder::from_memory(
"/nonexistent/model.mlmodelc",
real,
TextEmbedderOptions::new(),
) {
Err(Error::Load(_)) => {}
other => panic!("expected Error::Load past the guard, got {other:?}"),
}
}
#[test]
fn artifact_tokenizer_path_is_the_bundle_sibling() {
assert_eq!(
artifact_tokenizer_path(Path::new(
"/m/siglip2-base-patch16-naflex-512/siglip2_text_64.mlmodelc"
)),
Path::new("/m/siglip2-base-patch16-naflex-512/tokenizer.json"),
);
assert_eq!(
artifact_tokenizer_path(Path::new("siglip2_text_64.mlmodelc")),
Path::new("tokenizer.json"),
);
}
#[test]
fn load_reports_a_missing_artifact_tokenizer() {
match TextEmbedder::load("/nonexistent/model.mlmodelc", TextEmbedderOptions::new()) {
Err(Error::ArtifactTokenizerRead(e)) => {
assert_eq!(e.path(), Path::new("/nonexistent/tokenizer.json"));
}
other => panic!("expected ArtifactTokenizerRead, got {other:?}"),
}
}
#[test]
fn load_guards_the_sidecar_it_reads() {
let dir = tempfile::tempdir().expect("tempdir");
let model_path = dir.path().join("siglip2_text_64.mlmodelc");
let tokenizer_path = dir.path().join("tokenizer.json");
let mut placeholder = br#"{"junk":""#.to_vec();
placeholder.extend_from_slice(PLACEHOLDER_SENTINEL);
placeholder.extend_from_slice(br#""}"#);
std::fs::write(&tokenizer_path, &placeholder).expect("write placeholder sidecar");
match TextEmbedder::load(&model_path, TextEmbedderOptions::new()) {
Err(Error::TokenizerPlaceholder) => {}
other => panic!("expected TokenizerPlaceholder for a placeholder sidecar, got {other:?}"),
}
std::fs::write(&tokenizer_path, TINY_TOKENIZER.as_bytes()).expect("write foreign sidecar");
match TextEmbedder::load(&model_path, TextEmbedderOptions::new()) {
Err(Error::ArtifactTokenizerIdentity(e)) => {
assert_eq!(e.path(), tokenizer_path.as_path());
assert_eq!(e.expected(), contract::TOKENIZER_SHA256_HEX);
assert_ne!(e.actual(), contract::TOKENIZER_SHA256_HEX);
}
other => panic!("expected ArtifactTokenizerIdentity for a foreign sidecar, got {other:?}"),
}
}
#[test]
fn placeholder_guard_accepts_real_tokenizer_bytes() {
assert!(ensure_not_placeholder(TINY_TOKENIZER.as_bytes()).is_ok());
}
#[test]
fn placeholder_guard_rejects_the_sentinel_buffer() {
let mut buf = br#"{"junk":""#.to_vec();
buf.extend_from_slice(PLACEHOLDER_SENTINEL);
buf.extend_from_slice(br#""}"#);
match ensure_not_placeholder(&buf) {
Err(Error::TokenizerPlaceholder) => {}
other => panic!("expected TokenizerPlaceholder for a sentinel buffer, got {other:?}"),
}
}
const CASE_COLLISION_TOKENIZER: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": { "type": "Whitespace" },
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": { "<pad>": 0, "a": 1, "b": 2, "A": 6 },
"unk_token": "<pad>"
}
}"#;
const REPLACE_NORMALIZER_TOKENIZER: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": { "type": "Replace", "pattern": { "String": "x" }, "content": "a" },
"pre_tokenizer": { "type": "Whitespace" },
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": { "<pad>": 0, "a": 1, "b": 2 },
"unk_token": "<pad>"
}
}"#;
#[test]
fn configured_tokenizer_lowercases_before_lookup() {
let tok =
configured_tokenizer_from_bytes(TINY_TOKENIZER.as_bytes(), 8).expect("configure tokenizer");
let ids = tok.encode("A B", false).expect("encode").get_ids().to_vec();
assert_eq!(
ids,
vec![1u32, 2],
"mixed case must lowercase to [a, b] ids"
);
}
#[test]
fn configured_tokenizer_prefers_lowercase_vocab_entry() {
let tok = configured_tokenizer_from_bytes(CASE_COLLISION_TOKENIZER.as_bytes(), 8)
.expect("configure tokenizer");
let ids = tok.encode("A", false).expect("encode").get_ids().to_vec();
assert_eq!(ids, vec![1u32], "must pick the lowercase entry, not id 6");
}
#[test]
fn configured_tokenizer_composes_ahead_of_existing_normalizer() {
let tok = configured_tokenizer_from_bytes(REPLACE_NORMALIZER_TOKENIZER.as_bytes(), 8)
.expect("configure tokenizer");
let ids = tok.encode("X b", false).expect("encode").get_ids().to_vec();
assert_eq!(
ids,
vec![1u32, 2],
"Lowercase must run before the loaded Replace normalizer"
);
}
use crate::{
AxisRange, FeatureInfo, ModelDescription, embeddings::siglip::error::contract_violation,
model::RawShapeConstraint,
};
const STAGED_TEXT_WINDOW: usize = 64;
fn fixed(name: &str, shape: &[usize], dtype: DataType) -> FeatureInfo {
multi_array(name, shape, dtype, false, 2, vec![shape.to_vec()], shape)
}
fn multi_array(
name: &str,
shape: &[usize],
dtype: DataType,
optional: bool,
raw_type: isize,
enumerated: Vec<Vec<usize>>,
pinned: &[usize],
) -> FeatureInfo {
FeatureInfo::from_parts(
name.to_string(),
shape.to_vec(),
Some(dtype),
optional,
Some(RawShapeConstraint::new(
raw_type,
enumerated,
pinned.iter().map(|d| AxisRange::new(*d, 1)).collect(),
)),
)
}
fn text_description(t: usize) -> ModelDescription {
ModelDescription::from_parts(
vec![fixed(names::INPUT_IDS, &[1, t], DataType::I32)],
vec![fixed(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
)
}
fn check(description: &ModelDescription) -> Result<()> {
crate::model::contract::check_load_contract(description, &text_contract())
.map_err(contract_violation)
}
#[test]
fn the_contract_accepts_the_staged_geometry() {
assert!(check(&text_description(STAGED_TEXT_WINDOW)).is_ok());
}
#[test]
fn the_contract_reads_back_whatever_window_the_graph_pins() {
for t in [1usize, 16, STAGED_TEXT_WINDOW, 512] {
let description = text_description(t);
assert!(check(&description).is_ok(), "window {t}");
assert_eq!(
description
.input(names::INPUT_IDS)
.expect("input_ids")
.shape()[1],
t,
"the window read back must be the one the graph pins"
);
}
}
#[test]
fn the_contract_refuses_a_flexible_window() {
let description = ModelDescription::from_parts(
vec![multi_array(
names::INPUT_IDS,
&[1, STAGED_TEXT_WINDOW],
DataType::I32,
false,
3,
Vec::new(),
&[1, STAGED_TEXT_WINDOW],
)],
vec![fixed(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::INPUT_IDS),
"{err}"
);
}
#[test]
fn a_zero_window_is_refused_by_the_contract() {
let description = text_description(0);
let violation = crate::model::contract::check_load_contract(&description, &text_contract())
.expect_err("`AnyFixed`'s zero clause refuses a window pinned at zero");
assert!(
matches!(
&violation,
crate::model::contract::ContractViolation::ZeroSizedAxis(zero)
if zero.feature() == names::INPUT_IDS
),
"{violation}"
);
}
#[test]
fn the_contract_refuses_a_wrong_dtype_or_projection_width() {
let float_ids = ModelDescription::from_parts(
vec![fixed(
names::INPUT_IDS,
&[1, STAGED_TEXT_WINDOW],
DataType::F32,
)],
vec![fixed(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
let err = check(&float_ids).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::INPUT_IDS),
"{err}"
);
let wrong_width = ModelDescription::from_parts(
vec![fixed(
names::INPUT_IDS,
&[1, STAGED_TEXT_WINDOW],
DataType::I32,
)],
vec![fixed(names::TEXT_FEATURES, &[1, 512], DataType::F32)],
Vec::new(),
);
let err = check(&wrong_width).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::TEXT_FEATURES),
"{err}"
);
}
#[test]
fn the_contract_refuses_an_extra_required_input() {
let description = ModelDescription::from_parts(
vec![
fixed(names::INPUT_IDS, &[1, STAGED_TEXT_WINDOW], DataType::I32),
fixed("attention_mask", &[1, STAGED_TEXT_WINDOW], DataType::I32),
],
vec![fixed(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
assert!(
matches!(check(&description), Err(Error::UnsatisfiableInput(name)) if name == "attention_mask"),
"{:?}",
check(&description)
);
}
#[test]
fn the_contract_accepts_an_extra_optional_input() {
let description = ModelDescription::from_parts(
vec![
fixed(names::INPUT_IDS, &[1, STAGED_TEXT_WINDOW], DataType::I32),
multi_array(
"attention_mask",
&[1, STAGED_TEXT_WINDOW],
DataType::I32,
true,
2,
vec![vec![1, STAGED_TEXT_WINDOW]],
&[1, STAGED_TEXT_WINDOW],
),
],
vec![fixed(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
Vec::new(),
);
assert!(check(&description).is_ok());
}
#[test]
fn the_contract_refuses_an_optional_features_output() {
let description = ModelDescription::from_parts(
vec![fixed(
names::INPUT_IDS,
&[1, STAGED_TEXT_WINDOW],
DataType::I32,
)],
vec![multi_array(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
true,
2,
vec![vec![1, EMBEDDING_DIM]],
&[1, EMBEDDING_DIM],
)],
Vec::new(),
);
let err = check(&description).unwrap_err();
assert!(
matches!(&err, Error::ContractMismatch(m) if m.feature() == names::TEXT_FEATURES),
"{err}"
);
}
#[test]
fn the_contract_refuses_a_graph_that_declares_state() {
let description = ModelDescription::from_parts(
vec![fixed(
names::INPUT_IDS,
&[1, STAGED_TEXT_WINDOW],
DataType::I32,
)],
vec![fixed(
names::TEXT_FEATURES,
&[1, EMBEDDING_DIM],
DataType::F32,
)],
vec![fixed("kv_cache", &[1, 8], DataType::F32)],
);
assert!(
matches!(check(&description), Err(Error::UnsatisfiableState(name)) if name == "kv_cache")
);
}
#[test]
fn the_text_contract_refuses_the_vendored_silero_bundle() {
let bundle = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../Models/vadkit/silero-vad-unified-256ms-v6.2.1.mlmodelc");
assert!(
bundle.is_dir(),
"the vendored silero bundle is committed, so this gate is NOT model-gated; \
looked for {}",
bundle.display()
);
let model = Model::load(&bundle, ComputeUnits::CpuOnly).expect("the committed bundle loads");
assert!(
model.description().input(names::INPUT_IDS).is_none(),
"silero declares no `input_ids`, which is what makes it this gate's model"
);
let violation = Checked::new(model, &text_contract())
.expect_err("silero does not satisfy the siglip text contract");
assert!(
matches!(&violation, crate::model::contract::ContractViolation::Missing(m)
if m.feature() == names::INPUT_IDS),
"expected `input_ids` missing, got {violation}"
);
}
use crate::embeddings::siglip::error::PostProcessorTemplate;
const TEMPLATE_UNDECLARED: &str = r#"{"type":"TemplateProcessing","single":[{"SpecialToken":{"id":"<s>","type_id":0}},{"Sequence":{"id":"A","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{}}"#;
const TEMPLATE_PAIR_IN_SINGLE: &str = r#"{"type":"TemplateProcessing","single":[{"Sequence":{"id":"A","type_id":0}},{"Sequence":{"id":"B","type_id":1}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{}}"#;
const TEMPLATE_NO_INPUT: &str = r#"{"type":"TemplateProcessing","single":[{"SpecialToken":{"id":"<s>","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{"<s>":{"id":"<s>","ids":[0],"tokens":["<pad>"]}}}"#;
const SEQUENCE_FEEDING_THREE_ENCODINGS: &str = r#"{"type":"Sequence","processors":[{"type":"TemplateProcessing","single":[{"SpecialToken":{"id":"<s>","type_id":0}},{"Sequence":{"id":"A","type_id":0}},{"SpecialToken":{"id":"<s>","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{"<s>":{"id":"<s>","ids":[0],"tokens":["<pad>"]}}},{"type":"TemplateProcessing","single":[{"Sequence":{"id":"A","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{}}]}"#;
const TEMPLATE_PLACING_THE_TEXT_TWICE: &str = r#"{"type":"TemplateProcessing","single":[{"Sequence":{"id":"A","type_id":0}},{"Sequence":{"id":"A","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{}}"#;
const SEQUENCE_REACHING_A_PAIR_TEMPLATE: &str = r#"{"type":"Sequence","processors":[{"type":"TemplateProcessing","single":[{"Sequence":{"id":"A","type_id":0}},{"SpecialToken":{"id":"<s>","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{"<s>":{"id":"<s>","ids":[0],"tokens":["<pad>"]}}},{"type":"TemplateProcessing","single":[{"Sequence":{"id":"A","type_id":0}}],"pair":[],"special_tokens":{}}]}"#;
fn tiny_tokenizer_with_post_processor(post: &str) -> Vec<u8> {
TINY_TOKENIZER
.replace(
r#""post_processor": null"#,
&format!(r#""post_processor": {post}"#),
)
.into_bytes()
}
#[test]
fn configure_tokenizer_refuses_an_inconsistent_single_template() {
for (post, expected) in [
(
TEMPLATE_UNDECLARED,
PostProcessorTemplate::UndeclaredSpecialToken("<s>".to_string()),
),
(
TEMPLATE_PAIR_IN_SINGLE,
PostProcessorTemplate::PairSequenceInSingleTemplate,
),
(
TEMPLATE_NO_INPUT,
PostProcessorTemplate::NoInputSequenceInSingleTemplate,
),
] {
let bytes = tiny_tokenizer_with_post_processor(post);
match configured_tokenizer_from_bytes(&bytes, 64) {
Err(Error::PostProcessorTemplate(why)) => assert_eq!(why, expected),
other => panic!("expected PostProcessorTemplate({expected:?}), got {other:?}"),
}
}
}
#[test]
fn configure_tokenizer_refuses_a_defective_template_inside_a_sequence() {
let post = format!(r#"{{"type":"Sequence","processors":[{TEMPLATE_UNDECLARED}]}}"#);
match configured_tokenizer_from_bytes(&tiny_tokenizer_with_post_processor(&post), 64) {
Err(Error::PostProcessorTemplate(PostProcessorTemplate::UndeclaredSpecialToken(id))) => {
assert_eq!(id, "<s>");
}
other => panic!("expected UndeclaredSpecialToken, got {other:?}"),
}
}
#[test]
fn configure_tokenizer_accepts_a_well_formed_single_template() {
let post = r#"{"type":"TemplateProcessing","single":[{"SpecialToken":{"id":"<s>","type_id":0}},{"Sequence":{"id":"A","type_id":0}}],"pair":[{"Sequence":{"id":"A","type_id":0}}],"special_tokens":{"<s>":{"id":"<s>","ids":[0],"tokens":["<pad>"]}}}"#;
let tok = configured_tokenizer_from_bytes(&tiny_tokenizer_with_post_processor(post), 64)
.expect("a sound template configures");
assert_eq!(
tok.encode("a b", true).expect("encode").get_ids(),
&[0, 1, 2],
"the special token, then the text"
);
}
#[test]
fn the_overhead_reading_is_blind_to_undeclared_special_tokens() {
const T: usize = 8;
let single: Vec<String> = (0..T)
.map(|i| format!(r#"{{"SpecialToken":{{"id":"<undeclared{i}>","type_id":0}}}}"#))
.collect();
let post = format!(
r#"{{"type":"TemplateProcessing","single":[{}],"pair":[{{"Sequence":{{"id":"A","type_id":0}}}}],"special_tokens":{{}}}}"#,
single.join(",")
);
let bytes = tiny_tokenizer_with_post_processor(&post);
let raw = Tokenizer::from_bytes(&bytes).expect("parse");
assert_eq!(
raw
.get_post_processor()
.expect("has one")
.added_tokens(false),
0,
"the overhead reading is blind to undeclared ids"
);
match configured_tokenizer_from_bytes(&bytes, T) {
Err(Error::PostProcessorTemplate(PostProcessorTemplate::UndeclaredSpecialToken(id))) => {
assert_eq!(id, "<undeclared0>");
}
other => panic!("expected UndeclaredSpecialToken, got {other:?}"),
}
}
const FIRST_ID_PAST_I32: u32 = 2_147_483_648;
#[test]
fn resolve_pad_id_refuses_a_pad_token_past_int32() {
let bytes = TINY_TOKENIZER
.replace(r#""<pad>": 0"#, &format!(r#""<pad>": {FIRST_ID_PAST_I32}"#))
.into_bytes();
let tok = Tokenizer::from_bytes(&bytes).expect("parse");
match resolve_pad_id(&tok) {
Err(Error::TokenIdRange(id)) => assert_eq!(id, FIRST_ID_PAST_I32),
other => panic!("expected TokenIdRange, got {other:?}"),
}
}
#[test]
fn resolve_pad_id_reads_the_vocabulary_then_falls_back() {
let tok = Tokenizer::from_bytes(TINY_TOKENIZER.as_bytes()).expect("parse");
assert_eq!(resolve_pad_id(&tok).expect("in range"), 0);
let bytes = TINY_TOKENIZER
.replace(r#""<pad>": 0, "#, "")
.replace(r#""unk_token": "<pad>""#, r#""unk_token": "a""#)
.into_bytes();
let tok = Tokenizer::from_bytes(&bytes).expect("parse");
assert_eq!(resolve_pad_id(&tok).expect("no <pad> is not an error"), 0);
}
#[test]
fn load_hashes_the_sidecar_before_parsing_it() {
let dir = tempfile::tempdir().expect("tempdir");
let model_path = dir.path().join("siglip2_text_64.mlmodelc");
let tokenizer_path = dir.path().join("tokenizer.json");
std::fs::write(&tokenizer_path, b"this is not json at all").expect("write sidecar");
match TextEmbedder::load(&model_path, TextEmbedderOptions::new()) {
Err(Error::ArtifactTokenizerIdentity(e)) => {
assert_eq!(e.path(), tokenizer_path.as_path());
assert_eq!(e.expected(), contract::TOKENIZER_SHA256_HEX);
}
other => panic!("expected ArtifactTokenizerIdentity before any parse, got {other:?}"),
}
}
#[test]
fn a_structural_defect_outranks_an_overhead_that_would_also_refuse() {
const T: usize = 8;
let bytes = tokenizer_bytes_with_overhead_and_a_pair_sequence(T);
let raw = Tokenizer::from_bytes(&bytes).expect("parse");
assert_eq!(
raw
.get_post_processor()
.expect("has one")
.added_tokens(false),
T,
"the overhead guard would refuse this tokenizer on its own"
);
match configured_tokenizer_from_bytes(&bytes, T) {
Err(Error::PostProcessorTemplate(PostProcessorTemplate::PairSequenceInSingleTemplate)) => {}
other => panic!("expected the structural refusal to be the one reported, got {other:?}"),
}
}
fn tokenizer_bytes_with_overhead_and_a_pair_sequence(added: usize) -> Vec<u8> {
let mut json: serde_json::Value =
serde_json::from_slice(&tokenizer_bytes_with_special_overhead(added))
.expect("the fixture serializes as JSON");
json["post_processor"]["single"]
.as_array_mut()
.expect("the single template serializes as an array")
.push(serde_json::json!({"Sequence": {"id": "B", "type_id": 1}}));
serde_json::to_vec(&json).expect("serialize the spliced tokenizer")
}
#[test]
fn configure_tokenizer_refuses_a_chain_that_feeds_a_template_an_unsupported_count() {
let bytes = tiny_tokenizer_with_post_processor(SEQUENCE_FEEDING_THREE_ENCODINGS);
match configured_tokenizer_from_bytes(&bytes, 64) {
Err(Error::PostProcessorTemplate(PostProcessorTemplate::UnsupportedEncodingCount(n))) => {
assert_eq!(n, 3, "the count the second template would have received");
}
other => panic!("expected UnsupportedEncodingCount(3), got {other:?}"),
}
}
#[test]
fn configure_tokenizer_refuses_a_template_that_places_the_text_twice() {
const WINDOW: usize = 4;
let bytes = tiny_tokenizer_with_post_processor(TEMPLATE_PLACING_THE_TEXT_TWICE);
let mut raw = Tokenizer::from_bytes(&bytes).expect("parse");
assert_eq!(
raw
.get_post_processor()
.expect("has one")
.added_tokens(false),
0,
"a repeated `$A` reads as no overhead at all"
);
raw
.with_truncation(Some(TruncationParams {
max_length: WINDOW,
strategy: TruncationStrategy::LongestFirst,
stride: 0,
direction: TruncationDirection::Right,
}))
.expect("zero overhead does not overflow");
raw.with_padding(None);
let ids = raw
.encode("a b a", true)
.expect("encode")
.get_ids()
.to_vec();
assert_eq!(ids, vec![1, 2, 1, 1, 2, 1], "three tokens in, six out");
match build_window(&ids, 0, PadSide::Right, WINDOW) {
Err(Error::TokenCount(count)) => {
assert_eq!(count.got(), 6);
assert_eq!(count.max(), WINDOW);
}
other => panic!("expected the TokenCount backstop to fire, got {other:?}"),
}
match configured_tokenizer_from_bytes(&bytes, WINDOW) {
Err(Error::PostProcessorTemplate(PostProcessorTemplate::RepeatedInputSequence(n))) => {
assert_eq!(n, 2, "the number of `$A` placements");
}
other => panic!("expected RepeatedInputSequence(2), got {other:?}"),
}
}
#[test]
fn configure_tokenizer_refuses_a_chain_that_reaches_a_pair_template() {
let bytes = tiny_tokenizer_with_post_processor(SEQUENCE_REACHING_A_PAIR_TEMPLATE);
let raw = Tokenizer::from_bytes(&bytes).expect("parse");
for text in ["a b", "a b a b a b"] {
assert!(
raw.encode(text, true).expect("encode").get_ids().is_empty(),
"the pair template places no sequence, so the text is gone"
);
}
match configured_tokenizer_from_bytes(&bytes, 64) {
Err(Error::PostProcessorTemplate(PostProcessorTemplate::UnsupportedEncodingCount(n))) => {
assert_eq!(n, 2, "the count the second template would have received");
}
other => panic!("expected UnsupportedEncodingCount(2), got {other:?}"),
}
}