use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use crate::event::FlowRunId;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum FormKind {
Confirm {
prompt: String,
},
SingleSelect {
prompt: String,
options: Vec<String>,
},
MultiSelect {
prompt: String,
options: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
min: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
max: Option<usize>,
},
Text {
prompt: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
placeholder: Option<String>,
#[serde(default)]
multiline: bool,
},
}
impl FormKind {
pub fn prompt(&self) -> &str {
match self {
Self::Confirm { prompt }
| Self::SingleSelect { prompt, .. }
| Self::MultiSelect { prompt, .. }
| Self::Text { prompt, .. } => prompt,
}
}
pub fn discriminator(&self) -> &'static str {
match self {
Self::Confirm { .. } => "confirm",
Self::SingleSelect { .. } => "single_select",
Self::MultiSelect { .. } => "multi_select",
Self::Text { .. } => "text",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct FormQuestion {
pub id: String,
#[serde(flatten)]
pub kind: FormKind,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct CompositeForm {
pub questions: Vec<FormQuestion>,
}
impl CompositeForm {
pub fn accepts(&self, submission: &FormSubmission) -> bool {
let FormSubmission::Submitted { answers } = submission else {
return true;
};
answers.len() == self.questions.len()
&& self
.questions
.iter()
.zip(answers)
.all(|(question, answer)| match (&question.kind, answer) {
(_, FormAnswer::Cancelled) => true,
(FormKind::Confirm { .. }, FormAnswer::Confirmed { .. })
| (FormKind::Text { .. }, FormAnswer::TextEntered { .. }) => true,
(
FormKind::SingleSelect { options, .. },
FormAnswer::Selected { index, label },
) => options.get(*index) == Some(label),
(
FormKind::MultiSelect {
options, min, max, ..
},
FormAnswer::MultiSelected { indices, labels },
) => {
indices.len() == labels.len()
&& min.is_none_or(|min| indices.len() >= min)
&& max.is_none_or(|max| indices.len() <= max)
&& indices
.iter()
.zip(labels)
.all(|(index, label)| options.get(*index) == Some(label))
}
_ => false,
})
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct DeferredFormAnswer {
pub prompt_id: String,
pub form: CompositeForm,
pub submission: FormSubmission,
}
impl DeferredFormAnswer {
pub fn as_user_text(&self) -> String {
let FormSubmission::Submitted { answers } = &self.submission else {
return String::new();
};
let mut lines = vec!["Answer to an earlier form:".to_string()];
for (question, answer) in self.form.questions.iter().zip(answers) {
let value = match answer {
FormAnswer::Confirmed { value } => value.to_string(),
FormAnswer::Selected { label, .. } => label.clone(),
FormAnswer::MultiSelected { labels, .. } => labels.join(", "),
FormAnswer::TextEntered { text } => text.clone(),
FormAnswer::Cancelled => "Cancelled".into(),
};
lines.push(format!("{}: {}", question.kind.prompt(), value));
}
lines.join("\n")
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum FormSubmission {
Submitted { answers: Vec<FormAnswer> },
Rejected,
}
#[derive(Debug, Clone)]
pub struct PendingForm {
pub form_id: String,
pub run_id: FlowRunId,
pub tool_use_id: String,
pub form: CompositeForm,
pub kind: FormKind,
pub emitted_at: DateTime<Utc>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum FormAnswer {
Confirmed {
value: bool,
},
Selected {
index: usize,
label: String,
},
MultiSelected {
indices: Vec<usize>,
labels: Vec<String>,
},
TextEntered {
text: String,
},
Cancelled,
}
impl FormAnswer {
pub fn discriminator(&self) -> &'static str {
match self {
Self::Confirmed { .. } => "confirmed",
Self::Selected { .. } => "selected",
Self::MultiSelected { .. } => "multi_selected",
Self::TextEntered { .. } => "text_entered",
Self::Cancelled => "cancelled",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn form_kind_serializes_with_tag() {
let k = FormKind::SingleSelect {
prompt: "pick".into(),
options: vec!["a".into(), "b".into()],
};
let s = serde_json::to_string(&k).unwrap();
assert!(s.contains(r#""kind":"single_select""#));
assert!(s.contains(r#""prompt":"pick""#));
}
#[test]
fn submitted_answers_must_match_the_question_schema() {
let form = CompositeForm {
questions: vec![FormQuestion {
id: "pick".into(),
kind: FormKind::SingleSelect {
prompt: "Choose".into(),
options: vec!["A".into(), "B".into()],
},
}],
};
assert!(form.accepts(&FormSubmission::Submitted {
answers: vec![FormAnswer::Selected {
index: 1,
label: "B".into(),
}],
}));
assert!(!form.accepts(&FormSubmission::Submitted {
answers: vec![FormAnswer::Selected {
index: 1,
label: "A".into(),
}],
}));
}
#[test]
fn form_kind_round_trip_confirm() {
let k = FormKind::Confirm {
prompt: "sure?".into(),
};
let s = serde_json::to_string(&k).unwrap();
let back: FormKind = serde_json::from_str(&s).unwrap();
assert_eq!(back, k);
}
#[test]
fn form_kind_round_trip_multi_select_omits_empty_bounds() {
let k = FormKind::MultiSelect {
prompt: "tags".into(),
options: vec!["a".into()],
min: None,
max: Some(2),
};
let s = serde_json::to_string(&k).unwrap();
assert!(!s.contains("\"min\""));
assert!(s.contains("\"max\":2"));
let back: FormKind = serde_json::from_str(&s).unwrap();
assert_eq!(back, k);
}
#[test]
fn form_answer_cancelled_serializes_as_tag_only() {
let a = FormAnswer::Cancelled;
let s = serde_json::to_string(&a).unwrap();
assert_eq!(s, r#"{"kind":"cancelled"}"#);
}
#[test]
fn form_answer_round_trip_multi_selected() {
let a = FormAnswer::MultiSelected {
indices: vec![0, 2],
labels: vec!["a".into(), "c".into()],
};
let s = serde_json::to_string(&a).unwrap();
let back: FormAnswer = serde_json::from_str(&s).unwrap();
assert_eq!(back, a);
}
#[test]
fn discriminators_are_stable() {
assert_eq!(
FormKind::Text {
prompt: "".into(),
placeholder: None,
multiline: false,
}
.discriminator(),
"text"
);
assert_eq!(FormAnswer::Cancelled.discriminator(), "cancelled");
}
}