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}