@@ -42,7 +42,8 @@ def pressure(poll_id: str | None = None, *, reason: str = "storage_pressure", en
4242def 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