ubiquity_database/
embeddings.rs1use async_trait::async_trait;
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6use tokio::sync::RwLock;
7use std::sync::Arc;
8
9#[async_trait]
11pub trait EmbeddingGenerator: Send + Sync {
12 async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError>;
14
15 async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError>;
17
18 fn dimension(&self) -> usize;
20}
21
22pub 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 {
44 let cache = self.cache.read().await;
45 if let Some(embedding) = cache.get(text) {
46 return Ok(embedding.clone());
47 }
48 }
49
50 let embedding = self.inner.generate(text).await?;
52
53 {
55 let mut cache = self.cache.write().await;
56 if cache.len() >= self.max_cache_size {
57 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 {
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 if !uncached_texts.is_empty() {
89 let embeddings = self.inner.generate_batch(&uncached_texts).await?;
90
91 {
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 Ok(results.into_iter().map(|opt| opt.unwrap()).collect())
107 }
108
109 fn dimension(&self) -> usize {
110 self.inner.dimension()
111 }
112}
113
114pub 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
192pub 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#[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 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 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}