This is an automated email from the ASF dual-hosted git repository.
vincbeck 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 08153a88f83 Use the operator's AWS settings for deferred SageMaker
tasks (#71857)
08153a88f83 is described below
commit 08153a88f832bb553dd92153d66dd178f58683f9
Author: Sepuri Sai Krishna <[email protected]>
AuthorDate: Fri Aug 21 22:47:49 2026 +0530
Use the operator's AWS settings for deferred SageMaker tasks (#71857)
SageMaker operators inherit region_name, verify and botocore_config from
AwsBaseOperator, but did not hand them to the trigger they defer to. The
deferred half of the task then reached AWS with different settings than the
operator itself, so an explicitly requested region, a custom botocore config or
an explicit SSL verify setting stopped applying once the task was handed to the
triggerer.
---
.../providers/amazon/aws/operators/sagemaker.py | 21 +++++++++++
.../aws/operators/test_sagemaker_endpoint.py | 22 ++++++++++++
.../aws/operators/test_sagemaker_pipeline.py | 42 ++++++++++++++++++++++
.../aws/operators/test_sagemaker_processing.py | 35 ++++++++++++++++++
.../aws/operators/test_sagemaker_training.py | 36 +++++++++++++++++++
.../aws/operators/test_sagemaker_transform.py | 31 ++++++++++++++++
.../amazon/aws/operators/test_sagemaker_tuning.py | 20 +++++++++++
7 files changed, 207 insertions(+)
diff --git
a/providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py
b/providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py
index 066f63eca36..d8c7d0aa06b 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/sagemaker.py
@@ -379,6 +379,9 @@ class SageMakerProcessingOperator(SageMakerBaseOperator):
waiter_delay=self.check_interval,
waiter_max_attempts=self.max_attempts,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
@@ -696,6 +699,9 @@ class SageMakerEndpointOperator(SageMakerBaseOperator):
job_type="endpoint",
waiter_delay=self.check_interval,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
timeout=datetime.timedelta(seconds=self.max_ingestion_time),
@@ -924,6 +930,9 @@ class SageMakerTransformOperator(SageMakerBaseOperator):
waiter_delay=self.check_interval,
waiter_max_attempts=self.max_attempts,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
@@ -1100,6 +1109,9 @@ class SageMakerTuningOperator(SageMakerBaseOperator):
job_type="tuning",
waiter_delay=self.check_interval,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
timeout=(
@@ -1332,6 +1344,9 @@ class SageMakerTrainingOperator(SageMakerBaseOperator):
waiter_delay=self.check_interval,
waiter_max_attempts=self.max_attempts,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
@@ -1478,6 +1493,9 @@ class
SageMakerStartPipelineOperator(SageMakerBaseOperator):
waiter_delay=self.check_interval,
waiter_max_attempts=self.waiter_max_attempts,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
@@ -1577,6 +1595,9 @@ class
SageMakerStopPipelineOperator(SageMakerBaseOperator):
waiter_delay=self.check_interval,
waiter_max_attempts=self.waiter_max_attempts,
aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
),
method_name="execute_complete",
)
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_endpoint.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_endpoint.py
index 965bed467ef..dd894b0d9ed 100644
---
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_endpoint.py
+++
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_endpoint.py
@@ -57,6 +57,11 @@ CONFIG: dict = {
EXPECTED_INTEGER_FIELDS: list[list[str]] = [["EndpointConfig",
"ProductionVariants", "InitialInstanceCount"]]
+REGION_NAME = "eu-west-2"
+VERIFY = False
+BOTOCORE_CONFIG = {"read_timeout": 42}
+
+
class TestSageMakerEndpointOperator:
def setup_method(self):
self.sagemaker = SageMakerEndpointOperator(
@@ -176,5 +181,22 @@ class TestSageMakerEndpointOperator:
assert defer.value.trigger.job_name == "endpoint_name"
assert defer.value.trigger.job_type == "endpoint"
+ @mock.patch.object(SageMakerHook, "create_model")
+ @mock.patch.object(SageMakerHook, "create_endpoint_config")
+ @mock.patch.object(SageMakerHook, "create_endpoint")
+ def test_deferred_trigger_receives_hook_configuration(self,
mock_create_endpoint, _, __):
+ mock_create_endpoint.return_value = {"ResponseMetadata":
{"HTTPStatusCode": 200}}
+ self.sagemaker.deferrable = True
+ self.sagemaker.region_name = REGION_NAME
+ self.sagemaker.verify = VERIFY
+ self.sagemaker.botocore_config = BOTOCORE_CONFIG
+
+ with pytest.raises(TaskDeferred) as exc:
+ self.sagemaker.execute(None)
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
def test_template_fields(self):
validate_template_fields(self.sagemaker)
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_pipeline.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_pipeline.py
index 13f20115c65..fbfd2baafba 100644
---
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_pipeline.py
+++
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_pipeline.py
@@ -36,6 +36,11 @@ if TYPE_CHECKING:
from unittest.mock import MagicMock
+REGION_NAME = "eu-west-2"
+VERIFY = False
+BOTOCORE_CONFIG = {"read_timeout": 42}
+
+
class TestSageMakerStartPipelineOperator:
@mock.patch.object(SageMakerHook, "start_pipeline")
@mock.patch.object(SageMakerHook, "check_status")
@@ -73,6 +78,24 @@ class TestSageMakerStartPipelineOperator:
assert isinstance(defer.value.trigger, SageMakerPipelineTrigger)
assert defer.value.trigger.waiter_type ==
SageMakerPipelineTrigger.Type.COMPLETE
+ @mock.patch.object(SageMakerHook, "start_pipeline")
+ def test_deferred_trigger_receives_hook_configuration(self, start_mock):
+ op = SageMakerStartPipelineOperator(
+ task_id="test_sagemaker_operator",
+ pipeline_name="my_pipeline",
+ deferrable=True,
+ region_name=REGION_NAME,
+ verify=VERIFY,
+ botocore_config=BOTOCORE_CONFIG,
+ )
+
+ with pytest.raises(TaskDeferred) as exc:
+ op.execute({})
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
def test_template_fields(self):
operator = SageMakerStartPipelineOperator(
task_id="test_sagemaker_operator",
@@ -110,6 +133,25 @@ class TestSageMakerStopPipelineOperator:
assert isinstance(defer.value.trigger, SageMakerPipelineTrigger)
assert defer.value.trigger.waiter_type ==
SageMakerPipelineTrigger.Type.STOPPED
+ @mock.patch.object(SageMakerHook, "stop_pipeline")
+ def test_deferred_trigger_receives_hook_configuration(self, stop_mock:
MagicMock):
+ stop_mock.return_value = "Stopping"
+ op = SageMakerStopPipelineOperator(
+ task_id="test_sagemaker_operator",
+ pipeline_exec_arn="my_pipeline_arn",
+ deferrable=True,
+ region_name=REGION_NAME,
+ verify=VERIFY,
+ botocore_config=BOTOCORE_CONFIG,
+ )
+
+ with pytest.raises(TaskDeferred) as exc:
+ op.execute({})
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
def test_template_fields(self):
operator = SageMakerStopPipelineOperator(
task_id="test_sagemaker_operator",
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_processing.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_processing.py
index eb14ed4b3c6..6ce11d7a669 100644
---
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_processing.py
+++
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_processing.py
@@ -95,6 +95,11 @@ EXPECTED_INTEGER_FIELDS: list[list[str]] = [
EXPECTED_STOPPING_CONDITION_INTEGER_FIELDS: list[list[str]] =
[["StoppingCondition", "MaxRuntimeInSeconds"]]
+REGION_NAME = "eu-west-2"
+VERIFY = False
+BOTOCORE_CONFIG = {"read_timeout": 42}
+
+
class TestSageMakerProcessingOperator:
def setup_method(self):
self.processing_config_kwargs = dict(
@@ -271,6 +276,36 @@ class TestSageMakerProcessingOperator:
sagemaker_operator.execute(context=None)
assert isinstance(exc.value.trigger, SageMakerTrigger), "Trigger is
not a SagemakerTrigger"
+ @mock.patch.object(
+ SageMakerHook, "describe_processing_job",
return_value={"ProcessingJobStatus": "InProgress"}
+ )
+ @mock.patch.object(
+ SageMakerHook,
+ "create_processing_job",
+ return_value={
+ "ProcessingJobArn": "test_arn",
+ "ResponseMetadata": {"HTTPStatusCode": 200},
+ },
+ )
+ @mock.patch.object(SageMakerBaseOperator, "_check_if_job_exists",
return_value=False)
+ def test_deferred_trigger_receives_hook_configuration(
+ self, mock_job_exists, mock_processing, mock_describe
+ ):
+ sagemaker_operator = SageMakerProcessingOperator(
+ **self.defer_processing_config_kwargs,
+ config=CREATE_PROCESSING_PARAMS,
+ region_name=REGION_NAME,
+ verify=VERIFY,
+ botocore_config=BOTOCORE_CONFIG,
+ )
+
+ with pytest.raises(TaskDeferred) as exc:
+ sagemaker_operator.execute(context=None)
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
@mock.patch("airflow.providers.amazon.aws.operators.sagemaker.SageMakerProcessingOperator.defer")
@mock.patch.object(
SageMakerHook, "describe_processing_job",
return_value={"ProcessingJobStatus": "Completed"}
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_training.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_training.py
index 77db60fa5d4..abd1eccb848 100644
---
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_training.py
+++
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_training.py
@@ -65,6 +65,11 @@ CREATE_TRAINING_PARAMS = {
}
+REGION_NAME = "eu-west-2"
+VERIFY = False
+BOTOCORE_CONFIG = {"read_timeout": 42}
+
+
class TestSageMakerTrainingOperator:
def setup_method(self):
self.sagemaker = SageMakerTrainingOperator(
@@ -210,6 +215,37 @@ class TestSageMakerTrainingOperator:
self.sagemaker.execute(context=None)
assert isinstance(exc.value.trigger, SageMakerTrigger), "Trigger is
not a SagemakerTrigger"
+ @mock.patch.object(
+ SageMakerHook,
+ "describe_training_job",
+ return_value={
+ "TrainingJobStatus": "Training",
+ "ResourceConfig": {"InstanceCount": 1},
+ "TrainingEndTime": datetime(2023, 5, 15),
+ "TrainingStartTime": datetime(2023, 5, 16),
+ },
+ )
+ @mock.patch.object(SageMakerHook, "create_training_job")
+ def test_deferred_trigger_receives_hook_configuration(self, mock_training,
mock_describe_training_job):
+ mock_training.return_value = {
+ "TrainingJobArn": "test_arn",
+ "ResponseMetadata": {"HTTPStatusCode": 200},
+ }
+ self.sagemaker.deferrable = True
+ self.sagemaker.wait_for_completion = True
+ self.sagemaker.check_if_job_exists = False
+ self.sagemaker.print_log = False
+ self.sagemaker.region_name = REGION_NAME
+ self.sagemaker.verify = VERIFY
+ self.sagemaker.botocore_config = BOTOCORE_CONFIG
+
+ with pytest.raises(TaskDeferred) as exc:
+ self.sagemaker.execute(context=None)
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
@mock.patch.object(
SageMakerHook,
"describe_training_job",
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_transform.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_transform.py
index de8ed7841e8..bfc859d43ee 100644
---
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_transform.py
+++
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_transform.py
@@ -72,6 +72,11 @@ CONFIG: dict = {"Model": CREATE_MODEL_PARAMS, "Transform":
CREATE_TRANSFORM_PARA
MOCK_UNIX_TIME: int = 1234567890123456789 # reproducible time for testing
time.time_ns()
+REGION_NAME = "eu-west-2"
+VERIFY = False
+BOTOCORE_CONFIG = {"read_timeout": 42}
+
+
class TestSageMakerTransformOperator:
def setup_method(self):
self.sagemaker = SageMakerTransformOperator(
@@ -409,6 +414,32 @@ class TestSageMakerTransformOperator:
assert isinstance(exc.value.trigger, SageMakerTrigger), "Trigger is
not a SagemakerTrigger"
+ @mock.patch.object(
+ SageMakerHook, "describe_transform_job",
return_value={"TransformJobStatus": "InProgress"}
+ )
+ @mock.patch.object(SageMakerHook, "create_transform_job")
+ @mock.patch.object(SageMakerHook, "create_model")
+ def test_deferred_trigger_receives_hook_configuration(
+ self, _, mock_transform, mock_describe_transform_job
+ ):
+ mock_transform.return_value = {
+ "TransformJobArn": "test_arn",
+ "ResponseMetadata": {"HTTPStatusCode": 200},
+ }
+ self.sagemaker.deferrable = True
+ self.sagemaker.wait_for_completion = True
+ self.sagemaker.check_if_job_exists = False
+ self.sagemaker.region_name = REGION_NAME
+ self.sagemaker.verify = VERIFY
+ self.sagemaker.botocore_config = BOTOCORE_CONFIG
+
+ with pytest.raises(TaskDeferred) as exc:
+ self.sagemaker.execute(context=None)
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
@mock.patch.object(SageMakerHook, "describe_transform_job")
@mock.patch.object(SageMakerHook, "create_model")
@mock.patch.object(SageMakerHook, "describe_model")
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_tuning.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_tuning.py
index 39f1915f628..ef8e5889fed 100644
--- a/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_tuning.py
+++ b/providers/amazon/tests/unit/amazon/aws/operators/test_sagemaker_tuning.py
@@ -75,6 +75,11 @@ CREATE_TUNING_PARAMS: dict = {
}
+REGION_NAME = "eu-west-2"
+VERIFY = False
+BOTOCORE_CONFIG = {"read_timeout": 42}
+
+
class TestSageMakerTuningOperator:
def setup_method(self):
self.sagemaker = SageMakerTuningOperator(
@@ -123,5 +128,20 @@ class TestSageMakerTuningOperator:
assert defer.value.trigger.job_name == "job_name"
assert defer.value.trigger.job_type == "tuning"
+ @mock.patch.object(SageMakerHook, "create_tuning_job")
+ def test_deferred_trigger_receives_hook_configuration(self, create_mock):
+ create_mock.return_value = {"ResponseMetadata": {"HTTPStatusCode":
200}}
+ self.sagemaker.deferrable = True
+ self.sagemaker.region_name = REGION_NAME
+ self.sagemaker.verify = VERIFY
+ self.sagemaker.botocore_config = BOTOCORE_CONFIG
+
+ with pytest.raises(TaskDeferred) as exc:
+ self.sagemaker.execute(None)
+
+ assert exc.value.trigger.region_name == REGION_NAME
+ assert exc.value.trigger.verify == VERIFY
+ assert exc.value.trigger.botocore_config == BOTOCORE_CONFIG
+
def test_template_fields(self):
validate_template_fields(self.sagemaker)