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):