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

Reply via email to