Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ SnowflakeCortexAgentCreateOperator
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^

To create a Snowflake Cortex Agent you can use
:class:`~airflow.providers.snowflake.operators.snowflake_cortex_agent.SnowflakeCortexAgentCreateOperator`.
:class:`~airflow.providers.snowflake.operators.cortex_agent.SnowflakeCortexAgentCreateOperator`.

.. exampleinclude:: /../../snowflake/tests/system/snowflake/example_snowflake_cortex_agent.py
:language: python
Expand All @@ -80,7 +80,7 @@ SnowflakeCortexAgentUpdateOperator
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^

To update an existing Snowflake Cortex Agent you can use
:class:`~airflow.providers.snowflake.operators.snowflake_cortex_agent.SnowflakeCortexAgentUpdateOperator`.
:class:`~airflow.providers.snowflake.operators.cortex_agent.SnowflakeCortexAgentUpdateOperator`.

Only fields explicitly provided are updated. Optional fields left as ``None``
retain their existing values on the Cortex Agent.
Expand All @@ -96,7 +96,7 @@ retain their existing values on the Cortex Agent.
SnowflakeCortexAgentOperator
^^^^^^^^^^^^^^^^^^^^^^^^^^^^

Use the :class:`~airflow.providers.snowflake.operators.snowflake_cortex_agent.SnowflakeCortexAgentOperator`
Use the :class:`~airflow.providers.snowflake.operators.cortex_agent.SnowflakeCortexAgentOperator`
to execute an existing Snowflake Cortex Agent.

The operator wraps the Snowflake Cortex Agent Run API and executes an existing
Expand All @@ -118,7 +118,7 @@ SnowflakeCortexAgentDeleteOperator
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^

To delete a Snowflake Cortex Agent you can use
:class:`~airflow.providers.snowflake.operators.snowflake_cortex_agent.SnowflakeCortexAgentDeleteOperator`.
:class:`~airflow.providers.snowflake.operators.cortex_agent.SnowflakeCortexAgentDeleteOperator`.

.. exampleinclude:: /../../snowflake/tests/system/snowflake/example_snowflake_cortex_agent.py
:language: python
Expand Down
10 changes: 5 additions & 5 deletions providers/snowflake/provider.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -125,7 +125,7 @@ integrations:
- /docs/apache-airflow-providers-snowflake/operators/snowflake.rst
- /docs/apache-airflow-providers-snowflake/operators/snowpark.rst
- /docs/apache-airflow-providers-snowflake/operators/snowpark_containers.rst
- /docs/apache-airflow-providers-snowflake/operators/snowflake_cortex_agent.rst
- /docs/apache-airflow-providers-snowflake/operators/cortex_agent.rst
logo: /docs/integration-logos/Snowflake.png
tags: [service]

Expand All @@ -135,7 +135,7 @@ operators:
- airflow.providers.snowflake.operators.snowflake
- airflow.providers.snowflake.operators.snowpark
- airflow.providers.snowflake.operators.snowpark_containers
- airflow.providers.snowflake.operators.snowflake_cortex_agent
- airflow.providers.snowflake.operators.cortex_agent

task-decorators:
- class-name: airflow.providers.snowflake.decorators.snowpark.snowpark_task
Expand All @@ -158,10 +158,10 @@ dataset-uris:
hooks:
- integration-name: Snowflake
python-modules:
- airflow.providers.snowflake.hooks.snowflake
- airflow.providers.snowflake.hooks.snowflake_sql_api
- airflow.providers.snowflake.hooks.snowflake_cortex_agent
- airflow.providers.snowflake.hooks.cortex_model
- airflow.providers.snowflake.hooks.snowflake
- airflow.providers.snowflake.hooks.sql_api
- airflow.providers.snowflake.hooks.cortex_agent

transfers:
- source-integration-name: Amazon Simple Storage Service (S3)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ def get_provider_info():
"/docs/apache-airflow-providers-snowflake/operators/snowflake.rst",
"/docs/apache-airflow-providers-snowflake/operators/snowpark.rst",
"/docs/apache-airflow-providers-snowflake/operators/snowpark_containers.rst",
"/docs/apache-airflow-providers-snowflake/operators/snowflake_cortex_agent.rst",
"/docs/apache-airflow-providers-snowflake/operators/cortex_agent.rst",
],
"logo": "/docs/integration-logos/Snowflake.png",
"tags": ["service"],
Expand All @@ -47,7 +47,7 @@ def get_provider_info():
"airflow.providers.snowflake.operators.snowflake",
"airflow.providers.snowflake.operators.snowpark",
"airflow.providers.snowflake.operators.snowpark_containers",
"airflow.providers.snowflake.operators.snowflake_cortex_agent",
"airflow.providers.snowflake.operators.cortex_agent",
],
}
],
Expand Down Expand Up @@ -77,10 +77,10 @@ def get_provider_info():
{
"integration-name": "Snowflake",
"python-modules": [
"airflow.providers.snowflake.hooks.snowflake",
"airflow.providers.snowflake.hooks.snowflake_sql_api",
"airflow.providers.snowflake.hooks.snowflake_cortex_agent",
"airflow.providers.snowflake.hooks.cortex_model",
"airflow.providers.snowflake.hooks.snowflake",
"airflow.providers.snowflake.hooks.sql_api",
"airflow.providers.snowflake.hooks.cortex_agent",
],
}
],
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,3 +14,20 @@
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.

from __future__ import annotations

from airflow.utils.deprecation_tools import add_deprecated_classes

__deprecated_classes = {
"snowflake_cortex_agent": {
"SnowflakeCortexAgentHook": (
"airflow.providers.snowflake.hooks.cortex_agent.SnowflakeCortexAgentHook"
),
},
"snowflake_sql_api": {
"SnowflakeSqlApiHook": ("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook"),
},
}

add_deprecated_classes(__deprecated_classes, __name__)
Original file line number Diff line number Diff line change
Expand Up @@ -14,3 +14,17 @@
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.

from __future__ import annotations

from airflow.utils.deprecation_tools import add_deprecated_classes

__deprecated_classes = {
"snowflake_cortex_agent": {
"SnowflakeCortexAgentOperator": (
"airflow.providers.snowflake.operators.cortex_agent.SnowflakeCortexAgentOperator"
),
},
}

add_deprecated_classes(__deprecated_classes, __name__)
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
from typing import TYPE_CHECKING, Any

from airflow.providers.common.compat.sdk import BaseOperator
from airflow.providers.snowflake.hooks.snowflake_cortex_agent import CreateMode, SnowflakeCortexAgentHook
from airflow.providers.snowflake.hooks.cortex_agent import CreateMode, SnowflakeCortexAgentHook

if TYPE_CHECKING:
from airflow.providers.common.compat.sdk import Context
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@
SQLIntervalCheckOperator,
SQLValueCheckOperator,
)
from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.hooks.sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.triggers.snowflake_trigger import SnowflakeSqlApiTrigger

_DURABLE_UNSET = object()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@

from asgiref.sync import sync_to_async

from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.hooks.sql_api import SnowflakeSqlApiHook
from airflow.triggers.base import BaseTrigger, TriggerEvent

if TYPE_CHECKING:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
from openlineage.client.facet_v2 import JobFacet

from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook
from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.hooks.sql_api import SnowflakeSqlApiHook


log = logging.getLogger(__name__)
Expand Down Expand Up @@ -212,7 +212,7 @@ def _get_queries_details_from_snowflake(

try:
# Note: need to lazy import here to avoid circular imports
from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.hooks.sql_api import SnowflakeSqlApiHook

if isinstance(hook, SnowflakeSqlApiHook):
result = _run_single_query_with_api_hook(hook=hook, sql=query)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,8 +27,8 @@
from datetime import datetime

from airflow import DAG
from airflow.providers.snowflake.hooks.snowflake_cortex_agent import CreateMode
from airflow.providers.snowflake.operators.snowflake_cortex_agent import (
from airflow.providers.snowflake.hooks.cortex_agent import CreateMode
from airflow.providers.snowflake.operators.cortex_agent import (
SnowflakeCortexAgentCreateOperator,
SnowflakeCortexAgentDeleteOperator,
SnowflakeCortexAgentOperator,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,13 +23,13 @@
from cryptography.hazmat.backends import default_backend
from cryptography.hazmat.primitives.asymmetric import rsa

from airflow.providers.snowflake.hooks.snowflake_cortex_agent import (
from airflow.providers.snowflake.hooks.cortex_agent import (
CreateMode,
JsonResponse,
SnowflakeCortexAgentHook,
)

MODULE_PATH = "airflow.providers.snowflake.hooks.snowflake_cortex_agent"
MODULE_PATH = "airflow.providers.snowflake.hooks.cortex_agent"
HOOK_PATH = f"{MODULE_PATH}.SnowflakeCortexAgentHook"

ACCOUNT = "test-account"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@

from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.models import Connection
from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.hooks.sql_api import SnowflakeSqlApiHook

if TYPE_CHECKING:
from pathlib import Path
Expand Down Expand Up @@ -173,7 +173,7 @@

API_URL = "https://test.snowflakecomputing.com/api/v2/statements/test"

MODULE_PATH = "airflow.providers.snowflake.hooks.snowflake_sql_api"
MODULE_PATH = "airflow.providers.snowflake.hooks.sql_api"
HOOK_PATH = f"{MODULE_PATH}.SnowflakeSqlApiHook"


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,11 @@

import pytest

from airflow.providers.snowflake.hooks.snowflake_cortex_agent import (
from airflow.providers.snowflake.hooks.cortex_agent import (
CreateMode,
SnowflakeCortexAgentHook,
)
from airflow.providers.snowflake.operators.snowflake_cortex_agent import (
from airflow.providers.snowflake.operators.cortex_agent import (
SnowflakeCortexAgentCreateOperator,
SnowflakeCortexAgentDeleteOperator,
SnowflakeCortexAgentOperator,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@
TEST_SQL = "select * from any;"
NOTEBOOK = "MY_DB.MY_SCHEMA.MY_NOTEBOOK"

HOOK_MODULE = "airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook"
HOOK_MODULE = "airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook"

SQL_MULTIPLE_STMTS = (
"create or replace table user_test (i int); insert into user_test (i) "
Expand Down Expand Up @@ -438,7 +438,7 @@ def test_poll_on_queries_sleeps_once_per_cycle(self, mock_execute_query, mock_ge
("mock_sql", "statement_count"),
[pytest.param(SQL_MULTIPLE_STMTS, 4, id="multi"), pytest.param(SINGLE_STMT, 1, id="single")],
)
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.execute_query")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.execute_query")
def test_snowflake_sql_api_execute_operator_async(
self, mock_execute_query, mock_sql, statement_count, mock_get_sql_api_query_status
):
Expand Down Expand Up @@ -520,7 +520,7 @@ def test_snowflake_sql_api_execute_complete_failure(self):
({"status": "success", "statement_query_ids": ["uuid", "uuid"]}),
],
)
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.check_query_output")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.check_query_output")
def test_snowflake_sql_api_execute_complete(self, mock_conn, mock_event):
"""Tests execute_complete assert with successful message"""

Expand All @@ -543,7 +543,7 @@ def test_snowflake_sql_api_execute_complete(self, mock_conn, mock_event):
({"status": "success", "statement_query_ids": ["uuid", "uuid"]}),
],
)
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.check_query_output")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.check_query_output")
def test_snowflake_sql_api_execute_complete_reassigns_query_ids(self, mock_conn, mock_event):
"""Tests execute_complete assert with successful message"""

Expand Down Expand Up @@ -728,7 +728,7 @@ def test_poll_until_complete_pushes_query_ids_to_xcom_even_on_failure(

context["ti"].xcom_push.assert_called_once_with(key="query_ids", value=["uuid1"])

@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.cancel_queries")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.cancel_queries")
def test_snowflake_sql_api_on_kill_cancels_queries(self, mock_cancel_queries):
"""Test that on_kill cancels running queries."""
operator = SnowflakeSqlApiOperator(
Expand All @@ -743,7 +743,7 @@ def test_snowflake_sql_api_on_kill_cancels_queries(self, mock_cancel_queries):

mock_cancel_queries.assert_called_once_with(["uuid1", "uuid2"])

@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.cancel_queries")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.cancel_queries")
def test_snowflake_sql_api_on_kill_no_queries(self, mock_cancel_queries):
"""Test that on_kill does nothing when no query ids exist."""
operator = SnowflakeSqlApiOperator(
Expand All @@ -758,7 +758,7 @@ def test_snowflake_sql_api_on_kill_no_queries(self, mock_cancel_queries):

mock_cancel_queries.assert_not_called()

@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.cancel_queries")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.cancel_queries")
def test_snowflake_sql_api_on_kill_respects_cancel_on_kill_false(self, mock_cancel_queries):
"""on_kill does not cancel queries when cancel_on_kill is disabled."""
operator = SnowflakeSqlApiOperator(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,7 @@ async def test_on_kill_cancels_remaining_after_one_fails(self, mock_hook):

@pytest.mark.asyncio
@mock.patch(f"{MODULE}.triggers.snowflake_trigger.SnowflakeSqlApiTrigger.get_query_status")
@mock.patch(f"{MODULE}.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
@mock.patch(f"{MODULE}.hooks.sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
async def test_snowflake_sql_trigger_running(
self, mock_get_sql_api_query_status_async, mock_get_query_status
):
Expand All @@ -146,7 +146,7 @@ async def test_snowflake_sql_trigger_running(

@pytest.mark.asyncio
@mock.patch(f"{MODULE}.triggers.snowflake_trigger.SnowflakeSqlApiTrigger.get_query_status")
@mock.patch(f"{MODULE}.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
@mock.patch(f"{MODULE}.hooks.sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
async def test_snowflake_sql_trigger_completed(
self, mock_get_sql_api_query_status_async, mock_get_query_status
):
Expand All @@ -167,7 +167,7 @@ async def test_snowflake_sql_trigger_completed(
assert TriggerEvent({"status": "success", "statement_query_ids": QUERY_IDS}) == actual

@pytest.mark.asyncio
@mock.patch(f"{MODULE}.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
@mock.patch(f"{MODULE}.hooks.sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
async def test_snowflake_sql_trigger_failure_status(self, mock_get_sql_api_query_status_async):
"""Test SnowflakeSqlApiTrigger task is executed and triggered with failure status."""
mock_response = {
Expand All @@ -182,7 +182,7 @@ async def test_snowflake_sql_trigger_failure_status(self, mock_get_sql_api_query
assert TriggerEvent(mock_response) == actual

@pytest.mark.asyncio
@mock.patch(f"{MODULE}.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
@mock.patch(f"{MODULE}.hooks.sql_api.SnowflakeSqlApiHook.get_sql_api_query_status_async")
async def test_snowflake_sql_trigger_exception(self, mock_get_sql_api_query_status_async):
"""Tests the SnowflakeSqlApiTrigger does not fire if there is an exception."""
mock_get_sql_api_query_status_async.side_effect = Exception("Test exception")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException, timezone
from airflow.providers.openlineage.conf import namespace
from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook
from airflow.providers.snowflake.hooks.snowflake_sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.hooks.sql_api import SnowflakeSqlApiHook
from airflow.providers.snowflake.utils.openlineage import (
_create_snowflake_event_pair,
_get_parent_run_facet,
Expand Down Expand Up @@ -199,10 +199,10 @@ def test_run_single_query_with_hook(mock_get_cursor, mock_set_autocommit, mock_g


@mock.patch(
"airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_result_from_successful_sql_api_query"
"airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.get_result_from_successful_sql_api_query"
)
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.wait_for_query")
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.execute_query")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.wait_for_query")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.execute_query")
def test_run_single_query_with_api_hook_success(mock_execute, mock_wait, mock_get_result):
hook = SnowflakeSqlApiHook(snowflake_conn_id="test_conn")
hook.query_ids = ["old-id"]
Expand All @@ -225,10 +225,10 @@ def execute_query_side_effect(*args, **kwargs):


@mock.patch(
"airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.get_result_from_successful_sql_api_query"
"airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.get_result_from_successful_sql_api_query"
)
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.wait_for_query")
@mock.patch("airflow.providers.snowflake.hooks.snowflake_sql_api.SnowflakeSqlApiHook.execute_query")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.wait_for_query")
@mock.patch("airflow.providers.snowflake.hooks.sql_api.SnowflakeSqlApiHook.execute_query")
def test_run_single_query_exception_restores_query_ids(mock_execute, mock_wait, mock_get_result):
hook = SnowflakeSqlApiHook(snowflake_conn_id="test_conn")
hook.query_ids = ["persistent-id"]
Expand Down