Skip to content

Commit 1b8cab7

Browse files
author
Sameer Mesiah
committed
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.
1 parent d369e6f commit 1b8cab7

2 files changed

Lines changed: 265 additions & 7 deletions

File tree

providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py

Lines changed: 107 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717

1818
from __future__ import annotations
1919

20-
from typing import Any
20+
from typing import Any, cast
2121

2222
import requests
2323

@@ -54,8 +54,9 @@ def _request(
5454
method: str,
5555
endpoint: str,
5656
payload: dict[str, Any] | None = None,
57+
params: dict[str, Any] | None = None,
5758
timeout: int | None = None,
58-
) -> dict[str, Any]:
59+
) -> dict[str, Any] | list[dict[str, Any]]:
5960

6061
response = requests.request(
6162
method=method,
@@ -65,6 +66,7 @@ def _request(
6566
"Content-Type": "application/json",
6667
},
6768
json=payload,
69+
params=params,
6870
timeout=timeout,
6971
)
7072

@@ -159,11 +161,109 @@ def run_agent(
159161

160162
endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}:run"
161163

162-
return self._request(
163-
method="POST",
164-
endpoint=endpoint,
165-
payload=payload,
166-
timeout=timeout,
164+
return cast(
165+
"dict[str, Any]",
166+
self._request(
167+
method="POST",
168+
endpoint=endpoint,
169+
payload=payload,
170+
timeout=timeout,
171+
),
172+
)
173+
174+
def describe_agent(
175+
self,
176+
*,
177+
database: str,
178+
schema: str,
179+
agent_name: str,
180+
) -> dict[str, Any]:
181+
"""
182+
Describe a Snowflake Cortex Agent.
183+
184+
:param database: Database containing the Cortex Agent.
185+
:param schema: Schema containing the Cortex Agent.
186+
:param agent_name: Name of the Cortex Agent.
187+
:return: JSON description of the Cortex Agent.
188+
"""
189+
endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}"
190+
191+
return cast(
192+
"dict[str, Any]",
193+
self._request(
194+
method="GET",
195+
endpoint=endpoint,
196+
),
197+
)
198+
199+
def list_agents(
200+
self,
201+
*,
202+
database: str,
203+
schema: str,
204+
like: str | None = None,
205+
from_name: str | None = None,
206+
show_limit: int | None = None,
207+
) -> list[dict[str, Any]]:
208+
"""
209+
List Snowflake Cortex Agents.
210+
211+
:param database: Database containing the Cortex Agents.
212+
:param schema: Schema containing the Cortex Agents.
213+
:param like: Optional case-insensitive name filter.
214+
:param from_name: Optional pagination starting point.
215+
:param show_limit: Maximum number of agents to return.
216+
:return: List of Cortex Agents.
217+
"""
218+
endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents"
219+
220+
params: dict[str, Any] = {}
221+
222+
if like is not None:
223+
params["like"] = like
224+
225+
if from_name is not None:
226+
params["fromName"] = from_name
227+
228+
if show_limit is not None:
229+
params["showLimit"] = show_limit
230+
231+
return cast(
232+
"list[dict[str, Any]]",
233+
self._request(
234+
method="GET",
235+
endpoint=endpoint,
236+
params=params or None,
237+
),
238+
)
239+
240+
def delete_agent(
241+
self,
242+
*,
243+
database: str,
244+
schema: str,
245+
agent_name: str,
246+
if_exists: bool = False,
247+
) -> dict[str, Any]:
248+
"""
249+
Delete a Snowflake Cortex Agent.
250+
251+
:param database: Database containing the Cortex Agent.
252+
:param schema: Schema containing the Cortex Agent.
253+
:param agent_name: Name of the Cortex Agent.
254+
:param if_exists: If ``True``, do not fail when the agent does not exist.
255+
Defaults to ``False``.
256+
:return: JSON response confirming deletion.
257+
"""
258+
endpoint = f"/api/v2/databases/{database}/schemas/{schema}/agents/{agent_name}"
259+
260+
return cast(
261+
"dict[str, Any]",
262+
self._request(
263+
method="DELETE",
264+
endpoint=endpoint,
265+
params={"ifExists": str(if_exists).lower()},
266+
),
167267
)
168268

169269
@staticmethod

providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py

Lines changed: 158 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -127,6 +127,7 @@ def test_run_agent(
127127
],
128128
"stream": False,
129129
},
130+
params=None,
130131
timeout=REQUEST_TIMEOUT,
131132
)
132133

@@ -316,3 +317,160 @@ def test_get_text_response(
316317
expected,
317318
):
318319
assert SnowflakeCortexAgentHook.get_text_response(response) == expected
320+
321+
@mock.patch(f"{MODULE_PATH}.requests.request")
322+
@mock.patch(f"{HOOK_PATH}._get_conn_params")
323+
@mock.patch(
324+
f"{HOOK_PATH}._get_static_conn_params",
325+
new_callable=mock.PropertyMock,
326+
)
327+
def test_describe_agent(
328+
self,
329+
mock_static_conn_params,
330+
mock_conn_params,
331+
mock_request,
332+
):
333+
mock_conn_params.return_value = CONN_PARAMS
334+
mock_static_conn_params.return_value = STATIC_CONN_PARAMS
335+
mock_request.return_value = create_response(
336+
json_body={"name": AGENT_NAME},
337+
)
338+
339+
hook = SnowflakeCortexAgentHook(
340+
snowflake_conn_id="mock_conn_id",
341+
)
342+
343+
result = hook.describe_agent(
344+
database=DATABASE,
345+
schema=SCHEMA,
346+
agent_name=AGENT_NAME,
347+
)
348+
349+
assert result == {"name": AGENT_NAME}
350+
351+
mock_request.assert_called_once_with(
352+
method="GET",
353+
url=(
354+
f"https://{ACCOUNT}.snowflakecomputing.com"
355+
f"/api/v2/databases/{DATABASE}"
356+
f"/schemas/{SCHEMA}"
357+
f"/agents/{AGENT_NAME}"
358+
),
359+
headers={
360+
"Authorization": f"Bearer {ACCESS_TOKEN}",
361+
"Content-Type": "application/json",
362+
},
363+
json=None,
364+
params=None,
365+
timeout=None,
366+
)
367+
368+
@mock.patch(f"{MODULE_PATH}.requests.request")
369+
@mock.patch(f"{HOOK_PATH}._get_conn_params")
370+
@mock.patch(
371+
f"{HOOK_PATH}._get_static_conn_params",
372+
new_callable=mock.PropertyMock,
373+
)
374+
def test_list_agents(
375+
self,
376+
mock_static_conn_params,
377+
mock_conn_params,
378+
mock_request,
379+
):
380+
mock_conn_params.return_value = CONN_PARAMS
381+
mock_static_conn_params.return_value = STATIC_CONN_PARAMS
382+
mock_request.return_value = create_response(
383+
json_body=[{"name": AGENT_NAME}],
384+
)
385+
386+
hook = SnowflakeCortexAgentHook(
387+
snowflake_conn_id="mock_conn_id",
388+
)
389+
390+
result = hook.list_agents(
391+
database=DATABASE,
392+
schema=SCHEMA,
393+
like="AIRFLOW%",
394+
from_name="AIRFLOW_TEST",
395+
show_limit=10,
396+
)
397+
398+
assert result == [{"name": AGENT_NAME}]
399+
400+
mock_request.assert_called_once_with(
401+
method="GET",
402+
url=(
403+
f"https://{ACCOUNT}.snowflakecomputing.com"
404+
f"/api/v2/databases/{DATABASE}"
405+
f"/schemas/{SCHEMA}"
406+
f"/agents"
407+
),
408+
headers={
409+
"Authorization": f"Bearer {ACCESS_TOKEN}",
410+
"Content-Type": "application/json",
411+
},
412+
json=None,
413+
params={
414+
"like": "AIRFLOW%",
415+
"fromName": "AIRFLOW_TEST",
416+
"showLimit": 10,
417+
},
418+
timeout=None,
419+
)
420+
421+
@pytest.mark.parametrize(
422+
("if_exists", "expected"),
423+
[
424+
pytest.param(True, "true", id="if_exists"),
425+
pytest.param(False, "false", id="error_if_missing"),
426+
],
427+
)
428+
@mock.patch(f"{MODULE_PATH}.requests.request")
429+
@mock.patch(f"{HOOK_PATH}._get_conn_params")
430+
@mock.patch(
431+
f"{HOOK_PATH}._get_static_conn_params",
432+
new_callable=mock.PropertyMock,
433+
)
434+
def test_delete_agent(
435+
self,
436+
mock_static_conn_params,
437+
mock_conn_params,
438+
mock_request,
439+
if_exists,
440+
expected,
441+
):
442+
mock_conn_params.return_value = CONN_PARAMS
443+
mock_static_conn_params.return_value = STATIC_CONN_PARAMS
444+
mock_request.return_value = create_response(
445+
json_body={"status": "deleted"},
446+
)
447+
448+
hook = SnowflakeCortexAgentHook(
449+
snowflake_conn_id="mock_conn_id",
450+
)
451+
452+
result = hook.delete_agent(
453+
database=DATABASE,
454+
schema=SCHEMA,
455+
agent_name=AGENT_NAME,
456+
if_exists=if_exists,
457+
)
458+
459+
assert result == {"status": "deleted"}
460+
461+
mock_request.assert_called_once_with(
462+
method="DELETE",
463+
url=(
464+
f"https://{ACCOUNT}.snowflakecomputing.com"
465+
f"/api/v2/databases/{DATABASE}"
466+
f"/schemas/{SCHEMA}"
467+
f"/agents/{AGENT_NAME}"
468+
),
469+
headers={
470+
"Authorization": f"Bearer {ACCESS_TOKEN}",
471+
"Content-Type": "application/json",
472+
},
473+
json=None,
474+
params={"ifExists": expected},
475+
timeout=None,
476+
)

0 commit comments

Comments
 (0)