This is an automated email from the ASF dual-hosted git repository.

o-nikolas pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new 4ceb1fff01e Preserve SageMaker job lifecycle states in trigger events 
(#71653)
4ceb1fff01e is described below

commit 4ceb1fff01e8ee18f3fc7d210de7c19b989be142
Author: SameerMesiah97 <[email protected]>
AuthorDate: Sat Sep 19 00:59:35 2026 +0100

    Preserve SageMaker job lifecycle states in trigger events (#71653)
    
    Translate structured AWS waiter outcomes into failed, stopped, and 
timed-out SageMaker trigger events while retaining the shared waiter lifecycle 
handling. Add unit tests for the SageMaker-specific event translation.
---
 .../providers/amazon/aws/triggers/sagemaker.py     |  36 +++++++
 .../unit/amazon/aws/triggers/test_sagemaker.py     | 114 +++++++++++++++++++++
 2 files changed, 150 insertions(+)

diff --git 
a/providers/amazon/src/airflow/providers/amazon/aws/triggers/sagemaker.py 
b/providers/amazon/src/airflow/providers/amazon/aws/triggers/sagemaker.py
index 271fe8b51fd..cbaeb7a1c04 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/triggers/sagemaker.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/triggers/sagemaker.py
@@ -27,6 +27,7 @@ from typing import TYPE_CHECKING
 from botocore.exceptions import WaiterError
 
 from airflow.exceptions import AirflowProviderDeprecationWarning
+from airflow.providers.amazon.aws.exceptions import WaiterMaxAttemptsError, 
WaiterTerminalFailure
 from airflow.providers.amazon.aws.hooks.sagemaker import SageMakerHook
 from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger
 from airflow.providers.common.compat.sdk import AirflowException
@@ -108,6 +109,41 @@ class SageMakerTrigger(AwsBaseWaiterTrigger):
             config=self.botocore_config,
         )
 
+    def _event_from_exception(self, error: AirflowException) -> TriggerEvent:
+
+        if isinstance(error, WaiterMaxAttemptsError):
+            return TriggerEvent(
+                {
+                    "status": "timeout",
+                    "job_name": self.job_name,
+                    "message": str(error),
+                }
+            )
+
+        if isinstance(error, WaiterTerminalFailure):
+            response = error.last_response
+            status = response.get(self._get_response_status_key(self.job_type))
+
+            if status == "Failed":
+                return TriggerEvent(
+                    {
+                        "status": "failed",
+                        "job_name": self.job_name,
+                        "message": response.get("FailureReason") or str(error),
+                    }
+                )
+
+            if status == "Stopped":
+                return TriggerEvent(
+                    {
+                        "status": "stopped",
+                        "job_name": self.job_name,
+                        "message": str(error),
+                    }
+                )
+
+        return super()._event_from_exception(error)
+
     @staticmethod
     def _get_job_type_waiter(job_type: str) -> str:
         return {
diff --git a/providers/amazon/tests/unit/amazon/aws/triggers/test_sagemaker.py 
b/providers/amazon/tests/unit/amazon/aws/triggers/test_sagemaker.py
index 68cf554fbb1..b736505de3c 100644
--- a/providers/amazon/tests/unit/amazon/aws/triggers/test_sagemaker.py
+++ b/providers/amazon/tests/unit/amazon/aws/triggers/test_sagemaker.py
@@ -23,8 +23,10 @@ import pytest
 from botocore.exceptions import WaiterError
 
 from airflow.exceptions import AirflowProviderDeprecationWarning
+from airflow.providers.amazon.aws.exceptions import WaiterMaxAttemptsError, 
WaiterTerminalFailure
 from airflow.providers.amazon.aws.hooks.sagemaker import SageMakerHook
 from airflow.providers.amazon.aws.triggers.sagemaker import 
SageMakerPipelineTrigger, SageMakerTrigger
+from airflow.providers.common.compat.sdk import AirflowException
 from airflow.triggers.base import TriggerEvent
 
 JOB_NAME = "job_name"
@@ -121,6 +123,118 @@ class TestSagemakerTrigger:
 
         assert response == TriggerEvent({"status": "success", "job_name": 
JOB_NAME})
 
+    @pytest.mark.parametrize(
+        ("aws_status", "failure_reason", "expected_status", 
"expected_message"),
+        [
+            pytest.param("Failed", None, "failed", "SageMaker job failed", 
id="failed"),
+            pytest.param(
+                "Failed",
+                "Algorithm error",
+                "failed",
+                "Algorithm error",
+                id="failed-with-reason",
+            ),
+            pytest.param("Stopped", None, "stopped", "SageMaker job failed", 
id="stopped"),
+        ],
+    )
+    def test_event_from_exception_terminal_state(
+        self,
+        aws_status,
+        failure_reason,
+        expected_status,
+        expected_message,
+    ):
+        trigger = SageMakerTrigger(
+            job_name=JOB_NAME,
+            job_type=JOB_TYPE,
+            waiter_delay=WAITER_DELAY,
+            waiter_max_attempts=WAITER_MAX_ATTEMPTS,
+            aws_conn_id=AWS_CONN_ID,
+        )
+        error = WaiterTerminalFailure(
+            "SageMaker job failed",
+            last_response={"TrainingJobStatus": aws_status},
+        )
+
+        last_response = {"TrainingJobStatus": aws_status}
+        if failure_reason:
+            last_response["FailureReason"] = failure_reason
+
+        error = WaiterTerminalFailure(
+            "SageMaker job failed",
+            last_response=last_response,
+        )
+
+        response = trigger._event_from_exception(error)
+
+        assert response == TriggerEvent(
+            {
+                "status": expected_status,
+                "job_name": JOB_NAME,
+                "message": expected_message,
+            }
+        )
+
+    def test_event_from_exception_timeout(self):
+        trigger = SageMakerTrigger(
+            job_name=JOB_NAME,
+            job_type=JOB_TYPE,
+            waiter_delay=WAITER_DELAY,
+            waiter_max_attempts=WAITER_MAX_ATTEMPTS,
+            aws_conn_id=AWS_CONN_ID,
+        )
+        error = WaiterMaxAttemptsError("Waiter error: max attempts reached")
+
+        response = trigger._event_from_exception(error)
+
+        assert response == TriggerEvent(
+            {
+                "status": "timeout",
+                "job_name": JOB_NAME,
+                "message": "Waiter error: max attempts reached",
+            }
+        )
+
+    @pytest.mark.parametrize(
+        "error",
+        [
+            pytest.param(
+                WaiterTerminalFailure(
+                    "SageMaker job failed",
+                    last_response={},
+                ),
+                id="missing-status",
+            ),
+            pytest.param(
+                WaiterTerminalFailure(
+                    "SageMaker job failed",
+                    last_response={"TrainingJobStatus": "Unexpected"},
+                ),
+                id="unknown-status",
+            ),
+            pytest.param(
+                AirflowException("SageMaker job failed"),
+                id="generic-error",
+            ),
+        ],
+    )
+    def test_event_from_exception_falls_back_to_error(self, error):
+        trigger = SageMakerTrigger(
+            job_name=JOB_NAME,
+            job_type=JOB_TYPE,
+            waiter_delay=WAITER_DELAY,
+            waiter_max_attempts=WAITER_MAX_ATTEMPTS,
+            aws_conn_id=AWS_CONN_ID,
+        )
+
+        assert trigger._event_from_exception(error) == TriggerEvent(
+            {
+                "status": "error",
+                "message": "SageMaker job failed",
+                "job_name": JOB_NAME,
+            }
+        )
+
 
 class TestSagemakerPipelineTrigger:
     def test_serialize(self):

Reply via email to