Skip to main content

locus_sdk/infrastructure/
embeddings.rs

1#[cfg(feature = "local-embedding")]
2use std::sync::Arc;
3
4#[cfg(any(feature = "http-providers", feature = "local-embedding"))]
5use anyhow::{Result, anyhow};
6#[cfg(any(feature = "http-providers", feature = "local-embedding"))]
7use async_trait::async_trait;
8#[cfg(any(feature = "http-providers", feature = "local-embedding"))]
9use locus_core_rs::domain::contracts::EmbeddingProvider;
10#[cfg(feature = "http-providers")]
11use serde::{Deserialize, Serialize};
12
13#[cfg(feature = "http-providers")]
14#[derive(Debug, Serialize)]
15struct OllamaEmbeddingRequest<'a> {
16    model: &'a str,
17    prompt: &'a str,
18}
19
20#[cfg(feature = "http-providers")]
21#[derive(Debug, Deserialize)]
22struct OllamaEmbeddingResponse {
23    embedding: Option<Vec<f32>>,
24}
25
26#[cfg(feature = "http-providers")]
27#[derive(Clone)]
28pub struct OllamaEmbeddingProvider {
29    client: reqwest::Client,
30    endpoint: String,
31    model: String,
32}
33
34#[cfg(feature = "http-providers")]
35impl OllamaEmbeddingProvider {
36    pub fn new(endpoint: String, model: String) -> Self {
37        Self {
38            client: reqwest::Client::new(),
39            endpoint,
40            model,
41        }
42    }
43}
44
45#[cfg(feature = "http-providers")]
46#[async_trait]
47impl EmbeddingProvider for OllamaEmbeddingProvider {
48    fn model_name(&self) -> &str {
49        &self.model
50    }
51
52    async fn embed_async(&self, text: &str) -> Result<Vec<f32>> {
53        let response = self
54            .client
55            .post(&self.endpoint)
56            .json(&OllamaEmbeddingRequest {
57                model: &self.model,
58                prompt: text,
59            })
60            .send()
61            .await?
62            .error_for_status()?;
63
64        let body: OllamaEmbeddingResponse = response.json().await?;
65        match body.embedding {
66            Some(embedding) if !embedding.is_empty() => Ok(embedding),
67            _ => Err(anyhow!("embedding response missing vector")),
68        }
69    }
70}
71
72#[cfg(feature = "local-embedding")]
73pub struct LocalEmbeddingProvider {
74    model_name: String,
75    runtime: Arc<std::sync::Mutex<CandleRuntime>>,
76}
77
78#[cfg(feature = "local-embedding")]
79impl LocalEmbeddingProvider {
80    pub fn new(model_name: String, repo_id: String) -> Result<Self> {
81        let runtime = CandleRuntime::new(&repo_id)?;
82
83        Ok(Self {
84            model_name: format!("local-{}", model_name.trim().to_lowercase()),
85            runtime: Arc::new(std::sync::Mutex::new(runtime)),
86        })
87    }
88}
89
90#[cfg(feature = "local-embedding")]
91#[async_trait]
92impl EmbeddingProvider for LocalEmbeddingProvider {
93    fn model_name(&self) -> &str {
94        &self.model_name
95    }
96
97    async fn embed_async(&self, text: &str) -> Result<Vec<f32>> {
98        use anyhow::Context;
99
100        let runtime = Arc::clone(&self.runtime);
101        let input = text.to_string();
102
103        tokio::task::spawn_blocking(move || {
104            let runtime = runtime
105                .lock()
106                .map_err(|_| anyhow!("Local embedding runtime lock poisoned"))?;
107            runtime.embed(&input)
108        })
109        .await
110        .context("embedding worker join failure")?
111    }
112}
113
114#[cfg(feature = "local-embedding")]
115struct CandleRuntime {
116    model: candle_transformers::models::bert::BertModel,
117    tokenizer: tokenizers::Tokenizer,
118    device: candle_core::Device,
119}
120
121#[cfg(feature = "local-embedding")]
122impl CandleRuntime {
123    fn new(repo_id: &str) -> Result<Self> {
124        use anyhow::Context;
125        use candle_core::{DType, Device};
126        use candle_nn::VarBuilder;
127        use candle_transformers::models::bert::{BertModel, Config};
128        use hf_hub::{Repo, RepoType, api::sync::ApiBuilder};
129        use tokenizers::PaddingParams;
130
131        let device = Device::Cpu;
132
133        let api = ApiBuilder::new()
134            .with_endpoint("https://huggingface.co".to_string())
135            .build()
136            .context("failed to create HuggingFace API client")?;
137        let repo = api.repo(Repo::new(repo_id.to_string(), RepoType::Model));
138
139        let config_path = repo
140            .get("config.json")
141            .with_context(|| format!("failed to fetch config.json from {repo_id}"))?;
142        let tokenizer_path = repo
143            .get("tokenizer.json")
144            .with_context(|| format!("failed to fetch tokenizer.json from {repo_id}"))?;
145        let weights_path = repo
146            .get("model.safetensors")
147            .with_context(|| format!("failed to fetch model.safetensors from {repo_id}"))?;
148
149        let config: Config = serde_json::from_str(
150            &std::fs::read_to_string(&config_path)
151                .with_context(|| format!("failed to read {}", config_path.display()))?,
152        )
153        .with_context(|| format!("failed to parse {}", config_path.display()))?;
154
155        let mut tokenizer = tokenizers::Tokenizer::from_file(tokenizer_path)
156            .map_err(|err| anyhow!("tokenizer error: {err}"))?;
157        tokenizer.with_padding(Some(PaddingParams {
158            strategy: tokenizers::PaddingStrategy::BatchLongest,
159            ..Default::default()
160        }));
161
162        let vb = unsafe {
163            VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, &device)
164                .context("failed to map safetensors weights")?
165        };
166        let model = BertModel::load(vb, &config).context("failed to load BERT model")?;
167
168        Ok(Self {
169            model,
170            tokenizer,
171            device,
172        })
173    }
174
175    fn embed(&self, text: &str) -> Result<Vec<f32>> {
176        let embeddings = self.embed_batch(&[text])?;
177        embeddings
178            .into_iter()
179            .next()
180            .ok_or_else(|| anyhow!("empty embedding output"))
181    }
182
183    fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
184        use anyhow::Context;
185        use candle_core::{DType, Tensor};
186
187        if texts.is_empty() {
188            return Ok(Vec::new());
189        }
190
191        let encodings = self
192            .tokenizer
193            .encode_batch(texts.to_vec(), true)
194            .map_err(|err| anyhow!("tokenization failed: {err}"))?;
195
196        let seq_len = encodings[0].get_ids().len();
197        let batch_size = texts.len();
198
199        let input_ids: Vec<u32> = encodings.iter().flat_map(|e| e.get_ids().to_vec()).collect();
200        let attention_mask: Vec<u32> = encodings
201            .iter()
202            .flat_map(|e| e.get_attention_mask().to_vec())
203            .collect();
204        let token_type_ids: Vec<u32> = vec![0u32; batch_size * seq_len];
205
206        let input_ids = Tensor::from_vec(input_ids, (batch_size, seq_len), &self.device)?;
207        let attention_mask = Tensor::from_vec(attention_mask, (batch_size, seq_len), &self.device)?;
208        let token_type_ids = Tensor::from_vec(token_type_ids, (batch_size, seq_len), &self.device)?;
209
210        let output = self
211            .model
212            .forward(&input_ids, &token_type_ids, Some(&attention_mask))
213            .context("local embedding forward pass failed")?;
214
215        let mask_f32 = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?;
216        let masked = output.broadcast_mul(&mask_f32)?;
217        let summed = masked.sum(1)?;
218        let counts = mask_f32.sum(1)?;
219        let pooled = summed.broadcast_div(&counts)?;
220
221        let norm = pooled.sqr()?.sum_keepdim(1)?.sqrt()?;
222        let normalized = pooled.broadcast_div(&norm)?;
223
224        Ok(normalized.to_vec2::<f32>()?)
225    }
226}