use std::fmt;
use std::str::FromStr;
use crate::explore::reducer::names_a_stat;
use crate::export::StatColumns;
use crate::view::StatDescriptor;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Comparator {
Less,
LessOrEqual,
Greater,
GreaterOrEqual,
Equal,
NotEqual,
}
impl Comparator {
const PARSE_ORDER: [Self; 6] = [
Self::LessOrEqual,
Self::Less,
Self::GreaterOrEqual,
Self::Greater,
Self::Equal,
Self::NotEqual,
];
pub fn as_str(self) -> &'static str {
match self {
Self::Less => "<",
Self::LessOrEqual => "<=",
Self::Greater => ">",
Self::GreaterOrEqual => ">=",
Self::Equal => "==",
Self::NotEqual => "!=",
}
}
pub fn compare(self, value: f64, threshold: f64) -> bool {
match self {
Self::Less => value < threshold,
Self::LessOrEqual => value <= threshold,
Self::Greater => value > threshold,
Self::GreaterOrEqual => value >= threshold,
Self::Equal => value == threshold,
Self::NotEqual => value != threshold,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Comparison {
pub comparator: Comparator,
pub threshold: f64,
}
impl Comparison {
pub fn holds(self, value: f64) -> bool {
!value.is_nan() && self.comparator.compare(value, self.threshold)
}
pub fn check(self) -> Result<(), ComparisonError> {
if self.threshold.is_finite() {
Ok(())
} else {
Err(ComparisonError::BadThreshold {
raw: self.threshold.to_string(),
})
}
}
}
impl fmt::Display for Comparison {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}{}", self.comparator.as_str(), self.threshold)
}
}
impl FromStr for Comparison {
type Err = ComparisonError;
fn from_str(raw: &str) -> Result<Self, Self::Err> {
let text = raw.trim();
let (comparator, threshold) = Comparator::PARSE_ORDER
.iter()
.find_map(|&comparator| {
text.strip_prefix(comparator.as_str())
.map(|rest| (comparator, rest.trim()))
})
.ok_or_else(|| ComparisonError::MissingComparator { raw: text.to_owned() })?;
let threshold = threshold
.parse::<f64>()
.ok()
.filter(|number| number.is_finite())
.ok_or_else(|| ComparisonError::BadThreshold {
raw: threshold.to_owned(),
})?;
Ok(Self { comparator, threshold })
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ComparisonError {
MissingComparator {
raw: String,
},
BadThreshold {
raw: String,
},
}
impl fmt::Display for ComparisonError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::MissingComparator { raw } if raw.is_empty() => {
write!(f, "missing comparator, expected <, <=, >, >=, == or !=")
}
Self::MissingComparator { raw } => write!(f, "'{raw}' does not start with <, <=, >, >=, == or !="),
Self::BadThreshold { raw } => write!(f, "threshold '{raw}' is not a finite number"),
}
}
}
impl std::error::Error for ComparisonError {}
#[derive(Debug, Clone, PartialEq)]
pub struct StopSpec {
pub column: String,
pub comparison: Comparison,
pub min_tick: u64,
}
impl StopSpec {
pub fn parse(condition: &str, min_tick: u64) -> Result<Self, StopError> {
let is_comparator = |character: char| matches!(character, '<' | '>' | '=' | '!');
let split = condition.rfind(is_comparator).map_or(condition.len(), |last| {
condition[..last].trim_end_matches(is_comparator).len()
});
let (column, comparison) = condition.split_at(split);
let column = column.trim();
if column.is_empty() {
return Err(StopError::MissingColumn {
raw: condition.to_owned(),
});
}
let comparison = comparison.parse().map_err(|source| StopError::Comparison {
raw: condition.to_owned(),
source,
})?;
Ok(Self {
column: column.to_owned(),
comparison,
min_tick,
})
}
pub fn check_threshold(&self) -> Result<(), StopError> {
if self.comparison.threshold.is_finite() {
Ok(())
} else {
Err(StopError::NonFiniteThreshold { raw: self.to_string() })
}
}
pub fn check_label(&self, stats: &[StatDescriptor]) -> Result<(), StopError> {
if names_a_stat(&self.column, stats) {
Ok(())
} else {
Err(StopError::UnknownColumn {
column: self.column.clone(),
known: stats.iter().map(|stat| stat.label.to_owned()).collect(),
})
}
}
}
impl fmt::Display for StopSpec {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let Comparison { comparator, threshold } = self.comparison;
write!(f, "{} {} {threshold}", self.column, comparator.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StopError {
MissingColumn {
raw: String,
},
Comparison {
raw: String,
source: ComparisonError,
},
NonFiniteThreshold {
raw: String,
},
UnknownColumn {
column: String,
known: Vec<String>,
},
}
impl fmt::Display for StopError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::MissingColumn { raw } => write!(f, "stop condition '{raw}' names no column"),
Self::Comparison { raw, .. } => {
write!(
f,
"invalid stop condition '{raw}', expected COLUMN COMPARATOR THRESHOLD"
)
}
Self::NonFiniteThreshold { raw } => {
write!(f, "stop condition '{raw}' has a threshold that is not a finite number")
}
Self::UnknownColumn { column, known } => {
write!(
f,
"unknown stat column '{column}', expected one of {}",
known.join(", ")
)
}
}
}
}
impl std::error::Error for StopError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Comparison { source, .. } => Some(source),
Self::MissingColumn { .. } | Self::NonFiniteThreshold { .. } | Self::UnknownColumn { .. } => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct StopCondition {
column: usize,
comparison: Comparison,
min_tick: u64,
}
impl StopCondition {
pub fn bind(spec: &StopSpec, columns: &StatColumns) -> Result<Self, StopError> {
let column = columns.resolve(&spec.column).ok_or_else(|| StopError::UnknownColumn {
column: spec.column.clone(),
known: (0..columns.len())
.map(|column| columns.name(column).to_owned())
.collect(),
})?;
Ok(Self {
column,
comparison: spec.comparison,
min_tick: spec.min_tick,
})
}
pub fn column(&self) -> usize {
self.column
}
pub fn holds(&self, tick: u64, row: &[f64]) -> bool {
tick >= self.min_tick && self.comparison.holds(row[self.column])
}
}
#[cfg(test)]
mod tests {
use super::{Comparator, Comparison, ComparisonError, StopCondition, StopError, StopSpec};
use crate::export::StatColumns;
use crate::helpers::stat;
use crate::view::StatDescriptor;
const COLOR: [u8; 4] = [0, 0, 0, 255];
fn stop(condition: &str) -> StopSpec {
StopSpec::parse(condition, 0).expect("a well-formed condition")
}
#[test]
fn a_stop_condition_parses_labels_with_spaces() {
let spec = stop("Giant Component Share >= 0.5");
assert_eq!(spec.column, "Giant Component Share");
assert_eq!(
spec.comparison,
Comparison {
comparator: Comparator::GreaterOrEqual,
threshold: 0.5
}
);
assert_eq!(spec.to_string(), "Giant Component Share >= 0.5");
let tight = stop(" Infected<=0 ");
assert_eq!((tight.column.as_str(), tight.comparison.threshold), ("Infected", 0.0));
assert_eq!(tight.to_string(), "Infected <= 0");
}
#[test]
fn a_label_holding_comparator_characters_reads_back() {
for label in ["Agents (k=3)", "R>1 cells", "a<b", "Not!", "x >= y"] {
for comparator in Comparator::PARSE_ORDER {
let spec = StopSpec {
column: label.to_owned(),
comparison: Comparison {
comparator,
threshold: -0.5,
},
min_tick: 3,
};
assert_eq!(StopSpec::parse(&spec.to_string(), 3), Ok(spec.clone()), "{spec}");
}
}
assert_eq!(stop("R>1 cells<=0.5").column, "R>1 cells");
}
#[test]
fn every_comparator_parses() {
for comparator in Comparator::PARSE_ORDER {
let spec = stop(&format!("Infected {} -3.5", comparator.as_str()));
assert_eq!(spec.comparison.comparator, comparator);
assert_eq!(spec.comparison.threshold, -3.5);
let comparison = Comparison {
comparator,
threshold: 10.0,
};
assert_eq!(comparison.to_string().parse(), Ok(comparison), "{comparison}");
}
let holds = |comparator, value| {
Comparison {
comparator,
threshold: 1.0,
}
.holds(value)
};
assert!(holds(Comparator::Less, 0.5) && !holds(Comparator::Less, 1.0));
assert!(holds(Comparator::LessOrEqual, 1.0) && !holds(Comparator::LessOrEqual, 1.5));
assert!(holds(Comparator::Greater, 1.5) && !holds(Comparator::Greater, 1.0));
assert!(holds(Comparator::GreaterOrEqual, 1.0) && !holds(Comparator::GreaterOrEqual, 0.5));
assert!(holds(Comparator::Equal, 1.0) && !holds(Comparator::Equal, 0.5));
assert!(holds(Comparator::NotEqual, 0.5) && !holds(Comparator::NotEqual, 1.0));
}
#[test]
fn a_malformed_condition_is_refused() {
assert_eq!(
StopSpec::parse("<= 0", 0),
Err(StopError::MissingColumn { raw: "<= 0".to_owned() })
);
let comparison_error = |condition: &str| match StopSpec::parse(condition, 0) {
Err(StopError::Comparison { source, .. }) => source,
other => panic!("{condition} gave {other:?}"),
};
assert!(matches!(
comparison_error("Infected 0"),
ComparisonError::MissingComparator { .. }
));
assert!(matches!(
comparison_error("Infected => 0"),
ComparisonError::MissingComparator { .. }
));
assert!(matches!(
comparison_error("Infected = 0"),
ComparisonError::MissingComparator { .. }
));
assert_eq!(
comparison_error("Infected <= many"),
ComparisonError::BadThreshold { raw: "many".to_owned() }
);
assert!(matches!(
comparison_error("Infected < NaN"),
ComparisonError::BadThreshold { .. }
));
assert_eq!(
StopSpec::parse("Infected <= x", 0).map_err(|error| error.to_string()),
Err("invalid stop condition 'Infected <= x', expected COLUMN COMPARATOR THRESHOLD".to_owned())
);
}
#[test]
fn a_missing_comparator_quotes_the_trimmed_text() {
assert_eq!(
" 0 ".parse::<Comparison>(),
Err(ComparisonError::MissingComparator { raw: "0".to_owned() })
);
let error = "".parse::<Comparison>().expect_err("no comparator");
assert_eq!(error.to_string(), "missing comparator, expected <, <=, >, >=, == or !=");
}
#[test]
fn a_threshold_that_is_not_finite_is_refused() {
for threshold in [f64::INFINITY, f64::NEG_INFINITY, f64::NAN] {
let spec = StopSpec {
column: "Infected".to_owned(),
comparison: Comparison {
comparator: Comparator::LessOrEqual,
threshold,
},
min_tick: 0,
};
assert_eq!(
spec.check_threshold(),
Err(StopError::NonFiniteThreshold { raw: spec.to_string() })
);
assert!(StopSpec::parse(&spec.to_string(), 0).is_err(), "{spec}");
}
let infinite = StopSpec {
comparison: Comparison {
comparator: Comparator::LessOrEqual,
threshold: f64::INFINITY,
},
..stop("Infected <= 0")
};
assert_eq!(
infinite.check_threshold().map_err(|error| error.to_string()),
Err("stop condition 'Infected <= inf' has a threshold that is not a finite number".to_owned())
);
assert_eq!(stop("Infected <= 0").check_threshold(), Ok(()));
}
#[test]
fn nan_never_satisfies_a_stop_condition() {
for comparator in Comparator::PARSE_ORDER {
let comparison = Comparison {
comparator,
threshold: 0.0,
};
assert!(!comparison.holds(f64::NAN), "{comparison}");
}
let columns = StatColumns::plan(&[stat("Infected", 0.0, COLOR)]);
let condition = StopCondition::bind(&stop("Infected != 5"), &columns).expect("the column exists");
assert!(!condition.holds(10, &[f64::NAN]));
assert!(condition.holds(10, &[4.0]));
}
#[test]
fn a_stop_is_not_checked_before_its_min_tick() {
let columns = StatColumns::plan(&[stat("Recovered", 0.0, COLOR), stat("Infected", 0.0, COLOR)]);
let spec = StopSpec::parse("Infected <= 0", 20).expect("a well-formed condition");
let condition = StopCondition::bind(&spec, &columns).expect("the column exists");
assert_eq!(condition.column(), 1);
assert!(!condition.holds(0, &[5.0, 0.0]));
assert!(!condition.holds(19, &[5.0, 0.0]));
assert!(condition.holds(20, &[5.0, 0.0]));
assert!(!condition.holds(25, &[5.0, 1.0]));
}
#[test]
fn a_stop_over_an_unknown_column_is_refused() {
let stats = [StatDescriptor::new("Infected", COLOR)];
assert_eq!(stop("Infected <= 0").check_label(&stats), Ok(()));
assert!(stop("Recovered <= 0").check_label(&stats).is_err());
let columns = StatColumns::plan(&[stat("Infected", 0.0, COLOR)]);
let error = StopCondition::bind(&stop("Infected.x <= 0"), &columns).expect_err("no such column");
assert_eq!(
error.to_string(),
"unknown stat column 'Infected.x', expected one of Infected"
);
}
}