Skip to main content

ubiquity_database/
embeddings.rs

1//! Embedding generation and management
2
3use async_trait::async_trait;
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6use tokio::sync::RwLock;
7use std::sync::Arc;
8
9/// Embedding generator trait
10#[async_trait]
11pub trait EmbeddingGenerator: Send + Sync {
12    /// Generate embedding for text
13    async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError>;
14    
15    /// Generate embeddings for multiple texts
16    async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError>;
17    
18    /// Get embedding dimension
19    fn dimension(&self) -> usize;
20}
21
22/// Cached embedding generator
23pub struct CachedEmbeddingGenerator {
24    inner: Box<dyn EmbeddingGenerator>,
25    cache: Arc<RwLock<HashMap<String, Vec<f32>>>>,
26    max_cache_size: usize,
27}
28
29impl CachedEmbeddingGenerator {
30    pub fn new(inner: Box<dyn EmbeddingGenerator>, max_cache_size: usize) -> Self {
31        Self {
32            inner,
33            cache: Arc::new(RwLock::new(HashMap::new())),
34            max_cache_size,
35        }
36    }
37}
38
39#[async_trait]
40impl EmbeddingGenerator for CachedEmbeddingGenerator {
41    async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError> {
42        // Check cache first
43        {
44            let cache = self.cache.read().await;
45            if let Some(embedding) = cache.get(text) {
46                return Ok(embedding.clone());
47            }
48        }
49        
50        // Generate new embedding
51        let embedding = self.inner.generate(text).await?;
52        
53        // Store in cache
54        {
55            let mut cache = self.cache.write().await;
56            if cache.len() >= self.max_cache_size {
57                // Simple eviction: remove first item
58                if let Some(key) = cache.keys().next().cloned() {
59                    cache.remove(&key);
60                }
61            }
62            cache.insert(text.to_string(), embedding.clone());
63        }
64        
65        Ok(embedding)
66    }
67    
68    async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError> {
69        let mut results = Vec::with_capacity(texts.len());
70        let mut uncached_texts = Vec::new();
71        let mut uncached_indices = Vec::new();
72        
73        // Check cache for each text
74        {
75            let cache = self.cache.read().await;
76            for (i, text) in texts.iter().enumerate() {
77                if let Some(embedding) = cache.get(text) {
78                    results.push(Some(embedding.clone()));
79                } else {
80                    results.push(None);
81                    uncached_texts.push(text.clone());
82                    uncached_indices.push(i);
83                }
84            }
85        }
86        
87        // Generate embeddings for uncached texts
88        if !uncached_texts.is_empty() {
89            let embeddings = self.inner.generate_batch(&uncached_texts).await?;
90            
91            // Store in cache and results
92            {
93                let mut cache = self.cache.write().await;
94                for (i, (text, embedding)) in uncached_texts.iter().zip(embeddings.iter()).enumerate() {
95                    let result_idx = uncached_indices[i];
96                    results[result_idx] = Some(embedding.clone());
97                    
98                    if cache.len() < self.max_cache_size {
99                        cache.insert(text.clone(), embedding.clone());
100                    }
101                }
102            }
103        }
104        
105        // Convert Option<Vec<f32>> to Vec<Vec<f32>>
106        Ok(results.into_iter().map(|opt| opt.unwrap()).collect())
107    }
108    
109    fn dimension(&self) -> usize {
110        self.inner.dimension()
111    }
112}
113
114/// OpenAI-compatible embedding generator
115pub struct OpenAIEmbeddingGenerator {
116    client: reqwest::Client,
117    api_key: String,
118    model: String,
119    dimension: usize,
120}
121
122impl OpenAIEmbeddingGenerator {
123    pub fn new(api_key: String, model: String, dimension: usize) -> Self {
124        Self {
125            client: reqwest::Client::new(),
126            api_key,
127            model,
128            dimension,
129        }
130    }
131}
132
133#[async_trait]
134impl EmbeddingGenerator for OpenAIEmbeddingGenerator {
135    async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError> {
136        let response = self.client
137            .post("https://api.openai.com/v1/embeddings")
138            .header("Authorization", format!("Bearer {}", self.api_key))
139            .json(&serde_json::json!({
140                "model": self.model,
141                "input": text,
142                "dimensions": self.dimension,
143            }))
144            .send()
145            .await
146            .map_err(|e| crate::DatabaseError::Other(e.into()))?;
147        
148        let result: OpenAIEmbeddingResponse = response
149            .json()
150            .await
151            .map_err(|e| crate::DatabaseError::Other(e.into()))?;
152        
153        Ok(result.data[0].embedding.clone())
154    }
155    
156    async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError> {
157        let response = self.client
158            .post("https://api.openai.com/v1/embeddings")
159            .header("Authorization", format!("Bearer {}", self.api_key))
160            .json(&serde_json::json!({
161                "model": self.model,
162                "input": texts,
163                "dimensions": self.dimension,
164            }))
165            .send()
166            .await
167            .map_err(|e| crate::DatabaseError::Other(e.into()))?;
168        
169        let result: OpenAIEmbeddingResponse = response
170            .json()
171            .await
172            .map_err(|e| crate::DatabaseError::Other(e.into()))?;
173        
174        Ok(result.data.into_iter().map(|d| d.embedding).collect())
175    }
176    
177    fn dimension(&self) -> usize {
178        self.dimension
179    }
180}
181
182#[derive(Debug, Deserialize)]
183struct OpenAIEmbeddingResponse {
184    data: Vec<OpenAIEmbeddingData>,
185}
186
187#[derive(Debug, Deserialize)]
188struct OpenAIEmbeddingData {
189    embedding: Vec<f32>,
190}
191
192/// Create embedding generator based on configuration
193pub fn create_embedding_generator(config: &crate::config::EmbeddingConfig) -> Box<dyn EmbeddingGenerator> {
194    let generator: Box<dyn EmbeddingGenerator> = match config.provider.as_str() {
195        "openai" => {
196            let api_key = std::env::var("OPENAI_API_KEY")
197                .unwrap_or_else(|_| panic!("OPENAI_API_KEY environment variable not set"));
198            Box::new(OpenAIEmbeddingGenerator::new(
199                api_key,
200                config.model.clone(),
201                config.dimension,
202            ))
203        }
204        _ => panic!("Unsupported embedding provider: {}", config.provider),
205    };
206    
207    if config.cache_enabled {
208        Box::new(CachedEmbeddingGenerator::new(generator, config.cache_size))
209    } else {
210        generator
211    }
212}
213
214/// Consciousness embedding metadata
215#[derive(Debug, Clone, Serialize, Deserialize)]
216pub struct ConsciousnessEmbedding {
217    pub agent_id: String,
218    pub timestamp: chrono::DateTime<chrono::Utc>,
219    pub level: f64,
220    pub coherence: f64,
221    pub phase: ubiquity_core::DevelopmentPhase,
222    pub embedding: Vec<f32>,
223    pub text_representation: String,
224}
225
226impl ConsciousnessEmbedding {
227    /// Create text representation for embedding
228    pub fn create_text_representation(state: &ubiquity_core::ConsciousnessState) -> String {
229        format!(
230            "Agent {} consciousness at {:.2} coherence {:.2} phase {:?} breakthrough {}",
231            state.agent_id,
232            state.level.value(),
233            state.coherence,
234            state.phase,
235            state.breakthrough_detected
236        )
237    }
238    
239    /// Create from consciousness state
240    pub async fn from_state(
241        state: &ubiquity_core::ConsciousnessState,
242        generator: &dyn EmbeddingGenerator,
243    ) -> Result<Self, crate::DatabaseError> {
244        let text = Self::create_text_representation(state);
245        let embedding = generator.generate(&text).await?;
246        
247        Ok(Self {
248            agent_id: state.agent_id.clone(),
249            timestamp: state.timestamp,
250            level: state.level.value(),
251            coherence: state.coherence,
252            phase: state.phase,
253            embedding,
254            text_representation: text,
255        })
256    }
257}