use std::fmt;
use serde::Serialize;
use serde::de::DeserializeOwned;
use turnframe_provider::request::Message;
pub use turnframe_provider::purpose::ModelPurpose as TaskKind;
pub trait ModelTask: Send + Sync {
type Input: Send + Sync;
type Output: Serialize + DeserializeOwned + Clone + Send + Sync + 'static;
fn kind(&self) -> TaskKind;
fn prompt_name(&self) -> &str;
fn instructions(&self) -> &str;
fn schema(&self, input: &Self::Input) -> serde_json::Value;
fn render(&self, input: &Self::Input) -> Vec<Message>;
fn check(&self, _input: &Self::Input, _output: &Self::Output) -> Result<(), StructuralError> {
Ok(())
}
fn agree(&self, left: &Self::Output, right: &Self::Output) -> bool {
serde_json::to_value(left).ok() == serde_json::to_value(right).ok()
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error("{message}")]
pub struct StructuralError {
pub code: &'static str,
pub message: String,
}
impl StructuralError {
#[must_use]
pub fn new(code: &'static str, message: impl Into<String>) -> Self {
Self {
code,
message: message.into(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct TaskId(String);
impl TaskId {
#[must_use]
pub fn new(path: impl Into<String>) -> Self {
Self(path.into())
}
#[must_use]
pub fn child(&self, segment: impl fmt::Display) -> Self {
Self(format!("{}/{segment}", self.0))
}
#[must_use]
pub fn call(&self, suffix: impl fmt::Display) -> String {
format!("{}#{suffix}", self.0)
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for TaskId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn identifiers_read_as_the_graph_they_ran_as() {
let unit = TaskId::new("u2");
let extract = unit.child("extract");
assert_eq!(extract.as_str(), "u2/extract");
assert_eq!(extract.call("repair1"), "u2/extract#repair1");
}
}