Skip to content

Commit b60bf59

Browse files
Retry fenced workflow task heartbeat after backend loss
1 parent 20a0983 commit b60bf59

2 files changed

Lines changed: 93 additions & 4 deletions

File tree

‎src/durable_workflow/retry_policy.py‎

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,18 @@ def _backend_unavailable_refusal(exc: Exception) -> tuple[bool, int | None]:
8080
"/api/worker/heartbeat": "heartbeat_worker",
8181
}
8282
operation = next((name for path, name in operations.items() if request.url.path.endswith(path)), None)
83+
_, task_heartbeat_marker, task_heartbeat_tail = request.url.path.rpartition(
84+
"/api/worker/workflow-tasks/"
85+
)
86+
task_id = task_heartbeat_tail.removesuffix("/heartbeat")
87+
task_heartbeat = (
88+
bool(task_heartbeat_marker)
89+
and task_heartbeat_tail.endswith("/heartbeat")
90+
and bool(task_id)
91+
and "/" not in task_id
92+
)
93+
if task_heartbeat:
94+
operation = "heartbeat_workflow_task"
8395
if operation is None:
8496
return False, None
8597
try:
@@ -97,6 +109,24 @@ def _backend_unavailable_refusal(exc: Exception) -> tuple[bool, int | None]:
97109
worker_id = submitted.get("worker_id")
98110
queue = submitted.get("task_queue")
99111
delay = body.get("retry_after_seconds")
112+
if task_heartbeat:
113+
lease_owner = submitted.get("lease_owner")
114+
attempt = submitted.get("workflow_task_attempt")
115+
if (
116+
not isinstance(lease_owner, str) or not lease_owner
117+
or type(attempt) is not int or attempt <= 0
118+
or body.get("operation") != operation
119+
or body.get("outcome") != "unknown"
120+
or body.get("worker_id") != lease_owner
121+
or body.get("task_queue") is not None
122+
or body.get("task_id") != task_id
123+
or body.get("lease_owner") != lease_owner
124+
or body.get("workflow_task_attempt") != attempt
125+
or body.get("retryable") is not True
126+
or type(delay) is not int or delay <= 0
127+
):
128+
return True, None
129+
return True, delay
100130
if (
101131
not isinstance(worker_id, str) or not worker_id
102132
or body.get("operation") != operation

‎tests/test_storage_admission.py‎

Lines changed: 63 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,8 @@ def pressure(poll_id: str | None = None, *, reason: str = "storage_pressure", en
4242
def backend_unavailable(request: httpx.Request) -> dict[str, Any]:
4343
submitted = json.loads(request.content)
4444
path = request.url.path
45-
operation = {
45+
task_heartbeat = "/api/worker/workflow-tasks/" in path and path.endswith("/heartbeat")
46+
operation = "heartbeat_workflow_task" if task_heartbeat else {
4647
"/api/worker/workflow-tasks/poll": "poll_workflow_task",
4748
"/api/worker/activity-tasks/poll": "poll_activity_task",
4849
"/api/worker/query-tasks/poll": "poll_query_task",
@@ -51,9 +52,16 @@ def backend_unavailable(request: httpx.Request) -> dict[str, Any]:
5152
}[path]
5253
body: dict[str, Any] = {
5354
"reason": "backend_unavailable", "operation": operation, "outcome": "unknown",
54-
"worker_id": submitted["worker_id"], "task_queue": submitted.get("task_queue"),
55+
"worker_id": submitted.get("lease_owner") if task_heartbeat else submitted["worker_id"],
56+
"task_queue": submitted.get("task_queue"),
5557
"retryable": True, "retry_after_seconds": 1,
5658
}
59+
if task_heartbeat:
60+
body.update({
61+
"task_id": path.rpartition("/api/worker/workflow-tasks/")[2].removesuffix("/heartbeat"),
62+
"lease_owner": submitted["lease_owner"],
63+
"workflow_task_attempt": submitted["workflow_task_attempt"],
64+
})
5765
if path.endswith("/poll"):
5866
body.update({
5967
"task": None, "poll_status": "backend_unavailable",
@@ -85,9 +93,9 @@ async def sleep(delay: float) -> None:
8593
return sleeps
8694

8795

88-
def client_for(handler: Callable[..., Any]) -> Client:
96+
def client_for(handler: Callable[..., Any], base_url: str = "https://runtime.example") -> Client:
8997
client = Client(
90-
"https://runtime.example", token="test-runtime-token",
98+
base_url, token="test-runtime-token",
9199
retry_policy=TransportRetryPolicy(max_attempts=2, initial_backoff_seconds=0, jitter=False),
92100
)
93101
client._http = httpx.AsyncClient(base_url=client.base_url, transport=httpx.MockTransport(handler))
@@ -186,6 +194,57 @@ def handler(request: httpx.Request) -> httpx.Response:
186194
assert sum(retry_sleeps) == pytest.approx(4)
187195

188196

197+
@pytest.mark.parametrize("base_url", ["https://runtime.example", "https://runtime.example/managed/namespace"])
198+
async def test_backend_outage_retries_workflow_task_heartbeat_with_same_fence(
199+
base_url: str, retry_sleeps: list[float],
200+
) -> None:
201+
requests: list[bytes] = []
202+
203+
def handler(request: httpx.Request) -> httpx.Response:
204+
requests.append(request.content)
205+
if len(requests) <= 4:
206+
return httpx.Response(503, json=backend_unavailable(request))
207+
return httpx.Response(200, json={
208+
"task_id": "task-1", "lease_owner": "worker-1",
209+
"workflow_task_attempt": 3, "renewed": True,
210+
})
211+
212+
async with client_for(handler, base_url) as client:
213+
with worker_scope():
214+
result = await client.heartbeat_workflow_task(
215+
task_id="task-1", lease_owner="worker-1", workflow_task_attempt=3,
216+
)
217+
assert result["renewed"] is True
218+
assert len(requests) == 5
219+
assert len(set(requests)) == 1
220+
assert sum(retry_sleeps) == pytest.approx(4)
221+
222+
223+
@pytest.mark.parametrize("override", [
224+
{"task_id": "other-task"}, {"lease_owner": "other-worker"},
225+
{"worker_id": "other-worker"}, {"workflow_task_attempt": 4},
226+
{"operation": "heartbeat_worker"}, {"outcome": "failed"},
227+
{"retryable": False}, {"retry_after_seconds": 0},
228+
])
229+
async def test_invalid_workflow_task_heartbeat_backend_response_is_not_retried(
230+
override: dict[str, Any], retry_sleeps: list[float],
231+
) -> None:
232+
calls = 0
233+
234+
def handler(request: httpx.Request) -> httpx.Response:
235+
nonlocal calls
236+
calls += 1
237+
return httpx.Response(503, json={**backend_unavailable(request), **override})
238+
239+
async with client_for(handler) as client:
240+
with worker_scope(), pytest.raises(ServerError):
241+
await client.heartbeat_workflow_task(
242+
task_id="task-1", lease_owner="worker-1", workflow_task_attempt=3,
243+
)
244+
assert calls == 1
245+
assert not retry_sleeps
246+
247+
189248
@pytest.mark.parametrize("override", [
190249
{"operation": "poll_activity_task"},
191250
{"outcome": "failed"}, {"worker_id": "other-worker"},

0 commit comments

Comments
 (0)