use crate::core::language_models::BaseChatModel;
use crate::embeddings::Embeddings;
use crate::error::Error;
use crate::retrieval::RetrieverError;
use crate::schema::Message;
use crate::vector_stores::{Document, VectorStore, VectorStoreError};
use std::sync::Arc;
pub struct RAGPipeline {
llm: Arc<dyn BaseChatModel<Error = Error> + Send + Sync>,
embeddings: Arc<dyn Embeddings + Send + Sync>,
vector_store: Arc<dyn VectorStore + Send + Sync>,
retrieve_k: usize,
system_prompt: String,
}
impl std::fmt::Debug for RAGPipeline {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RAGPipeline")
.field("model_name", &self.llm.model_name())
.field("retrieve_k", &self.retrieve_k)
.field("system_prompt", &self.system_prompt)
.finish()
}
}
impl RAGPipeline {
pub async fn index_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
let texts: Vec<&str> = documents.iter().map(|d| d.page_content()).collect();
let embeddings = self
.embeddings
.embed_documents(&texts)
.await
.map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
self.vector_store
.add_documents(documents, embeddings)
.await
.map_err(RetrieverError::StoreError)?;
Ok(())
}
pub async fn query(&self, question: &str) -> Result<String, RetrieverError> {
let result = self.query_with_sources(question).await?;
Ok(result.answer)
}
pub async fn query_with_sources(
&self,
question: &str,
) -> Result<RAGQueryResult, RetrieverError> {
let query_embedding = self
.embeddings
.embed_query(question)
.await
.map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
let search_results = self
.vector_store
.similarity_search(&query_embedding, self.retrieve_k)
.await
.map_err(RetrieverError::StoreError)?;
let sources: Vec<Document> = search_results
.iter()
.map(|r| r.document.clone())
.collect();
let context = if sources.is_empty() {
"No relevant documents found.".to_string()
} else {
sources
.iter()
.enumerate()
.map(|(i, doc)| format!("[{}] {}", i + 1, doc.page_content()))
.collect::<Vec<_>>()
.join("\n\n")
};
let messages = vec![
Message::system(&format!(
"{}\n\nUse the following context to answer the question. If the context doesn't contain the answer, say so.",
self.system_prompt
)),
Message::human(&format!("Context:\n{}\n\nQuestion: {}", context, question)),
];
let llm_result = self
.llm
.chat(messages, None)
.await
.map_err(|e| RetrieverError::EmbeddingError(format!("LLM 调用失败: {}", e)))?;
Ok(RAGQueryResult {
answer: llm_result.content,
sources,
})
}
}
#[derive(Debug, Clone)]
pub struct RAGQueryResult {
pub answer: String,
pub sources: Vec<Document>,
}
pub struct RAGPipelineBuilder {
llm: Option<Arc<dyn BaseChatModel<Error = Error> + Send + Sync>>,
embeddings: Option<Arc<dyn Embeddings + Send + Sync>>,
vector_store: Option<Arc<dyn VectorStore + Send + Sync>>,
retrieve_k: usize,
system_prompt: Option<String>,
}
impl RAGPipelineBuilder {
pub fn new() -> Self {
Self {
llm: None,
embeddings: None,
vector_store: None,
retrieve_k: 4,
system_prompt: None,
}
}
pub fn llm<L>(mut self, llm: L) -> Self
where
L: BaseChatModel + Send + Sync + 'static,
L::Error: Into<Error>,
{
self.llm = Some(crate::core::language_models::wrap_chat_model(llm));
self
}
pub fn llm_from_arc(mut self, llm: Arc<dyn BaseChatModel<Error = Error> + Send + Sync>) -> Self {
self.llm = Some(llm);
self
}
pub fn llm_client(mut self, client: crate::language_models::LLMClient) -> Self {
self.llm = Some(client.into_inner());
self
}
pub fn embeddings<E: Embeddings + Send + Sync + 'static>(mut self, embeddings: E) -> Self {
self.embeddings = Some(Arc::new(embeddings));
self
}
pub fn vector_store<V: VectorStore + Send + Sync + 'static>(mut self, store: V) -> Self {
self.vector_store = Some(Arc::new(store));
self
}
pub fn retrieve_k(mut self, k: usize) -> Self {
self.retrieve_k = k;
self
}
pub fn system(mut self, prompt: impl Into<String>) -> Self {
self.system_prompt = Some(prompt.into());
self
}
pub fn build(self) -> Result<RAGPipeline, RetrieverError> {
let llm = self
.llm
.ok_or_else(|| RetrieverError::EmbeddingError("RAGPipelineBuilder: LLM is required. Call .llm() first.".into()))?;
let embeddings = self
.embeddings
.ok_or_else(|| RetrieverError::EmbeddingError("RAGPipelineBuilder: Embeddings is required. Call .embeddings() first.".into()))?;
let vector_store = self
.vector_store
.ok_or_else(|| RetrieverError::StoreError(VectorStoreError::StorageError("RAGPipelineBuilder: VectorStore is required. Call .vector_store() first.".into())))?;
Ok(RAGPipeline {
llm,
embeddings,
vector_store,
retrieve_k: self.retrieve_k,
system_prompt: self.system_prompt.unwrap_or_else(|| {
"You are a helpful assistant that answers questions based on the provided context.".to_string()
}),
})
}
}
impl Default for RAGPipelineBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embeddings::MockEmbeddings;
use crate::language_models::{OpenAIChat, OpenAIConfig};
use crate::vector_stores::InMemoryVectorStore;
#[test]
fn test_builder_missing_llm() {
let result = RAGPipelineBuilder::new()
.embeddings(MockEmbeddings::new(3))
.vector_store(InMemoryVectorStore::new())
.build();
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("LLM is required"));
}
#[test]
fn test_builder_missing_embeddings() {
let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
let result = RAGPipelineBuilder::new()
.llm(OpenAIChat::new(config))
.vector_store(InMemoryVectorStore::new())
.build();
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("Embeddings is required"));
}
#[test]
fn test_builder_missing_vector_store() {
let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
let result = RAGPipelineBuilder::new()
.llm(OpenAIChat::new(config))
.embeddings(MockEmbeddings::new(3))
.build();
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("VectorStore is required"));
}
#[test]
fn test_builder_success() {
let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
let result = RAGPipelineBuilder::new()
.llm(OpenAIChat::new(config))
.embeddings(MockEmbeddings::new(3))
.vector_store(InMemoryVectorStore::new())
.system("You are a test assistant.")
.retrieve_k(5)
.build();
assert!(result.is_ok());
let pipeline = result.unwrap();
assert_eq!(pipeline.retrieve_k, 5);
assert_eq!(pipeline.system_prompt, "You are a test assistant.");
}
#[test]
fn test_builder_default() {
let builder = RAGPipelineBuilder::default();
assert_eq!(builder.retrieve_k, 4);
assert!(builder.llm.is_none());
}
}