Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions crates/baked_potato/tests/agent/fixtures/prompt.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
model: gpt-4o
provider: OpenAI
messages:
- "You are processing a DAG task from a prompt file."
170 changes: 164 additions & 6 deletions crates/baked_potato/tests/agent/spec_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -87,11 +87,15 @@ fn spec_loader_builds_openai_agent_and_runs() {
.block_on(async { SpecLoader::from_spec(SINGLE_AGENT_YAML).await })
.unwrap();

let agent = loaded.agent("assistant").expect("agent 'assistant' not found");
let agent = loaded
.agent("assistant")
.expect("agent 'assistant' not found");

let prompt_msg = OpenAIChatMessage {
role: "user".to_string(),
content: vec![ContentPart::Text(TextContentPart::new("Hello!".to_string()))],
content: vec![ContentPart::Text(TextContentPart::new(
"Hello!".to_string(),
))],
name: None,
};
let prompt = Prompt::new_rs(
Expand Down Expand Up @@ -124,7 +128,9 @@ fn spec_loader_builds_sequential_workflow_and_runs() {
.block_on(async { SpecLoader::from_spec(SEQUENTIAL_YAML).await })
.unwrap();

let seq = loaded.sequential("pipeline").expect("sequential 'pipeline' not found");
let seq = loaded
.sequential("pipeline")
.expect("sequential 'pipeline' not found");

let mut session = SessionState::new();
let outcome = runtime
Expand Down Expand Up @@ -185,6 +191,81 @@ fn spec_loader_builds_dag_workflow_and_runs() {
mock.stop_server().unwrap();
}

#[test]
fn spec_loader_dag_task_with_prompt_file_runs() {
let runtime = tokio::runtime::Runtime::new().unwrap();
let mut mock = LLMTestServer::new();
mock.start_server().unwrap();

let prompt_path = format!(
"{}/tests/agent/fixtures/prompt.yaml",
env!("CARGO_MANIFEST_DIR")
);

let yaml = format!(
r#"
agents:
- id: worker
provider: openai
model: gpt-4o
max_iterations: 1
workflows:
- id: dag
type: workflow
tasks:
- id: t1
agent: worker
prompt:
path: "{}"
dependencies: []
"#,
prompt_path
);

let loaded = runtime
.block_on(async { SpecLoader::from_spec(&yaml).await })
.unwrap();

let prompt = Prompt::from_path(std::path::PathBuf::from(&prompt_path)).unwrap();
assert_eq!(prompt.model, "gpt-4o");
assert_eq!(prompt.provider, Provider::OpenAI);
assert!(!prompt.openai_messages().unwrap().messages.is_empty());

let wf = loaded.workflow("dag").expect("workflow 'dag' not found");
let result = runtime.block_on(async { wf.run(None).await });
assert!(result.is_ok());

mock.stop_server().unwrap();
}

#[test]
fn spec_loader_dag_task_prompt_file_not_found_returns_error() {
let runtime = tokio::runtime::Runtime::new().unwrap();

let yaml = r#"
agents:
- id: worker
provider: openai
model: gpt-4o
max_iterations: 1
workflows:
- id: dag
type: workflow
tasks:
- id: t1
agent: worker
prompt:
path: "/nonexistent/path/prompt.yaml"
dependencies: []
"#;

let result = runtime.block_on(async { SpecLoader::from_spec(yaml).await });
assert!(
matches!(result, Err(SpecError::PromptLoad { .. })),
"expected PromptLoad error for nonexistent prompt file"
);
}

#[test]
fn spec_loader_dag_agent_missing_model_returns_error() {
let runtime = tokio::runtime::Runtime::new().unwrap();
Expand All @@ -210,6 +291,84 @@ workflows:
);
}

#[test]
fn spec_loader_dag_file_prompt_agent_without_model_succeeds() {
let runtime = tokio::runtime::Runtime::new().unwrap();
let mut mock = LLMTestServer::new();
mock.start_server().unwrap();

let prompt_path = format!(
"{}/tests/agent/fixtures/prompt.yaml",
env!("CARGO_MANIFEST_DIR")
);

let yaml = format!(
r#"
agents:
- id: no_model_agent
provider: openai
max_iterations: 1
workflows:
- id: dag
type: workflow
tasks:
- id: t1
agent: no_model_agent
prompt:
path: "{}"
dependencies: []
"#,
prompt_path
);

let loaded = runtime
.block_on(async { SpecLoader::from_spec(&yaml).await })
.unwrap();

let wf = loaded.workflow("dag").expect("workflow 'dag' not found");
let result = runtime.block_on(async { wf.run(None).await });
assert!(result.is_ok());

mock.stop_server().unwrap();
}

#[test]
fn spec_loader_dag_file_prompt_provider_mismatch_returns_error() {
let runtime = tokio::runtime::Runtime::new().unwrap();

let prompt_path = format!(
"{}/tests/agent/fixtures/prompt.yaml",
env!("CARGO_MANIFEST_DIR")
);

// prompt.yaml specifies provider: OpenAI, but agent uses gemini
let yaml = format!(
r#"
agents:
- id: gemini_agent
provider: gemini
model: gemini-2.0-flash
max_iterations: 1
workflows:
- id: dag
type: workflow
tasks:
- id: t1
agent: gemini_agent
prompt:
path: "{}"
dependencies: []
"#,
prompt_path
);

let result = runtime.block_on(async { SpecLoader::from_spec(&yaml).await });
assert!(
matches!(result, Err(SpecError::WorkflowBuild { .. })),
"expected WorkflowBuild error for provider mismatch"
);
}

#[test]
fn spec_loader_loads_agent_from_file() {
let runtime = tokio::runtime::Runtime::new().unwrap();
Expand All @@ -229,9 +388,8 @@ fn spec_loader_loads_agent_from_file() {
fn spec_loader_file_not_found_returns_io_error() {
let runtime = tokio::runtime::Runtime::new().unwrap();

let result = runtime.block_on(async {
SpecLoader::from_spec_path("/nonexistent/path/spec.yaml").await
});
let result =
runtime.block_on(async { SpecLoader::from_spec_path("/nonexistent/path/spec.yaml").await });

assert!(
matches!(result, Err(SpecError::Io(_))),
Expand Down
2 changes: 2 additions & 0 deletions crates/potato_spec/src/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,4 +19,6 @@ pub enum SpecError {
InvalidProvider { value: String, reason: String },
#[error("workflow build error for '{id}': {reason}")]
WorkflowBuild { id: String, reason: String },
#[error("failed to load prompt from '{path}': {reason}")]
PromptLoad { path: String, reason: String },
}
2 changes: 1 addition & 1 deletion crates/potato_spec/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,4 +4,4 @@ pub mod spec;

pub use error::SpecError;
pub use loader::{LoadedSpec, SpecLoader};
pub use spec::PotatoSpec;
pub use spec::{PotatoSpec, PromptRef};
117 changes: 78 additions & 39 deletions crates/potato_spec/src/loader.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ use potato_agent::{
use potato_type::{prompt::Prompt, tools::AsyncTool, Provider};
use potato_workflow::{Task, Workflow};
use std::collections::{HashMap, HashSet};
use std::path::Path;
use std::path::{Component, Path, PathBuf};
use std::sync::Arc;

pub(crate) fn topo_sort_tasks(tasks: &[TaskSpec]) -> Result<Vec<&TaskSpec>, SpecError> {
Expand Down Expand Up @@ -254,35 +254,80 @@ impl SpecLoader {
id: task_spec.agent.clone(),
})?;

let provider = agent.provider.clone();
let model = agent
.model_override
.clone()
.ok_or_else(|| SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: format!(
"agent '{}' used in task '{}' has no model set",
task_spec.agent, task_spec.id
),
})?;

let config_value = serde_json::json!({
"model": model,
"provider": provider.as_str(),
"messages": [task_spec.prompt.clone()],
});
let prompt_config =
serde_json::from_value(config_value).map_err(|e| SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: e.to_string(),
})?;

let prompt = Prompt::from_generic_config(prompt_config).map_err(|e| {
SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: e.to_string(),
}
})?;
let prompt =
match &task_spec.prompt {
PromptRef::Inline(text) => {
let provider = agent.provider.clone();
let model = agent.model_override.clone().ok_or_else(|| {
SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: format!(
"agent '{}' used in task '{}' has no model set",
task_spec.agent, task_spec.id
),
}
})?;
let config_value = serde_json::json!({
"model": model,
"provider": provider.as_str(),
"messages": [text],
});
let prompt_config = serde_json::from_value(config_value).map_err(|e| {
SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: e.to_string(),
}
})?;
Prompt::from_generic_config(prompt_config).map_err(|e| {
SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: e.to_string(),
}
})?
}
PromptRef::File(path) => {
if Path::new(path)
.components()
.any(|c| c == Component::ParentDir)
{
return Err(SpecError::PromptLoad {
path: path.clone(),
reason: "path must not contain '..' components".into(),
});
}
let path_owned = path.clone();
let task_id = task_spec.id.clone();
let agent_provider = agent.provider.clone();
let prompt = tokio::task::spawn_blocking(move || {
Prompt::from_path(PathBuf::from(&path_owned)).map_err(|e| {
SpecError::PromptLoad {
path: path_owned,
reason: e.to_string(),
}
})
})
.await
.map_err(|e| SpecError::WorkflowBuild {
id: task_id,
reason: format!("spawn_blocking failed: {e}"),
})??;

if prompt.provider != agent_provider {
return Err(SpecError::WorkflowBuild {
id: task_spec.id.clone(),
reason: format!(
"prompt file '{}' specifies provider '{}' but agent '{}' uses '{}'",
path,
prompt.provider.as_str(),
task_spec.agent,
agent_provider.as_str(),
),
});
}

prompt
}
};

let task = Task::new(
&agent.id,
Expand Down Expand Up @@ -356,18 +401,15 @@ mod tests {
TaskSpec {
id: id.to_string(),
agent: "x".to_string(),
prompt: "p".to_string(),
prompt: PromptRef::Inline("p".to_string()),
dependencies: deps.into_iter().map(|s| s.to_string()).collect(),
max_retries: None,
}
}

#[test]
fn test_topo_sort_out_of_order() {
let tasks = vec![
make_task("t2", vec!["t1"]),
make_task("t1", vec![]),
];
let tasks = vec![make_task("t2", vec!["t1"]), make_task("t1", vec![])];
let sorted = topo_sort_tasks(&tasks).unwrap();
assert_eq!(sorted.len(), 2);
assert_eq!(sorted[0].id, "t1");
Expand All @@ -376,10 +418,7 @@ mod tests {

#[test]
fn test_topo_sort_cycle_returns_error() {
let tasks = vec![
make_task("a", vec!["b"]),
make_task("b", vec!["a"]),
];
let tasks = vec![make_task("a", vec!["b"]), make_task("b", vec!["a"])];
let result = topo_sort_tasks(&tasks);
assert!(result.is_err());
match result.unwrap_err() {
Expand Down
Loading
Loading