Skip to main content

lattice_embed/
model.rs

1//! Embedding model selection, runtime dimensions, and load provenance.
2//!
3//! `EmbeddingModel` defines model-specific limits, prompting, and inference wiring.
4//! `ModelConfig` validates optional Matryoshka output truncation for supported models.
5//!
6//! See docs/model.md for the model and cache design.
7
8use serde::{Deserialize, Serialize};
9use std::time::SystemTime;
10
11/// **Stable**: external consumers may depend on this; breaking changes require a SemVer bump.
12///
13/// Records the model source and metadata for a load event.
14///
15/// See [`docs/model.md`](../docs/model.md#modelprovenance-source-behavior) for hash semantics and verification boundaries.
16#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct ModelProvenance {
18    /// **Stable**: model variant that was loaded.
19    pub model: EmbeddingModel,
20    /// **Stable**: source identifier (HuggingFace ID, URL, or file path).
21    pub model_id: String,
22    /// **Stable**: metadata-derived BLAKE3 identifier for this load event, not a weight checksum.
23    pub hash: String,
24    /// **Stable**: when the model was loaded.
25    pub loaded_at: SystemTime,
26    /// **Stable**: formatted timestamp string for convenience.
27    pub loaded_at_iso: String,
28}
29
30impl ModelProvenance {
31    /// **Stable**: create new provenance information for a loaded model.
32    pub fn new(model: EmbeddingModel, model_id: String) -> Self {
33        let loaded_at = SystemTime::now();
34        let loaded_at_iso = {
35            let dt: chrono::DateTime<chrono::Utc> = loaded_at.into();
36            dt.to_rfc3339()
37        };
38
39        let hash_input = format!("{model_id}:{loaded_at_iso}:{model:?}");
40        let hash = blake3::hash(hash_input.as_bytes()).to_hex().to_string();
41
42        Self {
43            model,
44            model_id,
45            hash,
46            loaded_at,
47            loaded_at_iso,
48        }
49    }
50
51    /// **Stable**: get the model dimensions.
52    pub fn dimensions(&self) -> usize {
53        self.model.dimensions()
54    }
55
56    /// **Stable**: check if this provenance matches expected model.
57    pub fn matches_model(&self, expected: EmbeddingModel) -> bool {
58        self.model == expected
59    }
60}
61
62/// **Stable**: external consumers may depend on this; breaking changes require a SemVer bump.
63///
64/// Registry of supported local and remote embedding models.
65///
66/// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for model capabilities and identity rules.
67#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)]
68#[serde(rename_all = "snake_case")]
69#[non_exhaustive]
70pub enum EmbeddingModel {
71    /// BGE small English v1.5 (384 dimensions) - fast and efficient.
72    #[default]
73    #[serde(alias = "BgeSmallEnV15")]
74    BgeSmallEnV15,
75
76    /// BGE base English v1.5 (768 dimensions) - balanced quality/speed.
77    #[serde(alias = "BgeBaseEnV15")]
78    BgeBaseEnV15,
79
80    /// BGE large English v1.5 (1024 dimensions) - highest quality local.
81    #[serde(alias = "BgeLargeEnV15")]
82    BgeLargeEnV15,
83
84    /// Multilingual E5 small (384 dimensions) - multilingual, same arch as BGE.
85    #[serde(alias = "MultilingualE5Small")]
86    MultilingualE5Small,
87
88    /// Multilingual E5 base (768 dimensions) - best multilingual quality/speed.
89    #[serde(alias = "MultilingualE5Base")]
90    MultilingualE5Base,
91
92    /// Qwen3-Embedding-0.6B (1024 dimensions) - multilingual, decoder-only, GPU-accelerated.
93    #[serde(alias = "Qwen3Embedding0_6B")]
94    Qwen3Embedding0_6B,
95
96    /// Qwen3-Embedding-4B (2560 dimensions, MRL-capable) - multilingual, decoder-only, GPU-accelerated.
97    #[serde(alias = "Qwen3Embedding4B")]
98    Qwen3Embedding4B,
99
100    /// all-MiniLM-L6-v2 (384 dimensions) - BERT-class, WordPiece tokenizer, sentence-transformers.
101    #[serde(alias = "AllMiniLmL6V2")]
102    AllMiniLmL6V2,
103
104    /// paraphrase-multilingual-MiniLM-L12-v2 (384 dimensions) - multilingual, XLM-R base, sentence-transformers.
105    #[serde(alias = "ParaphraseMultilingualMiniLmL12V2")]
106    ParaphraseMultilingualMiniLmL12V2,
107
108    /// OpenAI text-embedding-3-small (1536 dimensions) - remote API.
109    #[serde(alias = "TextEmbedding3Small")]
110    TextEmbedding3Small,
111}
112
113impl EmbeddingModel {
114    /// **Stable**: get the native (full-resolution) output dimension of this model's embeddings.
115    ///
116    /// Returns the model's intrinsic dimension regardless of any MRL truncation.
117    /// For MRL-capable models with a configured truncation, use `ModelConfig::dimensions()`.
118    #[inline]
119    pub const fn native_dimensions(&self) -> usize {
120        match self {
121            EmbeddingModel::BgeSmallEnV15
122            | EmbeddingModel::MultilingualE5Small
123            | EmbeddingModel::AllMiniLmL6V2
124            | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => 384,
125            EmbeddingModel::BgeBaseEnV15 | EmbeddingModel::MultilingualE5Base => 768,
126            EmbeddingModel::BgeLargeEnV15 | EmbeddingModel::Qwen3Embedding0_6B => 1024,
127            EmbeddingModel::Qwen3Embedding4B => 2560,
128            EmbeddingModel::TextEmbedding3Small => 1536,
129        }
130    }
131
132    /// **Stable**: get this model's native output dimension.
133    ///
134    /// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for active-dimension selection.
135    #[inline]
136    pub const fn dimensions(&self) -> usize {
137        self.native_dimensions()
138    }
139
140    /// **Stable**: check if this model can run locally (via lattice-inference).
141    #[inline]
142    pub const fn is_local(&self) -> bool {
143        matches!(
144            self,
145            EmbeddingModel::BgeSmallEnV15
146                | EmbeddingModel::BgeBaseEnV15
147                | EmbeddingModel::BgeLargeEnV15
148                | EmbeddingModel::MultilingualE5Small
149                | EmbeddingModel::MultilingualE5Base
150                | EmbeddingModel::AllMiniLmL6V2
151                | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2
152                | EmbeddingModel::Qwen3Embedding0_6B
153                | EmbeddingModel::Qwen3Embedding4B
154        )
155    }
156
157    /// **Stable**: check if this model requires a remote API.
158    #[inline]
159    pub const fn is_remote(&self) -> bool {
160        matches!(self, EmbeddingModel::TextEmbedding3Small)
161    }
162
163    /// **Stable**: conservative maximum input tokens for chunking and truncation.
164    ///
165    /// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for per-model limits.
166    #[inline]
167    pub const fn max_input_tokens(&self) -> usize {
168        match self {
169            EmbeddingModel::BgeSmallEnV15 => 512,
170            EmbeddingModel::BgeBaseEnV15 => 512,
171            EmbeddingModel::BgeLargeEnV15 => 512,
172            EmbeddingModel::MultilingualE5Small => 512,
173            EmbeddingModel::MultilingualE5Base => 512,
174            EmbeddingModel::AllMiniLmL6V2 => 256,
175            EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => 128,
176            // Conservative cap; see docs/model.md.
177            EmbeddingModel::Qwen3Embedding0_6B => 8192,
178            EmbeddingModel::Qwen3Embedding4B => 8192,
179            EmbeddingModel::TextEmbedding3Small => 8191,
180        }
181    }
182
183    /// **Stable**: query instruction prefix for asymmetric retrieval, when required.
184    ///
185    /// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for prompt policy and vector-space implications.
186    #[inline]
187    pub const fn query_instruction(&self) -> Option<&'static str> {
188        match self {
189            EmbeddingModel::MultilingualE5Small | EmbeddingModel::MultilingualE5Base => {
190                Some("query: ")
191            }
192            EmbeddingModel::Qwen3Embedding0_6B | EmbeddingModel::Qwen3Embedding4B => Some(
193                "Instruct: Given a web search query, retrieve relevant passages that answer the query\nQuery: ",
194            ),
195            EmbeddingModel::BgeSmallEnV15
196            | EmbeddingModel::BgeBaseEnV15
197            | EmbeddingModel::BgeLargeEnV15 => {
198                Some("Represent this sentence for searching relevant passages: ")
199            }
200            _ => None,
201        }
202    }
203
204    /// **Stable**: document instruction prefix for asymmetric retrieval, when required.
205    ///
206    /// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for prompt policy and vector-space implications.
207    #[inline]
208    pub const fn document_instruction(&self) -> Option<&'static str> {
209        match self {
210            EmbeddingModel::MultilingualE5Small | EmbeddingModel::MultilingualE5Base => {
211                Some("passage: ")
212            }
213            _ => None,
214        }
215    }
216
217    /// **Stable**: byte length of this model's longest retrieval instruction.
218    ///
219    /// Zero for symmetric models, which take no prefix at all.
220    ///
221    /// Callers apply an instruction to caller-supplied text, so text that has
222    /// already been prepared is longer than what the caller passed by at most
223    /// this many bytes. Length guards that run on prepared text size themselves
224    /// with this; guards that run on caller text use [`MAX_TEXT_BYTES`] exactly.
225    ///
226    /// [`MAX_TEXT_BYTES`]: crate::service::MAX_TEXT_BYTES
227    #[inline]
228    pub const fn max_instruction_bytes(&self) -> usize {
229        let q = match self.query_instruction() {
230            Some(s) => s.len(),
231            None => 0,
232        };
233        let d = match self.document_instruction() {
234            Some(s) => s.len(),
235            None => 0,
236        };
237        if q > d { q } else { d }
238    }
239
240    /// **Stable**: get the model identifier (HuggingFace ID or provider/model).
241    #[inline]
242    pub const fn model_id(&self) -> &'static str {
243        match self {
244            EmbeddingModel::BgeSmallEnV15 => "BAAI/bge-small-en-v1.5",
245            EmbeddingModel::BgeBaseEnV15 => "BAAI/bge-base-en-v1.5",
246            EmbeddingModel::BgeLargeEnV15 => "BAAI/bge-large-en-v1.5",
247            EmbeddingModel::MultilingualE5Small => "intfloat/multilingual-e5-small",
248            EmbeddingModel::MultilingualE5Base => "intfloat/multilingual-e5-base",
249            EmbeddingModel::AllMiniLmL6V2 => "sentence-transformers/all-MiniLM-L6-v2",
250            EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
251                "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
252            }
253            EmbeddingModel::Qwen3Embedding0_6B => "Qwen/Qwen3-Embedding-0.6B",
254            EmbeddingModel::Qwen3Embedding4B => "Qwen/Qwen3-Embedding-4B",
255            EmbeddingModel::TextEmbedding3Small => "text-embedding-3-small",
256        }
257    }
258
259    /// **Stable**: whether this model supports configurable output dimensions (MRL/Matryoshka).
260    #[inline]
261    pub const fn supports_output_dim(&self) -> bool {
262        matches!(
263            self,
264            EmbeddingModel::Qwen3Embedding0_6B | EmbeddingModel::Qwen3Embedding4B
265        )
266    }
267
268    /// **Stable**: BERT pooling strategy for this model, or `None` for non-BERT paths.
269    ///
270    /// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for pooling routing.
271    #[cfg(feature = "native")]
272    #[inline]
273    pub const fn bert_pooling(&self) -> Option<lattice_inference::BertPooling> {
274        match self {
275            EmbeddingModel::BgeSmallEnV15
276            | EmbeddingModel::BgeBaseEnV15
277            | EmbeddingModel::BgeLargeEnV15 => Some(lattice_inference::BertPooling::CLS),
278            EmbeddingModel::MultilingualE5Small | EmbeddingModel::MultilingualE5Base => {
279                Some(lattice_inference::BertPooling::Mean)
280            }
281            EmbeddingModel::AllMiniLmL6V2 | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
282                Some(lattice_inference::BertPooling::Mean)
283            }
284            EmbeddingModel::Qwen3Embedding0_6B
285            | EmbeddingModel::Qwen3Embedding4B
286            | EmbeddingModel::TextEmbedding3Small => None,
287        }
288    }
289
290    /// **Stable**: embedding key revision string for this model family.
291    #[inline]
292    pub const fn key_version(&self) -> &'static str {
293        match self {
294            EmbeddingModel::TextEmbedding3Small
295            | EmbeddingModel::Qwen3Embedding0_6B
296            | EmbeddingModel::Qwen3Embedding4B => "v3",
297            EmbeddingModel::AllMiniLmL6V2 | EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
298                "v2"
299            }
300            _ => "v1.5",
301        }
302    }
303}
304
305impl std::fmt::Display for EmbeddingModel {
306    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
307        match self {
308            EmbeddingModel::BgeSmallEnV15 => write!(f, "bge-small-en-v1.5"),
309            EmbeddingModel::BgeBaseEnV15 => write!(f, "bge-base-en-v1.5"),
310            EmbeddingModel::BgeLargeEnV15 => write!(f, "bge-large-en-v1.5"),
311            EmbeddingModel::MultilingualE5Small => write!(f, "multilingual-e5-small"),
312            EmbeddingModel::MultilingualE5Base => write!(f, "multilingual-e5-base"),
313            EmbeddingModel::Qwen3Embedding0_6B => write!(f, "qwen3-embedding-0.6b"),
314            EmbeddingModel::Qwen3Embedding4B => write!(f, "qwen3-embedding-4b"),
315            EmbeddingModel::AllMiniLmL6V2 => write!(f, "all-minilm-l6-v2"),
316            EmbeddingModel::ParaphraseMultilingualMiniLmL12V2 => {
317                write!(f, "paraphrase-multilingual-minilm-l12-v2")
318            }
319            EmbeddingModel::TextEmbedding3Small => write!(f, "text-embedding-3-small"),
320        }
321    }
322}
323
324impl std::str::FromStr for EmbeddingModel {
325    type Err = String;
326
327    /// **Stable**: parse a normalized canonical name, alias, or supported provider identifier.
328    ///
329    /// See [`docs/model.md`](../docs/model.md#embeddingmodel-source-behavior) for accepted forms and persistence guidance.
330    fn from_str(s: &str) -> Result<Self, Self::Err> {
331        let lower = s.to_lowercase();
332        let normalized = lower.trim().replace("_", "-").replace("baai/", "");
333
334        match normalized.as_str() {
335            "bge-small-en-v1.5" | "bge-small-en" | "bge-small" | "small" => {
336                Ok(EmbeddingModel::BgeSmallEnV15)
337            }
338            "bge-base-en-v1.5" | "bge-base-en" | "bge-base" | "base" => {
339                Ok(EmbeddingModel::BgeBaseEnV15)
340            }
341            "bge-large-en-v1.5" | "bge-large-en" | "bge-large" | "large" => {
342                Ok(EmbeddingModel::BgeLargeEnV15)
343            }
344            "multilingual-e5-small" | "e5-small" | "intfloat/multilingual-e5-small" => {
345                Ok(EmbeddingModel::MultilingualE5Small)
346            }
347            "multilingual-e5-base" | "e5-base" | "intfloat/multilingual-e5-base" => {
348                Ok(EmbeddingModel::MultilingualE5Base)
349            }
350            "qwen3-embedding-0.6b" | "qwen3-embedding" | "qwen3" | "qwen/qwen3-embedding-0.6b" => {
351                Ok(EmbeddingModel::Qwen3Embedding0_6B)
352            }
353            "qwen3-embedding-4b" | "qwen3-4b" | "qwen/qwen3-embedding-4b" => {
354                Ok(EmbeddingModel::Qwen3Embedding4B)
355            }
356            "all-minilm-l6-v2"
357            | "minilm"
358            | "all-minilm"
359            | "sentence-transformers/all-minilm-l6-v2" => Ok(EmbeddingModel::AllMiniLmL6V2),
360            "paraphrase-multilingual-minilm-l12-v2"
361            | "paraphrase-multilingual"
362            | "multilingual-minilm"
363            | "sentence-transformers/paraphrase-multilingual-minilm-l12-v2" => {
364                Ok(EmbeddingModel::ParaphraseMultilingualMiniLmL12V2)
365            }
366            "text-embedding-3-small" | "openai-small" | "openai" => {
367                Ok(EmbeddingModel::TextEmbedding3Small)
368            }
369            _ => Err(format!(
370                "unknown embedding model: '{s}'. Valid: bge-small-en-v1.5, bge-base-en-v1.5, bge-large-en-v1.5, multilingual-e5-small, multilingual-e5-base, text-embedding-3-small"
371            )),
372        }
373    }
374}
375
376// ============================================================================
377// ModelConfig — runtime MRL dimension configuration
378// ============================================================================
379
380/// Minimum allowed MRL output dimension.
381pub const MIN_MRL_OUTPUT_DIM: usize = 32;
382
383/// Runtime model configuration with an optional MRL truncation dimension.
384///
385/// See [`docs/model.md`](../docs/model.md#modelconfig-source-behavior) for validation and namespace requirements.
386#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
387pub struct ModelConfig {
388    /// The underlying embedding model.
389    pub model: EmbeddingModel,
390    /// MRL truncation dimension. `None` uses the model's native dimension.
391    #[serde(default)]
392    pub output_dim: Option<usize>,
393}
394
395impl Default for ModelConfig {
396    fn default() -> Self {
397        Self::new(EmbeddingModel::default())
398    }
399}
400
401impl ModelConfig {
402    /// Create a config with no MRL truncation (native model dimension).
403    pub const fn new(model: EmbeddingModel) -> Self {
404        Self {
405            model,
406            output_dim: None,
407        }
408    }
409
410    /// Create and validate a config with an optional MRL truncation dimension.
411    pub fn try_new(
412        model: EmbeddingModel,
413        output_dim: Option<usize>,
414    ) -> std::result::Result<Self, crate::error::EmbedError> {
415        let config = Self { model, output_dim };
416        config.validate()?;
417        Ok(config)
418    }
419
420    /// Validate that the output dimension is consistent with the model.
421    pub fn validate(&self) -> std::result::Result<(), crate::error::EmbedError> {
422        let Some(dim) = self.output_dim else {
423            return Ok(());
424        };
425        if !self.model.supports_output_dim() {
426            return Err(crate::error::EmbedError::InvalidInput(format!(
427                "{} does not support configurable embedding dimensions",
428                self.model
429            )));
430        }
431        if dim < MIN_MRL_OUTPUT_DIM {
432            return Err(crate::error::EmbedError::InvalidInput(format!(
433                "embedding output dimension {dim} is below minimum {MIN_MRL_OUTPUT_DIM}"
434            )));
435        }
436        let native = self.model.native_dimensions();
437        if dim > native {
438            return Err(crate::error::EmbedError::InvalidInput(format!(
439                "embedding output dimension {dim} exceeds native dimension {native} for {}",
440                self.model
441            )));
442        }
443        Ok(())
444    }
445
446    /// Active output dimension: configured truncation if set, otherwise the model's native dimension.
447    pub fn dimensions(&self) -> usize {
448        self.output_dim
449            .unwrap_or_else(|| self.model.native_dimensions())
450    }
451}
452
453#[cfg(test)]
454mod tests {
455    use super::*;
456
457    #[test]
458    fn test_default_model() {
459        let model = EmbeddingModel::default();
460        assert_eq!(model, EmbeddingModel::BgeSmallEnV15);
461    }
462
463    #[test]
464    fn test_model_provenance_new() {
465        let provenance = ModelProvenance::new(
466            EmbeddingModel::BgeSmallEnV15,
467            "BAAI/bge-small-en-v1.5".into(),
468        );
469
470        assert_eq!(provenance.model, EmbeddingModel::BgeSmallEnV15);
471        assert_eq!(provenance.model_id, "BAAI/bge-small-en-v1.5");
472        assert!(!provenance.hash.is_empty());
473        assert_eq!(provenance.hash.len(), 64); // blake3 hex is 64 chars
474        assert!(!provenance.loaded_at_iso.is_empty());
475    }
476
477    #[test]
478    fn test_model_provenance_unique_hash() {
479        let p1 = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "model1".into());
480        std::thread::sleep(std::time::Duration::from_millis(10)); // Ensure different timestamp
481        let p2 = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "model1".into());
482
483        // Different timestamps should produce different hashes
484        assert_ne!(p1.hash, p2.hash);
485    }
486
487    #[test]
488    fn test_model_provenance_dimensions() {
489        let p1 = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "small".into());
490        assert_eq!(p1.dimensions(), 384);
491
492        let p2 = ModelProvenance::new(EmbeddingModel::BgeBaseEnV15, "base".into());
493        assert_eq!(p2.dimensions(), 768);
494
495        let p3 = ModelProvenance::new(EmbeddingModel::BgeLargeEnV15, "large".into());
496        assert_eq!(p3.dimensions(), 1024);
497    }
498
499    #[test]
500    fn test_model_provenance_matches_model() {
501        let provenance = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "test".into());
502
503        assert!(provenance.matches_model(EmbeddingModel::BgeSmallEnV15));
504        assert!(!provenance.matches_model(EmbeddingModel::BgeBaseEnV15));
505        assert!(!provenance.matches_model(EmbeddingModel::BgeLargeEnV15));
506    }
507
508    #[test]
509    fn test_model_provenance_serialization() {
510        let provenance = ModelProvenance::new(EmbeddingModel::BgeSmallEnV15, "test-model".into());
511
512        let json = serde_json::to_string(&provenance).unwrap();
513        // FP-037: EmbeddingModel has #[serde(rename_all = "snake_case")] so
514        // BgeSmallEnV15 serializes as "bge_small_en_v15", not "BgeSmallEnV15".
515        assert!(json.contains("bge_small_en_v15"), "json={json}");
516        assert!(json.contains("test-model"));
517        assert!(json.contains(&provenance.hash));
518
519        let parsed: ModelProvenance = serde_json::from_str(&json).unwrap();
520        assert_eq!(parsed.model, provenance.model);
521        assert_eq!(parsed.model_id, provenance.model_id);
522        assert_eq!(parsed.hash, provenance.hash);
523    }
524
525    #[test]
526    fn test_dimensions() {
527        assert_eq!(EmbeddingModel::BgeSmallEnV15.dimensions(), 384);
528        assert_eq!(EmbeddingModel::BgeBaseEnV15.dimensions(), 768);
529        assert_eq!(EmbeddingModel::BgeLargeEnV15.dimensions(), 1024);
530        assert_eq!(EmbeddingModel::Qwen3Embedding4B.dimensions(), 2560);
531    }
532
533    #[test]
534    fn test_model_config_native_dims() {
535        assert_eq!(
536            ModelConfig::new(EmbeddingModel::Qwen3Embedding4B).dimensions(),
537            2560
538        );
539        assert_eq!(
540            ModelConfig::new(EmbeddingModel::Qwen3Embedding0_6B).dimensions(),
541            1024
542        );
543        assert_eq!(
544            ModelConfig::new(EmbeddingModel::BgeSmallEnV15).dimensions(),
545            384
546        );
547    }
548
549    #[test]
550    fn test_model_config_configured_dim() {
551        let cfg = ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(1024)).unwrap();
552        assert_eq!(cfg.dimensions(), 1024);
553
554        let cfg = ModelConfig::try_new(EmbeddingModel::Qwen3Embedding0_6B, Some(512)).unwrap();
555        assert_eq!(cfg.dimensions(), 512);
556    }
557
558    #[test]
559    fn test_model_config_validation_below_min() {
560        assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(31)).is_err());
561        assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(0)).is_err());
562    }
563
564    #[test]
565    fn test_model_config_validation_above_native() {
566        assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, Some(2561)).is_err());
567        assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding0_6B, Some(1025)).is_err());
568    }
569
570    #[test]
571    fn test_model_config_validation_non_mrl_model() {
572        assert!(ModelConfig::try_new(EmbeddingModel::BgeSmallEnV15, Some(128)).is_err());
573        assert!(ModelConfig::try_new(EmbeddingModel::BgeBaseEnV15, Some(512)).is_err());
574    }
575
576    #[test]
577    fn test_model_config_none_output_dim_ok_for_any_model() {
578        assert!(ModelConfig::try_new(EmbeddingModel::BgeSmallEnV15, None).is_ok());
579        assert!(ModelConfig::try_new(EmbeddingModel::Qwen3Embedding4B, None).is_ok());
580    }
581
582    #[test]
583    fn test_is_local() {
584        assert!(EmbeddingModel::BgeSmallEnV15.is_local());
585        assert!(EmbeddingModel::BgeBaseEnV15.is_local());
586        assert!(EmbeddingModel::BgeLargeEnV15.is_local());
587    }
588
589    #[test]
590    fn test_display() {
591        assert_eq!(
592            EmbeddingModel::BgeSmallEnV15.to_string(),
593            "bge-small-en-v1.5"
594        );
595        assert_eq!(EmbeddingModel::BgeBaseEnV15.to_string(), "bge-base-en-v1.5");
596        assert_eq!(
597            EmbeddingModel::BgeLargeEnV15.to_string(),
598            "bge-large-en-v1.5"
599        );
600    }
601
602    #[test]
603    fn test_serialization_roundtrip() {
604        let model = EmbeddingModel::BgeSmallEnV15;
605        let json = serde_json::to_string(&model).unwrap();
606        let parsed: EmbeddingModel = serde_json::from_str(&json).unwrap();
607        assert_eq!(model, parsed);
608    }
609
610    #[test]
611    fn test_max_input_tokens() {
612        assert_eq!(EmbeddingModel::BgeSmallEnV15.max_input_tokens(), 512);
613        assert_eq!(EmbeddingModel::BgeBaseEnV15.max_input_tokens(), 512);
614        assert_eq!(EmbeddingModel::BgeLargeEnV15.max_input_tokens(), 512);
615    }
616
617    #[test]
618    fn test_from_str_display_names() {
619        assert_eq!(
620            "bge-small-en-v1.5".parse::<EmbeddingModel>().unwrap(),
621            EmbeddingModel::BgeSmallEnV15
622        );
623        assert_eq!(
624            "bge-base-en-v1.5".parse::<EmbeddingModel>().unwrap(),
625            EmbeddingModel::BgeBaseEnV15
626        );
627        assert_eq!(
628            "bge-large-en-v1.5".parse::<EmbeddingModel>().unwrap(),
629            EmbeddingModel::BgeLargeEnV15
630        );
631    }
632
633    #[test]
634    fn test_from_str_short_names() {
635        assert_eq!(
636            "small".parse::<EmbeddingModel>().unwrap(),
637            EmbeddingModel::BgeSmallEnV15
638        );
639        assert_eq!(
640            "bge-base".parse::<EmbeddingModel>().unwrap(),
641            EmbeddingModel::BgeBaseEnV15
642        );
643        assert_eq!(
644            "LARGE".parse::<EmbeddingModel>().unwrap(), // case insensitive
645            EmbeddingModel::BgeLargeEnV15
646        );
647    }
648
649    #[test]
650    fn test_from_str_huggingface_ids() {
651        assert_eq!(
652            "BAAI/bge-small-en-v1.5".parse::<EmbeddingModel>().unwrap(),
653            EmbeddingModel::BgeSmallEnV15
654        );
655    }
656
657    #[test]
658    fn test_from_str_invalid() {
659        let result = "unknown-model".parse::<EmbeddingModel>();
660        assert!(result.is_err());
661        assert!(result.unwrap_err().contains("unknown embedding model"));
662    }
663
664    // -------------------------------------------------------------------------
665    // bert_pooling() routing tests (P1-E3) — require `native` feature
666    // -------------------------------------------------------------------------
667
668    /// BGE small/base/large must use CLS pooling per their HF model cards.
669    #[cfg(feature = "native")]
670    #[test]
671    fn test_bge_models_use_cls_pooling() {
672        use lattice_inference::BertPooling;
673
674        assert_eq!(
675            EmbeddingModel::BgeSmallEnV15.bert_pooling(),
676            Some(BertPooling::CLS),
677            "BgeSmallEnV15 must use CLS pooling"
678        );
679        assert_eq!(
680            EmbeddingModel::BgeBaseEnV15.bert_pooling(),
681            Some(BertPooling::CLS),
682            "BgeBaseEnV15 must use CLS pooling"
683        );
684        assert_eq!(
685            EmbeddingModel::BgeLargeEnV15.bert_pooling(),
686            Some(BertPooling::CLS),
687            "BgeLargeEnV15 must use CLS pooling"
688        );
689    }
690
691    /// E5 models must use mean pooling per their HF model cards.
692    #[cfg(feature = "native")]
693    #[test]
694    fn test_e5_models_use_mean_pooling() {
695        use lattice_inference::BertPooling;
696
697        assert_eq!(
698            EmbeddingModel::MultilingualE5Small.bert_pooling(),
699            Some(BertPooling::Mean),
700            "MultilingualE5Small must use mean pooling"
701        );
702        assert_eq!(
703            EmbeddingModel::MultilingualE5Base.bert_pooling(),
704            Some(BertPooling::Mean),
705            "MultilingualE5Base must use mean pooling"
706        );
707    }
708
709    /// MiniLM models must use mean pooling per sentence-transformers convention.
710    #[cfg(feature = "native")]
711    #[test]
712    fn test_minilm_models_use_mean_pooling() {
713        use lattice_inference::BertPooling;
714
715        assert_eq!(
716            EmbeddingModel::AllMiniLmL6V2.bert_pooling(),
717            Some(BertPooling::Mean),
718            "AllMiniLmL6V2 must use mean pooling"
719        );
720        assert_eq!(
721            EmbeddingModel::ParaphraseMultilingualMiniLmL12V2.bert_pooling(),
722            Some(BertPooling::Mean),
723            "ParaphraseMultilingualMiniLmL12V2 must use mean pooling"
724        );
725    }
726
727    /// Qwen and remote models return None — they have separate pooling paths.
728    #[cfg(feature = "native")]
729    #[test]
730    fn test_non_bert_models_return_none_pooling() {
731        assert_eq!(
732            EmbeddingModel::Qwen3Embedding0_6B.bert_pooling(),
733            None,
734            "Qwen model must return None for bert_pooling()"
735        );
736        assert_eq!(
737            EmbeddingModel::Qwen3Embedding4B.bert_pooling(),
738            None,
739            "Qwen model must return None for bert_pooling()"
740        );
741        assert_eq!(
742            EmbeddingModel::TextEmbedding3Small.bert_pooling(),
743            None,
744            "Remote model must return None for bert_pooling()"
745        );
746    }
747
748    /// BGE and E5 use DIFFERENT pooling strategies — this is the key correctness distinction.
749    #[cfg(feature = "native")]
750    #[test]
751    fn test_bge_and_e5_use_different_pooling() {
752        assert_ne!(
753            EmbeddingModel::BgeSmallEnV15.bert_pooling(),
754            EmbeddingModel::MultilingualE5Small.bert_pooling(),
755            "BGE and E5 must use different pooling strategies"
756        );
757    }
758}