use std::path::{Path, PathBuf};
use serde::{Deserialize, Serialize};
use crate::corpus::{ItemId, Tag};
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields, default)]
pub struct EvalConfig {
pub execution: ExecutionConfig,
pub judging: JudgingConfig,
pub selection: SelectionConfig,
}
impl EvalConfig {
#[must_use]
pub fn single() -> Self {
Self::default()
}
#[must_use]
pub const fn with_samples_per_item(mut self, samples: u32) -> Self {
self.execution.samples_per_item = samples;
self
}
#[must_use]
pub const fn recording_answers(mut self) -> Self {
self.execution.record_answers = true;
self
}
#[must_use]
pub const fn with_votes_per_sample(mut self, votes: u32) -> Self {
self.judging.votes_per_sample = votes;
self
}
#[must_use]
pub const fn with_sample_concurrency(mut self, samples: u32) -> Self {
self.execution.sample_concurrency = samples;
self
}
#[must_use]
pub const fn with_max_observed_events(mut self, events: usize) -> Self {
self.execution.max_observed_events = Some(events);
self
}
#[must_use]
pub fn including_tag(mut self, tag: impl Into<Tag>) -> Self {
self.selection.include_tags.push(tag.into());
self
}
#[must_use]
pub fn excluding_tag(mut self, tag: impl Into<Tag>) -> Self {
self.selection.exclude_tags.push(tag.into());
self
}
pub fn from_toml_str(source: &str) -> Result<Self, ConfigError> {
toml::from_str(source).map_err(|error| ConfigError::Parse {
message: error.to_string(),
})
}
pub fn load(path: impl AsRef<Path>) -> Result<Self, ConfigError> {
let path = path.as_ref();
let source = std::fs::read_to_string(path).map_err(|error| ConfigError::Read {
path: path.to_path_buf(),
message: error.to_string(),
})?;
Self::from_toml_str(&source)
}
pub const fn validate(&self) -> Result<(), ConfigError> {
if self.execution.samples_per_item == 0 {
return Err(ConfigError::ZeroSamples);
}
if self.judging.votes_per_sample == 0 {
return Err(ConfigError::ZeroVotes);
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields, default)]
#[non_exhaustive]
pub struct ExecutionConfig {
pub samples_per_item: u32,
pub stop_after_failures: Option<u32>,
pub sample_concurrency: u32,
pub max_observed_events: Option<usize>,
pub record_answers: bool,
}
impl Default for ExecutionConfig {
fn default() -> Self {
Self {
samples_per_item: 1,
stop_after_failures: None,
sample_concurrency: 1,
max_observed_events: None,
record_answers: false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields, default)]
pub struct JudgingConfig {
pub votes_per_sample: u32,
pub judge_every_sample: bool,
}
impl Default for JudgingConfig {
fn default() -> Self {
Self {
votes_per_sample: 1,
judge_every_sample: false,
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields, default)]
pub struct SelectionConfig {
pub include_tags: Vec<Tag>,
pub exclude_tags: Vec<Tag>,
pub items: Vec<ItemId>,
}
impl SelectionConfig {
#[must_use]
pub fn is_empty(&self) -> bool {
self.include_tags.is_empty() && self.exclude_tags.is_empty() && self.items.is_empty()
}
#[must_use]
pub fn selects(&self, id: &ItemId, tags: &[Tag]) -> bool {
if !self.include_tags.is_empty() && !self.include_tags.iter().any(|t| tags.contains(t)) {
return false;
}
if self.exclude_tags.iter().any(|t| tags.contains(t)) {
return false;
}
self.items.is_empty() || self.items.contains(id)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum ConfigError {
#[error("evaluation configuration at {path} could not be read: {message}")]
Read {
path: PathBuf,
message: String,
},
#[error("evaluation configuration could not be parsed: {message}")]
Parse {
message: String,
},
#[error("samples_per_item must be at least 1: an item that never runs has no result")]
ZeroSamples,
#[error("votes_per_sample must be at least 1: a judge with no votes has no verdict")]
ZeroVotes,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_specification_snippet_parses() {
let config = EvalConfig::from_toml_str(
"[execution]\nsamples_per_item = 10\n\n[judging]\nvotes_per_sample = 3\n",
)
.unwrap();
assert_eq!(config.execution.samples_per_item, 10);
assert_eq!(config.judging.votes_per_sample, 3);
}
#[test]
fn an_unknown_key_is_refused_rather_than_ignored() {
let error =
EvalConfig::from_toml_str("[execution]\nsamples = 10\n").expect_err("unknown key");
assert!(matches!(error, ConfigError::Parse { .. }), "{error}");
}
#[test]
fn votes_cannot_stand_in_for_samples() {
assert!(matches!(
EvalConfig::default().with_samples_per_item(0).validate(),
Err(ConfigError::ZeroSamples)
));
assert!(matches!(
EvalConfig::default().with_votes_per_sample(0).validate(),
Err(ConfigError::ZeroVotes)
));
}
#[test]
fn selection_filters_by_tag_then_by_identifier() {
let selection = SelectionConfig {
include_tags: vec![Tag::new("trip")],
exclude_tags: vec![Tag::new("slow")],
items: Vec::new(),
};
assert!(selection.selects(&ItemId::new("a"), &[Tag::new("trip")]));
assert!(!selection.selects(&ItemId::new("a"), &[Tag::new("traveler")]));
assert!(!selection.selects(&ItemId::new("a"), &[Tag::new("trip"), Tag::new("slow")]));
}
}