Skip to main content

khive_retrieval/hybrid/
config.rs

1//! Hybrid search configuration types.
2
3use std::time::Duration;
4
5use khive_score::DeterministicScore;
6use serde::{Deserialize, Serialize};
7
8use khive_fusion::FusionStrategy;
9
10/// Default candidate pool multiplier over top_k.
11pub const DEFAULT_POOL_MULTIPLIER: usize = 5;
12
13/// Query for hybrid search.
14///
15/// Combines text for keyword search and optional embedding for vector search.
16#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct Query {
18    /// Text for keyword search (required).
19    pub text: String,
20
21    /// Pre-computed embedding for vector search (optional).
22    ///
23    /// If None, vector search is skipped or caller must provide.
24    pub embedding: Option<Vec<f32>>,
25
26    /// Optional filters to apply post-retrieval.
27    pub filters: Option<serde_json::Value>,
28}
29
30impl Query {
31    /// Create a new query with text only (keyword search).
32    pub fn text(text: impl Into<String>) -> Self {
33        Self {
34            text: text.into(),
35            embedding: None,
36            filters: None,
37        }
38    }
39
40    /// Create a query with both text and embedding (hybrid search).
41    pub fn hybrid(text: impl Into<String>, embedding: Vec<f32>) -> Self {
42        Self {
43            text: text.into(),
44            embedding: Some(embedding),
45            filters: None,
46        }
47    }
48
49    /// Add filters to the query.
50    #[must_use]
51    pub fn with_filters(mut self, filters: serde_json::Value) -> Self {
52        self.filters = Some(filters);
53        self
54    }
55
56    /// Check if this query supports vector search.
57    pub fn has_embedding(&self) -> bool {
58        self.embedding.is_some()
59    }
60}
61
62/// Raw wire format for [`HybridConfig`], used by `TryFrom` validation.
63#[derive(Deserialize)]
64struct RawHybridConfig {
65    fusion_strategy: FusionStrategy,
66    top_k: usize,
67    candidate_pool_size: usize,
68    min_score: Option<DeterministicScore>,
69    vector_weight: f64,
70    keyword_weight: f64,
71    #[serde(default, with = "crate::timeout::serde_opt_duration")]
72    timeout: Option<Duration>,
73}
74
75impl TryFrom<RawHybridConfig> for HybridConfig {
76    type Error = String;
77
78    fn try_from(raw: RawHybridConfig) -> Result<Self, Self::Error> {
79        if !raw.vector_weight.is_finite() {
80            return Err(format!(
81                "vector_weight must be finite, got {}",
82                raw.vector_weight
83            ));
84        }
85        if !raw.keyword_weight.is_finite() {
86            return Err(format!(
87                "keyword_weight must be finite, got {}",
88                raw.keyword_weight
89            ));
90        }
91        Ok(HybridConfig {
92            fusion_strategy: raw.fusion_strategy,
93            top_k: raw.top_k,
94            candidate_pool_size: raw.candidate_pool_size,
95            min_score: raw.min_score,
96            vector_weight: raw.vector_weight,
97            keyword_weight: raw.keyword_weight,
98            timeout: raw.timeout,
99        })
100    }
101}
102
103/// Configuration for hybrid search.
104#[derive(Debug, Clone, Serialize, Deserialize)]
105#[serde(try_from = "RawHybridConfig")]
106pub struct HybridConfig {
107    /// Fusion strategy to use (default: RRF with k=60).
108    pub fusion_strategy: FusionStrategy,
109
110    /// Number of results to return.
111    pub top_k: usize,
112
113    /// Candidates to fetch from each retriever before fusion.
114    ///
115    /// Should be >= 5 * top_k for quality fusion.
116    pub candidate_pool_size: usize,
117
118    /// Minimum score threshold (post-fusion).
119    pub min_score: Option<DeterministicScore>,
120
121    /// Weight for vector search results (0.0 to 1.0).
122    ///
123    /// Only used when fusion_strategy is Weighted.
124    pub vector_weight: f64,
125
126    /// Weight for keyword search results (0.0 to 1.0).
127    ///
128    /// Only used when fusion_strategy is Weighted.
129    pub keyword_weight: f64,
130
131    /// Optional timeout for the entire search operation.
132    ///
133    /// If set, the search will be cancelled if it exceeds this duration,
134    /// returning [`crate::error::RetrievalError::QueryTimeout`].
135    /// If None, no timeout is applied.
136    #[serde(
137        default,
138        skip_serializing_if = "Option::is_none",
139        with = "crate::timeout::serde_opt_duration"
140    )]
141    pub timeout: Option<Duration>,
142}
143
144impl Default for HybridConfig {
145    fn default() -> Self {
146        Self {
147            fusion_strategy: FusionStrategy::rrf(),
148            top_k: 10,
149            candidate_pool_size: 50, // 5 * top_k
150            min_score: None,
151            vector_weight: 0.7,
152            keyword_weight: 0.3,
153            timeout: None,
154        }
155    }
156}
157
158impl HybridConfig {
159    /// Create a new config with specified top_k.
160    ///
161    /// The candidate pool size is `top_k * DEFAULT_POOL_MULTIPLIER`, saturating
162    /// at `usize::MAX` on overflow (rather than wrapping or panicking in debug).
163    pub fn new(top_k: usize) -> Self {
164        Self {
165            top_k,
166            candidate_pool_size: top_k.saturating_mul(DEFAULT_POOL_MULTIPLIER),
167            ..Default::default()
168        }
169    }
170
171    /// Set the fusion strategy.
172    #[must_use]
173    pub fn with_fusion_strategy(mut self, strategy: FusionStrategy) -> Self {
174        self.fusion_strategy = strategy;
175        self
176    }
177
178    /// Set the candidate pool size.
179    #[must_use]
180    pub fn with_pool_size(mut self, size: usize) -> Self {
181        self.candidate_pool_size = size;
182        self
183    }
184
185    /// Set the minimum score threshold.
186    #[must_use]
187    pub fn with_min_score(mut self, score: DeterministicScore) -> Self {
188        self.min_score = Some(score);
189        self
190    }
191
192    /// Set weights for weighted fusion (clamped to [0.0, 1.0]). Debug-asserts both weights are finite.
193    #[must_use]
194    pub fn with_weights(mut self, vector: f64, keyword: f64) -> Self {
195        debug_assert!(
196            vector.is_finite(),
197            "vector weight must be finite, got {vector}"
198        );
199        debug_assert!(
200            keyword.is_finite(),
201            "keyword weight must be finite, got {keyword}"
202        );
203        self.vector_weight = vector.clamp(0.0, 1.0);
204        self.keyword_weight = keyword.clamp(0.0, 1.0);
205        self
206    }
207
208    /// Set the search timeout.
209    ///
210    /// If the search operation exceeds this duration, it will return
211    /// [`crate::error::RetrievalError::QueryTimeout`].
212    #[must_use]
213    pub fn with_timeout(mut self, timeout: Duration) -> Self {
214        self.timeout = Some(timeout);
215        self
216    }
217
218    /// Get normalized weights that sum to 1.0.
219    ///
220    /// If both weights are zero or their sum is non-finite, returns equal weights (0.5, 0.5).
221    pub fn normalized_weights(&self) -> (f64, f64) {
222        let sum = self.vector_weight + self.keyword_weight;
223        if sum <= 0.0 || !sum.is_finite() {
224            (0.5, 0.5)
225        } else {
226            (self.vector_weight / sum, self.keyword_weight / sum)
227        }
228    }
229}
230
231#[cfg(test)]
232mod tests {
233    use super::*;
234
235    #[test]
236    fn test_query_text_only() {
237        let q = Query::text("hello world");
238        assert_eq!(q.text, "hello world");
239        assert!(q.embedding.is_none());
240        assert!(!q.has_embedding());
241    }
242
243    #[test]
244    fn test_query_hybrid() {
245        let embedding = vec![0.1, 0.2, 0.3];
246        let q = Query::hybrid("hello", embedding.clone());
247        assert_eq!(q.text, "hello");
248        assert_eq!(q.embedding, Some(embedding));
249        assert!(q.has_embedding());
250    }
251
252    #[test]
253    fn test_query_with_filters() {
254        let q = Query::text("test").with_filters(serde_json::json!({"type": "memory"}));
255        assert!(q.filters.is_some());
256    }
257
258    #[test]
259    fn test_hybrid_config_default() {
260        let config = HybridConfig::default();
261        assert_eq!(config.top_k, 10);
262        assert_eq!(config.candidate_pool_size, 50);
263        assert!(matches!(
264            config.fusion_strategy,
265            FusionStrategy::Rrf { k: 60 }
266        ));
267        assert!(config.min_score.is_none());
268    }
269
270    #[test]
271    fn test_hybrid_config_new() {
272        let config = HybridConfig::new(20);
273        assert_eq!(config.top_k, 20);
274        assert_eq!(config.candidate_pool_size, 100); // 20 * 5
275    }
276
277    #[test]
278    fn test_hybrid_config_builder() {
279        let config = HybridConfig::new(10)
280            .with_fusion_strategy(FusionStrategy::union())
281            .with_pool_size(200)
282            .with_weights(0.6, 0.4);
283
284        assert_eq!(config.top_k, 10);
285        assert_eq!(config.candidate_pool_size, 200);
286        assert!(matches!(config.fusion_strategy, FusionStrategy::Union));
287        assert_eq!(config.vector_weight, 0.6);
288        assert_eq!(config.keyword_weight, 0.4);
289    }
290
291    #[test]
292    fn test_normalized_weights() {
293        let config = HybridConfig::default();
294        let (v, k) = config.normalized_weights();
295        assert!((v - 0.7).abs() < 0.01);
296        assert!((k - 0.3).abs() < 0.01);
297
298        // Zero weights -> equal
299        let config = HybridConfig::default().with_weights(0.0, 0.0);
300        let (v, k) = config.normalized_weights();
301        assert!((v - 0.5).abs() < 0.01);
302        assert!((k - 0.5).abs() < 0.01);
303    }
304
305    #[test]
306    fn test_weight_clamping() {
307        let config = HybridConfig::default().with_weights(1.5, -0.5);
308        assert_eq!(config.vector_weight, 1.0);
309        assert_eq!(config.keyword_weight, 0.0);
310    }
311
312    /// Parser-level rejection only; see `crates/khive-retrieval/docs/api/hybrid-config.md`
313    /// for why this is not the real `TryFrom` regression guard.
314    #[test]
315    fn test_serde_json_rejects_nan_literal_vector_weight() {
316        let json = r#"{
317            "fusion_strategy": {"rrf": {"k": 60}},
318            "top_k": 10,
319            "candidate_pool_size": 50,
320            "min_score": null,
321            "vector_weight": NaN,
322            "keyword_weight": 0.3
323        }"#;
324        let result: Result<HybridConfig, _> = serde_json::from_str(json);
325        assert!(
326            result.is_err(),
327            "JSON literal NaN must be rejected by the parser"
328        );
329    }
330
331    /// Same caveat as above: `Infinity` is not a valid JSON token either.
332    #[test]
333    fn test_serde_json_rejects_infinity_literal_vector_weight() {
334        let json = r#"{
335            "fusion_strategy": {"rrf": {"k": 60}},
336            "top_k": 10,
337            "candidate_pool_size": 50,
338            "min_score": null,
339            "vector_weight": Infinity,
340            "keyword_weight": 0.3
341        }"#;
342        let result: Result<HybridConfig, _> = serde_json::from_str(json);
343        assert!(
344            result.is_err(),
345            "JSON literal Infinity must be rejected by the parser"
346        );
347    }
348
349    /// The real `TryFrom<RawHybridConfig>` regression guard (constructed directly since
350    /// `serde_json` cannot encode a literal NaN). See
351    /// `crates/khive-retrieval/docs/api/hybrid-config.md` for the full rationale.
352    #[test]
353    fn test_try_from_rejects_nan_vector_weight() {
354        let raw = RawHybridConfig {
355            fusion_strategy: FusionStrategy::rrf(),
356            top_k: 10,
357            candidate_pool_size: 50,
358            min_score: None,
359            vector_weight: f64::NAN,
360            keyword_weight: 0.3,
361            timeout: None,
362        };
363        let result = HybridConfig::try_from(raw);
364        assert!(
365            result.is_err(),
366            "NaN vector_weight must be rejected via TryFrom"
367        );
368    }
369
370    #[test]
371    fn test_try_from_rejects_infinite_vector_weight() {
372        let raw = RawHybridConfig {
373            fusion_strategy: FusionStrategy::rrf(),
374            top_k: 10,
375            candidate_pool_size: 50,
376            min_score: None,
377            vector_weight: f64::INFINITY,
378            keyword_weight: 0.3,
379            timeout: None,
380        };
381        let result = HybridConfig::try_from(raw);
382        assert!(
383            result.is_err(),
384            "+Infinity vector_weight must be rejected via TryFrom"
385        );
386    }
387
388    #[test]
389    fn test_try_from_rejects_nan_keyword_weight() {
390        let raw = RawHybridConfig {
391            fusion_strategy: FusionStrategy::rrf(),
392            top_k: 10,
393            candidate_pool_size: 50,
394            min_score: None,
395            vector_weight: 0.7,
396            keyword_weight: f64::NAN,
397            timeout: None,
398        };
399        let result = HybridConfig::try_from(raw);
400        assert!(
401            result.is_err(),
402            "NaN keyword_weight must be rejected via TryFrom"
403        );
404    }
405
406    #[test]
407    fn test_try_from_rejects_negative_infinity_keyword_weight() {
408        let raw = RawHybridConfig {
409            fusion_strategy: FusionStrategy::rrf(),
410            top_k: 10,
411            candidate_pool_size: 50,
412            min_score: None,
413            vector_weight: 0.7,
414            keyword_weight: f64::NEG_INFINITY,
415            timeout: None,
416        };
417        let result = HybridConfig::try_from(raw);
418        assert!(
419            result.is_err(),
420            "-Infinity keyword_weight must be rejected via TryFrom"
421        );
422    }
423
424    /// Positive control: a valid finite config still deserializes correctly through the
425    /// `TryFrom` boundary — the fix must not reject legitimate input.
426    #[test]
427    fn test_serde_accepts_valid_finite_config() {
428        let json = r#"{
429            "fusion_strategy": {"rrf": {"k": 60}},
430            "top_k": 10,
431            "candidate_pool_size": 50,
432            "min_score": null,
433            "vector_weight": 0.7,
434            "keyword_weight": 0.3
435        }"#;
436        let config: HybridConfig = serde_json::from_str(json).expect("valid config");
437        assert_eq!(config.top_k, 10);
438        assert_eq!(config.candidate_pool_size, 50);
439        assert_eq!(config.vector_weight, 0.7);
440        assert_eq!(config.keyword_weight, 0.3);
441    }
442
443    #[test]
444    fn test_serde_roundtrip_preserves_default_config() {
445        let config = HybridConfig::default();
446        let json = serde_json::to_string(&config).unwrap();
447        let restored: HybridConfig = serde_json::from_str(&json).unwrap();
448        assert_eq!(restored.vector_weight, config.vector_weight);
449        assert_eq!(restored.keyword_weight, config.keyword_weight);
450        assert_eq!(restored.top_k, config.top_k);
451    }
452}