1use std::collections::BTreeMap;
18
19use serde::{Deserialize, Serialize};
20
21use turnframe_core::understanding::{
22 ActAction, ActTarget, ArgumentValue, UnderstoodAct, Unit, UnitKind as CoreUnitKind,
23};
24use turnframe_runtime::resume::CARD_UNIT;
25
26use crate::corpus::CorpusError;
27use crate::observation::Observation;
28
29#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
31#[serde(deny_unknown_fields)]
32#[non_exhaustive]
33pub struct UnderstandingExpectation {
34 #[serde(default)]
36 pub units: Vec<UnitExpectation>,
37}
38
39#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
41#[serde(deny_unknown_fields)]
42#[non_exhaustive]
43pub struct UnitExpectation {
44 pub kind: UnitKind,
46 #[serde(default, skip_serializing_if = "Option::is_none")]
48 pub words: Option<String>,
49 #[serde(default, skip_serializing_if = "Option::is_none")]
51 pub operation: Option<String>,
52 #[serde(default, skip_serializing_if = "Option::is_none")]
54 pub record: Option<String>,
55 #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
57 pub arguments: BTreeMap<String, ArgumentExpectation>,
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
62#[serde(rename_all = "snake_case")]
63#[non_exhaustive]
64pub enum UnitKind {
65 Request,
67 Question,
69 Constraint,
71 Correction,
73 Cancel,
75 CardAnswer,
77 Dispute,
79 ProvidesValue,
81 Chitchat,
83}
84
85impl UnitKind {
86 #[must_use]
88 pub const fn reaches_an_operation(self) -> bool {
89 matches!(
90 self,
91 Self::Request | Self::Correction | Self::Cancel | Self::ProvidesValue
92 )
93 }
94}
95
96#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
98#[serde(deny_unknown_fields)]
99#[non_exhaustive]
100pub struct ArgumentExpectation {
101 #[serde(default, skip_serializing_if = "Option::is_none")]
103 pub equals: Option<serde_json::Value>,
104 #[serde(default, skip_serializing_if = "std::ops::Not::not")]
106 pub not_given: bool,
107}
108
109impl UnderstandingExpectation {
110 pub fn validate(&self, text: Option<&str>) -> Result<(), CorpusError> {
117 for unit in &self.units {
118 unit.validate(text)?;
119 }
120 Ok(())
121 }
122}
123
124impl UnitExpectation {
125 fn validate(&self, text: Option<&str>) -> Result<(), CorpusError> {
126 const FIELD: &str = "expect.understanding.units";
127 let invalid = |reason: &str| CorpusError::Invalid {
128 field: FIELD.to_owned(),
129 reason: reason.to_owned(),
130 };
131 if let Some(words) = &self.words {
132 if words.trim().is_empty() {
133 return Err(invalid("`words` is empty"));
134 }
135 if !text.is_some_and(|text| text.contains(words.as_str())) {
136 return Err(invalid(&format!(
137 "`words` «{words}» is not in the turn's text"
138 )));
139 }
140 }
141 let routed = self.operation.is_some() || self.record.is_some();
142 if (routed || !self.arguments.is_empty()) && !self.kind.reaches_an_operation() {
143 return Err(invalid(&format!(
144 "a {:?} unit reaches no operation, so it has no operation, record or arguments",
145 self.kind
146 )));
147 }
148 if !self.arguments.is_empty() && self.operation.is_none() {
149 return Err(invalid("arguments belong to an operation; name it"));
150 }
151 for (name, argument) in &self.arguments {
152 if argument.not_given == argument.equals.is_some() {
153 return Err(invalid(&format!(
154 "argument `{name}` needs exactly one of `equals` and `not_given`"
155 )));
156 }
157 }
158 Ok(())
159 }
160}
161
162#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
165#[non_exhaustive]
166pub struct TaskScores {
167 pub segment: Option<bool>,
169 pub route: Option<bool>,
171 pub locate: Option<bool>,
173 pub extract: Option<bool>,
175}
176
177impl TaskScores {
178 #[must_use]
180 pub const fn by_task(&self) -> [(&'static str, Option<bool>); 4] {
181 [
182 ("segment", self.segment),
183 ("route", self.route),
184 ("locate", self.locate),
185 ("extract", self.extract),
186 ]
187 }
188}
189
190fn normalized(words: &str) -> String {
192 words
193 .trim_matches(|c: char| c.is_whitespace() || matches!(c, '.' | ',' | ';' | ':' | '!' | '?'))
194 .to_lowercase()
195}
196
197impl UnderstandingExpectation {
198 #[must_use]
202 pub fn score(&self, text: &str, observed: &Observation) -> TaskScores {
203 if self.units.is_empty() {
204 return TaskScores::default();
205 }
206 let failed = |stated: bool| stated.then_some(false);
207 let stated = |test: fn(&UnitExpectation) -> bool| self.units.iter().any(test);
208 let Some(understood) = observed.understanding.as_ref() else {
209 return TaskScores {
210 segment: Some(false),
211 route: failed(stated(|unit| unit.operation.is_some())),
212 locate: failed(stated(|unit| unit.record.is_some())),
213 extract: failed(stated(|unit| !unit.arguments.is_empty())),
214 };
215 };
216 let mut units: Vec<&Unit> = understood
217 .units
218 .iter()
219 .filter(|unit| unit.id != CARD_UNIT)
220 .collect();
221 units.sort_by_key(|unit| unit.words.start);
222 let said = |unit: &Unit| {
223 normalized(
224 text.get(unit.words.start..unit.words.end)
225 .unwrap_or_default(),
226 )
227 };
228 let kind_of = |kind: UnitKind| serde_json::to_value(kind).ok();
229 let same_kind = |expected: UnitKind, found: CoreUnitKind| {
230 kind_of(expected).is_some() && kind_of(expected) == serde_json::to_value(found).ok()
231 };
232 let segment = units.len() == self.units.len()
233 && self.units.iter().zip(&units).all(|(expected, found)| {
234 same_kind(expected.kind, found.kind)
235 && expected
236 .words
237 .as_deref()
238 .is_none_or(|words| normalized(words) == said(found))
239 });
240 let paired = |at: usize, expected: &UnitExpectation| -> Option<&Unit> {
241 match expected.words.as_deref() {
242 Some(words) => units
243 .iter()
244 .copied()
245 .find(|found| said(found) == normalized(words)),
246 None if units.len() == self.units.len() => units.get(at).copied(),
247 None => None,
248 }
249 };
250 let message_acts: Vec<&UnderstoodAct> = understood
251 .acts
252 .iter()
253 .filter(|act| act.id.unit != CARD_UNIT)
254 .collect();
255 let (mut route, mut locate, mut extract) = (None, None, None);
256 let record = |score: &mut Option<bool>, pass: bool| {
257 *score = Some(score.unwrap_or(true) && pass);
258 };
259 for (at, expected) in self.units.iter().enumerate() {
260 let does = |act: &UnderstoodAct, operation: &str| match &act.action {
261 ActAction::Apply { operation: found } => found.as_str() == operation,
262 ActAction::Start { workflow } => operation == format!("start:{workflow}"),
263 };
264 let act = paired(at, expected).and_then(|unit| {
267 let mut of_unit = message_acts
268 .iter()
269 .copied()
270 .filter(|act| act.id.unit == unit.id);
271 let first = of_unit.clone().next();
272 expected
273 .operation
274 .as_deref()
275 .and_then(|operation| of_unit.find(|act| does(act, operation)))
276 .or(first)
277 });
278 if let Some(operation) = &expected.operation {
279 let reached = act.is_some_and(|act| does(act, operation));
280 record(&mut route, reached);
281 }
282 if let Some(target) = &expected.record {
283 let aimed = act.is_some_and(|act| {
284 if target == "new" {
285 return matches!(act.target, ActTarget::New { .. });
286 }
287 let at = message_acts.iter().position(|other| other.id == act.id);
288 observed.target_resolutions.iter().any(|resolution| {
289 Some(resolution.act_index) == at
290 && resolution
291 .case_id
292 .as_ref()
293 .is_some_and(|id| id.as_str() == target)
294 })
295 });
296 record(&mut locate, aimed);
297 }
298 if !expected.arguments.is_empty() {
299 let given = act.is_some_and(|act| {
300 expected.arguments.iter().all(|(name, argument)| {
301 let found = act.arguments.get(name).map(|given| &given.value);
302 match (&argument.equals, found) {
303 (None, found) => found.is_none(),
304 (Some(value), Some(ArgumentValue::Json(found))) => found == value,
305 (Some(_), _) => false,
306 }
307 })
308 });
309 record(&mut extract, given);
310 }
311 }
312 TaskScores {
313 segment: Some(segment),
314 route,
315 locate,
316 extract,
317 }
318 }
319}
320
321#[cfg(test)]
322mod tests {
323 use super::*;
324
325 fn unit(toml: &str) -> UnitExpectation {
326 toml::from_str(toml).expect("the unit parses")
327 }
328
329 #[test]
330 fn a_request_with_a_missing_value_validates() {
331 let unit = unit(
332 r#"
333 kind = "request"
334 words = "the subject needs changing"
335 operation = "trip.set_name"
336 arguments = { value = { not_given = true } }
337 "#,
338 );
339 assert!(unit.validate(Some("the subject needs changing")).is_ok());
340 }
341
342 #[test]
343 fn words_the_turn_does_not_contain_are_refused() {
344 let unit = unit(
345 r#"kind = "question"
346words = "is it valid""#,
347 );
348 assert!(unit.validate(Some("that is wrong")).is_err());
349 }
350
351 #[test]
352 fn a_dispute_names_no_operation() {
353 let unit = unit(
354 r#"kind = "dispute"
355operation = "trip.set_name""#,
356 );
357 assert!(unit.validate(None).is_err());
358 }
359
360 #[test]
361 fn an_argument_is_given_or_not_never_both() {
362 let unit = unit(
363 r#"
364 kind = "request"
365 operation = "trip.set_name"
366 arguments = { value = { equals = "x", not_given = true } }
367 "#,
368 );
369 assert!(unit.validate(None).is_err());
370 }
371}