Skip to main content

turnframe_understand/tasks/
segment.rs

1//! `segment`: which units the message holds, what kind each is, and its words.
2
3use serde::{Deserialize, Serialize};
4use serde_json::{Value, json};
5use turnframe_core::plan::AnswerBasis;
6use turnframe_core::understanding::{ConstraintKind, UnitKind};
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, nullable, object, one_of, span, text, variant};
13use crate::tasks::{check_one_of, check_span};
14use crate::words::Span;
15
16/// Names a unit's workflow when the message does not say which.
17pub const UNKNOWN: &str = "unknown";
18
19const BUILT_IN: &str = include_str!("../../prompts/understand/segment.md");
20
21const BASES: [&str; 4] = [
22    "current_committed_state",
23    "proposed_state",
24    "committed_state_after_turn",
25    "general_domain_knowledge",
26];
27
28const CONSTRAINTS: [&str; 7] = [
29    "do_not_submit",
30    "do_not_delete",
31    "draft_only",
32    "ask_before_applying",
33    "apply_only_if",
34    "no_external_effects",
35    "keep_unchanged",
36];
37
38/// The segmentation task, over one turn.
39#[derive(Debug, Clone, Copy)]
40pub struct Segment<'a> {
41    turn: &'a UnderstandingInput,
42}
43
44impl<'a> Segment<'a> {
45    /// The task for `turn`.
46    #[must_use]
47    pub const fn new(turn: &'a UnderstandingInput) -> Self {
48        Self { turn }
49    }
50
51    fn workflows(&self) -> Vec<String> {
52        let mut names: Vec<String> = self
53            .turn
54            .workflows
55            .iter()
56            .map(|workflow| workflow.key.to_string())
57            .collect();
58        names.push(UNKNOWN.to_owned());
59        names
60    }
61
62    /// Whether the message can answer the assistant: it asked for a value, or it said
63    /// something last that a value can answer.
64    fn answers_the_assistant(&self) -> bool {
65        self.turn.expectation.is_some()
66            || self
67                .turn
68                .transcript
69                .iter()
70                .any(|message| message.speaker == crate::input::Speaker::Assistant)
71    }
72
73    fn options(&self) -> Vec<String> {
74        self.turn
75            .card
76            .as_ref()
77            .filter(|card| card.accepts_typed_answer)
78            .map(|card| card.options.iter().map(|o| o.id.to_string()).collect())
79            .unwrap_or_default()
80    }
81
82    fn receipts(&self) -> Vec<String> {
83        let mut keys: Vec<String> = self.turn.receipts.iter().map(|r| r.key.clone()).collect();
84        keys.push(UNKNOWN.to_owned());
85        keys
86    }
87}
88
89/// The units of a message, as the model found them.
90#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
91pub struct Segmentation {
92    /// What the message asks for, in order, in a sentence or two.
93    pub analysis: String,
94    /// The units, in the order listed.
95    pub units: Vec<SegmentedUnit>,
96}
97
98/// One unit as the model found it.
99#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
100#[serde(tag = "kind", rename_all = "snake_case")]
101#[allow(missing_docs)]
102pub enum SegmentedUnit {
103    Request {
104        words: Span,
105        workflow: String,
106    },
107    Question {
108        words: Span,
109        workflow: String,
110        basis: AnswerBasis,
111        continues_previous: bool,
112    },
113    Constraint {
114        words: Span,
115        constraint: ConstraintKind,
116    },
117    Correction {
118        words: Span,
119        workflow: String,
120        /// The unit number it changes, from 1; `None` for an earlier turn.
121        corrects: Option<usize>,
122    },
123    Cancel {
124        words: Span,
125        workflow: String,
126        /// The unit number it withdraws, from 1; `None` for an earlier turn.
127        cancels: Option<usize>,
128    },
129    CardAnswer {
130        words: Span,
131        option: String,
132    },
133    Dispute {
134        words: Span,
135        receipt: String,
136    },
137    ProvidesValue {
138        words: Span,
139    },
140    Chitchat {
141        words: Span,
142    },
143}
144
145impl SegmentedUnit {
146    /// Its words.
147    #[must_use]
148    pub const fn words(&self) -> Span {
149        match self {
150            Self::Request { words, .. }
151            | Self::Question { words, .. }
152            | Self::Constraint { words, .. }
153            | Self::Correction { words, .. }
154            | Self::Cancel { words, .. }
155            | Self::CardAnswer { words, .. }
156            | Self::Dispute { words, .. }
157            | Self::ProvidesValue { words }
158            | Self::Chitchat { words } => *words,
159        }
160    }
161
162    /// The same unit over `words`.
163    fn at(&self, words: Span) -> Self {
164        let mut unit = self.clone();
165        match &mut unit {
166            Self::Request { words: own, .. }
167            | Self::Question { words: own, .. }
168            | Self::Constraint { words: own, .. }
169            | Self::Correction { words: own, .. }
170            | Self::Cancel { words: own, .. }
171            | Self::CardAnswer { words: own, .. }
172            | Self::Dispute { words: own, .. }
173            | Self::ProvidesValue { words: own }
174            | Self::Chitchat { words: own } => *own = words,
175        }
176        unit
177    }
178
179    /// Its kind.
180    #[must_use]
181    pub const fn kind(&self) -> UnitKind {
182        match self {
183            Self::Request { .. } => UnitKind::Request,
184            Self::Question { .. } => UnitKind::Question,
185            Self::Constraint { .. } => UnitKind::Constraint,
186            Self::Correction { .. } => UnitKind::Correction,
187            Self::Cancel { .. } => UnitKind::Cancel,
188            Self::CardAnswer { .. } => UnitKind::CardAnswer,
189            Self::Dispute { .. } => UnitKind::Dispute,
190            Self::ProvidesValue { .. } => UnitKind::ProvidesValue,
191            Self::Chitchat { .. } => UnitKind::Chitchat,
192        }
193    }
194
195    /// The workflow it names, `unknown` included.
196    #[must_use]
197    pub fn workflow(&self) -> Option<&str> {
198        match self {
199            Self::Request { workflow, .. }
200            | Self::Question { workflow, .. }
201            | Self::Correction { workflow, .. }
202            | Self::Cancel { workflow, .. } => Some(workflow),
203            _ => None,
204        }
205    }
206
207    /// The unit number it corrects or cancels, when it is about this message.
208    #[must_use]
209    pub const fn refers_to(&self) -> Option<usize> {
210        match self {
211            Self::Correction { corrects, .. } => *corrects,
212            Self::Cancel { cancels, .. } => *cancels,
213            _ => None,
214        }
215    }
216}
217
218impl ModelTask for Segment<'_> {
219    type Input = ();
220    type Output = Segmentation;
221
222    fn kind(&self) -> TaskKind {
223        TaskKind::Segment
224    }
225
226    fn prompt_name(&self) -> &str {
227        "understand.segment"
228    }
229
230    fn instructions(&self) -> &str {
231        BUILT_IN
232    }
233
234    fn schema(&self, _input: &()) -> Value {
235        let workflow = one_of(self.workflows());
236        let words = || ("words", span());
237        let mut variants = vec![
238            variant("request", vec![words(), ("workflow", workflow.clone())]),
239            variant(
240                "question",
241                vec![
242                    words(),
243                    ("workflow", workflow.clone()),
244                    ("basis", one_of(BASES)),
245                    ("continues_previous", json!({ "type": "boolean" })),
246                ],
247            ),
248            variant(
249                "constraint",
250                vec![words(), ("constraint", one_of(CONSTRAINTS))],
251            ),
252            variant(
253                "correction",
254                vec![
255                    words(),
256                    ("workflow", workflow.clone()),
257                    ("corrects", nullable(unit_number())),
258                ],
259            ),
260            variant(
261                "cancel",
262                vec![
263                    words(),
264                    ("workflow", workflow),
265                    ("cancels", nullable(unit_number())),
266                ],
267            ),
268        ];
269        let options = self.options();
270        if !options.is_empty() {
271            variants.push(variant(
272                "card_answer",
273                vec![words(), ("option", one_of(options))],
274            ));
275        }
276        // Both answer something the assistant said: with nothing said, there is nothing
277        // to contest or to give a value for.
278        if self.answers_the_assistant() {
279            variants.push(variant(
280                "dispute",
281                vec![words(), ("receipt", one_of(self.receipts()))],
282            ));
283            variants.push(variant("provides_value", vec![words()]));
284        }
285        variants.push(variant("chitchat", vec![words()]));
286        object(vec![
287            (
288                "analysis",
289                text("What the message asks for, in order, in a sentence or two."),
290            ),
291            ("units", array(any_of(variants))),
292        ])
293    }
294
295    fn render(&self, _input: &()) -> Vec<Message> {
296        let turn = self.turn;
297        vec![Message::user(render::sections([
298            Some(render::workflows(turn)),
299            render::last_assistant(turn),
300            render::card(turn),
301            render::expectation(turn),
302            render::receipts(turn),
303            Some(render::message(&turn.message)),
304        ]))]
305    }
306
307    fn check(&self, _input: &(), output: &Segmentation) -> Result<(), StructuralError> {
308        if output.units.is_empty() {
309            return Err(StructuralError::new(
310                "no_units",
311                "`units` is empty; every message has at least one unit, chitchat included",
312            ));
313        }
314        let (workflows, options, receipts) = (self.workflows(), self.options(), self.receipts());
315        let mut card_answers = 0;
316        for (position, unit) in output.units.iter().enumerate() {
317            let number = position + 1;
318            check_span(&format!("unit {number}"), unit.words(), &self.turn.message)?;
319            if let Some(workflow) = unit.workflow() {
320                check_one_of("workflow", workflow, &workflows)?;
321            }
322            match unit {
323                SegmentedUnit::CardAnswer { option, .. } => {
324                    card_answers += 1;
325                    check_one_of("option", option, &options)?;
326                }
327                SegmentedUnit::Dispute { .. } if !self.answers_the_assistant() => {
328                    return Err(StructuralError::new(
329                        "nothing_said",
330                        "a dispute contests what the assistant did or said, and it has said nothing",
331                    ));
332                }
333                SegmentedUnit::Dispute { receipt, .. } => {
334                    check_one_of("receipt", receipt, &receipts)?;
335                }
336                SegmentedUnit::ProvidesValue { .. } if !self.answers_the_assistant() => {
337                    return Err(StructuralError::new(
338                        "no_expectation",
339                        "provides_value needs the assistant to have asked for a value",
340                    ));
341                }
342                _ => {}
343            }
344        }
345        if card_answers > 1 {
346            return Err(StructuralError::new(
347                "several_card_answers",
348                "the card on screen takes one answer; list one card_answer unit at most",
349            ));
350        }
351        check_disjoint(&output.units)
352    }
353
354    /// Readings that differ only by a single word at a unit's edge agree: one keeps a word
355    /// joining two parts that the other leaves out, and both say the same of the message.
356    fn agree(&self, left: &Segmentation, right: &Segmentation) -> bool {
357        left.units.len() == right.units.len()
358            && left.units.iter().zip(&right.units).all(|(a, b)| {
359                let (mine, theirs) = (a.words(), b.words());
360                a.at(theirs) == *b
361                    && mine.from.abs_diff(theirs.from) <= 1
362                    && mine.to.abs_diff(theirs.to) <= 1
363            })
364    }
365}
366
367fn unit_number() -> Value {
368    json!({ "type": "integer", "minimum": 1 })
369}
370
371fn check_disjoint(units: &[SegmentedUnit]) -> Result<(), StructuralError> {
372    let mut spans: Vec<(Span, usize)> = units
373        .iter()
374        .enumerate()
375        .map(|(index, unit)| (unit.words(), index + 1))
376        .collect();
377    spans.sort();
378    for pair in spans.windows(2) {
379        let ((first, a), (second, b)) = (pair[0], pair[1]);
380        if second.from <= first.to {
381            return Err(StructuralError::new(
382                "units_overlap",
383                format!("units {a} and {b} share words; each word belongs to one unit at most"),
384            ));
385        }
386    }
387    Ok(())
388}
389
390#[cfg(test)]
391mod tests {
392    use super::*;
393
394    #[test]
395    fn readings_that_differ_by_one_word_at_a_units_edge_agree() {
396        let turn = UnderstandingInput::new("x", "en-GB", chrono::NaiveDate::MIN);
397        let task = Segment::new(&turn);
398        let reading = |units: Vec<SegmentedUnit>| Segmentation {
399            analysis: String::new(),
400            units,
401        };
402        let request = |from, to| SegmentedUnit::Request {
403            words: Span::new(from, to),
404            workflow: "trip".to_owned(),
405        };
406        let kept = reading(vec![request(12, 24), request(26, 32)]);
407        assert!(task.agree(&kept, &reading(vec![request(12, 24), request(25, 32)])));
408        assert!(!task.agree(&kept, &reading(vec![request(12, 24), request(24, 32)])));
409        assert!(!task.agree(
410            &kept,
411            &reading(vec![
412                request(12, 24),
413                SegmentedUnit::Chitchat {
414                    words: Span::new(25, 32)
415                }
416            ])
417        ));
418    }
419
420    #[test]
421    fn the_listed_bases_and_constraints_are_the_ones_that_deserialize() {
422        for basis in BASES {
423            serde_json::from_value::<AnswerBasis>(json!(basis)).unwrap();
424        }
425        for constraint in CONSTRAINTS {
426            serde_json::from_value::<ConstraintKind>(json!(constraint)).unwrap();
427        }
428    }
429}