Skip to main content

potato_spec/
spec.rs

1use serde::Deserialize;
2
3#[derive(Debug, Clone)]
4pub enum PromptRef {
5    Inline(String),
6    File(String),
7}
8
9impl<'de> Deserialize<'de> for PromptRef {
10    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
11        use serde::de::{self, Visitor};
12        struct V;
13        impl<'de> Visitor<'de> for V {
14            type Value = PromptRef;
15            fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
16                write!(f, "a string or a map with a 'path' key")
17            }
18            fn visit_str<E: de::Error>(self, v: &str) -> Result<Self::Value, E> {
19                Ok(PromptRef::Inline(v.to_owned()))
20            }
21            fn visit_map<A: de::MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
22                let mut path: Option<String> = None;
23                while let Some(key) = map.next_key::<String>()? {
24                    if key == "path" {
25                        path = Some(map.next_value()?);
26                    } else {
27                        map.next_value::<serde::de::IgnoredAny>()?;
28                    }
29                }
30                path.map(PromptRef::File)
31                    .ok_or_else(|| de::Error::missing_field("path"))
32            }
33        }
34        d.deserialize_any(V)
35    }
36}
37
38#[derive(Debug, Deserialize)]
39pub struct PotatoSpec {
40    #[serde(default)]
41    pub agents: Vec<AgentSpec>,
42    #[serde(default)]
43    pub workflows: Vec<WorkflowSpec>,
44}
45
46#[derive(Debug, Deserialize, Clone)]
47pub struct AgentSpec {
48    pub id: String,
49    #[serde(default)]
50    pub provider: Option<String>,
51    pub model: Option<String>,
52    pub system_prompt: Option<String>,
53    pub max_iterations: Option<u32>,
54    pub memory: Option<MemorySpec>,
55    #[serde(default)]
56    pub criteria: Vec<CriteriaSpec>,
57    #[serde(default)]
58    pub callbacks: Vec<CallbackSpec>,
59    #[serde(default)]
60    pub tools: Vec<ToolRef>,
61}
62
63#[derive(Debug, Deserialize, Clone)]
64#[serde(tag = "type", rename_all = "snake_case")]
65pub enum MemorySpec {
66    InMemory,
67    Windowed { window_size: usize },
68}
69
70#[derive(Debug, Deserialize, Clone)]
71#[serde(tag = "type", rename_all = "snake_case")]
72pub enum CriteriaSpec {
73    MaxIterations { max: u32 },
74    Keyword { keyword: String },
75    StructuredOutput { schema: Option<serde_json::Value> },
76}
77
78#[derive(Debug, Deserialize, Clone)]
79#[serde(untagged)]
80pub enum CallbackSpec {
81    BuiltIn {
82        #[serde(rename = "type")]
83        kind: String,
84    },
85    Named {
86        name: String,
87    },
88}
89
90#[derive(Debug, Deserialize, Clone)]
91pub struct ToolRef {
92    pub name: String,
93}
94
95#[derive(Debug, Deserialize, Clone)]
96#[serde(tag = "type", rename_all = "snake_case")]
97pub enum WorkflowSpec {
98    Sequential {
99        id: String,
100        pass_output: Option<bool>,
101        steps: Vec<StepSpec>,
102    },
103    Parallel {
104        id: String,
105        merge_strategy: Option<MergeStrategySpec>,
106        steps: Vec<StepSpec>,
107    },
108    Workflow {
109        id: String,
110        tasks: Vec<TaskSpec>,
111    },
112}
113
114impl WorkflowSpec {
115    pub fn id(&self) -> &str {
116        match self {
117            Self::Sequential { id, .. } => id,
118            Self::Parallel { id, .. } => id,
119            Self::Workflow { id, .. } => id,
120        }
121    }
122}
123
124#[derive(Debug, Deserialize, Clone)]
125#[serde(untagged)]
126pub enum StepSpec {
127    Ref {
128        #[serde(rename = "ref")]
129        agent_ref: String,
130    },
131    Inline(AgentSpec),
132}
133
134#[derive(Debug, Deserialize, Clone, Default)]
135#[serde(rename_all = "snake_case")]
136pub enum MergeStrategySpec {
137    #[default]
138    CollectAll,
139    First,
140}
141
142#[derive(Debug, Deserialize, Clone)]
143pub struct TaskSpec {
144    pub id: String,
145    pub agent: String,
146    pub prompt: PromptRef,
147    #[serde(default)]
148    pub dependencies: Vec<String>,
149    pub max_retries: Option<u32>,
150}
151
152#[cfg(test)]
153mod tests {
154    use super::*;
155
156    #[test]
157    fn prompt_ref_deserializes_from_string() {
158        let yaml = "\"hello world\"";
159        let p: PromptRef = serde_yaml::from_str(yaml).unwrap();
160        assert!(matches!(p, PromptRef::Inline(s) if s == "hello world"));
161    }
162
163    #[test]
164    fn prompt_ref_deserializes_from_path_map() {
165        let yaml = "path: ./prompts/foo.yaml";
166        let p: PromptRef = serde_yaml::from_str(yaml).unwrap();
167        assert!(matches!(p, PromptRef::File(s) if s == "./prompts/foo.yaml"));
168    }
169
170    #[test]
171    fn prompt_ref_map_missing_path_returns_error() {
172        let yaml = "other_key: value";
173        let result: Result<PromptRef, _> = serde_yaml::from_str(yaml);
174        assert!(result.is_err());
175    }
176}