Skip to main content

heartbit_core/memory/
embedding.rs

1//! Embedding providers for semantic memory retrieval.
2
3use std::future::Future;
4use std::pin::Pin;
5use std::sync::Arc;
6
7use serde::Deserialize;
8
9use crate::auth::TenantScope;
10use crate::error::Error;
11
12use super::{Memory, MemoryEntry};
13
14/// Trait for generating text embeddings.
15#[allow(clippy::type_complexity)]
16pub trait EmbeddingProvider: Send + Sync {
17    /// Generate embeddings for each input string (parallel batch).
18    fn embed(
19        &self,
20        texts: &[&str],
21    ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>>;
22
23    /// Dimension of vectors returned by [`embed`](Self::embed).
24    fn dimension(&self) -> usize;
25}
26
27/// No-op embedding provider — returns empty results.
28/// Used when no embedding API is configured (graceful degradation).
29pub struct NoopEmbedding;
30
31impl EmbeddingProvider for NoopEmbedding {
32    fn embed(
33        &self,
34        texts: &[&str],
35    ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>> {
36        let len = texts.len();
37        Box::pin(async move { Ok(vec![vec![]; len]) })
38    }
39
40    fn dimension(&self) -> usize {
41        0
42    }
43}
44
45/// OpenAI-compatible embedding provider.
46///
47/// Calls `POST /v1/embeddings` with the configured model.
48/// Works with OpenAI API and compatible endpoints.
49pub struct OpenAiEmbedding {
50    client: reqwest::Client,
51    api_key: String,
52    model: String,
53    base_url: String,
54    dimension: usize,
55}
56
57impl OpenAiEmbedding {
58    /// Create an `OpenAiEmbedding` provider.
59    ///
60    /// SECURITY (F-MEM-4): the HTTP client is hardened with
61    /// `redirect::Policy::none()`, `https_only(true)`,
62    /// `connect_timeout(10s)`, `timeout(60s)`, and `.no_proxy()`. Without
63    /// these, a slow-loris embedding endpoint wedges every `memory_store`,
64    /// and a redirect to a non-HTTPS host would leak the Bearer API key.
65    pub fn new(api_key: impl Into<String>, model: impl Into<String>) -> Self {
66        let model = model.into();
67        let dimension = match model.as_str() {
68            "text-embedding-3-small" => 1536,
69            "text-embedding-3-large" => 3072,
70            "text-embedding-ada-002" => 1536,
71            _ => 1536, // default
72        };
73        let client = reqwest::Client::builder()
74            .redirect(reqwest::redirect::Policy::none())
75            .https_only(true)
76            .no_proxy()
77            .connect_timeout(std::time::Duration::from_secs(10))
78            .timeout(std::time::Duration::from_secs(60))
79            .build()
80            .expect("failed to build hardened HTTPS client for OpenAiEmbedding");
81        Self {
82            client,
83            api_key: api_key.into(),
84            model,
85            base_url: "https://api.openai.com".into(),
86            dimension,
87        }
88    }
89
90    /// Override the base URL.
91    ///
92    /// SECURITY (F-MEM-4): the URL must be HTTPS — `https_only(true)` is
93    /// already set on the client, so a plaintext URL here will fail at
94    /// request time. For local non-secret endpoints (Ollama embeddings,
95    /// vLLM), build a separate provider with a custom client.
96    pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
97        self.base_url = base_url.into();
98        self
99    }
100
101    /// Override the auto-detected embedding dimension.
102    pub fn with_dimension(mut self, dimension: usize) -> Self {
103        self.dimension = dimension;
104        self
105    }
106}
107
108#[derive(Deserialize)]
109struct EmbeddingResponse {
110    data: Vec<EmbeddingData>,
111}
112
113#[derive(Deserialize)]
114struct EmbeddingData {
115    embedding: Vec<f32>,
116}
117
118impl EmbeddingProvider for OpenAiEmbedding {
119    fn embed(
120        &self,
121        texts: &[&str],
122    ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>> {
123        let input: Vec<String> = texts.iter().map(|t| t.to_string()).collect();
124        Box::pin(async move {
125            if input.is_empty() {
126                return Ok(vec![]);
127            }
128
129            let body = serde_json::json!({
130                "model": self.model,
131                "input": input,
132            });
133
134            let resp = self
135                .client
136                .post(format!("{}/v1/embeddings", self.base_url))
137                .header("Authorization", format!("Bearer {}", self.api_key))
138                .header("Content-Type", "application/json")
139                .json(&body)
140                .send()
141                .await
142                .map_err(|e| Error::Memory(format!("embedding request failed: {e}")))?;
143
144            if !resp.status().is_success() {
145                let status = resp.status();
146                let text = resp.text().await.unwrap_or_else(|_| "unknown error".into());
147                return Err(Error::Memory(format!(
148                    "embedding API returned {status}: {text}"
149                )));
150            }
151
152            let response: EmbeddingResponse = resp
153                .json()
154                .await
155                .map_err(|e| Error::Memory(format!("failed to parse embedding response: {e}")))?;
156
157            Ok(response.data.into_iter().map(|d| d.embedding).collect())
158        })
159    }
160
161    fn dimension(&self) -> usize {
162        self.dimension
163    }
164}
165
166/// Decorator that generates embeddings on store and passes through to inner Memory.
167///
168/// When storing a `MemoryEntry` without an embedding, this wrapper generates
169/// one via the configured `EmbeddingProvider` before delegating to the inner store.
170/// All other operations pass through unchanged.
171pub struct EmbeddingMemory {
172    inner: Arc<dyn Memory>,
173    embedder: Arc<dyn EmbeddingProvider>,
174}
175
176impl EmbeddingMemory {
177    /// Wrap `inner` so that stored entries lacking an embedding are vectorized via `embedder` first.
178    pub fn new(inner: Arc<dyn Memory>, embedder: Arc<dyn EmbeddingProvider>) -> Self {
179        Self { inner, embedder }
180    }
181}
182
183impl Memory for EmbeddingMemory {
184    fn store(
185        &self,
186        scope: &TenantScope,
187        entry: MemoryEntry,
188    ) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send + '_>> {
189        let scope = scope.clone();
190        Box::pin(async move {
191            let mut entry = entry;
192            // Only generate embedding if not already present and embedder is real (dimension > 0)
193            if entry.embedding.is_none() && self.embedder.dimension() > 0 {
194                match self.embedder.embed(&[&entry.content]).await {
195                    Ok(mut embeddings) if !embeddings.is_empty() => {
196                        let emb = embeddings.swap_remove(0);
197                        if !emb.is_empty() {
198                            entry.embedding = Some(emb);
199                        }
200                    }
201                    Ok(_) => {} // empty result, skip
202                    Err(e) => {
203                        // Log but don't fail — embedding is optional
204                        tracing::warn!("failed to generate embedding for memory {}: {e}", entry.id);
205                    }
206                }
207            }
208            self.inner.store(&scope, entry).await
209        })
210    }
211
212    fn recall(
213        &self,
214        scope: &TenantScope,
215        query: super::MemoryQuery,
216    ) -> Pin<Box<dyn Future<Output = Result<Vec<MemoryEntry>, Error>> + Send + '_>> {
217        let scope = scope.clone();
218        Box::pin(async move {
219            let mut query = query;
220            // Generate query embedding for hybrid retrieval when text is present
221            // and no embedding was already provided.
222            if query.query_embedding.is_none()
223                && query.text.is_some()
224                && self.embedder.dimension() > 0
225            {
226                let text = query.text.as_deref().unwrap_or_default();
227                match self.embedder.embed(&[text]).await {
228                    Ok(mut embeddings) if !embeddings.is_empty() => {
229                        let emb = embeddings.swap_remove(0);
230                        if !emb.is_empty() {
231                            query.query_embedding = Some(emb);
232                        }
233                    }
234                    Ok(_) => {}
235                    Err(e) => {
236                        // Log but don't fail — fall back to BM25-only
237                        tracing::warn!("failed to generate query embedding: {e}");
238                    }
239                }
240            }
241            self.inner.recall(&scope, query).await
242        })
243    }
244
245    fn update(
246        &self,
247        scope: &TenantScope,
248        id: &str,
249        content: String,
250    ) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send + '_>> {
251        let scope = scope.clone();
252        let id = id.to_string();
253        Box::pin(async move { self.inner.update(&scope, &id, content).await })
254    }
255
256    fn forget(
257        &self,
258        scope: &TenantScope,
259        id: &str,
260    ) -> Pin<Box<dyn Future<Output = Result<bool, Error>> + Send + '_>> {
261        let scope = scope.clone();
262        let id = id.to_string();
263        Box::pin(async move { self.inner.forget(&scope, &id).await })
264    }
265
266    fn add_link(
267        &self,
268        scope: &TenantScope,
269        id: &str,
270        related_id: &str,
271    ) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send + '_>> {
272        let scope = scope.clone();
273        let id = id.to_string();
274        let related_id = related_id.to_string();
275        Box::pin(async move { self.inner.add_link(&scope, &id, &related_id).await })
276    }
277
278    fn prune(
279        &self,
280        scope: &TenantScope,
281        min_strength: f64,
282        min_age: chrono::Duration,
283        agent_prefix: Option<&str>,
284    ) -> Pin<Box<dyn Future<Output = Result<usize, Error>> + Send + '_>> {
285        let scope = scope.clone();
286        let agent_prefix = agent_prefix.map(String::from);
287        Box::pin(async move {
288            self.inner
289                .prune(&scope, min_strength, min_age, agent_prefix.as_deref())
290                .await
291        })
292    }
293}
294
295#[cfg(test)]
296mod tests {
297    use super::*;
298    use crate::memory::in_memory::InMemoryStore;
299    use crate::memory::{Confidentiality, MemoryEntry, MemoryQuery, MemoryType};
300    use chrono::Utc;
301
302    fn test_scope() -> TenantScope {
303        TenantScope::default()
304    }
305
306    fn make_entry(id: &str, content: &str) -> MemoryEntry {
307        MemoryEntry {
308            id: id.into(),
309            agent: "test".into(),
310            content: content.into(),
311            category: "fact".into(),
312            tags: vec![],
313            created_at: Utc::now(),
314            last_accessed: Utc::now(),
315            access_count: 0,
316            importance: 5,
317            memory_type: MemoryType::default(),
318            keywords: vec![],
319            summary: None,
320            strength: 1.0,
321            related_ids: vec![],
322            source_ids: vec![],
323            embedding: None,
324            confidentiality: Confidentiality::default(),
325            author_user_id: None,
326            author_tenant_id: None,
327        }
328    }
329
330    #[test]
331    fn noop_embedding_returns_empty() {
332        let noop = NoopEmbedding;
333        assert_eq!(noop.dimension(), 0);
334        let rt = tokio::runtime::Builder::new_current_thread()
335            .build()
336            .unwrap();
337        let result = rt.block_on(noop.embed(&["hello", "world"])).unwrap();
338        assert_eq!(result.len(), 2);
339        assert!(result[0].is_empty());
340        assert!(result[1].is_empty());
341    }
342
343    #[test]
344    fn embedding_provider_is_object_safe() {
345        fn _accepts_dyn(_p: &dyn EmbeddingProvider) {}
346    }
347
348    #[test]
349    fn embedding_memory_is_send_sync() {
350        fn assert_send_sync<T: Send + Sync>() {}
351        assert_send_sync::<EmbeddingMemory>();
352    }
353
354    #[tokio::test]
355    async fn noop_embedding_skips_embedding_on_store() {
356        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
357        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(NoopEmbedding);
358        let em = EmbeddingMemory::new(store.clone(), embedder);
359
360        em.store(&test_scope(), make_entry("m1", "test content"))
361            .await
362            .unwrap();
363
364        let results = store
365            .recall(
366                &test_scope(),
367                MemoryQuery {
368                    limit: 10,
369                    ..Default::default()
370                },
371            )
372            .await
373            .unwrap();
374        assert_eq!(results.len(), 1);
375        assert!(results[0].embedding.is_none());
376    }
377
378    /// Fake embedding provider for testing that returns deterministic vectors.
379    struct FakeEmbedding;
380
381    impl EmbeddingProvider for FakeEmbedding {
382        fn embed(
383            &self,
384            texts: &[&str],
385        ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>> {
386            let results: Vec<Vec<f32>> = texts
387                .iter()
388                .map(|t| {
389                    // Simple deterministic embedding: first 4 bytes as f32 values
390                    let bytes = t.as_bytes();
391                    vec![
392                        bytes.first().copied().unwrap_or(0) as f32 / 255.0,
393                        bytes.get(1).copied().unwrap_or(0) as f32 / 255.0,
394                        bytes.get(2).copied().unwrap_or(0) as f32 / 255.0,
395                    ]
396                })
397                .collect();
398            Box::pin(async move { Ok(results) })
399        }
400
401        fn dimension(&self) -> usize {
402            3
403        }
404    }
405
406    #[tokio::test]
407    async fn embedding_memory_generates_embedding_on_store() {
408        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
409        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(FakeEmbedding);
410        let em = EmbeddingMemory::new(store.clone(), embedder);
411
412        em.store(&test_scope(), make_entry("m1", "hello"))
413            .await
414            .unwrap();
415
416        let results = store
417            .recall(
418                &test_scope(),
419                MemoryQuery {
420                    limit: 10,
421                    ..Default::default()
422                },
423            )
424            .await
425            .unwrap();
426        assert_eq!(results.len(), 1);
427        let emb = results[0]
428            .embedding
429            .as_ref()
430            .expect("embedding should be set");
431        assert_eq!(emb.len(), 3);
432    }
433
434    #[tokio::test]
435    async fn embedding_memory_preserves_existing_embedding() {
436        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
437        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(FakeEmbedding);
438        let em = EmbeddingMemory::new(store.clone(), embedder);
439
440        let mut entry = make_entry("m1", "hello");
441        entry.embedding = Some(vec![9.0, 8.0, 7.0]);
442        em.store(&test_scope(), entry).await.unwrap();
443
444        let results = store
445            .recall(
446                &test_scope(),
447                MemoryQuery {
448                    limit: 10,
449                    ..Default::default()
450                },
451            )
452            .await
453            .unwrap();
454        let emb = results[0].embedding.as_ref().unwrap();
455        // Should keep original, not overwrite with FakeEmbedding output
456        assert!((emb[0] - 9.0).abs() < f32::EPSILON);
457    }
458
459    #[tokio::test]
460    async fn embedding_memory_delegates_recall() {
461        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
462        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(NoopEmbedding);
463        let em = EmbeddingMemory::new(store.clone(), embedder);
464
465        store
466            .store(&test_scope(), make_entry("m1", "test"))
467            .await
468            .unwrap();
469        let results = em
470            .recall(
471                &test_scope(),
472                MemoryQuery {
473                    limit: 10,
474                    ..Default::default()
475                },
476            )
477            .await
478            .unwrap();
479        assert_eq!(results.len(), 1);
480        assert_eq!(results[0].id, "m1");
481    }
482
483    #[tokio::test]
484    async fn embedding_memory_generates_query_embedding_on_recall() {
485        // When EmbeddingMemory wraps a store and query has text,
486        // it should generate a query embedding for hybrid retrieval.
487        use std::sync::atomic::{AtomicBool, Ordering};
488
489        // Tracking embedding provider that records whether embed() was called
490        struct TrackingEmbedding {
491            called: Arc<AtomicBool>,
492        }
493
494        impl EmbeddingProvider for TrackingEmbedding {
495            fn embed(
496                &self,
497                _texts: &[&str],
498            ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>>
499            {
500                self.called.store(true, Ordering::SeqCst);
501                Box::pin(async { Ok(vec![vec![0.5, 0.5, 0.5]]) })
502            }
503
504            fn dimension(&self) -> usize {
505                3
506            }
507        }
508
509        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
510        let called = Arc::new(AtomicBool::new(false));
511        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(TrackingEmbedding {
512            called: called.clone(),
513        });
514        let em = EmbeddingMemory::new(store.clone(), embedder);
515
516        store
517            .store(&test_scope(), make_entry("m1", "hello world"))
518            .await
519            .unwrap();
520
521        // Recall with text query should trigger embedding generation
522        let _results = em
523            .recall(
524                &test_scope(),
525                MemoryQuery {
526                    text: Some("hello".into()),
527                    limit: 10,
528                    ..Default::default()
529                },
530            )
531            .await
532            .unwrap();
533
534        assert!(
535            called.load(Ordering::SeqCst),
536            "embed() should have been called for query text"
537        );
538    }
539
540    #[tokio::test]
541    async fn embedding_memory_skips_query_embedding_without_text() {
542        use std::sync::atomic::{AtomicBool, Ordering};
543
544        struct TrackingEmbedding {
545            called: Arc<AtomicBool>,
546        }
547
548        impl EmbeddingProvider for TrackingEmbedding {
549            fn embed(
550                &self,
551                _texts: &[&str],
552            ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>>
553            {
554                self.called.store(true, Ordering::SeqCst);
555                Box::pin(async { Ok(vec![vec![0.5, 0.5, 0.5]]) })
556            }
557
558            fn dimension(&self) -> usize {
559                3
560            }
561        }
562
563        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
564        let called = Arc::new(AtomicBool::new(false));
565        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(TrackingEmbedding {
566            called: called.clone(),
567        });
568        let em = EmbeddingMemory::new(store.clone(), embedder);
569
570        store
571            .store(&test_scope(), make_entry("m1", "hello world"))
572            .await
573            .unwrap();
574
575        // Recall WITHOUT text query should NOT generate embedding
576        let _results = em
577            .recall(
578                &test_scope(),
579                MemoryQuery {
580                    limit: 10,
581                    ..Default::default()
582                },
583            )
584            .await
585            .unwrap();
586
587        assert!(
588            !called.load(Ordering::SeqCst),
589            "embed() should NOT be called when no text query"
590        );
591    }
592
593    #[tokio::test]
594    async fn embedding_memory_delegates_forget() {
595        let store: Arc<dyn Memory> = Arc::new(InMemoryStore::new());
596        let embedder: Arc<dyn EmbeddingProvider> = Arc::new(NoopEmbedding);
597        let em = EmbeddingMemory::new(store.clone(), embedder);
598
599        store
600            .store(&test_scope(), make_entry("m1", "test"))
601            .await
602            .unwrap();
603        let removed = em.forget(&test_scope(), "m1").await.unwrap();
604        assert!(removed);
605
606        let results = store
607            .recall(
608                &test_scope(),
609                MemoryQuery {
610                    limit: 10,
611                    ..Default::default()
612                },
613            )
614            .await
615            .unwrap();
616        assert!(results.is_empty());
617    }
618}