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, + )