use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use tokio::sync::RwLock;
use std::sync::Arc;
#[async_trait]
pub trait EmbeddingGenerator: Send + Sync {
async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError>;
async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError>;
fn dimension(&self) -> usize;
}
pub struct CachedEmbeddingGenerator {
inner: Box<dyn EmbeddingGenerator>,
cache: Arc<RwLock<HashMap<String, Vec<f32>>>>,
max_cache_size: usize,
}
impl CachedEmbeddingGenerator {
pub fn new(inner: Box<dyn EmbeddingGenerator>, max_cache_size: usize) -> Self {
Self {
inner,
cache: Arc::new(RwLock::new(HashMap::new())),
max_cache_size,
}
}
}
#[async_trait]
impl EmbeddingGenerator for CachedEmbeddingGenerator {
async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError> {
{
let cache = self.cache.read().await;
if let Some(embedding) = cache.get(text) {
return Ok(embedding.clone());
}
}
let embedding = self.inner.generate(text).await?;
{
let mut cache = self.cache.write().await;
if cache.len() >= self.max_cache_size {
if let Some(key) = cache.keys().next().cloned() {
cache.remove(&key);
}
}
cache.insert(text.to_string(), embedding.clone());
}
Ok(embedding)
}
async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError> {
let mut results = Vec::with_capacity(texts.len());
let mut uncached_texts = Vec::new();
let mut uncached_indices = Vec::new();
{
let cache = self.cache.read().await;
for (i, text) in texts.iter().enumerate() {
if let Some(embedding) = cache.get(text) {
results.push(Some(embedding.clone()));
} else {
results.push(None);
uncached_texts.push(text.clone());
uncached_indices.push(i);
}
}
}
if !uncached_texts.is_empty() {
let embeddings = self.inner.generate_batch(&uncached_texts).await?;
{
let mut cache = self.cache.write().await;
for (i, (text, embedding)) in uncached_texts.iter().zip(embeddings.iter()).enumerate() {
let result_idx = uncached_indices[i];
results[result_idx] = Some(embedding.clone());
if cache.len() < self.max_cache_size {
cache.insert(text.clone(), embedding.clone());
}
}
}
}
Ok(results.into_iter().map(|opt| opt.unwrap()).collect())
}
fn dimension(&self) -> usize {
self.inner.dimension()
}
}
pub struct OpenAIEmbeddingGenerator {
client: reqwest::Client,
api_key: String,
model: String,
dimension: usize,
}
impl OpenAIEmbeddingGenerator {
pub fn new(api_key: String, model: String, dimension: usize) -> Self {
Self {
client: reqwest::Client::new(),
api_key,
model,
dimension,
}
}
}
#[async_trait]
impl EmbeddingGenerator for OpenAIEmbeddingGenerator {
async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError> {
let response = self.client
.post("https://api.openai.com/v1/embeddings")
.header("Authorization", format!("Bearer {}", self.api_key))
.json(&serde_json::json!({
"model": self.model,
"input": text,
"dimensions": self.dimension,
}))
.send()
.await
.map_err(|e| crate::DatabaseError::Other(e.into()))?;
let result: OpenAIEmbeddingResponse = response
.json()
.await
.map_err(|e| crate::DatabaseError::Other(e.into()))?;
Ok(result.data[0].embedding.clone())
}
async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError> {
let response = self.client
.post("https://api.openai.com/v1/embeddings")
.header("Authorization", format!("Bearer {}", self.api_key))
.json(&serde_json::json!({
"model": self.model,
"input": texts,
"dimensions": self.dimension,
}))
.send()
.await
.map_err(|e| crate::DatabaseError::Other(e.into()))?;
let result: OpenAIEmbeddingResponse = response
.json()
.await
.map_err(|e| crate::DatabaseError::Other(e.into()))?;
Ok(result.data.into_iter().map(|d| d.embedding).collect())
}
fn dimension(&self) -> usize {
self.dimension
}
}
#[derive(Debug, Deserialize)]
struct OpenAIEmbeddingResponse {
data: Vec<OpenAIEmbeddingData>,
}
#[derive(Debug, Deserialize)]
struct OpenAIEmbeddingData {
embedding: Vec<f32>,
}
pub fn create_embedding_generator(config: &crate::config::EmbeddingConfig) -> Box<dyn EmbeddingGenerator> {
let generator: Box<dyn EmbeddingGenerator> = match config.provider.as_str() {
"openai" => {
let api_key = std::env::var("OPENAI_API_KEY")
.unwrap_or_else(|_| panic!("OPENAI_API_KEY environment variable not set"));
Box::new(OpenAIEmbeddingGenerator::new(
api_key,
config.model.clone(),
config.dimension,
))
}
_ => panic!("Unsupported embedding provider: {}", config.provider),
};
if config.cache_enabled {
Box::new(CachedEmbeddingGenerator::new(generator, config.cache_size))
} else {
generator
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConsciousnessEmbedding {
pub agent_id: String,
pub timestamp: chrono::DateTime<chrono::Utc>,
pub level: f64,
pub coherence: f64,
pub phase: ubiquity_core::DevelopmentPhase,
pub embedding: Vec<f32>,
pub text_representation: String,
}
impl ConsciousnessEmbedding {
pub fn create_text_representation(state: &ubiquity_core::ConsciousnessState) -> String {
format!(
"Agent {} consciousness at {:.2} coherence {:.2} phase {:?} breakthrough {}",
state.agent_id,
state.level.value(),
state.coherence,
state.phase,
state.breakthrough_detected
)
}
pub async fn from_state(
state: &ubiquity_core::ConsciousnessState,
generator: &dyn EmbeddingGenerator,
) -> Result<Self, crate::DatabaseError> {
let text = Self::create_text_representation(state);
let embedding = generator.generate(&text).await?;
Ok(Self {
agent_id: state.agent_id.clone(),
timestamp: state.timestamp,
level: state.level.value(),
coherence: state.coherence,
phase: state.phase,
embedding,
text_representation: text,
})
}
}