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