Files
windmill/python-client/wmill/tests/test_workflow.py
T
Ruben Fiszel aeaea57ca1 fix(wac): one failure record for tasks and steps, in every round (#10368)
* fix(wac): hand a caught task and step failure the same shape in every round

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* refactor(wac): decide the failure record once, server-side

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): leave a legacy SDK's failure marker untouched, and ship wacError to jsr

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): carry a step's custom error fields, and bound the stack in bytes

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): keep a step's extra fields serializable and bounded

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): record a non-Error throw the way a task records it

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): guard the last unguarded throw site in the step marker

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): make failure reporting non-throwing on both clients

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): take the step traceback the way the executor takes it

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): contain the reads that happen before a failure is checkpointed

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): fall back to the checkpointed marker, not the live one

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): keep non-finite fields and hostile proxies out of the checkpoint path

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* fix(wac): keep the snapshot that passed the serialization probe

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

* chore(wac): keep the failure-record module's surface to what is used

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-28 13:15:24 +02:00

1529 lines
52 KiB
Python

"""Tests for the Workflow-as-Code SDK."""
import asyncio
import json
import pytest
from datetime import datetime, timezone
from wmill.client import WorkflowCtx, _StepSuspend, TaskError, workflow, task, step, sleep, parallel, wait_for_approval, _run_workflow, _run_workflow_async
class _StubInlineClient:
"""Stands in for the httpx client the inline fast path POSTs with.
Decodes each request body, so a test sees exactly the JSON that reaches
``/jobs/wac/inline_checkpoint`` — and therefore what a replay reads back.
"""
def __init__(self):
self.posted = []
async def post(self, url, content=None):
self.posted.append(json.loads(content))
class _Response:
def raise_for_status(self):
pass
return _Response()
async def aclose(self):
pass
def _set_inline_fast_path_env(monkeypatch):
for var, val in (
("WM_JOB_ID", "job-1"),
("WM_WORKSPACE", "admins"),
("BASE_INTERNAL_URL", "http://localhost:8000"),
("WM_TOKEN", "tok"),
):
monkeypatch.setenv(var, val)
@task
async def extract_data(url: str):
pass # body unused in workflow context
@task
async def load_data(data=None):
pass
@task
async def clean_data(data=None):
pass
@task
async def compute_stats(data=None):
pass
@task
async def send_alert(msg: str = ""):
pass
@task
async def double(x: int):
return x * 2
@task
async def add_one(x: int):
return x + 1
@task
async def noop_task():
pass
# --- Module-level workflow definitions ---
@workflow
async def simple_workflow(url: str):
raw = await extract_data(url=url)
result = await load_data(data=raw)
return {"status": "done", "result": result}
@workflow
async def parallel_workflow(url: str):
raw = await extract_data(url=url)
cleaned, stats = await asyncio.gather(
clean_data(data=raw),
compute_stats(data=raw),
)
return {"cleaned": cleaned, "stats": stats}
@workflow
async def conditional_workflow(count: int):
if count > 100:
await send_alert(msg="large")
await load_data()
return {"done": True}
@workflow
async def step_workflow(x: int):
ts = await step("timestamp", lambda: 1234567890)
doubled = await double(x=x)
rid = await step("random_id", lambda: "abc-123")
return {"ts": ts, "doubled": doubled, "id": rid}
# Edge case workflows
@workflow
async def three_step_wf(n: int):
doubled = await double(x=n)
incremented = await add_one(x=doubled)
final = await double(x=incremented)
return {"doubled": doubled, "incremented": incremented, "final": final}
@workflow
async def seq_par_seq_wf(url: str):
raw = await extract_data(url=url)
cleaned, stats = await asyncio.gather(
clean_data(data=raw),
compute_stats(data=raw),
)
loaded = await load_data(data={"cleaned": cleaned, "stats": stats})
return loaded
@workflow
async def double_parallel_wf():
a, b = await asyncio.gather(double(x=1), double(x=2))
c, d = await asyncio.gather(add_one(x=a), add_one(x=b))
return {"a": a, "b": b, "c": c, "d": d}
@workflow
async def cond_on_result_wf():
val = await double(x=5)
if val > 8:
await send_alert(msg="big")
await load_data(data=val)
return {"val": val}
@workflow
async def empty_wf():
return {"status": "empty"}
@workflow
async def single_wf(x: int):
result = await double(x=x)
return result
@workflow
async def no_arg_wf():
result = await noop_task()
return result
@workflow
async def many_steps_wf(n: int):
val = n
for _ in range(10):
val = await add_one(x=val)
return val
@workflow
async def falsy_wf():
a = await double(x=0)
b = await load_data(data=a)
c = await extract_data(url="")
return {"a": a, "b": b, "c": c}
@task(path="f/external_script")
async def run_external(x: int):
return x * 3
@workflow
async def path_wf(x: int):
result = await run_external(x=x)
return result
@workflow
async def mixed_step_task_wf(x: int):
ts = await step("get_time", lambda: 999)
doubled = await double(x=x)
config = await step("get_config", lambda: {"retry": 3})
added = await add_one(x=doubled)
return {"ts": ts, "doubled": doubled, "config": config, "added": added}
@workflow
async def par_child_wf():
a, b = await asyncio.gather(double(x=3), add_one(x=7))
return {"a": a, "b": b}
@workflow
async def det_wf(n: int):
a = await double(x=n)
b = await add_one(x=a)
c = await double(x=b)
return c
@workflow
async def par_args_wf(x: int):
base = await double(x=x)
a, b = await asyncio.gather(add_one(x=base), double(x=base))
return {"a": a, "b": b}
@workflow
async def none_return_wf():
await double(x=1)
@workflow
async def large_par_wf():
results = await asyncio.gather(
double(x=1), double(x=2), double(x=3), double(x=4), double(x=5)
)
return list(results)
@workflow
async def complex_mixed_wf():
init = await extract_data(url="start")
a, b = await asyncio.gather(double(x=1), double(x=2))
mid = await load_data(data={"a": a, "b": b})
c, d = await asyncio.gather(add_one(x=3), add_one(x=4))
fin = await clean_data(data={"mid": mid, "c": c, "d": d})
return fin
@workflow
async def pre_par_child_wf(x: int):
base = await double(x=x)
a, b = await asyncio.gather(add_one(x=base), double(x=base))
return {"a": a, "b": b}
# --- Tests ---
# NOTE: Python SDK uses name-based keys (e.g. "double", "double_2")
# not index-based keys (e.g. "step_0", "step_1").
class TestWorkflowDecorator:
def test_marks_function(self):
assert hasattr(simple_workflow, "_is_workflow")
assert simple_workflow._is_workflow is True
class TestTaskDecorator:
def test_marks_function(self):
assert hasattr(extract_data, "_is_task")
assert extract_data._is_task is True
def test_standalone_execution(self):
"""Outside a workflow, @task runs the function body directly."""
result = asyncio.run(extract_data(url="https://example.com"))
assert result is None # body returns None
def test_preserves_function_name(self):
assert extract_data.__name__ == "extract_data"
assert double.__name__ == "double"
class TestFirstInvocation:
def test_dispatches_first_step(self):
result = _run_workflow(simple_workflow, {}, {"url": "https://example.com"})
assert result["type"] == "dispatch"
assert result["mode"] == "sequential"
assert len(result["steps"]) == 1
assert result["steps"][0]["name"] == "extract_data"
assert result["steps"][0]["script"] == "extract_data"
assert result["steps"][0]["key"] == "extract_data"
assert result["steps"][0]["args"] == {"url": "https://example.com"}
def test_positional_args_converted_to_kwargs(self):
"""Positional args should be mapped to parameter names in dispatch."""
@workflow
async def pos_workflow():
await extract_data("https://pos.example.com")
result = _run_workflow(pos_workflow, {}, {})
assert result["type"] == "dispatch"
assert result["steps"][0]["args"] == {"url": "https://pos.example.com"}
class TestReplayWithCheckpoint:
def test_second_invocation_dispatches_second_step(self):
checkpoint = {
"completed_steps": {
"extract_data": {"data": [1, 2, 3]},
}
}
result = _run_workflow(simple_workflow, checkpoint, {"url": "https://example.com"})
assert result["type"] == "dispatch"
assert result["mode"] == "sequential"
assert result["steps"][0]["name"] == "load_data"
assert result["steps"][0]["key"] == "load_data"
def test_all_steps_complete(self):
checkpoint = {
"completed_steps": {
"extract_data": {"data": [1, 2, 3]},
"load_data": {"loaded": True},
}
}
result = _run_workflow(simple_workflow, checkpoint, {"url": "https://example.com"})
assert result["type"] == "complete"
assert result["result"]["status"] == "done"
assert result["result"]["result"] == {"loaded": True}
class TestParallelDispatch:
def test_first_invocation(self):
result = _run_workflow(parallel_workflow, {}, {"url": "https://example.com"})
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "extract_data"
def test_parallel_dispatch(self):
checkpoint = {
"completed_steps": {
"extract_data": {"raw": "data"},
}
}
result = _run_workflow(parallel_workflow, checkpoint, {"url": "https://example.com"})
assert result["type"] == "dispatch"
assert result["mode"] == "parallel"
assert len(result["steps"]) == 2
assert result["steps"][0]["name"] == "clean_data"
assert result["steps"][1]["name"] == "compute_stats"
def test_parallel_complete(self):
checkpoint = {
"completed_steps": {
"extract_data": {"raw": "data"},
"clean_data": {"cleaned": True},
"compute_stats": {"count": 42},
}
}
result = _run_workflow(parallel_workflow, checkpoint, {"url": "https://example.com"})
assert result["type"] == "complete"
assert result["result"]["cleaned"] == {"cleaned": True}
assert result["result"]["stats"] == {"count": 42}
class TestConditionalWorkflow:
def test_condition_true(self):
result = _run_workflow(conditional_workflow, {}, {"count": 200})
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "send_alert"
def test_condition_false(self):
result = _run_workflow(conditional_workflow, {}, {"count": 50})
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "load_data"
class TestStepInlineCheckpoint:
def test_first_invocation_returns_inline_checkpoint(self):
result = _run_workflow(step_workflow, {}, {"x": 7})
assert result["type"] == "inline_checkpoint"
assert result["key"] == "timestamp"
assert result["result"] == 1234567890
def test_step_cached_then_task_dispatches(self):
checkpoint = {"completed_steps": {"timestamp": 1234567890}}
result = _run_workflow(step_workflow, checkpoint, {"x": 7})
assert result["type"] == "dispatch"
assert result["mode"] == "sequential"
assert result["steps"][0]["name"] == "double"
assert result["steps"][0]["key"] == "double"
def test_step_and_task_cached_then_second_step(self):
checkpoint = {"completed_steps": {"timestamp": 1234567890, "double": 14}}
result = _run_workflow(step_workflow, checkpoint, {"x": 7})
assert result["type"] == "inline_checkpoint"
assert result["key"] == "random_id"
assert result["result"] == "abc-123"
def test_all_complete(self):
checkpoint = {"completed_steps": {"timestamp": 1234567890, "double": 14, "random_id": "abc-123"}}
result = _run_workflow(step_workflow, checkpoint, {"x": 7})
assert result["type"] == "complete"
assert result["result"] == {"ts": 1234567890, "doubled": 14, "id": "abc-123"}
class TestUnawaitedTask:
def test_unawaited_last_task_is_flushed(self):
@workflow
async def unawaited_workflow():
await extract_data(url="x")
load_data(data="y")
checkpoint = {"completed_steps": {"extract_data": "raw"}}
result = _run_workflow(unawaited_workflow, checkpoint, {})
assert result["type"] == "dispatch"
assert result["mode"] == "sequential"
assert len(result["steps"]) == 1
assert result["steps"][0]["name"] == "load_data"
def test_unawaited_multiple_tasks_flushed_as_parallel(self):
@workflow
async def multi_unawaited_workflow():
await extract_data(url="x")
clean_data(data="y")
compute_stats(data="y")
checkpoint = {"completed_steps": {"extract_data": "raw"}}
result = _run_workflow(multi_unawaited_workflow, checkpoint, {})
assert result["type"] == "dispatch"
assert result["mode"] == "parallel"
assert len(result["steps"]) == 2
assert result["steps"][0]["name"] == "clean_data"
assert result["steps"][1]["name"] == "compute_stats"
class TestChildMode:
def test_child_executes_matching_task(self):
checkpoint = {"completed_steps": {"timestamp": 1234567890}, "_executing_key": "double"}
result = _run_workflow(step_workflow, checkpoint, {"x": 7})
assert result["type"] == "complete"
assert result["result"] == 14
def test_child_cannot_swallow_the_failure_of_the_step_it_executes(self):
# If `except Exception` could catch it, the child would report a success
# returning "swallowed" and the parent would record that as the step's value.
@task
async def boom():
raise ValueError("nope")
@workflow
async def wf():
try:
await boom()
except Exception:
return "swallowed"
return "unreachable"
with pytest.raises(ValueError, match="nope"):
_run_workflow(wf, {"_executing_key": "boom"}, {})
def test_child_replays_cached_steps(self):
checkpoint = {
"completed_steps": {"extract_data": {"data": [1, 2, 3]}},
"_executing_key": "load_data",
}
result = _run_workflow(simple_workflow, checkpoint, {"url": "https://example.com"})
assert result["type"] == "complete"
assert result["result"] is None
# =====================================================================
# EDGE CASE TESTS
# =====================================================================
class TestFullSequentialLifecycle:
def test_replay_0_dispatches_step_0(self):
result = _run_workflow(three_step_wf, {}, {"n": 5})
assert result["type"] == "dispatch"
assert result["steps"][0]["key"] == "double"
assert result["steps"][0]["name"] == "double"
assert result["steps"][0]["args"] == {"x": 5}
def test_replay_1_dispatches_step_1_with_step_0_result(self):
result = _run_workflow(three_step_wf, {"completed_steps": {"double": 10}}, {"n": 5})
assert result["type"] == "dispatch"
assert result["steps"][0]["key"] == "add_one"
assert result["steps"][0]["name"] == "add_one"
assert result["steps"][0]["args"] == {"x": 10}
def test_replay_2_dispatches_step_2_with_step_1_result(self):
result = _run_workflow(
three_step_wf, {"completed_steps": {"double": 10, "add_one": 11}}, {"n": 5}
)
assert result["type"] == "dispatch"
assert result["steps"][0]["key"] == "double_2"
assert result["steps"][0]["name"] == "double"
assert result["steps"][0]["args"] == {"x": 11}
def test_replay_3_all_complete(self):
result = _run_workflow(
three_step_wf,
{"completed_steps": {"double": 10, "add_one": 11, "double_2": 22}},
{"n": 5},
)
assert result["type"] == "complete"
assert result["result"] == {"doubled": 10, "incremented": 11, "final": 22}
class TestStepAfterParallelGroup:
def test_dispatches_first_sequential(self):
result = _run_workflow(seq_par_seq_wf, {}, {"url": "http://x"})
assert result["steps"][0]["name"] == "extract_data"
def test_dispatches_parallel_group(self):
result = _run_workflow(
seq_par_seq_wf, {"completed_steps": {"extract_data": "raw"}}, {"url": "http://x"}
)
assert result["mode"] == "parallel"
assert len(result["steps"]) == 2
def test_dispatches_final_step_after_parallel(self):
result = _run_workflow(
seq_par_seq_wf,
{"completed_steps": {"extract_data": "raw", "clean_data": "cleaned", "compute_stats": {"count": 5}}},
{"url": "http://x"},
)
assert result["mode"] == "sequential"
assert result["steps"][0]["name"] == "load_data"
assert result["steps"][0]["key"] == "load_data"
def test_completes_when_final_step_done(self):
result = _run_workflow(
seq_par_seq_wf,
{"completed_steps": {"extract_data": "raw", "clean_data": "cleaned", "compute_stats": {"count": 5}, "load_data": "final"}},
{"url": "http://x"},
)
assert result["type"] == "complete"
assert result["result"] == "final"
class TestParallelAfterParallel:
def test_dispatches_first_parallel(self):
result = _run_workflow(double_parallel_wf, {}, {})
assert result["mode"] == "parallel"
assert len(result["steps"]) == 2
assert result["steps"][0]["key"] == "double"
assert result["steps"][1]["key"] == "double_2"
def test_dispatches_second_parallel(self):
result = _run_workflow(
double_parallel_wf, {"completed_steps": {"double": 2, "double_2": 4}}, {}
)
assert result["mode"] == "parallel"
assert len(result["steps"]) == 2
assert result["steps"][0]["name"] == "add_one"
assert result["steps"][0]["args"] == {"x": 2}
assert result["steps"][1]["args"] == {"x": 4}
def test_completes_all_done(self):
result = _run_workflow(
double_parallel_wf,
{"completed_steps": {"double": 2, "double_2": 4, "add_one": 3, "add_one_2": 5}},
{},
)
assert result["type"] == "complete"
assert result["result"] == {"a": 2, "b": 4, "c": 3, "d": 5}
class TestConditionalBasedOnStepResult:
def test_condition_true_path(self):
result = _run_workflow(cond_on_result_wf, {"completed_steps": {"double": 10}}, {})
assert result["steps"][0]["name"] == "send_alert"
assert result["steps"][0]["key"] == "send_alert"
def test_condition_false_path(self):
result = _run_workflow(cond_on_result_wf, {"completed_steps": {"double": 4}}, {})
assert result["steps"][0]["name"] == "load_data"
assert result["steps"][0]["key"] == "load_data"
def test_condition_true_step_after_alert(self):
result = _run_workflow(
cond_on_result_wf, {"completed_steps": {"double": 10, "send_alert": "alerted"}}, {}
)
assert result["steps"][0]["name"] == "load_data"
assert result["steps"][0]["key"] == "load_data"
class TestEmptyWorkflow:
def test_completes_immediately(self):
result = _run_workflow(empty_wf, {}, {})
assert result["type"] == "complete"
assert result["result"] == {"status": "empty"}
class TestSingleTaskWorkflow:
def test_dispatches_single_step(self):
result = _run_workflow(single_wf, {}, {"x": 7})
assert result["type"] == "dispatch"
assert len(result["steps"]) == 1
assert result["steps"][0]["name"] == "double"
def test_completes_with_result(self):
result = _run_workflow(single_wf, {"completed_steps": {"double": 14}}, {"x": 7})
assert result["type"] == "complete"
assert result["result"] == 14
class TestTaskWithNoArgs:
def test_dispatches_with_empty_args(self):
result = _run_workflow(no_arg_wf, {}, {})
assert result["type"] == "dispatch"
assert result["steps"][0]["args"] == {}
class TestManySteps:
def test_first_dispatches_step_0(self):
result = _run_workflow(many_steps_wf, {}, {"n": 0})
assert result["steps"][0]["key"] == "add_one"
def test_with_5_complete_dispatches_step_5(self):
# add_one, add_one_2, add_one_3, add_one_4, add_one_5
completed = {}
for i in range(5):
key = "add_one" if i == 0 else f"add_one_{i + 1}"
completed[key] = i + 1
result = _run_workflow(many_steps_wf, {"completed_steps": completed}, {"n": 0})
assert result["steps"][0]["key"] == "add_one_6"
assert result["steps"][0]["args"] == {"x": 5}
def test_all_10_complete(self):
completed = {}
for i in range(10):
key = "add_one" if i == 0 else f"add_one_{i + 1}"
completed[key] = i + 1
result = _run_workflow(many_steps_wf, {"completed_steps": completed}, {"n": 0})
assert result["type"] == "complete"
assert result["result"] == 10
class TestFalsyValues:
def test_zero_preserved(self):
result = _run_workflow(falsy_wf, {"completed_steps": {"double": 0}}, {})
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "load_data"
assert result["steps"][0]["args"] == {"data": 0}
def test_none_preserved(self):
result = _run_workflow(falsy_wf, {"completed_steps": {"double": 0, "load_data": None}}, {})
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "extract_data"
def test_all_falsy_complete(self):
result = _run_workflow(
falsy_wf, {"completed_steps": {"double": 0, "load_data": None, "extract_data": ""}}, {}
)
assert result["type"] == "complete"
assert result["result"] == {"a": 0, "b": None, "c": ""}
def test_false_preserved(self):
@workflow
async def flag_wf():
val = await load_data(data="check")
if val:
await send_alert(msg="truthy")
return {"val": val}
result = _run_workflow(flag_wf, {"completed_steps": {"load_data": False}}, {})
assert result["type"] == "complete"
assert result["result"] == {"val": False}
class TestTaskWithExplicitPath:
def test_uses_path_as_script(self):
result = _run_workflow(path_wf, {}, {"x": 42})
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "run_external"
assert result["steps"][0]["script"] == "f/external_script"
assert result["steps"][0]["args"] == {"x": 42}
class TestMixedStepAndTask:
def test_step_0_inline(self):
result = _run_workflow(mixed_step_task_wf, {}, {"x": 5})
assert result["type"] == "inline_checkpoint"
assert result["key"] == "get_time"
assert result["result"] == 999
def test_step_1_task_dispatch(self):
result = _run_workflow(
mixed_step_task_wf, {"completed_steps": {"get_time": 999}}, {"x": 5}
)
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "double"
assert result["steps"][0]["key"] == "double"
def test_step_2_inline(self):
result = _run_workflow(
mixed_step_task_wf,
{"completed_steps": {"get_time": 999, "double": 10}},
{"x": 5},
)
assert result["type"] == "inline_checkpoint"
assert result["key"] == "get_config"
assert result["result"] == {"retry": 3}
def test_step_3_task_dispatch(self):
result = _run_workflow(
mixed_step_task_wf,
{"completed_steps": {"get_time": 999, "double": 10, "get_config": {"retry": 3}}},
{"x": 5},
)
assert result["type"] == "dispatch"
assert result["steps"][0]["name"] == "add_one"
assert result["steps"][0]["key"] == "add_one"
def test_all_complete(self):
result = _run_workflow(
mixed_step_task_wf,
{"completed_steps": {"get_time": 999, "double": 10, "get_config": {"retry": 3}, "add_one": 11}},
{"x": 5},
)
assert result["type"] == "complete"
assert result["result"] == {"ts": 999, "doubled": 10, "config": {"retry": 3}, "added": 11}
class TestChildModeParallel:
def test_child_executes_first_parallel_step(self):
result = _run_workflow(
par_child_wf, {"completed_steps": {}, "_executing_key": "double"}, {}
)
assert result["type"] == "complete"
assert result["result"] == 6
def test_child_executes_second_parallel_step(self):
result = _run_workflow(
par_child_wf, {"completed_steps": {}, "_executing_key": "add_one"}, {}
)
assert result["type"] == "complete"
assert result["result"] == 8
class TestKeyDeterminism:
def test_keys_consistent_across_replays(self):
r1 = _run_workflow(det_wf, {}, {"n": 3})
assert r1["steps"][0]["key"] == "double"
assert r1["steps"][0]["name"] == "double"
r2 = _run_workflow(det_wf, {"completed_steps": {"double": 6}}, {"n": 3})
assert r2["steps"][0]["key"] == "add_one"
assert r2["steps"][0]["name"] == "add_one"
r3 = _run_workflow(det_wf, {"completed_steps": {"double": 6, "add_one": 7}}, {"n": 3})
assert r3["steps"][0]["key"] == "double_2"
assert r3["steps"][0]["name"] == "double"
class TestParallelArgsFromCachedResult:
def test_parallel_steps_receive_cached_args(self):
result = _run_workflow(par_args_wf, {"completed_steps": {"double": 20}}, {"x": 10})
assert result["mode"] == "parallel"
assert result["steps"][0]["args"] == {"x": 20}
assert result["steps"][1]["args"] == {"x": 20}
class TestWorkflowReturningNone:
def test_none_return_captured(self):
result = _run_workflow(none_return_wf, {"completed_steps": {"double": 2}}, {})
assert result["type"] == "complete"
assert result["result"] is None
class TestLargeParallelGroup:
def test_dispatches_5_parallel(self):
result = _run_workflow(large_par_wf, {}, {})
assert result["mode"] == "parallel"
assert len(result["steps"]) == 5
keys = [result["steps"][i]["key"] for i in range(5)]
assert keys == ["double", "double_2", "double_3", "double_4", "double_5"]
for i in range(5):
assert result["steps"][i]["args"] == {"x": i + 1}
class TestComplexMixedWorkflow:
def test_replay_0_extract(self):
r = _run_workflow(complex_mixed_wf, {}, {})
assert r["steps"][0]["name"] == "extract_data"
def test_replay_1_parallel(self):
r = _run_workflow(complex_mixed_wf, {"completed_steps": {"extract_data": "init"}}, {})
assert r["mode"] == "parallel"
assert len(r["steps"]) == 2
def test_replay_2_load(self):
r = _run_workflow(
complex_mixed_wf,
{"completed_steps": {"extract_data": "init", "double": 2, "double_2": 4}},
{},
)
assert r["mode"] == "sequential"
assert r["steps"][0]["name"] == "load_data"
assert r["steps"][0]["key"] == "load_data"
def test_replay_3_second_parallel(self):
r = _run_workflow(
complex_mixed_wf,
{"completed_steps": {"extract_data": "init", "double": 2, "double_2": 4, "load_data": "mid"}},
{},
)
assert r["mode"] == "parallel"
assert len(r["steps"]) == 2
assert r["steps"][0]["name"] == "add_one"
def test_replay_4_clean(self):
r = _run_workflow(
complex_mixed_wf,
{"completed_steps": {
"extract_data": "init", "double": 2, "double_2": 4,
"load_data": "mid", "add_one": 4, "add_one_2": 5,
}},
{},
)
assert r["mode"] == "sequential"
assert r["steps"][0]["name"] == "clean_data"
assert r["steps"][0]["key"] == "clean_data"
def test_replay_5_all_complete(self):
r = _run_workflow(
complex_mixed_wf,
{"completed_steps": {
"extract_data": "init", "double": 2, "double_2": 4,
"load_data": "mid", "add_one": 4, "add_one_2": 5, "clean_data": "final",
}},
{},
)
assert r["type"] == "complete"
assert r["result"] == "final"
class TestChildModeWithCachedStepsBeforeParallel:
def test_child_executes_second_parallel_with_cached_base(self):
result = _run_workflow(
pre_par_child_wf,
{"completed_steps": {"double": 10}, "_executing_key": "double_2"},
{"x": 5},
)
assert result["type"] == "complete"
assert result["result"] == 20
def test_child_executes_first_parallel_with_cached_base(self):
result = _run_workflow(
pre_par_child_wf,
{"completed_steps": {"double": 10}, "_executing_key": "add_one"},
{"x": 5},
)
assert result["type"] == "complete"
assert result["result"] == 11
# =====================================================================
# ERROR PROPAGATION TESTS
# =====================================================================
class TestErrorPropagation:
def test_task_error_is_raised_on_replay(self):
@workflow
async def wf(x: int):
return await double(x=x)
with pytest.raises(TaskError, match="double"):
_run_workflow(
wf,
{
"completed_steps": {
"double": {
"__wmill_error": True,
"message": "Task 'double' failed",
"result": {"message": "boom"},
}
}
},
{"x": 5},
)
def test_error_catchable_with_try_except(self):
@workflow
async def wf(x: int):
try:
result = await double(x=x)
return {"success": True, "result": result}
except Exception as e:
return {"success": False, "error": str(e)}
r = _run_workflow(
wf,
{
"completed_steps": {
"double": {
"__wmill_error": True,
"message": "Task 'double' failed",
"result": {},
}
}
},
{"x": 5},
)
assert r["type"] == "complete"
assert r["result"]["success"] is False
assert "double" in r["result"]["error"]
def test_retry_pattern_with_try_except_loop(self):
@workflow
async def wf(x: int):
for i in range(3):
try:
result = await double(x=x)
return {"result": result, "attempts": i + 1}
except Exception:
if i == 2:
raise
# First double fails, second succeeds
r = _run_workflow(
wf,
{
"completed_steps": {
"double": {"__wmill_error": True, "message": "temporary", "result": {}},
"double_2": 10,
}
},
{"x": 5},
)
assert r["type"] == "complete"
assert r["result"]["result"] == 10
assert r["result"]["attempts"] == 2
def test_non_error_object_with_error_false(self):
@workflow
async def wf():
val = await double(x=5)
return val
r = _run_workflow(
wf,
{"completed_steps": {"double": {"__wmill_error": False, "data": "ok"}}},
{},
)
assert r["type"] == "complete"
assert r["result"] == {"__wmill_error": False, "data": "ok"}
def test_inline_step_error(self):
@workflow
async def wf():
try:
val = await step("risky", lambda: 42)
return {"val": val}
except Exception as e:
return {"caught": str(e)}
r = _run_workflow(
wf,
{"completed_steps": {"risky": {"__wmill_error": True, "message": "step failed", "result": {}}}},
{},
)
assert r["type"] == "complete"
assert "step failed" in r["result"]["caught"]
class TestRaisingInlineStepIsCheckpointed:
"""A ``step()`` whose body raises must still land in ``completed_steps``.
Otherwise a workflow that catches the exception and later dispatches a task
replays with ``_executing_key`` set, reaches the unrecorded key, and parks
on the never-resolving future forever.
"""
# A failed child job reports `{"error": {"name", "message", "stack"}}`, and a
# failed step() has to be indistinguishable from it. The stack is a traceback
# string, asserted separately.
MARKER = {
"__wmill_error": True,
"message": "boom",
"step_key": "risky",
"result": {"error": {"name": "ValueError", "message": "boom"}},
}
@staticmethod
def _without_stack(marker: dict) -> dict:
error = {k: v for k, v in marker["result"]["error"].items() if k != "stack"}
return {**marker, "result": {**marker["result"], "error": error}}
@staticmethod
def _boom():
raise ValueError("boom")
@classmethod
def _wf(cls):
@workflow
async def wf(x: int):
try:
await step("risky", cls._boom)
except Exception:
pass
return await double(x=x)
return wf
def test_first_run_emits_error_checkpoint(self):
r = _run_workflow(self._wf(), {}, {"x": 5})
assert r["type"] == "inline_checkpoint"
assert r["key"] == "risky"
assert self._without_stack(r["result"]) == self.MARKER
stack = r["result"]["result"]["error"]["stack"]
# Frames only and the SDK's own `result = fn()` frame dropped, the way
# the python executor formats a failed job's stack.
assert 'raise ValueError("boom")' in stack
assert "result = fn()" not in stack
assert not stack.startswith("Traceback")
def test_custom_exception_attributes_survive_under_extra(self):
"""A failed child job reports custom attributes under ``error.extra``;
a step dropping them would make the same exception carry less depending
on how it was run."""
from wmill.client import _step_error_marker
class HttpError(ValueError):
def __init__(self):
super().__init__("429")
self.code = 429
error = _step_error_marker("k", HttpError())["result"]["error"]
assert error["extra"] == {"code": 429}
assert error["name"] == "HttpError"
assert "extra" not in _step_error_marker("k", ValueError("plain"))["result"]["error"]
def test_unserializable_attributes_do_not_cost_the_fast_path(self):
"""The fast-path POST serializes strictly, so an exception holding a
live object — ``resp.raise_for_status()`` is the common one — must not
make the marker unserializable and drop the step onto the slow path."""
import json as _json
from wmill.client import _step_error_marker
class Boom(Exception):
def __init__(self):
super().__init__("boom")
self.response = object()
self.status = 429
marker = _step_error_marker("k", Boom())
_json.dumps(marker) # raises if an attribute leaked through unserialized
assert marker["result"]["error"]["extra"]["status"] == 429
def test_an_exception_whose_str_raises_still_reports(self):
"""Every coercion of the user's exception runs inside the ``except``
reporting it, so one that raises would replace their failure with an
unrelated one and leave the step uncheckpointed."""
import json as _json
from wmill.client import _step_error_marker
class Hostile(Exception):
def __str__(self):
raise RuntimeError("cannot be rendered")
marker = _step_error_marker("k", Hostile())
_json.dumps(marker)
assert marker["result"]["error"]["name"] == "Hostile"
assert "unrepresentable" in marker["result"]["error"]["message"]
# ...including one that makes reading its own traceback raise
class HostileTraceback(Exception):
def __getattribute__(self, item):
if item == "__traceback__":
raise RuntimeError("no traceback for you")
return super().__getattribute__(item)
marker = _step_error_marker("k", HostileTraceback())
_json.dumps(marker)
assert marker["result"]["error"]["name"] == "HostileTraceback"
# ...or reading its own attributes
class HostileDict(Exception):
def __getattribute__(self, item):
if item == "__dict__":
raise RuntimeError("no attributes for you")
return super().__getattribute__(item)
marker = _step_error_marker("k", HostileDict())
_json.dumps(marker)
assert marker["result"]["error"]["name"] == "HostileDict"
# A float is serializable, so `default=` never sees NaN — it would go out
# as a bare `NaN` literal, which is not JSON and which the backend
# rejects, so the step could not be checkpointed at all.
class NotFinite(Exception):
def __init__(self):
super().__init__("nan")
self.value = float("nan")
self.limit = float("inf")
marker = _step_error_marker("k", NotFinite())
_json.dumps(marker, allow_nan=False)
assert marker["result"]["error"]["extra"] == {"value": "NaN", "limit": "Infinity"}
def test_fast_path_posts_error_and_raises_the_replay_exception(self, monkeypatch):
"""The default path: the checkpoint is POSTed and the workflow body gets
the same ``TaskError`` a replay rebuilds from the marker — raising the
original ``ValueError`` here would make ``except ValueError:`` catch on
this run and miss on the next one. ``except`` is control flow, so
anything a handler can branch on has to be identical in both rounds."""
_set_inline_fast_path_env(monkeypatch)
class _EchoingStub(_StubInlineClient):
"""The endpoint normalizes the failure before storing it and echoes
back what it stored. The echo deliberately differs from what was
posted, so the assertions below can tell which copy was raised from."""
async def post(self, url, content=None):
await super().post(url, content=content)
stored = {**self.posted[-1]["result"], "message": "normalized by the backend"}
class _Response:
# the endpoint answers with a JSON body; a backend predating
# the echo answers without one, which is how the client tells
# "no echo" from "an echo it could not read"
headers = {"content-type": "application/json"}
def raise_for_status(self):
pass
def json(self):
return {"failure": stored}
return _Response()
stub = _EchoingStub()
posted = stub.posted
async def run():
ctx = WorkflowCtx({})
ctx._inline_http_client = stub
with pytest.raises(TaskError) as live:
await ctx._run_inline_step("risky", self._boom)
# The live round raised from the record the backend stored, not from
# the marker it posted: that is what keeps the two rounds identical
# even if the SDK and the backend ever build a record differently.
stored = {**posted[0]["result"], "message": "normalized by the backend"}
assert str(live.value) == "normalized by the backend"
# ...and the replay of that very record raises the same thing.
replayed = WorkflowCtx({"completed_steps": {"risky": stored}})
with pytest.raises(TaskError) as replay:
await replayed._run_inline_step("risky", self._boom)
assert type(live.value) is type(replay.value)
assert live.value.args == replay.value.args
assert live.value.result == replay.value.result == stored["result"]
assert live.value.step_key == replay.value.step_key == "risky"
# A step has no child job to name, and nothing hangs off __cause__:
# a replay has no original exception to chain, so neither round does.
assert live.value.child_job_id is replay.value.child_job_id is None
assert live.value.__cause__ is replay.value.__cause__ is None
asyncio.run(run())
assert len(posted) == 1
assert posted[0]["key"] == "risky"
assert self._without_stack(posted[0]["result"]) == self.MARKER
def test_a_missing_echo_is_not_the_same_as_an_unreadable_one(self, monkeypatch):
"""A backend predating the echo answers without a JSON body and the
locally checkpointed marker stands in. A JSON body that will not parse
means the stored record exists but is unknown, so the round has to end
and let the next one read whatever the backend actually kept."""
_set_inline_fast_path_env(monkeypatch)
def _client(headers, json_impl):
class _Response:
def raise_for_status(self):
pass
_Response.headers = headers
_Response.json = json_impl
class _Client(_StubInlineClient):
async def post(self, url, content=None):
await super().post(url, content=content)
return _Response()
return _Client()
def _boom_json(self):
raise ValueError("not json")
async def run():
# no JSON body: the fast path still completes, raising the failure
ctx = WorkflowCtx({})
ctx._inline_http_client = _client({}, _boom_json)
with pytest.raises(TaskError):
await ctx._run_inline_step("risky", self._boom)
# a JSON body that will not parse: fall through to the suspend path
ctx = WorkflowCtx({})
ctx._inline_http_client = _client(
{"content-type": "application/json"}, _boom_json
)
with pytest.raises(_StepSuspend) as suspend:
await ctx._run_inline_step("risky", self._boom)
assert suspend.value.dispatch_info["key"] == "risky"
asyncio.run(run())
def test_replay_reraises_and_does_not_hang(self):
checkpoint = {
"completed_steps": {"risky": self.MARKER},
"_executing_key": "double",
}
async def run():
return await asyncio.wait_for(
_run_workflow_async(self._wf(), checkpoint, {"x": 5}), timeout=5
)
r = asyncio.run(run())
assert r["type"] == "complete"
assert r["result"] == 10
class TestInlineStepRoundParity:
"""The round that runs a ``step()`` body must see what a replay sees.
The fast path returns the value it checkpointed, not the in-memory one:
a workflow branching on a datetime attribute or a tuple would otherwise
take one path on the round that ran the body and another on every replay,
which can change which tasks get dispatched, not just crash later.
"""
CASES = [
("dt", lambda: datetime(2026, 1, 1, tzinfo=timezone.utc), "2026-01-01 00:00:00+00:00"),
("pair", lambda: (1, 2), [1, 2]),
("intkeys", lambda: {1: "a"}, {"1": "a"}),
]
def test_outside_a_workflow_returns_the_same_shape(self):
"""No checkpoint, no replay — but a local run must not hand back a shape
a deployed one never produces, or testing a workflow locally proves
nothing. The async task path is the sharp edge: the wrapper is sync, so
the value has to be round-tripped after the await, not before."""
@task
async def make_pair():
return (1, datetime(2026, 1, 1, tzinfo=timezone.utc))
assert asyncio.run(step("pair", lambda: (1, 2))) == [1, 2]
assert asyncio.run(make_pair()) == [1, "2026-01-01 00:00:00+00:00"]
def test_live_round_matches_checkpoint_and_replay(self, monkeypatch):
_set_inline_fast_path_env(monkeypatch)
async def run():
for key, fn, expected in self.CASES:
stub = _StubInlineClient()
ctx = WorkflowCtx({})
ctx._inline_http_client = stub
live = await ctx._run_inline_step(key, fn)
checkpointed = stub.posted[0]["result"]
assert checkpointed == expected
assert live == expected and type(live) is type(expected)
replayed = WorkflowCtx({"completed_steps": {key: checkpointed}})
assert await replayed._run_inline_step(key, fn) == live
asyncio.run(run())
# =====================================================================
# TASK OPTIONS TESTS
# =====================================================================
class TestTaskOptions:
def test_options_forwarded_in_dispatch(self):
@task(timeout=600, tag="gpu", cache_ttl=3600, priority=10)
async def heavy(x: int):
return x
@workflow
async def wf(x: int):
return await heavy(x=x)
r = _run_workflow(wf, {}, {"x": 42})
assert r["type"] == "dispatch"
step_info = r["steps"][0]
assert step_info["timeout"] == 600
assert step_info["tag"] == "gpu"
assert step_info["cache_ttl"] == 3600
assert step_info["priority"] == 10
def test_task_without_options_has_no_extra_fields(self):
@task
async def simple(x: int):
return x
@workflow
async def wf(x: int):
return await simple(x=x)
r = _run_workflow(wf, {}, {"x": 1})
step_info = r["steps"][0]
assert "timeout" not in step_info
assert "tag" not in step_info
def test_concurrency_options_forwarded(self):
@task(concurrency_limit=5, concurrency_key="my-key", concurrency_time_window_s=60)
async def limited(x: int):
return x
@workflow
async def wf(x: int):
return await limited(x=x)
r = _run_workflow(wf, {}, {"x": 1})
step_info = r["steps"][0]
assert step_info["concurrent_limit"] == 5
assert step_info["concurrency_key"] == "my-key"
assert step_info["concurrency_time_window_s"] == 60
# =====================================================================
# SLEEP TESTS
# =====================================================================
class TestSleep:
def test_sleep_returns_sleep_output(self):
@workflow
async def wf():
await double(x=1)
await sleep(60)
await add_one(x=2)
return "done"
r = _run_workflow(
wf,
{"completed_steps": {"double": 2}},
{},
)
assert r["type"] == "sleep"
assert r["key"] == "sleep"
assert r["seconds"] == 60
def test_sleep_completes_on_replay(self):
@workflow
async def wf():
await double(x=1)
await sleep(60)
await add_one(x=2)
return "done"
r = _run_workflow(
wf,
{"completed_steps": {"double": 2, "sleep": True}},
{},
)
assert r["type"] == "dispatch"
assert r["steps"][0]["name"] == "add_one"
assert r["steps"][0]["key"] == "add_one"
def test_all_steps_with_sleep_complete(self):
@workflow
async def wf():
await double(x=1)
await sleep(60)
await add_one(x=2)
return "done"
r = _run_workflow(
wf,
{"completed_steps": {"double": 2, "sleep": True, "add_one": 3}},
{},
)
assert r["type"] == "complete"
assert r["result"] == "done"
def test_sleep_enforces_minimum(self):
@workflow
async def wf():
await sleep(0)
return "done"
r = _run_workflow(wf, {}, {})
assert r["seconds"] == 1
# =====================================================================
# PARALLEL UTILITY TESTS
# =====================================================================
class TestParallel:
def test_dispatches_all_items(self):
@workflow
async def wf():
results = await parallel([1, 2, 3], double)
return results
r = _run_workflow(wf, {}, {})
assert r["type"] == "dispatch"
assert r["mode"] == "parallel"
assert len(r["steps"]) == 3
def test_completes_with_all_results(self):
@workflow
async def wf():
results = await parallel([1, 2, 3], double)
return results
r = _run_workflow(
wf,
{"completed_steps": {"double": 2, "double_2": 4, "double_3": 6}},
{},
)
assert r["type"] == "complete"
assert r["result"] == [2, 4, 6]
def test_batched_dispatches_first_batch(self):
@workflow
async def wf():
results = await parallel([1, 2, 3, 4, 5], double, concurrency=2)
return results
r = _run_workflow(wf, {}, {})
assert r["type"] == "dispatch"
assert r["mode"] == "parallel"
assert len(r["steps"]) == 2
def test_batched_dispatches_second_batch(self):
@workflow
async def wf():
results = await parallel([1, 2, 3, 4, 5], double, concurrency=2)
return results
r = _run_workflow(
wf,
{"completed_steps": {"double": 2, "double_2": 4}},
{},
)
assert r["type"] == "dispatch"
assert len(r["steps"]) == 2
def test_batched_completes_with_all_results(self):
@workflow
async def wf():
results = await parallel([1, 2, 3, 4, 5], double, concurrency=2)
return results
r = _run_workflow(
wf,
{"completed_steps": {"double": 2, "double_2": 4, "double_3": 6, "double_4": 8, "double_5": 10}},
{},
)
assert r["type"] == "complete"
assert r["result"] == [2, 4, 6, 8, 10]
def test_empty_items_returns_empty(self):
@workflow
async def wf():
results = await parallel([], double)
return results
r = _run_workflow(wf, {}, {})
assert r["type"] == "complete"
assert r["result"] == []
class TestApprovalKeys:
"""`key` names the step that get_approval_urls() mints URLs against, so a
duplicate must fail rather than silently become `<key>_2` and leave the
caller holding a URL for the earlier step."""
def test_explicit_key_is_used_verbatim(self):
@workflow
async def wf():
return await wait_for_approval(key="manager")
assert _run_workflow(wf, {}, {})["key"] == "manager"
def test_duplicate_explicit_key_raises(self):
@workflow
async def wf():
await wait_for_approval(key="manager")
await wait_for_approval(key="manager")
with pytest.raises(RuntimeError, match="already used"):
_run_workflow(wf, {"completed_steps": {"manager": {"approved": True}}}, {})
def test_explicit_key_colliding_with_a_suffixed_step_key_raises(self):
"""`step("dup")` twice yields `dup`/`dup_2`, so an approval explicitly named
`dup_2` would alias the second step's key."""
@workflow
async def wf():
await step("dup", lambda: 1)
await step("dup", lambda: 2)
await wait_for_approval(key="dup_2")
with pytest.raises(RuntimeError, match="already used"):
_run_workflow(wf, {"completed_steps": {"dup": 1, "dup_2": 2}}, {})
@pytest.mark.parametrize("bad", ["", " ", ".", "..", "a/b"])
def test_unusable_key_raises(self, bad):
"""The key travels as one path segment when its URLs are minted, so anything
`get_approval_urls` could not address must be refused here too."""
@workflow
async def wf():
await wait_for_approval(key=bad)
with pytest.raises(RuntimeError, match="non-empty step name"):
_run_workflow(wf, {}, {})
def test_unnamed_approvals_still_auto_number(self):
@workflow
async def wf():
await wait_for_approval()
await wait_for_approval()
assert _run_workflow(wf, {"completed_steps": {"approval": {}}}, {})["key"] == "approval_2"