Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 13 additions & 9 deletions airflow/providers/amazon/aws/executors/ecs/ecs_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -243,20 +243,17 @@ def __update_running_task(self, task):
task_key = self.active_workers.arn_to_key[task.task_arn]

# Mark finished tasks as either a success/failure.
if task_state == State.FAILED:
self.fail(task_key)
if task_state == State.FAILED or task_state == State.REMOVED:
self.__log_container_failures(task_arn=task.task_arn)
elif task_state == State.SUCCESS:
self.success(task_key)
elif task_state == State.REMOVED:
self.__handle_failed_task(task.task_arn, task.stopped_reason)
Comment thread
o-nikolas marked this conversation as resolved.
if task_state in (State.FAILED, State.SUCCESS):
elif task_state == State.SUCCESS:
Comment thread
o-nikolas marked this conversation as resolved.
self.log.debug(
"Airflow task %s marked as %s after running on ECS Task (arn) %s",
task_key,
task_state,
task.task_arn,
)
self.success(task_key)
self.active_workers.pop_by_key(task_key)

def __describe_tasks(self, task_arns):
Expand Down Expand Up @@ -289,7 +286,14 @@ def __log_container_failures(self, task_arn: str):
)

def __handle_failed_task(self, task_arn: str, reason: str):
"""If an API failure occurs, the task is rescheduled."""
"""
If an API failure occurs, the task is rescheduled.

This function will determine whether the task has been attempted the appropriate number
of times, and determine whether the task should be marked failed or not. The task will
be removed active_workers, and marked as FAILED, or set into pending_tasks depending on
how many times it has been retried.
"""
task_key = self.active_workers.arn_to_key[task_arn]
task_info = self.active_workers.info_by_key(task_key)
task_cmd = task_info.cmd
Expand All @@ -305,7 +309,6 @@ def __handle_failed_task(self, task_arn: str, reason: str):
self.__class__.MAX_RUN_TASK_ATTEMPTS,
task_arn,
)
self.active_workers.increment_failure_count(task_key)
self.pending_tasks.append(
EcsQueuedTask(
task_key,
Expand All @@ -322,8 +325,8 @@ def __handle_failed_task(self, task_arn: str, reason: str):
task_key,
failure_count,
)
self.active_workers.pop_by_key(task_key)
self.fail(task_key)
self.active_workers.pop_by_key(task_key)

def attempt_task_runs(self):
"""
Expand All @@ -346,6 +349,7 @@ def attempt_task_runs(self):
attempt_number = ecs_task.attempt_number
_failure_reasons = []
if timezone.utcnow() < ecs_task.next_attempt_time:
self.pending_tasks.append(ecs_task)
continue
try:
run_task_response = self._run_task(task_key, cmd, queue, exec_config)
Expand Down
196 changes: 174 additions & 22 deletions tests/providers/amazon/aws/executors/ecs/test_ecs_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -598,6 +598,109 @@ def test_attempt_task_runs_attempts_when_some_tasks_fal(self, _, mock_executor,
== caplog.messages[0]
)

@mock.patch.object(ecs_executor, "calculate_next_attempt_delay", return_value=dt.timedelta(seconds=0))
def test_task_retry_on_api_failure_all_tasks_fail(self, _, mock_executor, caplog):
Comment thread
syedahsn marked this conversation as resolved.
"""
Test API failure retries.
"""
AwsEcsExecutor.MAX_RUN_TASK_ATTEMPTS = "2"
airflow_keys = ["TaskInstanceKey1", "TaskInstanceKey2"]
airflow_commands = [mock.Mock(spec=list), mock.Mock(spec=list)]

mock_executor.execute_async(airflow_keys[0], airflow_commands[0])
mock_executor.execute_async(airflow_keys[1], airflow_commands[1])
assert len(mock_executor.pending_tasks) == 2
caplog.set_level("WARNING")

describe_tasks = [
{
"taskArn": ARN1,
"desiredStatus": "STOPPED",
"lastStatus": "FAILED",
"startedAt": dt.datetime.now(),
"stoppedReason": "Task marked as FAILED",
"containers": [
{
"name": "some-ecs-container",
"lastStatus": "STOPPED",
"exitCode": 100,
}
],
},
{
"taskArn": ARN2,
"desiredStatus": "STOPPED",
"lastStatus": "FAILED",
"stoppedReason": "Task marked as REMOVED",
"containers": [
{
"name": "some-ecs-container",
"lastStatus": "STOPPED",
"exitCode": 100,
}
],
},
]
run_tasks = [
{
"taskArn": ARN1,
"lastStatus": "",
"desiredStatus": "",
"containers": [{"name": "some-ecs-container"}],
},
{
"taskArn": ARN2,
"lastStatus": "",
"desiredStatus": "",
"containers": [{"name": "some-ecs-container"}],
},
]
mock_executor.ecs.run_task.side_effect = [
{"tasks": [run_tasks[0]], "failures": []},
{"tasks": [run_tasks[1]], "failures": []},
]
mock_executor.ecs.describe_tasks.side_effect = [{"tasks": describe_tasks, "failures": []}]

mock_executor.attempt_task_runs()

for i in range(2):
RUN_TASK_KWARGS["overrides"]["containerOverrides"][0]["command"] = airflow_commands[i]
assert mock_executor.ecs.run_task.call_args_list[i].kwargs == RUN_TASK_KWARGS

assert len(mock_executor.pending_tasks) == 0
assert len(mock_executor.active_workers.get_all_arns()) == 2

mock_executor.sync_running_tasks()
for i in range(2):
assert (
f"Airflow task {airflow_keys[i]} failed due to {describe_tasks[i]['stoppedReason']}. Failure 1 out of 2"
in caplog.messages[i]
)

caplog.clear()
mock_executor.ecs.run_task.call_args_list.clear()

mock_executor.ecs.run_task.side_effect = [
{"tasks": [run_tasks[0]], "failures": []},
{"tasks": [run_tasks[1]], "failures": []},
]
mock_executor.ecs.describe_tasks.side_effect = [{"tasks": describe_tasks, "failures": []}]

mock_executor.attempt_task_runs()

mock_executor.attempt_task_runs()

for i in range(2):
RUN_TASK_KWARGS["overrides"]["containerOverrides"][0]["command"] = airflow_commands[i]
assert mock_executor.ecs.run_task.call_args_list[i].kwargs == RUN_TASK_KWARGS

mock_executor.sync_running_tasks()
for i in range(2):
assert (
f"Airflow task {airflow_keys[i]} has failed a maximum of 2 times. Marking as failed"
in caplog.messages[i]
)

@mock.patch.object(BaseExecutor, "fail")
@mock.patch.object(BaseExecutor, "success")
def test_sync(self, success_mock, fail_mock, mock_executor):
Expand Down Expand Up @@ -629,29 +732,31 @@ def test_sync_short_circuits_with_no_arns(self, _, success_mock, fail_mock, mock
@mock.patch.object(BaseExecutor, "success")
def test_failed_sync(self, success_mock, fail_mock, mock_executor):
"""Test success and failure states."""
AwsEcsExecutor.MAX_RUN_TASK_ATTEMPTS = "1"
self._mock_sync(mock_executor, State.FAILED)

mock_executor.sync()
mock_executor.ecs.describe_tasks.assert_called_once()

# Task is not stored in active workers.
assert len(mock_executor.active_workers) == 0
# Task is immediately succeeded.
# Task is immediately failed.
fail_mock.assert_called_once()
success_mock.assert_not_called()

@mock.patch.object(BaseExecutor, "fail")
@mock.patch.object(BaseExecutor, "success")
@mock.patch.object(BaseExecutor, "fail")
def test_removed_sync(self, fail_mock, success_mock, mock_executor):
"""A removed task will increment failure count but call neither fail() nor success()."""
"""A removed task will be treated as a failed task."""
AwsEcsExecutor.MAX_RUN_TASK_ATTEMPTS = "1"
self._mock_sync(mock_executor, expected_state=State.REMOVED, set_task_state=State.REMOVED)
task_instance_key = mock_executor.active_workers.arn_to_key[ARN1]

mock_executor.sync_running_tasks()

assert ARN1 in mock_executor.active_workers.get_all_arns()
assert mock_executor.active_workers.key_to_failure_counts[task_instance_key] == 2
fail_mock.assert_not_called()
# Task is not stored in active workers.
assert len(mock_executor.active_workers) == 0
# Task is immediately failed.
fail_mock.assert_called_once()
success_mock.assert_not_called()

@mock.patch.object(BaseExecutor, "fail")
Expand Down Expand Up @@ -698,19 +803,27 @@ def test_failed_sync_cumulative_fail(self, _, success_mock, fail_mock, mock_airf
],
}

# Call sync_running_tasks 2 times with failures.
# Call sync_running_tasks and attempt_task_runs 2 times with failures.
for _ in range(2):
mock_executor.sync_running_tasks()

# Ensure task arn is not removed from active.
# Ensure task gets removed from active_workers.
assert ARN1 not in mock_executor.active_workers.get_all_arns()
# Ensure task gets back on the pending_tasks queue
assert len(mock_executor.pending_tasks) == 1
keys = [task.key for task in mock_executor.pending_tasks]
assert task_key in keys

mock_executor.attempt_task_runs()
assert len(mock_executor.pending_tasks) == 0
assert ARN1 in mock_executor.active_workers.get_all_arns()

# Task is neither failed nor succeeded.
fail_mock.assert_not_called()
success_mock.assert_not_called()
# Task is neither failed nor succeeded.
fail_mock.assert_not_called()
success_mock.assert_not_called()

# run_task failed twice, and passed once
assert mock_executor.ecs.run_task.call_count == 3
# run_task failed twice, and passed 3 times
assert mock_executor.ecs.run_task.call_count == 5
# describe_tasks failed 2 times so far
assert mock_executor.ecs.describe_tasks.call_count == 2

Expand All @@ -731,27 +844,60 @@ def test_failed_sync_api_exception(self, mock_executor, caplog):

@mock.patch.object(BaseExecutor, "fail")
@mock.patch.object(BaseExecutor, "success")
def test_failed_sync_api(self, success_mock, fail_mock, mock_executor):
@mock.patch.object(ecs_executor, "calculate_next_attempt_delay", return_value=dt.timedelta(seconds=0))
def test_failed_sync_api(self, _, success_mock, fail_mock, mock_executor):
"""Test what happens when ECS sync fails for certain tasks repeatedly."""
self._mock_sync(mock_executor)
mock_executor.ecs.describe_tasks.return_value = {
airflow_key = "test-key"
airflow_cmd = mock.Mock(spec=list)
mock_executor.execute_async(airflow_key, airflow_cmd)
assert len(mock_executor.pending_tasks) == 1

run_task_ret_val = {
"taskArn": ARN1,
"desiredStatus": "STOPPED",
"lastStatus": "RUNNING",
"containers": [
{
"name": "some-ecs-container",
"lastStatus": "STOPPED",
"exitCode": 0,
}
],
}
mock_executor.ecs.run_task.return_value = {"tasks": [run_task_ret_val], "failures": []}
describe_tasks_ret_value = {
"tasks": [],
"failures": [
{"arn": ARN1, "reason": "Sample Failure", "detail": "UnitTest Failure - Please ignore"}
],
}
mock_executor.ecs.describe_tasks.return_value = describe_tasks_ret_value
mock_executor.attempt_task_runs()
assert len(mock_executor.pending_tasks) == 0
assert len(mock_executor.active_workers.get_all_arns()) == 1
task_key = mock_executor.active_workers.arn_to_key[ARN1]

# Call Sync 2 times with failures. The task can only fail MAX_RUN_TASK_ATTEMPTS times.
for check_count in range(1, int(AwsEcsExecutor.MAX_RUN_TASK_ATTEMPTS)):
mock_executor.sync_running_tasks()
assert mock_executor.ecs.describe_tasks.call_count == check_count

# Ensure task arn is not removed from active.
assert ARN1 in mock_executor.active_workers.get_all_arns()
# Ensure task gets removed from active_workers.
assert ARN1 not in mock_executor.active_workers.get_all_arns()
# Ensure task gets back on the pending_tasks queue
assert len(mock_executor.pending_tasks) == 1
keys = [task.key for task in mock_executor.pending_tasks]
assert task_key in keys

# Task is neither failed nor succeeded.
fail_mock.assert_not_called()
success_mock.assert_not_called()
mock_executor.attempt_task_runs()

assert len(mock_executor.pending_tasks) == 0
assert len(mock_executor.active_workers.get_all_arns()) == 1
assert ARN1 in mock_executor.active_workers.get_all_arns()
task_key = mock_executor.active_workers.arn_to_key[ARN1]

# Last call should fail the task.
mock_executor.sync_running_tasks()
Expand Down Expand Up @@ -878,6 +1024,7 @@ def _mock_sync(
set_task_state=TaskInstanceState.RUNNING,
) -> None:
"""Mock ECS to the expected state."""
executor.pending_tasks.clear()
self._add_mock_task(executor, ARN1, set_task_state)

response_task_json = {
Expand Down Expand Up @@ -938,9 +1085,13 @@ def test_update_running_tasks(
}
mock_executor.ecs.describe_tasks.return_value = {"tasks": [test_response_task_json], "failures": []}
mock_executor.sync_running_tasks()
assert mock_executor.active_workers.tasks["arn1"].get_task_state() == expected_status
# The task is not removed from active_workers in these states
assert len(mock_executor.active_workers) == 1
if expected_status != State.REMOVED:
assert mock_executor.active_workers.tasks["arn1"].get_task_state() == expected_status
# The task is not removed from active_workers in these states
assert len(mock_executor.active_workers) == 1
else:
# The task is removed from active_workers in this state
assert len(mock_executor.active_workers) == 0

def test_update_running_tasks_success(self, mock_executor):
self._add_mock_task(mock_executor, ARN1)
Expand All @@ -967,6 +1118,7 @@ def test_update_running_tasks_success(self, mock_executor):
mock_success_function.assert_called_once()

def test_update_running_tasks_failed(self, mock_executor, caplog):
AwsEcsExecutor.MAX_RUN_TASK_ATTEMPTS = "1"
caplog.set_level(logging.WARNING)
self._add_mock_task(mock_executor, ARN1)
test_response_task_json = {
Expand Down