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
Original file line number Diff line number Diff line change
Expand Up @@ -491,9 +491,6 @@ def __init__(
self.reattach = reattach
self.number_logs_exception = number_logs_exception

if self.awslogs_region is None:
self.awslogs_region = self.region_name

self.arn: str | None = None
self.container_name: str | None = container_name
self._started_by: str | None = None
Expand All @@ -513,18 +510,18 @@ def __init__(
)
self.stop_task_on_failure = stop_task_on_failure

if self._aws_logs_enabled() and not self.wait_for_completion:
self.log.warning(
"Trying to get logs without waiting for the task to complete is undefined behavior."
)

@staticmethod
def _get_ecs_task_id(task_arn: str | None) -> str | None:
if task_arn is None:
return None
return task_arn.split("/")[-1]

def execute(self, context):
if self._aws_logs_enabled() and not self.wait_for_completion:
self.log.warning(
"Trying to get logs without waiting for the task to complete is undefined behavior."
)

self.log.info(
"Running ECS Task - Task definition: %s - on cluster %s", self.task_definition, self.cluster
)
Expand Down Expand Up @@ -623,7 +620,9 @@ def execute_complete(self, context: Context, event: dict[str, Any] | None = None
self._after_execution()
if self._aws_logs_enabled():
# same behavior as non-deferrable mode, return last line of logs of the task.
logs_client = AwsLogsHook(aws_conn_id=self.aws_conn_id, region_name=self.region_name).conn
logs_client = AwsLogsHook(
aws_conn_id=self.aws_conn_id, region_name=self.resolve_awslogs_region()
).conn
one_log = logs_client.get_log_events(
logGroupName=self.awslogs_group,
logStreamName=self._get_logs_stream_name(),
Expand Down Expand Up @@ -737,13 +736,16 @@ def _get_logs_stream_name(self) -> str:
return f"{self.awslogs_stream_prefix}/{self.container_name}/{self._get_ecs_task_id(self.arn)}"
return f"{self.awslogs_stream_prefix}/{self._get_ecs_task_id(self.arn)}"

def resolve_awslogs_region(self) -> str | None:
return self.awslogs_region if self.awslogs_region is not None else self.region_name

def _get_task_log_fetcher(self) -> AwsTaskLogFetcher:
if not self.awslogs_group:
raise ValueError("must specify awslogs_group to fetch task logs")

return AwsTaskLogFetcher(
aws_conn_id=self.aws_conn_id,
region_name=self.awslogs_region,
region_name=self.resolve_awslogs_region(),
log_group=self.awslogs_group,
log_stream_name=self._get_logs_stream_name(),
fetch_interval=self.awslogs_fetch_interval,
Expand Down
54 changes: 54 additions & 0 deletions providers/amazon/tests/unit/amazon/aws/operators/test_ecs.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,6 +180,16 @@ def test_init(self):
assert self.ecs.task_definition == "t"
assert self.ecs.cluster == "c"
assert self.ecs.overrides == {}
assert self.ecs.awslogs_region is None

def test_get_task_log_fetcher_uses_region_name_when_awslogs_region_not_set(self):
self.set_up_operator(
awslogs_group="awslogs-group", awslogs_stream_prefix="prefix", region_name="region"
)

fetcher = self.ecs._get_task_log_fetcher()

assert fetcher.hook.region_name == "region"

def test_template_fields_overrides(self):
assert self.ecs.template_fields == (
Expand Down Expand Up @@ -406,6 +416,26 @@ def test_task_id_parsing(self):
id = EcsRunTaskOperator._get_ecs_task_id(f"arn:aws:ecs:us-east-1:012345678910:task/{TASK_ID}")
assert id == TASK_ID

@mock.patch.object(EcsBaseOperator, "client")
def test_execute_warns_when_fetching_logs_without_waiting(self, client_mock, caplog):
self.set_up_operator(
awslogs_group="awslogs-group",
awslogs_stream_prefix="prefix",
wait_for_completion=False,
)
caplog.clear()
client_mock.run_task.return_value = RESPONSE_WITHOUT_FAILURES
mock_ti = mock.MagicMock()
mock_context = {"ti": mock_ti, "task_instance": mock_ti}

result = self.ecs.execute(mock_context)

assert result is None
assert (
"Trying to get logs without waiting for the task to complete is undefined behavior."
in caplog.messages
)

@mock.patch.object(EcsBaseOperator, "client")
def test_execute_with_failures(self, client_mock):
resp_failures = deepcopy(RESPONSE_WITHOUT_FAILURES)
Expand Down Expand Up @@ -829,6 +859,30 @@ def test_execute_complete(self, client_mock):
# task gets described to assert its success
client_mock().describe_tasks.assert_called_once_with(cluster="test_cluster", tasks=["my_arn"])

@mock.patch("airflow.providers.amazon.aws.operators.ecs.AwsLogsHook")
@mock.patch.object(EcsRunTaskOperator, "_check_success_task")
def test_execute_complete_uses_awslogs_region(self, check_mock, logs_hook_mock):
self.set_up_operator(
awslogs_group="awslogs-group",
awslogs_region="logs-region",
awslogs_stream_prefix="prefix",
region_name="task-region",
)
logs_hook_mock.return_value.conn.get_log_events.return_value = {"events": [{"message": "Log output"}]}

result = self.ecs.execute_complete(
{},
{
"status": "success",
"task_arn": f"arn:aws:ecs:us-east-1:012345678910:task/{TASK_ID}",
"cluster": "test_cluster",
},
)

assert result == "Log output"
check_mock.assert_called_once_with()
logs_hook_mock.assert_called_once_with(aws_conn_id=self.ecs.aws_conn_id, region_name="logs-region")

@mock.patch.object(EcsBaseOperator, "client")
@mock.patch("airflow.providers.amazon.aws.utils.task_log_fetcher.AwsTaskLogFetcher")
def test_container_name_in_log_stream(self, client_mock, log_fetcher_mock):
Expand Down
1 change: 0 additions & 1 deletion scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
# execute()) MUST remove its entry in the same PR — the hook fails on stale entries.
# Burn-down tracked at https://github.com/apache/airflow/issues/70296
providers/amazon/src/airflow/providers/amazon/aws/operators/appflow.py::AppflowBaseOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py::EcsRunTaskOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py::EmrAddStepsOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStartDbClusterOperator
providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStopDbClusterOperator
Expand Down