1use crate::error::SpecError;
2use crate::spec::*;
3use potato_agent::agents::{agent::Agent, runner::AgentRunner};
4use potato_agent::{
5 AgentBuilder, AgentCallback, LoggingCallback, MergeStrategy, ParallelAgent,
6 ParallelAgentBuilder, SequentialAgent, SequentialAgentBuilder,
7};
8use potato_type::{prompt::Prompt, tools::AsyncTool, Provider};
9use potato_workflow::{Task, Workflow};
10use std::collections::{HashMap, HashSet};
11use std::path::{Component, Path, PathBuf};
12use std::sync::Arc;
13
14pub(crate) fn topo_sort_tasks(tasks: &[TaskSpec]) -> Result<Vec<&TaskSpec>, SpecError> {
15 let mut result: Vec<&TaskSpec> = Vec::with_capacity(tasks.len());
16 let mut remaining: Vec<&TaskSpec> = tasks.iter().collect();
17 let mut inserted_ids: HashSet<&str> = HashSet::new();
18
19 while !remaining.is_empty() {
20 let before = remaining.len();
21 remaining.retain(|task| {
22 let all_deps_inserted = task
23 .dependencies
24 .iter()
25 .all(|dep| inserted_ids.contains(dep.as_str()));
26 if all_deps_inserted {
27 inserted_ids.insert(task.id.as_str());
28 result.push(task);
29 false
30 } else {
31 true
32 }
33 });
34 if remaining.len() == before {
35 let cycle_ids: Vec<_> = remaining.iter().map(|t| t.id.as_str()).collect();
36 return Err(SpecError::WorkflowBuild {
37 id: "unknown".into(),
38 reason: format!("circular dependency or unresolvable tasks: {:?}", cycle_ids),
39 });
40 }
41 }
42 Ok(result)
43}
44
45pub struct SpecLoader {
46 async_tools: HashMap<String, Arc<dyn AsyncTool>>,
47 callbacks: HashMap<String, Arc<dyn AgentCallback>>,
48}
49
50impl Default for SpecLoader {
51 fn default() -> Self {
52 Self::new()
53 }
54}
55
56impl SpecLoader {
57 pub fn new() -> Self {
58 Self {
59 async_tools: HashMap::new(),
60 callbacks: HashMap::new(),
61 }
62 }
63
64 pub fn register_async_tool(mut self, name: &str, tool: Arc<dyn AsyncTool>) -> Self {
65 self.async_tools.insert(name.to_owned(), tool);
66 self
67 }
68
69 pub fn register_callback(mut self, name: &str, cb: Arc<dyn AgentCallback>) -> Self {
70 self.callbacks.insert(name.to_owned(), cb);
71 self
72 }
73
74 pub async fn from_spec(yaml: &str) -> Result<LoadedSpec, SpecError> {
76 Self::new().load_str(yaml).await
77 }
78
79 pub async fn from_spec_path(path: impl AsRef<Path>) -> Result<LoadedSpec, SpecError> {
81 Self::new().load_file(path).await
82 }
83
84 pub async fn load_file(&self, path: impl AsRef<Path>) -> Result<LoadedSpec, SpecError> {
85 let spec_path = path.as_ref().to_path_buf();
86 let content = tokio::fs::read_to_string(&spec_path).await?;
87 let base_dir = spec_path.parent().map(Path::to_path_buf);
88 self.load_str_with_base(&content, base_dir.as_deref()).await
89 }
90
91 pub async fn load_str(&self, yaml: &str) -> Result<LoadedSpec, SpecError> {
92 self.load_str_with_base(yaml, None).await
93 }
94
95 async fn load_str_with_base(
96 &self,
97 yaml: &str,
98 base_dir: Option<&Path>,
99 ) -> Result<LoadedSpec, SpecError> {
100 let spec: PotatoSpec = serde_yaml::from_str(yaml)?;
101 self.build_spec(spec, base_dir).await
102 }
103
104 async fn build_spec(
105 &self,
106 spec: PotatoSpec,
107 base_dir: Option<&Path>,
108 ) -> Result<LoadedSpec, SpecError> {
109 let mut agents: HashMap<String, Arc<Agent>> = HashMap::new();
110 for agent_spec in &spec.agents {
111 let agent = self.build_agent(agent_spec).await?;
112 agents.insert(agent_spec.id.clone(), agent);
113 }
114
115 let mut sequential: HashMap<String, Arc<SequentialAgent>> = HashMap::new();
116 let mut parallel: HashMap<String, Arc<ParallelAgent>> = HashMap::new();
117 let mut workflows: HashMap<String, Workflow> = HashMap::new();
118
119 for wf_spec in &spec.workflows {
120 match wf_spec {
121 WorkflowSpec::Sequential {
122 id,
123 pass_output,
124 steps,
125 } => {
126 let sa = self.build_sequential(*pass_output, steps, &agents).await?;
127 sequential.insert(id.clone(), sa);
128 }
129 WorkflowSpec::Parallel {
130 id,
131 merge_strategy,
132 steps,
133 } => {
134 let pa = self.build_parallel(merge_strategy, steps, &agents).await?;
135 parallel.insert(id.clone(), pa);
136 }
137 WorkflowSpec::Workflow { id, tasks } => {
138 let wf = self.build_workflow(id, tasks, &agents, base_dir).await?;
139 workflows.insert(id.clone(), wf);
140 }
141 }
142 }
143
144 Ok(LoadedSpec {
145 agents,
146 sequential,
147 parallel,
148 workflows,
149 })
150 }
151
152 async fn build_agent(&self, spec: &AgentSpec) -> Result<Arc<Agent>, SpecError> {
153 let provider = Provider::resolve(spec.provider.as_deref()).map_err(|e| {
154 SpecError::InvalidProvider {
155 value: spec.provider.clone().unwrap_or_default(),
156 reason: e.to_string(),
157 }
158 })?;
159
160 let mut builder = AgentBuilder::new().provider(provider);
161
162 if let Some(model) = &spec.model {
163 builder = builder.model(model.clone());
164 }
165 if let Some(sp) = &spec.system_prompt {
166 builder = builder.system_prompt(sp.clone());
167 }
168 let has_criteria_max_iterations = spec
169 .criteria
170 .iter()
171 .any(|c| matches!(c, CriteriaSpec::MaxIterations { .. }));
172
173 if let Some(max) = spec.max_iterations {
174 if !has_criteria_max_iterations {
175 builder = builder.max_iterations(max);
176 }
177 }
178
179 if let Some(mem) = &spec.memory {
180 builder = match mem {
181 MemorySpec::InMemory => builder.with_in_memory(),
182 MemorySpec::Windowed { window_size } => builder.with_windowed_memory(*window_size),
183 };
184 }
185
186 for criterion in &spec.criteria {
187 builder = match criterion {
188 CriteriaSpec::MaxIterations { max } => builder.max_iterations(*max),
189 CriteriaSpec::Keyword { keyword } => builder.stop_on_keyword(keyword.clone()),
190 CriteriaSpec::StructuredOutput { schema } => {
191 builder.stop_on_structured_output(schema.clone())
192 }
193 };
194 }
195
196 for cb_spec in &spec.callbacks {
197 let cb: Arc<dyn AgentCallback> = match cb_spec {
198 CallbackSpec::BuiltIn { kind } => match kind.as_str() {
199 "logging" => Arc::new(LoggingCallback),
200 other => return Err(SpecError::UnknownCallback { name: other.into() }),
201 },
202 CallbackSpec::Named { name } => self
203 .callbacks
204 .get(name)
205 .cloned()
206 .ok_or_else(|| SpecError::UnknownCallback { name: name.clone() })?,
207 };
208 builder = builder.with_callback(cb);
209 }
210
211 for tool_ref in &spec.tools {
212 if let Some(tool) = self.async_tools.get(&tool_ref.name) {
213 builder = builder.with_async_tool(Arc::clone(tool));
214 } else {
215 return Err(SpecError::UnknownTool {
216 name: tool_ref.name.clone(),
217 });
218 }
219 }
220
221 Ok(builder.build().await?)
222 }
223
224 async fn build_sequential(
225 &self,
226 pass_output: Option<bool>,
227 steps: &[StepSpec],
228 agents: &HashMap<String, Arc<Agent>>,
229 ) -> Result<Arc<SequentialAgent>, SpecError> {
230 let mut sb = SequentialAgentBuilder::new().pass_output(pass_output.unwrap_or(false));
231 for step in steps {
232 let runner = self.resolve_step(step, agents).await?;
233 sb = sb.then(runner);
234 }
235 Ok(sb.build())
236 }
237
238 async fn build_parallel(
239 &self,
240 merge_strategy: &Option<MergeStrategySpec>,
241 steps: &[StepSpec],
242 agents: &HashMap<String, Arc<Agent>>,
243 ) -> Result<Arc<ParallelAgent>, SpecError> {
244 let strategy = match merge_strategy {
245 None | Some(MergeStrategySpec::CollectAll) => MergeStrategy::CollectAll,
246 Some(MergeStrategySpec::First) => MergeStrategy::First,
247 };
248 let mut pb = ParallelAgentBuilder::new().merge_strategy(strategy);
249 for step in steps {
250 let runner = self.resolve_step(step, agents).await?;
251 pb = pb.with_agent(runner);
252 }
253 Ok(pb.build())
254 }
255
256 async fn build_workflow(
257 &self,
258 name: &str,
259 tasks: &[TaskSpec],
260 agents: &HashMap<String, Arc<Agent>>,
261 base_dir: Option<&Path>,
262 ) -> Result<Workflow, SpecError> {
263 let mut wf = Workflow::new(name);
264 let sorted = topo_sort_tasks(tasks)?;
265
266 for task_spec in sorted {
267 let agent = agents
268 .get(&task_spec.agent)
269 .ok_or_else(|| SpecError::UnknownAgentRef {
270 id: task_spec.agent.clone(),
271 })?;
272
273 let prompt = match &task_spec.prompt {
274 PromptRef::Inline(text) => {
275 let provider = agent.provider.clone();
276 let model =
277 agent
278 .model_override
279 .clone()
280 .ok_or_else(|| SpecError::WorkflowBuild {
281 id: task_spec.id.clone(),
282 reason: format!(
283 "agent '{}' used in task '{}' has no model set",
284 task_spec.agent, task_spec.id
285 ),
286 })?;
287 let config_value = serde_json::json!({
288 "model": model,
289 "provider": provider.as_str(),
290 "messages": [text],
291 });
292 let prompt_config = serde_json::from_value(config_value).map_err(|e| {
293 SpecError::WorkflowBuild {
294 id: task_spec.id.clone(),
295 reason: e.to_string(),
296 }
297 })?;
298 Prompt::from_generic_config(prompt_config).map_err(|e| {
299 SpecError::WorkflowBuild {
300 id: task_spec.id.clone(),
301 reason: e.to_string(),
302 }
303 })?
304 }
305 PromptRef::File(path) => {
306 if Path::new(path)
307 .components()
308 .any(|c| c == Component::ParentDir)
309 {
310 return Err(SpecError::PromptLoad {
311 path: path.clone(),
312 reason: "path must not contain '..' components".into(),
313 });
314 }
315 let path_owned = path.clone();
316 let base_dir_owned = base_dir.map(Path::to_path_buf);
317 let task_id = task_spec.id.clone();
318 let agent_provider = agent.provider.clone();
319 let prompt = tokio::task::spawn_blocking(move || {
320 let prompt_result = match &base_dir_owned {
321 Some(base_dir) => {
322 Prompt::from_path_with_base(PathBuf::from(&path_owned), base_dir)
323 }
324 None => Prompt::from_path(PathBuf::from(&path_owned)),
325 };
326
327 prompt_result.map_err(|e| SpecError::PromptLoad {
328 path: path_owned,
329 reason: e.to_string(),
330 })
331 })
332 .await
333 .map_err(|e| SpecError::WorkflowBuild {
334 id: task_id,
335 reason: format!("spawn_blocking failed: {e}"),
336 })??;
337
338 if prompt.provider != agent_provider {
339 return Err(SpecError::WorkflowBuild {
340 id: task_spec.id.clone(),
341 reason: format!(
342 "prompt file '{}' specifies provider '{}' but agent '{}' uses '{}'",
343 path,
344 prompt.provider.as_str(),
345 task_spec.agent,
346 agent_provider.as_str(),
347 ),
348 });
349 }
350
351 prompt
352 }
353 };
354
355 let task = Task::new(
356 &agent.id,
357 prompt,
358 &task_spec.id,
359 Some(task_spec.dependencies.clone()),
360 task_spec.max_retries,
361 )
362 .map_err(SpecError::AgentBuild)?;
363
364 wf.add_agent(agent);
365 wf.add_task(task).map_err(|e| SpecError::WorkflowBuild {
366 id: task_spec.id.clone(),
367 reason: e.to_string(),
368 })?;
369 }
370
371 Ok(wf)
372 }
373
374 async fn resolve_step(
375 &self,
376 step: &StepSpec,
377 agents: &HashMap<String, Arc<Agent>>,
378 ) -> Result<Arc<dyn AgentRunner>, SpecError> {
379 match step {
380 StepSpec::Ref { agent_ref } => agents
381 .get(agent_ref)
382 .map(|a| Arc::clone(a) as Arc<dyn AgentRunner>)
383 .ok_or_else(|| SpecError::UnknownAgentRef {
384 id: agent_ref.clone(),
385 }),
386 StepSpec::Inline(agent_spec) => {
387 let agent = self.build_agent(agent_spec).await?;
388 Ok(agent as Arc<dyn AgentRunner>)
389 }
390 }
391 }
392}
393
394pub struct LoadedSpec {
395 agents: HashMap<String, Arc<Agent>>,
396 sequential: HashMap<String, Arc<SequentialAgent>>,
397 parallel: HashMap<String, Arc<ParallelAgent>>,
398 workflows: HashMap<String, Workflow>,
399}
400
401impl LoadedSpec {
402 pub fn agent(&self, id: &str) -> Option<Arc<Agent>> {
403 self.agents.get(id).cloned()
404 }
405
406 pub fn sequential(&self, id: &str) -> Option<Arc<SequentialAgent>> {
407 self.sequential.get(id).cloned()
408 }
409
410 pub fn parallel(&self, id: &str) -> Option<Arc<ParallelAgent>> {
411 self.parallel.get(id).cloned()
412 }
413
414 pub fn workflow(&self, id: &str) -> Option<&Workflow> {
415 self.workflows.get(id)
416 }
417}
418
419#[cfg(test)]
420mod tests {
421 use super::*;
422 use std::fs;
423 use std::sync::Mutex;
424 use std::time::{SystemTime, UNIX_EPOCH};
425
426 static ENV_LOCK: Mutex<()> = Mutex::new(());
427
428 fn with_env_var<F: FnOnce()>(value: Option<&str>, f: F) {
429 let _guard = ENV_LOCK.lock().unwrap();
430 let prev = std::env::var(Provider::DEFAULT_ENV_VAR).ok();
431 match value {
432 Some(v) => std::env::set_var(Provider::DEFAULT_ENV_VAR, v),
433 None => std::env::remove_var(Provider::DEFAULT_ENV_VAR),
434 }
435 f();
436 match prev {
437 Some(v) => std::env::set_var(Provider::DEFAULT_ENV_VAR, v),
438 None => std::env::remove_var(Provider::DEFAULT_ENV_VAR),
439 }
440 }
441
442 fn create_temp_spec_dir() -> PathBuf {
443 let nanos = SystemTime::now()
444 .duration_since(UNIX_EPOCH)
445 .unwrap()
446 .as_nanos();
447 let dir = std::env::temp_dir().join(format!(
448 "potatohead-spec-tests-{}-{}",
449 std::process::id(),
450 nanos
451 ));
452 fs::create_dir_all(&dir).unwrap();
453 dir
454 }
455
456 fn make_task(id: &str, deps: Vec<&str>) -> TaskSpec {
457 TaskSpec {
458 id: id.to_string(),
459 agent: "x".to_string(),
460 prompt: PromptRef::Inline("p".to_string()),
461 dependencies: deps.into_iter().map(|s| s.to_string()).collect(),
462 max_retries: None,
463 }
464 }
465
466 #[test]
467 fn test_topo_sort_out_of_order() {
468 let tasks = vec![make_task("t2", vec!["t1"]), make_task("t1", vec![])];
469 let sorted = topo_sort_tasks(&tasks).unwrap();
470 assert_eq!(sorted.len(), 2);
471 assert_eq!(sorted[0].id, "t1");
472 assert_eq!(sorted[1].id, "t2");
473 }
474
475 #[test]
476 fn test_topo_sort_cycle_returns_error() {
477 let tasks = vec![make_task("a", vec!["b"]), make_task("b", vec!["a"])];
478 let result = topo_sort_tasks(&tasks);
479 assert!(result.is_err());
480 match result.unwrap_err() {
481 SpecError::WorkflowBuild { reason, .. } => {
482 assert!(reason.contains("circular dependency"));
483 assert!(reason.contains("a") || reason.contains("b"));
484 }
485 other => panic!("expected WorkflowBuild, got {:?}", other),
486 }
487 }
488
489 #[test]
490 fn test_from_spec_path_resolves_prompt_relative_to_spec_file() {
491 let runtime = tokio::runtime::Runtime::new().unwrap();
492 let temp_dir = create_temp_spec_dir();
493 let prompt_path = temp_dir.join("prompt.yaml");
494 let spec_path = temp_dir.join("workflow.yaml");
495
496 fs::write(
497 &prompt_path,
498 "model: gpt-4o\nprovider: openai\nmessages:\n - \"Hello ${name}\"\n",
499 )
500 .unwrap();
501 fs::write(
502 &spec_path,
503 r#"
504agents:
505 - id: worker
506 provider: openai
507 max_iterations: 1
508workflows:
509 - id: dag
510 type: workflow
511 tasks:
512 - id: t1
513 agent: worker
514 prompt:
515 path: "prompt"
516 dependencies: []
517"#,
518 )
519 .unwrap();
520
521 let loaded = runtime
522 .block_on(async { SpecLoader::from_spec_path(&spec_path).await })
523 .unwrap();
524
525 assert!(loaded.workflow("dag").is_some());
526
527 fs::remove_dir_all(temp_dir).unwrap();
528 }
529
530 #[test]
531 fn agent_spec_without_provider_uses_env_default() {
532 with_env_var(Some("openai"), || {
533 let runtime = tokio::runtime::Runtime::new().unwrap();
534 let spec = AgentSpec {
535 id: "worker".to_string(),
536 provider: None,
537 model: Some("gpt-4o".to_string()),
538 system_prompt: None,
539 max_iterations: Some(1),
540 memory: None,
541 criteria: Vec::new(),
542 callbacks: Vec::new(),
543 tools: Vec::new(),
544 };
545 let loader = SpecLoader::new();
546
547 let agent = runtime
548 .block_on(async { loader.build_agent(&spec).await })
549 .unwrap();
550
551 assert_eq!(agent.client_provider(), &Provider::OpenAI);
552 });
553 }
554
555 #[test]
556 fn agent_spec_without_provider_no_env_errors() {
557 with_env_var(None, || {
558 let runtime = tokio::runtime::Runtime::new().unwrap();
559 let spec = AgentSpec {
560 id: "worker".to_string(),
561 provider: None,
562 model: Some("gpt-4o".to_string()),
563 system_prompt: None,
564 max_iterations: Some(1),
565 memory: None,
566 criteria: Vec::new(),
567 callbacks: Vec::new(),
568 tools: Vec::new(),
569 };
570 let loader = SpecLoader::new();
571
572 let err = runtime
573 .block_on(async { loader.build_agent(&spec).await })
574 .unwrap_err();
575
576 match err {
577 SpecError::InvalidProvider { reason, .. } => {
578 assert!(reason.contains(Provider::DEFAULT_ENV_VAR));
579 }
580 other => panic!("expected InvalidProvider, got {:?}", other),
581 }
582 });
583 }
584
585 #[test]
586 fn test_load_str_without_base_does_not_resolve_relative_prompt_path() {
587 let runtime = tokio::runtime::Runtime::new().unwrap();
588
589 let yaml = r#"
590agents:
591 - id: worker
592 provider: openai
593 model: gpt-4o
594 max_iterations: 1
595workflows:
596 - id: dag
597 type: workflow
598 tasks:
599 - id: t1
600 agent: worker
601 prompt:
602 path: "definitely_missing_prompt"
603 dependencies: []
604"#;
605
606 let result = runtime.block_on(async { SpecLoader::from_spec(yaml).await });
607
608 match result {
609 Err(SpecError::PromptLoad { path, reason }) => {
610 assert_eq!(path, "definitely_missing_prompt");
611 assert!(reason.contains("definitely_missing_prompt"));
612 }
613 Ok(_) => panic!("expected PromptLoad, got Ok(..)"),
614 Err(other) => panic!("expected PromptLoad, got {other}"),
615 }
616 }
617}