Skip to content

Commit b963a1e

Browse files
authored
Fixed synchronization of multiple recovery workflows (#822)
This commit improves the synchronization of concurrent recovery workflows. Previously, an error occurred when a dependent workflow attempted to attach to an `InterWorkflowPort` of the dependee before it was created. This race condition happened because port creation occurred in a successor phase of job synchronization. This is now resolved by acquiring a lock to check the status and holding it until the recovery workflow responsible for the rollback is ready. Additionally, this commit fixes the `move_token_to_root` and `remove_port` methods. Token movement now correctly drives port deletion; previously, port deletion drove token deletion, which occasionally resulted in an empty token graph. Finally, this commit corrects the status precedence in the `_reduce_statuses` function. The `FAILED` and `CANCELLED` statuses now take priority over `RECOVERED`.
1 parent c53e069 commit b963a1e

7 files changed

Lines changed: 676 additions & 342 deletions

File tree

streamflow/core/recovery.py

Lines changed: 7 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,6 @@
33
import asyncio
44
import functools
55
from abc import ABC, abstractmethod
6-
from collections.abc import MutableMapping
7-
from enum import IntEnum
86
from typing import TYPE_CHECKING
97

108
from streamflow.core.context import SchemaEntity
@@ -86,10 +84,7 @@ async def close(self) -> None: ...
8684
def get_request(self, job_name: str) -> RetryRequest: ...
8785

8886
@abstractmethod
89-
async def recover(self, job: Job, step: Step, exception: BaseException) -> None: ...
90-
91-
@abstractmethod
92-
async def is_recovered(self, job_name: str) -> TokenAvailability: ...
87+
async def is_recovering(self, job_name: str) -> bool: ...
9388

9489
@abstractmethod
9590
async def notify(
@@ -99,6 +94,9 @@ async def notify(
9994
job_token: JobToken | None = None,
10095
) -> None: ...
10196

97+
@abstractmethod
98+
async def recover(self, job: Job, step: Step, exception: BaseException) -> None: ...
99+
102100
@abstractmethod
103101
async def update_request(self, job_name: str) -> None: ...
104102

@@ -112,17 +110,10 @@ async def recover(self, failed_job: Job, failed_step: Step) -> None: ...
112110

113111

114112
class RetryRequest:
115-
__slots__ = ("job_token", "lock", "output_tokens", "version", "workflow")
113+
__slots__ = ("lock", "name", "version", "workflow")
116114

117-
def __init__(self) -> None:
118-
self.job_token: JobToken | None = None
115+
def __init__(self, name: str) -> None:
119116
self.lock: asyncio.Lock = asyncio.Lock()
120-
self.output_tokens: MutableMapping[str, Token] = {}
117+
self.name: str = name
121118
self.version: int = 1
122119
self.workflow: Workflow | None = None
123-
124-
125-
class TokenAvailability(IntEnum):
126-
Unavailable = 0
127-
Available = 1
128-
FutureAvailable = 2

streamflow/recovery/failure_manager.py

Lines changed: 21 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -7,12 +7,7 @@
77

88
from streamflow.core.context import StreamFlowContext
99
from streamflow.core.exception import FailureHandlingException
10-
from streamflow.core.recovery import (
11-
FailureManager,
12-
RetryRequest,
13-
TokenAvailability,
14-
recoverable,
15-
)
10+
from streamflow.core.recovery import FailureManager, RetryRequest, recoverable
1611
from streamflow.core.workflow import Job, Status, Step, Token
1712
from streamflow.log_handler import logger
1813
from streamflow.recovery.policy.recovery import RollbackRecoveryPolicy
@@ -36,9 +31,14 @@ async def _do_handle_failure(self, job: Job, step: Step) -> None:
3631
# Delay rescheduling to manage temporary failures (e.g. connection lost)
3732
if self.retry_delay is not None:
3833
await asyncio.sleep(self.retry_delay)
39-
await RollbackRecoveryPolicy(self.context).recover(job, step)
40-
if logger.isEnabledFor(logging.INFO):
41-
logger.info(f"COMPLETED Recovery execution of failed job {job.name}")
34+
try:
35+
await RollbackRecoveryPolicy(self.context).recover(job, step)
36+
if logger.isEnabledFor(logging.INFO):
37+
logger.info(f"COMPLETED Recovery execution of failed job {job.name}")
38+
except FailureHandlingException as e:
39+
if logger.isEnabledFor(logging.INFO):
40+
logger.info(f"FAILED Recovery execution of failed job {job.name}")
41+
raise e
4242

4343
async def close(self) -> None:
4444
pass
@@ -47,7 +47,7 @@ def get_request(self, job_name: str) -> RetryRequest:
4747
if job_name in self._retry_requests.keys():
4848
return self._retry_requests[job_name]
4949
else:
50-
return self._retry_requests.setdefault(job_name, RetryRequest())
50+
return self._retry_requests.setdefault(job_name, RetryRequest(job_name))
5151

5252
@classmethod
5353
def get_schema(cls) -> str:
@@ -58,6 +58,13 @@ def get_schema(cls) -> str:
5858
.read_text("utf-8")
5959
)
6060

61+
async def is_recovering(self, job_name: str) -> bool:
62+
return self.context.scheduler.get_allocation(job_name).status in (
63+
Status.ROLLBACK,
64+
Status.RUNNING,
65+
Status.FIREABLE,
66+
)
67+
6168
async def recover(self, job: Job, step: Step, exception: BaseException) -> None:
6269
if logger.isEnabledFor(logging.INFO):
6370
logger.info(
@@ -66,46 +73,16 @@ async def recover(self, job: Job, step: Step, exception: BaseException) -> None:
6673
await self.context.scheduler.notify_status(job.name, Status.RECOVERY)
6774
await self._do_handle_failure(job, step)
6875

69-
async def is_recovered(self, job_name: str) -> TokenAvailability:
70-
if request := self._retry_requests.get(job_name):
71-
async with request.lock:
72-
if self.context.scheduler.get_allocation(job_name).status in (
73-
Status.ROLLBACK,
74-
Status.RUNNING,
75-
Status.FIREABLE,
76-
):
77-
return TokenAvailability.FutureAvailable
78-
elif len(request.output_tokens) > 0 and all(
79-
await asyncio.gather(
80-
*(
81-
asyncio.create_task(t.is_available(self.context))
82-
for t in request.output_tokens.values()
83-
)
84-
)
85-
):
86-
return TokenAvailability.Available
87-
return TokenAvailability.Unavailable
88-
8976
async def notify(
9077
self,
9178
output_port: str,
9279
output_token: Token,
9380
job_token: JobToken | None = None,
9481
) -> None:
95-
if job_token is not None:
96-
job_name = job_token.value.name
97-
if job_name in self._retry_requests.keys():
98-
async with self._retry_requests[job_name].lock:
99-
self._retry_requests[job_name].job_token = job_token
100-
self._retry_requests[job_name].output_tokens.setdefault(
101-
output_port, output_token
102-
)
82+
pass
10383

10484
async def update_request(self, job_name: str) -> None:
10585
retry_request = self._retry_requests[job_name]
106-
async with retry_request.lock:
107-
retry_request.job_token = None
108-
retry_request.output_tokens = {}
10986
if self.max_retries is None or retry_request.version < self.max_retries:
11087
retry_request.version += 1
11188
if logger.isEnabledFor(logging.DEBUG):
@@ -138,16 +115,16 @@ def get_schema(cls) -> str:
138115
def get_request(self, job_name: str) -> RetryRequest:
139116
pass
140117

118+
async def is_recovering(self, job_name: str) -> bool:
119+
return False
120+
141121
async def recover(self, job: Job, step: Step, exception: BaseException) -> None:
142122
if logger.isEnabledFor(logging.WARNING):
143123
logger.warning(
144124
f"Job {job.name} failure can not be recovered. Failure manager is not enabled."
145125
)
146126
raise exception
147127

148-
async def is_recovered(self, job_name: str) -> TokenAvailability:
149-
return TokenAvailability.Unavailable
150-
151128
async def notify(
152129
self,
153130
output_port: str,

streamflow/recovery/policy/recovery.py

Lines changed: 85 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -1,22 +1,18 @@
11
from __future__ import annotations
22

33
import asyncio
4+
import contextlib
45
import logging
56
from collections.abc import MutableSequence, MutableSet
67
from typing import cast
78

89
from streamflow.core.exception import FailureHandlingException
9-
from streamflow.core.recovery import RecoveryPolicy
10+
from streamflow.core.recovery import RecoveryPolicy, RetryRequest
1011
from streamflow.core.utils import get_tag
11-
from streamflow.core.workflow import Job, Step, Token, Workflow
12+
from streamflow.core.workflow import Job, Port, Step, Token, Workflow
1213
from streamflow.log_handler import logger
1314
from streamflow.persistence.loading_context import WorkflowBuilder
14-
from streamflow.recovery.utils import (
15-
GraphMapper,
16-
ProvenanceGraph,
17-
TokenAvailability,
18-
create_graph_mapper,
19-
)
15+
from streamflow.recovery.utils import GraphMapper, ProvenanceGraph, create_graph_mapper
2016
from streamflow.workflow.executor import StreamFlowExecutor
2117
from streamflow.workflow.port import (
2218
BoundaryAction,
@@ -29,6 +25,28 @@
2925
from streamflow.workflow.utils import get_job_token
3026

3127

28+
def _get_recovery_port(
29+
token_id: int,
30+
mapper: GraphMapper,
31+
original_workflow: Workflow,
32+
recovery_workflow: Workflow,
33+
) -> Port:
34+
port_name = next(
35+
curr_port
36+
for curr_port, curr_tokens in mapper.port_tokens.items()
37+
if token_id in curr_tokens
38+
)
39+
if port_name not in recovery_workflow.ports.keys() or not isinstance(
40+
recovery_workflow.ports[port_name], InterWorkflowPort
41+
):
42+
return recovery_workflow.create_port(
43+
cls=type(original_workflow.ports[port_name]),
44+
name=port_name,
45+
)
46+
else:
47+
return recovery_workflow.ports[port_name]
48+
49+
3250
async def _inject_tokens(
3351
failed_job: Job,
3452
failed_step: Step,
@@ -124,62 +142,48 @@ async def _populate_workflow(
124142

125143

126144
class RollbackRecoveryPolicy(RecoveryPolicy):
127-
async def _sync_workflows(
145+
async def _synchronize_workflows(
128146
self,
129-
job_names: MutableSet[str],
147+
failed_job: str,
130148
job_tokens: MutableSequence[Token],
131149
mapper: GraphMapper,
150+
retry_requests: MutableSequence[RetryRequest],
132151
workflow: Workflow,
133152
) -> None:
134-
for job_name in job_names:
135-
retry_request = self.context.failure_manager.get_request(job_name)
136-
if (
137-
is_available := await self.context.failure_manager.is_recovered(
138-
job_name
139-
)
140-
) == TokenAvailability.FutureAvailable:
153+
for retry_request in retry_requests:
154+
job_name = retry_request.name
155+
if await self.context.failure_manager.is_recovering(job_name):
141156
job_token = get_job_token(job_name, job_tokens)
142-
# The `retry_request` is the current job running, instead
143-
# the `job_token` is the token to remove in the graph because
144-
# the workflow will depend on the already running job
145157
if logger.isEnabledFor(logging.DEBUG):
146-
logger.debug(f"Synchronize rollbacks: job {job_name} is running")
147-
# todo: create a unit test for this case
148-
for port_name in await mapper.get_output_ports(job_token):
149-
if port_name in retry_request.workflow.ports.keys():
150-
cast(
151-
InterWorkflowJobPort,
152-
retry_request.workflow.ports[port_name],
153-
).add_inter_port(
154-
workflow.create_port(
155-
cls=InterWorkflowJobPort, name=port_name
156-
),
157-
boundary_tags=[job_token.tag],
158-
boundary_action=(
159-
BoundaryAction.PROPAGATE | BoundaryAction.TERMINATE
160-
),
161-
)
162-
# Remove tokens recovered in other workflows
163-
for token_id in await mapper.get_output_tokens(job_token.persistent_id):
158+
logger.debug(
159+
f"Synchronizing rollbacks for failed job {failed_job}: Job {job_name} is currently executing."
160+
)
161+
available_tokens = set()
162+
for token_id in (
163+
mapper.dag_tokens.successors(job_token.persistent_id)
164+
if mapper.dag_tokens.contains(job_token.persistent_id)
165+
else []
166+
):
164167
mapper.move_token_to_root(token_id)
165-
elif is_available == TokenAvailability.Available:
166-
job_token = get_job_token(job_name, job_tokens)
168+
available_tokens.add(token_id)
169+
# Some tokens could be discarded by the `move_token_to_root` method
170+
for token_id in available_tokens & mapper.token_instances.keys():
171+
new_port = _get_recovery_port(
172+
token_id, mapper, retry_request.workflow, workflow
173+
)
174+
cast(
175+
InterWorkflowPort,
176+
retry_request.workflow.ports[new_port.name],
177+
).add_inter_port(
178+
port=new_port,
179+
boundary_tags=[job_token.tag],
180+
boundary_action=BoundaryAction.PROPAGATE,
181+
)
182+
else:
167183
if logger.isEnabledFor(logging.DEBUG):
168184
logger.debug(
169-
f"Synchronize rollbacks: job {job_token.value.name} output available"
185+
f"Synchronizing rollbacks for failed job {failed_job}: Job {job_name} rollback"
170186
)
171-
# Search execute token after job token, replace this token with job_req token.
172-
# Then remove all the prev tokens
173-
for port_name in await mapper.get_output_ports(job_token):
174-
if port_name in retry_request.output_tokens.keys():
175-
new_token = retry_request.output_tokens[port_name]
176-
mapper.replace_token(
177-
port_name,
178-
new_token,
179-
True,
180-
)
181-
mapper.move_token_to_root(new_token.persistent_id)
182-
else:
183187
await self.context.failure_manager.update_request(job_name)
184188
retry_request.workflow = workflow
185189

@@ -211,24 +215,35 @@ async def recover(self, failed_job: Job, failed_step: Step) -> None:
211215
job_tokens = list(
212216
filter(lambda t: isinstance(t, JobToken), mapper.token_instances.values())
213217
)
214-
await self._sync_workflows(
215-
job_names={*(t.value.name for t in job_tokens), failed_job.name},
216-
job_tokens=job_tokens,
217-
mapper=mapper,
218-
workflow=new_workflow,
219-
)
220-
if mapper.dag_tokens.empty():
221-
raise FailureHandlingException(
222-
f"Impossible to recover {failed_job.name}: empty token graph"
218+
retry_requests = [
219+
self.context.failure_manager.get_request(job_name)
220+
for job_name in {*(t.value.name for t in job_tokens), failed_job.name}
221+
]
222+
async with contextlib.AsyncExitStack() as exit_stack:
223+
await asyncio.gather(
224+
*(
225+
asyncio.create_task(exit_stack.enter_async_context(request.lock))
226+
for request in retry_requests
227+
)
228+
)
229+
await self._synchronize_workflows(
230+
failed_job=failed_job.name,
231+
job_tokens=job_tokens,
232+
mapper=mapper,
233+
retry_requests=retry_requests,
234+
workflow=new_workflow,
235+
)
236+
if mapper.dag_tokens.empty():
237+
raise FailureHandlingException(
238+
f"Impossible to recover {failed_job.name}: empty token graph"
239+
)
240+
# Populate new workflow
241+
await _populate_workflow(
242+
failed_step=failed_step,
243+
step_ids=await mapper.get_step_ids(failed_step.output_ports.values()),
244+
workflow=new_workflow,
245+
workflow_builder=workflow_builder,
223246
)
224-
# Populate new workflow
225-
step_ids = await mapper.get_step_ids(failed_step.output_ports.values())
226-
await _populate_workflow(
227-
failed_step=failed_step,
228-
step_ids=step_ids,
229-
workflow=new_workflow,
230-
workflow_builder=workflow_builder,
231-
)
232247
await _inject_tokens(
233248
failed_job=failed_job,
234249
failed_step=failed_step,

0 commit comments

Comments
 (0)