diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py b/providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py index 4b71424e42d8a..3bdc3e63cc59b 100644 --- a/providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py +++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/ecs.py @@ -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 @@ -513,11 +510,6 @@ 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: @@ -525,6 +517,11 @@ def _get_ecs_task_id(task_arn: str | None) -> str | 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 ) @@ -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(), @@ -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, diff --git a/providers/amazon/tests/unit/amazon/aws/operators/test_ecs.py b/providers/amazon/tests/unit/amazon/aws/operators/test_ecs.py index 946b086e0fe32..346c29efbaf36 100644 --- a/providers/amazon/tests/unit/amazon/aws/operators/test_ecs.py +++ b/providers/amazon/tests/unit/amazon/aws/operators/test_ecs.py @@ -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 == ( @@ -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) @@ -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): diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index 378be2590c777..e192ef70bafb8 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -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