Skip to main content

roder_dynamic_workflows/
execution.rs

1mod executor;
2mod plan;
3mod script_stack;
4mod state;
5mod task;
6
7use std::collections::{BTreeMap, VecDeque};
8use std::sync::Arc;
9use std::sync::atomic::{AtomicBool, Ordering};
10
11use roder_api::dynamic_workflows::{
12    WorkflowApproval, WorkflowCheckpointRecorded, WorkflowRun, WorkflowRunId, WorkflowRunResumed,
13    WorkflowRunStatus, WorkflowScript,
14};
15use roder_api::events::{RoderEvent, ThreadId, TurnId};
16use roder_api::subagents::{SubagentExitReason, SubagentRequest, SubagentResult};
17use time::OffsetDateTime;
18use tokio::sync::{Mutex, Notify, broadcast};
19use tokio::task::{JoinHandle, JoinSet};
20
21pub use self::task::WorkflowTaskExecutor;
22use crate::host_api::{WorkflowAgentLaunch, WorkflowExecution};
23use crate::model::{WorkflowRunInput, WorkflowRuntimeOptions};
24use crate::runner::WorkflowScriptRuntime;
25use crate::store::{WorkflowCachedAgentResult, WorkflowCheckpointStore};
26
27pub use self::executor::{SubagentDispatcherWorkflowExecutor, WorkflowAgentExecutor};
28use self::plan::{
29    PlannedAgent, execution_context, mark_restarted_agent_completed, phases_for_execution,
30    plan_agents, render_final_report,
31};
32use self::script_stack::run_script_on_dedicated_stack;
33use self::state::{
34    mark_agent_completed, mark_agent_error, mark_agent_failed, mark_agent_started,
35    mark_phase_completed, mark_phase_started, mark_run_completed, mark_run_paused,
36    mark_run_started, mark_run_stopped,
37};
38
39#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
40#[serde(rename_all = "camelCase")]
41pub struct WorkflowRunRequest {
42    pub run_id: WorkflowRunId,
43    pub thread_id: Option<ThreadId>,
44    pub turn_id: Option<TurnId>,
45    pub script: WorkflowScript,
46    pub arguments: serde_json::Value,
47    pub start_paused: bool,
48    #[serde(default, skip_serializing_if = "Option::is_none")]
49    pub approval: Option<WorkflowApproval>,
50}
51
52#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq)]
53#[serde(rename_all = "camelCase")]
54pub struct WorkflowRunSnapshot {
55    pub run: WorkflowRun,
56    pub report: Option<String>,
57    pub reused_agent_results: u32,
58}
59
60#[derive(Debug, Clone)]
61pub struct WorkflowAgentExecutionContext {
62    pub run_id: WorkflowRunId,
63    pub phase_id: String,
64    pub agent_id: String,
65    pub thread_id: Option<ThreadId>,
66    pub turn_id: Option<TurnId>,
67    pub stopped: Arc<AtomicBool>,
68}
69
70#[derive(Debug, Clone)]
71pub struct WorkflowAgentExecutionRequest {
72    pub launch: WorkflowAgentLaunch,
73    pub subagent_request: SubagentRequest,
74    pub cache_key: crate::store::WorkflowAgentCacheKey,
75}
76
77#[derive(Clone)]
78pub struct WorkflowRunner {
79    executor: Arc<dyn WorkflowAgentExecutor>,
80    script_runtime: WorkflowScriptRuntime,
81    store: WorkflowCheckpointStore,
82    events: broadcast::Sender<RoderEvent>,
83}
84
85impl WorkflowRunner {
86    pub fn new(
87        executor: Arc<dyn WorkflowAgentExecutor>,
88        store: WorkflowCheckpointStore,
89        options: WorkflowRuntimeOptions,
90    ) -> Self {
91        let (events, _) = broadcast::channel(1024);
92        Self {
93            executor,
94            script_runtime: WorkflowScriptRuntime::new(options),
95            store,
96            events,
97        }
98    }
99
100    pub fn subscribe(&self) -> broadcast::Receiver<RoderEvent> {
101        self.events.subscribe()
102    }
103
104    pub async fn start(&self, request: WorkflowRunRequest) -> anyhow::Result<WorkflowRunHandle> {
105        let source = request
106            .script
107            .body
108            .as_deref()
109            .ok_or_else(|| anyhow::anyhow!("workflow script body is required to run"))?;
110        let checkpoints = self.store.read_checkpoints(&request.run_id)?;
111        let mut input = WorkflowRunInput::new(request.run_id.clone());
112        input.arguments = request.arguments.clone();
113        input.checkpoints = checkpoints;
114        let execution =
115            run_script_on_dedicated_stack(self.script_runtime.clone(), source.to_string(), input)
116                .await?;
117        let planned = plan_agents(
118            &request.run_id,
119            &request.thread_id,
120            &request.turn_id,
121            &execution,
122        );
123        let state = Arc::new(Mutex::new(WorkflowRunSnapshot {
124            run: initial_run(&request, &execution, &planned),
125            report: None,
126            reused_agent_results: 0,
127        }));
128        let control = Arc::new(WorkflowRunControl::new(request.start_paused));
129        let planned = Arc::new(planned);
130        let join = spawn_run(
131            self.clone(),
132            state.clone(),
133            control.clone(),
134            planned.clone(),
135            execution,
136        );
137
138        Ok(WorkflowRunHandle {
139            run_id: request.run_id,
140            state,
141            control,
142            planned,
143            executor: self.executor.clone(),
144            store: self.store.clone(),
145            events: self.events.clone(),
146            join: Arc::new(Mutex::new(Some(join))),
147        })
148    }
149
150    async fn run_to_completion(
151        self,
152        state: Arc<Mutex<WorkflowRunSnapshot>>,
153        control: Arc<WorkflowRunControl>,
154        planned: Arc<BTreeMap<String, PlannedAgent>>,
155        execution: WorkflowExecution,
156    ) -> anyhow::Result<WorkflowRunSnapshot> {
157        mark_run_started(&self.events, &state).await;
158        if control.paused.load(Ordering::SeqCst) {
159            mark_run_paused(&self.events, &state, Some("started paused".to_string())).await;
160        }
161        self.persist_script_checkpoints(&state, &execution).await?;
162
163        let phases = { state.lock().await.run.phases.clone() };
164        let mut ordered_results = Vec::new();
165        for phase in phases {
166            if self.wait_if_paused_or_stopped(&state, &control).await? {
167                return Ok(state.lock().await.clone());
168            }
169            mark_phase_started(&self.events, &state, &phase.phase_id).await;
170            ordered_results.extend(
171                self.run_phase_agents(&state, &control, &planned, &phase.phase_id)
172                    .await?,
173            );
174            mark_phase_completed(&self.events, &state, &phase.phase_id).await;
175        }
176
177        let report = render_final_report(&execution.report, &ordered_results);
178        mark_run_completed(&self.events, &state, report).await;
179        Ok(state.lock().await.clone())
180    }
181
182    async fn run_phase_agents(
183        &self,
184        state: &Arc<Mutex<WorkflowRunSnapshot>>,
185        control: &Arc<WorkflowRunControl>,
186        planned: &BTreeMap<String, PlannedAgent>,
187        phase_id: &str,
188    ) -> anyhow::Result<Vec<(String, SubagentResult)>> {
189        let max_concurrent = state.lock().await.run.limits.max_concurrent_agents.max(1) as usize;
190        let mut queue = planned
191            .values()
192            .filter(|agent| agent.phase_id == phase_id)
193            .cloned()
194            .collect::<VecDeque<_>>();
195        let mut active = JoinSet::new();
196        let mut results = Vec::new();
197
198        loop {
199            while active.len() < max_concurrent && !queue.is_empty() {
200                if self.wait_if_paused_or_stopped(state, control).await? {
201                    active.abort_all();
202                    return Ok(results);
203                }
204                let next = queue.pop_front().expect("queue is not empty");
205                let run_id = { state.lock().await.run.run_id.clone() };
206                if let Some(cached) = self.store.find_agent_result(&run_id, &next.cache_key)? {
207                    mark_agent_completed(
208                        &self.events,
209                        state,
210                        &next.agent_id,
211                        cached.result.clone(),
212                        true,
213                    )
214                    .await;
215                    results.push((next.agent_id, cached.result));
216                    continue;
217                }
218                mark_agent_started(&self.events, state, &next.agent_id).await;
219                active.spawn(execute_planned_agent(
220                    self.executor.clone(),
221                    state.clone(),
222                    control.stopped.clone(),
223                    next,
224                ));
225            }
226
227            if active.is_empty() {
228                break;
229            }
230            let (agent_id, result) = active.join_next().await.expect("active task exists")?;
231            self.record_agent_result(state, planned, &agent_id, result, &mut results)
232                .await?;
233            if control.stopped.load(Ordering::SeqCst) {
234                active.abort_all();
235                mark_run_stopped(&self.events, state, Some("stopped".to_string())).await;
236                return Ok(results);
237            }
238        }
239
240        Ok(results)
241    }
242
243    async fn record_agent_result(
244        &self,
245        state: &Arc<Mutex<WorkflowRunSnapshot>>,
246        planned: &BTreeMap<String, PlannedAgent>,
247        agent_id: &str,
248        result: anyhow::Result<SubagentResult>,
249        results: &mut Vec<(String, SubagentResult)>,
250    ) -> anyhow::Result<()> {
251        match result {
252            Ok(result) if result.exit_reason == SubagentExitReason::Completed => {
253                if let Some(planned) = planned.get(agent_id) {
254                    let run_id = { state.lock().await.run.run_id.clone() };
255                    self.store.append_agent_result(
256                        &run_id,
257                        &WorkflowCachedAgentResult {
258                            key: planned.cache_key.clone(),
259                            result: result.clone(),
260                            completed_at: OffsetDateTime::now_utc(),
261                        },
262                    )?;
263                }
264                mark_agent_completed(&self.events, state, agent_id, result.clone(), false).await;
265                results.push((agent_id.to_string(), result));
266            }
267            Ok(result) => {
268                mark_agent_failed(
269                    &self.events,
270                    state,
271                    agent_id,
272                    result.clone(),
273                    format!("subagent exited with {:?}", result.exit_reason),
274                )
275                .await;
276                results.push((agent_id.to_string(), result));
277            }
278            Err(error) => {
279                mark_agent_error(&self.events, state, agent_id, error.to_string()).await;
280            }
281        }
282        Ok(())
283    }
284
285    async fn persist_script_checkpoints(
286        &self,
287        state: &Arc<Mutex<WorkflowRunSnapshot>>,
288        execution: &WorkflowExecution,
289    ) -> anyhow::Result<()> {
290        for checkpoint in &execution.checkpoints {
291            let run = state.lock().await.run.clone();
292            self.store.append_checkpoint(&run.run_id, checkpoint)?;
293            let _ = self.events.send(RoderEvent::WorkflowCheckpointRecorded(
294                WorkflowCheckpointRecorded {
295                    run_id: run.run_id,
296                    thread_id: run.thread_id,
297                    turn_id: run.turn_id,
298                    phase_id: None,
299                    key: checkpoint.key.clone(),
300                    byte_count: checkpoint.byte_count,
301                    timestamp: OffsetDateTime::now_utc(),
302                },
303            ));
304        }
305        Ok(())
306    }
307
308    async fn wait_if_paused_or_stopped(
309        &self,
310        state: &Arc<Mutex<WorkflowRunSnapshot>>,
311        control: &Arc<WorkflowRunControl>,
312    ) -> anyhow::Result<bool> {
313        while control.paused.load(Ordering::SeqCst) {
314            if control.stopped.load(Ordering::SeqCst) {
315                mark_run_stopped(&self.events, state, Some("stopped".to_string())).await;
316                return Ok(true);
317            }
318            control.notify.notified().await;
319        }
320        if control.stopped.load(Ordering::SeqCst) {
321            mark_run_stopped(&self.events, state, Some("stopped".to_string())).await;
322            return Ok(true);
323        }
324        Ok(false)
325    }
326}
327
328pub struct WorkflowRunHandle {
329    run_id: WorkflowRunId,
330    state: Arc<Mutex<WorkflowRunSnapshot>>,
331    control: Arc<WorkflowRunControl>,
332    planned: Arc<BTreeMap<String, PlannedAgent>>,
333    executor: Arc<dyn WorkflowAgentExecutor>,
334    store: WorkflowCheckpointStore,
335    events: broadcast::Sender<RoderEvent>,
336    join: Arc<Mutex<Option<JoinHandle<anyhow::Result<WorkflowRunSnapshot>>>>>,
337}
338
339impl WorkflowRunHandle {
340    pub fn run_id(&self) -> &str {
341        &self.run_id
342    }
343
344    pub async fn snapshot(&self) -> WorkflowRunSnapshot {
345        self.state.lock().await.clone()
346    }
347
348    pub async fn wait(&self) -> anyhow::Result<WorkflowRunSnapshot> {
349        let join = self.join.lock().await.take();
350        if let Some(join) = join {
351            return join.await?;
352        }
353        Ok(self.snapshot().await)
354    }
355
356    pub async fn pause(&self, reason: Option<String>) {
357        self.control.paused.store(true, Ordering::SeqCst);
358        mark_run_paused(&self.events, &self.state, reason).await;
359    }
360
361    pub async fn resume(&self) {
362        self.control.paused.store(false, Ordering::SeqCst);
363        self.control.notify.notify_waiters();
364        let mut snapshot = self.state.lock().await;
365        snapshot.run.status = WorkflowRunStatus::Running;
366        let run = snapshot.run.clone();
367        drop(snapshot);
368        let _ = self
369            .events
370            .send(RoderEvent::WorkflowRunResumed(WorkflowRunResumed {
371                run_id: run.run_id,
372                thread_id: run.thread_id,
373                turn_id: run.turn_id,
374                status: WorkflowRunStatus::Running,
375                timestamp: OffsetDateTime::now_utc(),
376            }));
377    }
378
379    pub async fn stop(&self, reason: Option<String>) {
380        self.control.stopped.store(true, Ordering::SeqCst);
381        self.control.paused.store(false, Ordering::SeqCst);
382        self.control.notify.notify_waiters();
383        mark_run_stopped(&self.events, &self.state, reason).await;
384    }
385
386    pub async fn restart_agent(&self, agent_id: &str) -> anyhow::Result<bool> {
387        let Some(planned) = self.planned.get(agent_id).cloned() else {
388            return Ok(false);
389        };
390        self.store
391            .invalidate_agent_results(&self.run_id, agent_id)?;
392        let context = execution_context(&self.state, &planned, self.control.stopped.clone()).await;
393        let result = self
394            .executor
395            .execute_agent(
396                context,
397                WorkflowAgentExecutionRequest {
398                    launch: planned.launch.clone(),
399                    subagent_request: planned.request.clone(),
400                    cache_key: planned.cache_key.clone(),
401                },
402            )
403            .await?;
404        self.store.append_agent_result(
405            &self.run_id,
406            &WorkflowCachedAgentResult {
407                key: planned.cache_key,
408                result: result.clone(),
409                completed_at: OffsetDateTime::now_utc(),
410            },
411        )?;
412        mark_restarted_agent_completed(&self.state, agent_id, result).await;
413        Ok(true)
414    }
415}
416
417struct WorkflowRunControl {
418    paused: AtomicBool,
419    stopped: Arc<AtomicBool>,
420    notify: Notify,
421}
422
423impl WorkflowRunControl {
424    fn new(paused: bool) -> Self {
425        Self {
426            paused: AtomicBool::new(paused),
427            stopped: Arc::new(AtomicBool::new(false)),
428            notify: Notify::new(),
429        }
430    }
431}
432
433fn spawn_run(
434    runner: WorkflowRunner,
435    state: Arc<Mutex<WorkflowRunSnapshot>>,
436    control: Arc<WorkflowRunControl>,
437    planned: Arc<BTreeMap<String, PlannedAgent>>,
438    execution: WorkflowExecution,
439) -> JoinHandle<anyhow::Result<WorkflowRunSnapshot>> {
440    tokio::spawn(async move {
441        runner
442            .run_to_completion(state, control, planned, execution)
443            .await
444    })
445}
446
447async fn execute_planned_agent(
448    executor: Arc<dyn WorkflowAgentExecutor>,
449    state: Arc<Mutex<WorkflowRunSnapshot>>,
450    stopped: Arc<AtomicBool>,
451    planned: PlannedAgent,
452) -> (String, anyhow::Result<SubagentResult>) {
453    let context = execution_context(&state, &planned, stopped).await;
454    let request = WorkflowAgentExecutionRequest {
455        launch: planned.launch,
456        subagent_request: planned.request,
457        cache_key: planned.cache_key,
458    };
459    let agent_id = context.agent_id.clone();
460    (agent_id, executor.execute_agent(context, request).await)
461}
462
463fn initial_run(
464    request: &WorkflowRunRequest,
465    execution: &WorkflowExecution,
466    planned: &BTreeMap<String, PlannedAgent>,
467) -> WorkflowRun {
468    let now = OffsetDateTime::now_utc();
469    WorkflowRun {
470        run_id: request.run_id.clone(),
471        thread_id: request.thread_id.clone(),
472        turn_id: request.turn_id.clone(),
473        script: request.script.clone(),
474        status: WorkflowRunStatus::Queued,
475        limits: execution.definition.limits.clone(),
476        phases: phases_for_execution(execution),
477        agents: planned
478            .values()
479            .map(|planned| planned.agent.clone())
480            .collect(),
481        approval: request.approval.clone(),
482        cost_estimate: None,
483        summary: None,
484        error: None,
485        created_at: now,
486        updated_at: now,
487        started_at: None,
488        completed_at: None,
489    }
490}