use crate::core::language_models::BaseChatModel;
use crate::retrieval::{RetrieverError, RetrieverTrait};
use crate::schema::Message;
use crate::vector_stores::Document;
#[derive(Debug, Clone, PartialEq)]
pub enum RagDecision {
NoRetrieval,
SingleSearch,
MultiQuery,
}
impl std::fmt::Display for RagDecision {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RagDecision::NoRetrieval => write!(f, "no_retrieval"),
RagDecision::SingleSearch => write!(f, "single_search"),
RagDecision::MultiQuery => write!(f, "multi_query"),
}
}
}
#[derive(Debug, Clone)]
pub struct AdaptiveRAGResult {
pub answer: String,
pub decision: RagDecision,
pub sources: Vec<Document>,
}
#[derive(Debug, thiserror::Error)]
pub enum AdaptiveRAGError {
#[error("LLM error: {0}")]
Llm(String),
#[error("retrieval error: {0}")]
Retrieval(#[from] RetrieverError),
#[error("decision parse error: {0}")]
DecisionParse(String),
}
pub struct AdaptiveRAG<M: BaseChatModel, R: RetrieverTrait> {
llm: M,
retriever: R,
retrieve_k: usize,
multi_query_count: usize,
}
const ROUTING_PROMPT: &str = r#"Given the following query, decide whether retrieval is needed:
- "no_retrieval": The query can be answered from general knowledge
- "single_search": A single search is sufficient
- "multi_query": The query is complex and needs multiple search angles
Query: {query}
Respond with exactly one of: no_retrieval, single_search, multi_query"#;
const GENERATE_SYSTEM_PROMPT: &str = r#"You are a helpful assistant. Answer the user's question based on the provided context when available. If no context is provided, use your general knowledge. Be concise and accurate."#;
const MULTI_QUERY_PROMPT: &str = r#"You are an AI language model assistant. Your task is to generate {count} different versions of the given user question to retrieve relevant documents from a vector database.
By generating multiple perspectives on the user question, your goal is to help overcome some of the limitations of distance-based similarity search.
Provide these alternative questions separated by newlines.
Original question: {question}
Alternative questions:"#;
impl<M: BaseChatModel, R: RetrieverTrait> AdaptiveRAG<M, R> {
pub fn new(llm: M, retriever: R) -> Self {
Self {
llm,
retriever,
retrieve_k: 4,
multi_query_count: 3,
}
}
pub fn with_retrieve_k(mut self, k: usize) -> Self {
self.retrieve_k = k;
self
}
pub fn with_multi_query_count(mut self, count: usize) -> Self {
self.multi_query_count = count;
self
}
pub async fn invoke(&self, query: &str) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
let decision = self.route(query).await?;
match decision {
RagDecision::NoRetrieval => self.generate_no_retrieval(query).await,
RagDecision::SingleSearch => self.generate_single_search(query).await,
RagDecision::MultiQuery => self.generate_multi_query(query).await,
}
}
async fn route(&self, query: &str) -> Result<RagDecision, AdaptiveRAGError> {
let prompt = ROUTING_PROMPT.replace("{query}", query);
let messages = vec![Message::human(&prompt)];
let result = self
.llm
.chat(messages, None)
.await
.map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
parse_decision(&result.content)
}
async fn generate_no_retrieval(
&self,
query: &str,
) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
let messages = vec![Message::human(query)];
let result = self
.llm
.chat_with_system(GENERATE_SYSTEM_PROMPT.to_string(), messages)
.await
.map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
Ok(AdaptiveRAGResult {
answer: result.content,
decision: RagDecision::NoRetrieval,
sources: Vec::new(),
})
}
async fn generate_single_search(
&self,
query: &str,
) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
let docs = self.retriever.retrieve(query, self.retrieve_k).await?;
let context = build_context(&docs);
let user_msg = format!("Context:\n{}\n\nQuestion: {}\n\nAnswer:", context, query);
let messages = vec![Message::human(&user_msg)];
let result = self
.llm
.chat_with_system(GENERATE_SYSTEM_PROMPT.to_string(), messages)
.await
.map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
Ok(AdaptiveRAGResult {
answer: result.content,
decision: RagDecision::SingleSearch,
sources: docs,
})
}
async fn generate_multi_query(
&self,
query: &str,
) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
let alternative_queries = self.generate_queries(query).await?;
let all_queries: Vec<String> = std::iter::once(query.to_string())
.chain(alternative_queries)
.collect();
let docs = self.retrieve_and_merge(&all_queries).await?;
let context = build_context(&docs);
let user_msg = format!("Context:\n{}\n\nQuestion: {}\n\nAnswer:", context, query);
let messages = vec![Message::human(&user_msg)];
let result = self
.llm
.chat_with_system(GENERATE_SYSTEM_PROMPT.to_string(), messages)
.await
.map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
Ok(AdaptiveRAGResult {
answer: result.content,
decision: RagDecision::MultiQuery,
sources: docs,
})
}
async fn generate_queries(&self, query: &str) -> Result<Vec<String>, AdaptiveRAGError> {
let prompt = MULTI_QUERY_PROMPT
.replace("{count}", &self.multi_query_count.to_string())
.replace("{question}", query);
let messages = vec![Message::human(&prompt)];
let result = self
.llm
.chat(messages, None)
.await
.map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
let queries: Vec<String> = result
.content
.lines()
.filter(|line| !line.trim().is_empty())
.map(|line| line.trim().to_string())
.collect();
Ok(queries)
}
async fn retrieve_and_merge(
&self,
queries: &[String],
) -> Result<Vec<Document>, AdaptiveRAGError> {
let futures: Vec<_> = queries
.iter()
.map(|q| self.retriever.retrieve(q, self.retrieve_k))
.collect();
let all_results = futures_util::future::join_all(futures).await;
let mut seen_content: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut merged: Vec<Document> = Vec::new();
for result in all_results {
let docs = result?;
for doc in docs {
let key = doc.content.chars().take(80).collect::<String>();
if seen_content.insert(key) {
merged.push(doc);
}
}
}
Ok(merged)
}
}
fn parse_decision(response: &str) -> Result<RagDecision, AdaptiveRAGError> {
let lower = response.to_lowercase();
if lower.contains("no_retrieval") {
return Ok(RagDecision::NoRetrieval);
}
if lower.contains("multi_query") {
return Ok(RagDecision::MultiQuery);
}
if lower.contains("single_search") {
return Ok(RagDecision::SingleSearch);
}
Err(AdaptiveRAGError::DecisionParse(response.to_string()))
}
fn build_context(docs: &[Document]) -> String {
docs.iter()
.enumerate()
.map(|(i, doc)| format!("[Document {}]: {}", i + 1, doc.content))
.collect::<Vec<_>>()
.join("\n\n")
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult};
use crate::core::runnables::{Runnable, RunnableConfig};
use async_trait::async_trait;
use std::sync::{Arc, Mutex};
#[derive(Clone)]
struct MockLLM {
responses: Arc<Mutex<Vec<String>>>,
}
impl MockLLM {
fn new(responses: Vec<String>) -> Self {
Self {
responses: Arc::new(Mutex::new(responses)),
}
}
}
#[derive(Debug, thiserror::Error)]
#[error("mock llm error")]
struct MockLlmError;
#[async_trait]
impl Runnable<Vec<Message>, LLMResult> for MockLLM {
type Error = MockLlmError;
async fn invoke(
&self,
_input: Vec<Message>,
_config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
let content = {
let mut guard = self.responses.lock().unwrap();
if guard.is_empty() {
"mock response".to_string()
} else {
guard.remove(0)
}
};
Ok(LLMResult {
content,
model: "mock".to_string(),
token_usage: None,
tool_calls: None,
thinking_content: None,
})
}
}
#[async_trait]
impl BaseLanguageModel<Vec<Message>, LLMResult> for MockLLM {
fn model_name(&self) -> &str {
"mock"
}
fn get_num_tokens(&self, text: &str) -> usize {
text.split_whitespace().count()
}
fn with_temperature(self, _temp: f32) -> Self
where
Self: Sized,
{
self
}
fn with_max_tokens(self, _max: usize) -> Self
where
Self: Sized,
{
self
}
}
#[async_trait]
impl BaseChatModel for MockLLM {
async fn chat(
&self,
_messages: Vec<Message>,
_config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
let content = {
let mut guard = self.responses.lock().unwrap();
if guard.is_empty() {
"mock response".to_string()
} else {
guard.remove(0)
}
};
Ok(LLMResult {
content,
model: "mock".to_string(),
token_usage: None,
tool_calls: None,
thinking_content: None,
})
}
async fn stream_chat(
&self,
_messages: Vec<Message>,
_config: Option<RunnableConfig>,
) -> Result<
std::pin::Pin<Box<dyn futures_util::Stream<Item = Result<String, Self::Error>> + Send>>,
Self::Error,
> {
let content = {
let mut guard = self.responses.lock().unwrap();
if guard.is_empty() {
"mock response".to_string()
} else {
guard.remove(0)
}
};
let stream = futures_util::stream::once(async move { Ok(content) });
Ok(Box::pin(stream))
}
}
struct MockRetriever {
documents: Vec<Document>,
}
impl MockRetriever {
fn new(documents: Vec<Document>) -> Self {
Self { documents }
}
}
#[async_trait]
impl RetrieverTrait for MockRetriever {
async fn retrieve(&self, _query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
Ok(self.documents.iter().take(k).cloned().collect())
}
async fn retrieve_with_scores(
&self,
_query: &str,
k: usize,
) -> Result<Vec<crate::vector_stores::SearchResult>, RetrieverError> {
Ok(self
.documents
.iter()
.take(k)
.enumerate()
.map(|(i, doc)| crate::vector_stores::SearchResult {
document: doc.clone(),
score: 1.0 - i as f32 * 0.1,
})
.collect())
}
async fn add_documents(&self, _documents: Vec<Document>) -> Result<(), RetrieverError> {
Ok(())
}
}
#[test]
fn test_parse_decision_no_retrieval() {
let result = parse_decision("no_retrieval").unwrap();
assert_eq!(result, RagDecision::NoRetrieval);
}
#[test]
fn test_parse_decision_single_search() {
let result = parse_decision("single_search").unwrap();
assert_eq!(result, RagDecision::SingleSearch);
}
#[test]
fn test_parse_decision_multi_query() {
let result = parse_decision("multi_query").unwrap();
assert_eq!(result, RagDecision::MultiQuery);
}
#[test]
fn test_parse_decision_case_insensitive() {
assert_eq!(
parse_decision("NO_RETRIEVAL").unwrap(),
RagDecision::NoRetrieval
);
assert_eq!(
parse_decision("Single_Search").unwrap(),
RagDecision::SingleSearch
);
assert_eq!(
parse_decision("MULTI_QUERY").unwrap(),
RagDecision::MultiQuery
);
}
#[test]
fn test_parse_decision_embedded_in_text() {
assert_eq!(
parse_decision("The answer is: no_retrieval").unwrap(),
RagDecision::NoRetrieval
);
assert_eq!(
parse_decision("I think multi_query is best").unwrap(),
RagDecision::MultiQuery
);
}
#[test]
fn test_parse_decision_invalid() {
let result = parse_decision("something else entirely");
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
AdaptiveRAGError::DecisionParse(_)
));
}
#[test]
fn test_rag_decision_display() {
assert_eq!(format!("{}", RagDecision::NoRetrieval), "no_retrieval");
assert_eq!(format!("{}", RagDecision::SingleSearch), "single_search");
assert_eq!(format!("{}", RagDecision::MultiQuery), "multi_query");
}
#[test]
fn test_build_context() {
let docs = vec![
Document::new("First document content"),
Document::new("Second document content"),
];
let context = build_context(&docs);
assert!(context.contains("[Document 1]: First document content"));
assert!(context.contains("[Document 2]: Second document content"));
}
#[test]
fn test_build_context_empty() {
let context = build_context(&[]);
assert!(context.is_empty());
}
#[tokio::test]
async fn test_adaptive_rag_no_retrieval() {
let llm = MockLLM::new(vec![
"no_retrieval".to_string(),
"Paris is the capital of France.".to_string(),
]);
let retriever = MockRetriever::new(vec![Document::new("irrelevant doc")]);
let rag = AdaptiveRAG::new(llm, retriever);
let result = rag.invoke("What is the capital of France?").await.unwrap();
assert_eq!(result.decision, RagDecision::NoRetrieval);
assert!(result.answer.contains("Paris"));
assert!(result.sources.is_empty());
}
#[tokio::test]
async fn test_adaptive_rag_single_search() {
let llm = MockLLM::new(vec![
"single_search".to_string(),
"Rust is a systems programming language.".to_string(),
]);
let retriever = MockRetriever::new(vec![Document::new(
"Rust emphasizes safety and performance.",
)]);
let rag = AdaptiveRAG::new(llm, retriever);
let result = rag.invoke("Tell me about Rust").await.unwrap();
assert_eq!(result.decision, RagDecision::SingleSearch);
assert!(result.answer.contains("Rust"));
assert_eq!(result.sources.len(), 1);
}
#[tokio::test]
async fn test_adaptive_rag_multi_query() {
let llm = MockLLM::new(vec![
"multi_query".to_string(),
"How does Rust memory management work?\nWhat is ownership in Rust?".to_string(),
"Rust uses ownership and borrowing for memory management.".to_string(),
]);
let retriever = MockRetriever::new(vec![
Document::new("Rust ownership model"),
Document::new("Borrowing and lifetimes"),
]);
let rag = AdaptiveRAG::new(llm, retriever);
let result = rag.invoke("Explain Rust memory management").await.unwrap();
assert_eq!(result.decision, RagDecision::MultiQuery);
assert!(result.answer.contains("ownership"));
assert!(!result.sources.is_empty());
}
#[tokio::test]
async fn test_adaptive_rag_llm_error() {
let llm = MockLLM::new(vec!["something_unrelated".to_string()]);
let retriever = MockRetriever::new(vec![]);
let rag = AdaptiveRAG::new(llm, retriever);
let result = rag.invoke("test query").await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
AdaptiveRAGError::DecisionParse(_)
));
}
#[tokio::test]
async fn test_adaptive_rag_with_retrieve_k() {
let llm = MockLLM::new(vec![
"single_search".to_string(),
"Answer based on context.".to_string(),
]);
let retriever = MockRetriever::new(vec![
Document::new("Doc 1"),
Document::new("Doc 2"),
Document::new("Doc 3"),
]);
let rag = AdaptiveRAG::new(llm, retriever).with_retrieve_k(2);
let result = rag.invoke("test query").await.unwrap();
assert_eq!(result.decision, RagDecision::SingleSearch);
assert_eq!(result.sources.len(), 2); }
#[test]
fn test_adaptive_rag_error_display() {
let err = AdaptiveRAGError::Llm("timeout".to_string());
assert!(err.to_string().contains("LLM error"));
assert!(err.to_string().contains("timeout"));
let err = AdaptiveRAGError::DecisionParse("bad output".to_string());
assert!(err.to_string().contains("decision parse error"));
assert!(err.to_string().contains("bad output"));
}
}