This is an automated email from the ASF dual-hosted git repository.
potiuk 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 8daf7b102f8 Add CreateCustomJobOperator for VertexAI service (#71875)
8daf7b102f8 is described below
commit 8daf7b102f83032ee10166fd334d9dab2f541be8
Author: Maksim <[email protected]>
AuthorDate: Wed Sep 9 00:14:05 2026 +0200
Add CreateCustomJobOperator for VertexAI service (#71875)
---
generated/known_airflow_exceptions.txt | 6 +-
providers/google/provider.yaml | 1 +
.../google/cloud/hooks/vertex_ai/custom_job.py | 71 +++++++-
.../providers/google/cloud/links/vertex_ai.py | 11 ++
.../google/cloud/operators/vertex_ai/custom_job.py | 123 +++++++++++++
.../providers/google/cloud/triggers/vertex_ai.py | 67 +++++++
.../airflow/providers/google/get_provider_info.py | 1 +
.../example_vertex_ai_create_custom_job.py | 134 ++++++++++++++
.../cloud/hooks/vertex_ai/test_custom_job.py | 202 ++++++++++++++++++++-
.../unit/google/cloud/operators/test_vertex_ai.py | 49 +++++
.../unit/google/cloud/triggers/test_vertex_ai.py | 30 +++
11 files changed, 682 insertions(+), 13 deletions(-)
diff --git a/generated/known_airflow_exceptions.txt
b/generated/known_airflow_exceptions.txt
index 607b038da83..fa58bca8214 100644
--- a/generated/known_airflow_exceptions.txt
+++ b/generated/known_airflow_exceptions.txt
@@ -228,14 +228,13 @@
providers/google/src/airflow/providers/google/cloud/hooks/gdm.py::1
providers/google/src/airflow/providers/google/cloud/hooks/kubernetes_engine.py::2
providers/google/src/airflow/providers/google/cloud/hooks/looker.py::8
providers/google/src/airflow/providers/google/cloud/hooks/managed_kafka.py::1
-providers/google/src/airflow/providers/google/cloud/hooks/mlengine.py::2
providers/google/src/airflow/providers/google/cloud/hooks/spanner.py::5
providers/google/src/airflow/providers/google/cloud/hooks/stackdriver.py::2
providers/google/src/airflow/providers/google/cloud/hooks/tasks.py::3
providers/google/src/airflow/providers/google/cloud/hooks/translate.py::3
providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/auto_ml.py::3
providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/batch_prediction_job.py::1
-providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/custom_job.py::8
+providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/custom_job.py::9
providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/feature_store.py::2
providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/hyperparameter_tuning_job.py::1
providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/model_service.py::2
@@ -408,7 +407,8 @@
providers/trino/src/airflow/providers/trino/hooks/trino.py::1
providers/vespa/src/airflow/providers/vespa/operators/vespa_ingest.py::1
providers/ydb/src/airflow/providers/ydb/hooks/ydb.py::1
providers/ydb/src/airflow/providers/ydb/operators/ydb.py::1
-scripts/ci/prek/check_new_airflow_exception_usage.py::5
+scripts/ci/prek/check_new_airflow_exception_usage.py::4
+scripts/tests/ci/prek/test_check_new_airflow_exception_usage.py::9
task-sdk/src/airflow/sdk/bases/sensor.py::1
task-sdk/src/airflow/sdk/bases/skipmixin.py::3
task-sdk/src/airflow/sdk/crypto.py::1
diff --git a/providers/google/provider.yaml b/providers/google/provider.yaml
index 04db221c30d..b10937531fd 100644
--- a/providers/google/provider.yaml
+++ b/providers/google/provider.yaml
@@ -1439,6 +1439,7 @@ extra-links:
- airflow.providers.google.cloud.links.vertex_ai.VertexAIPipelineJobListLink
- airflow.providers.google.cloud.links.vertex_ai.VertexAIRayClusterLink
- airflow.providers.google.cloud.links.vertex_ai.VertexAIRayClusterListLink
+ - airflow.providers.google.cloud.links.vertex_ai.VertexAICustomJobLink
- airflow.providers.google.cloud.links.workflows.WorkflowsWorkflowDetailsLink
- airflow.providers.google.cloud.links.workflows.WorkflowsListOfWorkflowsLink
- airflow.providers.google.cloud.links.workflows.WorkflowsExecutionLink
diff --git
a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/custom_job.py
b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/custom_job.py
index ea446d3f5e2..ad949662e89 100644
---
a/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/custom_job.py
+++
b/providers/google/src/airflow/providers/google/cloud/hooks/vertex_ai/custom_job.py
@@ -20,6 +20,7 @@
from __future__ import annotations
import asyncio
+import time
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any
@@ -3093,17 +3094,63 @@ class CustomJobHook(GoogleBaseHook, OperationHelper):
)
return result
+ @GoogleBaseHook.fallback_to_default_project_id
+ def wait_for_custom_job(
+ self,
+ project_id: str,
+ region: str,
+ custom_job_id: str,
+ poll_interval: int = 10,
+ retry: Retry | _MethodDefault = DEFAULT,
+ timeout: float | None = None,
+ metadata: Sequence[tuple[str, str]] = (),
+ ) -> CustomJob:
+ """
+ Wait until an VertexAI Custom job completes.
+
+ :param project_id: Required. The ID of the Google Cloud project.
+ :param region: Required. The ID of the Google Cloud location that the
service belongs to.
+ :param custom_job_id: Required. The ID of the CustomJob to wait.
+ :param poll_interval: Time, in seconds, to wait between checks.
+ :param retry: Designation of what errors, if any, should be retried.
+ :param timeout: The timeout for this request.
+ :param metadata: Strings which should be sent along with the request
as metadata.
+ """
+ CUSTOB_JOB_FAILED_STATES = {
+ JobState.JOB_STATE_FAILED: lambda: custom_job.error.message, #
type: ignore
+ JobState.JOB_STATE_CANCELLED: lambda: "The CustomJob has been
cancelled.",
+ JobState.JOB_STATE_PAUSED: lambda: "The CustomJob has been
stopped, and can be resumed.",
+ JobState.JOB_STATE_EXPIRED: lambda: "The CustomJob has expired.",
+ JobState.JOB_STATE_PARTIALLY_SUCCEEDED: lambda:
custom_job.error.message, # type: ignore
+ }
+
+ while True:
+ self.log.info("Waiting for custom job with id %s", custom_job_id)
+ try:
+ custom_job = self.get_custom_job(
+ project_id=project_id,
+ region=region,
+ custom_job=custom_job_id,
+ retry=retry,
+ timeout=timeout,
+ metadata=metadata,
+ )
+ except Exception as ex:
+ self.log.exception("Exception occurred while waiting job %s",
custom_job_id)
+ raise AirflowException(ex)
+ self.log.info("Status of the custom job %s is %s",
custom_job.name, custom_job.state.name)
+ if custom_job.state == JobState.JOB_STATE_SUCCEEDED:
+ return custom_job
+ if custom_job.state in CUSTOB_JOB_FAILED_STATES:
+ raise
RuntimeError(CUSTOB_JOB_FAILED_STATES[custom_job.state]())
+ self.log.info("Sleeping for %s seconds.", poll_interval)
+ time.sleep(poll_interval)
+
class CustomJobAsyncHook(GoogleBaseAsyncHook):
"""Async hook for Custom Job Service Client."""
sync_hook_class = CustomJobHook
- JOB_COMPLETE_STATES = {
- JobState.JOB_STATE_CANCELLED,
- JobState.JOB_STATE_FAILED,
- JobState.JOB_STATE_PAUSED,
- JobState.JOB_STATE_SUCCEEDED,
- }
PIPELINE_COMPLETE_STATES = (
PipelineState.PIPELINE_STATE_CANCELLED,
PipelineState.PIPELINE_STATE_FAILED,
@@ -3238,6 +3285,14 @@ class CustomJobAsyncHook(GoogleBaseAsyncHook):
poll_interval: int = 10,
) -> types.CustomJob:
"""Make async calls to Vertex AI to check the custom job state until
it is complete."""
+ CUSTOB_JOB_FAILED_STATES = {
+ JobState.JOB_STATE_FAILED: lambda: job.error.message,
+ JobState.JOB_STATE_CANCELLED: lambda: "The CustomJob has been
cancelled.",
+ JobState.JOB_STATE_PAUSED: lambda: "The CustomJob has been
stopped, and can be resumed.",
+ JobState.JOB_STATE_EXPIRED: lambda: "The CustomJob has expired.",
+ JobState.JOB_STATE_PARTIALLY_SUCCEEDED: lambda: job.error.message,
+ }
+
client = await self.get_job_service_client(region=location)
while True:
try:
@@ -3255,8 +3310,10 @@ class CustomJobAsyncHook(GoogleBaseAsyncHook):
self.log.exception("Exception occurred while requesting job
%s", job_id)
raise AirflowException(ex)
self.log.info("Status of the custom job %s is %s", job.name,
job.state.name)
- if job.state in self.JOB_COMPLETE_STATES:
+ if job.state == JobState.JOB_STATE_SUCCEEDED:
return job
+ if job.state in CUSTOB_JOB_FAILED_STATES:
+ raise RuntimeError(CUSTOB_JOB_FAILED_STATES[job.state]())
self.log.info("Sleeping for %s seconds.", poll_interval)
await asyncio.sleep(poll_interval)
diff --git
a/providers/google/src/airflow/providers/google/cloud/links/vertex_ai.py
b/providers/google/src/airflow/providers/google/cloud/links/vertex_ai.py
index d749dda1e18..32d122f9a72 100644
--- a/providers/google/src/airflow/providers/google/cloud/links/vertex_ai.py
+++ b/providers/google/src/airflow/providers/google/cloud/links/vertex_ai.py
@@ -58,6 +58,9 @@ VERTEX_AI_RAY_CLUSTER_LINK = (
VERTEX_AI_BASE_LINK +
"/locations/{location}/ray-clusters/{cluster_id}?project={project_id}"
)
VERTEX_AI_RAY_CLUSTER_LIST_LINK = VERTEX_AI_BASE_LINK +
"/ray?project={project_id}"
+VERTEX_AI_CUSTOM_JOB_LINK = (
+ VERTEX_AI_BASE_LINK +
"/locations/{region}/training/{custom_job_id}/cpu?project={project_id}"
+)
class VertexAIModelLink(BaseGoogleLink):
@@ -202,3 +205,11 @@ class VertexAIRayClusterListLink(BaseGoogleLink):
name = "Ray Cluster List"
key = "ray_cluster_list_conf"
format_str = VERTEX_AI_RAY_CLUSTER_LIST_LINK
+
+
+class VertexAICustomJobLink(BaseGoogleLink):
+ """Helper class for constructing Vertex AI CustomJob link."""
+
+ name = "Vertex AI Custom Job"
+ key = "custom_job_conf"
+ format_str = VERTEX_AI_CUSTOM_JOB_LINK
diff --git
a/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/custom_job.py
b/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/custom_job.py
index ec88395d975..953fe8441ea 100644
---
a/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/custom_job.py
+++
b/providers/google/src/airflow/providers/google/cloud/operators/vertex_ai/custom_job.py
@@ -26,12 +26,14 @@ from typing import TYPE_CHECKING, Any
from google.api_core.exceptions import NotFound
from google.api_core.gapic_v1.method import DEFAULT, _MethodDefault
from google.cloud.aiplatform.models import Model
+from google.cloud.aiplatform_v1.types.custom_job import CustomJob
from google.cloud.aiplatform_v1.types.dataset import Dataset
from google.cloud.aiplatform_v1.types.training_pipeline import TrainingPipeline
from airflow.providers.common.compat.sdk import AirflowException, conf
from airflow.providers.google.cloud.hooks.vertex_ai.custom_job import
CustomJobHook
from airflow.providers.google.cloud.links.vertex_ai import (
+ VertexAICustomJobLink,
VertexAIModelLink,
VertexAITrainingLink,
VertexAITrainingPipelinesLink,
@@ -39,6 +41,7 @@ from airflow.providers.google.cloud.links.vertex_ai import (
from airflow.providers.google.cloud.operators.cloud_base import
GoogleCloudBaseOperator
from airflow.providers.google.cloud.triggers.vertex_ai import (
CustomContainerTrainingJobTrigger,
+ CustomJobTrigger,
CustomPythonPackageTrainingJobTrigger,
CustomTrainingJobTrigger,
)
@@ -1793,3 +1796,123 @@ class
ListCustomTrainingJobOperator(GoogleCloudBaseOperator):
)
VertexAITrainingPipelinesLink.persist(context=context)
return [TrainingPipeline.to_dict(result) for result in results]
+
+
+class CreateCustomJobOperator(GoogleCloudBaseOperator):
+ """
+ Create a CustomJob. A created CustomJob right away will be attempted to be
run.
+
+ :param project_id: Required. The ID of the Google Cloud project that the
service belongs to.
+ :param region: Required. The ID of the Google Cloud region that the
service belongs to.
+ :param custom_job: Required. The CustomJob to create.
+ :param retry: Designation of what errors, if any, should be retried.
+ :param timeout: The timeout for this request.
+ :param metadata: Strings which should be sent along with the request as
metadata.
+ """
+
+ template_fields = ("region", "project_id", "custom_job",
"impersonation_chain")
+ operator_extra_links = (VertexAICustomJobLink(),)
+
+ def __init__(
+ self,
+ *,
+ region: str,
+ project_id: str,
+ custom_job: CustomJob | dict,
+ retry: Retry | _MethodDefault = DEFAULT,
+ timeout: float | None = None,
+ metadata: Sequence[tuple[str, str]] = (),
+ gcp_conn_id: str = "google_cloud_default",
+ impersonation_chain: str | Sequence[str] | None = None,
+ deferrable: bool = conf.getboolean("operators", "default_deferrable",
fallback=False),
+ poll_interval: int = 10,
+ **kwargs,
+ ) -> None:
+ super().__init__(**kwargs)
+ self.region = region
+ self.project_id = project_id
+ self.custom_job = custom_job
+ self.retry = retry
+ self.timeout = timeout
+ self.metadata = metadata
+ self.gcp_conn_id = gcp_conn_id
+ self.impersonation_chain = impersonation_chain
+ self.deferrable = deferrable
+ self.poll_interval = poll_interval
+ self.custom_job_id: str | None = None
+
+ @property
+ def extra_links_params(self) -> dict[str, Any]:
+ return {
+ "region": self.region,
+ "project_id": self.project_id,
+ }
+
+ @cached_property
+ def hook(self) -> CustomJobHook:
+ return CustomJobHook(
+ gcp_conn_id=self.gcp_conn_id,
+ impersonation_chain=self.impersonation_chain,
+ )
+
+ def execute(self, context: Context):
+ self.log.info("Creating CustomJob")
+ custom_job_obj = self.hook.create_custom_job(
+ project_id=self.project_id,
+ region=self.region,
+ custom_job=self.custom_job,
+ retry=self.retry,
+ timeout=self.timeout,
+ metadata=self.metadata,
+ )
+ self.custom_job_id =
self.hook.extract_custom_job_id(custom_job_name=custom_job_obj.name)
+ self.log.info("Custom job was created. Job id: %s", self.custom_job_id)
+ context["ti"].xcom_push(key="custom_job_id", value=self.custom_job_id)
+ VertexAICustomJobLink.persist(context=context,
custom_job_id=self.custom_job_id)
+
+ if self.deferrable:
+ self.defer(
+ trigger=CustomJobTrigger(
+ gcp_conn_id=self.gcp_conn_id,
+ project_id=self.project_id,
+ location=self.region,
+ custom_job_id=self.custom_job_id,
+ poll_interval=self.poll_interval,
+ impersonation_chain=self.impersonation_chain,
+ ),
+ method_name="execute_complete",
+ )
+
+ custom_job_obj = self.hook.wait_for_custom_job(
+ project_id=self.project_id,
+ region=self.region,
+ custom_job_id=self.custom_job_id,
+ poll_interval=self.poll_interval,
+ retry=self.retry,
+ timeout=self.timeout,
+ metadata=self.metadata,
+ )
+
+ self.log.info("Custom Job %s completed.", self.custom_job_id)
+ custom_job = CustomJob.to_dict(custom_job_obj)
+ return custom_job
+
+ def execute_complete(self, context: Context, event: dict[str, Any]) ->
dict[str, Any]:
+ if event["status"] == "error":
+ raise RuntimeError(event["message"])
+ self.log.info(event["message"])
+ return event["custom_job"]
+
+ def on_kill(self) -> None:
+ """Act as a callback called when the operator is killed; cancel any
running job."""
+ if self.custom_job_id:
+ self.log.info("Cancelling CustomJob with JobID: %s",
self.custom_job_id)
+ self.hook.cancel_custom_job(
+ project_id=self.project_id,
+ region=self.region,
+ custom_job=self.custom_job_id,
+ retry=self.retry,
+ timeout=self.timeout,
+ metadata=self.metadata,
+ )
+ self.log.info("Custom Job %s canceled.", self.custom_job_id)
diff --git
a/providers/google/src/airflow/providers/google/cloud/triggers/vertex_ai.py
b/providers/google/src/airflow/providers/google/cloud/triggers/vertex_ai.py
index 6e7ecff02fc..1ba0819ac4d 100644
--- a/providers/google/src/airflow/providers/google/cloud/triggers/vertex_ai.py
+++ b/providers/google/src/airflow/providers/google/cloud/triggers/vertex_ai.py
@@ -401,3 +401,70 @@ class
CustomPythonPackageTrainingJobTrigger(BaseVertexAIJobTrigger):
poll_interval=self.poll_interval,
)
return pipeline
+
+
+class CustomJobTrigger(BaseTrigger):
+ """Trigger that waits until a Vertex AI Custom job completes."""
+
+ def __init__(
+ self,
+ project_id: str,
+ location: str,
+ custom_job_id: str,
+ gcp_conn_id: str = "google_cloud_default",
+ impersonation_chain: str | Sequence[str] | None = None,
+ poll_interval: int = 10,
+ ):
+ super().__init__()
+ self.project_id = project_id
+ self.location = location
+ self.custom_job_id = custom_job_id
+ self.gcp_conn_id = gcp_conn_id
+ self.impersonation_chain = impersonation_chain
+ self.poll_interval = poll_interval
+
+ def serialize(self) -> tuple[str, dict[str, Any]]:
+ return (
+
"airflow.providers.google.cloud.triggers.vertex_ai.CustomJobTrigger",
+ {
+ "project_id": self.project_id,
+ "location": self.location,
+ "custom_job_id": self.custom_job_id,
+ "gcp_conn_id": self.gcp_conn_id,
+ "impersonation_chain": self.impersonation_chain,
+ "poll_interval": self.poll_interval,
+ },
+ )
+
+ @cached_property
+ def async_hook(self) -> CustomJobAsyncHook:
+ return CustomJobAsyncHook(
+ gcp_conn_id=self.gcp_conn_id,
+ impersonation_chain=self.impersonation_chain,
+ )
+
+ async def run(self) -> AsyncIterator[TriggerEvent]:
+ try:
+ custom_job = await self.async_hook.wait_for_custom_job(
+ project_id=self.project_id,
+ location=self.location,
+ job_id=self.custom_job_id,
+ poll_interval=self.poll_interval,
+ )
+ except (AirflowException, RuntimeError) as ex:
+ yield TriggerEvent(
+ {
+ "status": "error",
+ "message": str(ex),
+ }
+ )
+ return
+
+ message = f"Custom Job {custom_job.name} completed with status
{custom_job.state.name}"
+ yield TriggerEvent(
+ {
+ "status": "success",
+ "message": message,
+ "custom_job": types.custom_job.CustomJob.to_dict(custom_job),
+ }
+ )
diff --git a/providers/google/src/airflow/providers/google/get_provider_info.py
b/providers/google/src/airflow/providers/google/get_provider_info.py
index d4f607dc85a..048f5023f99 100644
--- a/providers/google/src/airflow/providers/google/get_provider_info.py
+++ b/providers/google/src/airflow/providers/google/get_provider_info.py
@@ -1649,6 +1649,7 @@ def get_provider_info():
"airflow.providers.google.cloud.links.vertex_ai.VertexAIPipelineJobListLink",
"airflow.providers.google.cloud.links.vertex_ai.VertexAIRayClusterLink",
"airflow.providers.google.cloud.links.vertex_ai.VertexAIRayClusterListLink",
+
"airflow.providers.google.cloud.links.vertex_ai.VertexAICustomJobLink",
"airflow.providers.google.cloud.links.workflows.WorkflowsWorkflowDetailsLink",
"airflow.providers.google.cloud.links.workflows.WorkflowsListOfWorkflowsLink",
"airflow.providers.google.cloud.links.workflows.WorkflowsExecutionLink",
diff --git
a/providers/google/tests/system/google/cloud/vertex_ai/example_vertex_ai_create_custom_job.py
b/providers/google/tests/system/google/cloud/vertex_ai/example_vertex_ai_create_custom_job.py
new file mode 100644
index 00000000000..85fc8d97542
--- /dev/null
+++
b/providers/google/tests/system/google/cloud/vertex_ai/example_vertex_ai_create_custom_job.py
@@ -0,0 +1,134 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+
+"""Example Airflow DAG for Google Vertex AI service testing create Custom Jobs
operator."""
+
+from __future__ import annotations
+
+import os
+from datetime import datetime
+
+from airflow.models.dag import DAG
+from airflow.providers.google.cloud.operators.vertex_ai.custom_job import (
+ CreateCustomJobOperator,
+)
+
+ENV_ID = os.environ.get("SYSTEM_TESTS_ENV_ID", "default")
+PROJECT_ID = os.environ.get("SYSTEM_TESTS_GCP_PROJECT", "default")
+REGION = "us-central1"
+DAG_ID = "vertex_ai_create_custom_job"
+
+REPLICA_COUNT = 1
+MACHINE_TYPE = "n1-standard-4"
+ACCELERATOR_TYPE = "ACCELERATOR_TYPE_UNSPECIFIED"
+ACCELERATOR_COUNT = 0
+IMAGE_URI = "us-docker.pkg.dev/vertex-ai/training/tf-cpu.2-16.py310:latest"
+
+test_python_code = (
+ "import sys; "
+ "print('=== VERTEX AI RAW CUSTOM_JOB RUNNING SUCCESSFULLY ==='); "
+ "import math; "
+ "print(f'Sanity check calculation (pi): {math.pi}'); "
+ "print('=== TEST COMPLETED CLEANLY ==='); "
+ "sys.exit(0);"
+)
+
+
+with DAG(
+ DAG_ID,
+ schedule="@once",
+ start_date=datetime(2021, 1, 1),
+ catchup=False,
+ tags=["example", "vertex_ai", "custom_job"],
+) as dag:
+ create_custom_job = CreateCustomJobOperator(
+ task_id="create_custom_job",
+ region=REGION,
+ project_id=PROJECT_ID,
+ custom_job={
+ "display_name": f"{DAG_ID}_{ENV_ID}",
+ "labels": {
+ "vertex_pipelines": "",
+ "airflow_dag_id": DAG_ID,
+ "airflow_task_id": DAG_ID,
+ },
+ "job_spec": {
+ "scheduling": {"disable_retries": True},
+ "worker_pool_specs": [
+ {
+ "machine_spec": {
+ "machine_type": MACHINE_TYPE,
+ "accelerator_type": ACCELERATOR_TYPE,
+ "accelerator_count": ACCELERATOR_COUNT,
+ },
+ "replica_count": REPLICA_COUNT,
+ "container_spec": {
+ "image_uri": IMAGE_URI,
+ "command": ["python3", "-c", test_python_code],
+ },
+ }
+ ],
+ },
+ },
+ )
+
+ create_custom_job_def = CreateCustomJobOperator(
+ task_id="create_custom_job_def",
+ region=REGION,
+ project_id=PROJECT_ID,
+ custom_job={
+ "display_name": f"{DAG_ID}_{ENV_ID}_def",
+ "labels": {
+ "vertex_pipelines": "",
+ "airflow_dag_id": DAG_ID,
+ "airflow_task_id": DAG_ID,
+ },
+ "job_spec": {
+ "scheduling": {"disable_retries": True},
+ "worker_pool_specs": [
+ {
+ "machine_spec": {
+ "machine_type": MACHINE_TYPE,
+ "accelerator_type": ACCELERATOR_TYPE,
+ "accelerator_count": ACCELERATOR_COUNT,
+ },
+ "replica_count": REPLICA_COUNT,
+ "container_spec": {
+ "image_uri": IMAGE_URI,
+ "command": ["python3", "-c", test_python_code],
+ },
+ }
+ ],
+ },
+ },
+ deferrable=True,
+ )
+
+ # ### Everything below this line is not part of example ###
+ # ### Just for system tests purpose ###
+ from tests_common.test_utils.watcher import watcher
+
+ # This test needs watcher in order to properly mark success/failure
+ # when "tearDown" task with trigger rule is part of the DAG
+ list(dag.tasks) >> watcher()
+
+from tests_common.test_utils.system_tests import get_test_run # noqa: E402
+
+# Needed to run the example DAG with pytest (see:
contributing-docs/testing/system_tests.rst)
+test_run = get_test_run(dag)
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_custom_job.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_custom_job.py
index 38599467317..fb949a3ec73 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_custom_job.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_custom_job.py
@@ -206,6 +206,79 @@ class TestCustomJobWithDefaultProjectIdHook:
)
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
+ @pytest.mark.parametrize(
+ "job_state_value",
+ [
+ JobState.JOB_STATE_SUCCEEDED,
+ ],
+ )
+ @mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_custom_job"))
+ def test_wait_for_custom_job(
+ self,
+ mock_get_custom_job,
+ job_state_value,
+ test_custom_job_name,
+ ):
+ expected_obj = types.CustomJob(
+ state=job_state_value,
+ name=test_custom_job_name,
+ )
+ mock_get_custom_job.return_value = expected_obj
+ actual_obj = self.hook.wait_for_custom_job(
+ project_id=TEST_PROJECT_ID,
+ region=TEST_REGION,
+ custom_job_id=TEST_PIPELINE_JOB_ID,
+ )
+ assert actual_obj == expected_obj
+
+ @pytest.mark.parametrize(
+ ("job_state_value", "error_message"),
+ [
+ (
+ JobState.JOB_STATE_FAILED,
+ "Error message from VertexAI",
+ ),
+ (
+ JobState.JOB_STATE_CANCELLED,
+ "The CustomJob has been cancelled.",
+ ),
+ (
+ JobState.JOB_STATE_PAUSED,
+ "The CustomJob has been stopped, and can be resumed.",
+ ),
+ (
+ JobState.JOB_STATE_EXPIRED,
+ "The CustomJob has expired.",
+ ),
+ (
+ JobState.JOB_STATE_PARTIALLY_SUCCEEDED,
+ "Error message from VertexAI",
+ ),
+ ],
+ )
+ @mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_custom_job"))
+ def test_wait_for_custom_job_failed_states(
+ self,
+ mock_get_custom_job,
+ job_state_value,
+ error_message,
+ test_custom_job_name,
+ ):
+ expected_obj = types.CustomJob(
+ state=job_state_value,
+ name=test_custom_job_name,
+ error={
+ "message": "Error message from VertexAI",
+ },
+ )
+ mock_get_custom_job.return_value = expected_obj
+ with pytest.raises(RuntimeError, match=error_message):
+ self.hook.wait_for_custom_job(
+ project_id=TEST_PROJECT_ID,
+ region=TEST_REGION,
+ custom_job_id=TEST_PIPELINE_JOB_ID,
+ )
+
class TestCustomJobWithoutDefaultProjectIdHook:
def setup_method(self):
@@ -314,6 +387,79 @@ class TestCustomJobWithoutDefaultProjectIdHook:
)
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
+ @pytest.mark.parametrize(
+ "job_state_value",
+ [
+ JobState.JOB_STATE_SUCCEEDED,
+ ],
+ )
+ @mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_custom_job"))
+ def test_wait_for_custom_job(
+ self,
+ mock_get_custom_job,
+ job_state_value,
+ test_custom_job_name,
+ ):
+ expected_obj = types.CustomJob(
+ state=job_state_value,
+ name=test_custom_job_name,
+ )
+ mock_get_custom_job.return_value = expected_obj
+ actual_obj = self.hook.wait_for_custom_job(
+ project_id=TEST_PROJECT_ID,
+ region=TEST_REGION,
+ custom_job_id=TEST_PIPELINE_JOB_ID,
+ )
+ assert actual_obj == expected_obj
+
+ @pytest.mark.parametrize(
+ ("job_state_value", "error_message"),
+ [
+ (
+ JobState.JOB_STATE_FAILED,
+ "Error message from VertexAI",
+ ),
+ (
+ JobState.JOB_STATE_CANCELLED,
+ "The CustomJob has been cancelled.",
+ ),
+ (
+ JobState.JOB_STATE_PAUSED,
+ "The CustomJob has been stopped, and can be resumed.",
+ ),
+ (
+ JobState.JOB_STATE_EXPIRED,
+ "The CustomJob has expired.",
+ ),
+ (
+ JobState.JOB_STATE_PARTIALLY_SUCCEEDED,
+ "Error message from VertexAI",
+ ),
+ ],
+ )
+ @mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_custom_job"))
+ def test_wait_for_custom_job_failed_states(
+ self,
+ mock_get_custom_job,
+ job_state_value,
+ error_message,
+ test_custom_job_name,
+ ):
+ expected_obj = types.CustomJob(
+ state=job_state_value,
+ name=test_custom_job_name,
+ error={
+ "message": "Error message from VertexAI",
+ },
+ )
+ mock_get_custom_job.return_value = expected_obj
+ with pytest.raises(RuntimeError, match=error_message):
+ self.hook.wait_for_custom_job(
+ project_id=TEST_PROJECT_ID,
+ region=TEST_REGION,
+ custom_job_id=TEST_PIPELINE_JOB_ID,
+ )
+
class TestCustomJobAsyncHook:
@pytest.mark.asyncio
@@ -398,9 +544,6 @@ class TestCustomJobAsyncHook:
@pytest.mark.parametrize(
"job_state_value",
[
- JobState.JOB_STATE_CANCELLED,
- JobState.JOB_STATE_FAILED,
- JobState.JOB_STATE_PAUSED,
JobState.JOB_STATE_SUCCEEDED,
],
)
@@ -427,6 +570,59 @@ class TestCustomJobAsyncHook:
mock_get_job_service_client.assert_awaited_once_with(region=TEST_REGION)
assert actual_obj == expected_obj
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("job_state_value", "error_message"),
+ [
+ (
+ JobState.JOB_STATE_FAILED,
+ "Error message from VertexAI",
+ ),
+ (
+ JobState.JOB_STATE_CANCELLED,
+ "The CustomJob has been cancelled.",
+ ),
+ (
+ JobState.JOB_STATE_PAUSED,
+ "The CustomJob has been stopped, and can be resumed.",
+ ),
+ (
+ JobState.JOB_STATE_EXPIRED,
+ "The CustomJob has expired.",
+ ),
+ (
+ JobState.JOB_STATE_PARTIALLY_SUCCEEDED,
+ "Error message from VertexAI",
+ ),
+ ],
+ )
+ @mock.patch(CUSTOM_JOB_STRING.format("CustomJobAsyncHook.get_custom_job"))
+
@mock.patch(CUSTOM_JOB_STRING.format("CustomJobAsyncHook.get_job_service_client"))
+ async def test_wait_for_custom_job_failed_states(
+ self,
+ mock_get_job_service_client,
+ mock_get_custom_job,
+ job_state_value,
+ error_message,
+ test_async_hook,
+ test_custom_job_name,
+ ):
+ expected_obj = types.CustomJob(
+ state=job_state_value,
+ name=test_custom_job_name,
+ error={
+ "message": "Error message from VertexAI",
+ },
+ )
+ mock_get_custom_job.return_value = expected_obj
+ with pytest.raises(RuntimeError, match=error_message):
+ await test_async_hook.wait_for_custom_job(
+ project_id=TEST_PROJECT_ID,
+ location=TEST_REGION,
+ job_id=TEST_PIPELINE_JOB_ID,
+ )
+
mock_get_job_service_client.assert_awaited_once_with(region=TEST_REGION)
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"pipeline_state_value",
diff --git
a/providers/google/tests/unit/google/cloud/operators/test_vertex_ai.py
b/providers/google/tests/unit/google/cloud/operators/test_vertex_ai.py
index 762648f6192..7fee751c47d 100644
--- a/providers/google/tests/unit/google/cloud/operators/test_vertex_ai.py
+++ b/providers/google/tests/unit/google/cloud/operators/test_vertex_ai.py
@@ -45,6 +45,7 @@ from
airflow.providers.google.cloud.operators.vertex_ai.batch_prediction_job imp
)
from airflow.providers.google.cloud.operators.vertex_ai.custom_job import (
CreateCustomContainerTrainingJobOperator,
+ CreateCustomJobOperator,
CreateCustomPythonPackageTrainingJobOperator,
CreateCustomTrainingJobOperator,
DeleteCustomTrainingJobOperator,
@@ -233,6 +234,27 @@ TEST_RAY_VERSION: str = "2.33"
TEST_CLUSTER_NAME: str = "test-cluster-name"
TEST_CLUSTER_ID: str = "test-cluster-id"
+TEST_CUSTOM_JOB = {
+ "display_name": DISPLAY_NAME,
+ "job_spec": {
+ "scheduling": {"disable_retries": True},
+ "worker_pool_specs": [
+ {
+ "machine_spec": {
+ "machine_type": MACHINE_TYPE,
+ "accelerator_type": ACCELERATOR_TYPE,
+ "accelerator_count": ACCELERATOR_COUNT,
+ },
+ "replica_count": REPLICA_COUNT,
+ "container_spec": {
+ "image_uri": "test_image_uri",
+ "command": ["python3", "-c", "test_python_code"],
+ },
+ }
+ ],
+ },
+}
+
class TestVertexAICreateCustomContainerTrainingJobOperator:
@mock.patch(VERTEX_AI_PATH.format("custom_job.Dataset"))
@@ -3172,3 +3194,30 @@ class TestVertexAIDeleteRayClusterOperator:
project_id=GCP_PROJECT,
cluster_id=TEST_CLUSTER_ID,
)
+
+
+class TestVertexAICreateCustomJobOperator:
+ @mock.patch(VERTEX_AI_PATH.format("custom_job.CustomJob.to_dict"))
+ @mock.patch(VERTEX_AI_PATH.format("custom_job.CustomJobHook"))
+ def test_execute(self, mock_hook, to_dict_mock):
+ op = CreateCustomJobOperator(
+ task_id=TASK_ID,
+ gcp_conn_id=GCP_CONN_ID,
+ impersonation_chain=IMPERSONATION_CHAIN,
+ region=GCP_LOCATION,
+ project_id=GCP_PROJECT,
+ custom_job=TEST_CUSTOM_JOB,
+ retry=RETRY,
+ timeout=TIMEOUT,
+ metadata=METADATA,
+ )
+ op.execute(context={"ti": mock.MagicMock(), "task": mock.MagicMock()})
+ mock_hook.assert_called_once_with(gcp_conn_id=GCP_CONN_ID,
impersonation_chain=IMPERSONATION_CHAIN)
+ mock_hook.return_value.create_custom_job.assert_called_once_with(
+ region=GCP_LOCATION,
+ project_id=GCP_PROJECT,
+ custom_job=TEST_CUSTOM_JOB,
+ retry=RETRY,
+ timeout=TIMEOUT,
+ metadata=METADATA,
+ )
diff --git
a/providers/google/tests/unit/google/cloud/triggers/test_vertex_ai.py
b/providers/google/tests/unit/google/cloud/triggers/test_vertex_ai.py
index 1fb59d6ea6b..b34c9b875e4 100644
--- a/providers/google/tests/unit/google/cloud/triggers/test_vertex_ai.py
+++ b/providers/google/tests/unit/google/cloud/triggers/test_vertex_ai.py
@@ -42,6 +42,7 @@ from airflow.providers.google.cloud.triggers.vertex_ai import
(
CreateBatchPredictionJobTrigger,
CreateHyperparameterTuningJobTrigger,
CustomContainerTrainingJobTrigger,
+ CustomJobTrigger,
CustomPythonPackageTrainingJobTrigger,
CustomTrainingJobTrigger,
RunPipelineJobTrigger,
@@ -113,6 +114,18 @@ def custom_training_job_trigger():
)
[email protected]
+def custom_job_trigger():
+ return CustomJobTrigger(
+ gcp_conn_id=TEST_CONN_ID,
+ project_id=TEST_PROJECT_ID,
+ location=TEST_LOCATION,
+ custom_job_id=TEST_HPT_JOB_ID,
+ poll_interval=TEST_POLL_INTERVAL,
+ impersonation_chain=TEST_IMPERSONATION_CHAIN,
+ )
+
+
@pytest_asyncio.fixture
async def custom_job_async_hook():
return CustomJobAsyncHook(
@@ -998,3 +1011,20 @@ class TestCustomPythonPackageTrainingJobTrigger:
pipeline_id=custom_python_package_training_job_trigger.job_id,
poll_interval=custom_python_package_training_job_trigger.poll_interval,
)
+
+
+class TestCustomJobTrigger:
+ def test_serialize(self, custom_job_trigger):
+ actual_data = custom_job_trigger.serialize()
+ expected_data = (
+
"airflow.providers.google.cloud.triggers.vertex_ai.CustomJobTrigger",
+ {
+ "project_id": TEST_PROJECT_ID,
+ "location": TEST_LOCATION,
+ "custom_job_id": TEST_HPT_JOB_ID,
+ "gcp_conn_id": TEST_CONN_ID,
+ "impersonation_chain": TEST_IMPERSONATION_CHAIN,
+ "poll_interval": TEST_POLL_INTERVAL,
+ },
+ )
+ assert actual_data == expected_data