This is an automated email from the ASF dual-hosted git repository.
eladkal 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 2e803e58997 Add key-pair JWT authentication to Snowflake Cortex Agent
hook (#73815)
2e803e58997 is described below
commit 2e803e58997ababd9bb19a74acd806a939c564c1
Author: SameerMesiah97 <[email protected]>
AuthorDate: Wed Sep 30 20:02:40 2026 +0100
Add key-pair JWT authentication to Snowflake Cortex Agent hook (#73815)
Generate JWT access tokens from Snowflake connection private keys and set
the
appropriate authorization token type for OAuth and key-pair authentication.
Add unit tests covering both authentication paths and credential failures.
Co-authored-by: Sameer Mesiah <[email protected]>
---
.../snowflake/hooks/snowflake_cortex_agent.py | 42 ++++++--
.../snowflake/hooks/test_snowflake_cortex_agent.py | 106 ++++++++++++++++++++-
2 files changed, 134 insertions(+), 14 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 aa13b081ca1..065b3ca80a3 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
@@ -23,6 +23,7 @@ from urllib.parse import quote
import requests
from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook,
_validate_account_component
+from airflow.providers.snowflake.utils.sql_api_generate_jwt import JWTGenerator
JsonDict = dict[str, Any]
JsonList = list[JsonDict]
@@ -42,17 +43,41 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
account = _validate_account_component(conn_config["account"],
"account")
return f"https://{account}.snowflakecomputing.com"
- def _get_access_token(self) -> str:
+ def _get_auth_headers(self) -> dict[str, str]:
+ """Build authentication headers using OAuth or key-pair
authentication."""
conn_config = self._get_conn_params()
- token = conn_config.get("token")
- if not token:
+ if token := conn_config.get("token"):
+ return {
+ "Authorization": f"Bearer {token}",
+ "Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "OAUTH",
+ }
+
+ account = conn_config.get("account")
+ user = conn_config.get("user")
+ private_key = self.get_private_key()
+
+ if not account or not user or private_key is None:
raise ValueError(
- "Snowflake connection does not provide an OAuth access token. "
- "This hook currently requires an OAuth access token."
+ "Snowflake connection must provide either OAuth credentials or
"
+ "an account, user, and private key for key-pair
authentication."
)
- return token
+ token = JWTGenerator(
+ account=account,
+ user=user,
+ private_key=private_key,
+ ).get_token()
+
+ if token is None:
+ raise RuntimeError("Failed to generate a Snowflake key-pair JWT.")
+
+ return {
+ "Authorization": f"Bearer {token}",
+ "Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "KEYPAIR_JWT",
+ }
@overload
def _request(
@@ -92,10 +117,7 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
response = requests.request(
method=method,
url=f"{self._get_base_url()}{endpoint}",
- headers={
- "Authorization": f"Bearer {self._get_access_token()}",
- "Content-Type": "application/json",
- },
+ headers=self._get_auth_headers(),
json=payload,
params=params,
timeout=timeout,
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 cd454ddba57..efc8e095018 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
@@ -35,6 +35,15 @@ DATABASE = "TEST/DATABASE"
SCHEMA = "TEST?SCHEMA"
AGENT_NAME = "TEST#AGENT"
+USER = "test-user"
+PRIVATE_KEY = mock.sentinel.private_key
+KEYPAIR_TOKEN = "test-keypair-token"
+
+KEYPAIR_CONN_PARAMS = {
+ "account": ACCOUNT,
+ "user": USER,
+}
+
ENCODED_DATABASE = "TEST%2FDATABASE"
ENCODED_SCHEMA = "TEST%3FSCHEMA"
ENCODED_AGENT_NAME = "TEST%23AGENT"
@@ -175,6 +184,7 @@ class TestSnowflakeCortexAgentHook:
headers={
"Authorization": f"Bearer {ACCESS_TOKEN}",
"Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "OAUTH",
},
json={
"messages": [
@@ -312,20 +322,105 @@ class TestSnowflakeCortexAgentHook:
messages=[],
)
+ @mock.patch(f"{MODULE_PATH}.JWTGenerator")
+ @mock.patch(f"{HOOK_PATH}.get_private_key")
+ @mock.patch(f"{HOOK_PATH}._get_conn_params")
+ def test_get_auth_headers_uses_oauth(
+ self,
+ mock_conn_params,
+ mock_get_private_key,
+ mock_jwt_generator,
+ ):
+ mock_conn_params.return_value = CONN_PARAMS
+
+ hook = SnowflakeCortexAgentHook(snowflake_conn_id="mock_conn_id")
+
+ assert hook._get_auth_headers() == {
+ "Authorization": f"Bearer {ACCESS_TOKEN}",
+ "Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "OAUTH",
+ }
+
+ mock_get_private_key.assert_not_called()
+ mock_jwt_generator.assert_not_called()
+
+ @mock.patch(f"{MODULE_PATH}.JWTGenerator")
+ @mock.patch(f"{HOOK_PATH}.get_private_key")
+ @mock.patch(f"{HOOK_PATH}._get_conn_params")
+ def test_get_auth_headers_uses_keypair_jwt(
+ self,
+ mock_conn_params,
+ mock_get_private_key,
+ mock_jwt_generator,
+ ):
+ mock_conn_params.return_value = KEYPAIR_CONN_PARAMS
+ mock_get_private_key.return_value = PRIVATE_KEY
+ mock_jwt_generator.return_value.get_token.return_value = KEYPAIR_TOKEN
+
+ hook = SnowflakeCortexAgentHook(snowflake_conn_id="mock_conn_id")
+
+ assert hook._get_auth_headers() == {
+ "Authorization": f"Bearer {KEYPAIR_TOKEN}",
+ "Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "KEYPAIR_JWT",
+ }
+
+ mock_get_private_key.assert_called_once_with()
+ mock_jwt_generator.assert_called_once_with(
+ account=ACCOUNT,
+ user=USER,
+ private_key=PRIVATE_KEY,
+ )
+ mock_jwt_generator.return_value.get_token.assert_called_once_with()
+
+ @mock.patch(f"{MODULE_PATH}.JWTGenerator")
+ @mock.patch(f"{HOOK_PATH}.get_private_key")
@mock.patch(f"{HOOK_PATH}._get_conn_params")
- def test_get_access_token_raises_when_token_missing(
+ def test_get_auth_headers_raises_when_credentials_missing(
self,
mock_conn_params,
+ mock_get_private_key,
+ mock_jwt_generator,
):
- mock_conn_params.return_value = {}
+ mock_conn_params.return_value = KEYPAIR_CONN_PARAMS
+ mock_get_private_key.return_value = None
hook = SnowflakeCortexAgentHook(snowflake_conn_id="mock_conn_id")
with pytest.raises(
ValueError,
- match="access token",
+ match="Snowflake connection must provide either OAuth credentials
or an account, user, and private key for key-pair authentication.",
+ ):
+ hook._get_auth_headers()
+
+ mock_jwt_generator.assert_not_called()
+
+ @mock.patch(f"{MODULE_PATH}.JWTGenerator")
+ @mock.patch(f"{HOOK_PATH}.get_private_key")
+ @mock.patch(f"{HOOK_PATH}._get_conn_params")
+ def test_get_auth_headers_raises_when_jwt_generation_fails(
+ self,
+ mock_conn_params,
+ mock_get_private_key,
+ mock_jwt_generator,
+ ):
+ mock_conn_params.return_value = KEYPAIR_CONN_PARAMS
+ mock_get_private_key.return_value = PRIVATE_KEY
+ mock_jwt_generator.return_value.get_token.return_value = None
+
+ hook = SnowflakeCortexAgentHook(snowflake_conn_id="mock_conn_id")
+
+ with pytest.raises(
+ RuntimeError,
+ match="Failed to generate a Snowflake key-pair JWT",
):
- hook._get_access_token()
+ hook._get_auth_headers()
+
+ mock_jwt_generator.assert_called_once_with(
+ account=ACCOUNT,
+ user=USER,
+ private_key=PRIVATE_KEY,
+ )
@pytest.mark.parametrize(
("response", "expected"),
@@ -422,6 +517,7 @@ class TestSnowflakeCortexAgentHook:
headers={
"Authorization": f"Bearer {ACCESS_TOKEN}",
"Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "OAUTH",
},
json=None,
params=None,
@@ -471,6 +567,7 @@ class TestSnowflakeCortexAgentHook:
headers={
"Authorization": f"Bearer {ACCESS_TOKEN}",
"Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "OAUTH",
},
json=None,
params={
@@ -532,6 +629,7 @@ class TestSnowflakeCortexAgentHook:
headers={
"Authorization": f"Bearer {ACCESS_TOKEN}",
"Content-Type": "application/json",
+ "X-Snowflake-Authorization-Token-Type": "OAUTH",
},
json=None,
params={"ifExists": expected},