xz-rag 0.1.1

Multi-channel Retrieval-Augmented Generation engine
Documentation
use async_trait::async_trait;
use std::sync::Arc;

use crate::error::RagError;
use crate::pipeline::channel::{ChannelConfig, ChannelType};
use crate::types::chunk::ChunkMetadata;
use crate::types::retrieval::{RetrievedChunk, StructuredFilter};

/// Trait for semantic vector search.
/// Returns: Vec of (chunk_id, score, metadata, content, document_id)
#[async_trait]
pub trait SemanticSearch: Send + Sync {
    /// Search the vector store for chunks similar to the query embedding.
    async fn search(
        &self,
        query_embedding: &[f32],
        top_k: usize,
        namespace: Option<&str>,
    ) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError>;
}

/// Trait for text embedding.
#[async_trait]
pub trait Embedder: Send + Sync {
    /// Embed a batch of texts into vectors.
    async fn embed(&self, text: &[String]) -> Result<Vec<Vec<f32>>, RagError>;
    /// Return the dimensionality of the embedding vectors.
    fn dimensions(&self) -> usize;
}

/// Semantic channel executor using vector similarity search.
pub struct SemanticChannelExecutor {
    embedder: Arc<dyn Embedder>,
    store: Arc<dyn SemanticSearch>,
}

impl SemanticChannelExecutor {
    /// Create a new semantic channel executor with the given embedder and store.
    pub fn new(embedder: Arc<dyn Embedder>, store: Arc<dyn SemanticSearch>) -> Self {
        Self { embedder, store }
    }

    /// Execute semantic search: embed the query, search the vector store, and filter results.
    pub async fn execute(
        &self,
        query: &str,
        config: &ChannelConfig,
        _global_filters: &[StructuredFilter],
        namespace: Option<&str>,
    ) -> Result<Vec<RetrievedChunk>, RagError> {
        let embeddings = self
            .embedder
            .embed(&[query.to_string()])
            .await
            .map_err(|e| RagError::Embedding(e.to_string()))?;

        let query_embedding = embeddings
            .into_iter()
            .next()
            .ok_or_else(|| RagError::Embedding("no embedding returned".into()))?;

        let results = self
            .store
            .search(&query_embedding, config.top_k, namespace)
            .await
            .map_err(|e| RagError::Store(e.to_string()))?;

        let hits: Vec<RetrievedChunk> =
            results
                .into_iter()
                .filter(|(_, score, _, _, _)| {
                    if let Some(min_score) = config.min_score { *score >= min_score } else { true }
                })
                .map(|(id, score, metadata, content, document_id)| RetrievedChunk {
                    chunk_id: id.clone(),
                    document_id,
                    content,
                    score,
                    channel: ChannelType::Semantic.as_str().to_string(),
                    channel_score: score,
                    metadata,
                    embedding: None,
                })
                .collect();

        Ok(hits)
    }
}