1use std::time::Duration;
4
5use khive_score::DeterministicScore;
6use serde::{Deserialize, Serialize};
7
8use khive_fusion::FusionStrategy;
9
10pub const DEFAULT_POOL_MULTIPLIER: usize = 5;
12
13#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct Query {
18 pub text: String,
20
21 pub embedding: Option<Vec<f32>>,
25
26 pub filters: Option<serde_json::Value>,
28}
29
30impl Query {
31 pub fn text(text: impl Into<String>) -> Self {
33 Self {
34 text: text.into(),
35 embedding: None,
36 filters: None,
37 }
38 }
39
40 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 #[must_use]
51 pub fn with_filters(mut self, filters: serde_json::Value) -> Self {
52 self.filters = Some(filters);
53 self
54 }
55
56 pub fn has_embedding(&self) -> bool {
58 self.embedding.is_some()
59 }
60}
61
62#[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#[derive(Debug, Clone, Serialize, Deserialize)]
105#[serde(try_from = "RawHybridConfig")]
106pub struct HybridConfig {
107 pub fusion_strategy: FusionStrategy,
109
110 pub top_k: usize,
112
113 pub candidate_pool_size: usize,
117
118 pub min_score: Option<DeterministicScore>,
120
121 pub vector_weight: f64,
125
126 pub keyword_weight: f64,
130
131 #[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, min_score: None,
151 vector_weight: 0.7,
152 keyword_weight: 0.3,
153 timeout: None,
154 }
155 }
156}
157
158impl HybridConfig {
159 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 #[must_use]
173 pub fn with_fusion_strategy(mut self, strategy: FusionStrategy) -> Self {
174 self.fusion_strategy = strategy;
175 self
176 }
177
178 #[must_use]
180 pub fn with_pool_size(mut self, size: usize) -> Self {
181 self.candidate_pool_size = size;
182 self
183 }
184
185 #[must_use]
187 pub fn with_min_score(mut self, score: DeterministicScore) -> Self {
188 self.min_score = Some(score);
189 self
190 }
191
192 #[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 #[must_use]
213 pub fn with_timeout(mut self, timeout: Duration) -> Self {
214 self.timeout = Some(timeout);
215 self
216 }
217
218 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); }
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 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 #[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 #[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 #[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 #[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}