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::Custom(name) => name.clone(),
111    }
112}
113
114impl From<ConditionResult> for ConditionResponse {
115    fn from(condition: ConditionResult) -> Self {
116        match condition {
117            ConditionResult::Evaluated { expression, value } => Self {
118                state: "evaluated".to_string(),
119                expression: Some(expression),
120                value: Some(value),
121                reason: None,
122            },
123            ConditionResult::Skipped { reason } => Self {
124                state: "skipped".to_string(),
125                expression: None,
126                value: None,
127                reason: Some(reason),
128            },
129            ConditionResult::Unevaluable { expression, reason } => Self {
130                state: "unevaluable".to_string(),
131                expression: Some(expression),
132                value: None,
133                reason: Some(reason),
134            },
135        }
136    }
137}
138
139impl From<PlannedStep> for PlannedStepResponse {
140    fn from(step: PlannedStep) -> Self {
141        let estimated_duration_ms = step.estimated_duration_ms();
142        Self {
143            name: step.name,
144            kind: kind_label(&step.kind),
145            workflow: step.workflow,
146            depth: step.depth,
147            depends_on: step.depends_on,
148            condition: step.condition.map(ConditionResponse::from),
149            parallel_group: step.parallel_group,
150            estimated_duration_ms,
151        }
152    }
153}
154
155impl From<ExecutionPlan> for ExecutionPlanResponse {
156    fn from(plan: ExecutionPlan) -> Self {
157        let estimated_duration_ms = plan.estimated_duration_ms();
158        Self {
159            workflow: plan.workflow,
160            steps: plan
161                .steps
162                .into_iter()
163                .map(PlannedStepResponse::from)
164                .collect(),
165            estimated_duration_ms,
166            max_depth: plan.max_depth,
167            truncated: plan.truncated,
168            incomplete_reason: plan.incomplete_reason,
169        }
170    }
171}
172
173/// Build a workflow execution plan without running it.
174///
175/// # Errors
176///
177/// - 400 if `max_depth` is out of range or the payload is not a JSON object
178/// - 404 if the workflow is not registered
179#[cfg_attr(
180    feature = "openapi",
181    utoipa::path(
182        post,
183        path = "/api/v1/workflows/{name}/plan",
184        tags = ["workflows"],
185        params(("name" = String, Path, description = "Workflow name")),
186        request_body(content = PlanWorkflowRequest, description = "Input payload and planning options"),
187        responses(
188            (status = 200, description = "Execution plan", body = ExecutionPlanResponse),
189            (status = 400, description = "Invalid payload or max_depth"),
190            (status = 401, description = "Unauthorized"),
191            (status = 404, description = "Workflow not found")
192        ),
193        security(("Bearer" = []))
194    )
195)]
196pub async fn plan_workflow(
197    _auth: Authenticated,
198    State(state): State<AppState>,
199    Path(name): Path<String>,
200    Json(req): Json<PlanWorkflowRequest>,
201) -> Result<impl IntoResponse, ApiError> {
202    if !state.engine.handler_names().contains(&name.as_str()) {
203        return Err(ApiError::WorkflowNotFound(name));
204    }
205
206    let max_depth = req.max_depth.unwrap_or(DEFAULT_PLAN_MAX_DEPTH);
207    if max_depth == 0 || max_depth > MAX_PLAN_DEPTH {
208        return Err(ApiError::BadRequest(format!(
209            "max_depth must be between 1 and {MAX_PLAN_DEPTH}"
210        )));
211    }
212
213    let payload = req.payload.unwrap_or_else(|| json!({}));
214    if !payload.is_object() {
215        return Err(ApiError::BadRequest(
216            "payload must be a JSON object".to_string(),
217        ));
218    }
219
220    let options = PlanOptions {
221        max_depth,
222        estimate_durations: req.estimate_durations.unwrap_or(true),
223        sample_runs: DEFAULT_ESTIMATE_SAMPLE_RUNS,
224    };
225
226    let plan = state
227        .engine
228        .plan_handler(&name, payload, options)
229        .await
230        .map_err(|e| match e {
231            EngineError::InvalidWorkflow(msg) => ApiError::BadRequest(msg),
232            EngineError::Store(err) => ApiError::from(err),
233            other => ApiError::Internal(other.to_string()),
234        })?;
235
236    Ok(ok(ExecutionPlanResponse::from(plan)))
237}
238
239#[cfg(test)]
240mod tests {
241    use std::sync::Arc;
242
243    use axum::Router;
244    use axum::body::Body;
245    use axum::http::{Request, StatusCode};
246    use axum::response::Response;
247    use axum::routing::post;
248    use http_body_util::BodyExt;
249    use ironflow_auth::jwt::{AccessToken, JwtConfig};
250    use ironflow_core::providers::claude::ClaudeCodeProvider;
251    use ironflow_engine::config::{ShellConfig, StepConfig};
252    use ironflow_engine::context::WorkflowContext;
253    use ironflow_engine::engine::Engine;
254    use ironflow_engine::handler::{HandlerFuture, WorkflowHandler};
255    use ironflow_engine::notify::Event;
256    use ironflow_store::memory::InMemoryStore;
257    use ironflow_store::models::RunFilter;
258    use serde_json::{Value as JsonValue, from_slice};
259    use tokio::sync::broadcast;
260    use tower::ServiceExt;
261    use uuid::Uuid;
262
263    use super::*;
264
265    struct PlannedWorkflow;
266
267    impl WorkflowHandler for PlannedWorkflow {
268        fn name(&self) -> &str {
269            "planned"
270        }
271
272        fn execute<'a>(&'a self, ctx: &'a mut WorkflowContext) -> HandlerFuture<'a> {
273            Box::pin(async move {
274                ctx.shell("build", ShellConfig::new("echo build")).await?;
275                ctx.parallel(
276                    vec![
277                        ("test", StepConfig::Shell(ShellConfig::new("echo test"))),
278                        ("lint", StepConfig::Shell(ShellConfig::new("echo lint"))),
279                    ],
280                    true,
281                )
282                .await?;
283                Ok(())
284            })
285        }
286    }
287
288    struct ConditionalWorkflow;
289
290    impl WorkflowHandler for ConditionalWorkflow {
291        fn name(&self) -> &str {
292            "conditional"
293        }
294
295        fn execute<'a>(&'a self, ctx: &'a mut WorkflowContext) -> HandlerFuture<'a> {
296            Box::pin(async move {
297                if ctx.when("env == prod", |p| p["env"] == "prod").await? {
298                    ctx.shell("deploy", ShellConfig::new("echo deploy")).await?;
299                } else {
300                    ctx.skip("deploy", "not prod").await?;
301                }
302                Ok(())
303            })
304        }
305    }
306
307    fn test_state() -> AppState {
308        let store = Arc::new(InMemoryStore::new());
309        let provider = Arc::new(ClaudeCodeProvider::new());
310        let mut engine = Engine::new(store.clone(), provider);
311        engine.register(PlannedWorkflow).unwrap();
312        engine.register(ConditionalWorkflow).unwrap();
313        let jwt_config = Arc::new(JwtConfig {
314            secret: "test-secret".to_string(),
315            access_token_ttl_secs: 900,
316            refresh_token_ttl_secs: 604800,
317            cookie_domain: None,
318            cookie_secure: false,
319        });
320        let (event_sender, _) = broadcast::channel::<Event>(1);
321        AppState::new(
322            store,
323            Arc::new(engine),
324            jwt_config,
325            "test-worker-token".to_string(),
326            event_sender,
327        )
328    }
329
330    fn make_auth_header(state: &AppState) -> String {
331        let user_id = Uuid::now_v7();
332        let token = AccessToken::for_user(user_id, "testuser", false, &state.jwt_config).unwrap();
333        format!("Bearer {}", token.0)
334    }
335
336    fn app(state: AppState) -> Router {
337        Router::new()
338            .route("/api/v1/workflows/{name}/plan", post(plan_workflow))
339            .with_state(state)
340    }
341
342    fn plan_request(name: &str, auth: Option<&str>, body: JsonValue) -> Request<Body> {
343        let mut builder = Request::builder()
344            .method("POST")
345            .uri(format!("/api/v1/workflows/{name}/plan"))
346            .header("content-type", "application/json");
347        if let Some(header) = auth {
348            builder = builder.header("authorization", header);
349        }
350        builder.body(Body::from(body.to_string())).unwrap()
351    }
352
353    async fn body_json(response: Response) -> JsonValue {
354        let bytes = response.into_body().collect().await.unwrap().to_bytes();
355        from_slice(&bytes).unwrap()
356    }
357
358    #[tokio::test]
359    async fn plan_returns_steps_for_registered_workflow() {
360        let state = test_state();
361        let auth = make_auth_header(&state);
362        let response = app(state)
363            .oneshot(plan_request("planned", Some(&auth), json!({})))
364            .await
365            .unwrap();
366
367        assert_eq!(response.status(), StatusCode::OK);
368        let body = body_json(response).await;
369        assert_eq!(body["data"]["workflow"], "planned");
370        let steps = body["data"]["steps"].as_array().unwrap();
371        assert_eq!(steps.len(), 3);
372        assert_eq!(steps[0]["name"], "build");
373        assert_eq!(steps[0]["kind"], "shell");
374        assert_eq!(steps[1]["name"], "test");
375        assert_eq!(steps[2]["name"], "lint");
376        assert_eq!(body["data"]["truncated"], false);
377    }
378
379    #[tokio::test]
380    async fn plan_marks_parallel_group() {
381        let state = test_state();
382        let auth = make_auth_header(&state);
383        let response = app(state)
384            .oneshot(plan_request("planned", Some(&auth), json!({})))
385            .await
386            .unwrap();
387
388        let body = body_json(response).await;
389        let steps = body["data"]["steps"].as_array().unwrap();
390        assert!(steps[0].get("parallel_group").is_none());
391        assert_eq!(steps[1]["parallel_group"], "parallel-1");
392        assert_eq!(steps[2]["parallel_group"], "parallel-1");
393    }
394
395    #[tokio::test]
396    async fn plan_evaluates_condition_from_payload() {
397        let state = test_state();
398        let auth = make_auth_header(&state);
399        let router = app(state);
400
401        let response = router
402            .clone()
403            .oneshot(plan_request(
404                "conditional",
405                Some(&auth),
406                json!({"payload": {"env": "prod"}}),
407            ))
408            .await
409            .unwrap();
410        let body = body_json(response).await;
411        let step = &body["data"]["steps"][0];
412        assert_eq!(step["kind"], "shell");
413        assert_eq!(step["condition"]["state"], "evaluated");
414        assert_eq!(step["condition"]["value"], true);
415
416        let response = router
417            .oneshot(plan_request(
418                "conditional",
419                Some(&auth),
420                json!({"payload": {"env": "dev"}}),
421            ))
422            .await
423            .unwrap();
424        let body = body_json(response).await;
425        let step = &body["data"]["steps"][0];
426        assert_eq!(step["kind"], "skip");
427        assert_eq!(step["condition"]["state"], "skipped");
428        assert_eq!(step["condition"]["reason"], "not prod");
429    }
430
431    #[tokio::test]
432    async fn plan_unknown_workflow_returns_404() {
433        let state = test_state();
434        let auth = make_auth_header(&state);
435        let response = app(state)
436            .oneshot(plan_request("nonexistent", Some(&auth), json!({})))
437            .await
438            .unwrap();
439
440        assert_eq!(response.status(), StatusCode::NOT_FOUND);
441    }
442
443    #[tokio::test]
444    async fn plan_rejects_zero_and_excessive_max_depth() {
445        let state = test_state();
446        let auth = make_auth_header(&state);
447        let router = app(state);
448
449        let response = router
450            .clone()
451            .oneshot(plan_request(
452                "planned",
453                Some(&auth),
454                json!({"max_depth": 0}),
455            ))
456            .await
457            .unwrap();
458        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
459
460        let response = router
461            .oneshot(plan_request(
462                "planned",
463                Some(&auth),
464                json!({"max_depth": 11}),
465            ))
466            .await
467            .unwrap();
468        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
469    }
470
471    #[tokio::test]
472    async fn plan_rejects_non_object_payload() {
473        let state = test_state();
474        let auth = make_auth_header(&state);
475        let response = app(state)
476            .oneshot(plan_request(
477                "planned",
478                Some(&auth),
479                json!({"payload": [1, 2, 3]}),
480            ))
481            .await
482            .unwrap();
483
484        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
485    }
486
487    #[tokio::test]
488    async fn plan_requires_authentication() {
489        let state = test_state();
490        let response = app(state)
491            .oneshot(plan_request("planned", None, json!({})))
492            .await
493            .unwrap();
494
495        assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
496    }
497
498    #[tokio::test]
499    async fn plan_does_not_create_a_run() {
500        let state = test_state();
501        let store = state.store.clone();
502        let auth = make_auth_header(&state);
503        let response = app(state)
504            .oneshot(plan_request("planned", Some(&auth), json!({})))
505            .await
506            .unwrap();
507        assert_eq!(response.status(), StatusCode::OK);
508
509        let runs = store.list_runs(RunFilter::default(), 1, 10).await.unwrap();
510        assert!(runs.items.is_empty());
511    }
512}