turnframe_understand/tasks/
question_frame.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 turnframe_core::understanding::QuestionTopic;
11
12use crate::input::{RecordBrief, UnderstandingInput};
13use crate::render;
14use crate::schema::{array, object, one_of};
15use crate::tasks::check_one_of;
16use crate::words::Span;
17
18pub const NONE: &str = "none";
20
21pub const TOPICS: [&str; 5] = [
23 "record_state",
24 "accepted_values",
25 "capabilities",
26 "ability",
27 "knowledge",
28];
29
30const BUILT_IN: &str = include_str!("../../prompts/understand/question_frame.md");
31
32#[derive(Debug, Clone, Copy)]
34pub struct QuestionFrame<'a> {
35 turn: &'a UnderstandingInput,
36}
37
38impl<'a> QuestionFrame<'a> {
39 #[must_use]
41 pub const fn new(turn: &'a UnderstandingInput) -> Self {
42 Self { turn }
43 }
44
45 fn topics(&self) -> Vec<String> {
47 TOPICS
48 .iter()
49 .filter(|topic| self.turn.knowledge || **topic != "knowledge")
50 .map(|topic| (*topic).to_owned())
51 .collect()
52 }
53}
54
55#[derive(Debug, Clone)]
57pub struct QuestionInput<'a> {
58 pub words: Span,
60 pub records: Vec<&'a RecordBrief>,
62 pub subjects: Vec<&'a str>,
64}
65
66impl QuestionInput<'_> {
67 #[must_use]
69 pub fn choices(&self) -> Vec<String> {
70 let mut choices: Vec<String> = (1..=self.records.len()).map(|n| format!("r{n}")).collect();
71 choices.push(NONE.to_owned());
72 choices
73 }
74
75 #[must_use]
77 pub fn record(&self, handle: &str) -> Option<&RecordBrief> {
78 let position: usize = handle.strip_prefix('r')?.parse().ok()?;
79 self.records.get(position.checked_sub(1)?).copied()
80 }
81}
82
83#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
85pub struct Framing {
86 pub topic: String,
88 pub record: String,
90 #[serde(default)]
92 pub subjects: Vec<String>,
93}
94
95impl Framing {
96 #[must_use]
98 pub fn asks_for_it(&self) -> bool {
99 self.topic == "ability"
100 }
101
102 #[must_use]
104 pub fn topic(&self) -> QuestionTopic {
105 match self.topic.as_str() {
106 "accepted_values" => QuestionTopic::AcceptedValues,
107 "capabilities" | "ability" => QuestionTopic::Capabilities,
108 "knowledge" => QuestionTopic::Knowledge,
109 _ => QuestionTopic::RecordState,
110 }
111 }
112}
113
114impl<'a> ModelTask for QuestionFrame<'a> {
115 type Input = QuestionInput<'a>;
116 type Output = Framing;
117
118 fn kind(&self) -> TaskKind {
119 TaskKind::QuestionFrame
120 }
121
122 fn prompt_name(&self) -> &str {
123 "understand.question_frame"
124 }
125
126 fn instructions(&self) -> &str {
127 BUILT_IN
128 }
129
130 fn schema(&self, input: &QuestionInput<'a>) -> Value {
131 let mut properties = vec![
132 ("topic", one_of(self.topics())),
133 ("record", one_of(input.choices())),
134 ];
135 if !input.subjects.is_empty() {
136 properties.push(("subjects", array(one_of(input.subjects.iter().copied()))));
137 }
138 object(properties)
139 }
140
141 fn render(&self, input: &QuestionInput<'a>) -> Vec<Message> {
142 let mut records = String::from("Records:");
143 let mut last_about = Vec::new();
144 for (position, record) in input.records.iter().enumerate() {
145 let handle = format!("r{}", position + 1);
146 let _ = write!(
147 records,
148 "\n- {handle}: {}",
149 render::record_line(record, true)
150 );
151 if !record.obligations.is_empty() {
152 let _ = write!(records, " ยท still needs: {}", record.obligations.join("; "));
153 }
154 if self.turn.last_subjects.contains(&record.token) {
155 last_about.push(handle);
156 }
157 }
158 let _ = write!(records, "\n- {NONE}: no record");
159 let conversation = (!last_about.is_empty()).then(|| {
160 format!(
161 "The last assistant message was about {}.",
162 last_about.join(" and ")
163 )
164 });
165 let subjects = (!input.subjects.is_empty())
166 .then(|| format!("Subjects: {}", input.subjects.join(", ")));
167 let words = &self.turn.message;
168 vec![Message::user(render::sections([
169 Some(records),
170 subjects,
171 conversation,
172 render::last_assistant(self.turn),
173 Some(render::message(words)),
174 Some(render::unit("Question", words, input.words)),
175 ]))]
176 }
177
178 fn check(&self, input: &QuestionInput<'a>, output: &Framing) -> Result<(), StructuralError> {
179 check_one_of("topic", &output.topic, &self.topics())?;
180 check_one_of("record", &output.record, &input.choices())?;
181 let subjects: Vec<String> = input.subjects.iter().map(|s| (*s).to_owned()).collect();
182 for subject in &output.subjects {
183 check_one_of("subjects", subject, &subjects)?;
184 }
185 Ok(())
186 }
187}
188
189#[cfg(test)]
190mod tests {
191 use super::*;
192
193 fn topics(turn: &UnderstandingInput) -> Vec<String> {
194 let input = QuestionInput {
195 words: Span::new(0, 1),
196 records: Vec::new(),
197 subjects: Vec::new(),
198 };
199 let schema = QuestionFrame::new(turn).schema(&input);
200 schema["properties"]["topic"]["enum"]
201 .as_array()
202 .map(|values| {
203 values
204 .iter()
205 .filter_map(|value| value.as_str().map(str::to_owned))
206 .collect()
207 })
208 .unwrap_or_default()
209 }
210
211 #[test]
212 fn knowledge_is_offered_only_where_a_source_holds_some() {
213 let turn = UnderstandingInput::new("what is it?", "en-GB", chrono::NaiveDate::MIN);
214 assert!(topics(&turn).contains(&"knowledge".to_owned()));
215 let without = turn.with_knowledge(false);
216 assert!(!topics(&without).contains(&"knowledge".to_owned()));
217 assert!(topics(&without).contains(&"record_state".to_owned()));
218 }
219}