apollo/memory/
embeddings.rs1use anyhow::{Context, Result};
2use async_trait::async_trait;
3use serde::{Deserialize, Serialize};
4use std::sync::Arc;
5
6#[async_trait]
7pub trait EmbeddingProvider: Send + Sync {
8 fn name(&self) -> &str;
9 fn dimensions(&self) -> usize;
10 async fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>>;
11
12 async fn embed_one(&self, text: &str) -> Result<Vec<f32>> {
13 let mut results = self.embed(&[text]).await?;
14 results.pop().context("No embedding returned")
15 }
16}
17
18pub struct NoopEmbedding;
20
21#[async_trait]
22impl EmbeddingProvider for NoopEmbedding {
23 fn name(&self) -> &str {
24 "noop"
25 }
26
27 fn dimensions(&self) -> usize {
28 0
29 }
30
31 async fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
32 Ok(vec![vec![]; texts.len()])
33 }
34}
35
36pub struct OpenAiEmbedding {
38 client: reqwest::Client,
39 api_key: Option<String>,
40 base_url: String,
41 model: String,
42 dimensions: usize,
43}
44
45pub struct GeminiEmbedding {
46 client: reqwest::Client,
47 api_key: String,
48 model: String,
49 dimensions: usize,
50}
51
52impl GeminiEmbedding {
53 pub fn new(api_key: String, model: Option<String>) -> Self {
54 let model = model.unwrap_or_else(|| "text-embedding-004".to_string());
55 Self {
56 client: reqwest::Client::new(),
57 api_key,
58 model,
59 dimensions: 768,
60 }
61 }
62}
63
64impl OpenAiEmbedding {
65 pub fn new(api_key: Option<String>, model: Option<String>, base_url: Option<String>) -> Self {
66 let model = model.unwrap_or_else(|| "text-embedding-3-small".to_string());
67 let dimensions = if model.contains("text-embedding-3-small") {
68 1536
69 } else if model.contains("text-embedding-3-large") {
70 3072
71 } else {
72 1536 };
74
75 Self {
76 client: reqwest::Client::new(),
77 api_key,
78 base_url: base_url.unwrap_or_else(|| "https://api.openai.com".to_string()),
79 model,
80 dimensions,
81 }
82 }
83}
84
85#[derive(Serialize)]
86struct OpenAiEmbeddingRequest {
87 input: Vec<String>,
88 model: String,
89}
90
91#[derive(Deserialize)]
92struct OpenAiEmbeddingResponse {
93 data: Vec<OpenAiEmbeddingData>,
94}
95
96#[derive(Deserialize)]
97struct OpenAiEmbeddingData {
98 embedding: Vec<f32>,
99}
100
101#[async_trait]
102impl EmbeddingProvider for OpenAiEmbedding {
103 fn name(&self) -> &str {
104 &self.model
105 }
106
107 fn dimensions(&self) -> usize {
108 self.dimensions
109 }
110
111 async fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
112 let url = format!("{}/v1/embeddings", self.base_url);
113
114 let request = OpenAiEmbeddingRequest {
115 input: texts.iter().map(|s| s.to_string()).collect(),
116 model: self.model.clone(),
117 };
118
119 let mut request_builder = self
120 .client
121 .post(&url)
122 .header("Content-Type", "application/json");
123 if let Some(api_key) = &self.api_key {
124 request_builder =
125 request_builder.header("Authorization", format!("Bearer {}", api_key));
126 }
127 let response = request_builder
128 .json(&request)
129 .send()
130 .await
131 .context("Failed to send embedding request")?;
132
133 if !response.status().is_success() {
134 let status = response.status();
135 let body = response.text().await.unwrap_or_default();
136 anyhow::bail!("OpenAI API error {}: {}", status, body);
137 }
138
139 let response: OpenAiEmbeddingResponse = response
140 .json()
141 .await
142 .context("Failed to parse embedding response")?;
143
144 Ok(response.data.into_iter().map(|d| d.embedding).collect())
145 }
146}
147
148#[derive(Serialize)]
149struct GeminiPart {
150 text: String,
151}
152
153#[derive(Serialize)]
154struct GeminiContent {
155 parts: Vec<GeminiPart>,
156}
157
158#[derive(Serialize)]
159struct GeminiEmbeddingRequest {
160 model: String,
161 content: GeminiContent,
162}
163
164#[derive(Deserialize)]
165struct GeminiEmbeddingResponse {
166 embedding: Option<GeminiEmbeddingValues>,
167}
168
169#[derive(Deserialize)]
170struct GeminiEmbeddingValues {
171 values: Vec<f32>,
172}
173
174#[async_trait]
175impl EmbeddingProvider for GeminiEmbedding {
176 fn name(&self) -> &str {
177 &self.model
178 }
179
180 fn dimensions(&self) -> usize {
181 self.dimensions
182 }
183
184 async fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
185 let mut results = Vec::with_capacity(texts.len());
186 for text in texts {
187 let url = format!(
188 "https://generativelanguage.googleapis.com/v1beta/models/{}:embedContent",
189 self.model
190 );
191 let request = GeminiEmbeddingRequest {
192 model: format!("models/{}", self.model),
193 content: GeminiContent {
194 parts: vec![GeminiPart {
195 text: (*text).to_string(),
196 }],
197 },
198 };
199
200 let response = self
201 .client
202 .post(&url)
203 .header("x-goog-api-key", &self.api_key)
204 .json(&request)
205 .send()
206 .await
207 .context("Failed to send Gemini embedding request")?;
208
209 if !response.status().is_success() {
210 let status = response.status();
211 let body = response.text().await.unwrap_or_default();
212 anyhow::bail!("Gemini API error {}: {}", status, body);
213 }
214
215 let response: GeminiEmbeddingResponse = response
216 .json()
217 .await
218 .context("Failed to parse Gemini embedding response")?;
219
220 let values = response
221 .embedding
222 .map(|embedding| embedding.values)
223 .context("Gemini response did not include embedding values")?;
224 results.push(values);
225 }
226 Ok(results)
227 }
228}
229
230pub fn create_embedding_provider(
232 provider_type: &str,
233 api_key: Option<String>,
234 model: Option<String>,
235 base_url: Option<String>,
236) -> Result<Arc<dyn EmbeddingProvider>> {
237 match provider_type.to_lowercase().as_str() {
238 "noop" | "none" | "keyword" => Ok(Arc::new(NoopEmbedding)),
239 "openai" => {
240 let api_key = api_key.context("OpenAI API key required")?;
241 Ok(Arc::new(OpenAiEmbedding::new(
242 Some(api_key),
243 model,
244 base_url,
245 )))
246 }
247 "openai_compat" => Ok(Arc::new(OpenAiEmbedding::new(api_key, model, base_url))),
248 "ollama" | "local" => {
249 let model = model.or_else(|| Some("nomic-embed-text".to_string()));
250 let base_url = base_url.or_else(|| Some("http://localhost:11434".to_string()));
251 Ok(Arc::new(OpenAiEmbedding::new(None, model, base_url)))
252 }
253 "gemini" => {
254 let api_key = api_key.context("Gemini API key required")?;
255 Ok(Arc::new(GeminiEmbedding::new(api_key, model)))
256 }
257 _ => anyhow::bail!("Unknown embedding provider: {}", provider_type),
258 }
259}