use std::sync::Arc;
use async_trait::async_trait;
use xz_rag::{
ChannelConfig, ChannelPipeline, ChunkMetadata, ContextBuilder, DefaultRagEngineBuilder,
Embedder, KnowledgeGraphSearch, MetadataStore, PromptTemplate, RagEngine, RagError, RagRequest,
RetrieveRequest, SemanticSearch, StructuredFilter,
};
#[tokio::test]
async fn test_engine_builder() {
let engine = DefaultRagEngineBuilder::default()
.name("test-engine")
.version("0.1.0")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.build();
let info = engine.engine_info();
assert_eq!(info.name, "test-engine");
assert!(info.supported_channels.contains(&"semantic".to_string()));
}
#[tokio::test]
async fn test_retrieve_empty_returns_error() {
let engine = DefaultRagEngineBuilder::default().pipeline(ChannelPipeline::new(vec![])).build();
let request = RetrieveRequest::builder("test query").build();
let result = engine.retrieve(&request).await.unwrap();
assert!(result.hits.is_empty());
assert!(result.channel_report.is_empty());
}
#[tokio::test]
async fn test_retrieve_and_generate_no_hits() {
let engine = DefaultRagEngineBuilder::default().pipeline(ChannelPipeline::new(vec![])).build();
let request = RagRequest::builder("test query").build();
let result = engine.retrieve_and_generate(&request).await;
assert!(result.is_err());
}
#[test]
fn test_context_builder() {
let builder = ContextBuilder::new(4096);
let budget = builder.context_budget();
assert!(budget > 0);
}
#[test]
fn test_prompt_template_render() {
let template = PromptTemplate::default_qa();
let rendered = template.render("What is Rust?", "Rust is a systems programming language.");
assert!(rendered.contains("What is Rust?"));
assert!(rendered.contains("Rust is a systems programming language"));
}
struct MockEmbedder;
#[async_trait]
impl Embedder for MockEmbedder {
async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, RagError> {
Ok(texts.iter().map(|_| vec![0.1, 0.2, 0.3]).collect())
}
fn dimensions(&self) -> usize {
3
}
}
struct MockSemanticStore;
#[async_trait]
impl SemanticSearch for MockSemanticStore {
async fn search(
&self,
_query_embedding: &[f32],
_top_k: usize,
_namespace: Option<&str>,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
Ok(vec![(
"chunk-sem-1".into(),
0.95,
ChunkMetadata::default(),
"Rust is a systems programming language focused on safety and performance.".into(),
"doc-1".into(),
)])
}
}
struct MockMetadataStore;
#[async_trait]
impl MetadataStore for MockMetadataStore {
async fn search_by_metadata(
&self,
_query: &str,
_filters: &[StructuredFilter],
_top_k: usize,
_namespace: Option<&str>,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
Ok(vec![(
"chunk-meta-1".into(),
0.87,
ChunkMetadata::default(),
"Tokio is an async runtime for Rust providing a multi-threaded work-stealing scheduler."
.into(),
"doc-2".into(),
)])
}
}
struct MockGraphStore;
#[async_trait]
impl KnowledgeGraphSearch for MockGraphStore {
async fn search(
&self,
_query: &str,
_top_k: usize,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
Ok(vec![(
"graph-node-1".into(),
0.82,
ChunkMetadata::default(),
"Entity: Rust — a systems programming language; relates-to: performance, safety."
.into(),
"kg-doc-1".into(),
)])
}
}
struct MultiResultSemanticStore;
#[async_trait]
impl SemanticSearch for MultiResultSemanticStore {
async fn search(
&self,
_query_embedding: &[f32],
top_k: usize,
_namespace: Option<&str>,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
let hits = vec![
("chunk-a".into(), 0.95, ChunkMetadata::default(), "Content A".into(), "doc-a".into()),
("chunk-b".into(), 0.85, ChunkMetadata::default(), "Content B".into(), "doc-b".into()),
("chunk-c".into(), 0.75, ChunkMetadata::default(), "Content C".into(), "doc-c".into()),
];
Ok(hits.into_iter().take(top_k).collect())
}
}
#[tokio::test]
async fn test_retrieve_with_semantic_store() {
let engine = DefaultRagEngineBuilder::default()
.name("test-e2e")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.embedder(Arc::new(MockEmbedder))
.semantic_store(Arc::new(MockSemanticStore))
.build();
let retrieve_req = RetrieveRequest::builder("What is Rust?")
.channels(vec![ChannelConfig::semantic(0.5, 10)])
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert_eq!(result.hits.len(), 1);
assert!(result.hits[0].content.contains("Rust is a systems"));
assert!(!result.channel_report.is_empty());
assert_eq!(result.effective_query, "What is Rust?");
}
#[tokio::test]
async fn test_multi_channel_retrieval_fusion() {
let engine = DefaultRagEngineBuilder::default()
.name("test-multi")
.pipeline(
ChannelPipeline::new(vec![
ChannelConfig::semantic(0.5, 10),
ChannelConfig::metadata(0.3, 5),
])
.with_rrf_k(60)
.with_normalize(true),
)
.embedder(Arc::new(MockEmbedder))
.semantic_store(Arc::new(MultiResultSemanticStore))
.metadata_store(Arc::new(MockMetadataStore))
.build();
let retrieve_req = RetrieveRequest::builder("What is async Rust?")
.channels(vec![ChannelConfig::semantic(0.5, 10), ChannelConfig::metadata(0.3, 5)])
.top_k(5)
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert!(result.channel_report.len() >= 2);
assert!(result.hits.len() <= 5);
let all_content: String =
result.hits.iter().map(|h| &h.content[..]).collect::<Vec<_>>().concat();
assert!(
all_content.contains("Content") && all_content.contains("Tokio"),
"Expected content from both channels in fused results"
);
assert_eq!(result.effective_query, "What is async Rust?");
}
#[tokio::test]
async fn test_retrieve_with_graph_channel() {
let engine = DefaultRagEngineBuilder::default()
.name("test-graph")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::new(
xz_rag::ChannelType::Graph,
0.4,
10,
)]))
.graph_store(Arc::new(MockGraphStore))
.build();
let retrieve_req = RetrieveRequest::builder("rust language")
.channels(vec![ChannelConfig::new(xz_rag::ChannelType::Graph, 0.4, 10)])
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert_eq!(result.hits.len(), 1);
assert_eq!(result.hits[0].chunk_id, "graph-node-1");
assert!(result.hits[0].content.contains("Entity"));
}
#[tokio::test]
async fn test_retrieve_and_generate_empty_hits_error() {
let engine = DefaultRagEngineBuilder::default()
.name("test-empty")
.pipeline(ChannelPipeline::new(vec![]))
.build();
let rag_req = RagRequest::builder("nobody has this").build();
let result = engine.retrieve_and_generate(&rag_req).await;
assert!(result.is_err());
match result {
Err(RagError::NoResults(msg)) => {
assert!(msg.contains("nobody has this"), "Expected query in error, got: {msg}");
}
other => panic!("Expected NoResults error, got: {other:?}"),
}
}
#[cfg(feature = "llm-generation")]
mod streaming_tests {
use std::pin::Pin;
use std::sync::Arc;
use async_trait::async_trait;
use futures::stream::{self, Stream, StreamExt};
use xz_provider::{
CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelInfo, ProviderError,
RequestOptions as ProviderRequestOptions, StreamEvent, TokenUsage,
};
use xz_rag::{
ChannelConfig, ChannelPipeline, ChunkMetadata, DefaultRagEngineBuilder, Embedder,
RagEngine, RagError, RagRequest, RagStreamEvent, RetrieveRequest, SemanticSearch,
};
struct StreamEmbedder;
#[async_trait]
impl Embedder for StreamEmbedder {
async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, RagError> {
Ok(texts.iter().map(|_| vec![0.1, 0.2, 0.3]).collect())
}
fn dimensions(&self) -> usize {
3
}
}
struct StreamSemanticStore;
#[async_trait]
impl SemanticSearch for StreamSemanticStore {
async fn search(
&self,
_query_embedding: &[f32],
_top_k: usize,
_namespace: Option<&str>,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
Ok(vec![(
"stream-chunk".into(),
0.92,
ChunkMetadata::default(),
"Streaming is an efficient way to handle large data flows.".into(),
"doc-stream".into(),
)])
}
}
#[derive(Debug)]
struct StreamingMockProvider {
deltas: Vec<String>,
}
impl StreamingMockProvider {
fn new(deltas: Vec<String>) -> Self {
Self { deltas }
}
}
#[async_trait]
impl LlmProvider for StreamingMockProvider {
async fn complete(
&self,
_request: CompletionRequest,
_options: ProviderRequestOptions,
) -> Result<CompletionResponse, ProviderError> {
let all = self.deltas.concat();
Ok(CompletionResponse {
content: Some(all),
thinking: None,
tool_calls: vec![],
usage: TokenUsage::new(10, 20),
model: "mock-model".into(),
finish_reason: FinishReason::Stop,
latency_ms: 0,
cache_info: None,
})
}
async fn complete_stream(
&self,
_request: CompletionRequest,
_options: ProviderRequestOptions,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>,
ProviderError,
> {
let deltas = self.deltas.clone();
let stream = stream::iter(
deltas
.into_iter()
.map(|d| Ok(StreamEvent::ContentDelta { delta: d }))
.chain(std::iter::once(Ok(StreamEvent::Done {
finish_reason: FinishReason::Stop,
usage: Some(TokenUsage::new(10, 20)),
})))
.collect::<Vec<_>>(),
);
Ok(Box::pin(stream))
}
fn models(&self) -> &[ModelInfo] {
&[]
}
fn name(&self) -> &str {
"streaming-mock"
}
}
#[tokio::test]
async fn test_retrieve_and_generate_stream_success() {
let engine = DefaultRagEngineBuilder::default()
.name("test-stream")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.embedder(Arc::new(StreamEmbedder))
.semantic_store(Arc::new(StreamSemanticStore))
.provider(Arc::new(StreamingMockProvider::new(vec![
"Hello ".into(),
"streaming ".into(),
"world!".into(),
])))
.build();
let retrieve_req = RetrieveRequest::builder("stream test")
.channels(vec![ChannelConfig::semantic(0.5, 10)])
.build();
let rag_req = RagRequest::builder("stream test").retrieve_config(retrieve_req).build();
let mut stream = engine.retrieve_and_generate_stream(&rag_req).await.unwrap();
let mut events: Vec<RagStreamEvent> = vec![];
while let Some(event) = stream.next().await {
events.push(event.unwrap());
}
assert!(events.len() >= 3, "Expected at least 3 events, got {}", events.len());
let started = events.iter().any(|e| matches!(e, RagStreamEvent::GenerationStarted { .. }));
assert!(started, "Missing GenerationStarted event");
let deltas: Vec<&str> = events
.iter()
.filter_map(|e| {
if let RagStreamEvent::ContentDelta { delta } = e {
Some(delta.as_str())
} else {
None
}
})
.collect();
let full: String = deltas.concat();
assert_eq!(full, "Hello streaming world!");
let done = events.iter().any(|e| matches!(e, RagStreamEvent::Done { .. }));
assert!(done, "Missing Done event");
}
#[tokio::test]
async fn test_retrieve_and_generate_stream_empty_hits() {
let engine = DefaultRagEngineBuilder::default()
.name("test-stream-empty")
.pipeline(ChannelPipeline::new(vec![]))
.provider(Arc::new(StreamingMockProvider::new(vec![])))
.build();
let rag_req = RagRequest::builder("no results").build();
let result = engine.retrieve_and_generate_stream(&rag_req).await;
assert!(result.is_err());
match result {
Err(RagError::NoResults(msg)) => {
assert!(msg.contains("no results"));
}
_other => panic!("Expected NoResults, got unexpected result"),
}
}
}
#[tokio::test]
async fn test_retrieve_effective_query_tracking() {
let engine = DefaultRagEngineBuilder::default()
.name("test-eff-query")
.pipeline(ChannelPipeline::new(vec![]))
.build();
let retrieve_req = RetrieveRequest::builder("original search phrase").build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert_eq!(result.effective_query, "original search phrase");
assert!(result.hits.is_empty());
}
#[tokio::test]
async fn test_retrieve_with_chat_history_setup() {
let engine = DefaultRagEngineBuilder::default()
.name("test-history")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.embedder(Arc::new(MockEmbedder))
.semantic_store(Arc::new(MockSemanticStore))
.build();
let retrieve_req = RetrieveRequest::builder("What about memory safety?")
.channels(vec![ChannelConfig::semantic(0.5, 10)])
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert_eq!(result.hits.len(), 1);
assert!(result.hits[0].content.contains("Rust is a systems"));
assert_eq!(result.effective_query, "What about memory safety?");
}
#[tokio::test]
async fn test_retrieve_with_hyde_preprocessing() {
use xz_rag::QueryPreprocessing;
let engine = DefaultRagEngineBuilder::default()
.name("test-hyde")
.pipeline(ChannelPipeline::new(vec![]))
.build();
let retrieve_req = RetrieveRequest {
query: "explain Rust ownership".into(),
channels: vec![],
global_filters: vec![],
top_k: 10,
namespace: None,
include_embeddings: false,
query_preprocessing: Some(QueryPreprocessing::Hyde),
};
let result = engine.retrieve(&retrieve_req).await;
match result {
Ok(retrieve_result) => {
assert_eq!(retrieve_result.effective_query, "explain Rust ownership");
assert!(retrieve_result.hits.is_empty());
}
Err(RagError::QueryPreprocessing(_)) => {
}
Err(other) => panic!("Unexpected error: {other:?}"),
}
}
#[tokio::test]
async fn test_engine_builder_full_config() {
let engine = DefaultRagEngineBuilder::default()
.name("full-engine")
.version("1.0.0-test")
.pipeline(
ChannelPipeline::new(vec![
ChannelConfig::semantic(0.5, 10).with_min_score(0.1),
ChannelConfig::metadata(0.3, 5),
ChannelConfig::new(xz_rag::ChannelType::Bm25, 0.2, 3),
])
.with_rrf_k(60),
)
.embedder(Arc::new(MockEmbedder))
.semantic_store(Arc::new(MockSemanticStore))
.metadata_store(Arc::new(MockMetadataStore))
.context_builder(ContextBuilder::new(8192))
.prompt_template(PromptTemplate::default_qa())
.build();
let info = engine.engine_info();
assert_eq!(info.name, "full-engine");
assert_eq!(info.version, "1.0.0-test");
assert!(info.supported_channels.contains(&"semantic".to_string()));
assert!(info.supported_channels.contains(&"metadata".to_string()));
assert_eq!(info.max_context_window, 7296);
let retrieve_req = RetrieveRequest::builder("test all channels")
.channels(vec![ChannelConfig::semantic(0.5, 10), ChannelConfig::metadata(0.3, 5)])
.top_k(5)
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert!(!result.hits.is_empty());
assert!(result.channel_report.len() >= 2);
}
#[tokio::test]
async fn test_retrieve_with_system_prompt_setup() {
let engine = DefaultRagEngineBuilder::default()
.name("test-custom-system")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.embedder(Arc::new(MockEmbedder))
.semantic_store(Arc::new(MockSemanticStore))
.build();
let retrieve_req = RetrieveRequest::builder("Rust traits")
.channels(vec![ChannelConfig::semantic(0.5, 10)])
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert_eq!(result.hits.len(), 1);
assert!(result.hits[0].content.contains("Rust is a systems"));
}
#[tokio::test]
async fn test_retrieve_result_includes_stats() {
let engine = DefaultRagEngineBuilder::default()
.name("test-stats")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.embedder(Arc::new(MockEmbedder))
.semantic_store(Arc::new(MockSemanticStore))
.build();
let retrieve_req = RetrieveRequest::builder("stats test")
.channels(vec![ChannelConfig::semantic(0.5, 10)])
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert!(!result.channel_report.is_empty());
assert_eq!(result.hits.len(), 1);
assert_eq!(result.effective_query, "stats test");
let _ = result.latency_ms; }
#[cfg(feature = "rerank")]
mod rerank_tests {
use std::sync::Arc;
use xz_rag::{
ChannelConfig, ChannelPipeline, ChunkMetadata, DefaultRagEngineBuilder, Embedder,
RagEngine, RagError, RetrieveRequest, SemanticSearch,
};
use xz_rerank::MockReranker;
struct RerankEmbedder;
#[async_trait::async_trait]
impl Embedder for RerankEmbedder {
async fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, RagError> {
Ok(texts.iter().map(|_| vec![0.1, 0.2, 0.3]).collect())
}
fn dimensions(&self) -> usize {
3
}
}
struct RerankSemanticStore;
#[async_trait::async_trait]
impl SemanticSearch for RerankSemanticStore {
async fn search(
&self,
_query_embedding: &[f32],
top_k: usize,
_namespace: Option<&str>,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
let all = vec![
(
"r1".into(),
0.60,
ChunkMetadata::default(),
"First candidate content".into(),
"doc-r1".into(),
),
(
"r2".into(),
0.85,
ChunkMetadata::default(),
"Second candidate content".into(),
"doc-r2".into(),
),
(
"r3".into(),
0.50,
ChunkMetadata::default(),
"Third candidate content".into(),
"doc-r3".into(),
),
];
Ok(all.into_iter().take(top_k).collect())
}
}
#[tokio::test]
async fn test_rerank_integration_with_mock_reranker() {
let engine = DefaultRagEngineBuilder::default()
.name("test-rerank")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.embedder(Arc::new(RerankEmbedder))
.semantic_store(Arc::new(RerankSemanticStore))
.reranker(Arc::new(MockReranker::new("mock-reranker")))
.build();
let retrieve_req = RetrieveRequest::builder("important query")
.channels(vec![ChannelConfig::semantic(0.5, 10)])
.top_k(2)
.build();
let result = engine.retrieve(&retrieve_req).await.unwrap();
assert_eq!(result.hits.len(), 2);
let hit_ids: Vec<&str> = result.hits.iter().map(|h| h.chunk_id.as_str()).collect();
assert!(hit_ids.contains(&"r2"), "Expected r2 in top results, got: {hit_ids:?}");
let info = engine.engine_info();
assert!(info.reranking_enabled);
}
#[tokio::test]
async fn test_rerank_engine_info_reflects_reranker() {
let engine = DefaultRagEngineBuilder::default()
.name("test-rerank-info")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::semantic(0.5, 10)]))
.reranker(Arc::new(MockReranker::new("mock-reranker")))
.build();
let info = engine.engine_info();
assert!(info.reranking_enabled);
assert_eq!(info.name, "test-rerank-info");
}
}
#[cfg(feature = "llm-generation")]
mod llm_tests {
use std::pin::Pin;
use std::sync::Arc;
use async_trait::async_trait;
use futures::Stream;
use xz_provider::{
CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelInfo, ProviderError,
RequestOptions, StreamEvent, TokenUsage,
};
use xz_rag::{
ChannelConfig, ChannelPipeline, ChunkMetadata, DefaultRagEngineBuilder, MetadataStore,
RagEngine, RagError, RagRequest, RetrieveRequest, StructuredFilter,
};
#[derive(Debug)]
struct FailingProvider;
#[async_trait]
impl LlmProvider for FailingProvider {
async fn complete(
&self,
_request: CompletionRequest,
_options: RequestOptions,
) -> Result<CompletionResponse, ProviderError> {
Err(ProviderError::Config("mock LLM failure".into()))
}
async fn complete_stream(
&self,
_request: CompletionRequest,
_options: RequestOptions,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>,
ProviderError,
> {
Err(ProviderError::Config("mock LLM stream failure".into()))
}
fn models(&self) -> &[ModelInfo] {
&[]
}
fn name(&self) -> &str {
"failing-mock"
}
}
#[derive(Debug)]
struct SuccessMockProvider;
#[async_trait]
impl LlmProvider for SuccessMockProvider {
async fn complete(
&self,
_request: CompletionRequest,
_options: RequestOptions,
) -> Result<CompletionResponse, ProviderError> {
Ok(CompletionResponse {
content: Some("Rust is a modern systems programming language.".into()),
thinking: None,
tool_calls: vec![],
usage: TokenUsage::new(10, 20),
model: "mock-model".into(),
finish_reason: FinishReason::Stop,
latency_ms: 0,
cache_info: None,
})
}
async fn complete_stream(
&self,
_request: CompletionRequest,
_options: RequestOptions,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>,
ProviderError,
> {
Err(ProviderError::Config("streaming not supported".into()))
}
fn models(&self) -> &[ModelInfo] {
&[]
}
fn name(&self) -> &str {
"success-mock"
}
}
#[derive(Debug)]
struct MockMetadataStore;
#[async_trait]
impl MetadataStore for MockMetadataStore {
async fn search_by_metadata(
&self,
_query: &str,
_filters: &[StructuredFilter],
_top_k: usize,
_namespace: Option<&str>,
) -> Result<Vec<(String, f32, ChunkMetadata, String, String)>, RagError> {
Ok(vec![(
"chunk-1".into(),
0.9,
ChunkMetadata::default(),
"Rust is a systems programming language for performance and safety.".into(),
"doc-1".into(),
)])
}
}
#[tokio::test]
async fn test_llm_failure_propagates_error() {
let engine = DefaultRagEngineBuilder::default()
.name("test-engine")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::metadata(0.5, 10)]))
.metadata_store(Arc::new(MockMetadataStore))
.provider(Arc::new(FailingProvider))
.build();
let retrieve_req = RetrieveRequest::builder("What is Rust?")
.channels(vec![ChannelConfig::metadata(0.5, 10)])
.build();
let rag_request =
RagRequest::builder("What is Rust?").retrieve_config(retrieve_req).build();
let result = engine.retrieve_and_generate(&rag_request).await;
assert!(result.is_err(), "Expected error from failing LLM provider, but got success");
match result {
Err(RagError::Provider(msg)) => {
assert!(
msg.contains("LLM generation failed"),
"Expected Provider error about LLM generation failure, got: {}",
msg
);
}
other => panic!("Expected Err(RagError::Provider(...)), got: {:?}", other),
}
}
#[tokio::test]
async fn test_no_provider_configured_returns_error() {
let engine = DefaultRagEngineBuilder::default()
.name("test-engine")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::metadata(0.5, 10)]))
.metadata_store(Arc::new(MockMetadataStore))
.build();
let retrieve_req = RetrieveRequest::builder("What is Rust?")
.channels(vec![ChannelConfig::metadata(0.5, 10)])
.build();
let rag_request =
RagRequest::builder("What is Rust?").retrieve_config(retrieve_req).build();
let result = engine.retrieve_and_generate(&rag_request).await;
assert!(result.is_err(), "Expected error when no provider configured, but got success");
match result {
Err(RagError::Provider(msg)) => {
assert!(
msg.contains("no LLM provider configured"),
"Expected 'no LLM provider configured', got: {}",
msg
);
}
other => panic!("Expected Err(RagError::Provider(...)), got: {:?}", other),
}
}
#[tokio::test]
async fn test_retrieve_and_generate_with_mock_provider() {
let engine = DefaultRagEngineBuilder::default()
.name("test-e2e-gen")
.pipeline(ChannelPipeline::new(vec![ChannelConfig::metadata(0.5, 10)]))
.metadata_store(Arc::new(MockMetadataStore))
.provider(Arc::new(SuccessMockProvider))
.build();
let retrieve_req = RetrieveRequest::builder("What is Rust?")
.channels(vec![ChannelConfig::metadata(0.5, 10)])
.build();
let rag_req = RagRequest::builder("What is Rust?").retrieve_config(retrieve_req).build();
let response = engine.retrieve_and_generate(&rag_req).await.unwrap();
assert!(response.answer.contains("Rust is"));
assert_eq!(response.retrieve_stats.hits.len(), 1);
}
}