1use std::collections::HashMap;
2use std::pin::Pin;
3use std::sync::Arc;
4use std::time::Instant;
5
6use async_trait::async_trait;
7use futures::stream::Stream;
8use tracing::info;
9
10use crate::channels::graph::{GraphChannelExecutor, KnowledgeGraphSearch};
11use crate::channels::metadata::{MetadataChannelExecutor, MetadataStore};
12use crate::channels::semantic::{Embedder, SemanticChannelExecutor, SemanticSearch};
13use crate::context::token_budget::ContextBuilder;
14use crate::error::RagError;
15use crate::pipeline::channel::{ChannelConfig, ChannelPipeline, ChannelType};
16use crate::pipeline::fusion::RrfFusion;
17use crate::pipeline::normalize::MinMaxNormalizer;
18use crate::traits::RagEngine;
19use crate::types::config::RagEngineInfo;
20use crate::types::rag::{
21 BuiltContext, ChatMessage, ChatRole, PromptTemplate, RagRequest, RagResponse, RagStreamEvent,
22 RagTokenUsage,
23};
24use crate::types::retrieval::{
25 ChannelStats, QueryPreprocessing, RetrieveRequest, RetrieveResult, RetrievedChunk,
26};
27
28#[cfg(feature = "rerank")]
29use xz_rerank::{RerankCandidate, RerankConfig, traits::Reranker};
30
31#[cfg(feature = "llm-generation")]
32use xz_provider::{
33 CompletionRequest as ProviderCompletionRequest, LlmProvider,
34 RequestOptions as ProviderRequestOptions, StreamEvent,
35 types::message::Message as ProviderMessage,
36};
37
38pub struct DefaultRagEngine {
40 info: RagEngineInfo,
41 pipeline: ChannelPipeline,
42 context_builder: ContextBuilder,
43 prompt_template: PromptTemplate,
44 embedder: Option<Arc<dyn Embedder>>,
46 semantic_store: Option<Arc<dyn SemanticSearch>>,
47 metadata_store: Option<Arc<dyn MetadataStore>>,
48 graph_store: Option<Arc<dyn KnowledgeGraphSearch>>,
49 #[cfg(feature = "rerank")]
50 reranker: Option<Arc<dyn Reranker>>,
51 #[cfg(feature = "llm-generation")]
52 provider: Option<Arc<dyn LlmProvider>>,
53 #[cfg(feature = "caching")]
54 cache: Option<crate::cache::memory_cache::RagMemoryCache>,
55}
56
57impl DefaultRagEngine {
58 pub fn builder() -> DefaultRagEngineBuilder {
60 DefaultRagEngineBuilder::default()
61 }
62}
63
64impl DefaultRagEngine {
65 async fn do_retrieve(&self, request: &RetrieveRequest) -> Result<RetrieveResult, RagError> {
67 #[cfg(feature = "caching")]
69 if let Some(ref cache) = self.cache {
70 let cache_key = build_cache_key(&request.query, &request.namespace);
71 if let Some(cached) = cache.get(&cache_key).await {
72 info!(query = %request.query, "RAG cache hit");
73 return Ok(cached);
74 }
75 }
76
77 let effective_query = self.preprocess_query(request).await?;
79
80 let start = Instant::now();
81 let mut channel_results: HashMap<String, Vec<RetrievedChunk>> = HashMap::new();
82 let mut channel_report: HashMap<String, ChannelStats> = HashMap::new();
83
84 for (channel_idx, channel_config) in self.pipeline.channels.iter().enumerate() {
85 let channel_start = Instant::now();
86 let namespace = request.namespace.as_deref();
87
88 let hits = match &channel_config.channel_type {
89 ChannelType::Semantic => {
90 if let (Some(emb), Some(store)) = (&self.embedder, &self.semantic_store) {
91 let executor = SemanticChannelExecutor::new(emb.clone(), store.clone());
92 executor
93 .execute(
94 &effective_query,
95 channel_config,
96 &request.global_filters,
97 namespace,
98 )
99 .await?
100 } else {
101 vec![]
102 }
103 }
104 ChannelType::Metadata => {
105 if let Some(store) = &self.metadata_store {
106 let executor = MetadataChannelExecutor::new(store.clone());
107 executor
108 .execute(
109 &effective_query,
110 channel_config,
111 &request.global_filters,
112 namespace,
113 )
114 .await?
115 } else {
116 vec![]
117 }
118 }
119 ChannelType::Bm25 => {
120 #[cfg(feature = "bm25")]
121 {
122 let executor = crate::channels::bm25::Bm25ChannelExecutor::new();
123 executor
124 .execute(
125 &effective_query,
126 channel_config,
127 &request.global_filters,
128 namespace,
129 )
130 .await?
131 }
132 #[cfg(not(feature = "bm25"))]
133 vec![]
134 }
135 ChannelType::Graph => {
136 if let Some(store) = &self.graph_store {
137 let executor = GraphChannelExecutor::new(store.clone());
138 executor
139 .execute(
140 &effective_query,
141 channel_config,
142 &request.global_filters,
143 namespace,
144 )
145 .await?
146 } else {
147 vec![]
148 }
149 }
150 _ => vec![],
151 };
152
153 let latency = channel_start.elapsed().as_millis() as u64;
154 let min_score = hits.iter().map(|h| h.score).fold(f32::INFINITY, f32::min);
155 let max_score = hits.iter().map(|h| h.score).fold(0.0_f32, f32::max);
156
157 channel_report.insert(
158 channel_config.channel_type.as_str().to_string(),
159 ChannelStats {
160 channel_type: format!(
161 "{}#{}",
162 channel_config.channel_type.as_str(),
163 channel_idx
164 ),
165 hits: hits.len(),
166 latency_ms: latency,
167 min_score,
168 max_score,
169 },
170 );
171
172 channel_results
173 .insert(format!("{}#{}", channel_config.channel_type.as_str(), channel_idx), hits);
174 }
175
176 if self.pipeline.normalize_scores {
178 for hits in channel_results.values_mut() {
179 let mut scores: Vec<f32> = hits.iter().map(|h| h.score).collect();
180 MinMaxNormalizer::normalize(&mut scores);
181 for (hit, score) in hits.iter_mut().zip(scores) {
182 hit.score = score;
183 }
184 }
185 }
186
187 let fusion = RrfFusion::new(self.pipeline.rrf_k);
189 let mut fused = fusion.fuse(channel_results);
190
191 #[cfg(feature = "rerank")]
192 if let Some(reranker) = &self.reranker
193 && !fused.is_empty()
194 {
195 let original_hits: HashMap<String, RetrievedChunk> =
196 fused.iter().cloned().map(|chunk| (chunk.chunk_id.clone(), chunk)).collect();
197
198 let candidates: Vec<RerankCandidate> =
199 fused.iter().map(retrieved_chunk_to_rerank_candidate).collect();
200
201 let rerank_result = reranker
202 .rerank(
203 &effective_query,
204 candidates,
205 &RerankConfig {
206 top_k: request.top_k,
207 min_score: None,
208 include_score_breakdown: false,
209 recency_mode: None,
210 query_embedding: None,
211 },
212 )
213 .await
214 .map_err(|e| RagError::Rerank(e.to_string()))?;
215
216 fused = rerank_result
217 .hits
218 .into_iter()
219 .filter_map(|hit| {
220 original_hits.get(&hit.candidate_id).cloned().map(|mut chunk| {
221 chunk.score = hit.score;
222 chunk
223 })
224 })
225 .collect();
226 }
227
228 fused.truncate(request.top_k);
230
231 let latency_ms = start.elapsed().as_millis() as u64;
232
233 let result = RetrieveResult { hits: fused, channel_report, latency_ms, effective_query };
234
235 #[cfg(feature = "caching")]
237 if let Some(ref cache) = self.cache {
238 let cache_key = build_cache_key(&request.query, &request.namespace);
239 cache.set(&cache_key, result.clone()).await;
240 }
241
242 Ok(result)
243 }
244
245 fn build_context(&self, chunks: &[RetrievedChunk], query: &str) -> BuiltContext {
247 self.context_builder.build(chunks, query)
248 }
249
250 async fn preprocess_query(&self, request: &RetrieveRequest) -> Result<String, RagError> {
252 match &request.query_preprocessing {
253 Some(QueryPreprocessing::Hyde) => {
254 #[cfg(feature = "hyde")]
255 {
256 let provider = self.provider.as_ref().ok_or_else(|| {
257 RagError::QueryPreprocessing("No LLM provider configured for HYDE".into())
258 })?;
259 let expander = crate::preprocessing::hyde::HydeExpander;
260 expander.expand(&request.query, provider.as_ref()).await
261 }
262 #[cfg(not(feature = "hyde"))]
263 Ok(request.query.clone())
264 }
265 Some(QueryPreprocessing::QueryExpansion { count }) => {
266 #[cfg(feature = "hyde")]
267 {
268 let provider = self.provider.as_ref().ok_or_else(|| {
269 RagError::QueryPreprocessing(
270 "No LLM provider configured for expansion".into(),
271 )
272 })?;
273 let expander = crate::preprocessing::hyde::HydeExpander;
274 let expanded = expander.expand(&request.query, provider.as_ref()).await?;
275 Ok(format!("{} {}", request.query, expanded))
277 }
278 #[cfg(not(feature = "hyde"))]
279 {
280 let _ = count;
281 Ok(request.query.clone())
282 }
283 }
284 _ => Ok(request.query.clone()),
285 }
286 }
287
288 fn assemble_prompt(
290 &self,
291 query: &str,
292 context: &BuiltContext,
293 system_prompt: Option<&str>,
294 history: &[ChatMessage],
295 ) -> String {
296 let system = system_prompt.unwrap_or(&self.prompt_template.system);
297 let context_block = format!(
298 "{}{}{}",
299 self.prompt_template.context_prefix,
300 context.context_text,
301 self.prompt_template.context_suffix
302 );
303
304 let mut history_block = String::new();
306 for msg in history {
307 let role = match msg.role {
308 ChatRole::System => "System",
309 ChatRole::User => "User",
310 ChatRole::Assistant => "Assistant",
311 };
312 history_block.push_str(&format!("{}: {}\n", role, msg.content));
313 }
314
315 let user = if history_block.is_empty() {
316 self.prompt_template.render(query, &context_block)
317 } else {
318 let with_history = format!(
319 "Previous conversation:\n{}\n\nContext:\n{}\n\nQuestion: {}",
320 history_block, context_block, query
321 );
322 with_history
323 };
324
325 format!("{}\n\n{}", system, user)
326 }
327}
328
329#[async_trait]
330impl RagEngine for DefaultRagEngine {
331 async fn retrieve(&self, request: &RetrieveRequest) -> Result<RetrieveResult, RagError> {
332 self.do_retrieve(request).await
333 }
334
335 async fn retrieve_and_generate(&self, request: &RagRequest) -> Result<RagResponse, RagError> {
336 let start = Instant::now();
337
338 let retrieve_result = self.do_retrieve(&request.retrieve_config).await?;
340
341 if retrieve_result.hits.is_empty() {
342 return Err(RagError::NoResults(request.query.clone()));
343 }
344
345 let built = self.build_context(&retrieve_result.hits, &request.query);
347
348 let prompt = self.assemble_prompt(
350 &request.query,
351 &built,
352 request.system_prompt.as_deref(),
353 &request.history,
354 );
355
356 let (answer, _llm_usage) = {
358 #[cfg(feature = "llm-generation")]
359 {
360 let provider = self
361 .provider
362 .as_ref()
363 .ok_or_else(|| RagError::Provider("no LLM provider configured".into()))?;
364 crate::generation::generate_response(
365 provider,
366 &prompt,
367 request.generation.model.as_deref(),
368 request.generation.temperature,
369 request.generation.max_output_tokens,
370 )
371 .await
372 .map_err(|e| {
373 tracing::error!("LLM generation failed: {}", e);
374 e
375 })?
376 }
377 #[cfg(not(feature = "llm-generation"))]
378 {
379 (
380 format!(
381 "RAG Response for: '{}'\n\nBased on {} context chunks:\n{}",
382 request.query, built.chunks_used, built.context_text
383 ),
384 RagTokenUsage::default(),
385 )
386 }
387 };
388
389 let usage = RagTokenUsage {
390 context_tokens: built.tokens_used,
391 prompt_tokens: prompt.len() / 4,
392 completion_tokens: answer.len() / 4,
393 total_tokens: (prompt.len() + answer.len()) / 4,
394 chunks_used: built.chunks_used,
395 chunks_dropped: built.chunks_dropped,
396 };
397
398 let total_latency_ms = start.elapsed().as_millis() as u64;
399
400 info!(
401 query = %request.query,
402 hits = %retrieve_result.hits.len(),
403 chunks_used = %built.chunks_used,
404 latency_ms = %total_latency_ms,
405 "RAG retrieval and generation complete"
406 );
407
408 Ok(RagResponse {
409 answer,
410 citations: built.citations,
411 usage,
412 retrieve_stats: retrieve_result,
413 total_latency_ms,
414 model: request.generation.model.clone(),
415 })
416 }
417
418 async fn retrieve_and_generate_stream(
419 &self,
420 request: &RagRequest,
421 ) -> Result<Pin<Box<dyn Stream<Item = Result<RagStreamEvent, RagError>> + Send>>, RagError>
422 {
423 let retrieve_result = self.do_retrieve(&request.retrieve_config).await?;
424
425 if retrieve_result.hits.is_empty() {
426 return Err(RagError::NoResults(request.query.clone()));
427 }
428
429 let built = self.build_context(&retrieve_result.hits, &request.query);
430 #[cfg_attr(not(feature = "llm-generation"), allow(unused_variables))]
431 let prompt = self.assemble_prompt(
432 &request.query,
433 &built,
434 request.system_prompt.as_deref(),
435 &request.history,
436 );
437
438 #[cfg(feature = "llm-generation")]
439 {
440 let provider = self
441 .provider
442 .as_ref()
443 .ok_or_else(|| RagError::Provider("no LLM provider configured".into()))?;
444
445 let model_name = request
446 .generation
447 .model
448 .clone()
449 .or_else(|| provider.default_model().map(|s| s.to_string()))
450 .unwrap_or_else(|| "default".to_string());
451
452 let provider_request = ProviderCompletionRequest {
453 model: Some(model_name),
454 messages: vec![ProviderMessage::user(&prompt)],
455 temperature: request.generation.temperature,
456 max_tokens: request.generation.max_output_tokens,
457 stop: None,
458 frequency_penalty: None,
459 presence_penalty: None,
460 tools: None,
461 tool_choice: None,
462 response_format: None,
463 max_completion_tokens: None,
464 top_p: None,
465 top_k: None,
466 seed: None,
467 reasoning_effort: None,
468 logprobs: None,
469 logit_bias: None,
470 stream_include_usage: None,
471 request_id: String::new(),
472 };
473
474 let stream = provider
475 .complete_stream(provider_request, ProviderRequestOptions::default())
476 .await
477 .map_err(|e| RagError::Provider(format!("Stream failed: {}", e)))?;
478
479 let citations = built.citations.clone();
480 let context_chunks = built.chunks_used;
481 let context_tokens = built.tokens_used;
482
483 let mapped = futures::stream::unfold(
484 (stream, citations, context_chunks, context_tokens, true),
485 |(mut stream, citations, context_chunks, context_tokens, sent_start)| async move {
486 if sent_start {
487 Some((
488 Ok(RagStreamEvent::GenerationStarted {
489 context_chunks,
490 context_tokens,
491 }),
492 (stream, citations, context_chunks, context_tokens, false),
493 ))
494 } else {
495 match futures::StreamExt::next(&mut stream).await {
496 Some(Ok(StreamEvent::ContentDelta { delta })) => Some((
497 Ok(RagStreamEvent::ContentDelta { delta }),
498 (stream, citations, context_chunks, context_tokens, false),
499 )),
500 Some(Ok(StreamEvent::Done { .. })) => {
501 let done_event = Ok(RagStreamEvent::Done {
502 total_latency_ms: 0,
503 citations: citations.clone(),
504 usage: RagTokenUsage::default(),
505 });
506 Some((
507 done_event,
508 (stream, citations, context_chunks, context_tokens, false),
509 ))
510 }
511 Some(Ok(StreamEvent::Usage { .. })) => Some((
512 Ok(RagStreamEvent::ContentDelta { delta: String::new() }),
513 (stream, citations, context_chunks, context_tokens, false),
514 )),
515 Some(Ok(_other)) => Some((
516 Ok(RagStreamEvent::ContentDelta { delta: String::new() }),
517 (stream, citations, context_chunks, context_tokens, false),
518 )),
519 Some(Err(e)) => Some((
520 Err(RagError::Provider(format!("Stream error: {}", e))),
521 (stream, citations, context_chunks, context_tokens, false),
522 )),
523 None => None,
524 }
525 }
526 },
527 );
528
529 return Ok(Box::pin(mapped));
530 }
531
532 #[cfg(not(feature = "llm-generation"))]
533 {
534 let stream = futures::stream::iter(vec![
535 Ok(RagStreamEvent::GenerationStarted {
536 context_chunks: built.chunks_used,
537 context_tokens: built.tokens_used,
538 }),
539 Ok(RagStreamEvent::ContentDelta {
540 delta: format!(
541 "RAG Response for: '{}'\n\nBased on {} context chunks",
542 request.query, built.chunks_used
543 ),
544 }),
545 Ok(RagStreamEvent::Done {
546 total_latency_ms: 0,
547 citations: built.citations.clone(),
548 usage: RagTokenUsage {
549 context_tokens: built.tokens_used,
550 chunks_used: built.chunks_used,
551 chunks_dropped: built.chunks_dropped,
552 ..RagTokenUsage::default()
553 },
554 }),
555 ]);
556
557 Ok(Box::pin(stream))
558 }
559 }
560
561 fn engine_info(&self) -> RagEngineInfo {
562 self.info.clone()
563 }
564}
565
566#[derive(Default)]
568pub struct DefaultRagEngineBuilder {
569 name: Option<String>,
570 version: Option<String>,
571 pipeline: Option<ChannelPipeline>,
572 context_builder: Option<ContextBuilder>,
573 prompt_template: Option<PromptTemplate>,
574 embedder: Option<Arc<dyn Embedder>>,
575 semantic_store: Option<Arc<dyn SemanticSearch>>,
576 metadata_store: Option<Arc<dyn MetadataStore>>,
577 graph_store: Option<Arc<dyn KnowledgeGraphSearch>>,
578 #[cfg(feature = "rerank")]
579 reranker: Option<Arc<dyn Reranker>>,
580 #[cfg(feature = "llm-generation")]
581 provider: Option<Arc<dyn LlmProvider>>,
582 #[cfg(feature = "caching")]
583 cache: Option<crate::cache::memory_cache::RagMemoryCache>,
584}
585
586impl DefaultRagEngineBuilder {
587 pub fn name(mut self, name: impl Into<String>) -> Self {
589 self.name = Some(name.into());
590 self
591 }
592
593 pub fn version(mut self, version: impl Into<String>) -> Self {
595 self.version = Some(version.into());
596 self
597 }
598
599 pub fn pipeline(mut self, pipeline: ChannelPipeline) -> Self {
601 self.pipeline = Some(pipeline);
602 self
603 }
604
605 pub fn context_builder(mut self, cb: ContextBuilder) -> Self {
607 self.context_builder = Some(cb);
608 self
609 }
610
611 pub fn prompt_template(mut self, pt: PromptTemplate) -> Self {
613 self.prompt_template = Some(pt);
614 self
615 }
616
617 pub fn embedder(mut self, embedder: Arc<dyn Embedder>) -> Self {
619 self.embedder = Some(embedder);
620 self
621 }
622
623 pub fn semantic_store(mut self, store: Arc<dyn SemanticSearch>) -> Self {
625 self.semantic_store = Some(store);
626 self
627 }
628
629 pub fn metadata_store(mut self, store: Arc<dyn MetadataStore>) -> Self {
631 self.metadata_store = Some(store);
632 self
633 }
634
635 pub fn graph_store(mut self, store: Arc<dyn KnowledgeGraphSearch>) -> Self {
637 self.graph_store = Some(store);
638 self
639 }
640
641 #[cfg(feature = "rerank")]
643 pub fn reranker(mut self, reranker: Arc<dyn Reranker>) -> Self {
644 self.reranker = Some(reranker);
645 self
646 }
647
648 #[cfg(feature = "caching")]
650 pub fn cache(mut self, cache: crate::cache::memory_cache::RagMemoryCache) -> Self {
651 self.cache = Some(cache);
652 self
653 }
654
655 #[cfg(feature = "llm-generation")]
657 pub fn provider(mut self, provider: Arc<dyn LlmProvider>) -> Self {
658 self.provider = Some(provider);
659 self
660 }
661
662 pub fn build(self) -> DefaultRagEngine {
664 let pipeline = self.pipeline.unwrap_or_else(|| {
665 ChannelPipeline::new(vec![
666 ChannelConfig::semantic(0.5, 10).with_min_score(0.1),
667 ChannelConfig::metadata(0.3, 5),
668 ])
669 });
670
671 let mut supported_channels: Vec<String> =
672 pipeline.channels.iter().map(|c| c.channel_type.as_str().to_string()).collect();
673
674 if self.graph_store.is_some() && !supported_channels.iter().any(|c| c == "graph") {
675 supported_channels.push("graph".to_string());
676 }
677
678 DefaultRagEngine {
679 info: RagEngineInfo {
680 name: self.name.unwrap_or_else(|| "default".into()),
681 version: self.version.unwrap_or_else(|| "0.1.0".into()),
682 supported_channels,
683 #[cfg(feature = "llm-generation")]
684 supports_streaming: self.provider.is_some(),
685 #[cfg(not(feature = "llm-generation"))]
686 supports_streaming: false,
687 reranking_enabled: {
688 #[cfg(feature = "rerank")]
689 {
690 self.reranker.is_some()
691 }
692 #[cfg(not(feature = "rerank"))]
693 {
694 false
695 }
696 },
697 max_context_window: self
698 .context_builder
699 .as_ref()
700 .map(|cb| cb.context_budget())
701 .unwrap_or(4096),
702 },
703 pipeline,
704 context_builder: self.context_builder.unwrap_or_else(|| ContextBuilder::new(4096)),
705 prompt_template: self.prompt_template.unwrap_or_else(PromptTemplate::default_qa),
706 embedder: self.embedder,
707 semantic_store: self.semantic_store,
708 metadata_store: self.metadata_store,
709 graph_store: self.graph_store,
710 #[cfg(feature = "rerank")]
711 reranker: self.reranker,
712 #[cfg(feature = "llm-generation")]
713 provider: self.provider,
714 #[cfg(feature = "caching")]
715 cache: self.cache,
716 }
717 }
718}
719
720#[cfg(feature = "rerank")]
721fn retrieved_chunk_to_rerank_candidate(chunk: &RetrievedChunk) -> RerankCandidate {
722 RerankCandidate {
723 id: chunk.chunk_id.clone(),
724 content: chunk.content.clone(),
725 metadata: chunk_metadata_to_map(&chunk.metadata, &chunk.document_id),
726 retrieval_score: Some(chunk.score),
727 channel: Some(chunk.channel.clone()),
728 created_at: chunk.metadata.created_at,
729 embedding: chunk.embedding.clone(),
730 }
731}
732
733#[cfg(feature = "rerank")]
734fn chunk_metadata_to_map(
735 metadata: &crate::types::chunk::ChunkMetadata,
736 document_id: &str,
737) -> std::collections::HashMap<String, String> {
738 let mut map = metadata.extra.clone();
739
740 map.insert("document_id".to_string(), document_id.to_string());
741
742 if let Some(source) = &metadata.source {
743 map.insert("source".to_string(), source.clone());
744 }
745 if let Some(document_title) = &metadata.document_title {
746 map.insert("document_title".to_string(), document_title.clone());
747 }
748 if let Some(author) = &metadata.author {
749 map.insert("author".to_string(), author.clone());
750 }
751 if let Some(created_at) = metadata.created_at {
752 map.insert("created_at".to_string(), created_at.to_string());
753 }
754 if !metadata.tags.is_empty() {
755 map.insert("tags".to_string(), metadata.tags.join(","));
756 }
757 if let Some(namespace) = &metadata.namespace {
758 map.insert("namespace".to_string(), namespace.clone());
759 }
760
761 map
762}
763
764#[allow(unused)]
766fn build_cache_key(query: &str, namespace: &Option<String>) -> String {
767 if let Some(ns) = namespace {
768 format!("rag:{}:{}", ns, query)
769 } else {
770 format!("rag:default:{}", query)
771 }
772}