locus_sdk/infrastructure/
embeddings.rs1#[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}