1use serde::{Deserialize, Serialize};
7
8use super::error::EvalError;
9use super::types::ParameterKind;
10
11#[derive(Debug, Clone, Serialize, Deserialize)]
35#[serde(try_from = "RawParameterRange")]
36pub struct ParameterRange {
37 kind: ParameterKind,
38 min: f64,
39 max: f64,
40 step: Option<f64>,
41 default: f64,
42}
43
44#[derive(Deserialize)]
51struct RawParameterRange {
52 kind: ParameterKind,
53 min: f64,
54 max: f64,
55 step: Option<f64>,
56 default: f64,
57}
58
59impl TryFrom<RawParameterRange> for ParameterRange {
60 type Error = EvalError;
61
62 fn try_from(raw: RawParameterRange) -> Result<Self, Self::Error> {
63 Self::new(raw.kind, raw.min, raw.max, raw.step, raw.default)
64 }
65}
66
67impl ParameterRange {
68 const DEFAULT_STEP_DIVISIONS: f64 = 20.0;
74
75 pub fn new(
108 kind: ParameterKind,
109 min: f64,
110 max: f64,
111 step: Option<f64>,
112 default: f64,
113 ) -> Result<Self, EvalError> {
114 if !min.is_finite() || !max.is_finite() || min >= max {
115 return Err(EvalError::InvalidRange { min, max });
116 }
117 if !default.is_finite() || default < min || default > max {
118 return Err(EvalError::DefaultOutOfRange { default, min, max });
119 }
120 Ok(Self {
121 kind,
122 min,
123 max,
124 step,
125 default,
126 })
127 }
128
129 #[must_use]
131 pub fn kind(&self) -> ParameterKind {
132 self.kind
133 }
134
135 #[must_use]
137 pub fn min(&self) -> f64 {
138 self.min
139 }
140
141 #[must_use]
143 pub fn max(&self) -> f64 {
144 self.max
145 }
146
147 #[must_use]
149 pub fn step(&self) -> Option<f64> {
150 self.step
151 }
152
153 #[must_use]
173 pub fn effective_step(&self) -> f64 {
174 self.step
175 .unwrap_or_else(|| (self.max - self.min) / Self::DEFAULT_STEP_DIVISIONS)
176 }
177
178 #[must_use]
182 pub fn default_value(&self) -> f64 {
183 self.default
184 }
185
186 #[must_use]
202 pub fn step_count(&self) -> Option<usize> {
203 let step = self.step?;
204 if step <= 0.0 {
205 return None;
206 }
207 #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
208 Some(((self.max - self.min) / step).floor() as usize + 1)
209 }
210
211 #[must_use]
223 pub fn clamp(&self, value: f64) -> f64 {
224 value.clamp(self.min, self.max)
225 }
226
227 #[must_use]
239 pub fn contains(&self, value: f64) -> bool {
240 (self.min..=self.max).contains(&value)
241 }
242
243 #[must_use]
248 pub fn quantize(&self, value: f64) -> f64 {
249 if let Some(step) = self.step
250 && step > 0.0
251 {
252 let quantized = self.min + ((value - self.min) / step).round() * step;
253 return self.clamp((quantized * 100.0).round() / 100.0);
254 }
255 value
256 }
257}
258
259#[derive(Debug, Clone, Serialize, Deserialize)]
278#[serde(default)]
279pub struct SearchSpace {
280 pub parameters: Vec<ParameterRange>,
282}
283
284impl Default for SearchSpace {
285 fn default() -> Self {
286 Self {
287 parameters: vec![
288 ParameterRange::new(ParameterKind::Temperature, 0.0, 1.0, Some(0.1), 0.7)
289 .expect("default Temperature range is valid"),
290 ParameterRange::new(ParameterKind::TopP, 0.1, 1.0, Some(0.05), 0.9)
291 .expect("default TopP range is valid"),
292 ParameterRange::new(ParameterKind::TopK, 1.0, 100.0, Some(5.0), 40.0)
293 .expect("default TopK range is valid"),
294 ParameterRange::new(ParameterKind::FrequencyPenalty, -2.0, 2.0, Some(0.2), 0.0)
295 .expect("default FrequencyPenalty range is valid"),
296 ParameterRange::new(ParameterKind::PresencePenalty, -2.0, 2.0, Some(0.2), 0.0)
297 .expect("default PresencePenalty range is valid"),
298 ],
299 }
300 }
301}
302
303impl SearchSpace {
304 #[must_use]
321 pub fn range_for(&self, kind: ParameterKind) -> Option<&ParameterRange> {
322 self.parameters.iter().find(|r| r.kind() == kind)
323 }
324
325 #[must_use]
343 pub fn grid_size(&self) -> usize {
344 self.parameters
345 .iter()
346 .filter_map(ParameterRange::step_count)
347 .sum()
348 }
349}
350
351#[cfg(test)]
352mod tests {
353 use super::*;
354 use std::assert_matches;
355
356 fn make_range(
357 kind: ParameterKind,
358 min: f64,
359 max: f64,
360 step: Option<f64>,
361 default: f64,
362 ) -> ParameterRange {
363 ParameterRange::new(kind, min, max, step, default).unwrap()
364 }
365
366 #[test]
367 fn new_valid_range() {
368 let r = make_range(ParameterKind::Temperature, 0.0, 1.0, Some(0.5), 0.5);
369 assert_eq!(r.kind(), ParameterKind::Temperature);
370 assert!((r.min() - 0.0).abs() < f64::EPSILON);
371 assert!((r.max() - 1.0).abs() < f64::EPSILON);
372 assert!((r.default_value() - 0.5).abs() < f64::EPSILON);
373 assert_eq!(r.step(), Some(0.5));
374 }
375
376 #[test]
377 fn new_invalid_range_min_ge_max() {
378 assert_matches!(
379 ParameterRange::new(ParameterKind::Temperature, 1.0, 0.0, None, 0.5),
380 Err(EvalError::InvalidRange { .. })
381 );
382 assert_matches!(
384 ParameterRange::new(ParameterKind::Temperature, 0.5, 0.5, None, 0.5),
385 Err(EvalError::InvalidRange { .. })
386 );
387 }
388
389 #[test]
390 fn new_invalid_range_nonfinite_bounds() {
391 assert_matches!(
392 ParameterRange::new(ParameterKind::Temperature, f64::NAN, 1.0, None, 0.5),
393 Err(EvalError::InvalidRange { .. })
394 );
395 assert_matches!(
396 ParameterRange::new(ParameterKind::Temperature, 0.0, f64::INFINITY, None, 0.5),
397 Err(EvalError::InvalidRange { .. })
398 );
399 }
400
401 #[test]
402 fn new_invalid_default_out_of_range() {
403 assert_matches!(
404 ParameterRange::new(ParameterKind::Temperature, 0.0, 1.0, None, 2.0),
405 Err(EvalError::DefaultOutOfRange { .. })
406 );
407 assert_matches!(
408 ParameterRange::new(ParameterKind::Temperature, 0.0, 1.0, None, -0.1),
409 Err(EvalError::DefaultOutOfRange { .. })
410 );
411 }
412
413 #[test]
414 fn step_count_with_step() {
415 let r = make_range(ParameterKind::Temperature, 0.0, 1.0, Some(0.5), 0.5);
416 assert_eq!(r.step_count(), Some(3)); }
418
419 #[test]
420 fn step_count_no_step() {
421 let r = make_range(ParameterKind::Temperature, 0.0, 1.0, None, 0.5);
422 assert_eq!(r.step_count(), None);
423 }
424
425 #[test]
426 fn step_count_zero_step() {
427 let mut r = make_range(ParameterKind::Temperature, 0.0, 1.0, None, 0.5);
429 r.step = Some(0.0);
430 assert_eq!(r.step_count(), None);
431 }
432
433 #[test]
434 fn clamp_below_min() {
435 let r = make_range(ParameterKind::TopP, 0.1, 1.0, Some(0.1), 0.9);
436 assert!((r.clamp(-1.0) - 0.1).abs() < f64::EPSILON);
437 }
438
439 #[test]
440 fn clamp_above_max() {
441 let r = make_range(ParameterKind::TopP, 0.1, 1.0, Some(0.1), 0.9);
442 assert!((r.clamp(2.0) - 1.0).abs() < f64::EPSILON);
443 }
444
445 #[test]
446 fn clamp_within_range() {
447 let r = make_range(ParameterKind::Temperature, 0.0, 2.0, Some(0.1), 0.7);
448 assert!((r.clamp(1.0) - 1.0).abs() < f64::EPSILON);
449 }
450
451 #[test]
452 fn contains_within_range() {
453 let r = make_range(ParameterKind::Temperature, 0.0, 2.0, Some(0.1), 0.7);
454 assert!(r.contains(1.0));
455 assert!(r.contains(0.0));
456 assert!(r.contains(2.0));
457 assert!(!r.contains(-0.1));
458 assert!(!r.contains(2.1));
459 }
460
461 #[test]
462 fn quantize_snaps_to_nearest_step() {
463 let r = make_range(ParameterKind::Temperature, 0.0, 2.0, Some(0.1), 0.7);
464 let q = r.quantize(0.73);
465 assert!((q - 0.7).abs() < 1e-10, "expected 0.7, got {q}");
466 }
467
468 #[test]
469 fn quantize_no_step_returns_value_unchanged() {
470 let r = make_range(ParameterKind::Temperature, 0.0, 2.0, None, 0.7);
471 assert!((r.quantize(1.234) - 1.234).abs() < f64::EPSILON);
472 }
473
474 #[test]
475 fn quantize_clamps_result() {
476 let r = make_range(ParameterKind::Temperature, 0.0, 1.0, Some(0.1), 0.5);
477 let q = r.quantize(100.0);
478 assert!(q <= 1.0, "quantize must clamp to max");
479 }
480
481 #[test]
482 fn quantize_avoids_fp_accumulation() {
483 let r = make_range(ParameterKind::Temperature, 0.0, 2.0, Some(0.1), 0.7);
484 let accumulated = 0.1_f64 * 7.0;
485 let q = r.quantize(accumulated);
486 assert!(
487 (q - 0.7).abs() < 1e-10,
488 "expected 0.7, got {q} (accumulated={accumulated})"
489 );
490 }
491
492 #[test]
493 fn default_search_space_has_five_parameters() {
494 let space = SearchSpace::default();
495 assert_eq!(space.parameters.len(), 5);
496 }
497
498 #[test]
499 fn default_grid_size_is_reasonable() {
500 let space = SearchSpace::default();
501 let size = space.grid_size();
502 assert!(size > 0);
504 assert!(size < 200);
505 }
506
507 #[test]
508 fn range_for_finds_temperature() {
509 let space = SearchSpace::default();
510 let range = space.range_for(ParameterKind::Temperature);
511 assert!(range.is_some());
512 assert!((range.unwrap().default_value() - 0.7).abs() < f64::EPSILON);
513 }
514
515 #[test]
516 fn range_for_missing_returns_none() {
517 let space = SearchSpace::default();
518 let range = space.range_for(ParameterKind::RetrievalTopK);
519 assert!(range.is_none());
520 }
521
522 #[test]
523 fn grid_size_empty_space_is_zero() {
524 let space = SearchSpace { parameters: vec![] };
525 assert_eq!(space.grid_size(), 0);
526 }
527
528 #[test]
529 fn quantize_with_nonzero_min_anchors_to_min() {
530 let r = make_range(ParameterKind::TopK, 1.0, 100.0, Some(5.0), 40.0);
531 let q = r.quantize(6.0);
532 assert!(
533 (q - 6.0).abs() < 1e-10,
534 "expected 6.0 (min-anchored grid), got {q}"
535 );
536 let q2 = r.quantize(3.0);
537 assert!((q2 - 1.0).abs() < 1e-10, "expected 1.0, got {q2}");
538 }
539
540 #[test]
541 fn quantize_negative_step_returns_unchanged() {
542 let mut r = make_range(ParameterKind::Temperature, 0.0, 2.0, None, 0.7);
543 r.step = Some(-0.1);
544 assert!((r.quantize(0.75) - 0.75).abs() < f64::EPSILON);
545 }
546
547 #[test]
548 fn parameter_range_is_valid_for_default() {
549 for r in &SearchSpace::default().parameters {
550 assert!(
552 r.min() < r.max(),
553 "default range {:?} has min >= max",
554 r.kind()
555 );
556 }
557 }
558
559 #[test]
560 fn deserialize_rejects_inverted_range() {
561 let json = r#"{"kind":"temperature","min":2.0,"max":0.0,"step":null,"default":1.0}"#;
562 let result: Result<ParameterRange, _> = serde_json::from_str(json);
563 assert!(result.is_err(), "min > max must fail deserialization");
564 }
565
566 #[test]
567 fn deserialize_rejects_default_out_of_range() {
568 let json = r#"{"kind":"temperature","min":0.0,"max":1.0,"step":null,"default":5.0}"#;
569 let result: Result<ParameterRange, _> = serde_json::from_str(json);
570 assert!(
571 result.is_err(),
572 "default outside [min, max] must fail deserialization"
573 );
574 }
575
576 #[test]
577 fn deserialize_rejects_nonfinite_bounds() {
578 let toml_src = "kind = \"temperature\"\nmin = nan\nmax = 1.0\nstep = 0.1\ndefault = 0.5\n";
581 let result: Result<ParameterRange, _> = toml::from_str(toml_src);
582 assert!(result.is_err(), "non-finite min must fail deserialization");
583 }
584
585 #[test]
586 fn deserialize_accepts_valid_range() {
587 let json = r#"{"kind":"temperature","min":0.0,"max":1.0,"step":0.1,"default":0.5}"#;
588 let r: ParameterRange = serde_json::from_str(json).unwrap();
589 assert!((r.min() - 0.0).abs() < f64::EPSILON);
590 assert!((r.max() - 1.0).abs() < f64::EPSILON);
591 assert!((r.default_value() - 0.5).abs() < f64::EPSILON);
592 }
593
594 #[test]
595 fn deserialize_search_space_rejects_invalid_range() {
596 let json = r#"{"parameters":[{"kind":"temperature","min":1.0,"max":0.0,"step":null,"default":0.5}]}"#;
597 let result: Result<SearchSpace, _> = serde_json::from_str(json);
598 assert!(
599 result.is_err(),
600 "SearchSpace must reject an inverted range in any member"
601 );
602 }
603
604 #[test]
605 fn roundtrip_serialize_deserialize_preserves_range() {
606 let r = make_range(ParameterKind::TopP, 0.1, 1.0, Some(0.05), 0.9);
607 let json = serde_json::to_string(&r).unwrap();
608 let r2: ParameterRange = serde_json::from_str(&json).unwrap();
609 assert_eq!(r.kind(), r2.kind());
610 assert!((r.min() - r2.min()).abs() < f64::EPSILON);
611 assert!((r.max() - r2.max()).abs() < f64::EPSILON);
612 assert!((r.default_value() - r2.default_value()).abs() < f64::EPSILON);
613 }
614}