Skip to main content

turnframe_understand/tasks/
cross_check.rs

1//! `cross_check`: whether what was understood of a message says what the message says.
2
3use std::fmt::Write as _;
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7use turnframe_provider::request::Message;
8use turnframe_tasks::{ModelTask, StructuralError, TaskKind};
9
10use crate::input::UnderstandingInput;
11use crate::render;
12use crate::schema::{any_of, array, object, one_of, span, variant};
13use crate::tasks::{check_one_of, check_span};
14use crate::words::Span;
15
16const BUILT_IN: &str = include_str!("../../prompts/understand/cross_check.md");
17
18/// The whole-turn check, over one turn.
19#[derive(Debug, Clone, Copy)]
20pub struct CrossCheck<'a> {
21    turn: &'a UnderstandingInput,
22}
23
24impl<'a> CrossCheck<'a> {
25    /// The task for `turn`.
26    #[must_use]
27    pub const fn new(turn: &'a UnderstandingInput) -> Self {
28        Self { turn }
29    }
30}
31
32/// One act as the check is shown it.
33#[derive(Debug, Clone)]
34pub struct ShownAct {
35    /// Its id, `u1.a1`.
36    pub id: String,
37    /// Its operation, record and values, each value with the words it came from.
38    pub line: String,
39    /// Its argument names.
40    pub arguments: Vec<String>,
41}
42
43/// What the check is shown.
44#[derive(Debug, Clone, Default)]
45pub struct CrossCheckInput {
46    /// Every act understood.
47    pub acts: Vec<ShownAct>,
48    /// Every question, by its words.
49    pub questions: Vec<String>,
50    /// Every constraint, by its words.
51    pub constraints: Vec<String>,
52    /// Words read as nothing to act on.
53    pub unread: Vec<String>,
54    /// Words already read as an act, a question, a constraint, small talk or a value.
55    pub held: Vec<Span>,
56}
57
58/// One place the reading does not say what the message says.
59#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
60#[serde(tag = "kind", rename_all = "snake_case")]
61pub enum Finding {
62    /// Words asking for something nothing holds.
63    Missing {
64        /// The words.
65        words: Span,
66    },
67    /// A value the message does not give, and the words it should come from.
68    WrongValue {
69        /// The act.
70        act: String,
71        /// Its argument.
72        argument: String,
73        /// The words the value should come from.
74        words: Span,
75    },
76    /// A record the message does not mean, and the words naming the one it does.
77    WrongRecord {
78        /// The act.
79        act: String,
80        /// The words naming the record meant.
81        words: Span,
82    },
83    /// An act the message does not ask for.
84    NotAsked {
85        /// The act.
86        act: String,
87    },
88}
89
90/// The check's answer.
91#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
92pub struct CrossChecked {
93    /// Empty when the reading says what the message says.
94    pub findings: Vec<Finding>,
95}
96
97impl ModelTask for CrossCheck<'_> {
98    type Input = CrossCheckInput;
99    type Output = CrossChecked;
100
101    fn kind(&self) -> TaskKind {
102        TaskKind::CrossCheck
103    }
104
105    fn prompt_name(&self) -> &str {
106        "understand.cross_check"
107    }
108
109    fn instructions(&self) -> &str {
110        BUILT_IN
111    }
112
113    fn schema(&self, input: &CrossCheckInput) -> Value {
114        let acts: Vec<String> = input.acts.iter().map(|act| act.id.clone()).collect();
115        let mut arguments: Vec<String> = input
116            .acts
117            .iter()
118            .flat_map(|act| act.arguments.clone())
119            .collect();
120        arguments.sort();
121        arguments.dedup();
122        let mut kinds = vec![variant("missing", vec![("words", span())])];
123        if !acts.is_empty() {
124            if !arguments.is_empty() {
125                kinds.push(variant(
126                    "wrong_value",
127                    vec![
128                        ("act", one_of(acts.clone())),
129                        ("argument", one_of(arguments)),
130                        ("words", span()),
131                    ],
132                ));
133            }
134            kinds.push(variant(
135                "wrong_record",
136                vec![("act", one_of(acts.clone())), ("words", span())],
137            ));
138            kinds.push(variant("not_asked", vec![("act", one_of(acts))]));
139        }
140        object(vec![("findings", array(any_of(kinds)))])
141    }
142
143    fn render(&self, input: &CrossCheckInput) -> Vec<Message> {
144        let list = |title: &str, lines: &[String]| {
145            (!lines.is_empty()).then(|| {
146                let mut out = format!("{title}:");
147                for line in lines {
148                    let _ = write!(out, "\n- {line}");
149                }
150                out
151            })
152        };
153        let acts: Vec<String> = input.acts.iter().map(|act| act.line.clone()).collect();
154        vec![Message::user(render::sections([
155            render::last_assistant(self.turn),
156            Some(render::message(&self.turn.message)),
157            list("Acts understood", &acts).or_else(|| Some("Acts understood: none".to_owned())),
158            list("Questions", &input.questions),
159            list("Constraints", &input.constraints),
160            list("Read as nothing to act on", &input.unread),
161        ]))]
162    }
163
164    fn check(&self, input: &CrossCheckInput, output: &CrossChecked) -> Result<(), StructuralError> {
165        let acts: Vec<String> = input.acts.iter().map(|act| act.id.clone()).collect();
166        let words = &self.turn.message;
167        for (position, finding) in output.findings.iter().enumerate() {
168            let what = format!("finding {}", position + 1);
169            match finding {
170                Finding::Missing { words: span } => {
171                    check_span(&what, *span, words)?;
172                    let read = input
173                        .held
174                        .iter()
175                        .any(|held| held.from <= span.to && span.from <= held.to);
176                    if read {
177                        return Err(StructuralError::new(
178                            "words_already_read",
179                            format!(
180                                "{what}: those words are already read; missing words are \
181                                 words nothing holds"
182                            ),
183                        ));
184                    }
185                }
186                Finding::WrongValue {
187                    act,
188                    argument,
189                    words: span,
190                } => {
191                    check_one_of("act", act, &acts)?;
192                    check_span(&what, *span, words)?;
193                    let known = input
194                        .acts
195                        .iter()
196                        .find(|shown| &shown.id == act)
197                        .is_some_and(|shown| shown.arguments.contains(argument));
198                    if !known {
199                        return Err(StructuralError::new(
200                            "not_an_argument_of_the_act",
201                            format!("{what}: {act} has no argument `{argument}`"),
202                        ));
203                    }
204                }
205                Finding::WrongRecord { act, words: span } => {
206                    check_one_of("act", act, &acts)?;
207                    check_span(&what, *span, words)?;
208                }
209                Finding::NotAsked { act } => check_one_of("act", act, &acts)?,
210            }
211        }
212        Ok(())
213    }
214}
215
216#[cfg(test)]
217mod tests {
218    use super::*;
219    use crate::input::UnderstandingInput;
220
221    fn input() -> CrossCheckInput {
222        CrossCheckInput {
223            acts: vec![ShownAct {
224                id: "u1.a1".to_owned(),
225                line: "u1.a1 trip.set_name on Trip 1: value «Lisbon» (words 2 to 2)".to_owned(),
226                arguments: vec!["value".to_owned()],
227            }],
228            questions: Vec::new(),
229            constraints: Vec::new(),
230            unread: Vec::new(),
231            held: vec![Span::new(1, 1)],
232        }
233    }
234
235    fn turn() -> UnderstandingInput {
236        // [1]name [2]Lisbon [3]and [4]meals [5]too
237        UnderstandingInput::new("name Lisbon and meals too", "en-GB", chrono::NaiveDate::MIN)
238    }
239
240    #[test]
241    fn a_finding_must_name_an_act_and_an_argument_it_has() {
242        let turn = turn();
243        let task = CrossCheck::new(&turn);
244        let wrong = CrossChecked {
245            findings: vec![Finding::WrongValue {
246                act: "u1.a1".to_owned(),
247                argument: "due".to_owned(),
248                words: Span::new(0, 0),
249            }],
250        };
251        assert_eq!(
252            task.check(&input(), &wrong).unwrap_err().code,
253            "not_an_argument_of_the_act"
254        );
255        let unknown = CrossChecked {
256            findings: vec![Finding::NotAsked {
257                act: "u9.a1".to_owned(),
258            }],
259        };
260        assert_eq!(
261            task.check(&input(), &unknown).unwrap_err().code,
262            "not_in_set"
263        );
264    }
265
266    #[test]
267    fn missing_words_lie_outside_what_was_read() {
268        let turn = turn();
269        let task = CrossCheck::new(&turn);
270        let read = CrossChecked {
271            findings: vec![Finding::Missing {
272                words: Span::new(1, 2),
273            }],
274        };
275        assert_eq!(
276            task.check(&input(), &read).unwrap_err().code,
277            "words_already_read"
278        );
279        let fresh = CrossChecked {
280            findings: vec![Finding::Missing {
281                words: Span::new(3, 4),
282            }],
283        };
284        assert!(task.check(&input(), &fresh).is_ok());
285    }
286
287    #[test]
288    fn an_empty_answer_is_an_answer() {
289        let turn = turn();
290        let task = CrossCheck::new(&turn);
291        assert!(
292            task.check(
293                &input(),
294                &CrossChecked {
295                    findings: Vec::new()
296                }
297            )
298            .is_ok()
299        );
300    }
301}