lunaris_retrieve/operators/
combinators.rs1use std::any::Any;
20use std::collections::HashMap;
21
22use async_trait::async_trait;
23use lunaris_core::LunarisError;
24use lunaris_core::storage::types::Filter;
25
26use super::{QueryContext, Retriever};
27use crate::types::RawHit;
28
29#[must_use = "AndRetriever is a query node — pass it to RetrievalBuilder::with_root() or wrap it via .fuse_rrf / .or, otherwise it never executes"]
31pub struct AndRetriever {
32 pub(crate) left: Box<dyn Retriever>,
33 pub(crate) right: Box<dyn Retriever>,
34}
35
36impl AndRetriever {
37 pub fn branches(&self) -> (&dyn Retriever, &dyn Retriever) {
41 (self.left.as_ref(), self.right.as_ref())
42 }
43}
44
45impl AndRetriever {
46 pub fn new(left: Box<dyn Retriever>, right: Box<dyn Retriever>) -> Self {
47 Self { left, right }
48 }
49
50 pub fn fuse_rrf(self, k: u32) -> super::fuse::FuseRrfRetriever {
51 super::fuse::FuseRrfRetriever::new(Box::new(self), k as usize)
52 }
53
54 pub fn top(self, n: usize) -> super::modifiers::TopRetriever {
55 super::modifiers::TopRetriever::new(Box::new(self), n)
56 }
57
58 pub fn rerank(
60 self,
61 reranker: std::sync::Arc<dyn lunaris_rerank::Reranker>,
62 ) -> super::rerank::RerankRetriever {
63 super::rerank::RerankRetriever::new(Box::new(self), reranker)
64 }
65
66 pub fn degraded_fallback<R: Retriever + 'static>(
69 self,
70 fallback: R,
71 ) -> super::degraded::DegradedFallbackRetriever {
72 super::degraded::DegradedFallbackRetriever::new(Box::new(self), Box::new(fallback))
73 }
74}
75
76#[async_trait]
77impl Retriever for AndRetriever {
78 async fn retrieve(&self, ctx: &QueryContext) -> Result<Vec<RawHit>, LunarisError> {
79 let (left_res, right_res) = tokio::join!(self.left.retrieve(ctx), self.right.retrieve(ctx));
80 let mut out = left_res?;
81 out.extend(right_res?);
82 Ok(out)
83 }
84
85 fn as_any(&self) -> &dyn Any {
86 self
87 }
88}
89
90#[must_use = "OrRetriever is a query node — pass it to RetrievalBuilder::with_root() or wrap further, otherwise it never executes"]
92pub struct OrRetriever {
93 pub(crate) left: Box<dyn Retriever>,
94 pub(crate) right: Box<dyn Retriever>,
95}
96
97impl OrRetriever {
98 pub fn new(left: Box<dyn Retriever>, right: Box<dyn Retriever>) -> Self {
99 Self { left, right }
100 }
101
102 pub fn top(self, n: usize) -> super::modifiers::TopRetriever {
103 super::modifiers::TopRetriever::new(Box::new(self), n)
104 }
105}
106
107#[async_trait]
108impl Retriever for OrRetriever {
109 async fn retrieve(&self, ctx: &QueryContext) -> Result<Vec<RawHit>, LunarisError> {
110 let (left_res, right_res) = tokio::join!(self.left.retrieve(ctx), self.right.retrieve(ctx));
111 let mut by_id: HashMap<Vec<u8>, RawHit> = HashMap::new();
112 for h in left_res?.into_iter().chain(right_res?) {
113 by_id
114 .entry(h.id.clone())
115 .and_modify(|existing| {
116 if h.score > existing.score {
117 *existing = h.clone();
118 }
119 })
120 .or_insert(h);
121 }
122 let mut out: Vec<RawHit> = by_id.into_values().collect();
123 out.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
125 Ok(out)
126 }
127
128 fn as_any(&self) -> &dyn Any {
129 self
130 }
131}
132
133#[must_use = "ThenRetriever is a query node — pass it to RetrievalBuilder::with_root() or wrap further, otherwise it never executes"]
135pub struct ThenRetriever {
136 pub(crate) first: Box<dyn Retriever>,
137 pub(crate) second: Box<dyn Retriever>,
138}
139
140impl ThenRetriever {
141 pub fn new(first: Box<dyn Retriever>, second: Box<dyn Retriever>) -> Self {
142 Self { first, second }
143 }
144
145 pub fn top(self, n: usize) -> super::modifiers::TopRetriever {
146 super::modifiers::TopRetriever::new(Box::new(self), n)
147 }
148}
149
150pub fn then(first: Box<dyn Retriever>, second: Box<dyn Retriever>) -> ThenRetriever {
153 ThenRetriever::new(first, second)
154}
155
156#[async_trait]
157impl Retriever for ThenRetriever {
158 async fn retrieve(&self, ctx: &QueryContext) -> Result<Vec<RawHit>, LunarisError> {
159 let firsts = self.first.retrieve(ctx).await?;
160 if firsts.is_empty() {
161 return Ok(Vec::new());
162 }
163 let id_filter = Filter::Or(
165 firsts
166 .iter()
167 .map(|h| Filter::Eq {
168 field: "id".to_string(),
169 value: serde_json::Value::String(String::from_utf8_lossy(&h.id).into_owned()),
170 })
171 .collect(),
172 );
173
174 let new_filter = match &ctx.query.filter {
176 Some(existing) => Filter::And(vec![existing.clone(), id_filter]),
177 None => id_filter,
178 };
179
180 let mut narrowed_query = ctx.query.clone();
189 narrowed_query.filter = Some(new_filter);
190
191 let narrowed_embedding = tokio::sync::OnceCell::new();
192 if let Some(existing) = ctx.query_embedding.get().cloned() {
193 let _ = narrowed_embedding.set(existing);
194 }
195
196 let narrowed_ctx = QueryContext {
197 query: narrowed_query,
198 scope: ctx.scope.clone(),
199 embedder: ctx.embedder.clone(),
200 storage: ctx.storage.clone(),
201 keyword: ctx.keyword.clone(),
202 query_embedding: narrowed_embedding,
203 moon_storage: ctx.moon_storage.clone(),
204 };
205
206 self.second.retrieve(&narrowed_ctx).await
207 }
208
209 fn as_any(&self) -> &dyn Any {
210 self
211 }
212}