Skip to main content

xz_rag/
engine.rs

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
38/// Default multi-channel RAG engine implementation.
39pub struct DefaultRagEngine {
40    info: RagEngineInfo,
41    pipeline: ChannelPipeline,
42    context_builder: ContextBuilder,
43    prompt_template: PromptTemplate,
44    // Pluggable components
45    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    /// Create a new builder for configuring a `DefaultRagEngine`.
59    pub fn builder() -> DefaultRagEngineBuilder {
60        DefaultRagEngineBuilder::default()
61    }
62}
63
64impl DefaultRagEngine {
65    /// Execute retrieval across all configured channels.
66    async fn do_retrieve(&self, request: &RetrieveRequest) -> Result<RetrieveResult, RagError> {
67        // Check cache first
68        #[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        // Preprocess query (HYDE, expansion, translation)
78        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        // Normalize scores per channel
177        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        // RRF fusion
188        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        // Apply global top_k
229        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        // Store in cache
236        #[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    /// Build context from retrieved chunks.
246    fn build_context(&self, chunks: &[RetrievedChunk], query: &str) -> BuiltContext {
247        self.context_builder.build(chunks, query)
248    }
249
250    /// Preprocess query with HYDE, expansion, or translation.
251    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                    // Concatenate original with expanded for richer retrieval
276                    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    /// Assemble the full prompt for LLM generation with optional chat history.
289    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        // Format chat history
305        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        // Step 1: Retrieve
339        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        // Step 2: Build context
346        let built = self.build_context(&retrieve_result.hits, &request.query);
347
348        // Step 3: Assemble prompt with history
349        let prompt = self.assemble_prompt(
350            &request.query,
351            &built,
352            request.system_prompt.as_deref(),
353            &request.history,
354        );
355
356        // Step 4: Generate via LLM (or fallback to placeholder when feature disabled)
357        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/// Builder for DefaultRagEngine.
567#[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    /// Set the engine name.
588    pub fn name(mut self, name: impl Into<String>) -> Self {
589        self.name = Some(name.into());
590        self
591    }
592
593    /// Set the engine version.
594    pub fn version(mut self, version: impl Into<String>) -> Self {
595        self.version = Some(version.into());
596        self
597    }
598
599    /// Set the channel pipeline configuration.
600    pub fn pipeline(mut self, pipeline: ChannelPipeline) -> Self {
601        self.pipeline = Some(pipeline);
602        self
603    }
604
605    /// Set the context builder for token budget management.
606    pub fn context_builder(mut self, cb: ContextBuilder) -> Self {
607        self.context_builder = Some(cb);
608        self
609    }
610
611    /// Set the prompt template for LLM generation.
612    pub fn prompt_template(mut self, pt: PromptTemplate) -> Self {
613        self.prompt_template = Some(pt);
614        self
615    }
616
617    /// Set the embedder for semantic search.
618    pub fn embedder(mut self, embedder: Arc<dyn Embedder>) -> Self {
619        self.embedder = Some(embedder);
620        self
621    }
622
623    /// Set the semantic vector store.
624    pub fn semantic_store(mut self, store: Arc<dyn SemanticSearch>) -> Self {
625        self.semantic_store = Some(store);
626        self
627    }
628
629    /// Set the metadata store for metadata-based search.
630    pub fn metadata_store(mut self, store: Arc<dyn MetadataStore>) -> Self {
631        self.metadata_store = Some(store);
632        self
633    }
634
635    /// Set the knowledge graph store.
636    pub fn graph_store(mut self, store: Arc<dyn KnowledgeGraphSearch>) -> Self {
637        self.graph_store = Some(store);
638        self
639    }
640
641    /// Set the reranker for post-fusion reranking.
642    #[cfg(feature = "rerank")]
643    pub fn reranker(mut self, reranker: Arc<dyn Reranker>) -> Self {
644        self.reranker = Some(reranker);
645        self
646    }
647
648    /// Set the in-memory result cache.
649    #[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    /// Set the LLM provider for generation.
656    #[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    /// Build the `DefaultRagEngine` with the configured components.
663    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/// Build a cache key from query and optional namespace.
765#[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}