1use 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
16pub 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#[derive(Debug, Clone, Copy)]
40pub struct Segment<'a> {
41 turn: &'a UnderstandingInput,
42}
43
44impl<'a> Segment<'a> {
45 #[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 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#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
91pub struct Segmentation {
92 pub analysis: String,
94 pub units: Vec<SegmentedUnit>,
96}
97
98#[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 corrects: Option<usize>,
122 },
123 Cancel {
124 words: Span,
125 workflow: String,
126 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 #[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 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 #[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 #[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 #[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 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 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}