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},

Reply via email to