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 @@ -17,7 +17,7 @@

from __future__ import annotations

from typing import Any
from typing import Any, cast

import requests

Expand Down Expand Up @@ -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,
Expand All @@ -65,6 +66,7 @@ def _request(
"Content-Type": "application/json",
},
json=payload,
params=params,
timeout=timeout,
)

Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,7 @@ def test_run_agent(
],
"stream": False,
},
params=None,
timeout=REQUEST_TIMEOUT,
)

Expand Down Expand Up @@ -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,
)