Skip to main content

turnframe_eval/
understanding.rs

1//! What an item expects of each understanding task, beside what it expects of the turn.
2//!
3//! A turn-level expectation says what happened; these say which task got it right or
4//! wrong: how the message split into units, which operation and record each request
5//! reached, and which arguments were given or not. They are parsed and validated with
6//! the rest of the item, and scored against the task records a turn leaves behind.
7//!
8//! ```toml
9//! [[expect.understanding.units]]
10//! kind = "request"
11//! words = "the subject needs changing"
12//! operation = "trip.set_name"
13//! record = "trip-1"
14//! arguments = { value = { not_given = true } }
15//! ```
16
17use 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/// The units a message splits into, in message order.
30#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
31#[serde(deny_unknown_fields)]
32#[non_exhaustive]
33pub struct UnderstandingExpectation {
34    /// Every unit the message carries, in the order the user wrote them.
35    #[serde(default)]
36    pub units: Vec<UnitExpectation>,
37}
38
39/// One unit of a message, and what the tasks after segmentation must make of it.
40#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
41#[serde(deny_unknown_fields)]
42#[non_exhaustive]
43pub struct UnitExpectation {
44    /// What kind of thing the user did in these words.
45    pub kind: UnitKind,
46    /// The user's words the unit covers, exactly as written.
47    #[serde(default, skip_serializing_if = "Option::is_none")]
48    pub words: Option<String>,
49    /// The operation a request reaches.
50    #[serde(default, skip_serializing_if = "Option::is_none")]
51    pub operation: Option<String>,
52    /// The record it aims at: a seeded case id, or `new`.
53    #[serde(default, skip_serializing_if = "Option::is_none")]
54    pub record: Option<String>,
55    /// Expected arguments, by name.
56    #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
57    pub arguments: BTreeMap<String, ArgumentExpectation>,
58}
59
60/// The kinds of unit segmentation distinguishes.
61#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
62#[serde(rename_all = "snake_case")]
63#[non_exhaustive]
64pub enum UnitKind {
65    /// Something to do.
66    Request,
67    /// Something asked.
68    Question,
69    /// A condition on the whole turn.
70    Constraint,
71    /// A change to something said earlier in the same message.
72    Correction,
73    /// Withdrawing something.
74    Cancel,
75    /// An answer to the card on screen.
76    CardAnswer,
77    /// Contesting what the assistant did or said.
78    Dispute,
79    /// The value the assistant asked for.
80    ProvidesValue,
81    /// Nothing to act on.
82    Chitchat,
83}
84
85impl UnitKind {
86    /// Whether units of this kind are routed to an operation.
87    #[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/// What one argument must come out as: a value, or explicitly not given.
97#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
98#[serde(deny_unknown_fields)]
99#[non_exhaustive]
100pub struct ArgumentExpectation {
101    /// The value it must equal.
102    #[serde(default, skip_serializing_if = "Option::is_none")]
103    pub equals: Option<serde_json::Value>,
104    /// The user did not give it.
105    #[serde(default, skip_serializing_if = "std::ops::Not::not")]
106    pub not_given: bool,
107}
108
109impl UnderstandingExpectation {
110    /// Checks each unit against the turn text it is about.
111    ///
112    /// # Errors
113    ///
114    /// [`CorpusError::Invalid`] when a unit quotes words the turn does not contain,
115    /// names an operation its kind cannot reach, or states an argument both ways.
116    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/// How each understanding task did on one sample: `None` where the item states nothing
163/// that task decides.
164#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
165#[non_exhaustive]
166pub struct TaskScores {
167    /// The message split into the units expected, of the kinds expected.
168    pub segment: Option<bool>,
169    /// Each request reached the operation expected.
170    pub route: Option<bool>,
171    /// Each request aimed at the record expected.
172    pub locate: Option<bool>,
173    /// Each argument came out as expected: that value, or not given.
174    pub extract: Option<bool>,
175}
176
177impl TaskScores {
178    /// Each task's score, by the task's name.
179    #[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
190/// Words compared as a reader compares them: trimmed, without the punctuation around.
191fn normalized(words: &str) -> String {
192    words
193        .trim_matches(|c: char| c.is_whitespace() || matches!(c, '.' | ',' | ';' | ':' | '!' | '?'))
194        .to_lowercase()
195}
196
197impl UnderstandingExpectation {
198    /// Scores the understanding a turn recorded against what this item expects of
199    /// each task. A unit is paired with the understood unit quoting the same words,
200    /// else with the one at the same position.
201    #[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            // A unit may ask for several things: the act scored is the one doing the
265            // expected operation, when there is one.
266            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}