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 ecf2842c30f Change google sheets operator and connected system tests
(#66929)
ecf2842c30f is described below
commit ecf2842c30f870d234f7526c3e613058d23ad0cd
Author: Nitochkin <[email protected]>
AuthorDate: Thu Aug 27 12:54:37 2026 +0200
Change google sheets operator and connected system tests (#66929)
* Change google sheets operator and connected system tests
* WIP: Fix documentation issues
* Refactor tests to not use render_template_as_native_obj
* Cast XComArg to expected types to satisfy mypy
---------
Co-authored-by: Anton Nitochkin <[email protected]>
Co-authored-by: Marcin Lubimow <[email protected]>
---
docs/spelling_wordlist.txt | 1 +
.../airflow/providers/google/suite/hooks/drive.py | 26 ++++++++++++
.../providers/google/suite/operators/sheets.py | 44 +++++++++++++++++--
.../google/cloud/gcs/example_gcs_to_sheets.py | 31 ++++++++++----
.../system/google/cloud/gcs/example_sheets.py | 46 +++++++++++++-------
.../google/cloud/gcs/example_sheets_to_gcs.py | 23 +++++++---
.../cloud/sql_to_sheets/example_sql_to_sheets.py | 23 +++++++---
.../tests/unit/google/suite/hooks/test_drive.py | 19 +++++++++
.../unit/google/suite/operators/test_sheets.py | 49 ++++++++++++++++++----
9 files changed, 219 insertions(+), 43 deletions(-)
diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt
index a153988c461..b0cd8ce92ea 100644
--- a/docs/spelling_wordlist.txt
+++ b/docs/spelling_wordlist.txt
@@ -1930,6 +1930,7 @@ webpage
Webserver
webserver
webservers
+webViewLink
Werkzeug
werkzeug
whitespace
diff --git a/providers/google/src/airflow/providers/google/suite/hooks/drive.py
b/providers/google/src/airflow/providers/google/suite/hooks/drive.py
index 7d1257cf4ed..7139234281a 100644
--- a/providers/google/src/airflow/providers/google/suite/hooks/drive.py
+++ b/providers/google/src/airflow/providers/google/suite/hooks/drive.py
@@ -322,3 +322,29 @@ class GoogleDriveHook(GoogleBaseHook):
"""
request = self.get_media_request(file_id=file_id)
self.download_content_from_request(file_handle=file_handle,
request=request, chunk_size=chunk_size)
+
+ def create_file(
+ self,
+ file_metadata: dict[str, Any],
+ fields: str = "id, webViewLink",
+ supports_all_drives: bool = True,
+ **kwargs: Any,
+ ) -> dict[str, Any]:
+ """
+ Create a file on Google Drive.
+
+ :param file_metadata: Metadata of the file that will be created, e.g.
name, mime type, etc.
+ :param fields: Selector specifying which fields to include in a
partial response.
+ Default is "id, webViewLink".
+ :param supports_all_drives: Whether the requesting application
supports both My Drive and shared drives.
+ Default is True.
+ """
+ service = self.get_conn()
+
+ response = (
+ service.files()
+ .create(body=file_metadata, fields=fields,
supportsAllDrives=supports_all_drives, **kwargs)
+ .execute(num_retries=self.num_retries)
+ )
+
+ return response
diff --git
a/providers/google/src/airflow/providers/google/suite/operators/sheets.py
b/providers/google/src/airflow/providers/google/suite/operators/sheets.py
index d19a05a7557..856ed2f7471 100644
--- a/providers/google/src/airflow/providers/google/suite/operators/sheets.py
+++ b/providers/google/src/airflow/providers/google/suite/operators/sheets.py
@@ -19,6 +19,7 @@ from __future__ import annotations
from collections.abc import Sequence
from typing import Any
+from airflow.providers.google.suite.hooks.drive import GoogleDriveHook
from airflow.providers.google.suite.hooks.sheets import GSheetsHook
from airflow.providers.google.version_compat import BaseOperator
@@ -34,6 +35,11 @@ class GoogleSheetsCreateSpreadsheetOperator(BaseOperator):
:param spreadsheet: an instance of Spreadsheet
https://developers.google.com/sheets/api/reference/rest/v4/spreadsheets#Spreadsheet
:param gcp_conn_id: The connection ID to use when fetching connection info.
+ :param drive_id: Shared Drive ID where the spreadsheet should be created.
+ This is useful when using a service account, since service accounts
+ do not have personal Drive storage.
+ When drive_id is set, only spreadsheet["properties"]["title"] is used.
+ All other properties are ignored. (templated)
:param impersonation_chain: Optional service account to impersonate using
short-term
credentials, or chained list of accounts required to get the
access_token
of the last account in the list, which will be impersonated in the
request.
@@ -49,6 +55,7 @@ class GoogleSheetsCreateSpreadsheetOperator(BaseOperator):
template_fields: Sequence[str] = (
"spreadsheet",
"impersonation_chain",
+ "drive_id",
)
def __init__(
@@ -57,6 +64,7 @@ class GoogleSheetsCreateSpreadsheetOperator(BaseOperator):
spreadsheet: dict[str, Any],
gcp_conn_id: str = "google_cloud_default",
impersonation_chain: str | Sequence[str] | None = None,
+ drive_id: str | None = None,
api_endpoint: str | None = None,
**kwargs,
) -> None:
@@ -65,14 +73,42 @@ class GoogleSheetsCreateSpreadsheetOperator(BaseOperator):
self.spreadsheet = spreadsheet
self.impersonation_chain = impersonation_chain
self.api_endpoint = api_endpoint
+ self.drive_id = drive_id
def execute(self, context: Any) -> dict[str, Any]:
+ if self.drive_id:
+ spreadsheet = self._create_spreadsheet_via_drive_api()
+ else:
+ spreadsheet = self._create_spreadsheet_via_sheets_api()
+ context["task_instance"].xcom_push(key="spreadsheet_id",
value=spreadsheet["spreadsheetId"])
+ context["task_instance"].xcom_push(key="spreadsheet_url",
value=spreadsheet["spreadsheetUrl"])
+ return spreadsheet
+
+ def _construct_spreadsheet_metadata(self, spreadsheet) -> dict:
+ return {
+ "name": spreadsheet["properties"]["title"],
+ "mimeType": "application/vnd.google-apps.spreadsheet",
+ "parents": [self.drive_id],
+ }
+
+ def _create_spreadsheet_via_drive_api(self) -> dict:
+ spreadsheet_metadata =
self._construct_spreadsheet_metadata(self.spreadsheet)
+ hook = GoogleDriveHook(
+ gcp_conn_id=self.gcp_conn_id,
+ impersonation_chain=self.impersonation_chain,
+ )
+ spreadsheet = hook.create_file(file_metadata=spreadsheet_metadata)
+
+ response = {
+ "spreadsheetId": spreadsheet.get("id"),
+ "spreadsheetUrl": spreadsheet.get("webViewLink"),
+ }
+ return response
+
+ def _create_spreadsheet_via_sheets_api(self) -> dict:
hook = GSheetsHook(
gcp_conn_id=self.gcp_conn_id,
impersonation_chain=self.impersonation_chain,
api_endpoint=self.api_endpoint,
)
- spreadsheet = hook.create_spreadsheet(spreadsheet=self.spreadsheet)
- context["task_instance"].xcom_push(key="spreadsheet_id",
value=spreadsheet["spreadsheetId"])
- context["task_instance"].xcom_push(key="spreadsheet_url",
value=spreadsheet["spreadsheetUrl"])
- return spreadsheet
+ return hook.create_spreadsheet(spreadsheet=self.spreadsheet)
diff --git
a/providers/google/tests/system/google/cloud/gcs/example_gcs_to_sheets.py
b/providers/google/tests/system/google/cloud/gcs/example_gcs_to_sheets.py
index 39abcd3c6ad..dea3026f811 100644
--- a/providers/google/tests/system/google/cloud/gcs/example_gcs_to_sheets.py
+++ b/providers/google/tests/system/google/cloud/gcs/example_gcs_to_sheets.py
@@ -21,7 +21,7 @@ import json
import logging
import os
from datetime import datetime
-from typing import Any
+from typing import Any, cast
from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS
@@ -31,8 +31,10 @@ else:
# Airflow 2 path
from airflow.decorators import task # type: ignore[attr-defined,no-redef]
from airflow.models.dag import DAG
+from airflow.models.xcom_arg import XComArg
from airflow.providers.google.cloud.operators.gcs import
GCSCreateBucketOperator, GCSDeleteBucketOperator
from airflow.providers.google.cloud.transfers.sheets_to_gcs import
GoogleSheetsToGCSOperator
+from airflow.providers.google.common.utils.get_secret import get_secret
from airflow.providers.google.suite.operators.sheets import
GoogleSheetsCreateSpreadsheetOperator
from airflow.providers.google.suite.transfers.gcs_to_sheets import
GCSToGoogleSheetsOperator
@@ -56,6 +58,7 @@ SPREADSHEET = {
"sheets": [{"properties": {"title": "Sheet1"}}],
}
CONNECTION_ID = f"connection_{DAG_ID}_{ENV_ID}"
+GDRIVE_SECRET_ID = "gdrive_shared_folder_id"
log = logging.getLogger(__name__)
@@ -66,6 +69,13 @@ with DAG(
catchup=False,
tags=["example", "gcs"],
) as dag:
+
+ @task
+ def get_shared_drive_id() -> str:
+ return get_secret(secret_id=GDRIVE_SECRET_ID).strip()
+
+ get_shared_drive_id_task = get_shared_drive_id()
+
create_bucket = GCSCreateBucketOperator(
task_id="create_bucket", bucket_name=BUCKET_NAME, project_id=PROJECT_ID
)
@@ -73,7 +83,7 @@ with DAG(
@task
def create_connection(connection_id: str):
conn_extra = {
- "scope":
"https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
+ "scope":
"https://www.googleapis.com/auth/drive,https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
"project": PROJECT_ID,
"keyfile_dict": "", # Override to match your needs
}
@@ -88,22 +98,29 @@ with DAG(
create_connection_task = create_connection(connection_id=CONNECTION_ID)
create_spreadsheet = GoogleSheetsCreateSpreadsheetOperator(
- task_id="create_spreadsheet", spreadsheet=SPREADSHEET,
gcp_conn_id=CONNECTION_ID
+ task_id="create_spreadsheet",
+ spreadsheet=SPREADSHEET,
+ gcp_conn_id=CONNECTION_ID,
+ drive_id=get_shared_drive_id_task,
)
upload_sheet_to_gcs = GoogleSheetsToGCSOperator(
task_id="upload_sheet_to_gcs",
destination_bucket=BUCKET_NAME,
- spreadsheet_id="{{
task_instance.xcom_pull(task_ids='create_spreadsheet', key='spreadsheet_id')
}}",
+ spreadsheet_id=cast("str", XComArg(create_spreadsheet,
key="spreadsheet_id")),
gcp_conn_id=CONNECTION_ID,
)
+ @task
+ def get_first_item(items: list[Any]) -> Any:
+ return items[0]
+
# [START upload_gcs_to_sheets]
upload_gcs_to_sheet = GCSToGoogleSheetsOperator(
task_id="upload_gcs_to_sheet",
bucket_name=BUCKET_NAME,
- object_name="{{ task_instance.xcom_pull('upload_sheet_to_gcs')[0] }}",
- spreadsheet_id="{{
task_instance.xcom_pull(task_ids='create_spreadsheet', key='spreadsheet_id')
}}",
+ object_name=get_first_item(upload_sheet_to_gcs.output),
+ spreadsheet_id=cast("str", XComArg(create_spreadsheet,
key="spreadsheet_id")),
gcp_conn_id=CONNECTION_ID,
)
# [END upload_gcs_to_sheets]
@@ -120,7 +137,7 @@ with DAG(
(
# TEST SETUP
- [create_bucket, create_connection_task]
+ [get_shared_drive_id_task, create_bucket, create_connection_task]
>> create_spreadsheet
>> upload_sheet_to_gcs
# TEST BODY
diff --git a/providers/google/tests/system/google/cloud/gcs/example_sheets.py
b/providers/google/tests/system/google/cloud/gcs/example_sheets.py
index 66158f8595a..08923284672 100644
--- a/providers/google/tests/system/google/cloud/gcs/example_sheets.py
+++ b/providers/google/tests/system/google/cloud/gcs/example_sheets.py
@@ -21,7 +21,7 @@ import json
import logging
import os
from datetime import datetime
-from typing import Any
+from typing import Any, cast
from airflow.models.dag import DAG
@@ -35,6 +35,7 @@ else:
from airflow.models.xcom_arg import XComArg
from airflow.providers.google.cloud.operators.gcs import
GCSCreateBucketOperator, GCSDeleteBucketOperator
from airflow.providers.google.cloud.transfers.sheets_to_gcs import
GoogleSheetsToGCSOperator
+from airflow.providers.google.common.utils.get_secret import get_secret
from airflow.providers.google.suite.operators.sheets import
GoogleSheetsCreateSpreadsheetOperator
from airflow.providers.google.suite.transfers.gcs_to_sheets import
GCSToGoogleSheetsOperator
from airflow.providers.standard.operators.bash import BashOperator
@@ -59,6 +60,7 @@ SPREADSHEET = {
"sheets": [{"properties": {"title": "Sheet1"}}],
}
CONNECTION_ID = f"connection_{DAG_ID}_{ENV_ID}"
+GDRIVE_SECRET_ID = "gdrive_shared_folder_id"
log = logging.getLogger(__name__)
@@ -69,6 +71,13 @@ with DAG(
catchup=False,
tags=["example", "sheets"],
) as dag:
+
+ @task
+ def get_shared_drive_id() -> str:
+ return get_secret(secret_id=GDRIVE_SECRET_ID).strip()
+
+ get_shared_drive_id_task = get_shared_drive_id()
+
create_bucket = GCSCreateBucketOperator(
task_id="create_bucket", bucket_name=BUCKET_NAME, project_id=PROJECT_ID
)
@@ -76,7 +85,7 @@ with DAG(
@task
def create_connection(connection_id: str):
conn_extra = {
- "scope":
"https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
+ "scope":
"https://www.googleapis.com/auth/drive,https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
"project": PROJECT_ID,
"keyfile_dict": "", # Override to match your needs
}
@@ -90,18 +99,12 @@ with DAG(
create_connection_task = create_connection(connection_id=CONNECTION_ID)
- # [START upload_sheet_to_gcs]
- upload_sheet_to_gcs = GoogleSheetsToGCSOperator(
- task_id="upload_sheet_to_gcs",
- destination_bucket=BUCKET_NAME,
- spreadsheet_id="{{
task_instance.xcom_pull(task_ids='create_spreadsheet', key='spreadsheet_id')
}}",
- gcp_conn_id=CONNECTION_ID,
- )
- # [END upload_sheet_to_gcs]
-
# [START create_spreadsheet]
create_spreadsheet = GoogleSheetsCreateSpreadsheetOperator(
- task_id="create_spreadsheet", spreadsheet=SPREADSHEET,
gcp_conn_id=CONNECTION_ID
+ task_id="create_spreadsheet",
+ spreadsheet=SPREADSHEET,
+ gcp_conn_id=CONNECTION_ID,
+ drive_id=get_shared_drive_id_task,
)
# [END create_spreadsheet]
@@ -112,12 +115,25 @@ with DAG(
)
# [END print_spreadsheet_url]
+ # [START upload_sheet_to_gcs]
+ upload_sheet_to_gcs = GoogleSheetsToGCSOperator(
+ task_id="upload_sheet_to_gcs",
+ destination_bucket=BUCKET_NAME,
+ spreadsheet_id=cast("str", XComArg(create_spreadsheet,
key="spreadsheet_id")),
+ gcp_conn_id=CONNECTION_ID,
+ )
+ # [END upload_sheet_to_gcs]
+
+ @task
+ def get_first_item(items: list[Any]) -> Any:
+ return items[0]
+
# [START upload_gcs_to_sheet]
upload_gcs_to_sheet = GCSToGoogleSheetsOperator(
task_id="upload_gcs_to_sheet",
bucket_name=BUCKET_NAME,
- object_name="{{ task_instance.xcom_pull('upload_sheet_to_gcs')[0] }}",
- spreadsheet_id="{{
task_instance.xcom_pull(task_ids='create_spreadsheet', key='spreadsheet_id')
}}",
+ object_name=get_first_item(upload_sheet_to_gcs.output),
+ spreadsheet_id=cast("str", XComArg(create_spreadsheet,
key="spreadsheet_id")),
gcp_conn_id=CONNECTION_ID,
)
# [END upload_gcs_to_sheet]
@@ -134,7 +150,7 @@ with DAG(
(
# TEST SETUP
- [create_bucket, create_connection_task]
+ [get_shared_drive_id_task, create_bucket, create_connection_task]
# TEST BODY
>> create_spreadsheet
>> print_spreadsheet_url
diff --git
a/providers/google/tests/system/google/cloud/gcs/example_sheets_to_gcs.py
b/providers/google/tests/system/google/cloud/gcs/example_sheets_to_gcs.py
index ad70627c08b..b23a9d988e8 100644
--- a/providers/google/tests/system/google/cloud/gcs/example_sheets_to_gcs.py
+++ b/providers/google/tests/system/google/cloud/gcs/example_sheets_to_gcs.py
@@ -21,9 +21,10 @@ import json
import logging
import os
from datetime import datetime
-from typing import Any
+from typing import Any, cast
from airflow.models.dag import DAG
+from airflow.models.xcom_arg import XComArg
from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS
@@ -34,6 +35,7 @@ else:
from airflow.decorators import task # type: ignore[attr-defined,no-redef]
from airflow.providers.google.cloud.operators.gcs import
GCSCreateBucketOperator, GCSDeleteBucketOperator
from airflow.providers.google.cloud.transfers.sheets_to_gcs import
GoogleSheetsToGCSOperator
+from airflow.providers.google.common.utils.get_secret import get_secret
from airflow.providers.google.suite.operators.sheets import
GoogleSheetsCreateSpreadsheetOperator
try:
@@ -56,6 +58,7 @@ SPREADSHEET = {
"sheets": [{"properties": {"title": "Sheet1"}}],
}
CONNECTION_ID = f"connection_{DAG_ID}_{ENV_ID}"
+GDRIVE_SECRET_ID = "gdrive_shared_folder_id"
log = logging.getLogger(__name__)
@@ -66,6 +69,13 @@ with DAG(
catchup=False,
tags=["example", "sheets"],
) as dag:
+
+ @task
+ def get_shared_drive_id() -> str:
+ return get_secret(secret_id=GDRIVE_SECRET_ID).strip()
+
+ get_shared_drive_id_task = get_shared_drive_id()
+
create_bucket = GCSCreateBucketOperator(
task_id="create_bucket", bucket_name=BUCKET_NAME, project_id=PROJECT_ID
)
@@ -73,7 +83,7 @@ with DAG(
@task
def create_connection(connection_id: str):
conn_extra = {
- "scope":
"https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
+ "scope":
"https://www.googleapis.com/auth/drive,https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
"project": PROJECT_ID,
"keyfile_dict": "", # Override to match your needs
}
@@ -88,14 +98,17 @@ with DAG(
create_connection_task = create_connection(connection_id=CONNECTION_ID)
create_spreadsheet = GoogleSheetsCreateSpreadsheetOperator(
- task_id="create_spreadsheet", spreadsheet=SPREADSHEET,
gcp_conn_id=CONNECTION_ID
+ task_id="create_spreadsheet",
+ spreadsheet=SPREADSHEET,
+ gcp_conn_id=CONNECTION_ID,
+ drive_id=get_shared_drive_id_task,
)
# [START upload_sheet_to_gcs]
upload_sheet_to_gcs = GoogleSheetsToGCSOperator(
task_id="upload_sheet_to_gcs",
destination_bucket=BUCKET_NAME,
- spreadsheet_id="{{
task_instance.xcom_pull(task_ids='create_spreadsheet', key='spreadsheet_id')
}}",
+ spreadsheet_id=cast("str", XComArg(create_spreadsheet,
key="spreadsheet_id")),
gcp_conn_id=CONNECTION_ID,
)
# [END upload_sheet_to_gcs]
@@ -112,7 +125,7 @@ with DAG(
(
# TEST SETUP
- [create_bucket, create_connection_task]
+ [get_shared_drive_id_task, create_bucket, create_connection_task]
>> create_spreadsheet
# TEST BODY
>> upload_sheet_to_gcs
diff --git
a/providers/google/tests/system/google/cloud/sql_to_sheets/example_sql_to_sheets.py
b/providers/google/tests/system/google/cloud/sql_to_sheets/example_sql_to_sheets.py
index 7cc510a04a0..00a9c5263fd 100644
---
a/providers/google/tests/system/google/cloud/sql_to_sheets/example_sql_to_sheets.py
+++
b/providers/google/tests/system/google/cloud/sql_to_sheets/example_sql_to_sheets.py
@@ -31,7 +31,7 @@ import json
import logging
import os
from datetime import datetime
-from typing import Any
+from typing import Any, cast
from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS
@@ -41,6 +41,7 @@ else:
# Airflow 2 path
from airflow.decorators import task # type: ignore[attr-defined,no-redef]
from airflow.models.dag import DAG
+from airflow.models.xcom_arg import XComArg
from airflow.providers.common.sql.operators.sql import SQLExecuteQueryOperator
from airflow.providers.google.cloud.hooks.compute import ComputeEngineHook
from airflow.providers.google.cloud.hooks.compute_ssh import
ComputeEngineSSHHook
@@ -48,6 +49,7 @@ from airflow.providers.google.cloud.operators.compute import (
ComputeEngineDeleteInstanceOperator,
ComputeEngineInsertInstanceOperator,
)
+from airflow.providers.google.common.utils.get_secret import get_secret
from airflow.providers.google.suite.operators.sheets import
GoogleSheetsCreateSpreadsheetOperator
from airflow.providers.google.suite.transfers.sql_to_sheets import
SQLToGoogleSheetsOperator
from airflow.providers.ssh.operators.ssh import SSHOperator
@@ -153,6 +155,7 @@ SPREADSHEET = {
"properties": {"title": "Test1"},
"sheets": [{"properties": {"title": "Sheet1"}}],
}
+GDRIVE_SECRET_ID = "gdrive_shared_folder_id"
log = logging.getLogger(__name__)
@@ -163,6 +166,13 @@ with DAG(
catchup=False,
tags=["example", "postgres", "gcs"],
) as dag:
+
+ @task
+ def get_shared_drive_id() -> str:
+ return get_secret(secret_id=GDRIVE_SECRET_ID).strip()
+
+ get_shared_drive_id_task = get_shared_drive_id()
+
create_gce_instance = ComputeEngineInsertInstanceOperator(
task_id="create_gce_instance",
project_id=PROJECT_ID,
@@ -233,7 +243,7 @@ with DAG(
@task
def setup_sheets_connection():
conn_extra = {
- "scope":
"https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
+ "scope":
"https://www.googleapis.com/auth/drive,https://www.googleapis.com/auth/spreadsheets,https://www.googleapis.com/auth/cloud-platform",
"project": PROJECT_ID,
"keyfile_dict": "", # Override to match your needs
}
@@ -248,7 +258,10 @@ with DAG(
setup_sheets_connection_task = setup_sheets_connection()
create_spreadsheet = GoogleSheetsCreateSpreadsheetOperator(
- task_id="create_spreadsheet", spreadsheet=SPREADSHEET,
gcp_conn_id=SHEETS_CONNECTION_ID
+ task_id="create_spreadsheet",
+ spreadsheet=SPREADSHEET,
+ gcp_conn_id=SHEETS_CONNECTION_ID,
+ drive_id=get_shared_drive_id_task,
)
# [START upload_sql_to_sheets]
@@ -257,7 +270,7 @@ with DAG(
sql=SQL_SELECT,
sql_conn_id=CONNECTION_ID,
database=DB_NAME,
- spreadsheet_id="{{
task_instance.xcom_pull(task_ids='create_spreadsheet', key='spreadsheet_id')
}}",
+ spreadsheet_id=cast("str", XComArg(create_spreadsheet,
key="spreadsheet_id")),
gcp_conn_id=SHEETS_CONNECTION_ID,
)
# [END upload_sql_to_sheets]
@@ -287,7 +300,7 @@ with DAG(
create_gce_instance >> get_public_ip_task >> create_connection_task
[create_gce_instance, create_firewall_rule] >> setup_postgres
[setup_postgres, create_connection_task, create_firewall_rule] >>
create_sql_table >> insert_sql_data
- setup_sheets_connection_task >> create_spreadsheet
+ get_shared_drive_id_task >> setup_sheets_connection_task >>
create_spreadsheet
(
[create_spreadsheet, insert_sql_data]
diff --git a/providers/google/tests/unit/google/suite/hooks/test_drive.py
b/providers/google/tests/unit/google/suite/hooks/test_drive.py
index 2eb282e964b..35eb6102081 100644
--- a/providers/google/tests/unit/google/suite/hooks/test_drive.py
+++ b/providers/google/tests/unit/google/suite/hooks/test_drive.py
@@ -428,3 +428,22 @@ class TestGoogleDriveHook:
]
)
assert return_value == file_id
+
+
@mock.patch("airflow.providers.google.suite.hooks.drive.GoogleDriveHook.get_conn")
+ def test_create_file(self, mock_get_conn):
+ file_metadata = {
+ "name": "Test File",
+ "mimeType": "application/vnd.google-apps.spreadsheet",
+ "parents": ["shared_drive_123"],
+ }
+
+ mock_execute =
mock_get_conn.return_value.files.return_value.create.return_value.execute
+ mock_execute.return_value = {"id": "NEW_FILE_ID", "webViewLink":
"https://example.com/view"}
+
+ result = self.gdrive_hook.create_file(file_metadata=file_metadata)
+
+
mock_get_conn.return_value.files.return_value.create.assert_called_once_with(
+ body=file_metadata, fields="id, webViewLink",
supportsAllDrives=True
+ )
+ mock_execute.assert_called_once()
+ assert result == {"id": "NEW_FILE_ID", "webViewLink":
"https://example.com/view"}
diff --git a/providers/google/tests/unit/google/suite/operators/test_sheets.py
b/providers/google/tests/unit/google/suite/operators/test_sheets.py
index 1d0fc216c03..5d6e1fd9884 100644
--- a/providers/google/tests/unit/google/suite/operators/test_sheets.py
+++ b/providers/google/tests/unit/google/suite/operators/test_sheets.py
@@ -23,29 +23,64 @@ from airflow.providers.google.suite.operators.sheets import
GoogleSheetsCreateSp
GCP_CONN_ID = "test"
SPREADSHEET_URL = "https://example/sheets"
SPREADSHEET_ID = "1234567890"
+DRIVE_ID = "shared_drive_123"
+SPREADSHEET_DATA = {"properties": {"title": "My Test Spreadsheet"}}
class TestGoogleSheetsCreateSpreadsheet:
@mock.patch("airflow.providers.google.suite.operators.sheets.GSheetsHook")
- def test_execute(self, mock_hook):
+ def test_execute_via_sheets_api(self, mock_sheets_hook):
+ """Test spreadsheet creation using the standard Sheets API, no
drive_id provided."""
mock_task_instance = mock.MagicMock()
context = {"task_instance": mock_task_instance}
- spreadsheet = mock.MagicMock()
- mock_hook.return_value.create_spreadsheet.return_value = {
+
+ mock_sheets_hook.return_value.create_spreadsheet.return_value = {
"spreadsheetId": SPREADSHEET_ID,
"spreadsheetUrl": SPREADSHEET_URL,
}
+
op = GoogleSheetsCreateSpreadsheetOperator(
- task_id="test_task", spreadsheet=spreadsheet,
gcp_conn_id=GCP_CONN_ID
+ task_id="test_task", spreadsheet=SPREADSHEET_DATA,
gcp_conn_id=GCP_CONN_ID
)
op_execute_result = op.execute(context)
-
mock_hook.return_value.create_spreadsheet.assert_called_once_with(spreadsheet=spreadsheet)
+
mock_sheets_hook.return_value.create_spreadsheet.assert_called_once_with(spreadsheet=SPREADSHEET_DATA)
# Verify xcom_push was called with correct arguments
assert mock_task_instance.xcom_push.call_count == 2
mock_task_instance.xcom_push.assert_any_call(key="spreadsheet_id",
value=SPREADSHEET_ID)
mock_task_instance.xcom_push.assert_any_call(key="spreadsheet_url",
value=SPREADSHEET_URL)
- assert op_execute_result["spreadsheetId"] == "1234567890"
- assert op_execute_result["spreadsheetUrl"] == "https://example/sheets"
+ assert op_execute_result["spreadsheetId"] == SPREADSHEET_ID
+ assert op_execute_result["spreadsheetUrl"] == SPREADSHEET_URL
+
+
@mock.patch("airflow.providers.google.suite.operators.sheets.GoogleDriveHook")
+ def test_execute_via_drive_api(self, mock_drive_hook):
+ """Test spreadsheet creation using the Drive API with drive_id
provided."""
+ mock_task_instance = mock.MagicMock()
+ context = {"task_instance": mock_task_instance}
+
+ mock_drive_hook.return_value.create_file.return_value = {
+ "id": SPREADSHEET_ID,
+ "webViewLink": SPREADSHEET_URL,
+ }
+
+ op = GoogleSheetsCreateSpreadsheetOperator(
+ task_id="test_task", spreadsheet=SPREADSHEET_DATA,
gcp_conn_id=GCP_CONN_ID, drive_id=DRIVE_ID
+ )
+ op_execute_result = op.execute(context)
+
+ expected_metadata = {
+ "name": "My Test Spreadsheet",
+ "mimeType": "application/vnd.google-apps.spreadsheet",
+ "parents": [DRIVE_ID],
+ }
+
+
mock_drive_hook.return_value.create_file.assert_called_once_with(file_metadata=expected_metadata)
+
+ assert mock_task_instance.xcom_push.call_count == 2
+ mock_task_instance.xcom_push.assert_any_call(key="spreadsheet_id",
value=SPREADSHEET_ID)
+ mock_task_instance.xcom_push.assert_any_call(key="spreadsheet_url",
value=SPREADSHEET_URL)
+
+ assert op_execute_result["spreadsheetId"] == SPREADSHEET_ID
+ assert op_execute_result["spreadsheetUrl"] == SPREADSHEET_URL