1use std::fmt;
4
5use serde::Serialize;
6use serde::de::DeserializeOwned;
7use turnframe_provider::request::Message;
8
9pub use turnframe_provider::purpose::ModelPurpose as TaskKind;
11
12pub trait ModelTask: Send + Sync {
17 type Input: Send + Sync;
19 type Output: Serialize + DeserializeOwned + Clone + Send + Sync + 'static;
21
22 fn kind(&self) -> TaskKind;
24
25 fn prompt_name(&self) -> &str;
27
28 fn instructions(&self) -> &str;
30
31 fn schema(&self, input: &Self::Input) -> serde_json::Value;
33
34 fn render(&self, input: &Self::Input) -> Vec<Message>;
36
37 fn check(&self, _input: &Self::Input, _output: &Self::Output) -> Result<(), StructuralError> {
43 Ok(())
44 }
45
46 fn agree(&self, left: &Self::Output, right: &Self::Output) -> bool {
48 serde_json::to_value(left).ok() == serde_json::to_value(right).ok()
49 }
50}
51
52#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
54#[error("{message}")]
55pub struct StructuralError {
56 pub code: &'static str,
58 pub message: String,
60}
61
62impl StructuralError {
63 #[must_use]
65 pub fn new(code: &'static str, message: impl Into<String>) -> Self {
66 Self {
67 code,
68 message: message.into(),
69 }
70 }
71}
72
73#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
75pub struct TaskId(String);
76
77impl TaskId {
78 #[must_use]
80 pub fn new(path: impl Into<String>) -> Self {
81 Self(path.into())
82 }
83
84 #[must_use]
86 pub fn child(&self, segment: impl fmt::Display) -> Self {
87 Self(format!("{}/{segment}", self.0))
88 }
89
90 #[must_use]
92 pub fn call(&self, suffix: impl fmt::Display) -> String {
93 format!("{}#{suffix}", self.0)
94 }
95
96 #[must_use]
98 pub fn as_str(&self) -> &str {
99 &self.0
100 }
101}
102
103impl fmt::Display for TaskId {
104 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
105 f.write_str(&self.0)
106 }
107}
108
109#[cfg(test)]
110mod tests {
111 use super::*;
112
113 #[test]
114 fn identifiers_read_as_the_graph_they_ran_as() {
115 let unit = TaskId::new("u2");
116 let extract = unit.child("extract");
117 assert_eq!(extract.as_str(), "u2/extract");
118 assert_eq!(extract.call("repair1"), "u2/extract#repair1");
119 }
120}