1use crate::{EmbeddableContent, EmbeddingConfig, Vector};
4use anyhow::{anyhow, Result};
5use scirs2_core::random::Random;
6use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9#[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 #[serde(skip_serializing_if = "Option::is_none", default)]
24 pub api_token: Option<String>,
25 #[serde(default)]
30 pub use_inference_api: bool,
31 #[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#[derive(Debug, Clone, Serialize, Deserialize)]
43pub enum PoolingStrategy {
44 Cls,
46 Mean,
48 Max,
50 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#[derive(Debug)]
73pub struct HuggingFaceEmbedder {
74 config: HuggingFaceConfig,
75 model_cache: HashMap<String, ModelInfo>,
76}
77
78#[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 pub fn new(config: HuggingFaceConfig) -> Result<Self> {
90 Ok(Self {
91 config,
92 model_cache: HashMap::new(),
93 })
94 }
95
96 pub fn with_default_config() -> Result<Self> {
98 Self::new(HuggingFaceConfig::default())
99 }
100
101 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 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 async fn get_model_info(&self, model_name: &str) -> Result<ModelInfo> {
117 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, };
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 pub async fn embed_batch(&mut self, contents: &[EmbeddableContent]) -> Result<Vec<Vector>> {
138 if contents.is_empty() {
139 return Ok(vec![]);
140 }
141
142 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 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 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 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 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 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 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 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 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 fn deterministic_mock_embedding(&self, text: &str, dimensions: usize) -> Result<Vector> {
373 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); }
387
388 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 pub fn get_cached_models(&self) -> Vec<String> {
403 self.model_cache.keys().cloned().collect()
404 }
405
406 pub fn clear_cache(&mut self) {
408 self.model_cache.clear();
409 }
410
411 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#[derive(Debug)]
419pub struct HuggingFaceModelManager {
420 embedders: HashMap<String, HuggingFaceEmbedder>,
421 default_model: String,
422}
423
424impl HuggingFaceModelManager {
425 pub fn new(default_model: String) -> Self {
427 Self {
428 embedders: HashMap::new(),
429 default_model,
430 }
431 }
432
433 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 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 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 pub fn list_models(&self) -> Vec<String> {
461 self.embedders.keys().cloned().collect()
462 }
463}
464
465impl 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 #[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 #[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 assert_eq!(a.as_f32(), b.as_f32());
621 Ok(())
622 }
623
624 #[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 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 #[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}