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 f2112c85667 Add Cortex Agent management methods to 
SnowflakeCortexAgentHook (#70101)
f2112c85667 is described below

commit f2112c856670d96f31f3a0e71235538ff95b816e
Author: SameerMesiah97 <[email protected]>
AuthorDate: Sat Sep 19 07:41:28 2026 +0100

    Add Cortex Agent management methods to SnowflakeCortexAgentHook (#70101)
---
 .../snowflake/hooks/snowflake_cortex_agent.py      | 184 +++++++++++++++-
 .../snowflake/hooks/test_snowflake_cortex_agent.py | 237 ++++++++++++++++++++-
 2 files changed, 402 insertions(+), 19 deletions(-)

diff --git 
a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
 
b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
index 1413d7fee71..59a4ac0bfa1 100644
--- 
a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
+++ 
b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
@@ -17,12 +17,17 @@
 
 from __future__ import annotations
 
-from typing import Any
+from typing import Any, Literal, overload
+from urllib.parse import quote
 
 import requests
 
 from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook
 
+JsonDict = dict[str, Any]
+JsonList = list[JsonDict]
+JsonResponse = JsonDict | JsonList
+
 
 class SnowflakeCortexAgentHook(SnowflakeHook):
     """Hook for interacting with Snowflake Cortex Agents."""
@@ -48,14 +53,40 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
 
         return token
 
+    @overload
+    def _request(
+        self,
+        *,
+        method: str,
+        endpoint: str,
+        payload: JsonDict | None = None,
+        params: JsonDict | None = None,
+        timeout: int | None = None,
+        response_type: Literal["dict"],
+    ) -> JsonDict: ...
+
+    @overload
+    def _request(
+        self,
+        *,
+        method: str,
+        endpoint: str,
+        payload: JsonDict | None = None,
+        params: JsonDict | None = None,
+        timeout: int | None = None,
+        response_type: Literal["list"],
+    ) -> JsonList: ...
+
     def _request(
         self,
         *,
         method: str,
         endpoint: str,
-        payload: dict[str, Any] | None = None,
+        payload: JsonDict | None = None,
+        params: JsonDict | None = None,
         timeout: int | None = None,
-    ) -> dict[str, Any]:
+        response_type: Literal["dict", "list"],
+    ) -> JsonResponse:
 
         response = requests.request(
             method=method,
@@ -65,6 +96,7 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
                 "Content-Type": "application/json",
             },
             json=payload,
+            params=params,
             timeout=timeout,
         )
 
@@ -77,7 +109,20 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
 
         response.raise_for_status()
 
-        return response.json()
+        data = response.json()
+
+        if response_type == "dict":
+            if not isinstance(data, dict):
+                raise TypeError(f"Expected dict response, got 
{type(data).__name__}")
+            return data
+
+        if not isinstance(data, list):
+            raise TypeError(f"Expected list[dict] response, got 
{type(data).__name__}")
+
+        if not all(isinstance(item, dict) for item in data):
+            raise TypeError("Expected list[dict] response, got list containing 
non-dict elements")
+
+        return data
 
     def run_agent(
         self,
@@ -95,7 +140,7 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
         tools: list[dict[str, Any]] | None = None,
         tool_resources: dict[str, Any] | None = None,
         timeout: int | None = 600,
-    ) -> dict[str, Any]:
+    ) -> JsonDict:
         """
         Execute a Snowflake Cortex Agent and return the response payload.
 
@@ -106,11 +151,10 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
             conversation, this should contain the conversation history and the
             current user message. When ``thread_id`` and ``parent_message_id``
             are provided, this should contain only the current user message.
-        :param thread_id: Existing conversation thread identifier. Optional.
-            When provided, ``parent_message_id`` must also be supplied.
-            Defaults to ``None``.
+        :param thread_id: Existing conversation thread identifier. When 
provided,
+            ``parent_message_id`` must also be supplied. Optional. Defaults to 
``None``.
         :param parent_message_id: Parent message identifier within the 
specified
-            thread. Required when ``thread_id`` is provided. Defaults to 
``None``.
+            thread. Required when ``thread_id`` is provided. Optional. 
Defaults to ``None``.
         :param tool_choice: Tool selection configuration for the agent. 
Optional.
             Defaults to ``None``.
         :param models: Model configuration for the agent. Optional. Defaults to
@@ -124,7 +168,7 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
         :param tool_resources: Configuration for tools specified in ``tools``.
             Optional. Defaults to ``None``.
         :param timeout: Maximum time in seconds to wait for the Cortex Agent 
request
-            to complete. Defaults to ``600``.
+            to complete. Optional. Defaults to ``600``.
         :return: JSON response returned by the Cortex Agent.
         """
         if thread_id is not None and parent_message_id is None:
@@ -157,13 +201,131 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
         if tool_resources is not None:
             payload["tool_resources"] = tool_resources
 
-        endpoint = 
f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}:run"
+        endpoint = (
+            f"/api/v2/databases/{quote(database, safe='')}"
+            f"/schemas/{quote(schema, safe='')}"
+            f"/agents/{quote(agent_name, safe='')}:run"
+        )
 
         return self._request(
             method="POST",
             endpoint=endpoint,
             payload=payload,
             timeout=timeout,
+            response_type="dict",
+        )
+
+    def describe_agent(
+        self,
+        *,
+        database: str,
+        schema: str,
+        agent_name: str,
+        timeout: int | None = 600,
+    ) -> JsonDict:
+        """
+        Describe a Snowflake Cortex Agent.
+
+        :param database: Database containing the Cortex Agent.
+        :param schema: Schema containing the Cortex Agent.
+        :param agent_name: Name of the Cortex Agent.
+        :param timeout: Maximum time in seconds to wait for the Cortex Agent
+            request to complete. Optional. Defaults to ``600``.
+        :return: JSON description of the Cortex Agent.
+        """
+        endpoint = (
+            f"/api/v2/databases/{quote(database, safe='')}"
+            f"/schemas/{quote(schema, safe='')}"
+            f"/agents/{quote(agent_name, safe='')}"
+        )
+
+        return self._request(
+            method="GET",
+            endpoint=endpoint,
+            timeout=timeout,
+            response_type="dict",
+        )
+
+    def list_agents(
+        self,
+        *,
+        database: str,
+        schema: str,
+        like: str | None = None,
+        from_name: str | None = None,
+        show_limit: int | None = None,
+        timeout: int | None = 600,
+    ) -> JsonList:
+        """
+        List one page of Snowflake Cortex Agents.
+
+        :param database: Database containing the Cortex Agents.
+        :param schema: Schema containing the Cortex Agents.
+        :param like: Case-insensitive name filter. Optional.
+            Defaults to ``None``.
+        :param from_name: Agent name from which to continue listing results. 
Pass the
+            pagination value from the preceding page to retrieve the next 
page. Optional.
+            Defaults to ``None``.
+        :param show_limit: Maximum number of agents to include in this page. 
Optional.
+            Defaults to ``None``.
+        :param timeout: Maximum time in seconds to wait for the Cortex Agent
+            request to complete. Optional. Defaults to ``600``.
+        :return: One page of Cortex Agents.
+        """
+        endpoint = f"/api/v2/databases/{quote(database, 
safe='')}/schemas/{quote(schema, safe='')}/agents"
+
+        params: dict[str, Any] = {}
+
+        if like is not None:
+            params["like"] = like
+
+        if from_name is not None:
+            params["fromName"] = from_name
+
+        if show_limit is not None:
+            params["showLimit"] = show_limit
+
+        return self._request(
+            method="GET",
+            endpoint=endpoint,
+            params=params or None,
+            timeout=timeout,
+            response_type="list",
+        )
+
+    def delete_agent(
+        self,
+        *,
+        database: str,
+        schema: str,
+        agent_name: str,
+        if_exists: bool = False,
+        timeout: int | None = 600,
+    ) -> JsonDict:
+        """
+        Delete a Snowflake Cortex Agent.
+
+        :param database: Database containing the Cortex Agent.
+        :param schema: Schema containing the Cortex Agent.
+        :param agent_name: Name of the Cortex Agent.
+        :param if_exists: If ``True``, do not fail when the agent does not 
exist.
+            Optional. Defaults to ``False``.
+        :param timeout: Maximum time in seconds to wait for the Cortex Agent 
request
+            to complete. Optional. Defaults to ``600``.
+        :return: JSON response confirming deletion.
+        """
+        endpoint = (
+            f"/api/v2/databases/{quote(database, safe='')}"
+            f"/schemas/{quote(schema, safe='')}"
+            f"/agents/{quote(agent_name, safe='')}"
+        )
+
+        return self._request(
+            method="DELETE",
+            endpoint=endpoint,
+            params={"ifExists": str(if_exists).lower()},
+            timeout=timeout,
+            response_type="dict",
         )
 
     @staticmethod
diff --git 
a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py 
b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
index e686b698574..8fb60fcee0a 100644
--- 
a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
+++ 
b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
@@ -22,6 +22,7 @@ import pytest
 import requests
 
 from airflow.providers.snowflake.hooks.snowflake_cortex_agent import (
+    JsonResponse,
     SnowflakeCortexAgentHook,
 )
 
@@ -30,9 +31,13 @@ HOOK_PATH = f"{MODULE_PATH}.SnowflakeCortexAgentHook"
 
 ACCOUNT = "test-account"
 ACCESS_TOKEN = "test-token"
-DATABASE = "TEST_DATABASE"
-SCHEMA = "TEST_SCHEMA"
-AGENT_NAME = "TEST_AGENT"
+DATABASE = "TEST/DATABASE"
+SCHEMA = "TEST?SCHEMA"
+AGENT_NAME = "TEST#AGENT"
+
+ENCODED_DATABASE = "TEST%2FDATABASE"
+ENCODED_SCHEMA = "TEST%3FSCHEMA"
+ENCODED_AGENT_NAME = "TEST%23AGENT"
 
 CONN_PARAMS = {
     "account": ACCOUNT,
@@ -49,11 +54,11 @@ REQUEST_TIMEOUT = 600
 def create_response(
     status_code: int = 200,
     *,
-    json_body: dict | None = None,
+    json_body: JsonResponse | None = None,
 ):
     response = mock.MagicMock()
     response.status_code = status_code
-    response.json.return_value = json_body or {}
+    response.json.return_value = {} if json_body is None else json_body
 
     if status_code >= 400:
         response.raise_for_status.side_effect = 
requests.exceptions.HTTPError(response=response)
@@ -64,6 +69,64 @@ def create_response(
 
 
 class TestSnowflakeCortexAgentHook:
+    @pytest.mark.parametrize(
+        ("method_name", "method_kwargs", "json_body", "expected_error"),
+        [
+            pytest.param(
+                "describe_agent",
+                {
+                    "database": DATABASE,
+                    "schema": SCHEMA,
+                    "agent_name": AGENT_NAME,
+                },
+                [{"name": AGENT_NAME}],
+                "Expected dict response, got list",
+                id="describe_agent_expected_dict_got_list",
+            ),
+            pytest.param(
+                "list_agents",
+                {
+                    "database": DATABASE,
+                    "schema": SCHEMA,
+                },
+                {"name": AGENT_NAME},
+                r"Expected list\[dict\] response, got dict",
+                id="list_agents_expected_list_got_dict",
+            ),
+            pytest.param(
+                "list_agents",
+                {
+                    "database": DATABASE,
+                    "schema": SCHEMA,
+                },
+                [{"name": AGENT_NAME}, 1],
+                r"Expected list\[dict\] response, got list containing non-dict 
elements",
+                id="list_agents_contains_non_dict_element",
+            ),
+        ],
+    )
+    @mock.patch(f"{MODULE_PATH}.requests.request")
+    @mock.patch(f"{HOOK_PATH}._get_conn_params")
+    @mock.patch(f"{HOOK_PATH}._get_static_conn_params", 
new_callable=mock.PropertyMock)
+    def test_agent_methods_raise_for_unexpected_response_shape(
+        self,
+        mock_static_conn_params,
+        mock_conn_params,
+        mock_request,
+        method_name,
+        method_kwargs,
+        json_body,
+        expected_error,
+    ):
+        mock_conn_params.return_value = CONN_PARAMS
+        mock_static_conn_params.return_value = STATIC_CONN_PARAMS
+        mock_request.return_value = create_response(json_body=json_body)
+
+        hook = SnowflakeCortexAgentHook(snowflake_conn_id="mock_conn_id")
+
+        with pytest.raises(TypeError, match=expected_error):
+            getattr(hook, method_name)(**method_kwargs)
+
     @mock.patch(f"{MODULE_PATH}.requests.request")
     @mock.patch(f"{HOOK_PATH}._get_conn_params")
     @mock.patch(
@@ -105,9 +168,9 @@ class TestSnowflakeCortexAgentHook:
             method="POST",
             url=(
                 f"https://{ACCOUNT}.snowflakecomputing.com";
-                f"/api/v2/databases/{DATABASE}"
-                f"/schemas/{SCHEMA}"
-                f"/agents/{AGENT_NAME}:run"
+                f"/api/v2/databases/{ENCODED_DATABASE}"
+                f"/schemas/{ENCODED_SCHEMA}"
+                f"/agents/{ENCODED_AGENT_NAME}:run"
             ),
             headers={
                 "Authorization": f"Bearer {ACCESS_TOKEN}",
@@ -127,6 +190,7 @@ class TestSnowflakeCortexAgentHook:
                 ],
                 "stream": False,
             },
+            params=None,
             timeout=REQUEST_TIMEOUT,
         )
 
@@ -316,3 +380,160 @@ class TestSnowflakeCortexAgentHook:
         expected,
     ):
         assert SnowflakeCortexAgentHook.get_text_response(response) == expected
+
+    @mock.patch(f"{MODULE_PATH}.requests.request")
+    @mock.patch(f"{HOOK_PATH}._get_conn_params")
+    @mock.patch(
+        f"{HOOK_PATH}._get_static_conn_params",
+        new_callable=mock.PropertyMock,
+    )
+    def test_describe_agent(
+        self,
+        mock_static_conn_params,
+        mock_conn_params,
+        mock_request,
+    ):
+        mock_conn_params.return_value = CONN_PARAMS
+        mock_static_conn_params.return_value = STATIC_CONN_PARAMS
+        mock_request.return_value = create_response(
+            json_body={"name": AGENT_NAME},
+        )
+
+        hook = SnowflakeCortexAgentHook(
+            snowflake_conn_id="mock_conn_id",
+        )
+
+        result = hook.describe_agent(
+            database=DATABASE,
+            schema=SCHEMA,
+            agent_name=AGENT_NAME,
+        )
+
+        assert result == {"name": AGENT_NAME}
+
+        mock_request.assert_called_once_with(
+            method="GET",
+            url=(
+                f"https://{ACCOUNT}.snowflakecomputing.com";
+                f"/api/v2/databases/{ENCODED_DATABASE}"
+                f"/schemas/{ENCODED_SCHEMA}"
+                f"/agents/{ENCODED_AGENT_NAME}"
+            ),
+            headers={
+                "Authorization": f"Bearer {ACCESS_TOKEN}",
+                "Content-Type": "application/json",
+            },
+            json=None,
+            params=None,
+            timeout=REQUEST_TIMEOUT,
+        )
+
+    @mock.patch(f"{MODULE_PATH}.requests.request")
+    @mock.patch(f"{HOOK_PATH}._get_conn_params")
+    @mock.patch(
+        f"{HOOK_PATH}._get_static_conn_params",
+        new_callable=mock.PropertyMock,
+    )
+    def test_list_agents(
+        self,
+        mock_static_conn_params,
+        mock_conn_params,
+        mock_request,
+    ):
+        mock_conn_params.return_value = CONN_PARAMS
+        mock_static_conn_params.return_value = STATIC_CONN_PARAMS
+        mock_request.return_value = create_response(
+            json_body=[{"name": AGENT_NAME}],
+        )
+
+        hook = SnowflakeCortexAgentHook(
+            snowflake_conn_id="mock_conn_id",
+        )
+
+        result = hook.list_agents(
+            database=DATABASE,
+            schema=SCHEMA,
+            like="AIRFLOW%",
+            from_name="AIRFLOW_TEST",
+            show_limit=10,
+        )
+
+        assert result == [{"name": AGENT_NAME}]
+
+        mock_request.assert_called_once_with(
+            method="GET",
+            url=(
+                f"https://{ACCOUNT}.snowflakecomputing.com";
+                f"/api/v2/databases/{ENCODED_DATABASE}"
+                f"/schemas/{ENCODED_SCHEMA}"
+                f"/agents"
+            ),
+            headers={
+                "Authorization": f"Bearer {ACCESS_TOKEN}",
+                "Content-Type": "application/json",
+            },
+            json=None,
+            params={
+                "like": "AIRFLOW%",
+                "fromName": "AIRFLOW_TEST",
+                "showLimit": 10,
+            },
+            timeout=REQUEST_TIMEOUT,
+        )
+
+    @pytest.mark.parametrize(
+        ("if_exists", "expected"),
+        [
+            pytest.param(True, "true", id="if_exists"),
+            pytest.param(False, "false", id="error_if_missing"),
+        ],
+    )
+    @mock.patch(f"{MODULE_PATH}.requests.request")
+    @mock.patch(f"{HOOK_PATH}._get_conn_params")
+    @mock.patch(
+        f"{HOOK_PATH}._get_static_conn_params",
+        new_callable=mock.PropertyMock,
+    )
+    def test_delete_agent(
+        self,
+        mock_static_conn_params,
+        mock_conn_params,
+        mock_request,
+        if_exists,
+        expected,
+    ):
+        mock_conn_params.return_value = CONN_PARAMS
+        mock_static_conn_params.return_value = STATIC_CONN_PARAMS
+        mock_request.return_value = create_response(
+            json_body={"status": "deleted"},
+        )
+
+        hook = SnowflakeCortexAgentHook(
+            snowflake_conn_id="mock_conn_id",
+        )
+
+        result = hook.delete_agent(
+            database=DATABASE,
+            schema=SCHEMA,
+            agent_name=AGENT_NAME,
+            if_exists=if_exists,
+        )
+
+        assert result == {"status": "deleted"}
+
+        mock_request.assert_called_once_with(
+            method="DELETE",
+            url=(
+                f"https://{ACCOUNT}.snowflakecomputing.com";
+                f"/api/v2/databases/{ENCODED_DATABASE}"
+                f"/schemas/{ENCODED_SCHEMA}"
+                f"/agents/{ENCODED_AGENT_NAME}"
+            ),
+            headers={
+                "Authorization": f"Bearer {ACCESS_TOKEN}",
+                "Content-Type": "application/json",
+            },
+            json=None,
+            params={"ifExists": expected},
+            timeout=REQUEST_TIMEOUT,
+        )

Reply via email to