use std::path::Path;
use serde::{Deserialize, Serialize};
use turnframe_core::locale::Locale;
use crate::corpus::{
CaseCountExpectation, CorpusError, EvalItem, Expectations, ItemId, Setup, StateExpectation,
TurnSpec, WorkflowStateExpectation,
};
pub const MAX_TURNS: u32 = 40;
const fn ten() -> u32 {
10
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Goal {
pub id: String,
pub name: String,
pub want: String,
pub manners: Vec<String>,
#[serde(default = "ten")]
pub max_turns: u32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub locale: Option<Locale>,
#[serde(default)]
pub setup: Setup,
pub reached: Reached,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields, default)]
pub struct Reached {
pub case_state: Vec<StateExpectation>,
pub workflow_state: Vec<WorkflowStateExpectation>,
pub case_count: Vec<CaseCountExpectation>,
}
impl Reached {
#[must_use]
pub fn is_empty(&self) -> bool {
self.case_state.is_empty() && self.workflow_state.is_empty() && self.case_count.is_empty()
}
#[must_use]
pub fn expectations(&self) -> Expectations {
Expectations {
case_state: self.case_state.clone(),
workflow_state: self.workflow_state.clone(),
case_count: self.case_count.clone(),
..Expectations::default()
}
}
}
impl Goal {
pub fn validate(&self) -> Result<(), CorpusError> {
let blank = |text: &str| text.trim().is_empty();
if blank(&self.id) || blank(&self.name) || blank(&self.want) {
return Err(CorpusError::invalid(
"goal",
"a goal needs an id, a name and a want",
));
}
if self.manners.is_empty() || self.manners.iter().any(|manner| blank(manner)) {
return Err(CorpusError::invalid(
"manners",
"a goal needs at least one manner",
));
}
if !(1..=MAX_TURNS).contains(&self.max_turns) {
return Err(CorpusError::invalid("max_turns", "between 1 and 40"));
}
if self.reached.is_empty() {
return Err(CorpusError::invalid(
"reached",
"a goal with no state to reach scores every conversation reached",
));
}
self.reached.expectations().validate()?;
for seed in &self.setup.cases {
seed.validate()?;
}
Ok(())
}
#[must_use]
pub fn item(&self) -> EvalItem {
EvalItem {
id: ItemId(self.id.clone()),
name: self.name.clone(),
description: None,
tags: Vec::new(),
setup: self.setup.clone(),
before: Vec::new(),
turn: TurnSpec {
text: Some(self.want.clone()),
locale: self.locale.clone(),
..TurnSpec::default()
},
expect: self.reached.expectations(),
judge: Vec::new(),
provenance: crate::corpus::Provenance::default(),
}
}
pub fn load_dir(dir: impl AsRef<Path>) -> Result<Vec<Self>, CorpusError> {
let dir = dir.as_ref();
let read = |path: &Path, error: std::io::Error| CorpusError::Read {
path: path.to_path_buf(),
message: error.to_string(),
};
let mut goals = Vec::new();
for entry in std::fs::read_dir(dir).map_err(|error| read(dir, error))? {
let path = entry.map_err(|error| read(dir, error))?.path();
if path.extension().and_then(|extension| extension.to_str()) != Some("toml") {
continue;
}
let text = std::fs::read_to_string(&path).map_err(|error| read(&path, error))?;
let goal: Self = toml::from_str(&text).map_err(|error| CorpusError::Parse {
path: path.clone(),
message: error.to_string(),
})?;
goal.validate()?;
goals.push(goal);
}
goals.sort_by(|one, other| one.id.cmp(&other.id));
if let Some(pair) = goals.windows(2).find(|pair| pair[0].id == pair[1].id) {
return Err(CorpusError::DuplicateId {
id: ItemId(pair[0].id.clone()),
});
}
Ok(goals)
}
}