diff --git a/airflow/providers/amazon/aws/log/s3_task_handler.py b/airflow/providers/amazon/aws/log/s3_task_handler.py index 831c86417185f..098f17a28a95c 100644 --- a/airflow/providers/amazon/aws/log/s3_task_handler.py +++ b/airflow/providers/amazon/aws/log/s3_task_handler.py @@ -128,8 +128,12 @@ def _read(self, ti, try_number, metadata=None): if logs: return "".join(f"*** {x}\n" for x in messages) + "\n".join(logs), {"end_of_log": True} else: + if metadata and metadata.get("log_pos", 0) > 0: + log_prefix = "" + else: + log_prefix = "*** Falling back to local log\n" local_log, metadata = super()._read(ti, try_number, metadata) - return "*** Falling back to local log\n" + local_log, metadata + return f"{log_prefix}{local_log}", metadata def s3_log_exists(self, remote_log_location: str) -> bool: """ diff --git a/tests/providers/amazon/aws/log/test_s3_task_handler.py b/tests/providers/amazon/aws/log/test_s3_task_handler.py index 4fbe453274a9b..aeca09d36d6a7 100644 --- a/tests/providers/amazon/aws/log/test_s3_task_handler.py +++ b/tests/providers/amazon/aws/log/test_s3_task_handler.py @@ -146,6 +146,33 @@ def test_read_when_s3_log_missing(self): assert actual == expected assert {"end_of_log": True, "log_pos": 0} == metadata[0] + def test_read_when_s3_log_missing_and_log_pos_missing_pre_26(self): + ti = copy.copy(self.ti) + ti.state = TaskInstanceState.SUCCESS + # mock that super class has no _read_remote_logs method + with mock.patch("airflow.providers.amazon.aws.log.s3_task_handler.hasattr", return_value=False): + log, metadata = self.s3_task_handler.read(ti) + assert 1 == len(log) + assert log[0][0][-1].startswith("*** Falling back to local log") + + def test_read_when_s3_log_missing_and_log_pos_zero_pre_26(self): + ti = copy.copy(self.ti) + ti.state = TaskInstanceState.SUCCESS + # mock that super class has no _read_remote_logs method + with mock.patch("airflow.providers.amazon.aws.log.s3_task_handler.hasattr", return_value=False): + log, metadata = self.s3_task_handler.read(ti, metadata={"log_pos": 0}) + assert 1 == len(log) + assert log[0][0][-1].startswith("*** Falling back to local log") + + def test_read_when_s3_log_missing_and_log_pos_over_zero_pre_26(self): + ti = copy.copy(self.ti) + ti.state = TaskInstanceState.SUCCESS + # mock that super class has no _read_remote_logs method + with mock.patch("airflow.providers.amazon.aws.log.s3_task_handler.hasattr", return_value=False): + log, metadata = self.s3_task_handler.read(ti, metadata={"log_pos": 1}) + assert 1 == len(log) + assert not log[0][0][-1].startswith("*** Falling back to local log") + def test_s3_read_when_log_missing(self): handler = self.s3_task_handler url = "s3://bucket/foo"