This is an automated email from the ASF dual-hosted git repository.
kaxil 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 092444019b5 Remove Google hook tests duplicated across default-project
test classes (#74362)
092444019b5 is described below
commit 092444019b59407287a0f42767c526a0c429b887
Author: Kaxil Naik <[email protected]>
AuthorDate: Tue Oct 6 18:39:07 2026 +0100
Remove Google hook tests duplicated across default-project test classes
(#74362)
Each removed test in a *WithoutDefaultProjectIdHook class is identical to
a test in its *WithDefaultProjectIdHook sibling and passes an explicit
project_id, so fallback_to_default_project_id never reads the hook default
and both copies run the same code path. Tests that omit project_id or
expect a raise still cover the no-default path and are kept.
---
.../google/cloud/hooks/test_dataproc_metastore.py | 268 -----------------
.../unit/google/cloud/hooks/test_managed_kafka.py | 331 ---------------------
.../google/cloud/hooks/vertex_ai/test_auto_ml.py | 70 -----
.../hooks/vertex_ai/test_batch_prediction_job.py | 65 ----
.../cloud/hooks/vertex_ai/test_custom_job.py | 125 --------
.../google/cloud/hooks/vertex_ai/test_dataset.py | 231 --------------
.../cloud/hooks/vertex_ai/test_endpoint_service.py | 172 -----------
.../hooks/vertex_ai/test_experiment_service.py | 87 ------
.../vertex_ai/test_hyperparameter_tuning_job.py | 74 -----
.../cloud/hooks/vertex_ai/test_model_service.py | 205 -------------
.../cloud/hooks/vertex_ai/test_pipeline_job.py | 82 -----
.../hooks/vertex_ai/test_prediction_service.py | 29 --
.../unit/google/cloud/hooks/vertex_ai/test_ray.py | 122 --------
13 files changed, 1861 deletions(-)
diff --git
a/providers/google/tests/unit/google/cloud/hooks/test_dataproc_metastore.py
b/providers/google/tests/unit/google/cloud/hooks/test_dataproc_metastore.py
index 4bc89792aa5..cf9145a9c63 100644
--- a/providers/google/tests/unit/google/cloud/hooks/test_dataproc_metastore.py
+++ b/providers/google/tests/unit/google/cloud/hooks/test_dataproc_metastore.py
@@ -26,7 +26,6 @@ from airflow.providers.google.cloud.hooks.dataproc_metastore
import DataprocMeta
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -358,270 +357,3 @@ class TestDataprocMetastoreWithDefaultProjectIdHook:
query=TEST_PARTITIONS_QUERY_ALL.format(TEST_TABLE_ID),
),
)
-
-
-class TestDataprocMetastoreWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = DataprocMetastoreHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_create_backup(self, mock_client) -> None:
- self.hook.create_backup(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- backup=TEST_BACKUP,
- backup_id=TEST_BACKUP_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.create_backup.assert_called_once_with(
- request=dict(
- parent=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- backup=TEST_BACKUP,
- backup_id=TEST_BACKUP_ID,
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_create_metadata_import(self, mock_client) -> None:
- self.hook.create_metadata_import(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- metadata_import=TEST_METADATA_IMPORT,
- metadata_import_id=TEST_METADATA_IMPORT_ID,
- )
- mock_client.assert_called_once()
-
mock_client.return_value.create_metadata_import.assert_called_once_with(
- request=dict(
- parent=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- metadata_import=TEST_METADATA_IMPORT,
- metadata_import_id=TEST_METADATA_IMPORT_ID,
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_create_service(self, mock_client) -> None:
- self.hook.create_service(
- region=TEST_REGION,
- project_id=TEST_PROJECT_ID,
- service=TEST_SERVICE,
- service_id=TEST_SERVICE_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.create_service.assert_called_once_with(
- request=dict(
- parent=TEST_PARENT.format(TEST_PROJECT_ID, TEST_REGION),
- service_id=TEST_SERVICE_ID,
- service=TEST_SERVICE,
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_delete_backup(self, mock_client) -> None:
- self.hook.delete_backup(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- backup_id=TEST_BACKUP_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.delete_backup.assert_called_once_with(
- request=dict(
- name=TEST_NAME_BACKUPS.format(TEST_PROJECT_ID, TEST_REGION,
TEST_SERVICE_ID, TEST_BACKUP_ID),
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_delete_service(self, mock_client) -> None:
- self.hook.delete_service(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.delete_service.assert_called_once_with(
- request=dict(
- name=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID, TEST_REGION,
TEST_SERVICE_ID),
- request_id=None,
- ),
- retry=DEFAULT,
- timeout=None,
- metadata=(),
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_export_metadata(self, mock_client) -> None:
- self.hook.export_metadata(
- destination_gcs_folder=TEST_DESTINATION_GCS_FOLDER,
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.export_metadata.assert_called_once_with(
- request=dict(
- destination_gcs_folder=TEST_DESTINATION_GCS_FOLDER,
- service=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- request_id=None,
- database_dump_type=None,
- ),
- retry=DEFAULT,
- timeout=None,
- metadata=(),
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_get_service(self, mock_client) -> None:
- self.hook.get_service(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.get_service.assert_called_once_with(
- request=dict(
- name=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID, TEST_REGION,
TEST_SERVICE_ID),
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_list_backups(self, mock_client) -> None:
- self.hook.list_backups(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.list_backups.assert_called_once_with(
- request=dict(
- parent=TEST_PARENT_BACKUPS.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- page_size=None,
- page_token=None,
- filter=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_restore_service(self, mock_client) -> None:
- self.hook.restore_service(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- backup_project_id=TEST_PROJECT_ID,
- backup_region=TEST_REGION,
- backup_service_id=TEST_SERVICE_ID,
- backup_id=TEST_BACKUP_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.restore_service.assert_called_once_with(
- request=dict(
- service=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- backup=TEST_NAME_BACKUPS.format(
- TEST_PROJECT_ID, TEST_REGION, TEST_SERVICE_ID,
TEST_BACKUP_ID
- ),
- restore_type=None,
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client"))
- def test_update_service(self, mock_client) -> None:
- self.hook.update_service(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- service_id=TEST_SERVICE_ID,
- service=TEST_SERVICE_TO_UPDATE,
- update_mask=TEST_UPDATE_MASK,
- )
- mock_client.assert_called_once()
- mock_client.return_value.update_service.assert_called_once_with(
- request=dict(
- service=TEST_SERVICE_TO_UPDATE,
- update_mask=TEST_UPDATE_MASK,
- request_id=None,
- ),
- retry=DEFAULT,
- timeout=None,
- metadata=(),
- )
-
- @pytest.mark.parametrize(
- ("partitions_input", "partitions"),
- [
- ([TEST_PARTITION_NAME], f"'{TEST_PARTITION_NAME}'"),
- ([TEST_SUBPARTITION_NAME], f"'{TEST_SUBPARTITION_NAME}'"),
- (
- [TEST_PARTITION_NAME, TEST_SUBPARTITION_NAME],
- f"'{TEST_PARTITION_NAME}', '{TEST_SUBPARTITION_NAME}'",
- ),
- ([TEST_PARTITION_NAME, TEST_PARTITION_NAME],
f"'{TEST_PARTITION_NAME}'"),
- ],
- )
- @mock.patch(
-
DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client_v1beta")
- )
- def test_list_hive_partitions(self, mock_client, partitions_input,
partitions) -> None:
- self.hook.list_hive_partitions(
- project_id=TEST_PROJECT_ID,
- service_id=TEST_SERVICE_ID,
- region=TEST_REGION,
- table=TEST_TABLE_ID,
- partition_names=partitions_input,
- )
- mock_client.assert_called_once()
- mock_client.return_value.query_metadata.assert_called_once_with(
- request=dict(
- service=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- query=TEST_PARTITIONS_QUERY.format(TEST_TABLE_ID, partitions),
- ),
- )
-
- @pytest.mark.parametrize("partitions", [[], None])
- @mock.patch(
-
DATAPROC_METASTORE_STRING.format("DataprocMetastoreHook.get_dataproc_metastore_client_v1beta")
- )
- def test_list_hive_partitions_empty_list(self, mock_client, partitions) ->
None:
- self.hook.list_hive_partitions(
- project_id=TEST_PROJECT_ID,
- service_id=TEST_SERVICE_ID,
- region=TEST_REGION,
- table=TEST_TABLE_ID,
- partition_names=partitions,
- )
- mock_client.assert_called_once()
- mock_client.return_value.query_metadata.assert_called_once_with(
- request=dict(
- service=TEST_PARENT_SERVICES.format(TEST_PROJECT_ID,
TEST_REGION, TEST_SERVICE_ID),
- query=TEST_PARTITIONS_QUERY_ALL.format(TEST_TABLE_ID),
- ),
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/test_managed_kafka.py
b/providers/google/tests/unit/google/cloud/hooks/test_managed_kafka.py
index 4aeff93ef62..9ba7e6fad41 100644
--- a/providers/google/tests/unit/google/cloud/hooks/test_managed_kafka.py
+++ b/providers/google/tests/unit/google/cloud/hooks/test_managed_kafka.py
@@ -25,7 +25,6 @@ from airflow.providers.google.cloud.hooks.managed_kafka
import ManagedKafkaHook
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -395,333 +394,3 @@ class TestManagedKafkaWithDefaultProjectIdHook:
mock_client.return_value.cluster_path.assert_called_once_with(
TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
)
-
-
-class TestManagedKafkaWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = ManagedKafkaHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_create_cluster(self, mock_client) -> None:
- self.hook.create_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster=TEST_CLUSTER,
- cluster_id=TEST_CLUSTER_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.create_cluster.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- cluster=TEST_CLUSTER,
- cluster_id=TEST_CLUSTER_ID,
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_LOCATION)
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_delete_cluster(self, mock_client) -> None:
- self.hook.delete_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.delete_cluster.assert_called_once_with(
-
request=dict(name=mock_client.return_value.cluster_path.return_value,
request_id=None),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.cluster_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_get_cluster(self, mock_client) -> None:
- self.hook.get_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.get_cluster.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.cluster_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.cluster_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_update_cluster(self, mock_client) -> None:
- self.hook.update_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster=TEST_UPDATED_CLUSTER,
- cluster_id=TEST_CLUSTER_ID,
- update_mask=TEST_CLUSTER_UPDATE_MASK,
- )
- mock_client.assert_called_once()
- mock_client.return_value.update_cluster.assert_called_once_with(
- request=dict(
- update_mask=TEST_CLUSTER_UPDATE_MASK,
- cluster={
- "name": mock_client.return_value.cluster_path.return_value,
- **TEST_UPDATED_CLUSTER,
- },
- request_id=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.cluster_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_list_clusters(self, mock_client) -> None:
- self.hook.list_clusters(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- )
- mock_client.assert_called_once()
- mock_client.return_value.list_clusters.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- page_size=None,
- page_token=None,
- filter=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_LOCATION)
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_create_topic(self, mock_client) -> None:
- self.hook.create_topic(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- topic_id=TEST_TOPIC_ID,
- topic=TEST_TOPIC,
- )
- mock_client.assert_called_once()
- mock_client.return_value.create_topic.assert_called_once_with(
- request=dict(
- parent=mock_client.return_value.cluster_path.return_value,
- topic_id=TEST_TOPIC_ID,
- topic=TEST_TOPIC,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.cluster_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_delete_topic(self, mock_client) -> None:
- self.hook.delete_topic(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- topic_id=TEST_TOPIC_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.delete_topic.assert_called_once_with(
-
request=dict(name=mock_client.return_value.topic_path.return_value),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.topic_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_LOCATION,
- TEST_CLUSTER_ID,
- TEST_TOPIC_ID,
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_get_topic(self, mock_client) -> None:
- self.hook.get_topic(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- topic_id=TEST_TOPIC_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.get_topic.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.topic_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.topic_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_LOCATION,
- TEST_CLUSTER_ID,
- TEST_TOPIC_ID,
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_update_topic(self, mock_client) -> None:
- self.hook.update_topic(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- topic_id=TEST_TOPIC_ID,
- topic=TEST_UPDATED_TOPIC,
- update_mask=TEST_TOPIC_UPDATE_MASK,
- )
- mock_client.assert_called_once()
- mock_client.return_value.update_topic.assert_called_once_with(
- request=dict(
- update_mask=TEST_TOPIC_UPDATE_MASK,
- topic={
- "name": mock_client.return_value.topic_path.return_value,
- **TEST_UPDATED_TOPIC,
- },
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.topic_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID, TEST_TOPIC_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_list_topics(self, mock_client) -> None:
- self.hook.list_topics(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.list_topics.assert_called_once_with(
- request=dict(
- parent=mock_client.return_value.cluster_path.return_value,
- page_size=None,
- page_token=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.cluster_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_delete_consumer_group(self, mock_client) -> None:
- self.hook.delete_consumer_group(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- consumer_group_id=TEST_CONSUMER_GROUP_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.delete_consumer_group.assert_called_once_with(
-
request=dict(name=mock_client.return_value.consumer_group_path.return_value),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.consumer_group_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_LOCATION,
- TEST_CLUSTER_ID,
- TEST_CONSUMER_GROUP_ID,
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_get_consumer_group(self, mock_client) -> None:
- self.hook.get_consumer_group(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- consumer_group_id=TEST_CONSUMER_GROUP_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.get_consumer_group.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.consumer_group_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.consumer_group_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_LOCATION,
- TEST_CLUSTER_ID,
- TEST_CONSUMER_GROUP_ID,
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_update_consumer_group(self, mock_client) -> None:
- self.hook.update_consumer_group(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- consumer_group_id=TEST_CONSUMER_GROUP_ID,
- consumer_group={},
- update_mask={},
- )
- mock_client.assert_called_once()
- mock_client.return_value.update_consumer_group.assert_called_once_with(
- request=dict(
- update_mask={},
- consumer_group={
- "name":
mock_client.return_value.consumer_group_path.return_value,
- **{},
- },
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.consumer_group_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID,
TEST_CONSUMER_GROUP_ID
- )
-
-
@mock.patch(MANAGED_KAFKA_STRING.format("ManagedKafkaHook.get_managed_kafka_client"))
- def test_list_consumer_groups(self, mock_client) -> None:
- self.hook.list_consumer_groups(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_ID,
- )
- mock_client.assert_called_once()
- mock_client.return_value.list_consumer_groups.assert_called_once_with(
- request=dict(
- parent=mock_client.return_value.cluster_path.return_value,
- page_size=None,
- page_token=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.cluster_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_LOCATION, TEST_CLUSTER_ID
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_auto_ml.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_auto_ml.py
index 8759e85a341..be47d1f9561 100644
--- a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_auto_ml.py
+++ b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_auto_ml.py
@@ -30,7 +30,6 @@ from airflow.providers.google.cloud.hooks.vertex_ai.auto_ml
import AutoMLHook
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -112,72 +111,3 @@ class TestAutoMLWithDefaultProjectIdHook:
timeout=None,
)
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
-class TestAutoMLWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = AutoMLHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(CUSTOM_JOB_STRING.format("AutoMLHook.get_pipeline_service_client"))
- def test_delete_training_pipeline(self, mock_client) -> None:
- self.hook.delete_training_pipeline(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- training_pipeline=TEST_TRAINING_PIPELINE_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.delete_training_pipeline.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.training_pipeline_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.training_pipeline_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_TRAINING_PIPELINE_NAME
- )
-
-
@mock.patch(CUSTOM_JOB_STRING.format("AutoMLHook.get_pipeline_service_client"))
- def test_get_training_pipeline(self, mock_client) -> None:
- self.hook.get_training_pipeline(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- training_pipeline=TEST_TRAINING_PIPELINE_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_training_pipeline.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.training_pipeline_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.training_pipeline_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_TRAINING_PIPELINE_NAME
- )
-
-
@mock.patch(CUSTOM_JOB_STRING.format("AutoMLHook.get_pipeline_service_client"))
- def test_list_training_pipelines(self, mock_client) -> None:
- self.hook.list_training_pipelines(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.list_training_pipelines.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- page_size=None,
- page_token=None,
- filter=None,
- read_mask=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_batch_prediction_job.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_batch_prediction_job.py
index 650ee52cd74..cef0f0730c3 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_batch_prediction_job.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_batch_prediction_job.py
@@ -195,71 +195,6 @@ class TestBatchPredictionJobWithoutDefaultProjectIdHook:
mock_create.assert_called_once_with(**expected_params)
assert actual_job == expected_job
-
@mock.patch(BATCH_PREDICTION_JOB_STRING.format("BatchPredictionJobHook.get_job_service_client"))
- def test_delete_batch_prediction_job(self, mock_client) -> None:
- self.hook.delete_batch_prediction_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- batch_prediction_job=TEST_BATCH_PREDICTION_JOB,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.delete_batch_prediction_job.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.batch_prediction_job_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.batch_prediction_job_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_BATCH_PREDICTION_JOB,
- )
-
-
@mock.patch(BATCH_PREDICTION_JOB_STRING.format("BatchPredictionJobHook.get_job_service_client"))
- def test_get_batch_prediction_job(self, mock_client) -> None:
- self.hook.get_batch_prediction_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- batch_prediction_job=TEST_BATCH_PREDICTION_JOB,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.get_batch_prediction_job.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.batch_prediction_job_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.batch_prediction_job_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_BATCH_PREDICTION_JOB,
- )
-
-
@mock.patch(BATCH_PREDICTION_JOB_STRING.format("BatchPredictionJobHook.get_job_service_client"))
- def test_list_batch_prediction_jobs(self, mock_client) -> None:
- self.hook.list_batch_prediction_jobs(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.list_batch_prediction_jobs.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
class TestBatchPredictionJobAsyncHook:
def setup_method(self):
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 fb949a3ec73..a797e2a0654 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
@@ -287,131 +287,6 @@ class TestCustomJobWithoutDefaultProjectIdHook:
):
self.hook = CustomJobHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
@mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_pipeline_service_client"))
- def test_cancel_training_pipeline(self, mock_client) -> None:
- self.hook.cancel_training_pipeline(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- training_pipeline=TEST_TRAINING_PIPELINE_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.cancel_training_pipeline.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.training_pipeline_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.training_pipeline_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_TRAINING_PIPELINE_NAME
- )
-
-
@mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_pipeline_service_client"))
- def test_create_training_pipeline(self, mock_client) -> None:
- self.hook.create_training_pipeline(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- training_pipeline=TEST_TRAINING_PIPELINE,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.create_training_pipeline.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- training_pipeline=TEST_TRAINING_PIPELINE,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_pipeline_service_client"))
- def test_delete_training_pipeline(self, mock_client) -> None:
- self.hook.delete_training_pipeline(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- training_pipeline=TEST_TRAINING_PIPELINE_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.delete_training_pipeline.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.training_pipeline_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.training_pipeline_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_TRAINING_PIPELINE_NAME
- )
-
-
@mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_pipeline_service_client"))
- def test_get_training_pipeline(self, mock_client) -> None:
- self.hook.get_training_pipeline(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- training_pipeline=TEST_TRAINING_PIPELINE_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_training_pipeline.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.training_pipeline_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.training_pipeline_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_TRAINING_PIPELINE_NAME
- )
-
-
@mock.patch(CUSTOM_JOB_STRING.format("CustomJobHook.get_pipeline_service_client"))
- def test_list_training_pipelines(self, mock_client) -> None:
- self.hook.list_training_pipelines(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.list_training_pipelines.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- page_size=None,
- page_token=None,
- filter=None,
- read_mask=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
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"),
[
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_dataset.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_dataset.py
index 63fbc6d3bd5..bd65bffc754 100644
--- a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_dataset.py
+++ b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_dataset.py
@@ -30,7 +30,6 @@ from airflow.providers.google.cloud.hooks.vertex_ai.dataset
import DatasetHook
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -280,233 +279,3 @@ class TestVertexAIWithDefaultProjectIdHook:
mock_client.return_value.dataset_path.assert_called_once_with(
TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
)
-
-
-class TestVertexAIWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = DatasetHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_create_dataset(self, mock_client) -> None:
- self.hook.create_dataset(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.create_dataset.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- dataset=TEST_DATASET,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_delete_dataset(self, mock_client) -> None:
- self.hook.delete_dataset(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.delete_dataset.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.dataset_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.dataset_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_export_data(self, mock_client) -> None:
- self.hook.export_data(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- export_config=TEST_EXPORT_CONFIG,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.export_data.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.dataset_path.return_value,
- export_config=TEST_EXPORT_CONFIG,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.dataset_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_get_annotation_spec(self, mock_client) -> None:
- self.hook.get_annotation_spec(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- annotation_spec=TEST_ANNOTATION_SPEC,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_annotation_spec.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.annotation_spec_path.return_value,
- read_mask=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.annotation_spec_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID, TEST_ANNOTATION_SPEC
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_get_dataset(self, mock_client) -> None:
- self.hook.get_dataset(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_dataset.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.dataset_path.return_value,
- read_mask=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.dataset_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_import_data(self, mock_client) -> None:
- self.hook.import_data(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- import_configs=TEST_IMPORT_CONFIGS,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.import_data.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.dataset_path.return_value,
- import_configs=TEST_IMPORT_CONFIGS,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.dataset_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_list_annotations(self, mock_client) -> None:
- self.hook.list_annotations(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- data_item=TEST_DATA_ITEM,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_annotations.assert_called_once_with(
- request=dict(
- parent=mock_client.return_value.data_item_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.data_item_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID, TEST_DATA_ITEM
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_list_data_items(self, mock_client) -> None:
- self.hook.list_data_items(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset=TEST_DATASET_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_data_items.assert_called_once_with(
- request=dict(
- parent=mock_client.return_value.dataset_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.dataset_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
- )
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_list_datasets(self, mock_client) -> None:
- self.hook.list_datasets(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_datasets.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(DATASET_STRING.format("DatasetHook.get_dataset_service_client"))
- def test_update_dataset(self, mock_client) -> None:
- self.hook.update_dataset(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- dataset_id=TEST_DATASET_ID,
- dataset=TEST_DATASET,
- update_mask=TEST_UPDATE_MASK,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.update_dataset.assert_called_once_with(
- request=dict(
- dataset=TEST_DATASET,
- update_mask=TEST_UPDATE_MASK,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.dataset_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_DATASET_ID
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_endpoint_service.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_endpoint_service.py
index 6feac8aba17..e988f951777 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_endpoint_service.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_endpoint_service.py
@@ -30,7 +30,6 @@ from
airflow.providers.google.cloud.hooks.vertex_ai.endpoint_service import Endp
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -217,174 +216,3 @@ class TestEndpointServiceWithDefaultProjectIdHook:
retry=DEFAULT,
timeout=None,
)
-
-
-class TestEndpointServiceWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = EndpointServiceHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_create_endpoint(self, mock_client) -> None:
- self.hook.create_endpoint(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- endpoint=TEST_ENDPOINT,
- endpoint_id=TEST_ENDPOINT_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.create_endpoint.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- endpoint=TEST_ENDPOINT,
- endpoint_id=TEST_ENDPOINT_ID,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.common_location_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- )
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_delete_endpoint(self, mock_client) -> None:
- self.hook.delete_endpoint(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- endpoint=TEST_ENDPOINT_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.delete_endpoint.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.endpoint_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.endpoint_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_ENDPOINT_NAME,
- )
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_deploy_model(self, mock_client) -> None:
- self.hook.deploy_model(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- endpoint=TEST_ENDPOINT_NAME,
- deployed_model=TEST_DEPLOYED_MODEL,
- traffic_split=TEST_TRAFFIC_SPLIT,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.deploy_model.assert_called_once_with(
- request=dict(
- endpoint=mock_client.return_value.endpoint_path.return_value,
- deployed_model=TEST_DEPLOYED_MODEL,
- traffic_split=TEST_TRAFFIC_SPLIT,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.endpoint_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_ENDPOINT_NAME,
- )
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_get_endpoint(self, mock_client) -> None:
- self.hook.get_endpoint(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- endpoint=TEST_ENDPOINT_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_endpoint.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.endpoint_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.endpoint_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_ENDPOINT_NAME,
- )
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_list_endpoints(self, mock_client) -> None:
- self.hook.list_endpoints(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_endpoints.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.common_location_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- )
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_undeploy_model(self, mock_client) -> None:
- self.hook.undeploy_model(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- endpoint=TEST_ENDPOINT_NAME,
- deployed_model_id=TEST_DEPLOYED_MODEL_ID,
- traffic_split=TEST_TRAFFIC_SPLIT,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.undeploy_model.assert_called_once_with(
- request=dict(
- endpoint=mock_client.return_value.endpoint_path.return_value,
- deployed_model_id=TEST_DEPLOYED_MODEL_ID,
- traffic_split=TEST_TRAFFIC_SPLIT,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.endpoint_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_ENDPOINT_NAME
- )
-
-
@mock.patch(ENDPOINT_SERVICE_STRING.format("EndpointServiceHook.get_endpoint_service_client"))
- def test_update_endpoint(self, mock_client) -> None:
- self.hook.update_endpoint(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- endpoint_id=TEST_ENDPOINT_NAME,
- endpoint=TEST_ENDPOINT,
- update_mask=TEST_UPDATE_MASK,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.update_endpoint.assert_called_once_with(
- request=dict(
- endpoint=TEST_ENDPOINT,
- update_mask=TEST_UPDATE_MASK,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_experiment_service.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_experiment_service.py
index f7bf7dd2f0d..8a17fff9736 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_experiment_service.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_experiment_service.py
@@ -28,7 +28,6 @@ from
airflow.providers.google.cloud.hooks.vertex_ai.experiment_service import (
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -87,48 +86,6 @@ class TestExperimentWithDefaultProjectIdHook:
)
-class TestExperimentWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = ExperimentHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
- @mock.patch(EXPERIMENT_SERVICE_STRING.format("aiplatform.init"))
- def test_create_experiment(self, mock_init) -> None:
- self.hook.create_experiment(
- project_id=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment_name=TEST_EXPERIMENT_NAME,
- experiment_description=TEST_EXPERIMENT_DESCRIPTION,
- experiment_tensorboard=TEST_TENSORBOARD,
- )
- mock_init.assert_called_with(
- project=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment=TEST_EXPERIMENT_NAME,
- experiment_description=TEST_EXPERIMENT_DESCRIPTION,
- experiment_tensorboard=TEST_TENSORBOARD,
- )
-
- @mock.patch(EXPERIMENT_SERVICE_STRING.format("aiplatform.Experiment"))
- def test_delete_experiment(self, mock_experiment) -> None:
- self.hook.delete_experiment(
- project_id=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment_name=TEST_EXPERIMENT_NAME,
-
delete_backing_tensorboard_runs=TEST_DELETE_BACKING_TENSORBOARD_RUNS,
- )
- mock_experiment.assert_called_with(
- project=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment_name=TEST_EXPERIMENT_NAME,
- )
- mock_experiment.return_value.delete.assert_called_with(
-
delete_backing_tensorboard_runs=TEST_DELETE_BACKING_TENSORBOARD_RUNS
- )
-
-
class TestExperimentRunWithDefaultProjectIdHook:
def setup_method(self):
with mock.patch(
@@ -171,47 +128,3 @@ class TestExperimentRunWithDefaultProjectIdHook:
mock_experiment_run.return_value.delete.assert_called_with(
delete_backing_tensorboard_run=TEST_DELETE_BACKING_TENSORBOARD_RUNS
)
-
-
-class TestExperimentRunWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = ExperimentRunHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
- @mock.patch(EXPERIMENT_SERVICE_STRING.format("aiplatform.ExperimentRun"))
- def test_create_experiment_run(self, mock_experiment_run) -> None:
- self.hook.create_experiment_run(
- project_id=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment_name=TEST_EXPERIMENT_NAME,
- experiment_run_name=TEST_EXPERIMENT_RUN_NAME,
- experiment_run_tensorboard=TEST_TENSORBOARD,
- )
- mock_experiment_run.create.assert_called_with(
- project=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment=TEST_EXPERIMENT_NAME,
- run_name=TEST_EXPERIMENT_RUN_NAME,
- state=aiplatform.gapic.Execution.State.NEW,
- tensorboard=TEST_TENSORBOARD,
- )
-
- @mock.patch(EXPERIMENT_SERVICE_STRING.format("aiplatform.ExperimentRun"))
- def test_delete_experiment_run(self, mock_experiment_run) -> None:
- self.hook.delete_experiment_run(
- project_id=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment_name=TEST_EXPERIMENT_NAME,
- experiment_run_name=TEST_EXPERIMENT_RUN_NAME,
- )
- mock_experiment_run.assert_called_with(
- project=TEST_PROJECT_ID,
- location=TEST_REGION,
- experiment=TEST_EXPERIMENT_NAME,
- run_name=TEST_EXPERIMENT_RUN_NAME,
- )
- mock_experiment_run.return_value.delete.assert_called_with(
- delete_backing_tensorboard_run=TEST_DELETE_BACKING_TENSORBOARD_RUNS
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_hyperparameter_tuning_job.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_hyperparameter_tuning_job.py
index 21c4f0177b0..708c631908a 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_hyperparameter_tuning_job.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_hyperparameter_tuning_job.py
@@ -40,7 +40,6 @@ from
airflow.providers.google.cloud.hooks.vertex_ai.hyperparameter_tuning_job im
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -203,79 +202,6 @@ class TestHyperparameterTuningJobWithDefaultProjectIdHook:
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-class TestHyperparameterTuningJobWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook =
HyperparameterTuningJobHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(HYPERPARAMETER_TUNING_JOB_HOOK_STRING.format("get_job_service_client"))
- def test_delete_hyperparameter_tuning_job(self, mock_client) -> None:
- self.hook.delete_hyperparameter_tuning_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- hyperparameter_tuning_job=TEST_HYPERPARAMETER_TUNING_JOB_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.delete_hyperparameter_tuning_job.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.hyperparameter_tuning_job_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.hyperparameter_tuning_job_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_HYPERPARAMETER_TUNING_JOB_ID,
- )
-
-
@mock.patch(HYPERPARAMETER_TUNING_JOB_HOOK_STRING.format("get_job_service_client"))
- def test_get_hyperparameter_tuning_job(self, mock_client) -> None:
- self.hook.get_hyperparameter_tuning_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- hyperparameter_tuning_job=TEST_HYPERPARAMETER_TUNING_JOB_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.get_hyperparameter_tuning_job.assert_called_once_with(
- request=dict(
-
name=mock_client.return_value.hyperparameter_tuning_job_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.hyperparameter_tuning_job_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_HYPERPARAMETER_TUNING_JOB_ID,
- )
-
-
@mock.patch(HYPERPARAMETER_TUNING_JOB_HOOK_STRING.format("get_job_service_client"))
- def test_list_hyperparameter_tuning_jobs(self, mock_client) -> None:
- self.hook.list_hyperparameter_tuning_jobs(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
-
mock_client.return_value.list_hyperparameter_tuning_jobs.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
class TestHyperparameterTuningJobHook:
def setup_method(self):
with mock.patch(
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_model_service.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_model_service.py
index d9ba9c31f31..47a57a39776 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_model_service.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_model_service.py
@@ -30,7 +30,6 @@ from
airflow.providers.google.cloud.hooks.vertex_ai.model_service import ModelSe
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -248,207 +247,3 @@ class TestModelServiceWithDefaultProjectIdHook:
retry=DEFAULT,
timeout=None,
)
-
-
-class TestModelServiceWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = ModelServiceHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_delete_model(self, mock_client) -> None:
- self.hook.delete_model(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- model=TEST_MODEL_NAME,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.delete_model.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.model_path.assert_called_once_with(
- TEST_PROJECT_ID,
- TEST_REGION,
- TEST_MODEL_NAME,
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_export_model(self, mock_client) -> None:
- self.hook.export_model(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- model=TEST_MODEL_NAME,
- output_config=TEST_OUTPUT_CONFIG,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.export_model.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- output_config=TEST_OUTPUT_CONFIG,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.model_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_MODEL_NAME
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_list_models(self, mock_client) -> None:
- self.hook.list_models(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_models.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- filter=None,
- page_size=None,
- page_token=None,
- read_mask=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_upload_model(self, mock_client) -> None:
- self.hook.upload_model(project_id=TEST_PROJECT_ID, region=TEST_REGION,
model=TEST_MODEL)
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.upload_model.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- model=TEST_MODEL,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_upload_model_with_parent_model(self, mock_client) -> None:
- self.hook.upload_model(
- project_id=TEST_PROJECT_ID, region=TEST_REGION, model=TEST_MODEL,
parent_model=TEST_PARENT_MODEL
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.upload_model.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- model=TEST_MODEL,
- parent_model=mock_client.return_value.model_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_list_model_versions(self, mock_client) -> None:
- self.hook.list_model_versions(
- project_id=TEST_PROJECT_ID, region=TEST_REGION,
model_id=TEST_MODEL_NAME
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_model_versions.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_delete_model_version(self, mock_client) -> None:
- self.hook.delete_model_version(
- project_id=TEST_PROJECT_ID, region=TEST_REGION,
model_id=TEST_MODEL_NAME
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.delete_model_version.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_get_model(self, mock_client) -> None:
- self.hook.get_model(project_id=TEST_PROJECT_ID, region=TEST_REGION,
model_id=TEST_MODEL_NAME)
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_model.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_set_version_as_default(self, mock_client) -> None:
- self.hook.set_version_as_default(
- project_id=TEST_PROJECT_ID, region=TEST_REGION,
model_id=TEST_MODEL_NAME
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.merge_version_aliases.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- version_aliases=["default"],
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_add_version_aliases(self, mock_client) -> None:
- self.hook.add_version_aliases(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- model_id=TEST_MODEL_NAME,
- version_aliases=TEST_VERSION_ALIASES,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.merge_version_aliases.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- version_aliases=TEST_VERSION_ALIASES,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
-
@mock.patch(MODEL_SERVICE_STRING.format("ModelServiceHook.get_model_service_client"))
- def test_delete_version_aliases(self, mock_client) -> None:
- self.hook.delete_version_aliases(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- model_id=TEST_MODEL_NAME,
- version_aliases=TEST_VERSION_ALIASES,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.merge_version_aliases.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.model_path.return_value,
- version_aliases=["-" + alias for alias in
TEST_VERSION_ALIASES],
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_pipeline_job.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_pipeline_job.py
index f5a9ec016b6..70d4fb83022 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_pipeline_job.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_pipeline_job.py
@@ -164,88 +164,6 @@ class TestPipelineJobWithoutDefaultProjectIdHook:
):
self.hook = PipelineJobHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
@mock.patch(PIPELINE_JOB_STRING.format("PipelineJobHook.get_pipeline_service_client"))
- def test_create_pipeline_job(self, mock_client) -> None:
- self.hook.create_pipeline_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- pipeline_job=TEST_PIPELINE_JOB,
- pipeline_job_id=TEST_PIPELINE_JOB_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.create_pipeline_job.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- pipeline_job=TEST_PIPELINE_JOB,
- pipeline_job_id=TEST_PIPELINE_JOB_ID,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
-
@mock.patch(PIPELINE_JOB_STRING.format("PipelineJobHook.get_pipeline_service_client"))
- def test_delete_pipeline_job(self, mock_client) -> None:
- self.hook.delete_pipeline_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- pipeline_job_id=TEST_PIPELINE_JOB_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.delete_pipeline_job.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.pipeline_job_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.pipeline_job_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_PIPELINE_JOB_ID
- )
-
-
@mock.patch(PIPELINE_JOB_STRING.format("PipelineJobHook.get_pipeline_service_client"))
- def test_get_pipeline_job(self, mock_client) -> None:
- self.hook.get_pipeline_job(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- pipeline_job_id=TEST_PIPELINE_JOB_ID,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.get_pipeline_job.assert_called_once_with(
- request=dict(
- name=mock_client.return_value.pipeline_job_path.return_value,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
- mock_client.return_value.pipeline_job_path.assert_called_once_with(
- TEST_PROJECT_ID, TEST_REGION, TEST_PIPELINE_JOB_ID
- )
-
-
@mock.patch(PIPELINE_JOB_STRING.format("PipelineJobHook.get_pipeline_service_client"))
- def test_list_pipeline_jobs(self, mock_client) -> None:
- self.hook.list_pipeline_jobs(
- project_id=TEST_PROJECT_ID,
- region=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.list_pipeline_jobs.assert_called_once_with(
- request=dict(
-
parent=mock_client.return_value.common_location_path.return_value,
- page_size=None,
- page_token=None,
- filter=None,
- order_by=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
-
mock_client.return_value.common_location_path.assert_called_once_with(TEST_PROJECT_ID,
TEST_REGION)
-
@pytest.mark.parametrize(
"reserved_ip_ranges", [None, [], ["range-1", "range-2"]], ids=["none",
"empty", "multiple"]
)
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_prediction_service.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_prediction_service.py
index 26b24668944..9a6c347620c 100644
---
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_prediction_service.py
+++
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_prediction_service.py
@@ -31,7 +31,6 @@ from
airflow.providers.google.cloud.hooks.vertex_ai.prediction_service import (
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -70,31 +69,3 @@ class TestPredictionServiceWithDefaultProjectIdHook:
retry=DEFAULT,
timeout=None,
)
-
-
-class TestPredictionServiceWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = PredictionServiceHook(gcp_conn_id=TEST_GCP_CONN_ID)
-
-
@mock.patch(PREDICTION_SERVICE_STRING.format("PredictionServiceHook.get_prediction_service_client"))
- def test_predict(self, mock_client):
- self.hook.predict(
- endpoint_id=TEST_ENDPOINT_ID,
- instances=["instance1", "instance2"],
- project_id=TEST_PROJECT_ID,
- location=TEST_REGION,
- )
- mock_client.assert_called_once_with(TEST_REGION)
- mock_client.return_value.predict.assert_called_once_with(
- request=dict(
-
endpoint=f"projects/{TEST_PROJECT_ID}/locations/{TEST_REGION}/endpoints/{TEST_ENDPOINT_ID}",
- instances=["instance1", "instance2"],
- parameters=None,
- ),
- metadata=(),
- retry=DEFAULT,
- timeout=None,
- )
diff --git
a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py
b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py
index 26205939a85..dac8819fb23 100644
--- a/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py
+++ b/providers/google/tests/unit/google/cloud/hooks/vertex_ai/test_ray.py
@@ -28,7 +28,6 @@ from airflow.providers.google.cloud.hooks.vertex_ai.ray
import RayHook
from unit.google.cloud.utils.base_gcp_mock import (
mock_base_gcp_hook_default_project_id,
- mock_base_gcp_hook_no_default_project_id,
)
TEST_GCP_CONN_ID: str = "test-gcp-conn-id"
@@ -234,124 +233,3 @@ class TestRayWithDefaultProjectIdHook:
)
assert self.hook.serialize_cluster_obj(cluster_obj) ==
SAMPLE_CLUSTER_SERIALIZED
-
-
-class TestRayWithoutDefaultProjectIdHook:
- def setup_method(self):
- with mock.patch(
- BASE_STRING.format("GoogleBaseHook.__init__"),
new=mock_base_gcp_hook_no_default_project_id
- ):
- self.hook = RayHook(gcp_conn_id=TEST_GCP_CONN_ID)
- self.hook.get_credentials = mock.MagicMock()
-
- @mock.patch(RAY_STRING.format("vertex_ray.create_ray_cluster"))
- @mock.patch(RAY_STRING.format("aiplatform.init"))
- def test_create_ray_cluster(self, mock_aiplatform_init,
mock_create_ray_cluster) -> None:
- self.hook.create_ray_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- head_node_type=TEST_NODE_RESOURCES,
- python_version=TEST_PYTHON_VERSION,
- ray_version=TEST_RAY_VERSION,
- network=None,
- service_account=None,
- cluster_name=TEST_CLUSTER_NAME,
- worker_node_types=[TEST_NODE_RESOURCES],
- custom_images=None,
- enable_metrics_collection=True,
- enable_logging=True,
- psc_interface_config=None,
- reserved_ip_ranges=None,
- labels=None,
- )
- mock_aiplatform_init.assert_called_once()
- mock_create_ray_cluster.assert_called_once_with(
- head_node_type=TEST_NODE_RESOURCES,
- python_version=TEST_PYTHON_VERSION,
- ray_version=TEST_RAY_VERSION,
- network=None,
- service_account=None,
- cluster_name=TEST_CLUSTER_NAME,
- worker_node_types=[TEST_NODE_RESOURCES],
- custom_images=None,
- enable_metrics_collection=True,
- enable_logging=True,
- psc_interface_config=None,
- reserved_ip_ranges=None,
- labels=None,
- )
-
- @mock.patch(RAY_STRING.format("vertex_ray.delete_ray_cluster"))
- @mock.patch(RAY_STRING.format("aiplatform.init"))
-
@mock.patch(RAY_STRING.format("PersistentResourceServiceClient.persistent_resource_path"))
- def test_delete_ray_cluster(
- self, mock_persistent_resource_path, mock_aiplatform_init,
mock_delete_ray_cluster
- ) -> None:
- self.hook.delete_ray_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_NAME,
- )
- mock_aiplatform_init.assert_called_once()
- mock_persistent_resource_path.assert_called_once_with(
- project=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- persistent_resource=TEST_CLUSTER_NAME,
- )
- mock_delete_ray_cluster.assert_called_once_with(
- cluster_resource_name=mock_persistent_resource_path.return_value,
- )
-
- @mock.patch(RAY_STRING.format("vertex_ray.get_ray_cluster"))
- @mock.patch(RAY_STRING.format("aiplatform.init"))
-
@mock.patch(RAY_STRING.format("PersistentResourceServiceClient.persistent_resource_path"))
- def test_get_ray_cluster(
- self, mock_persistent_resource_path, mock_aiplatform_init,
mock_get_ray_cluster
- ) -> None:
- self.hook.get_ray_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_NAME,
- )
- mock_aiplatform_init.assert_called_once()
- mock_persistent_resource_path.assert_called_once_with(
- project=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- persistent_resource=TEST_CLUSTER_NAME,
- )
- mock_get_ray_cluster.assert_called_once_with(
- cluster_resource_name=mock_persistent_resource_path.return_value,
- )
-
- @mock.patch(RAY_STRING.format("vertex_ray.update_ray_cluster"))
- @mock.patch(RAY_STRING.format("aiplatform.init"))
-
@mock.patch(RAY_STRING.format("PersistentResourceServiceClient.persistent_resource_path"))
- def test_update_ray_cluster(
- self, mock_persistent_resource_path, mock_aiplatform_init,
mock_update_ray_cluster
- ) -> None:
- self.hook.update_ray_cluster(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- cluster_id=TEST_CLUSTER_NAME,
- worker_node_types=[TEST_NODE_RESOURCES],
- )
- mock_aiplatform_init.assert_called_once()
- mock_persistent_resource_path.assert_called_once_with(
- project=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- persistent_resource=TEST_CLUSTER_NAME,
- )
- mock_update_ray_cluster.assert_called_once_with(
- cluster_resource_name=mock_persistent_resource_path.return_value,
- worker_node_types=[TEST_NODE_RESOURCES],
- )
-
- @mock.patch(RAY_STRING.format("vertex_ray.list_ray_clusters"))
- @mock.patch(RAY_STRING.format("aiplatform.init"))
- def test_list_ray_clusters(self, mock_aiplatform_init,
mock_list_ray_clusters) -> None:
- self.hook.list_ray_clusters(
- project_id=TEST_PROJECT_ID,
- location=TEST_LOCATION,
- )
- mock_aiplatform_init.assert_called_once()
- mock_list_ray_clusters.assert_called_once()