Skip to main content

henad_core/explore/
factor.rs

1//! Factors of a sweep, as a spec writes them and as a plan resolves them.
2//!
3//! A factor is one parameter, or the tick of one action, that a sweep varies. Each value it takes is a level. A
4//! random or Latin hypercube design samples its factors, and can draw a factor's value from a whole range instead
5//! of a list of levels.
6
7use std::fmt;
8use std::num::ParseFloatError;
9
10use crate::explore::design::DesignKind;
11use crate::explore::plan::Config;
12use crate::explore::spec::ActionSpec;
13use crate::explore::value::{ValueError, check_value, parse_value};
14use crate::params::{ParamDescriptor, ParamKind, ParamValue};
15
16/// Maximum number of levels in one factor.
17pub const MAX_LEVELS: u64 = 1 << 24;
18
19/// Relative amount by which a range may miss a whole number of steps and still end on its `max`.
20const ROUNDING: f64 = 1e-9;
21
22/// Levels of a factor, as a spec writes them.
23#[derive(Debug, Clone, PartialEq)]
24pub enum LevelSpec {
25    /// Values in the text form that [`parse_value`] accepts, or ticks for an action.
26    ///
27    /// A plan trims the spaces around each value before it reads it.
28    Values(Vec<String>),
29    /// Every value from `min` to `max` inclusive, `step` apart.
30    ///
31    /// A whole-number factor takes a step of 1 when `step` is `None`. An `F32` parameter needs a step, except in a
32    /// sampled design. A sampled design draws from the whole range.
33    Range {
34        /// Lowest value of the range.
35        min: f64,
36        /// Highest value the range can reach.
37        max: f64,
38        /// Distance between two consecutive values.
39        step: Option<f64>,
40    },
41    /// Every value of a `Bool` or `Choice` parameter.
42    All,
43}
44
45impl LevelSpec {
46    /// Reads levels as `--vary` writes them: `all`, a range `min:max:step` or `min:max`, or a list `v1,v2,...`.
47    ///
48    /// Spaces around the text, a listed value or a part of a range are ignored. The model checks the levels when
49    /// the sweep is planned.
50    ///
51    /// # Errors
52    ///
53    /// Returns [`LevelSpecError`] for a range with other than two or three parts, or a part that is not a number.
54    pub fn parse(text: &str) -> Result<Self, LevelSpecError> {
55        let text = text.trim();
56        if text == "all" {
57            return Ok(Self::All);
58        }
59        if !text.contains(':') {
60            return Ok(Self::Values(
61                text.split(',').map(|value| value.trim().to_owned()).collect(),
62            ));
63        }
64        let number = |part: &str| {
65            part.trim().parse::<f64>().map_err(|source| LevelSpecError::NotANumber {
66                part: part.to_owned(),
67                range: text.to_owned(),
68                source,
69            })
70        };
71        match text.split(':').collect::<Vec<_>>().as_slice() {
72            [min, max] => Ok(Self::Range {
73                min: number(min)?,
74                max: number(max)?,
75                step: None,
76            }),
77            [min, max, step] => Ok(Self::Range {
78                min: number(min)?,
79                max: number(max)?,
80                step: Some(number(step)?),
81            }),
82            _ => Err(LevelSpecError::BadRange { range: text.to_owned() }),
83        }
84    }
85}
86
87/// Text that cannot be parsed as a [`LevelSpec`].
88#[derive(Debug, Clone, PartialEq, Eq)]
89pub enum LevelSpecError {
90    /// A part of range `range` that is not a number.
91    NotANumber {
92        /// Part that is not a number, as written.
93        part: String,
94        /// Range the part belongs to, as written, after trimming.
95        range: String,
96        /// Reason the part cannot be parsed as an `f64`.
97        source: ParseFloatError,
98    },
99    /// A range with other than two or three parts.
100    BadRange {
101        /// Range as written, after trimming.
102        range: String,
103    },
104}
105
106impl fmt::Display for LevelSpecError {
107    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108        match self {
109            Self::NotANumber { part, range, .. } => write!(f, "'{part}' in range '{range}' is not a number"),
110            Self::BadRange { range } => write!(f, "invalid range '{range}', expected MIN:MAX:STEP or MIN:MAX"),
111        }
112    }
113}
114
115impl std::error::Error for LevelSpecError {
116    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
117        match self {
118            Self::NotANumber { source, .. } => Some(source),
119            Self::BadRange { .. } => None,
120        }
121    }
122}
123
124/// Part of a config that a factor varies, identified by the id or name that a spec uses.
125#[derive(Debug, Clone, PartialEq, Eq)]
126pub enum FactorTarget {
127    /// Parameter with this id.
128    Param(String),
129    /// Tick of the action with this [`ActionSpec::name`].
130    Action(String),
131}
132
133impl fmt::Display for FactorTarget {
134    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
135        match self {
136            Self::Param(id) => write!(f, "parameter '{id}'"),
137            Self::Action(name) => write!(f, "action '{name}'"),
138        }
139    }
140}
141
142/// A factor as a spec writes it.
143#[derive(Debug, Clone, PartialEq)]
144pub struct FactorSpec {
145    /// Parameter or action tick the factor varies.
146    pub target: FactorTarget,
147    /// Levels the factor takes, as written.
148    pub levels: LevelSpec,
149}
150
151impl FactorSpec {
152    /// Returns a factor over parameter `id`.
153    pub fn param(id: impl Into<String>, levels: LevelSpec) -> Self {
154        Self {
155            target: FactorTarget::Param(id.into()),
156            levels,
157        }
158    }
159
160    /// Returns a factor over the tick of the action named `name`.
161    pub fn action(name: impl Into<String>, levels: LevelSpec) -> Self {
162        Self {
163            target: FactorTarget::Action(name.into()),
164            levels,
165        }
166    }
167
168    /// Resolves the target against `params` and `actions`, and checks every level against the target.
169    ///
170    /// In a random or Latin hypercube design, a range with no step stands for every value from `min` to `max`.
171    /// Everywhere else the factor lists its levels.
172    ///
173    /// # Errors
174    ///
175    /// Returns [`FactorError`] for a parameter or action the sweep does not have, levels the target cannot take, or
176    /// a level the target rejects.
177    pub fn resolve(
178        &self,
179        params: &[ParamDescriptor],
180        actions: &[ActionSpec],
181        design: &DesignKind,
182    ) -> Result<Factor, FactorError> {
183        self.resolve_domain(params, actions, design.is_sampled())
184    }
185
186    /// Resolves the factor as [`Self::resolve`] does for a random or Latin hypercube design.
187    ///
188    /// # Errors
189    ///
190    /// Returns [`FactorError`] for a parameter or action the sweep does not have, levels the target cannot take, or
191    /// a level the target rejects.
192    pub fn resolve_sampled(&self, params: &[ParamDescriptor], actions: &[ActionSpec]) -> Result<Factor, FactorError> {
193        self.resolve_domain(params, actions, true)
194    }
195
196    fn resolve_domain(
197        &self,
198        params: &[ParamDescriptor],
199        actions: &[ActionSpec],
200        sampled: bool,
201    ) -> Result<Factor, FactorError> {
202        match &self.target {
203            FactorTarget::Param(id) => {
204                let Some(index) = params.iter().position(|descriptor| descriptor.id == *id) else {
205                    return Err(FactorError::UnknownParam {
206                        id: id.clone(),
207                        known: params.iter().map(|descriptor| descriptor.id).collect(),
208                    });
209                };
210                let kind = &params[index].kind;
211                let domain = match &self.levels {
212                    LevelSpec::Values(raw) => FactorDomain::Levels(parsed_levels(id, kind, raw)?),
213                    &LevelSpec::Range { min, max, step } => range_domain(id, kind, min, max, step, sampled)?,
214                    LevelSpec::All => FactorDomain::Levels(every_level(id, kind)?),
215                };
216                Ok(Factor {
217                    slot: FactorSlot::Param(index),
218                    domain,
219                })
220            }
221            FactorTarget::Action(name) => {
222                let Some(index) = actions.iter().position(|action| action.name == *name) else {
223                    return Err(FactorError::UnknownAction {
224                        name: name.clone(),
225                        known: actions.iter().map(|action| action.name.clone()).collect(),
226                    });
227                };
228                let domain = match &self.levels {
229                    LevelSpec::Values(raw) => FactorDomain::Levels(tick_levels(name, raw)?),
230                    &LevelSpec::Range { min, max, step } => {
231                        let min = whole_tick(name, min)?;
232                        let max = whole_tick(name, max)?;
233                        whole_domain(&self.target, min, max, step, sampled, FactorLevel::Tick)?
234                    }
235                    LevelSpec::All => {
236                        return Err(FactorError::AllOverNumber {
237                            target: self.target.clone(),
238                        });
239                    }
240                };
241                Ok(Factor {
242                    slot: FactorSlot::Action(index),
243                    domain,
244                })
245            }
246        }
247    }
248}
249
250/// Part of a config a resolved factor writes.
251#[derive(Debug, Clone, Copy, PartialEq, Eq)]
252pub enum FactorSlot {
253    /// Index into [`Config::params`].
254    Param(usize),
255    /// Index into [`Config::action_ticks`].
256    Action(usize),
257}
258
259/// One level of a resolved factor.
260#[derive(Debug, Clone, PartialEq)]
261pub enum FactorLevel {
262    /// Value of a parameter.
263    Param(ParamValue),
264    /// Tick of an action.
265    Tick(u64),
266}
267
268/// Values a resolved factor can take.
269#[derive(Debug, Clone, PartialEq)]
270pub enum FactorDomain {
271    /// Levels listed one by one.
272    Levels(Vec<FactorLevel>),
273    /// Every `f32` value from `min` to `max`, for a sampled design to draw from.
274    Continuous {
275        /// Lower bound of the draws.
276        min: f64,
277        /// Upper bound of the draws.
278        max: f64,
279    },
280    /// Every whole number from `min` to `max` inclusive, for a sampled design to draw from.
281    WholeNumbers {
282        /// Lowest whole number a draw can take.
283        min: u64,
284        /// Highest whole number a draw can take.
285        max: u64,
286    },
287}
288
289/// A factor resolved against a model's parameters and a spec's actions, with every level checked.
290#[derive(Debug, Clone, PartialEq)]
291pub struct Factor {
292    /// Part of a config the factor writes.
293    pub slot: FactorSlot,
294    /// Values the factor takes.
295    pub domain: FactorDomain,
296}
297
298impl Factor {
299    /// Returns the listed levels, or `None` for a factor that takes a whole range.
300    pub fn levels(&self) -> Option<&[FactorLevel]> {
301        match &self.domain {
302            FactorDomain::Levels(levels) => Some(levels),
303            FactorDomain::Continuous { .. } | FactorDomain::WholeNumbers { .. } => None,
304        }
305    }
306
307    /// Writes `level` into `config`.
308    ///
309    /// # Panics
310    ///
311    /// Panics when the kind of `level` does not match the slot, or the slot is past the end of `config`.
312    pub fn apply(&self, level: &FactorLevel, config: &mut Config) {
313        match (self.slot, level) {
314            (FactorSlot::Param(index), FactorLevel::Param(value)) => config.params[index] = value.clone(),
315            (FactorSlot::Action(index), &FactorLevel::Tick(tick)) => config.action_ticks[index] = tick,
316            (slot, level) => panic!("level {level:?} does not fit slot {slot:?}"),
317        }
318    }
319}
320
321/// A range that yields no values.
322#[derive(Debug, Clone, Copy, PartialEq)]
323pub enum RangeError {
324    /// A range with an end that is not finite.
325    NotFinite {
326        /// End the range starts from, as given.
327        min: f64,
328        /// End the range runs to, as given.
329        max: f64,
330    },
331    /// A step that is zero, negative or not finite.
332    BadStep {
333        /// Step as given.
334        step: f64,
335    },
336    /// A `min` above `max`.
337    Reversed {
338        /// End the range starts from, as given.
339        min: f64,
340        /// End the range runs to, as given.
341        max: f64,
342    },
343    /// A range of more than [`MAX_LEVELS`] values.
344    TooManyLevels,
345}
346
347impl fmt::Display for RangeError {
348    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
349        match self {
350            Self::NotFinite { min, max } => write!(f, "range ends must be finite, got {min} and {max}"),
351            Self::BadStep { step } => write!(f, "step {step} is not a positive number"),
352            Self::Reversed { min, max } => write!(f, "min {min} is greater than max {max}"),
353            Self::TooManyLevels => write!(f, "range has more than {MAX_LEVELS} values"),
354        }
355    }
356}
357
358impl std::error::Error for RangeError {}
359
360/// A factor that cannot be resolved against a model's parameters and a spec's actions.
361#[derive(Debug, Clone, PartialEq)]
362pub enum FactorError {
363    /// An id no descriptor has.
364    UnknownParam {
365        /// Id as given.
366        id: String,
367        /// Ids the descriptors have.
368        known: Vec<&'static str>,
369    },
370    /// A name that no action of the spec has.
371    UnknownAction {
372        /// Name as given.
373        name: String,
374        /// Names the spec's actions have.
375        known: Vec<String>,
376    },
377    /// A level that parameter `id` rejects.
378    Level {
379        /// Id of the parameter.
380        id: String,
381        /// Reason the descriptor rejects the level.
382        source: ValueError,
383    },
384    /// A tick of action `name` that is not a non-negative whole number.
385    BadTick {
386        /// Name of the action.
387        name: String,
388        /// Tick as listed, or a range end as `f64`'s `Display` writes it.
389        raw: String,
390    },
391    /// A range over `target` that yields no values.
392    Range {
393        /// Parameter or action tick the range varies.
394        target: FactorTarget,
395        /// Reason the range yields no values.
396        source: RangeError,
397    },
398    /// A range with no step over `F32` parameter `id`, in a design that lists levels.
399    MissingStep {
400        /// Id of the parameter.
401        id: String,
402    },
403    /// A range end or step for a whole-number `target` that is not a whole number.
404    NotWhole {
405        /// Parameter or action tick the range varies.
406        target: FactorTarget,
407        /// Range end or step as given.
408        value: f64,
409    },
410    /// `all` over a numeric `target`.
411    AllOverNumber {
412        /// Parameter or action tick `all` was given for.
413        target: FactorTarget,
414    },
415    /// A range over `Bool` or `Choice` parameter `id`.
416    RangeOverOptions {
417        /// Id of the parameter.
418        id: String,
419    },
420    /// An empty list of values for `target`.
421    NoLevels {
422        /// Parameter or action tick with no values.
423        target: FactorTarget,
424    },
425}
426
427impl fmt::Display for FactorError {
428    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
429        match self {
430            Self::UnknownParam { id, known } => {
431                write!(f, "unknown parameter '{id}', expected one of {}", known.join(", "))
432            }
433            Self::UnknownAction { name, known } if known.is_empty() => {
434                write!(f, "unknown action '{name}' (spec declares no actions)")
435            }
436            Self::UnknownAction { name, known } => {
437                write!(f, "unknown action '{name}', expected one of {}", known.join(", "))
438            }
439            Self::Level { id, .. } => write!(f, "parameter '{id}'"),
440            Self::BadTick { name, raw } => write!(f, "action '{name}' takes a non-negative integer tick, got '{raw}'"),
441            Self::Range { target, .. } => write!(f, "range of {target}"),
442            Self::MissingStep { id } => write!(
443                f,
444                "range of parameter '{id}' needs a step, except in a random or Latin hypercube design"
445            ),
446            Self::NotWhole { target, value } => write!(f, "{target} takes integers, got {value}"),
447            Self::AllOverNumber { target } => write!(f, "{target} is a number, and 'all' needs a bool or choice"),
448            Self::RangeOverOptions { id } => {
449                write!(
450                    f,
451                    "parameter '{id}' is a bool or choice, and takes values or 'all' but no range"
452                )
453            }
454            Self::NoLevels { target } => write!(f, "{target} has no values to vary over"),
455        }
456    }
457}
458
459impl std::error::Error for FactorError {
460    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
461        match self {
462            Self::Level { source, .. } => Some(source),
463            Self::Range { source, .. } => Some(source),
464            _ => None,
465        }
466    }
467}
468
469/// Returns `min + i * step` for every `i` from 0 up to the last value not past `max`.
470///
471/// Each value is computed from its index, so no rounding error accumulates across steps. A last value within a
472/// rounding error of `max`, on either side, is `max` itself.
473///
474/// # Errors
475///
476/// Returns [`RangeError`] for an end that is not finite, a step that is not positive, a `min` above `max`, or more
477/// than [`MAX_LEVELS`] values.
478pub fn inclusive_steps(min: f64, max: f64, step: f64) -> Result<Vec<f64>, RangeError> {
479    if !min.is_finite() || !max.is_finite() {
480        return Err(RangeError::NotFinite { min, max });
481    }
482    if !(step.is_finite() && step > 0.0) {
483        return Err(RangeError::BadStep { step });
484    }
485    if min > max {
486        return Err(RangeError::Reversed { min, max });
487    }
488    let span = (max - min) / step;
489    let tolerance = span.max(1.0) * ROUNDING;
490    let last = (span + tolerance).floor();
491    if last >= MAX_LEVELS as f64 {
492        return Err(RangeError::TooManyLevels);
493    }
494    let ends_on_max = span - last <= tolerance;
495    let last = last as u64;
496    Ok((0..=last)
497        .map(|i| {
498            if i == last && ends_on_max {
499                max
500            } else {
501                min + i as f64 * step
502            }
503        })
504        .collect())
505}
506
507fn parsed_levels(id: &str, kind: &ParamKind, raw: &[String]) -> Result<Vec<FactorLevel>, FactorError> {
508    if raw.is_empty() {
509        return Err(FactorError::NoLevels {
510            target: FactorTarget::Param(id.to_owned()),
511        });
512    }
513    raw.iter()
514        .map(|text| {
515            parse_value(kind, text.trim())
516                .map(FactorLevel::Param)
517                .map_err(|source| FactorError::Level {
518                    id: id.to_owned(),
519                    source,
520                })
521        })
522        .collect()
523}
524
525fn tick_levels(name: &str, raw: &[String]) -> Result<Vec<FactorLevel>, FactorError> {
526    if raw.is_empty() {
527        return Err(FactorError::NoLevels {
528            target: FactorTarget::Action(name.to_owned()),
529        });
530    }
531    raw.iter()
532        .map(|text| {
533            text.trim()
534                .parse()
535                .ok()
536                .map(FactorLevel::Tick)
537                .ok_or_else(|| FactorError::BadTick {
538                    name: name.to_owned(),
539                    raw: text.clone(),
540                })
541        })
542        .collect()
543}
544
545fn range_domain(
546    id: &str,
547    kind: &ParamKind,
548    min: f64,
549    max: f64,
550    step: Option<f64>,
551    sampled: bool,
552) -> Result<FactorDomain, FactorError> {
553    let target = FactorTarget::Param(id.to_owned());
554    match *kind {
555        ParamKind::F32 { .. } => {
556            // Every value lies between the two ends, and rounding to f32 keeps that order.
557            for end in [min, max] {
558                check_value(kind, &ParamValue::F32(end as f32)).map_err(|source| FactorError::Level {
559                    id: id.to_owned(),
560                    source,
561                })?;
562            }
563            let Some(step) = step else {
564                if !sampled {
565                    return Err(FactorError::MissingStep { id: id.to_owned() });
566                }
567                if min > max {
568                    return Err(FactorError::Range {
569                        target,
570                        source: RangeError::Reversed { min, max },
571                    });
572                }
573                return Ok(FactorDomain::Continuous { min, max });
574            };
575            let values = inclusive_steps(min, max, step).map_err(|source| FactorError::Range { target, source })?;
576            Ok(FactorDomain::Levels(
577                values
578                    .into_iter()
579                    .map(|value| FactorLevel::Param(ParamValue::F32(value as f32)))
580                    .collect(),
581            ))
582        }
583        ParamKind::U32 {
584            min: lowest,
585            max: highest,
586            ..
587        } => {
588            let min = whole_level(id, min, lowest, highest)?;
589            let max = whole_level(id, max, lowest, highest)?;
590            whole_domain(&target, min, max, step, sampled, |value| {
591                FactorLevel::Param(ParamValue::U32(value as u32))
592            })
593        }
594        ParamKind::Bool { .. } | ParamKind::Choice { .. } => Err(FactorError::RangeOverOptions { id: id.to_owned() }),
595    }
596}
597
598/// Returns the whole numbers from `min` to `max` inclusive, `step` apart, each made a level by `level`.
599///
600/// A sampled design with no step takes the whole range instead.
601fn whole_domain(
602    target: &FactorTarget,
603    min: u64,
604    max: u64,
605    step: Option<f64>,
606    sampled: bool,
607    level: impl Fn(u64) -> FactorLevel,
608) -> Result<FactorDomain, FactorError> {
609    let range_error = |source| FactorError::Range {
610        target: target.clone(),
611        source,
612    };
613    if let Some(step) = step {
614        if !(step.is_finite() && step > 0.0) {
615            return Err(range_error(RangeError::BadStep { step }));
616        }
617        if step.fract() != 0.0 {
618            return Err(FactorError::NotWhole {
619                target: target.clone(),
620                value: step,
621            });
622        }
623    }
624    if min > max {
625        return Err(range_error(RangeError::Reversed {
626            min: min as f64,
627            max: max as f64,
628        }));
629    }
630    let Some(step) = step else {
631        if sampled {
632            return Ok(FactorDomain::WholeNumbers { min, max });
633        }
634        return levels_by_step(min, max, 1, &level).map_err(range_error);
635    };
636    levels_by_step(min, max, step as u64, &level).map_err(range_error)
637}
638
639/// Returns `min + i * step` for every `i` up to the last value not past `max`, each made a level by `level`.
640fn levels_by_step(
641    min: u64,
642    max: u64,
643    step: u64,
644    level: &impl Fn(u64) -> FactorLevel,
645) -> Result<FactorDomain, RangeError> {
646    let count = (max - min) / step + 1;
647    if count > MAX_LEVELS {
648        return Err(RangeError::TooManyLevels);
649    }
650    Ok(FactorDomain::Levels(
651        (0..count).map(|i| level(min + i * step)).collect(),
652    ))
653}
654
655/// Returns `value` as a level of a `U32` parameter bounded by `lowest` and `highest`.
656fn whole_level(id: &str, value: f64, lowest: u32, highest: u32) -> Result<u64, FactorError> {
657    if !value.is_finite() || value.fract() != 0.0 {
658        return Err(FactorError::NotWhole {
659            target: FactorTarget::Param(id.to_owned()),
660            value,
661        });
662    }
663    if value < f64::from(lowest) || value > f64::from(highest) {
664        return Err(FactorError::Level {
665            id: id.to_owned(),
666            source: ValueError::OutOfRange {
667                value: value.to_string(),
668                min: lowest.to_string(),
669                max: highest.to_string(),
670            },
671        });
672    }
673    Ok(value as u64)
674}
675
676/// Returns `value` as a tick of action `name`.
677fn whole_tick(name: &str, value: f64) -> Result<u64, FactorError> {
678    // `u64::MAX` rounds up to 2^64, the first whole number a tick cannot hold.
679    if value.is_finite() && value.fract() == 0.0 && value >= 0.0 && value < u64::MAX as f64 {
680        Ok(value as u64)
681    } else {
682        Err(FactorError::BadTick {
683            name: name.to_owned(),
684            raw: value.to_string(),
685        })
686    }
687}
688
689fn every_level(id: &str, kind: &ParamKind) -> Result<Vec<FactorLevel>, FactorError> {
690    let values = match kind {
691        ParamKind::Bool { .. } => vec![ParamValue::Bool(false), ParamValue::Bool(true)],
692        ParamKind::Choice { options: [], .. } => {
693            return Err(FactorError::NoLevels {
694                target: FactorTarget::Param(id.to_owned()),
695            });
696        }
697        ParamKind::Choice { options, .. } => (0..options.len()).map(ParamValue::Choice).collect(),
698        ParamKind::F32 { .. } | ParamKind::U32 { .. } => {
699            return Err(FactorError::AllOverNumber {
700                target: FactorTarget::Param(id.to_owned()),
701            });
702        }
703    };
704    Ok(values.into_iter().map(FactorLevel::Param).collect())
705}
706
707#[cfg(test)]
708mod tests {
709    use super::{
710        Factor, FactorDomain, FactorError, FactorLevel, FactorSlot, FactorSpec, FactorTarget, LevelSpec,
711        LevelSpecError, MAX_LEVELS, RangeError, inclusive_steps,
712    };
713    use crate::explore::design::DesignKind;
714    use crate::explore::spec::ActionSpec;
715    use crate::helpers::{bool_param, choice_param, f32_param, u32_param};
716    use crate::params::{ParamDescriptor, ParamValue};
717
718    const NEIGHBORHOODS: &[&str] = &["moore", "von_neumann", "hexagonal"];
719
720    fn descriptors() -> Vec<ParamDescriptor> {
721        vec![
722            f32_param("rate", "Rate", 0.5, 0.0, 1.0, Some(0.01)),
723            u32_param("size", "Size", 10, 1, 100),
724            bool_param("wrap", "Wrap", true),
725            choice_param("neighborhood", "Neighborhood", NEIGHBORHOODS, 0),
726        ]
727    }
728
729    fn actions() -> Vec<ActionSpec> {
730        vec![ActionSpec::new("seed_outbreak", 40), ActionSpec::new("clear", 0)]
731    }
732
733    fn resolve(factor: &FactorSpec, design: &DesignKind) -> Result<Factor, FactorError> {
734        factor.resolve(&descriptors(), &actions(), design)
735    }
736
737    fn levels(id: &str, levels: LevelSpec) -> Result<Vec<ParamValue>, FactorError> {
738        let factor = resolve(&FactorSpec::param(id, levels), &DesignKind::Factorial)?;
739        Ok(factor
740            .levels()
741            .expect("a factorial lists its levels")
742            .iter()
743            .map(|level| match level {
744                FactorLevel::Param(value) => value.clone(),
745                FactorLevel::Tick(tick) => panic!("a parameter factor gave tick {tick}"),
746            })
747            .collect())
748    }
749
750    fn ticks(levels: LevelSpec) -> Result<Vec<u64>, FactorError> {
751        let factor = resolve(&FactorSpec::action("clear", levels), &DesignKind::Zip)?;
752        assert_eq!(factor.slot, FactorSlot::Action(1));
753        Ok(factor
754            .levels()
755            .expect("a zip lists its levels")
756            .iter()
757            .map(|level| match level {
758                FactorLevel::Tick(tick) => *tick,
759                FactorLevel::Param(value) => panic!("an action factor gave value {value:?}"),
760            })
761            .collect())
762    }
763
764    fn range(min: f64, max: f64, step: Option<f64>) -> LevelSpec {
765        LevelSpec::Range { min, max, step }
766    }
767
768    #[test]
769    fn a_stepped_range_is_inclusive_and_does_not_drift() {
770        let values = levels("rate", range(0.0, 1.0, Some(0.1))).expect("0 to 1 is in bounds");
771        assert_eq!(values.len(), 11, "both ends are levels");
772        for (i, value) in values.iter().enumerate() {
773            assert_eq!(*value, ParamValue::F32((i as f64 * 0.1) as f32), "level {i}");
774        }
775        assert_eq!(values[10], ParamValue::F32(1.0), "the last level is exactly the max");
776
777        assert_eq!(inclusive_steps(0.1, 0.5, 0.1).map(|values| values.len()), Ok(5));
778        let ends = inclusive_steps(0.1, 0.7, 0.2).expect("a valid range");
779        assert_eq!(ends.last(), Some(&0.7), "0.1 + 3 * 0.2 is a rounding error past 0.7");
780        let short = inclusive_steps(0.0, 0.6, 0.2).expect("a valid range");
781        assert_eq!(short.len(), 4, "0.6 / 0.2 lands a rounding error short of 3");
782        assert_eq!(short.last(), Some(&0.6));
783    }
784
785    #[test]
786    fn a_range_whose_step_overshoots_stops_before_the_end() {
787        let values = levels("rate", range(0.0, 1.0, Some(0.3))).expect("0 to 1 is in bounds");
788        let expected: Vec<ParamValue> = (0..4).map(|i| ParamValue::F32((f64::from(i) * 0.3) as f32)).collect();
789        assert_eq!(values, expected);
790    }
791
792    #[test]
793    fn a_single_point_range_has_one_level() {
794        assert_eq!(
795            levels("rate", range(0.25, 0.25, Some(0.1))),
796            Ok(vec![ParamValue::F32(0.25)])
797        );
798        assert_eq!(levels("size", range(7.0, 7.0, None)), Ok(vec![ParamValue::U32(7)]));
799    }
800
801    #[test]
802    fn an_integer_range_counts_every_value() {
803        let u32_levels = |values: &[u32]| values.iter().map(|&value| ParamValue::U32(value)).collect::<Vec<_>>();
804        assert_eq!(levels("size", range(1.0, 5.0, None)), Ok(u32_levels(&[1, 2, 3, 4, 5])));
805        assert_eq!(levels("size", range(2.0, 10.0, Some(4.0))), Ok(u32_levels(&[2, 6, 10])));
806        assert_eq!(
807            levels("size", range(1.0, 10.0, Some(3.0))),
808            Ok(u32_levels(&[1, 4, 7, 10]))
809        );
810        assert_eq!(levels("size", range(1.0, 9.0, Some(3.0))), Ok(u32_levels(&[1, 4, 7])));
811        assert_eq!(
812            levels("size", range(1.0, 100.0, None)).map(|values| values.len()),
813            Ok(100)
814        );
815        assert!(matches!(
816            levels("size", range(1.5, 5.0, None)),
817            Err(FactorError::NotWhole { .. })
818        ));
819        assert!(matches!(
820            levels("size", range(1.0, 5.0, Some(1.5))),
821            Err(FactorError::NotWhole { .. })
822        ));
823    }
824
825    #[test]
826    fn a_range_outside_the_bounds_is_refused() {
827        for spec in [
828            range(0.0, 1.5, Some(0.5)),
829            range(-0.5, 1.0, Some(0.5)),
830            range(0.0, 1.05, Some(0.1)),
831            range(0.0, f64::INFINITY, Some(0.1)),
832        ] {
833            let error = levels("rate", spec.clone()).expect_err("an end outside 0..=1");
834            assert!(matches!(error, FactorError::Level { .. }), "{spec:?} gave {error:?}");
835        }
836        for spec in [range(0.0, 5.0, None), range(1.0, 101.0, None), range(1.0, 5e9, None)] {
837            let error = levels("size", spec.clone()).expect_err("an end outside 1..=100");
838            assert!(matches!(error, FactorError::Level { .. }), "{spec:?} gave {error:?}");
839        }
840        let error = levels("rate", LevelSpec::Values(vec!["0.5".to_owned(), "2".to_owned()])).expect_err("2 > 1");
841        assert!(matches!(error, FactorError::Level { .. }), "{error:?}");
842    }
843
844    #[test]
845    fn a_zero_or_negative_step_is_refused() {
846        for step in [0.0, -0.1, f64::NAN, f64::INFINITY] {
847            let error = levels("rate", range(0.0, 1.0, Some(step))).expect_err("no usable step");
848            assert!(
849                matches!(
850                    error,
851                    FactorError::Range {
852                        source: RangeError::BadStep { .. },
853                        ..
854                    }
855                ),
856                "step {step} gave {error:?}"
857            );
858            let error = levels("size", range(1.0, 5.0, Some(step))).expect_err("no usable step");
859            assert!(
860                matches!(
861                    error,
862                    FactorError::Range {
863                        source: RangeError::BadStep { .. },
864                        ..
865                    }
866                ),
867                "step {step} gave {error:?}"
868            );
869        }
870    }
871
872    #[test]
873    fn a_reversed_or_oversized_range_is_refused() {
874        assert_eq!(
875            inclusive_steps(1.0, 0.0, 0.1),
876            Err(RangeError::Reversed { min: 1.0, max: 0.0 })
877        );
878        assert_eq!(inclusive_steps(0.0, 1.0, 1e-9), Err(RangeError::TooManyLevels));
879        assert_eq!(
880            inclusive_steps(0.0, MAX_LEVELS as f64, 1.0),
881            Err(RangeError::TooManyLevels),
882            "one value past the limit"
883        );
884        assert!(matches!(
885            levels("size", range(5.0, 1.0, None)),
886            Err(FactorError::Range {
887                source: RangeError::Reversed { .. },
888                ..
889            })
890        ));
891    }
892
893    #[test]
894    fn every_option_of_a_choice_is_a_level() {
895        let choices = levels("neighborhood", LevelSpec::All).expect("a choice has options");
896        assert_eq!(
897            choices,
898            [ParamValue::Choice(0), ParamValue::Choice(1), ParamValue::Choice(2)]
899        );
900        let flags = levels("wrap", LevelSpec::All).expect("a bool has two values");
901        assert_eq!(flags, [ParamValue::Bool(false), ParamValue::Bool(true)]);
902        let named = levels(
903            "neighborhood",
904            LevelSpec::Values(vec!["hexagonal".to_owned(), "0".to_owned()]),
905        );
906        assert_eq!(
907            named,
908            Ok(vec![ParamValue::Choice(2), ParamValue::Choice(0)]),
909            "by name or index"
910        );
911    }
912
913    #[test]
914    fn levels_a_kind_cannot_take_are_refused() {
915        assert!(matches!(
916            levels("rate", range(0.0, 1.0, None)),
917            Err(FactorError::MissingStep { .. })
918        ));
919        assert!(matches!(
920            levels("rate", LevelSpec::All),
921            Err(FactorError::AllOverNumber { .. })
922        ));
923        assert!(matches!(
924            levels("wrap", range(0.0, 1.0, Some(1.0))),
925            Err(FactorError::RangeOverOptions { .. })
926        ));
927        assert!(matches!(
928            levels("rate", LevelSpec::Values(Vec::new())),
929            Err(FactorError::NoLevels { .. })
930        ));
931        let error = levels("speed", LevelSpec::All).expect_err("no such parameter");
932        assert_eq!(
933            error.to_string(),
934            "unknown parameter 'speed', expected one of rate, size, wrap, neighborhood"
935        );
936    }
937
938    #[test]
939    fn a_resolved_factor_names_its_slot() {
940        let factor = resolve(&FactorSpec::param("wrap", LevelSpec::All), &DesignKind::Factorial).expect("a bool");
941        assert_eq!(factor.slot, FactorSlot::Param(2));
942    }
943
944    #[test]
945    fn level_text_parses_lists_ranges_and_all() {
946        assert_eq!(LevelSpec::parse("0.1:0.5:0.1"), Ok(range(0.1, 0.5, Some(0.1))));
947        assert_eq!(LevelSpec::parse("16:64"), Ok(range(16.0, 64.0, None)));
948        assert_eq!(LevelSpec::parse(" -1 : 1 : 0.5"), Ok(range(-1.0, 1.0, Some(0.5))));
949        assert_eq!(
950            LevelSpec::parse("Random,Geometric"),
951            Ok(LevelSpec::Values(vec!["Random".to_owned(), "Geometric".to_owned()]))
952        );
953        assert_eq!(LevelSpec::parse("0.05"), Ok(LevelSpec::Values(vec!["0.05".to_owned()])));
954        assert_eq!(LevelSpec::parse("all"), Ok(LevelSpec::All));
955        assert!(matches!(
956            LevelSpec::parse("0.1:x:0.1"),
957            Err(LevelSpecError::NotANumber { part, .. }) if part == "x"
958        ));
959        assert_eq!(
960            LevelSpec::parse("0:1:0.1:2").map_err(|error| error.to_string()),
961            Err("invalid range '0:1:0.1:2', expected MIN:MAX:STEP or MIN:MAX".to_owned())
962        );
963    }
964
965    /// A listed value drops its surrounding spaces, as a range and a tick do.
966    #[test]
967    fn spaces_around_a_listed_value_are_ignored() {
968        assert_eq!(
969            LevelSpec::parse(" moore, von_neumann "),
970            Ok(LevelSpec::Values(vec!["moore".to_owned(), "von_neumann".to_owned()]))
971        );
972        assert_eq!(LevelSpec::parse(" all "), Ok(LevelSpec::All));
973        let spaced = |raw: &[&str]| LevelSpec::Values(raw.iter().map(|&text| text.to_owned()).collect());
974        assert_eq!(
975            levels("neighborhood", spaced(&["moore", " hexagonal "])),
976            Ok(vec![ParamValue::Choice(0), ParamValue::Choice(2)]),
977            "a spec file's values are trimmed when planned"
978        );
979        assert_eq!(
980            levels("rate", spaced(&[" 0.1", "0.2 "])),
981            Ok(vec![ParamValue::F32(0.1), ParamValue::F32(0.2)])
982        );
983    }
984
985    #[test]
986    fn an_action_factor_takes_ticks() {
987        let values = |raw: &[&str]| LevelSpec::Values(raw.iter().map(|&text| text.to_owned()).collect());
988        assert_eq!(ticks(values(&["200", "400", " 600"])), Ok(vec![200, 400, 600]));
989        assert_eq!(ticks(range(0.0, 400.0, Some(200.0))), Ok(vec![0, 200, 400]));
990        assert_eq!(ticks(range(3.0, 5.0, None)), Ok(vec![3, 4, 5]));
991        for bad in [
992            values(&["-1"]),
993            values(&["1.5"]),
994            values(&["soon"]),
995            range(-2.0, 4.0, None),
996        ] {
997            assert!(
998                matches!(ticks(bad.clone()), Err(FactorError::BadTick { .. })),
999                "{bad:?}"
1000            );
1001        }
1002        assert!(matches!(
1003            ticks(range(0.0, 4.0, Some(0.5))),
1004            Err(FactorError::NotWhole { .. })
1005        ));
1006        assert!(matches!(ticks(LevelSpec::All), Err(FactorError::AllOverNumber { .. })));
1007        let error =
1008            resolve(&FactorSpec::action("wave", LevelSpec::All), &DesignKind::Factorial).expect_err("no such action");
1009        assert_eq!(
1010            error.to_string(),
1011            "unknown action 'wave', expected one of seed_outbreak, clear"
1012        );
1013        let error = FactorSpec::action("wave", range(0.0, 1.0, None))
1014            .resolve(&descriptors(), &[], &DesignKind::Factorial)
1015            .expect_err("no actions");
1016        assert_eq!(error.to_string(), "unknown action 'wave' (spec declares no actions)");
1017        assert_eq!(FactorTarget::Action("wave".to_owned()).to_string(), "action 'wave'");
1018    }
1019
1020    #[test]
1021    fn a_sampled_design_takes_a_whole_range_without_a_step() {
1022        let sampled = DesignKind::LatinHypercube { samples: 10 };
1023        let domain = |factor: FactorSpec| resolve(&factor, &sampled).map(|factor| factor.domain);
1024        assert_eq!(
1025            domain(FactorSpec::param("rate", range(0.25, 0.75, None))),
1026            Ok(FactorDomain::Continuous { min: 0.25, max: 0.75 })
1027        );
1028        assert_eq!(
1029            domain(FactorSpec::param("size", range(1.0, 100.0, None))),
1030            Ok(FactorDomain::WholeNumbers { min: 1, max: 100 })
1031        );
1032        assert_eq!(
1033            domain(FactorSpec::action("seed_outbreak", range(0.0, 1e12, None))),
1034            Ok(FactorDomain::WholeNumbers {
1035                min: 0,
1036                max: 1_000_000_000_000
1037            }),
1038            "a whole range is not listed, so it can pass the level limit"
1039        );
1040        let stepped = domain(FactorSpec::param("size", range(1.0, 9.0, Some(4.0)))).expect("a stepped range");
1041        assert_eq!(
1042            stepped,
1043            FactorDomain::Levels(vec![
1044                FactorLevel::Param(ParamValue::U32(1)),
1045                FactorLevel::Param(ParamValue::U32(5)),
1046                FactorLevel::Param(ParamValue::U32(9)),
1047            ]),
1048            "a stepped range stays a list"
1049        );
1050        assert!(matches!(
1051            domain(FactorSpec::param("rate", range(0.75, 0.25, None))),
1052            Err(FactorError::Range {
1053                source: RangeError::Reversed { .. },
1054                ..
1055            })
1056        ));
1057        assert!(matches!(
1058            domain(FactorSpec::param("rate", range(0.5, 1.5, None))),
1059            Err(FactorError::Level { .. })
1060        ));
1061    }
1062}