diff --git a/backend/tests/worker.rs b/backend/tests/worker.rs index 506ce3e848..de4268866c 100644 --- a/backend/tests/worker.rs +++ b/backend/tests/worker.rs @@ -28,7 +28,7 @@ use windmill_api_client::types::{EditSchedule, NewSchedule, ScriptArgs}; use windmill_common::worker::{WORKER_CONFIG, PriorityTags}; use windmill_common::{ - flow_status::{FlowStatus, FlowStatusModule}, + flow_status::{FlowStatus, FlowStatusModule, RestartedFrom}, flows::{FlowModule, FlowModuleValue, FlowValue, InputTransform}, jobs::{JobPayload, RawCode, JobKind}, scripts::{ScriptLang, ScriptHash} @@ -2610,6 +2610,196 @@ async fn test_flow_lock_all(db: Pool) { }); } +#[sqlx::test(fixtures("base"))] + +async fn test_complex_flow_restart(db: Pool) { + initialize_tracing().await; + let server = ApiServer::start(db.clone()).await; + let port = server.addr.port(); + + let flow: FlowValue = serde_json::from_value(json!({ + "modules": [ + { + "id": "a", + "value": { + "type": "rawscript", + "language": "go", + "content": "package inner\nimport (\n\t\"fmt\"\n\t\"math/rand\"\n)\nfunc main(max int) (interface{}, error) {\n\tresult := rand.Intn(max) + 1\n\tfmt.Printf(\"Number generated: '%d'\", result)\n\treturn result, nil\n}", + "input_transforms": { + "max": { + "type": "static", + "value": json!(20), + }, + } + }, + "summary": "Generate random number in [1, 20]" + }, + { + "id": "b", + "value": + { + "type": "branchall", + "branches": + [ + { + "modules": + [ + { + "id": "d", + "value": + { + "type": "branchone", + "default": + [ + { + "id": "f", + "value": + { + "type": "rawscript", + "content": "package inner\nimport \"math/rand\"\nfunc main(x int) (interface{}, error) {\n\treturn rand.Intn(x) + 1, nil\n}", + "language": "go", + "input_transforms": + { + "x": + { + "expr": "results.a", + "type": "javascript" + } + } + }, + "summary": "Rand N in [1; x]" + } + ], + "branches": + [ + { + "expr": "results.a < flow_input.max / 2", + "modules": + [ + { + "id": "e", + "value": + { + "type": "rawscript", + "content": "package inner\nimport \"math/rand\"\nfunc main(x int) (interface{}, error) {\n\treturn rand.Intn(x * 2) + 1, nil\n}\n", + "language": "go", + "input_transforms": + { + "x": + { + "expr": "results.a", + "type": "javascript" + } + } + }, + "summary": "Rand N in [1; x*2]" + } + ], + "summary": "N in first half" + } + ] + }, + "summary": "" + } + ], + "summary": "Process x", + "parallel": true, + "skip_failure": false + }, + { + "modules": + [ + { + "id": "c", + "value": + { + "type": "rawscript", + "content": "package inner\nfunc main(x int) (interface{}, error) {\n\treturn x, nil\n}", + "language": "go", + "input_transforms": + { + "x": + { + "expr": "results.a", + "type": "javascript" + } + } + }, + "summary": "Identity" + } + ], + "summary": "Do nothing", + "parallel": true, + "skip_failure": false + } + ], + "parallel": false + }, + "summary": "" + }, + { + "id": "g", + "value": + { + "tag": "", + "type": "rawscript", + "content": "package inner\nimport \"fmt\"\nfunc main(x []int) (interface{}, error) {\n\tfmt.Printf(\"Results: %v\", x)\n\treturn x, nil\n}\n", + "language": "go", + "input_transforms": + { + "x": + { + "expr": "results.b", + "type": "javascript" + } + } + }, + "summary": "Print results - This will get the results from the prior step directly" + }, + { + "id": "h", + "value": + { + "tag": "", + "type": "rawscript", + "content": "package inner\nimport (\n\t\"fmt\"\n\t\"slices\"\n)\nfunc main(x []int) (interface{}, error) {\n\tresult := slices.Max(x)\n\tfmt.Printf(\"Result is %d\", result)\n\treturn result, nil\n}", + "language": "go", + "input_transforms": + { + "x": + { + "expr": "results.b", + "type": "javascript" + } + } + }, + "summary": "Choose max - this will get results.b querying get_result_by_id on the backend" + } + ], + })) + .unwrap(); + + let first_run_result = RunJob::from(JobPayload::RawFlow { + value: flow.clone(), + path: None, + restarted_from: None + }).run_until_complete(&db, port).await; + + let restarted_flow_result = RunJob::from(JobPayload::RawFlow { + value: flow.clone(), + path: None, + restarted_from: Some(RestartedFrom { + flow_job_id: first_run_result.id, + step_id: "h".to_owned(), + }), + }).run_until_complete(&db, port).await; + + assert_eq!( + first_run_result.json_result().unwrap(), + restarted_flow_result.json_result().unwrap() + ); +} + #[sqlx::test(fixtures("base"))] async fn test_rust_client(db: Pool) { initialize_tracing().await; diff --git a/backend/windmill-queue/src/jobs.rs b/backend/windmill-queue/src/jobs.rs index 3d7f39600b..6e2ce42ed8 100644 --- a/backend/windmill-queue/src/jobs.rs +++ b/backend/windmill-queue/src/jobs.rs @@ -6,7 +6,11 @@ * LICENSE-AGPL for a copy of the license. */ -use std::{collections::HashMap, iter, sync::Arc, vec}; +use std::{ + collections::{HashMap, HashSet}, + sync::Arc, + vec, +}; use anyhow::Context; use async_recursion::async_recursion; @@ -19,7 +23,6 @@ use axum::{ }; use bigdecimal::ToPrimitive; use chrono::{DateTime, Duration, Utc}; -use itertools::Itertools; use prometheus::IntCounter; use reqwest::{ header::{HeaderMap, CONTENT_TYPE}, @@ -1681,19 +1684,18 @@ async fn compute_leaf_jobs_for_completed_flow( // we add the module as an element of ListJob for this step ID and recursiively extract leaf job of the sub-flow match module { FlowStatusModule::Success { flow_jobs: Some(jobs), .. } => { - for job in jobs { - if *job == child_job_id { - let new_list_job = match recursive_result.get(&module.id()) { - Some(JobResult::ListJob(jobs_list)) => jobs_list - .into_iter() - .chain(iter::once(&child_job_id)) - .cloned() - .collect_vec(), - _ => iter::once(&child_job_id).cloned().collect_vec(), - }; - recursive_result - .insert(module.id(), JobResult::ListJob(new_list_job)); - } + let jobs_set: HashSet = HashSet::from_iter(jobs.iter().cloned()); + if jobs_set.contains(&child_job_id) { + let new_list_job = match recursive_result.get(&module.id()) { + Some(JobResult::ListJob(jobs_list)) => { + let mut jobs_list_c = jobs_list.clone(); + jobs_list_c.push(child_job_id); + jobs_list_c + } + _ => vec![child_job_id], + }; + recursive_result + .insert(module.id(), JobResult::ListJob(new_list_job)); } } _ => {}