Skip to main content

turnframe_understand/tasks/
question_frame.rs

1//! `question_frame`: which record a question is about, and which declared subjects.
2
3use 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
18/// The answer for a question about no record.
19pub const NONE: &str = "none";
20
21/// The topics a question may have, as a model names them.
22pub 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/// The question-framing task, over one turn.
33#[derive(Debug, Clone, Copy)]
34pub struct QuestionFrame<'a> {
35    turn: &'a UnderstandingInput,
36}
37
38impl<'a> QuestionFrame<'a> {
39    /// The task for `turn`.
40    #[must_use]
41    pub const fn new(turn: &'a UnderstandingInput) -> Self {
42        Self { turn }
43    }
44
45    /// The topics on offer: knowledge only where a source can answer it.
46    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/// One question and what it may be about.
56#[derive(Debug, Clone)]
57pub struct QuestionInput<'a> {
58    /// Its words.
59    pub words: Span,
60    /// The records it may be about, listed as `r1`, `r2`...
61    pub records: Vec<&'a RecordBrief>,
62    /// The subjects it may be about.
63    pub subjects: Vec<&'a str>,
64}
65
66impl QuestionInput<'_> {
67    /// Every record answer allowed.
68    #[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    /// The record listed under `handle`.
76    #[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/// What a question is about.
84#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
85pub struct Framing {
86    /// What kind of thing it asks, one of [`TOPICS`].
87    pub topic: String,
88    /// A record handle, or `none`.
89    pub record: String,
90    /// The subjects, among those listed.
91    #[serde(default)]
92    pub subjects: Vec<String>,
93}
94
95impl Framing {
96    /// Whether the question asks if one particular thing can be done, which asks for it.
97    #[must_use]
98    pub fn asks_for_it(&self) -> bool {
99        self.topic == "ability"
100    }
101
102    /// The topic, as the understanding records it.
103    #[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}