Skip to main content

apollo/memory/
embeddings.rs

1use 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
18// Noop provider for keyword-only fallback
19pub 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
36// OpenAI provider
37pub 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 // default
73        };
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
230// Factory function
231pub 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}