1use 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
20const MAX_PLAN_DEPTH: u32 = 10;
22
23#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
25#[derive(Debug, Default, Deserialize, Serialize)]
26pub struct PlanWorkflowRequest {
27 #[cfg_attr(feature = "openapi", schema(value_type = Option<std::collections::HashMap<String, serde_json::Value>>))]
29 #[serde(default)]
30 pub payload: Option<Value>,
31 #[serde(default)]
33 pub max_depth: Option<u32>,
34 #[serde(default)]
36 pub estimate_durations: Option<bool>,
37}
38
39#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
41#[derive(Debug, Serialize)]
42pub struct ConditionResponse {
43 pub state: String,
45 #[serde(skip_serializing_if = "Option::is_none")]
47 pub expression: Option<String>,
48 #[serde(skip_serializing_if = "Option::is_none")]
50 pub value: Option<bool>,
51 #[serde(skip_serializing_if = "Option::is_none")]
53 pub reason: Option<String>,
54}
55
56#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
58#[derive(Debug, Serialize)]
59pub struct PlannedStepResponse {
60 pub name: String,
62 pub kind: String,
64 pub workflow: String,
66 pub depth: u32,
68 pub depends_on: Vec<String>,
70 #[serde(skip_serializing_if = "Option::is_none")]
72 pub condition: Option<ConditionResponse>,
73 #[serde(skip_serializing_if = "Option::is_none")]
75 pub parallel_group: Option<String>,
76 #[serde(skip_serializing_if = "Option::is_none")]
78 pub estimated_duration_ms: Option<u64>,
79}
80
81#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
83#[derive(Debug, Serialize)]
84pub struct ExecutionPlanResponse {
85 pub workflow: String,
87 pub steps: Vec<PlannedStepResponse>,
89 #[serde(skip_serializing_if = "Option::is_none")]
91 pub estimated_duration_ms: Option<u64>,
92 pub max_depth: u32,
94 pub truncated: bool,
96 #[serde(skip_serializing_if = "Option::is_none")]
98 pub incomplete_reason: Option<String>,
99}
100
101fn 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#[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}