1use 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#[allow(clippy::type_complexity)]
16pub trait EmbeddingProvider: Send + Sync {
17 fn embed(
19 &self,
20 texts: &[&str],
21 ) -> Pin<Box<dyn Future<Output = Result<Vec<Vec<f32>>, Error>> + Send + '_>>;
22
23 fn dimension(&self) -> usize;
25}
26
27pub 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
45pub 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 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, };
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 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 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
166pub struct EmbeddingMemory {
172 inner: Arc<dyn Memory>,
173 embedder: Arc<dyn EmbeddingProvider>,
174}
175
176impl EmbeddingMemory {
177 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 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(_) => {} Err(e) => {
203 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 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 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 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 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 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 use std::sync::atomic::{AtomicBool, Ordering};
488
489 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 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 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}