Skip to main content

oxirs_vec/
huggingface.rs

1//! HuggingFace Transformers integration for embedding generation
2
3use crate::{EmbeddableContent, EmbeddingConfig, Vector};
4use anyhow::{anyhow, Result};
5use scirs2_core::random::Random;
6use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9/// HuggingFace model configuration
10#[derive(Debug, Clone, Serialize, Deserialize)]
11pub struct HuggingFaceConfig {
12    pub model_name: String,
13    pub cache_dir: Option<String>,
14    pub device: String,
15    pub batch_size: usize,
16    pub max_length: usize,
17    pub pooling_strategy: PoolingStrategy,
18    pub trust_remote_code: bool,
19    /// HuggingFace Inference API token (https://huggingface.co/settings/tokens).
20    /// When set and [`HuggingFaceConfig::use_inference_api`] is `true`,
21    /// [`HuggingFaceEmbedder`] calls the real HuggingFace Inference API over
22    /// HTTP instead of falling back to [`HuggingFaceEmbedder::deterministic_mock_embedding`].
23    #[serde(skip_serializing_if = "Option::is_none", default)]
24    pub api_token: Option<String>,
25    /// Whether to call the real HuggingFace Inference API (requires
26    /// `api_token` and network access). Defaults to `false`: without an
27    /// explicit opt-in, this type never masquerades as real HF inference —
28    /// it uses an honestly-labeled deterministic offline mock instead.
29    #[serde(default)]
30    pub use_inference_api: bool,
31    /// Base URL for the HuggingFace Inference API (overridable for testing
32    /// against a mock server).
33    #[serde(default = "default_inference_api_base_url")]
34    pub inference_api_base_url: String,
35}
36
37fn default_inference_api_base_url() -> String {
38    "https://api-inference.huggingface.co/models".to_string()
39}
40
41/// Pooling strategies for transformer outputs
42#[derive(Debug, Clone, Serialize, Deserialize)]
43pub enum PoolingStrategy {
44    /// Use `[CLS]` token embedding
45    Cls,
46    /// Mean pooling of all token embeddings
47    Mean,
48    /// Max pooling of all token embeddings
49    Max,
50    /// Weighted mean pooling based on attention weights
51    AttentionWeighted,
52}
53
54impl Default for HuggingFaceConfig {
55    fn default() -> Self {
56        Self {
57            model_name: "sentence-transformers/all-MiniLM-L6-v2".to_string(),
58            cache_dir: None,
59            device: "cpu".to_string(),
60            batch_size: 32,
61            max_length: 512,
62            pooling_strategy: PoolingStrategy::Mean,
63            trust_remote_code: false,
64            api_token: None,
65            use_inference_api: false,
66            inference_api_base_url: default_inference_api_base_url(),
67        }
68    }
69}
70
71/// HuggingFace transformer model for embedding generation
72#[derive(Debug)]
73pub struct HuggingFaceEmbedder {
74    config: HuggingFaceConfig,
75    model_cache: HashMap<String, ModelInfo>,
76}
77
78/// Model information and metadata
79#[derive(Debug, Clone)]
80struct ModelInfo {
81    dimensions: usize,
82    max_sequence_length: usize,
83    model_type: String,
84    loaded: bool,
85}
86
87impl HuggingFaceEmbedder {
88    /// Create a new HuggingFace embedder
89    pub fn new(config: HuggingFaceConfig) -> Result<Self> {
90        Ok(Self {
91            config,
92            model_cache: HashMap::new(),
93        })
94    }
95
96    /// Create embedder with default configuration
97    pub fn with_default_config() -> Result<Self> {
98        Self::new(HuggingFaceConfig::default())
99    }
100
101    /// Load a model and prepare it for inference
102    pub async fn load_model(&mut self, model_name: &str) -> Result<()> {
103        if self.model_cache.contains_key(model_name) {
104            return Ok(());
105        }
106
107        // Check if model exists in cache directory
108        let model_info = self.get_model_info(model_name).await?;
109        self.model_cache.insert(model_name.to_string(), model_info);
110
111        tracing::info!("Loaded HuggingFace model: {}", model_name);
112        Ok(())
113    }
114
115    /// Get model information from HuggingFace Hub
116    async fn get_model_info(&self, model_name: &str) -> Result<ModelInfo> {
117        // Simulate fetching model info from HuggingFace Hub
118        // In a real implementation, this would use the HuggingFace API
119        let dimensions = match model_name {
120            "sentence-transformers/all-MiniLM-L6-v2" => 384,
121            "sentence-transformers/all-mpnet-base-v2" => 768,
122            "microsoft/DialoGPT-medium" => 1024,
123            "bert-base-uncased" => 768,
124            "distilbert-base-uncased" => 768,
125            _ => 768, // Default dimension
126        };
127
128        Ok(ModelInfo {
129            dimensions,
130            max_sequence_length: self.config.max_length,
131            model_type: "transformer".to_string(),
132            loaded: true,
133        })
134    }
135
136    /// Generate embeddings for a batch of content
137    pub async fn embed_batch(&mut self, contents: &[EmbeddableContent]) -> Result<Vec<Vector>> {
138        if contents.is_empty() {
139            return Ok(vec![]);
140        }
141
142        // Load model if not already loaded
143        let model_name = self.config.model_name.clone();
144        self.load_model(&model_name).await?;
145
146        let model_info = self
147            .model_cache
148            .get(&self.config.model_name)
149            .ok_or_else(|| anyhow!("Model not loaded: {}", self.config.model_name))?;
150
151        let mut embeddings = Vec::with_capacity(contents.len());
152
153        // Process in batches
154        for chunk in contents.chunks(self.config.batch_size) {
155            let texts: Vec<String> = chunk
156                .iter()
157                .map(|content| self.content_to_text(content))
158                .collect();
159
160            let batch_embeddings = self.generate_embeddings(&texts, model_info).await?;
161            embeddings.extend(batch_embeddings);
162        }
163
164        Ok(embeddings)
165    }
166
167    /// Generate a single embedding
168    pub async fn embed(&mut self, content: &EmbeddableContent) -> Result<Vector> {
169        let embeddings = self.embed_batch(std::slice::from_ref(content)).await?;
170        embeddings
171            .into_iter()
172            .next()
173            .ok_or_else(|| anyhow!("Failed to generate embedding"))
174    }
175
176    /// Convert embeddable content to text
177    fn content_to_text(&self, content: &EmbeddableContent) -> String {
178        match content {
179            EmbeddableContent::Text(text) => text.clone(),
180            EmbeddableContent::RdfResource {
181                uri,
182                label,
183                description,
184                properties,
185            } => {
186                let mut text_parts = vec![uri.clone()];
187
188                if let Some(label) = label {
189                    text_parts.push(label.clone());
190                }
191
192                if let Some(desc) = description {
193                    text_parts.push(desc.clone());
194                }
195
196                for (prop, values) in properties {
197                    text_parts.push(format!("{}: {}", prop, values.join(", ")));
198                }
199
200                text_parts.join(" ")
201            }
202            EmbeddableContent::SparqlQuery(query) => query.clone(),
203            EmbeddableContent::GraphPattern(pattern) => pattern.clone(),
204        }
205    }
206
207    /// Generate embeddings using the configured backend.
208    ///
209    /// Calls the real HuggingFace Inference API over HTTP when
210    /// [`HuggingFaceConfig::use_inference_api`] is `true` (and an
211    /// `api_token` is configured); otherwise falls back to
212    /// [`Self::deterministic_mock_embedding`] with a logged warning, since
213    /// that fallback is NOT derived from any real transformer inference.
214    async fn generate_embeddings(
215        &self,
216        texts: &[String],
217        model_info: &ModelInfo,
218    ) -> Result<Vec<Vector>> {
219        if self.config.use_inference_api {
220            let token = self.config.api_token.as_ref().ok_or_else(|| {
221                anyhow!(
222                    "HuggingFaceConfig.use_inference_api is true but no api_token is \
223                     configured; set api_token (see https://huggingface.co/settings/tokens) \
224                     to call the real HuggingFace Inference API"
225                )
226            })?;
227            return self.call_inference_api(texts, token).await;
228        }
229
230        tracing::warn!(
231            "HuggingFaceEmbedder: use_inference_api is false, so '{}' embeddings are \
232             produced by deterministic_mock_embedding — a hash-seeded pseudo-random \
233             vector that is NOT related to text semantics and NOT real HuggingFace \
234             inference. Set HuggingFaceConfig::use_inference_api=true with an api_token \
235             for real embeddings.",
236            self.config.model_name
237        );
238
239        let mut embeddings = Vec::with_capacity(texts.len());
240        for text in texts {
241            let embedding = self.deterministic_mock_embedding(text, model_info.dimensions)?;
242            embeddings.push(embedding);
243        }
244
245        Ok(embeddings)
246    }
247
248    /// Call the real HuggingFace Inference API's feature-extraction endpoint.
249    async fn call_inference_api(&self, texts: &[String], token: &str) -> Result<Vec<Vector>> {
250        let url = format!(
251            "{}/{}",
252            self.config.inference_api_base_url, self.config.model_name
253        );
254
255        let client = reqwest::Client::new();
256        let response = client
257            .post(&url)
258            .bearer_auth(token)
259            .json(&serde_json::json!({
260                "inputs": texts,
261                "options": { "wait_for_model": true },
262            }))
263            .send()
264            .await
265            .map_err(|e| anyhow!("HuggingFace Inference API request to {} failed: {}", url, e))?;
266
267        if !response.status().is_success() {
268            let status = response.status();
269            let body = response.text().await.unwrap_or_default();
270            return Err(anyhow!(
271                "HuggingFace Inference API returned {} for {}: {}",
272                status,
273                url,
274                body
275            ));
276        }
277
278        let raw: serde_json::Value = response.json().await.map_err(|e| {
279            anyhow!(
280                "Failed to parse HuggingFace Inference API response as JSON: {}",
281                e
282            )
283        })?;
284
285        Self::parse_inference_response(&raw, texts.len())
286    }
287
288    /// Parse a HuggingFace feature-extraction response into per-text
289    /// vectors. Handles both already-pooled (`[[f32; d]; n]`) and
290    /// token-level (`[[[f32; d]; tokens]; n]`, mean-pooled here) response
291    /// shapes, since different models' pipeline tags return either.
292    fn parse_inference_response(
293        raw: &serde_json::Value,
294        expected_count: usize,
295    ) -> Result<Vec<Vector>> {
296        let items = raw.as_array().ok_or_else(|| {
297            anyhow!("Unexpected HuggingFace Inference API response shape: expected a JSON array")
298        })?;
299
300        let vectors = items
301            .iter()
302            .map(Self::parse_single_embedding)
303            .collect::<Result<Vec<Vector>>>()?;
304
305        if vectors.len() != expected_count {
306            return Err(anyhow!(
307                "HuggingFace Inference API returned {} embeddings for {} input texts",
308                vectors.len(),
309                expected_count
310            ));
311        }
312
313        Ok(vectors)
314    }
315
316    /// Parse one response item, which is either an already-pooled flat
317    /// array of numbers, or a nested array of per-token vectors (mean-pooled
318    /// here into a single sentence vector).
319    fn parse_single_embedding(value: &serde_json::Value) -> Result<Vector> {
320        let arr = value
321            .as_array()
322            .ok_or_else(|| anyhow!("Unexpected embedding item shape in HF response"))?;
323
324        if arr.iter().all(|v| v.is_number()) {
325            let values: Vec<f32> = arr
326                .iter()
327                .map(|n| n.as_f64().unwrap_or(0.0) as f32)
328                .collect();
329            return Ok(Vector::new(values));
330        }
331
332        // Token-level response: mean-pool across tokens.
333        let mut sum: Option<Vec<f32>> = None;
334        let mut count = 0usize;
335        for token in arr {
336            let token_values: Vec<f32> = token
337                .as_array()
338                .ok_or_else(|| anyhow!("Unexpected token embedding shape in HF response"))?
339                .iter()
340                .map(|n| n.as_f64().unwrap_or(0.0) as f32)
341                .collect();
342            sum = Some(match sum {
343                None => token_values,
344                Some(mut acc) => {
345                    for (a, b) in acc.iter_mut().zip(token_values.iter()) {
346                        *a += b;
347                    }
348                    acc
349                }
350            });
351            count += 1;
352        }
353
354        let mut pooled =
355            sum.ok_or_else(|| anyhow!("Empty token embedding array in HF response"))?;
356        if count > 0 {
357            for v in pooled.iter_mut() {
358                *v /= count as f32;
359            }
360        }
361        Ok(Vector::new(pooled))
362    }
363
364    /// Deterministic, hash-seeded pseudo-random embedding.
365    ///
366    /// **This is NOT real HuggingFace inference.** It is an offline mock
367    /// used only when [`HuggingFaceConfig::use_inference_api`] is `false`
368    /// (the default), so tests and offline development don't require
369    /// network access. The output is unrelated to the text's semantics —
370    /// callers that need real embeddings must set `use_inference_api = true`
371    /// with a valid `api_token`.
372    fn deterministic_mock_embedding(&self, text: &str, dimensions: usize) -> Result<Vector> {
373        // Simple hash-based embedding simulation
374        use std::collections::hash_map::DefaultHasher;
375        use std::hash::{Hash, Hasher};
376
377        let mut hasher = DefaultHasher::new();
378        text.hash(&mut hasher);
379        let seed = hasher.finish();
380
381        let mut rng = Random::seed(seed);
382
383        let mut embedding = vec![0.0f32; dimensions];
384        for value in embedding.iter_mut().take(dimensions) {
385            *value = rng.gen_range(-1.0..1.0); // Random values between -1 and 1
386        }
387
388        // Normalize if required
389        if matches!(self.config.pooling_strategy, PoolingStrategy::Mean) {
390            let norm = embedding.iter().map(|x| x * x).sum::<f32>().sqrt();
391            if norm > 0.0 {
392                for x in &mut embedding {
393                    *x /= norm;
394                }
395            }
396        }
397
398        Ok(Vector::new(embedding))
399    }
400
401    /// Get available models from cache
402    pub fn get_cached_models(&self) -> Vec<String> {
403        self.model_cache.keys().cloned().collect()
404    }
405
406    /// Clear model cache
407    pub fn clear_cache(&mut self) {
408        self.model_cache.clear();
409    }
410
411    /// Get model dimensions
412    pub fn get_model_dimensions(&self, model_name: &str) -> Option<usize> {
413        self.model_cache.get(model_name).map(|info| info.dimensions)
414    }
415}
416
417/// HuggingFace model manager for multiple models
418#[derive(Debug)]
419pub struct HuggingFaceModelManager {
420    embedders: HashMap<String, HuggingFaceEmbedder>,
421    default_model: String,
422}
423
424impl HuggingFaceModelManager {
425    /// Create a new model manager
426    pub fn new(default_model: String) -> Self {
427        Self {
428            embedders: HashMap::new(),
429            default_model,
430        }
431    }
432
433    /// Add a model to the manager
434    pub fn add_model(&mut self, name: String, config: HuggingFaceConfig) -> Result<()> {
435        let embedder = HuggingFaceEmbedder::new(config)?;
436        self.embedders.insert(name, embedder);
437        Ok(())
438    }
439
440    /// Get embeddings using specified model
441    pub async fn embed_with_model(
442        &mut self,
443        model_name: &str,
444        content: &EmbeddableContent,
445    ) -> Result<Vector> {
446        let embedder = self
447            .embedders
448            .get_mut(model_name)
449            .ok_or_else(|| anyhow!("Model not found: {}", model_name))?;
450        embedder.embed(content).await
451    }
452
453    /// Get embeddings using default model
454    pub async fn embed(&mut self, content: &EmbeddableContent) -> Result<Vector> {
455        self.embed_with_model(&self.default_model.clone(), content)
456            .await
457    }
458
459    /// List available models
460    pub fn list_models(&self) -> Vec<String> {
461        self.embedders.keys().cloned().collect()
462    }
463}
464
465/// Integration with existing embedding config
466impl From<EmbeddingConfig> for HuggingFaceConfig {
467    fn from(config: EmbeddingConfig) -> Self {
468        Self {
469            model_name: config.model_name,
470            cache_dir: None,
471            device: "cpu".to_string(),
472            batch_size: 32,
473            max_length: config.max_sequence_length,
474            pooling_strategy: if config.normalize {
475                PoolingStrategy::Mean
476            } else {
477                PoolingStrategy::Cls
478            },
479            trust_remote_code: false,
480            api_token: None,
481            use_inference_api: false,
482            inference_api_base_url: default_inference_api_base_url(),
483        }
484    }
485}
486
487#[cfg(test)]
488mod tests {
489    use super::*;
490    use anyhow::Result;
491
492    #[tokio::test]
493    async fn test_huggingface_embedder_creation() {
494        let embedder = HuggingFaceEmbedder::with_default_config();
495        assert!(embedder.is_ok());
496    }
497
498    #[tokio::test]
499    async fn test_model_loading() -> Result<()> {
500        let mut embedder = HuggingFaceEmbedder::with_default_config()?;
501        let result = embedder
502            .load_model("sentence-transformers/all-MiniLM-L6-v2")
503            .await;
504        assert!(result.is_ok());
505
506        let dimensions = embedder.get_model_dimensions("sentence-transformers/all-MiniLM-L6-v2");
507        assert_eq!(dimensions, Some(384));
508        Ok(())
509    }
510
511    #[tokio::test]
512    async fn test_text_embedding() -> Result<()> {
513        let mut embedder = HuggingFaceEmbedder::with_default_config()?;
514        let content = EmbeddableContent::Text("Hello, world!".to_string());
515
516        let result = embedder.embed(&content).await;
517        assert!(result.is_ok());
518
519        let embedding = result?;
520        assert_eq!(embedding.dimensions, 384);
521        Ok(())
522    }
523
524    #[tokio::test]
525    async fn test_rdf_resource_embedding() -> Result<()> {
526        let mut embedder = HuggingFaceEmbedder::with_default_config()?;
527        let mut properties = HashMap::new();
528        properties.insert("type".to_string(), vec!["Person".to_string()]);
529
530        let content = EmbeddableContent::RdfResource {
531            uri: "http://example.org/person/1".to_string(),
532            label: Some("John Doe".to_string()),
533            description: Some("A person in the knowledge graph".to_string()),
534            properties,
535        };
536
537        let result = embedder.embed(&content).await;
538        assert!(result.is_ok());
539        Ok(())
540    }
541
542    #[tokio::test]
543    async fn test_batch_embedding() -> Result<()> {
544        let mut embedder = HuggingFaceEmbedder::with_default_config()?;
545        let contents = vec![
546            EmbeddableContent::Text("First text".to_string()),
547            EmbeddableContent::Text("Second text".to_string()),
548            EmbeddableContent::Text("Third text".to_string()),
549        ];
550
551        let result = embedder.embed_batch(&contents).await;
552        assert!(result.is_ok());
553
554        let embeddings = result?;
555        assert_eq!(embeddings.len(), 3);
556        Ok(())
557    }
558
559    #[tokio::test]
560    async fn test_model_manager() {
561        let mut manager = HuggingFaceModelManager::new("default".to_string());
562        let config = HuggingFaceConfig::default();
563
564        let result = manager.add_model("default".to_string(), config);
565        assert!(result.is_ok());
566
567        let models = manager.list_models();
568        assert!(models.contains(&"default".to_string()));
569    }
570
571    #[test]
572    fn test_config_conversion() {
573        let embedding_config = EmbeddingConfig {
574            model_name: "test-model".to_string(),
575            dimensions: 768,
576            max_sequence_length: 512,
577            normalize: true,
578        };
579
580        let hf_config: HuggingFaceConfig = embedding_config.into();
581        assert_eq!(hf_config.model_name, "test-model");
582        assert_eq!(hf_config.max_length, 512);
583        assert!(matches!(hf_config.pooling_strategy, PoolingStrategy::Mean));
584    }
585
586    /// Regression test for the P1 finding: `use_inference_api` requires an
587    /// `api_token`; without one it must fail loudly instead of silently
588    /// falling back to the mock while `use_inference_api` is `true`.
589    #[tokio::test]
590    async fn test_use_inference_api_without_token_errors() {
591        let config = HuggingFaceConfig {
592            use_inference_api: true,
593            api_token: None,
594            ..Default::default()
595        };
596        let mut embedder = HuggingFaceEmbedder::new(config).expect("embedder should construct");
597        let content = EmbeddableContent::Text("hello".to_string());
598        let result = embedder.embed(&content).await;
599        assert!(result.is_err(), "missing api_token must be a hard error");
600    }
601
602    /// Regression test for the P1 finding: `HuggingFaceEmbedder` used to
603    /// silently return hash-noise labeled as if it were a transformer
604    /// embedding. The default config (no `use_inference_api`) must still
605    /// work offline via the explicitly-named `deterministic_mock_embedding`,
606    /// but must never be mistaken for real inference by identical inputs
607    /// producing identical (deterministic) — not text-semantic — output.
608    #[tokio::test]
609    async fn test_default_config_uses_deterministic_mock_not_real_api() -> Result<()> {
610        let mut embedder = HuggingFaceEmbedder::with_default_config()?;
611        assert!(!embedder.config.use_inference_api);
612
613        let a = embedder
614            .embed(&EmbeddableContent::Text("hello world".to_string()))
615            .await?;
616        let b = embedder
617            .embed(&EmbeddableContent::Text("hello world".to_string()))
618            .await?;
619        // Deterministic (same hash seed) but NOT claiming semantic meaning.
620        assert_eq!(a.as_f32(), b.as_f32());
621        Ok(())
622    }
623
624    /// `parse_single_embedding` must correctly handle both the pooled
625    /// (`[f32; d]`) and token-level (`[[f32; d]; tokens]`, mean-pooled)
626    /// response shapes returned by different HF feature-extraction models.
627    #[test]
628    fn test_parse_single_embedding_pooled_and_token_level() -> Result<()> {
629        let pooled = serde_json::json!([0.1, 0.2, 0.3]);
630        let pooled_vec = HuggingFaceEmbedder::parse_single_embedding(&pooled)?;
631        assert_eq!(pooled_vec.as_f32(), vec![0.1, 0.2, 0.3]);
632
633        // Two tokens: [1.0, 1.0] and [3.0, 3.0] -> mean = [2.0, 2.0]
634        let token_level = serde_json::json!([[1.0, 1.0], [3.0, 3.0]]);
635        let pooled_from_tokens = HuggingFaceEmbedder::parse_single_embedding(&token_level)?;
636        assert_eq!(pooled_from_tokens.as_f32(), vec![2.0, 2.0]);
637        Ok(())
638    }
639
640    /// `parse_inference_response` must error (not silently truncate/pad) on
641    /// a count mismatch between the number of input texts and returned
642    /// embeddings.
643    #[test]
644    fn test_parse_inference_response_count_mismatch_errors() {
645        let raw = serde_json::json!([[0.1, 0.2], [0.3, 0.4]]);
646        let result = HuggingFaceEmbedder::parse_inference_response(&raw, 3);
647        assert!(result.is_err());
648    }
649
650    #[test]
651    fn test_pooling_strategies() {
652        let strategies = vec![
653            PoolingStrategy::Cls,
654            PoolingStrategy::Mean,
655            PoolingStrategy::Max,
656            PoolingStrategy::AttentionWeighted,
657        ];
658
659        for strategy in strategies {
660            let config = HuggingFaceConfig {
661                pooling_strategy: strategy,
662                ..Default::default()
663            };
664            assert!(matches!(
665                config.pooling_strategy,
666                PoolingStrategy::Cls
667                    | PoolingStrategy::Mean
668                    | PoolingStrategy::Max
669                    | PoolingStrategy::AttentionWeighted
670            ));
671        }
672    }
673}