Skip to main content

khive_retrieval/
query_ir.rs

1//! Query Intermediate Representation for the retrieval pipeline.
2//!
3//! Composable IR tree representing how sub-queries are structured, filtered, and fused,
4//! separate from the Query struct which captures what to search.
5
6use khive_score::DeterministicScore;
7use serde::{Deserialize, Serialize};
8
9// ---------------------------------------------------------------------------
10// Core IR node
11// ---------------------------------------------------------------------------
12
13/// A node in the Query IR tree.
14///
15/// Each variant represents a single retrieval operation or combinator.
16/// Nodes compose recursively -- a `Fuse` holds children, a `Filter` wraps
17/// a single child, and leaf nodes (`Vector`, `Keyword`, `Empty`) terminate
18/// the tree.
19#[derive(Debug, Clone, Serialize, Deserialize)]
20#[serde(rename_all = "snake_case")]
21pub enum QueryNode {
22    /// Vector similarity search (e.g. HNSW nearest-neighbor).
23    Vector {
24        /// Pre-computed query embedding.
25        embedding: Vec<f32>,
26        /// Number of results to return.
27        top_k: usize,
28        /// Optional minimum similarity threshold.
29        min_score: Option<DeterministicScore>,
30    },
31
32    /// Keyword / BM25 text search.
33    Keyword {
34        /// Query text.
35        text: String,
36        /// Number of results to return.
37        top_k: usize,
38        /// Optional minimum relevance threshold.
39        min_score: Option<DeterministicScore>,
40    },
41
42    /// Fuse multiple sub-queries into a single ranked list.
43    Fuse {
44        /// Sub-queries to fuse.
45        children: Vec<QueryNode>,
46        /// Strategy for combining ranked lists.
47        strategy: FuseStrategy,
48        /// Number of results after fusion.
49        top_k: usize,
50    },
51
52    /// Filter the results of a sub-query.
53    Filter {
54        /// The sub-query whose results are filtered.
55        child: Box<QueryNode>,
56        /// Predicate to apply.
57        predicate: FilterPredicate,
58    },
59
60    /// Rerank the results of a sub-query.
61    Rerank {
62        /// The sub-query whose results are reranked.
63        child: Box<QueryNode>,
64        /// Reranking method.
65        method: RerankMethod,
66        /// Number of results after reranking.
67        top_k: usize,
68    },
69
70    /// An empty query that is guaranteed to produce no results.
71    ///
72    /// Useful as the result of constant-folding provably-empty sub-trees.
73    Empty,
74}
75
76// ---------------------------------------------------------------------------
77// Supporting enums
78// ---------------------------------------------------------------------------
79
80/// Fusion strategy for combining sub-query result lists.
81///
82/// Mirrors [`FusionStrategy`](crate::FusionStrategy) at the IR level
83/// so that the query plan is self-contained and serialisable without depending
84/// on runtime fusion internals.
85#[derive(Debug, Clone, Serialize, Deserialize)]
86#[serde(rename_all = "snake_case")]
87pub enum FuseStrategy {
88    /// Reciprocal Rank Fusion with smoothing constant `k`.
89    ///
90    /// Standard default: k = 60 (Cormack et al., 2009).
91    Rrf {
92        /// Smoothing constant.
93        k: usize,
94    },
95
96    /// Weighted linear combination of scores.
97    ///
98    /// One weight per child; weights are normalised at execution time.
99    Weighted {
100        /// Per-child weights (will be normalised to sum to 1.0).
101        weights: Vec<f64>,
102    },
103
104    /// Union with max-score-per-document semantics.
105    Union,
106}
107
108/// Predicate for post-retrieval filtering.
109#[derive(Debug, Clone, Serialize, Deserialize)]
110#[serde(rename_all = "snake_case")]
111pub enum FilterPredicate {
112    /// Keep only results whose score meets a minimum threshold.
113    MinScore(DeterministicScore),
114
115    /// Keep at most `k` results (top-k truncation).
116    TopK(usize),
117
118    /// Keep results where a metadata field equals a given value.
119    MetadataEquals {
120        /// Metadata field name.
121        field: String,
122        /// Expected value (JSON).
123        value: serde_json::Value,
124    },
125
126    /// All contained predicates must hold (conjunction).
127    And(Vec<FilterPredicate>),
128
129    /// At least one contained predicate must hold (disjunction).
130    Or(Vec<FilterPredicate>),
131}
132
133/// Method for reranking search results.
134#[derive(Debug, Clone, Serialize, Deserialize)]
135#[serde(rename_all = "snake_case")]
136pub enum RerankMethod {
137    /// Cross-encoder neural reranking (placeholder for future integration).
138    CrossEncoder {
139        /// Model identifier.
140        model: String,
141    },
142
143    /// Score-based reranking with custom per-signal weights.
144    ScoreWeighted {
145        /// Weights applied to each scoring signal.
146        weights: Vec<f64>,
147    },
148}
149
150// ---------------------------------------------------------------------------
151// Construction helpers
152// ---------------------------------------------------------------------------
153
154impl QueryNode {
155    /// Create a vector search leaf node.
156    pub fn vector(embedding: Vec<f32>, top_k: usize) -> Self {
157        QueryNode::Vector {
158            embedding,
159            top_k,
160            min_score: None,
161        }
162    }
163
164    /// Create a keyword search leaf node.
165    pub fn keyword(text: impl Into<String>, top_k: usize) -> Self {
166        QueryNode::Keyword {
167            text: text.into(),
168            top_k,
169            min_score: None,
170        }
171    }
172
173    /// Create a hybrid query (vector + keyword with RRF fusion, `top_k * 3` candidates each).
174    pub fn hybrid(embedding: Vec<f32>, text: impl Into<String>, top_k: usize) -> Self {
175        let candidate_k = top_k.saturating_mul(3);
176        QueryNode::Fuse {
177            children: vec![
178                QueryNode::vector(embedding, candidate_k),
179                QueryNode::keyword(text, candidate_k),
180            ],
181            strategy: FuseStrategy::Rrf { k: 60 },
182            top_k,
183        }
184    }
185
186    /// Wrap this node with a minimum-score filter.
187    #[must_use]
188    pub fn with_min_score(self, min_score: DeterministicScore) -> Self {
189        QueryNode::Filter {
190            child: Box::new(self),
191            predicate: FilterPredicate::MinScore(min_score),
192        }
193    }
194
195    /// Wrap this node with a top-k truncation filter.
196    #[must_use]
197    pub fn with_top_k(self, k: usize) -> Self {
198        QueryNode::Filter {
199            child: Box::new(self),
200            predicate: FilterPredicate::TopK(k),
201        }
202    }
203
204    // -----------------------------------------------------------------------
205    // Analysis helpers
206    // -----------------------------------------------------------------------
207
208    /// Returns `true` if this query is provably empty (no results possible).
209    ///
210    /// A query is provably empty when:
211    /// - It is the `Empty` variant.
212    /// - A leaf has `top_k == 0`.
213    /// - A keyword leaf has empty text.
214    /// - A fuse node has no children.
215    /// - A filter/rerank wraps a provably-empty child.
216    pub fn is_empty(&self) -> bool {
217        match self {
218            QueryNode::Empty => true,
219            QueryNode::Vector { top_k: 0, .. } => true,
220            QueryNode::Keyword { top_k: 0, .. } => true,
221            QueryNode::Keyword { text, .. } if text.is_empty() => true,
222            QueryNode::Fuse { children, .. } if children.is_empty() => true,
223            QueryNode::Filter { child, .. } => child.is_empty(),
224            QueryNode::Rerank { child, .. } => child.is_empty(),
225            _ => false,
226        }
227    }
228
229    /// Count the total number of leaf search operations in the tree.
230    ///
231    /// `Vector` and `Keyword` nodes each count as 1.  `Empty` counts as 0.
232    /// Combinators recurse into their children.
233    pub fn leaf_count(&self) -> usize {
234        match self {
235            QueryNode::Vector { .. } | QueryNode::Keyword { .. } => 1,
236            QueryNode::Fuse { children, .. } => children.iter().map(|c| c.leaf_count()).sum(),
237            QueryNode::Filter { child, .. } | QueryNode::Rerank { child, .. } => child.leaf_count(),
238            QueryNode::Empty => 0,
239        }
240    }
241
242    /// Return the effective `top_k` requested by this node.
243    ///
244    /// For `Filter` nodes with a `TopK` predicate, the predicate's value is
245    /// returned.  Otherwise the child's `top_k` propagates upward.
246    pub fn top_k(&self) -> usize {
247        match self {
248            QueryNode::Vector { top_k, .. } => *top_k,
249            QueryNode::Keyword { top_k, .. } => *top_k,
250            QueryNode::Fuse { top_k, .. } => *top_k,
251            QueryNode::Filter { child, predicate } => match predicate {
252                FilterPredicate::TopK(k) => *k,
253                _ => child.top_k(),
254            },
255            QueryNode::Rerank { top_k, .. } => *top_k,
256            QueryNode::Empty => 0,
257        }
258    }
259
260    /// Return the depth of the IR tree (longest root-to-leaf path).
261    ///
262    /// Leaf nodes have depth 1.  `Empty` has depth 0.
263    pub fn depth(&self) -> usize {
264        match self {
265            QueryNode::Empty => 0,
266            QueryNode::Vector { .. } | QueryNode::Keyword { .. } => 1,
267            QueryNode::Fuse { children, .. } => {
268                1 + children.iter().map(|c| c.depth()).max().unwrap_or(0)
269            }
270            QueryNode::Filter { child, .. } | QueryNode::Rerank { child, .. } => 1 + child.depth(),
271        }
272    }
273}
274
275// ---------------------------------------------------------------------------
276// Tests
277// ---------------------------------------------------------------------------
278
279#[cfg(test)]
280#[allow(clippy::uninlined_format_args)]
281mod tests {
282    use super::*;
283
284    // -- Construction -------------------------------------------------------
285
286    #[test]
287    fn test_vector_construction() {
288        let emb = vec![0.1, 0.2, 0.3];
289        let node = QueryNode::vector(emb.clone(), 10);
290        match &node {
291            QueryNode::Vector {
292                embedding,
293                top_k,
294                min_score,
295            } => {
296                assert_eq!(embedding, &emb);
297                assert_eq!(*top_k, 10);
298                assert!(min_score.is_none());
299            }
300            other => panic!("expected Vector, got {:?}", other),
301        }
302    }
303
304    #[test]
305    fn test_keyword_construction() {
306        let node = QueryNode::keyword("hello world", 5);
307        match &node {
308            QueryNode::Keyword {
309                text,
310                top_k,
311                min_score,
312            } => {
313                assert_eq!(text, "hello world");
314                assert_eq!(*top_k, 5);
315                assert!(min_score.is_none());
316            }
317            other => panic!("expected Keyword, got {:?}", other),
318        }
319    }
320
321    #[test]
322    fn test_hybrid_construction() {
323        let emb = vec![0.1_f32; 128];
324        let node = QueryNode::hybrid(emb, "distributed consensus", 10);
325        match &node {
326            QueryNode::Fuse {
327                children,
328                strategy,
329                top_k,
330            } => {
331                assert_eq!(children.len(), 2);
332                assert_eq!(*top_k, 10);
333                // Sub-queries should request 3x candidates.
334                assert_eq!(children[0].top_k(), 30);
335                assert_eq!(children[1].top_k(), 30);
336                assert!(matches!(strategy, FuseStrategy::Rrf { k: 60 }));
337            }
338            other => panic!("expected Fuse, got {:?}", other),
339        }
340    }
341
342    // -- is_empty -----------------------------------------------------------
343
344    #[test]
345    fn test_empty_variant() {
346        assert!(QueryNode::Empty.is_empty());
347        assert_eq!(QueryNode::Empty.leaf_count(), 0);
348        assert_eq!(QueryNode::Empty.top_k(), 0);
349        assert_eq!(QueryNode::Empty.depth(), 0);
350    }
351
352    #[test]
353    fn test_vector_top_k_zero_is_empty() {
354        let node = QueryNode::vector(vec![1.0], 0);
355        assert!(node.is_empty());
356    }
357
358    #[test]
359    fn test_keyword_top_k_zero_is_empty() {
360        let node = QueryNode::keyword("hello", 0);
361        assert!(node.is_empty());
362    }
363
364    #[test]
365    fn test_keyword_empty_text_is_empty() {
366        let node = QueryNode::keyword("", 10);
367        assert!(node.is_empty());
368    }
369
370    #[test]
371    fn test_fuse_no_children_is_empty() {
372        let node = QueryNode::Fuse {
373            children: vec![],
374            strategy: FuseStrategy::Rrf { k: 60 },
375            top_k: 10,
376        };
377        assert!(node.is_empty());
378    }
379
380    #[test]
381    fn test_filter_of_empty_is_empty() {
382        let node = QueryNode::Empty.with_min_score(DeterministicScore::from_f64(0.5));
383        assert!(node.is_empty());
384    }
385
386    #[test]
387    fn test_rerank_of_empty_is_empty() {
388        let node = QueryNode::Rerank {
389            child: Box::new(QueryNode::Empty),
390            method: RerankMethod::ScoreWeighted { weights: vec![1.0] },
391            top_k: 10,
392        };
393        assert!(node.is_empty());
394    }
395
396    #[test]
397    fn test_non_empty_query() {
398        let node = QueryNode::keyword("hello", 5);
399        assert!(!node.is_empty());
400    }
401
402    // -- leaf_count ---------------------------------------------------------
403
404    #[test]
405    fn test_leaf_count_single() {
406        assert_eq!(QueryNode::vector(vec![1.0], 5).leaf_count(), 1);
407        assert_eq!(QueryNode::keyword("q", 5).leaf_count(), 1);
408    }
409
410    #[test]
411    fn test_leaf_count_hybrid() {
412        let q = QueryNode::hybrid(vec![1.0], "q", 10);
413        assert_eq!(q.leaf_count(), 2);
414    }
415
416    #[test]
417    fn test_leaf_count_nested() {
418        // Fuse(Fuse(vec, kw), kw) = 3 leaves
419        let inner = QueryNode::hybrid(vec![1.0], "inner", 10);
420        let outer = QueryNode::Fuse {
421            children: vec![inner, QueryNode::keyword("outer", 10)],
422            strategy: FuseStrategy::Union,
423            top_k: 10,
424        };
425        assert_eq!(outer.leaf_count(), 3);
426    }
427
428    // -- top_k --------------------------------------------------------------
429
430    #[test]
431    fn test_top_k_leaf() {
432        assert_eq!(QueryNode::vector(vec![], 7).top_k(), 7);
433        assert_eq!(QueryNode::keyword("q", 3).top_k(), 3);
434    }
435
436    #[test]
437    fn test_top_k_fuse() {
438        let q = QueryNode::hybrid(vec![1.0], "q", 15);
439        assert_eq!(q.top_k(), 15);
440    }
441
442    #[test]
443    fn test_top_k_filter_topk_predicate() {
444        let node = QueryNode::keyword("q", 100).with_top_k(5);
445        assert_eq!(node.top_k(), 5);
446    }
447
448    #[test]
449    fn test_top_k_filter_non_topk_predicate() {
450        let node = QueryNode::keyword("q", 20).with_min_score(DeterministicScore::from_f64(0.5));
451        // min_score filter doesn't change top_k -- falls through to child.
452        assert_eq!(node.top_k(), 20);
453    }
454
455    // -- depth --------------------------------------------------------------
456
457    #[test]
458    fn test_depth_leaf() {
459        assert_eq!(QueryNode::vector(vec![1.0], 5).depth(), 1);
460        assert_eq!(QueryNode::keyword("q", 5).depth(), 1);
461    }
462
463    #[test]
464    fn test_depth_hybrid() {
465        let q = QueryNode::hybrid(vec![1.0], "q", 10);
466        // Fuse -> leaf = depth 2
467        assert_eq!(q.depth(), 2);
468    }
469
470    #[test]
471    fn test_depth_chained_filters() {
472        let q = QueryNode::keyword("q", 10)
473            .with_min_score(DeterministicScore::from_f64(0.5))
474            .with_top_k(5);
475        // TopK(Filter(MinScore(Keyword))) = 3 wrappers + 1 leaf = depth 3
476        assert_eq!(q.depth(), 3);
477    }
478
479    // -- with_min_score / with_top_k chaining -------------------------------
480
481    #[test]
482    fn test_builder_chaining() {
483        let node = QueryNode::keyword("rust async patterns", 20)
484            .with_min_score(DeterministicScore::from_f64(0.3))
485            .with_top_k(10);
486
487        assert_eq!(node.top_k(), 10);
488        assert_eq!(node.leaf_count(), 1);
489        assert!(!node.is_empty());
490    }
491
492    // -- Serde round-trip ---------------------------------------------------
493
494    #[test]
495    fn test_serde_roundtrip_vector() {
496        let node = QueryNode::vector(vec![0.1, 0.2, 0.3], 10);
497        let json = serde_json::to_string(&node).expect("serialize");
498        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
499        assert_eq!(back.top_k(), 10);
500        assert_eq!(back.leaf_count(), 1);
501    }
502
503    #[test]
504    fn test_serde_roundtrip_keyword() {
505        let node = QueryNode::keyword("hello world", 5);
506        let json = serde_json::to_string(&node).expect("serialize");
507        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
508        assert_eq!(back.top_k(), 5);
509    }
510
511    #[test]
512    fn test_serde_roundtrip_hybrid() {
513        let node = QueryNode::hybrid(vec![1.0, 2.0], "search query", 10);
514        let json = serde_json::to_string(&node).expect("serialize");
515        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
516        assert_eq!(back.top_k(), 10);
517        assert_eq!(back.leaf_count(), 2);
518    }
519
520    #[test]
521    fn test_serde_roundtrip_complex() {
522        let node = QueryNode::hybrid(vec![0.5; 4], "complex query", 10)
523            .with_min_score(DeterministicScore::from_f64(0.2))
524            .with_top_k(5);
525
526        let json = serde_json::to_string_pretty(&node).expect("serialize");
527        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
528
529        assert_eq!(back.top_k(), 5);
530        assert_eq!(back.leaf_count(), 2);
531        assert!(!back.is_empty());
532    }
533
534    #[test]
535    fn test_serde_roundtrip_empty() {
536        let node = QueryNode::Empty;
537        let json = serde_json::to_string(&node).expect("serialize");
538        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
539        assert!(back.is_empty());
540    }
541
542    #[test]
543    fn test_serde_roundtrip_filter_metadata() {
544        let node = QueryNode::Filter {
545            child: Box::new(QueryNode::keyword("docs", 10)),
546            predicate: FilterPredicate::MetadataEquals {
547                field: "type".to_string(),
548                value: serde_json::json!("memory"),
549            },
550        };
551        let json = serde_json::to_string(&node).expect("serialize");
552        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
553        assert_eq!(back.leaf_count(), 1);
554    }
555
556    #[test]
557    fn test_serde_roundtrip_rerank() {
558        let node = QueryNode::Rerank {
559            child: Box::new(QueryNode::keyword("rerank me", 20)),
560            method: RerankMethod::CrossEncoder {
561                model: "ms-marco-MiniLM".to_string(),
562            },
563            top_k: 10,
564        };
565        let json = serde_json::to_string(&node).expect("serialize");
566        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
567        assert_eq!(back.top_k(), 10);
568    }
569
570    #[test]
571    fn test_serde_roundtrip_compound_predicate() {
572        let pred = FilterPredicate::And(vec![
573            FilterPredicate::MinScore(DeterministicScore::from_f64(0.3)),
574            FilterPredicate::Or(vec![
575                FilterPredicate::MetadataEquals {
576                    field: "lang".to_string(),
577                    value: serde_json::json!("en"),
578                },
579                FilterPredicate::MetadataEquals {
580                    field: "lang".to_string(),
581                    value: serde_json::json!("zh"),
582                },
583            ]),
584        ]);
585        let node = QueryNode::Filter {
586            child: Box::new(QueryNode::keyword("test", 10)),
587            predicate: pred,
588        };
589        let json = serde_json::to_string(&node).expect("serialize");
590        let back: QueryNode = serde_json::from_str(&json).expect("deserialize");
591        assert_eq!(back.leaf_count(), 1);
592    }
593}