turnframe_understand/tasks/
cross_check.rs1use 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#[derive(Debug, Clone, Copy)]
20pub struct CrossCheck<'a> {
21 turn: &'a UnderstandingInput,
22}
23
24impl<'a> CrossCheck<'a> {
25 #[must_use]
27 pub const fn new(turn: &'a UnderstandingInput) -> Self {
28 Self { turn }
29 }
30}
31
32#[derive(Debug, Clone)]
34pub struct ShownAct {
35 pub id: String,
37 pub line: String,
39 pub arguments: Vec<String>,
41}
42
43#[derive(Debug, Clone, Default)]
45pub struct CrossCheckInput {
46 pub acts: Vec<ShownAct>,
48 pub questions: Vec<String>,
50 pub constraints: Vec<String>,
52 pub unread: Vec<String>,
54 pub held: Vec<Span>,
56}
57
58#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
60#[serde(tag = "kind", rename_all = "snake_case")]
61pub enum Finding {
62 Missing {
64 words: Span,
66 },
67 WrongValue {
69 act: String,
71 argument: String,
73 words: Span,
75 },
76 WrongRecord {
78 act: String,
80 words: Span,
82 },
83 NotAsked {
85 act: String,
87 },
88}
89
90#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
92pub struct CrossChecked {
93 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 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}