ubiquity-database 0.1.1

Database abstraction layer for Ubiquity supporting SQLite and Astra DB
Documentation
//! Embedding generation and management

use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use tokio::sync::RwLock;
use std::sync::Arc;

/// Embedding generator trait
#[async_trait]
pub trait EmbeddingGenerator: Send + Sync {
    /// Generate embedding for text
    async fn generate(&self, text: &str) -> Result<Vec<f32>, crate::DatabaseError>;
    
    /// Generate embeddings for multiple texts
    async fn generate_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, crate::DatabaseError>;
    
    /// Get embedding dimension
    fn dimension(&self) -> usize;
}

/// Cached embedding generator
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> {
        // Check cache first
        {
            let cache = self.cache.read().await;
            if let Some(embedding) = cache.get(text) {
                return Ok(embedding.clone());
            }
        }
        
        // Generate new embedding
        let embedding = self.inner.generate(text).await?;
        
        // Store in cache
        {
            let mut cache = self.cache.write().await;
            if cache.len() >= self.max_cache_size {
                // Simple eviction: remove first item
                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();
        
        // Check cache for each text
        {
            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);
                }
            }
        }
        
        // Generate embeddings for uncached texts
        if !uncached_texts.is_empty() {
            let embeddings = self.inner.generate_batch(&uncached_texts).await?;
            
            // Store in cache and results
            {
                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());
                    }
                }
            }
        }
        
        // Convert Option<Vec<f32>> to Vec<Vec<f32>>
        Ok(results.into_iter().map(|opt| opt.unwrap()).collect())
    }
    
    fn dimension(&self) -> usize {
        self.inner.dimension()
    }
}

/// OpenAI-compatible embedding generator
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>,
}

/// Create embedding generator based on configuration
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
    }
}

/// Consciousness embedding metadata
#[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 {
    /// Create text representation for embedding
    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
        )
    }
    
    /// Create from consciousness state
    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,
        })
    }
}