From 1b8cab740d6e9aeb48500d11e3557cebfd46a061 Mon Sep 17 00:00:00 2001 From: Sameer Mesiah Date: Sun, 19 Jul 2026 21:11:36 +0100 Subject: [PATCH] Add Cortex Agent management methods to SnowflakeCortexAgentHook This change extends SnowflakeCortexAgentHook with support for managing Cortex Agent Objects through the Snowflake REST API. The hook now supports describing, listing and deleting Cortex Agents in addition to executing them via run_agent(). The internal request helper has also been enhanced to support query parameters, enabling endpoints such as list_agents() and delete_agent() to pass optional REST query parameters while reusing the existing request implementation. --- .../snowflake/hooks/snowflake_cortex_agent.py | 114 ++++++++++++- .../hooks/test_snowflake_cortex_agent.py | 158 ++++++++++++++++++ 2 files changed, 265 insertions(+), 7 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 1413d7fee7135..e0c47acfc8c13 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,7 +17,7 @@ from __future__ import annotations -from typing import Any +from typing import Any, cast import requests @@ -54,8 +54,9 @@ def _request( method: str, endpoint: str, payload: dict[str, Any] | None = None, + params: dict[str, Any] | None = None, timeout: int | None = None, - ) -> dict[str, Any]: + ) -> dict[str, Any] | list[dict[str, Any]]: response = requests.request( method=method, @@ -65,6 +66,7 @@ def _request( "Content-Type": "application/json", }, json=payload, + params=params, timeout=timeout, ) @@ -159,11 +161,109 @@ def run_agent( endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}:run" - return self._request( - method="POST", - endpoint=endpoint, - payload=payload, - timeout=timeout, + return cast( + "dict[str, Any]", + self._request( + method="POST", + endpoint=endpoint, + payload=payload, + timeout=timeout, + ), + ) + + def describe_agent( + self, + *, + database: str, + schema: str, + agent_name: str, + ) -> dict[str, Any]: + """ + 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. + :return: JSON description of the Cortex Agent. + """ + endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}" + + return cast( + "dict[str, Any]", + self._request( + method="GET", + endpoint=endpoint, + ), + ) + + def list_agents( + self, + *, + database: str, + schema: str, + like: str | None = None, + from_name: str | None = None, + show_limit: int | None = None, + ) -> list[dict[str, Any]]: + """ + List Snowflake Cortex Agents. + + :param database: Database containing the Cortex Agents. + :param schema: Schema containing the Cortex Agents. + :param like: Optional case-insensitive name filter. + :param from_name: Optional pagination starting point. + :param show_limit: Maximum number of agents to return. + :return: List of Cortex Agents. + """ + endpoint = f"/api/v2/databases/{database}/schemas/{schema}/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 cast( + "list[dict[str, Any]]", + self._request( + method="GET", + endpoint=endpoint, + params=params or None, + ), + ) + + def delete_agent( + self, + *, + database: str, + schema: str, + agent_name: str, + if_exists: bool = False, + ) -> dict[str, Any]: + """ + 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. + Defaults to ``False``. + :return: JSON response confirming deletion. + """ + endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}" + + return cast( + "dict[str, Any]", + self._request( + method="DELETE", + endpoint=endpoint, + params={"ifExists": str(if_exists).lower()}, + ), ) @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 e686b698574ad..ea62c1cb16ccd 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 @@ -127,6 +127,7 @@ def test_run_agent( ], "stream": False, }, + params=None, timeout=REQUEST_TIMEOUT, ) @@ -316,3 +317,160 @@ def test_get_text_response( 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/{DATABASE}" + f"/schemas/{SCHEMA}" + f"/agents/{AGENT_NAME}" + ), + headers={ + "Authorization": f"Bearer {ACCESS_TOKEN}", + "Content-Type": "application/json", + }, + json=None, + params=None, + timeout=None, + ) + + @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/{DATABASE}" + f"/schemas/{SCHEMA}" + f"/agents" + ), + headers={ + "Authorization": f"Bearer {ACCESS_TOKEN}", + "Content-Type": "application/json", + }, + json=None, + params={ + "like": "AIRFLOW%", + "fromName": "AIRFLOW_TEST", + "showLimit": 10, + }, + timeout=None, + ) + + @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/{DATABASE}" + f"/schemas/{SCHEMA}" + f"/agents/{AGENT_NAME}" + ), + headers={ + "Authorization": f"Bearer {ACCESS_TOKEN}", + "Content-Type": "application/json", + }, + json=None, + params={"ifExists": expected}, + timeout=None, + )