use std::fmt::Write as _;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use turnframe_core::ids::WorkflowKey;
use turnframe_core::operation::OperationSpec;
use turnframe_core::understanding::ActId;
use turnframe_provider::request::Message;
use turnframe_tasks::{ModelTask, StructuralError, TaskKind};
use crate::input::{RecordBrief, UnderstandingInput};
use crate::render;
use crate::schema::{nullable, object, one_of, span};
use crate::tasks::{check_one_of, check_span};
use crate::words::Span;
pub const NEW: &str = "new";
pub const BY_NAME: &str = "by_name";
pub const AMBIGUOUS: &str = "ambiguous";
const BUILT_IN: &str = include_str!("../../prompts/understand/locate.md");
#[derive(Debug, Clone, Copy)]
pub struct Locate<'a> {
turn: &'a UnderstandingInput,
}
impl<'a> Locate<'a> {
#[must_use]
pub const fn new(turn: &'a UnderstandingInput) -> Self {
Self { turn }
}
}
#[derive(Debug, Clone, Copy)]
pub enum Candidate<'a> {
Record(&'a RecordBrief),
SameTurn {
act: ActId,
words: Span,
},
}
#[derive(Debug, Clone)]
pub struct LocateInput<'a> {
pub label: &'static str,
pub words: Span,
pub spec: &'a OperationSpec,
pub workflow: &'a WorkflowKey,
pub candidates: Vec<Candidate<'a>>,
pub allow_new: bool,
pub allow_not_listed: bool,
pub note: Option<String>,
}
impl LocateInput<'_> {
#[must_use]
pub fn handles(&self) -> Vec<(String, &Candidate<'_>)> {
let (mut records, mut same_turn) = (0, 0);
self.candidates
.iter()
.map(|candidate| {
let handle = match candidate {
Candidate::Record(_) => {
records += 1;
format!("r{records}")
}
Candidate::SameTurn { .. } => {
same_turn += 1;
format!("s{same_turn}")
}
};
(handle, candidate)
})
.collect()
}
#[must_use]
pub fn choices(&self) -> Vec<String> {
let mut choices: Vec<String> = self
.handles()
.into_iter()
.map(|(handle, _)| handle)
.collect();
if self.allow_new {
choices.push(NEW.to_owned());
}
if self.allow_not_listed {
choices.push(BY_NAME.to_owned());
}
choices.push(AMBIGUOUS.to_owned());
choices
}
#[must_use]
pub fn candidate(&self, handle: &str) -> Option<&Candidate<'_>> {
self.handles()
.into_iter()
.find(|(listed, _)| listed == handle)
.map(|(_, candidate)| candidate)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Location {
pub record: String,
pub named: Option<Span>,
}
impl<'a> ModelTask for Locate<'a> {
type Input = LocateInput<'a>;
type Output = Location;
fn kind(&self) -> TaskKind {
TaskKind::Locate
}
fn prompt_name(&self) -> &str {
"understand.locate"
}
fn instructions(&self) -> &str {
BUILT_IN
}
fn schema(&self, input: &LocateInput<'a>) -> Value {
object(vec![
("record", one_of(input.choices())),
("named", nullable(span())),
])
}
fn render(&self, input: &LocateInput<'a>) -> Vec<Message> {
let words = &self.turn.message;
let mut records = String::from("Records:");
let mut card_handle = None;
let mut last_about = Vec::new();
for (handle, candidate) in input.handles() {
match candidate {
Candidate::Record(record) => {
let _ = write!(
records,
"\n- {handle}: {}",
render::record_line(record, true)
);
let on_card = self.turn.card.as_ref().and_then(|c| c.record.as_ref());
if on_card == Some(&record.token) {
card_handle = Some(handle.clone());
}
if self.turn.last_subjects.contains(&record.token) {
last_about.push(handle);
}
}
Candidate::SameTurn { words: span, .. } => {
let said = words.slice(*span).unwrap_or_default();
let _ = write!(
records,
"\n- {handle}: the {} record this message creates, words {} to {}, {}",
input.workflow,
span.shown().0,
span.shown().1,
render::quoted(said)
);
}
}
}
if input.allow_new {
let _ = write!(records, "\n- {NEW}: a new {} record", input.workflow);
}
if input.allow_not_listed {
let _ = write!(
records,
"\n- {BY_NAME}: a record the user names that is not listed"
);
}
let _ = write!(records, "\n- {AMBIGUOUS}: more than one listed record fits");
vec![Message::user(render::sections([
Some(format!(
"Operation: {}",
render::operation_line(input.spec, self.turn)
)),
Some(records),
card_handle.map(|handle| format!("The card on screen is about {handle}.")),
render::last_assistant(self.turn),
(!last_about.is_empty()).then(|| {
format!(
"The last assistant message was about {}.",
last_about.join(" and ")
)
}),
Some(render::message(words)),
Some(render::unit(input.label, words, input.words)),
input
.note
.as_ref()
.map(|note| format!("A check of the whole message found: {note}")),
]))]
}
fn check(&self, input: &LocateInput<'a>, output: &Location) -> Result<(), StructuralError> {
check_one_of("record", &output.record, &input.choices())?;
if output.record == BY_NAME {
let Some(named) = output.named else {
return Err(StructuralError::new(
"missing_named",
"`named` must point at the words naming the record when `record` is by_name",
));
};
check_span("named", named, &self.turn.message)?;
}
Ok(())
}
fn agree(&self, left: &Location, right: &Location) -> bool {
left.record == right.record
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_note_from_the_whole_turn_check_is_shown() {
let turn = UnderstandingInput::new("name Lisbon", "en-GB", chrono::NaiveDate::MIN);
let spec = OperationSpec::new("trip.set_name").summary("Name the trip.");
let key = WorkflowKey::from("trip");
let input = LocateInput {
label: "Request",
words: Span::new(0, 1),
spec: &spec,
workflow: &key,
candidates: Vec::new(),
allow_new: true,
allow_not_listed: false,
note: Some("the record meant is named in «rent»".to_owned()),
};
let rendered = format!("{:?}", Locate::new(&turn).render(&input));
assert!(
rendered.contains(
"A check of the whole message found: the record meant is named in «rent»"
),
"{rendered}"
);
}
}