turnframe_understand/tasks/
coverage.rs1use 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#[derive(Debug, Clone, Copy)]
22pub struct Coverage<'a> {
23 turn: &'a UnderstandingInput,
24}
25
26impl<'a> Coverage<'a> {
27 #[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#[derive(Debug, Clone)]
47pub struct Found {
48 pub units: Vec<(UnitKind, Span)>,
50}
51
52#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
54pub struct Missed {
55 pub missed: Vec<MissedUnit>,
57}
58
59#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
61pub struct MissedUnit {
62 pub kind: MissedKind,
64 pub words: Span,
66 pub workflow: String,
68}
69
70#[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}