Skip to main content

turnframe_understand/tasks/
coverage.rs

1//! `coverage`: which requests, questions or constraints the segmentation missed.
2
3use std::fmt::Write as _;
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7use turnframe_core::understanding::UnitKind;
8use turnframe_provider::request::Message;
9use turnframe_tasks::{ModelTask, StructuralError, TaskKind};
10
11use crate::input::UnderstandingInput;
12use crate::render;
13use crate::schema::{array, object, one_of, span};
14use crate::tasks::segment::UNKNOWN;
15use crate::tasks::{check_one_of, check_span};
16use crate::words::Span;
17
18const BUILT_IN: &str = include_str!("../../prompts/understand/coverage.md");
19
20/// The coverage task, over one turn.
21#[derive(Debug, Clone, Copy)]
22pub struct Coverage<'a> {
23    turn: &'a UnderstandingInput,
24}
25
26impl<'a> Coverage<'a> {
27    /// The task for `turn`.
28    #[must_use]
29    pub const fn new(turn: &'a UnderstandingInput) -> Self {
30        Self { turn }
31    }
32
33    fn workflows(&self) -> Vec<String> {
34        let mut names: Vec<String> = self
35            .turn
36            .workflows
37            .iter()
38            .map(|workflow| workflow.key.to_string())
39            .collect();
40        names.push(UNKNOWN.to_owned());
41        names
42    }
43}
44
45/// The units already found, which coverage looks past.
46#[derive(Debug, Clone)]
47pub struct Found {
48    /// Kind and words of each unit, in message order.
49    pub units: Vec<(UnitKind, Span)>,
50}
51
52/// What the segmentation missed.
53#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
54pub struct Missed {
55    /// The missed units.
56    pub missed: Vec<MissedUnit>,
57}
58
59/// One missed unit.
60#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
61pub struct MissedUnit {
62    /// What it is.
63    pub kind: MissedKind,
64    /// Its words.
65    pub words: Span,
66    /// The workflow it is about, or `unknown`.
67    pub workflow: String,
68}
69
70/// The kinds a missed unit may be.
71#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
72#[serde(rename_all = "snake_case")]
73#[allow(missing_docs)]
74pub enum MissedKind {
75    Request,
76    Question,
77    Constraint,
78    Correction,
79    Cancel,
80}
81
82impl From<MissedKind> for UnitKind {
83    fn from(kind: MissedKind) -> Self {
84        match kind {
85            MissedKind::Request => Self::Request,
86            MissedKind::Question => Self::Question,
87            MissedKind::Constraint => Self::Constraint,
88            MissedKind::Correction => Self::Correction,
89            MissedKind::Cancel => Self::Cancel,
90        }
91    }
92}
93
94pub(crate) fn kind_name(kind: UnitKind) -> &'static str {
95    match kind {
96        UnitKind::Request => "request",
97        UnitKind::Question => "question",
98        UnitKind::Constraint => "constraint",
99        UnitKind::Correction => "correction",
100        UnitKind::Cancel => "cancel",
101        UnitKind::CardAnswer => "card_answer",
102        UnitKind::Dispute => "dispute",
103        UnitKind::ProvidesValue => "provides_value",
104        _ => "chitchat",
105    }
106}
107
108impl ModelTask for Coverage<'_> {
109    type Input = Found;
110    type Output = Missed;
111
112    fn kind(&self) -> TaskKind {
113        TaskKind::Coverage
114    }
115
116    fn prompt_name(&self) -> &str {
117        "understand.coverage"
118    }
119
120    fn instructions(&self) -> &str {
121        BUILT_IN
122    }
123
124    fn schema(&self, _input: &Found) -> Value {
125        let kinds = one_of(["request", "question", "constraint", "correction", "cancel"]);
126        object(vec![(
127            "missed",
128            array(object(vec![
129                ("kind", kinds),
130                ("words", span()),
131                ("workflow", one_of(self.workflows())),
132            ])),
133        )])
134    }
135
136    fn render(&self, input: &Found) -> Vec<Message> {
137        let words = &self.turn.message;
138        let mut found = String::from("Units found:");
139        for (position, (kind, span)) in input.units.iter().enumerate() {
140            let text = words.slice(*span).unwrap_or_default();
141            let _ = write!(
142                found,
143                "\n{}. {}: words {} to {}, {}",
144                position + 1,
145                kind_name(*kind),
146                span.shown().0,
147                span.shown().1,
148                render::quoted(text)
149            );
150        }
151        vec![Message::user(render::sections([
152            Some(render::workflows(self.turn)),
153            Some(render::message(words)),
154            Some(found),
155        ]))]
156    }
157
158    fn check(&self, _input: &Found, output: &Missed) -> Result<(), StructuralError> {
159        let workflows = self.workflows();
160        for (position, unit) in output.missed.iter().enumerate() {
161            check_span(
162                &format!("missed unit {}", position + 1),
163                unit.words,
164                &self.turn.message,
165            )?;
166            check_one_of("workflow", &unit.workflow, &workflows)?;
167        }
168        Ok(())
169    }
170}