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::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#[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}