Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -337,7 +337,6 @@ repos:
^airflow\/providers\/google\/cloud\/operators\/mlengine.py$|
^airflow\/providers\/google\/cloud\/operators\/cloud_storage_transfer_service.py$|
^airflow\/providers\/apache\/spark\/operators\/spark_submit.py\.py$|
^airflow\/providers\/google\/cloud\/operators\/vertex_ai\/auto_ml\.py$|
^airflow\/providers\/apache\/spark\/operators\/spark_submit\.py$|
^airflow\/providers\/databricks\/operators\/databricks_sql\.py$|
)$
Expand Down
16 changes: 14 additions & 2 deletions airflow/providers/google/cloud/operators/vertex_ai/auto_ml.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,14 @@

from typing import TYPE_CHECKING, Sequence

from deprecated import deprecated
from google.api_core.exceptions import NotFound
from google.api_core.gapic_v1.method import DEFAULT, _MethodDefault
from google.cloud.aiplatform import datasets
from google.cloud.aiplatform.models import Model
from google.cloud.aiplatform_v1.types.training_pipeline import TrainingPipeline

from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.providers.google.cloud.hooks.vertex_ai.auto_ml import AutoMLHook
from airflow.providers.google.cloud.links.vertex_ai import (
VertexAIModelLink,
Expand Down Expand Up @@ -607,7 +609,7 @@ class DeleteAutoMLTrainingJobOperator(GoogleCloudBaseOperator):
AutoMLTabularTrainingJob, AutoMLTextTrainingJob, or AutoMLVideoTrainingJob.
"""

template_fields = ("training_pipeline", "region", "project_id", "impersonation_chain")
template_fields = ("training_pipeline_id", "region", "project_id", "impersonation_chain")
Comment thread
shahar1 marked this conversation as resolved.
Outdated

def __init__(
self,
Expand All @@ -623,7 +625,7 @@ def __init__(
**kwargs,
) -> None:
super().__init__(**kwargs)
self.training_pipeline = training_pipeline_id
self.training_pipeline_id = training_pipeline_id
self.region = region
self.project_id = project_id
self.retry = retry
Expand All @@ -632,6 +634,16 @@ def __init__(
self.gcp_conn_id = gcp_conn_id
self.impersonation_chain = impersonation_chain

@property
@deprecated(
reason="`training_pipeline` is deprecated and will be removed in the future. "
"Please use `training_pipeline_id` instead.",
category=AirflowProviderDeprecationWarning,
)
def training_pipeline(self):
"""Alias for ``training_pipeline_id``, used for compatibility (deprecated)."""
return self.training_pipeline_id

def execute(self, context: Context):
hook = AutoMLHook(
gcp_conn_id=self.gcp_conn_id,
Expand Down
24 changes: 24 additions & 0 deletions tests/providers/google/cloud/operators/test_vertex_ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -1028,6 +1028,30 @@ def test_execute(self, mock_hook):
metadata=METADATA,
)

@pytest.mark.db_test
def test_templating(self, create_task_instance_of_operator):
ti = create_task_instance_of_operator(
DeleteAutoMLTrainingJobOperator,
# Templated fields
training_pipeline_id="{{ 'training-pipeline-id' }}",
region="{{ 'region' }}",
project_id="{{ 'project-id' }}",
impersonation_chain="{{ 'impersonation-chain' }}",
# Other parameters
dag_id="test_template_body_templating_dag",
task_id="test_template_body_templating_task",
execution_date=timezone.datetime(2024, 2, 1, tzinfo=timezone.utc),
)
ti.render_templates()
task: DeleteAutoMLTrainingJobOperator = ti.task
assert task.training_pipeline_id == "training-pipeline-id"
assert task.region == "region"
assert task.project_id == "project-id"
assert task.impersonation_chain == "impersonation-chain"

with pytest.warns(AirflowProviderDeprecationWarning):
assert task.training_pipeline == "training-pipeline-id"


class TestVertexAIListAutoMLTrainingJobOperator:
@mock.patch(VERTEX_AI_PATH.format("auto_ml.AutoMLHook"))
Expand Down