Skip to main content

potato_spec/
loader.rs

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    /// Load from a YAML string with a default (no-registry) loader.
75    pub async fn from_spec(yaml: &str) -> Result<LoadedSpec, SpecError> {
76        Self::new().load_str(yaml).await
77    }
78
79    /// Load from a YAML file with a default (no-registry) loader.
80    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}