use crate::tasklist::TaskList;
use crate::{
events::{EventTracker, TaskEvent},
workflow::error::WorkflowError,
};
pub use potato_agent::agents::{
agent::{Agent, PyAgent},
task::{Task, TaskStatus, WorkflowTask},
};
use potato_agent::{AgentError, PyAgentResponse};
use potato_state::block_on;
use potato_type::prompt::{parse_response_to_json, MessageNum};
use potato_type::Provider;
use potato_util::utils::depythonize_object_to_value;
use potato_util::{create_uuid7, utils::update_serde_map_with, PyHelperFuncs};
use pyo3::prelude::*;
use pythonize::pythonize;
use serde::{
de::{self, MapAccess, Visitor},
ser::SerializeStruct,
Deserialize, Deserializer, Serialize, Serializer,
};
use serde_json::Map;
use serde_json::Value;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::sync::RwLock;
use tracing::instrument;
use tracing::{debug, error, info, warn};
use pyo3::types::PyDict;
pub type Context = (HashMap<String, Vec<MessageNum>>, Value, Option<Arc<Value>>);
#[derive(Debug)]
#[pyclass(skip_from_py_object)]
pub struct WorkflowResult {
#[pyo3(get)]
pub tasks: HashMap<String, Py<WorkflowTask>>,
#[pyo3(get)]
pub events: Vec<TaskEvent>,
last_task_id: Option<String>,
}
impl WorkflowResult {
pub fn new(
py: Python,
tasks: HashMap<String, Task>,
output_types: &HashMap<String, Arc<Py<PyAny>>>,
events: Vec<TaskEvent>,
last_task_id: Option<String>,
) -> Self {
let py_tasks = tasks
.into_iter()
.map(|(id, task)| {
let py_agent_response = if let Some(result) = task.result {
let output_type = output_types.get(&id).map(|arc| arc.as_ref().clone_ref(py));
Some(PyAgentResponse::new(result, output_type))
} else {
None
};
let py_task = WorkflowTask {
id: task.id.clone(),
prompt: task.prompt,
dependencies: task.dependencies,
status: task.status,
agent_id: task.agent_id,
result: py_agent_response,
max_retries: task.max_retries,
retry_count: task.retry_count,
};
(id, Py::new(py, py_task).unwrap())
})
.collect::<HashMap<_, _>>();
Self {
tasks: py_tasks,
events,
last_task_id,
}
}
}
#[pymethods]
impl WorkflowResult {
pub fn __str__(&self) -> String {
let json = serde_json::json!({
"tasks": serde_json::to_value(&self.tasks).unwrap_or(Value::Null),
"events": serde_json::to_value(&self.events).unwrap_or(Value::Null)
});
PyHelperFuncs::__str__(&json)
}
#[getter]
pub fn result<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, AgentError> {
if let Some(last_task_id) = &self.last_task_id {
if let Some(task) = self.tasks.get(last_task_id) {
let result = task.bind(py).getattr("result")?;
return Ok(result);
}
}
Ok(py.None().bind(py).clone())
}
}
#[derive(Debug, Clone)]
pub struct Workflow {
pub id: String,
pub name: String,
pub task_list: TaskList,
pub agents: HashMap<String, Arc<Agent>>,
pub event_tracker: Arc<RwLock<EventTracker>>,
pub global_context: Option<Arc<Value>>,
}
impl PartialEq for Workflow {
fn eq(&self, other: &Self) -> bool {
self.id == other.id && self.name == other.name
}
}
impl Workflow {
pub async fn reset_agents(&mut self) -> Result<(), WorkflowError> {
let mut agents_map = self.agents.clone();
for agent in self.agents.values_mut() {
agents_map.insert(agent.id.clone(), Arc::new(agent.rebuild_client().await?));
}
self.agents = agents_map;
Ok(())
}
pub fn new(name: &str) -> Self {
debug!("Creating new workflow: {}", name);
let id = create_uuid7();
Self {
id: id.clone(),
name: name.to_string(),
task_list: TaskList::new(),
agents: HashMap::new(),
event_tracker: Arc::new(RwLock::new(EventTracker::new(id))),
global_context: None, }
}
pub fn events(&self) -> Vec<TaskEvent> {
let tracker = self.event_tracker.read().unwrap();
let events = tracker.events.read().unwrap().clone();
events
}
pub fn total_duration(&self) -> i32 {
let tracker = self.event_tracker.read().unwrap();
if tracker.is_empty() {
0
} else {
let mut total_duration = chrono::Duration::zero();
for event in tracker.events.read().unwrap().iter() {
total_duration += event.details.duration.unwrap_or(chrono::Duration::zero());
}
total_duration.subsec_millis()
}
}
pub fn get_new_workflow(
&self,
global_context: Option<Arc<Value>>,
) -> Result<Self, WorkflowError> {
let id = create_uuid7();
let task_list = self.task_list.deep_clone()?;
Ok(Workflow {
id: id.clone(),
name: self.name.clone(),
task_list,
agents: self.agents.clone(), event_tracker: Arc::new(RwLock::new(EventTracker::new(id))),
global_context,
})
}
pub async fn run(
&self,
global_context: Option<Value>,
) -> Result<Arc<RwLock<Workflow>>, WorkflowError> {
debug!("Running workflow: {}", self.name);
let global_context = global_context.map(Arc::new);
let run_workflow = Arc::new(RwLock::new(self.get_new_workflow(global_context)?));
execute_workflow(&run_workflow).await?;
Ok(run_workflow)
}
pub fn is_complete(&self) -> bool {
self.task_list.is_complete()
}
pub fn pending_count(&self) -> usize {
self.task_list.pending_count()
}
pub fn add_task(&mut self, task: Task) -> Result<(), WorkflowError> {
self.task_list.add_task(task)
}
pub fn add_tasks(&mut self, tasks: Vec<Task>) -> Result<(), WorkflowError> {
for task in tasks {
self.task_list.add_task(task)?;
}
Ok(())
}
pub fn add_agent(&mut self, agent: &Agent) {
self.agents
.insert(agent.id.clone(), Arc::new(agent.clone()));
}
pub fn add_agents(&mut self, agents: &[&Agent]) {
for agent in agents {
self.add_agent(agent);
}
}
pub async fn execute_task(&self, task: &str, context: &Value) -> Result<Value, WorkflowError> {
let task = self
.task_list
.get_task(task)
.ok_or_else(|| WorkflowError::TaskNotFound(task.to_string()))?;
let agent = {
let task_guard = task.read().map_err(|_| WorkflowError::TaskLockError)?;
self.agents
.get(&task_guard.agent_id)
.ok_or_else(|| WorkflowError::AgentNotFound(task_guard.agent_id.clone()))?
.clone()
};
let max_retries = {
let task_guard = task.read().map_err(|_| WorkflowError::TaskLockError)?;
task_guard.max_retries
};
for attempt in 0..=max_retries {
match agent.execute_task_with_context(&task, context).await {
Ok(response) => {
let is_valid = validate_response_schema(&task, &response);
if !is_valid {
if attempt == max_retries {
let (task_id, expected_schema, received_response) = {
let task_guard = task.read().unwrap();
(
task_guard.id.clone(),
task_guard
.prompt
.response_json_schema()
.map(|s| s.to_string())
.unwrap_or_else(|| "No schema".to_string()),
response
.response_value()
.map(|v| v.to_string())
.unwrap_or_else(|| "No response".to_string()),
)
};
error!(
"Task {} response validation failed after {} attempts",
task_id,
attempt + 1
);
return Err(WorkflowError::ResponseValidationFailed {
task_id,
expected_schema,
received_response,
});
}
warn!(
"Task validation failed (attempt {}/{}), retrying...",
attempt + 1,
max_retries + 1
);
continue;
}
let value = match response.response_value() {
Some(v) => v,
None => Value::String(response.response_text()),
};
return Ok(value);
}
Err(e) => {
let task_id = { task.read().unwrap().id.clone() };
warn!(
"Task {} execution failed (attempt {}/{}): {}",
task_id,
attempt + 1,
max_retries + 1,
e
);
if attempt == max_retries {
error!("Task {} exceeded max retries ({})", task_id, max_retries);
return Err(WorkflowError::MaxRetriesExceeded(task_id));
}
}
}
}
unreachable!("Loop should always return via Ok or error")
}
pub fn execution_plan(&self) -> Result<HashMap<i32, HashSet<String>>, WorkflowError> {
let mut remaining: HashMap<String, HashSet<String>> = self
.task_list
.tasks
.iter()
.map(|(id, task)| {
(
id.clone(),
task.read().unwrap().dependencies.iter().cloned().collect(),
)
})
.collect();
let mut executed = HashSet::new();
let mut plan = HashMap::new();
let mut step = 1;
while !remaining.is_empty() {
let ready_keys: Vec<String> = remaining
.iter()
.filter(|(_, deps)| deps.is_subset(&executed))
.map(|(id, _)| id.to_string())
.collect();
if ready_keys.is_empty() {
break;
}
let mut ready_set = HashSet::with_capacity(ready_keys.len());
for key in ready_keys {
executed.insert(key.clone());
remaining.remove(&key);
ready_set.insert(key);
}
plan.insert(step, ready_set);
step += 1;
}
Ok(plan)
}
pub fn __str__(&self) -> String {
PyHelperFuncs::__str__(&self.task_list)
}
pub fn serialize(&self) -> Result<String, serde_json::Error> {
let json = serde_json::to_string(self).unwrap();
Ok(json)
}
pub fn from_json(json: &str) -> Result<Self, WorkflowError> {
Ok(serde_json::from_str(json)?)
}
pub fn task_names(&self) -> Vec<String> {
self.task_list
.tasks
.keys()
.cloned()
.collect::<Vec<String>>()
}
pub fn last_task_id(&self) -> Option<String> {
self.task_list.get_last_task_id()
}
}
fn is_workflow_complete(workflow: &Arc<RwLock<Workflow>>) -> bool {
workflow.read().unwrap().is_complete()
}
fn reset_failed_workflow_tasks(workflow: &Arc<RwLock<Workflow>>) -> Result<(), WorkflowError> {
match workflow.write().unwrap().task_list.reset_failed_tasks() {
Ok(_) => Ok(()),
Err(e) => {
warn!("Failed to reset failed tasks: {}", e);
Err(e)
}
}
}
fn get_ready_tasks(workflow: &Arc<RwLock<Workflow>>) -> Vec<Arc<RwLock<Task>>> {
workflow.read().unwrap().task_list.get_ready_tasks()
}
fn check_for_circular_dependencies(workflow: &Arc<RwLock<Workflow>>) -> bool {
let pending_count = workflow.read().unwrap().pending_count();
if pending_count > 0 {
warn!(
"No ready tasks found but {} pending tasks remain. Possible circular dependency.",
pending_count
);
return true;
}
false
}
fn mark_task_as_running(task: Arc<RwLock<Task>>, event_tracker: &Arc<RwLock<EventTracker>>) {
let mut task = task.write().unwrap();
task.set_status(TaskStatus::Running);
event_tracker.write().unwrap().record_task_started(&task.id);
}
fn get_agent_for_task(
workflow: &Arc<RwLock<Workflow>>,
agent_id: &str,
) -> Result<Arc<Agent>, WorkflowError> {
let wf = workflow.read().unwrap();
match wf.agents.get(agent_id) {
Some(agent) => Ok(agent.clone()),
None => Err(WorkflowError::AgentNotFound(agent_id.to_string())),
}
}
#[instrument(skip_all)]
fn build_task_context(
workflow: &Arc<RwLock<Workflow>>,
task_dependencies: &Vec<String>,
provider: &Provider,
) -> Result<Context, WorkflowError> {
let wf = workflow.read().unwrap();
let mut ctx = HashMap::new();
let mut param_ctx: Value = Value::Object(Map::new());
for dep_id in task_dependencies {
debug!("Building context for task dependency: {}", dep_id);
if let Some(dep) = wf.task_list.get_task(dep_id) {
if let Some(result) = &dep.read().unwrap().result {
let msg_to_insert = result.response.to_message_num(provider);
match msg_to_insert {
Ok(message) => {
ctx.insert(dep_id.clone(), message);
}
Err(e) => {
warn!("Failed to convert response to message: {}", e);
}
}
if let Some(structure_output) = result.response.extract_structured_data() {
if structure_output.is_object() {
update_serde_map_with(&mut param_ctx, &structure_output)?;
}
}
}
}
}
debug!("Built context for task dependencies: {:?}", ctx);
let global_context = workflow
.read()
.unwrap()
.global_context
.as_ref()
.map(Arc::clone);
Ok((ctx, param_ctx, global_context))
}
fn validate_response_schema(
task: &Arc<RwLock<Task>>,
response: &potato_agent::AgentResponse,
) -> bool {
task.read()
.ok()
.and_then(|t| {
response
.response_value()
.map(|value| t.validate_output(&value).is_ok())
})
.unwrap_or(true)
}
fn spawn_task_execution(
event_tracker: Arc<RwLock<EventTracker>>,
task: Arc<RwLock<Task>>,
task_id: String,
agent: Arc<Agent>,
context: HashMap<String, Vec<MessageNum>>,
parameter_context: Value,
global_context: Option<Arc<Value>>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let result = agent
.execute_task_with_context_message(&task, context, parameter_context, global_context)
.await;
match result {
Ok(response) => {
info!("Task {} completed successfully", task_id);
let is_valid = validate_response_schema(&task, &response); if !is_valid {
error!(
"Task {} response validation against JSON schema failed",
task_id
);
if let Ok(mut write_task) = task.write() {
write_task.set_status(TaskStatus::Failed);
if let Ok(tracker) = event_tracker.write() {
tracker.record_task_failed(
&write_task.id,
"Response JSON schema validation failed",
&write_task.prompt,
);
}
}
return;
}
if let Ok(mut write_task) = task.write() {
write_task.set_status(TaskStatus::Completed);
write_task.set_result(response.clone());
if let Ok(tracker) = event_tracker.write() {
tracker.record_task_completed(&write_task.id, &write_task.prompt, response);
}
}
}
Err(e) => {
error!("Task {} failed: {}", task_id, e);
if let Ok(mut write_task) = task.write() {
write_task.set_status(TaskStatus::Failed);
if let Ok(tracker) = event_tracker.write() {
tracker.record_task_failed(
&write_task.id,
&e.to_string(),
&write_task.prompt,
);
}
}
}
}
})
}
fn get_parameters_from_context(task: Arc<RwLock<Task>>) -> (String, Vec<String>, String, Provider) {
let (task_id, dependencies, agent_id, provider) = {
let task_guard = task.read().unwrap();
(
task_guard.id.clone(),
task_guard.dependencies.clone(),
task_guard.agent_id.clone(),
task_guard.prompt.provider.clone(),
)
};
(task_id, dependencies, agent_id, provider)
}
fn spawn_task_executions(
workflow: &Arc<RwLock<Workflow>>,
ready_tasks: Vec<Arc<RwLock<Task>>>,
) -> Result<Vec<tokio::task::JoinHandle<()>>, WorkflowError> {
let mut handles = Vec::with_capacity(ready_tasks.len());
let event_tracker = workflow.read().unwrap().event_tracker.clone();
for task in ready_tasks {
let (task_id, dependencies, agent_id, provider) = get_parameters_from_context(task.clone());
mark_task_as_running(task.clone(), &event_tracker);
let (context, parameter_context, global_context) =
build_task_context(workflow, &dependencies, &provider)?;
let agent = get_agent_for_task(workflow, &agent_id)?;
let handle = spawn_task_execution(
event_tracker.clone(),
task.clone(),
task_id,
agent,
context,
parameter_context,
global_context,
);
handles.push(handle);
}
Ok(handles)
}
async fn await_task_completions(handles: Vec<tokio::task::JoinHandle<()>>) {
for handle in handles {
if let Err(e) = handle.await {
warn!("Task execution failed: {}", e);
}
}
}
#[instrument(skip_all)]
pub async fn execute_workflow(workflow: &Arc<RwLock<Workflow>>) -> Result<(), WorkflowError> {
debug!("Starting workflow execution");
while !is_workflow_complete(workflow) {
reset_failed_workflow_tasks(workflow)?;
let ready_tasks = get_ready_tasks(workflow);
debug!("Found {} ready tasks for execution", ready_tasks.len());
if ready_tasks.is_empty() {
if check_for_circular_dependencies(workflow) {
break;
}
continue;
}
let handles = spawn_task_executions(workflow, ready_tasks)?;
await_task_completions(handles).await;
}
debug!("Workflow execution completed");
Ok(())
}
impl Serialize for Workflow {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut state = serializer.serialize_struct("Workflow", 4)?;
state.serialize_field("id", &self.id)?;
state.serialize_field("name", &self.name)?;
state.serialize_field("task_list", &self.task_list)?;
state.serialize_field("agents", &self.agents)?;
state.end()
}
}
impl<'de> Deserialize<'de> for Workflow {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(field_identifier, rename_all = "snake_case")]
enum Field {
Id,
Name,
TaskList,
Agents,
}
struct WorkflowVisitor;
impl<'de> Visitor<'de> for WorkflowVisitor {
type Value = Workflow;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("struct Workflow")
}
fn visit_map<V>(self, mut map: V) -> Result<Workflow, V::Error>
where
V: MapAccess<'de>,
{
let mut id = None;
let mut name = None;
let mut task_list_data = None;
let mut agents: Option<HashMap<String, Agent>> = None;
while let Some(key) = map.next_key()? {
match key {
Field::Id => {
let value: String = map.next_value().map_err(|e| {
error!("Failed to deserialize field 'id': {e}");
de::Error::custom(format!("Failed to deserialize field 'id': {e}"))
})?;
id = Some(value);
}
Field::TaskList => {
let value: TaskList = map.next_value().map_err(|e| {
error!("Failed to deserialize field 'task_list': {e}");
de::Error::custom(format!(
"Failed to deserialize field 'task_list': {e}",
))
})?;
task_list_data = Some(value);
}
Field::Name => {
let value: String = map.next_value().map_err(|e| {
error!("Failed to deserialize field 'name': {e}");
de::Error::custom(format!(
"Failed to deserialize field 'name': {e}",
))
})?;
name = Some(value);
}
Field::Agents => {
let value: HashMap<String, Agent> = map.next_value().map_err(|e| {
error!("Failed to deserialize field 'agents': {e}");
de::Error::custom(format!(
"Failed to deserialize field 'agents': {e}"
))
})?;
agents = Some(value);
}
}
}
let id = id.ok_or_else(|| de::Error::missing_field("id"))?;
let name = name.ok_or_else(|| de::Error::missing_field("name"))?;
let task_list_data =
task_list_data.ok_or_else(|| de::Error::missing_field("task_list"))?;
let agents = agents.ok_or_else(|| de::Error::missing_field("agents"))?;
let event_tracker = Arc::new(RwLock::new(EventTracker::new(create_uuid7())));
let agents = agents
.into_iter()
.map(|(id, agent)| (id, Arc::new(agent)))
.collect();
Ok(Workflow {
id,
name,
task_list: task_list_data,
agents,
event_tracker,
global_context: None, })
}
}
const FIELDS: &[&str] = &["id", "name", "task_list", "agents"];
deserializer.deserialize_struct("Workflow", FIELDS, WorkflowVisitor)
}
}
#[pyclass(skip_from_py_object, name = "Workflow")]
#[derive(Debug, Clone)]
pub struct PyWorkflow {
workflow: Workflow,
output_types: HashMap<String, Arc<Py<PyAny>>>,
}
#[pymethods]
impl PyWorkflow {
#[new]
#[pyo3(signature = (name))]
pub fn new(name: &str) -> Result<Self, WorkflowError> {
debug!("Creating new workflow: {}", name);
Ok(Self {
workflow: Workflow::new(name),
output_types: HashMap::new(),
})
}
#[getter]
pub fn name(&self) -> String {
self.workflow.name.clone()
}
#[getter]
pub fn task_list(&self) -> TaskList {
self.workflow.task_list.clone()
}
#[getter]
pub fn is_workflow(&self) -> bool {
true
}
#[getter]
pub fn __workflow__(&self) -> Result<String, WorkflowError> {
self.model_dump_json()
}
#[getter]
pub fn agents(&self) -> Result<HashMap<String, PyAgent>, WorkflowError> {
self.workflow
.agents
.iter()
.map(|(id, agent)| {
Ok((
id.clone(),
PyAgent {
agent: agent.clone(),
},
))
})
.collect::<Result<HashMap<_, _>, _>>()
}
#[pyo3(signature = (task_output_types))]
pub fn add_task_output_types<'py>(
&mut self,
task_output_types: Bound<'py, PyDict>,
) -> PyResult<()> {
let converted: HashMap<String, Arc<Py<PyAny>>> = task_output_types
.iter()
.map(|(k, v)| -> PyResult<(String, Arc<Py<PyAny>>)> {
let key = k.extract::<String>()?;
let value = v.clone().unbind();
Ok((key, Arc::new(value)))
})
.collect::<PyResult<_>>()?;
self.output_types.extend(converted);
Ok(())
}
#[pyo3(signature = (task, output_type = None))]
pub fn add_task(
&mut self,
py: Python<'_>,
mut task: Task,
output_type: Option<Bound<'_, PyAny>>,
) -> Result<(), WorkflowError> {
if let Some(output_type) = output_type {
let (response_type, response_json_schema) = parse_response_to_json(py, &output_type)
.map_err(|e| WorkflowError::InvalidOutputType(e.to_string()))?;
task.prompt
.set_response_json_schema(response_json_schema, response_type);
self.output_types
.insert(task.id.clone(), Arc::new(output_type.unbind()));
}
self.workflow.task_list.add_task(task)?;
Ok(())
}
pub fn add_tasks(&mut self, tasks: Vec<Task>) -> Result<(), WorkflowError> {
for task in tasks {
self.workflow.task_list.add_task(task)?;
}
Ok(())
}
pub fn add_agent(&mut self, agent: &Bound<'_, PyAgent>) {
let agent = agent.extract::<PyAgent>().unwrap().agent.clone();
self.workflow.agents.insert(agent.id.clone(), agent);
}
pub fn add_agents(&mut self, agents: Vec<Bound<'_, PyAgent>>) {
for agent in agents {
self.add_agent(&agent);
}
}
pub fn is_complete(&self) -> bool {
self.workflow.task_list.is_complete()
}
pub fn pending_count(&self) -> usize {
self.workflow.task_list.pending_count()
}
pub fn execution_plan<'py>(&self, py: Python<'py>) -> Result<Bound<'py, PyAny>, WorkflowError> {
let plan = self.workflow.execution_plan()?;
debug!("Execution plan: {:?}", plan);
let json = serde_json::to_value(plan).map_err(|e| {
error!("Failed to serialize execution plan to JSON: {}", e);
e
})?;
Ok(pythonize(py, &json)?)
}
#[pyo3(signature = (global_context=None))]
pub fn run(
&self,
py: Python,
global_context: Option<Bound<'_, PyAny>>,
) -> Result<WorkflowResult, WorkflowError> {
debug!("Running workflow: {}", self.workflow.name);
let global_context = if let Some(context) = global_context {
let json_value = depythonize_object_to_value(py, &context)?;
Some(json_value)
} else {
None
};
let workflow: Arc<RwLock<Workflow>> =
block_on(async { self.workflow.run(global_context).await })?;
let workflow_result = match Arc::try_unwrap(workflow) {
Ok(rwlock) => {
let workflow = rwlock
.into_inner()
.map_err(|_| WorkflowError::LockAcquireError)?;
let events = workflow
.event_tracker
.read()
.unwrap()
.events
.read()
.unwrap()
.clone();
WorkflowResult::new(
py,
workflow.task_list.tasks(),
&self.output_types,
events,
workflow.task_list.get_last_task_id(),
)
}
Err(arc) => {
error!("Workflow still has other references, reading instead of consuming.");
let workflow = arc
.read()
.map_err(|_| WorkflowError::ReadLockAcquireError)?;
let events = workflow
.event_tracker
.read()
.unwrap()
.events
.read()
.unwrap()
.clone();
WorkflowResult::new(
py,
workflow.task_list.tasks(),
&self.output_types,
events,
workflow.task_list.get_last_task_id(),
)
}
};
info!("Workflow execution completed successfully.");
Ok(workflow_result)
}
#[pyo3(signature = (task_id, context=None))]
pub fn execute_task<'py>(
&self,
py: Python<'py>,
task_id: String,
context: Option<Bound<'py, PyAny>>,
) -> Result<Bound<'py, PyAny>, WorkflowError> {
let context_value = if let Some(ctx) = context {
depythonize_object_to_value(py, &ctx)?
} else {
Value::Null
};
let response_value =
block_on(async { self.workflow.execute_task(&task_id, &context_value).await })?;
let py_response = pythonize(py, &response_value)?;
Ok(py_response)
}
pub fn model_dump_json(&self) -> Result<String, WorkflowError> {
Ok(self.workflow.serialize()?)
}
#[staticmethod]
#[pyo3(signature = (json_string, output_types=None))]
pub fn model_validate_json(
json_string: String,
output_types: Option<Bound<'_, PyDict>>,
) -> Result<Self, WorkflowError> {
let mut workflow: Workflow = Workflow::from_json(&json_string)?;
workflow.task_list.rebuild_task_validators()?;
block_on(async { workflow.reset_agents().await })?;
let output_types = match output_types {
Some(output_types) => output_types
.iter()
.map(|(k, v)| -> PyResult<(String, Arc<Py<PyAny>>)> {
let key = k.extract::<String>()?;
let value = v.clone().unbind();
Ok((key, Arc::new(value)))
})
.collect::<PyResult<HashMap<String, Arc<Py<PyAny>>>>>()?,
None => HashMap::new(),
};
let py_workflow = PyWorkflow {
workflow,
output_types,
};
Ok(py_workflow)
}
}
#[cfg(test)]
mod tests {
use super::*;
use potato_type::openai::v1::chat::request::{
ChatMessage as OpenAIChatMessage, ContentPart, TextContentPart,
};
use potato_type::prompt::Prompt;
use potato_type::prompt::ResponseType;
fn create_openai_chat_message_num() -> MessageNum {
let text_part = TextContentPart::new("What company is this logo from?".to_string());
let text_content_part = ContentPart::Text(text_part);
let text_message = OpenAIChatMessage {
role: "user".to_string(),
content: vec![text_content_part],
name: None,
};
MessageNum::OpenAIMessageV1(text_message)
}
#[test]
fn test_workflow_creation() {
let workflow = Workflow::new("Test Workflow");
assert_eq!(workflow.name, "Test Workflow");
assert_eq!(workflow.id.len(), 36); }
#[test]
fn test_task_list_add_and_get() {
let mut task_list = TaskList::new();
let prompt = Prompt::new_rs(
vec![create_openai_chat_message_num()],
"gpt-4o",
potato_type::Provider::OpenAI,
vec![],
None,
None,
ResponseType::Null,
)
.unwrap();
let task = Task::new("task1", prompt, "task1", None, None).unwrap();
task_list.add_task(task.clone()).unwrap();
assert_eq!(
task_list.get_task(&task.id).unwrap().read().unwrap().id,
task.id
);
task_list.reset_failed_tasks().unwrap();
}
}