1use khive_score::DeterministicScore;
7use serde::{Deserialize, Serialize};
8
9#[derive(Debug, Clone, Serialize, Deserialize)]
20#[serde(rename_all = "snake_case")]
21pub enum QueryNode {
22 Vector {
24 embedding: Vec<f32>,
26 top_k: usize,
28 min_score: Option<DeterministicScore>,
30 },
31
32 Keyword {
34 text: String,
36 top_k: usize,
38 min_score: Option<DeterministicScore>,
40 },
41
42 Fuse {
44 children: Vec<QueryNode>,
46 strategy: FuseStrategy,
48 top_k: usize,
50 },
51
52 Filter {
54 child: Box<QueryNode>,
56 predicate: FilterPredicate,
58 },
59
60 Rerank {
62 child: Box<QueryNode>,
64 method: RerankMethod,
66 top_k: usize,
68 },
69
70 Empty,
74}
75
76#[derive(Debug, Clone, Serialize, Deserialize)]
86#[serde(rename_all = "snake_case")]
87pub enum FuseStrategy {
88 Rrf {
92 k: usize,
94 },
95
96 Weighted {
100 weights: Vec<f64>,
102 },
103
104 Union,
106}
107
108#[derive(Debug, Clone, Serialize, Deserialize)]
110#[serde(rename_all = "snake_case")]
111pub enum FilterPredicate {
112 MinScore(DeterministicScore),
114
115 TopK(usize),
117
118 MetadataEquals {
120 field: String,
122 value: serde_json::Value,
124 },
125
126 And(Vec<FilterPredicate>),
128
129 Or(Vec<FilterPredicate>),
131}
132
133#[derive(Debug, Clone, Serialize, Deserialize)]
135#[serde(rename_all = "snake_case")]
136pub enum RerankMethod {
137 CrossEncoder {
139 model: String,
141 },
142
143 ScoreWeighted {
145 weights: Vec<f64>,
147 },
148}
149
150impl QueryNode {
155 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 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 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 #[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 #[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 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 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 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 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#[cfg(test)]
280#[allow(clippy::uninlined_format_args)]
281mod tests {
282 use super::*;
283
284 #[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 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 #[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 #[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 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 #[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 assert_eq!(node.top_k(), 20);
453 }
454
455 #[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 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 assert_eq!(q.depth(), 3);
477 }
478
479 #[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 #[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}