Skip to main content

ironflow_api/routes/
plan_workflow.rs

1//! `POST /api/v1/workflows/:name/plan` — Build a workflow execution plan.
2
3use axum::Json;
4use axum::extract::{Path, State};
5use axum::response::IntoResponse;
6use ironflow_auth::extractor::Authenticated;
7use ironflow_engine::error::EngineError;
8use ironflow_engine::plan::{
9    ConditionResult, DEFAULT_ESTIMATE_SAMPLE_RUNS, DEFAULT_PLAN_MAX_DEPTH, ExecutionPlan,
10    PlanOptions, PlannedStep,
11};
12use ironflow_store::entities::StepKind;
13use serde::{Deserialize, Serialize};
14use serde_json::{Value, json};
15
16use crate::error::ApiError;
17use crate::response::ok;
18use crate::state::AppState;
19
20/// Highest sub-workflow expansion depth the API accepts.
21const MAX_PLAN_DEPTH: u32 = 10;
22
23/// Request body for building a workflow execution plan.
24#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
25#[derive(Debug, Default, Deserialize, Serialize)]
26pub struct PlanWorkflowRequest {
27    /// Input payload the plan is computed for. Defaults to `{}`.
28    #[cfg_attr(feature = "openapi", schema(value_type = Option<std::collections::HashMap<String, serde_json::Value>>))]
29    #[serde(default)]
30    pub payload: Option<Value>,
31    /// How deep sub-workflows are expanded. Defaults to 3, capped at 10.
32    #[serde(default)]
33    pub max_depth: Option<u32>,
34    /// Estimate step durations from run history. Defaults to `true`.
35    #[serde(default)]
36    pub estimate_durations: Option<bool>,
37}
38
39/// Outcome of a branch condition as recorded by the planner.
40#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
41#[derive(Debug, Serialize)]
42pub struct ConditionResponse {
43    /// Condition state: `evaluated`, `skipped` or `unevaluable`.
44    pub state: String,
45    /// Expression the handler declared, when the planner knows one.
46    #[serde(skip_serializing_if = "Option::is_none")]
47    pub expression: Option<String>,
48    /// What the expression evaluated to, for an `evaluated` condition.
49    #[serde(skip_serializing_if = "Option::is_none")]
50    pub value: Option<bool>,
51    /// Why the step is skipped, or why the condition cannot be evaluated.
52    #[serde(skip_serializing_if = "Option::is_none")]
53    pub reason: Option<String>,
54}
55
56/// One step the planner expects the run to create.
57#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
58#[derive(Debug, Serialize)]
59pub struct PlannedStepResponse {
60    /// Step name as the handler declares it.
61    pub name: String,
62    /// Step kind, or the name of a custom operation.
63    pub kind: String,
64    /// Workflow that owns this step.
65    pub workflow: String,
66    /// Sub-workflow nesting depth; `0` for the top-level workflow.
67    pub depth: u32,
68    /// Names of the steps this one runs after.
69    pub depends_on: Vec<String>,
70    /// Branch condition recorded for this step, when the handler declared one.
71    #[serde(skip_serializing_if = "Option::is_none")]
72    pub condition: Option<ConditionResponse>,
73    /// Parallel wave this step belongs to, when it runs concurrently.
74    #[serde(skip_serializing_if = "Option::is_none")]
75    pub parallel_group: Option<String>,
76    /// Average duration of this step in past completed runs, in milliseconds.
77    #[serde(skip_serializing_if = "Option::is_none")]
78    pub estimated_duration_ms: Option<u64>,
79}
80
81/// The execution plan of one workflow for one input payload.
82#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
83#[derive(Debug, Serialize)]
84pub struct ExecutionPlanResponse {
85    /// Workflow the plan was built for.
86    pub workflow: String,
87    /// Steps the run is expected to create, in execution order.
88    pub steps: Vec<PlannedStepResponse>,
89    /// Sum of the step estimates, counting each parallel wave once.
90    #[serde(skip_serializing_if = "Option::is_none")]
91    pub estimated_duration_ms: Option<u64>,
92    /// Sub-workflow expansion depth used for this plan.
93    pub max_depth: u32,
94    /// `true` when the step cap or the depth limit cut the plan short.
95    pub truncated: bool,
96    /// Why the plan stopped early, when it did.
97    #[serde(skip_serializing_if = "Option::is_none")]
98    pub incomplete_reason: Option<String>,
99}
100
101/// Render a [`StepKind`] the way the rest of the API does: a plain string.
102fn kind_label(kind: &StepKind) -> String {
103    match kind {
104        StepKind::Shell => "shell".to_string(),
105        StepKind::Http => "http".to_string(),
106        StepKind::Agent => "agent".to_string(),
107        StepKind::Workflow => "workflow".to_string(),
108        StepKind::Approval => "approval".to_string(),
109        StepKind::Decision => "decision".to_string(),
110        StepKind::HumanInput => "human_input".to_string(),
111        StepKind::Custom(name) => name.clone(),
112    }
113}
114
115impl From<ConditionResult> for ConditionResponse {
116    fn from(condition: ConditionResult) -> Self {
117        match condition {
118            ConditionResult::Evaluated { expression, value } => Self {
119                state: "evaluated".to_string(),
120                expression: Some(expression),
121                value: Some(value),
122                reason: None,
123            },
124            ConditionResult::Skipped { reason } => Self {
125                state: "skipped".to_string(),
126                expression: None,
127                value: None,
128                reason: Some(reason),
129            },
130            ConditionResult::Unevaluable { expression, reason } => Self {
131                state: "unevaluable".to_string(),
132                expression: Some(expression),
133                value: None,
134                reason: Some(reason),
135            },
136        }
137    }
138}
139
140impl From<PlannedStep> for PlannedStepResponse {
141    fn from(step: PlannedStep) -> Self {
142        let estimated_duration_ms = step.estimated_duration_ms();
143        Self {
144            name: step.name,
145            kind: kind_label(&step.kind),
146            workflow: step.workflow,
147            depth: step.depth,
148            depends_on: step.depends_on,
149            condition: step.condition.map(ConditionResponse::from),
150            parallel_group: step.parallel_group,
151            estimated_duration_ms,
152        }
153    }
154}
155
156impl From<ExecutionPlan> for ExecutionPlanResponse {
157    fn from(plan: ExecutionPlan) -> Self {
158        let estimated_duration_ms = plan.estimated_duration_ms();
159        Self {
160            workflow: plan.workflow,
161            steps: plan
162                .steps
163                .into_iter()
164                .map(PlannedStepResponse::from)
165                .collect(),
166            estimated_duration_ms,
167            max_depth: plan.max_depth,
168            truncated: plan.truncated,
169            incomplete_reason: plan.incomplete_reason,
170        }
171    }
172}
173
174/// Build a workflow execution plan without running it.
175///
176/// # Errors
177///
178/// - 400 if `max_depth` is out of range or the payload is not a JSON object
179/// - 404 if the workflow is not registered
180#[cfg_attr(
181    feature = "openapi",
182    utoipa::path(
183        post,
184        path = "/api/v1/workflows/{name}/plan",
185        tags = ["workflows"],
186        params(("name" = String, Path, description = "Workflow name")),
187        request_body(content = PlanWorkflowRequest, description = "Input payload and planning options"),
188        responses(
189            (status = 200, description = "Execution plan", body = ExecutionPlanResponse),
190            (status = 400, description = "Invalid payload or max_depth"),
191            (status = 401, description = "Unauthorized"),
192            (status = 404, description = "Workflow not found")
193        ),
194        security(("Bearer" = []))
195    )
196)]
197pub async fn plan_workflow(
198    _auth: Authenticated,
199    State(state): State<AppState>,
200    Path(name): Path<String>,
201    Json(req): Json<PlanWorkflowRequest>,
202) -> Result<impl IntoResponse, ApiError> {
203    if !state.engine.handler_names().contains(&name.as_str()) {
204        return Err(ApiError::WorkflowNotFound(name));
205    }
206
207    let max_depth = req.max_depth.unwrap_or(DEFAULT_PLAN_MAX_DEPTH);
208    if max_depth == 0 || max_depth > MAX_PLAN_DEPTH {
209        return Err(ApiError::BadRequest(format!(
210            "max_depth must be between 1 and {MAX_PLAN_DEPTH}"
211        )));
212    }
213
214    let payload = req.payload.unwrap_or_else(|| json!({}));
215    if !payload.is_object() {
216        return Err(ApiError::BadRequest(
217            "payload must be a JSON object".to_string(),
218        ));
219    }
220
221    let options = PlanOptions {
222        max_depth,
223        estimate_durations: req.estimate_durations.unwrap_or(true),
224        sample_runs: DEFAULT_ESTIMATE_SAMPLE_RUNS,
225    };
226
227    let plan = state
228        .engine
229        .plan_handler(&name, payload, options)
230        .await
231        .map_err(|e| match e {
232            EngineError::InvalidWorkflow(msg) => ApiError::BadRequest(msg),
233            EngineError::Store(err) => ApiError::from(err),
234            other => ApiError::Internal(other.to_string()),
235        })?;
236
237    Ok(ok(ExecutionPlanResponse::from(plan)))
238}
239
240#[cfg(test)]
241mod tests {
242    use std::sync::Arc;
243
244    use axum::Router;
245    use axum::body::Body;
246    use axum::http::{Request, StatusCode};
247    use axum::response::Response;
248    use axum::routing::post;
249    use http_body_util::BodyExt;
250    use ironflow_auth::jwt::{AccessToken, JwtConfig};
251    use ironflow_core::providers::claude::ClaudeCodeProvider;
252    use ironflow_engine::config::{ShellConfig, StepConfig};
253    use ironflow_engine::context::WorkflowContext;
254    use ironflow_engine::engine::Engine;
255    use ironflow_engine::handler::{HandlerFuture, WorkflowHandler};
256    use ironflow_engine::notify::Event;
257    use ironflow_store::memory::InMemoryStore;
258    use ironflow_store::models::RunFilter;
259    use serde_json::{Value as JsonValue, from_slice};
260    use tokio::sync::broadcast;
261    use tower::ServiceExt;
262    use uuid::Uuid;
263
264    use super::*;
265
266    struct PlannedWorkflow;
267
268    impl WorkflowHandler for PlannedWorkflow {
269        fn name(&self) -> &str {
270            "planned"
271        }
272
273        fn execute<'a>(&'a self, ctx: &'a mut WorkflowContext) -> HandlerFuture<'a> {
274            Box::pin(async move {
275                ctx.shell("build", ShellConfig::new("echo build")).await?;
276                ctx.parallel(
277                    vec![
278                        ("test", StepConfig::Shell(ShellConfig::new("echo test"))),
279                        ("lint", StepConfig::Shell(ShellConfig::new("echo lint"))),
280                    ],
281                    true,
282                )
283                .await?;
284                Ok(())
285            })
286        }
287    }
288
289    #[derive(Deserialize)]
290    struct DeployInput {
291        env: String,
292    }
293
294    struct ConditionalWorkflow;
295
296    impl WorkflowHandler for ConditionalWorkflow {
297        fn name(&self) -> &str {
298            "conditional"
299        }
300
301        fn execute<'a>(&'a self, ctx: &'a mut WorkflowContext) -> HandlerFuture<'a> {
302            Box::pin(async move {
303                if ctx
304                    .when("production run", |i: &DeployInput| i.env == "prod")
305                    .await?
306                {
307                    ctx.shell("deploy", ShellConfig::new("echo deploy")).await?;
308                } else {
309                    ctx.skip("deploy", "not prod").await?;
310                }
311                Ok(())
312            })
313        }
314    }
315
316    fn test_state() -> AppState {
317        let store = Arc::new(InMemoryStore::new());
318        let provider = Arc::new(ClaudeCodeProvider::new());
319        let mut engine = Engine::new(store.clone(), provider);
320        engine.register(PlannedWorkflow).unwrap();
321        engine.register(ConditionalWorkflow).unwrap();
322        let jwt_config = Arc::new(JwtConfig {
323            secret: "test-secret".to_string(),
324            access_token_ttl_secs: 900,
325            refresh_token_ttl_secs: 604800,
326            cookie_domain: None,
327            cookie_secure: false,
328        });
329        let (event_sender, _) = broadcast::channel::<Event>(1);
330        AppState::new(
331            store,
332            Arc::new(engine),
333            jwt_config,
334            "test-worker-token".to_string(),
335            event_sender,
336        )
337    }
338
339    fn make_auth_header(state: &AppState) -> String {
340        let user_id = Uuid::now_v7();
341        let token = AccessToken::for_user(user_id, "testuser", false, &state.jwt_config).unwrap();
342        format!("Bearer {}", token.0)
343    }
344
345    fn app(state: AppState) -> Router {
346        Router::new()
347            .route("/api/v1/workflows/{name}/plan", post(plan_workflow))
348            .with_state(state)
349    }
350
351    fn plan_request(name: &str, auth: Option<&str>, body: JsonValue) -> Request<Body> {
352        let mut builder = Request::builder()
353            .method("POST")
354            .uri(format!("/api/v1/workflows/{name}/plan"))
355            .header("content-type", "application/json");
356        if let Some(header) = auth {
357            builder = builder.header("authorization", header);
358        }
359        builder.body(Body::from(body.to_string())).unwrap()
360    }
361
362    async fn body_json(response: Response) -> JsonValue {
363        let bytes = response.into_body().collect().await.unwrap().to_bytes();
364        from_slice(&bytes).unwrap()
365    }
366
367    #[tokio::test]
368    async fn plan_returns_steps_for_registered_workflow() {
369        let state = test_state();
370        let auth = make_auth_header(&state);
371        let response = app(state)
372            .oneshot(plan_request("planned", Some(&auth), json!({})))
373            .await
374            .unwrap();
375
376        assert_eq!(response.status(), StatusCode::OK);
377        let body = body_json(response).await;
378        assert_eq!(body["data"]["workflow"], "planned");
379        let steps = body["data"]["steps"].as_array().unwrap();
380        assert_eq!(steps.len(), 3);
381        assert_eq!(steps[0]["name"], "build");
382        assert_eq!(steps[0]["kind"], "shell");
383        assert_eq!(steps[1]["name"], "test");
384        assert_eq!(steps[2]["name"], "lint");
385        assert_eq!(body["data"]["truncated"], false);
386    }
387
388    #[tokio::test]
389    async fn plan_marks_parallel_group() {
390        let state = test_state();
391        let auth = make_auth_header(&state);
392        let response = app(state)
393            .oneshot(plan_request("planned", Some(&auth), json!({})))
394            .await
395            .unwrap();
396
397        let body = body_json(response).await;
398        let steps = body["data"]["steps"].as_array().unwrap();
399        assert!(steps[0].get("parallel_group").is_none());
400        assert_eq!(steps[1]["parallel_group"], "parallel-1");
401        assert_eq!(steps[2]["parallel_group"], "parallel-1");
402    }
403
404    #[tokio::test]
405    async fn plan_evaluates_condition_from_payload() {
406        let state = test_state();
407        let auth = make_auth_header(&state);
408        let router = app(state);
409
410        let response = router
411            .clone()
412            .oneshot(plan_request(
413                "conditional",
414                Some(&auth),
415                json!({"payload": {"env": "prod"}}),
416            ))
417            .await
418            .unwrap();
419        let body = body_json(response).await;
420        let step = &body["data"]["steps"][0];
421        assert_eq!(step["kind"], "shell");
422        assert_eq!(step["condition"]["state"], "evaluated");
423        assert_eq!(step["condition"]["expression"], "production run");
424        assert_eq!(step["condition"]["value"], true);
425
426        let response = router
427            .oneshot(plan_request(
428                "conditional",
429                Some(&auth),
430                json!({"payload": {"env": "dev"}}),
431            ))
432            .await
433            .unwrap();
434        let body = body_json(response).await;
435        let step = &body["data"]["steps"][0];
436        assert_eq!(step["kind"], "skip");
437        assert_eq!(step["condition"]["state"], "skipped");
438        assert_eq!(step["condition"]["reason"], "not prod");
439    }
440
441    #[tokio::test]
442    async fn plan_unknown_workflow_returns_404() {
443        let state = test_state();
444        let auth = make_auth_header(&state);
445        let response = app(state)
446            .oneshot(plan_request("nonexistent", Some(&auth), json!({})))
447            .await
448            .unwrap();
449
450        assert_eq!(response.status(), StatusCode::NOT_FOUND);
451    }
452
453    #[tokio::test]
454    async fn plan_rejects_zero_and_excessive_max_depth() {
455        let state = test_state();
456        let auth = make_auth_header(&state);
457        let router = app(state);
458
459        let response = router
460            .clone()
461            .oneshot(plan_request(
462                "planned",
463                Some(&auth),
464                json!({"max_depth": 0}),
465            ))
466            .await
467            .unwrap();
468        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
469
470        let response = router
471            .oneshot(plan_request(
472                "planned",
473                Some(&auth),
474                json!({"max_depth": 11}),
475            ))
476            .await
477            .unwrap();
478        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
479    }
480
481    #[tokio::test]
482    async fn plan_rejects_non_object_payload() {
483        let state = test_state();
484        let auth = make_auth_header(&state);
485        let response = app(state)
486            .oneshot(plan_request(
487                "planned",
488                Some(&auth),
489                json!({"payload": [1, 2, 3]}),
490            ))
491            .await
492            .unwrap();
493
494        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
495    }
496
497    #[tokio::test]
498    async fn plan_requires_authentication() {
499        let state = test_state();
500        let response = app(state)
501            .oneshot(plan_request("planned", None, json!({})))
502            .await
503            .unwrap();
504
505        assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
506    }
507
508    #[tokio::test]
509    async fn plan_does_not_create_a_run() {
510        let state = test_state();
511        let store = state.store.clone();
512        let auth = make_auth_header(&state);
513        let response = app(state)
514            .oneshot(plan_request("planned", Some(&auth), json!({})))
515            .await
516            .unwrap();
517        assert_eq!(response.status(), StatusCode::OK);
518
519        let runs = store.list_runs(RunFilter::default(), 1, 10).await.unwrap();
520        assert!(runs.items.is_empty());
521    }
522}