This is an automated email from the ASF dual-hosted git repository.
shahar1 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 6ede6646097 Add SnowflakeNotebookOperator for executing Snowflake
Notebooks (#63470)
6ede6646097 is described below
commit 6ede6646097523559c075898b8e6ae97ee824b7a
Author: Jacob Beaudin <[email protected]>
AuthorDate: Sun Sep 20 02:39:35 2026 -0700
Add SnowflakeNotebookOperator for executing Snowflake Notebooks (#63470)
---
providers/snowflake/docs/operators/snowflake.rst | 36 +++
.../providers/snowflake/operators/snowflake.py | 61 +++-
.../system/snowflake/example_snowflake_notebook.py | 75 +++++
.../unit/snowflake/operators/test_snowflake.py | 334 +++++++++++++++++++++
4 files changed, 505 insertions(+), 1 deletion(-)
diff --git a/providers/snowflake/docs/operators/snowflake.rst
b/providers/snowflake/docs/operators/snowflake.rst
index 4cb3059718c..fb58281364a 100644
--- a/providers/snowflake/docs/operators/snowflake.rst
+++ b/providers/snowflake/docs/operators/snowflake.rst
@@ -218,3 +218,39 @@ already tracks the statement handles across the wait, so
deferrable mode takes p
Durable execution requires Airflow 3.3 or newer, since it relies on the task
state store. Below
3.3, ``durable`` has no effect either way: setting it explicitly only emits a
warning, and the
operator always submits fresh SQL on retry, exactly as before this feature
existed.
+
+
+SnowflakeNotebookOperator
+=========================
+
+Use the :class:`SnowflakeNotebookOperator
<airflow.providers.snowflake.operators.snowflake.SnowflakeNotebookOperator>`
+to execute a `Snowflake Notebook
<https://docs.snowflake.com/en/sql-reference/sql/execute-notebook>`__
+via the Snowflake SQL API.
+
+This operator builds an ``EXECUTE NOTEBOOK`` statement and delegates execution
to
+:class:`SnowflakeSqlApiOperator
<airflow.providers.snowflake.operators.snowflake.SnowflakeSqlApiOperator>`.
+
+Using the Operator
+^^^^^^^^^^^^^^^^^^
+
+.. exampleinclude::
/../../snowflake/tests/system/snowflake/example_snowflake_notebook.py
+ :language: python
+ :start-after: [START howto_operator_snowflake_notebook]
+ :end-before: [END howto_operator_snowflake_notebook]
+ :dedent: 4
+
+You can pass parameters to the notebook:
+
+.. exampleinclude::
/../../snowflake/tests/system/snowflake/example_snowflake_notebook.py
+ :language: python
+ :start-after: [START howto_operator_snowflake_notebook_with_params]
+ :end-before: [END howto_operator_snowflake_notebook_with_params]
+ :dedent: 4
+
+You can also run the operator in deferrable mode:
+
+.. exampleinclude::
/../../snowflake/tests/system/snowflake/example_snowflake_notebook.py
+ :language: python
+ :start-after: [START howto_operator_snowflake_notebook_deferrable]
+ :end-before: [END howto_operator_snowflake_notebook_deferrable]
+ :dedent: 4
diff --git
a/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py
b/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py
index f716a88ef71..0d4adbb0115 100644
--- a/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py
+++ b/providers/snowflake/src/airflow/providers/snowflake/operators/snowflake.py
@@ -22,7 +22,7 @@ import warnings
from collections.abc import Iterable, Mapping, Sequence
from datetime import timedelta
from functools import cached_property
-from typing import TYPE_CHECKING, Any, SupportsAbs, cast
+from typing import TYPE_CHECKING, Any, ClassVar, SupportsAbs, cast
import requests
@@ -649,3 +649,62 @@ class SnowflakeSqlApiOperator(ResumableJobMixin,
SQLExecuteQueryOperator):
self.log.info("Cancelling the query ids %s", self.query_ids)
self._hook.cancel_queries(self.query_ids)
self.log.info("Query ids %s cancelled successfully",
self.query_ids)
+
+
+class SnowflakeNotebookOperator(SnowflakeSqlApiOperator):
+ """
+ Execute a Snowflake Notebook via the Snowflake SQL API.
+
+ Builds an ``EXECUTE NOTEBOOK`` statement and delegates execution to
+
:class:`~airflow.providers.snowflake.operators.snowflake.SnowflakeSqlApiOperator`,
+ which handles query submission, polling, deferral, and cancellation.
+
+ .. seealso::
+ `Snowflake EXECUTE NOTEBOOK
+ <https://docs.snowflake.com/en/sql-reference/sql/execute-notebook>`_
+
+ :param notebook: Fully-qualified notebook name
+ (e.g. ``MY_DB.MY_SCHEMA.MY_NOTEBOOK``).
+ :param notebook_parameters: Optional list of string parameters to pass to
+ the notebook. Values must be strings (the type hint declares
+ ``list[str]``). Parameters are accessible in the notebook via
+ ``sys.argv``.
+ """
+
+ template_fields: Sequence[str] = tuple(
+ set(SnowflakeSqlApiOperator.template_fields) | {"notebook",
"notebook_parameters"}
+ )
+ # The SQL is generated from `notebook`/`notebook_parameters`, never loaded
from a
+ # file, so the inherited `.sql`/`.json` extensions would only cause harm:
a notebook
+ # parameter that happens to end in one gets replaced by the contents of a
file of
+ # that name.
+ template_ext: Sequence[str] = ()
+ # Same reason the parent's `parameters` renderer is dropped: it describes
SQL bind
+ # parameters, not notebook arguments.
+ template_fields_renderers: ClassVar[dict] = {"sql": "sql"}
+
+ def __init__(
+ self,
+ *,
+ notebook: str,
+ notebook_parameters: list[str] | None = None,
+ **kwargs: Any,
+ ) -> None:
+ self.notebook = notebook
+ self.notebook_parameters = notebook_parameters
+ super().__init__(sql=self._build_execute_notebook_query(),
statement_count=1, **kwargs)
+
+ def execute(self, context: Context) -> None:
+ """Rebuild SQL from rendered template fields, then execute."""
+ self.sql = self._build_execute_notebook_query()
+ return super().execute(context)
+
+ def _build_execute_notebook_query(self) -> str:
+ """Build the ``EXECUTE NOTEBOOK`` SQL statement."""
+ params_clause = ""
+ if self.notebook_parameters:
+ # Escape backslashes first (Snowflake interprets `\` in string
literals),
+ # then single quotes.
+ sanitized = [p.replace("\\", "\\\\").replace("'", "''") for p in
self.notebook_parameters]
+ params_clause = ", ".join(f"'{p}'" for p in sanitized)
+ return f"EXECUTE NOTEBOOK {self.notebook}({params_clause})"
diff --git
a/providers/snowflake/tests/system/snowflake/example_snowflake_notebook.py
b/providers/snowflake/tests/system/snowflake/example_snowflake_notebook.py
new file mode 100644
index 00000000000..0185db0fdf7
--- /dev/null
+++ b/providers/snowflake/tests/system/snowflake/example_snowflake_notebook.py
@@ -0,0 +1,75 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""
+Example use of SnowflakeNotebookOperator.
+"""
+
+from __future__ import annotations
+
+import os
+from datetime import datetime
+
+from airflow import DAG
+from airflow.providers.snowflake.operators.snowflake import
SnowflakeNotebookOperator
+
+SNOWFLAKE_CONN_ID = "my_snowflake_conn"
+SNOWFLAKE_NOTEBOOK = os.environ.get("SNOWFLAKE_NOTEBOOK",
"MY_DB.MY_SCHEMA.MY_NOTEBOOK")
+ENV_ID = os.environ.get("SYSTEM_TESTS_ENV_ID")
+DAG_ID = "example_snowflake_notebook"
+
+with DAG(
+ DAG_ID,
+ start_date=datetime(2021, 1, 1),
+ default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID},
+ tags=["example"],
+ schedule="@once",
+ catchup=False,
+) as dag:
+ # [START howto_operator_snowflake_notebook]
+ execute_notebook = SnowflakeNotebookOperator(
+ task_id="execute_notebook",
+ notebook=SNOWFLAKE_NOTEBOOK,
+ snowflake_conn_id=SNOWFLAKE_CONN_ID,
+ )
+ # [END howto_operator_snowflake_notebook]
+
+ # [START howto_operator_snowflake_notebook_with_params]
+ execute_notebook_with_params = SnowflakeNotebookOperator(
+ task_id="execute_notebook_with_params",
+ notebook=SNOWFLAKE_NOTEBOOK,
+ snowflake_conn_id=SNOWFLAKE_CONN_ID,
+ notebook_parameters=["param1", "target_db=PROD"],
+ )
+ # [END howto_operator_snowflake_notebook_with_params]
+
+ # [START howto_operator_snowflake_notebook_deferrable]
+ execute_notebook_deferrable = SnowflakeNotebookOperator(
+ task_id="execute_notebook_deferrable",
+ notebook=SNOWFLAKE_NOTEBOOK,
+ snowflake_conn_id=SNOWFLAKE_CONN_ID,
+ deferrable=True,
+ )
+ # [END howto_operator_snowflake_notebook_deferrable]
+
+ execute_notebook >> execute_notebook_with_params >>
execute_notebook_deferrable
+
+
+from tests_common.test_utils.system_tests import get_test_run # noqa: E402
+
+# Needed to run the example DAG with pytest (see:
tests/system/README.md#run_via_pytest)
+test_run = get_test_run(dag)
diff --git
a/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py
b/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py
index 517691fed3e..c8a75157fe8 100644
--- a/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py
+++ b/providers/snowflake/tests/unit/snowflake/operators/test_snowflake.py
@@ -35,6 +35,7 @@ from airflow.providers.snowflake.operators.snowflake import (
_DURABLE_UNSET,
SnowflakeCheckOperator,
SnowflakeIntervalCheckOperator,
+ SnowflakeNotebookOperator,
SnowflakeSqlApiOperator,
SnowflakeValueCheckOperator,
_warn_and_disable_durable_pre_3_3,
@@ -55,6 +56,9 @@ TEST_DAG_ID = "unit_test_dag"
TASK_ID = "snowflake_check"
CONN_ID = "my_snowflake_conn"
TEST_SQL = "select * from any;"
+NOTEBOOK = "MY_DB.MY_SCHEMA.MY_NOTEBOOK"
+
+HOOK_MODULE =
"airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook"
SQL_MULTIPLE_STMTS = (
"create or replace table user_test (i int); insert into user_test (i) "
@@ -1006,3 +1010,333 @@ class TestWarnAndDisableDurableAirflowPre3_3:
with pytest.warns(UserWarning, match="durable.*no effect"):
result = _warn_and_disable_durable_pre_3_3(value)
assert result is False
+
+
+class TestSnowflakeNotebookOperatorSQL:
+ """Tests for SQL query building."""
+
+ def test_build_sql_no_params(self):
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ )
+ assert operator.sql == "EXECUTE NOTEBOOK MY_DB.MY_SCHEMA.MY_NOTEBOOK()"
+
+ def test_build_sql_with_params(self):
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=["param1", "target_db=PROD"],
+ )
+ assert operator.sql == "EXECUTE NOTEBOOK
MY_DB.MY_SCHEMA.MY_NOTEBOOK('param1', 'target_db=PROD')"
+
+ def test_build_sql_escapes_single_quotes(self):
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=["O'Brien", "it's"],
+ )
+ assert operator.sql == "EXECUTE NOTEBOOK
MY_DB.MY_SCHEMA.MY_NOTEBOOK('O''Brien', 'it''s')"
+
+ def test_build_sql_escapes_backslashes(self):
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=["C:\\data", "a\\'b"],
+ )
+ assert operator.sql == "EXECUTE NOTEBOOK
MY_DB.MY_SCHEMA.MY_NOTEBOOK('C:\\\\data', 'a\\\\''b')"
+
+ def
test_notebook_parameters_do_not_collide_with_parent_bind_parameters(self):
+ """`notebook_parameters` must stay distinct from the parent's SQL bind
`parameters`."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=["a", "b"],
+ )
+ assert operator.notebook_parameters == ["a", "b"]
+ assert operator.parameters is None
+
+ def
test_parameter_ending_in_sql_or_json_is_not_replaced_by_file_contents(self,
tmp_path):
+ """A parameter that looks like a filename stays a literal value."""
+ (tmp_path / "config.json").write_text('{"secret": "from file"}')
+ (tmp_path / "query.sql").write_text("SELECT 1")
+ with DAG(
+ "test_notebook_template_ext",
+ schedule=None,
+ start_date=DEFAULT_DATE,
+ template_searchpath=str(tmp_path),
+ ) as dag:
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=["config.json", "query.sql", "plain"],
+ dag=dag,
+ )
+ operator.resolve_template_files()
+ assert operator.notebook_parameters == ["config.json", "query.sql",
"plain"]
+ assert operator._build_execute_notebook_query() == (
+ "EXECUTE NOTEBOOK MY_DB.MY_SCHEMA.MY_NOTEBOOK('config.json',
'query.sql', 'plain')"
+ )
+
+ def test_execute_rebuilds_sql_from_rendered_parameters(self):
+ """Simulate template rendering mutating parameters; execute() should
rebuild SQL with escaping."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=["{{ params.name }}"],
+ )
+ # Simulate Airflow template rendering mutating the attribute in-place
+ operator.notebook_parameters = ["O'Brien"]
+ operator.sql = operator._build_execute_notebook_query()
+ assert operator.sql == "EXECUTE NOTEBOOK
MY_DB.MY_SCHEMA.MY_NOTEBOOK('O''Brien')"
+
+ def test_real_template_rendering_escapes_rendered_value(self):
+ """Rendering mutates the fields in place, so the rebuild must escape
the rendered value."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook="{{ params.db }}.NB",
+ notebook_parameters=["{{ params.name }}", "plain"],
+ )
+ operator.render_template_fields({"params": {"db": "MY_DB.MY_SCHEMA",
"name": "O'Brien"}})
+ assert operator.notebook == "MY_DB.MY_SCHEMA.NB"
+ assert operator.notebook_parameters == ["O'Brien", "plain"]
+ assert operator._build_execute_notebook_query() == (
+ "EXECUTE NOTEBOOK MY_DB.MY_SCHEMA.NB('O''Brien', 'plain')"
+ )
+
+ def test_build_sql_empty_params(self):
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ notebook_parameters=[],
+ )
+ assert operator.sql == "EXECUTE NOTEBOOK MY_DB.MY_SCHEMA.MY_NOTEBOOK()"
+
+ def test_sql_renderer_is_preserved(self):
+ """Dropping the parent's `parameters` renderer must not lose SQL
highlighting."""
+ assert SnowflakeNotebookOperator.template_fields_renderers == {"sql":
"sql"}
+
+ def test_template_fields(self):
+ assert "notebook" in SnowflakeNotebookOperator.template_fields
+ assert "notebook_parameters" in
SnowflakeNotebookOperator.template_fields
+ assert "snowflake_conn_id" in SnowflakeNotebookOperator.template_fields
+
+ def test_statement_count_is_one(self):
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ notebook=NOTEBOOK,
+ )
+ assert operator.statement_count == 1
+
+ def test_is_subclass_of_snowflake_sql_api_operator(self):
+ assert issubclass(SnowflakeNotebookOperator, SnowflakeSqlApiOperator)
+
+
[email protected]_test
+class TestSnowflakeNotebookOperator:
+ @pytest.fixture(autouse=True)
+ def setup_tests(self):
+ clear_db_dags()
+ clear_db_runs()
+ if AIRFLOW_V_3_0_PLUS:
+ clear_db_dag_bundles()
+
+ yield
+
+ clear_db_dags()
+ clear_db_runs()
+ if AIRFLOW_V_3_0_PLUS:
+ clear_db_dag_bundles()
+
+ def test_execute_success_immediate(
+ self, mock_execute_query, mock_get_sql_api_query_status,
mock_check_query_output
+ ):
+ """Notebook completes on the first status check, without deferring."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id="snowflake_default",
+ notebook=NOTEBOOK,
+ do_xcom_push=False,
+ durable=False,
+ )
+ mock_execute_query.return_value = ["uuid1"]
+ mock_get_sql_api_query_status.side_effect = [{"status": "success"}]
+ operator.execute(context=None)
+
+ def test_execute_failure_immediate(self, mock_execute_query,
mock_get_sql_api_query_status):
+ """Notebook that fails on the first status check raises."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id="snowflake_default",
+ notebook=NOTEBOOK,
+ do_xcom_push=False,
+ durable=False,
+ )
+ mock_execute_query.return_value = ["uuid1"]
+ mock_get_sql_api_query_status.side_effect = [{"status": "error",
"message": "Notebook failed"}]
+ with pytest.raises(RuntimeError):
+ operator.execute(context=None)
+
+ @mock.patch(f"{HOOK_MODULE}.execute_query")
+ def test_execute_deferred(self, mock_execute_query,
mock_get_sql_api_query_status):
+ """Running notebook with deferrable=True raises TaskDeferred."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ deferrable=True,
+ )
+ mock_execute_query.return_value = ["uuid1"]
+ mock_get_sql_api_query_status.side_effect = [{"status": "running"}]
+
+ with pytest.raises(TaskDeferred) as exc:
+ operator.execute(create_context(operator))
+
+ assert isinstance(exc.value.trigger, SnowflakeSqlApiTrigger)
+
+ def test_execute_polling_success(
+ self, mock_execute_query, mock_get_sql_api_query_status,
mock_check_query_output
+ ):
+ """Non-deferrable mode polls until success."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ do_xcom_push=False,
+ deferrable=False,
+ durable=False,
+ )
+ mock_execute_query.return_value = ["uuid1"]
+ mock_get_sql_api_query_status.side_effect = [
+ {"status": "running"},
+ {"status": "running"},
+ {"status": "success"},
+ ]
+
+ with mock.patch("time.sleep"):
+ operator.execute(context=None)
+
+ def test_execute_polling_failure(self, mock_execute_query,
mock_get_sql_api_query_status):
+ """Non-deferrable mode raises when polling finds error."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ do_xcom_push=False,
+ deferrable=False,
+ durable=False,
+ )
+ mock_execute_query.return_value = ["uuid1"]
+ mock_get_sql_api_query_status.side_effect = [
+ {"status": "running"},
+ {"status": "error", "message": "Notebook execution failed"},
+ ]
+
+ with mock.patch("time.sleep"), pytest.raises(RuntimeError):
+ operator.execute(context=None)
+
+ def test_execute_xcom_push(
+ self, mock_execute_query, mock_get_sql_api_query_status,
mock_check_query_output
+ ):
+ """XCom push stores query_ids when do_xcom_push is True."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id="snowflake_default",
+ notebook=NOTEBOOK,
+ do_xcom_push=True,
+ durable=False,
+ )
+ mock_execute_query.return_value = ["uuid1"]
+ mock_get_sql_api_query_status.side_effect = [{"status": "success"}]
+
+ mock_ti = mock.Mock(spec=TaskInstance)
+ context = {"ti": mock_ti}
+ operator.execute(context=context)
+ mock_ti.xcom_push.assert_called_once_with(key="query_ids",
value=["uuid1"])
+
+ def test_execute_complete_success(self):
+ """execute_complete handles success event."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ deferrable=True,
+ )
+ event = {"status": "success", "statement_query_ids": ["uuid1"]}
+ with mock.patch(f"{HOOK_MODULE}.check_query_output"):
+ operator.execute_complete(context=None, event=event)
+
+ def test_execute_complete_failure(self):
+ """execute_complete raises on an error event."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ deferrable=True,
+ )
+ with pytest.raises(RuntimeError):
+ operator.execute_complete(
+ context=None,
+ event={"status": "error", "message": "Notebook failed"},
+ )
+
+ def test_execute_complete_none_event(self):
+ """execute_complete handles None event gracefully."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ deferrable=True,
+ )
+ operator.execute_complete(context=None, event=None)
+
+ def test_execute_complete_reassigns_query_ids(self):
+ """execute_complete sets query_ids from event."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ deferrable=True,
+ )
+ assert operator.query_ids == []
+ with mock.patch(f"{HOOK_MODULE}.check_query_output"):
+ operator.execute_complete(
+ context=None,
+ event={"status": "success", "statement_query_ids": ["uuid1",
"uuid2"]},
+ )
+ assert operator.query_ids == ["uuid1", "uuid2"]
+
+ @mock.patch(f"{HOOK_MODULE}.cancel_queries")
+ def test_on_kill_with_queries(self, mock_cancel_queries):
+ """on_kill cancels running queries."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ )
+ operator.query_ids = ["uuid1", "uuid2"]
+ operator.on_kill()
+ mock_cancel_queries.assert_called_once_with(["uuid1", "uuid2"])
+
+ @mock.patch(f"{HOOK_MODULE}.cancel_queries")
+ def test_on_kill_no_queries(self, mock_cancel_queries):
+ """on_kill does nothing when no query ids exist."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ )
+ operator.query_ids = []
+ operator.on_kill()
+ mock_cancel_queries.assert_not_called()
+
+ def test_hook_caching(self):
+ """_hook property returns the same instance on repeated access."""
+ operator = SnowflakeNotebookOperator(
+ task_id=TASK_ID,
+ snowflake_conn_id=CONN_ID,
+ notebook=NOTEBOOK,
+ )
+ hook1 = operator._hook
+ hook2 = operator._hook
+ assert hook1 is hook2