Lee-W commented on code in PR #73932:
URL: https://github.com/apache/airflow/pull/73932#discussion_r4177990404


##########
providers/snowflake/src/airflow/providers/snowflake/utils/rest_auth.py:
##########
@@ -0,0 +1,180 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""
+Shared Snowflake REST API authentication (OAuth, PAT, key-pair JWT).
+
+Every Snowflake REST caller (the SQL API, Cortex Agents, the Cortex 
chat-completions
+endpoint used by pydantic-ai) authenticates the same three ways: an OAuth 
access token,
+a Programmatic Access Token (PAT), or a JWT signed with the connection's 
private key.
+:class:`SnowflakeRestTokenProvider` produces those headers from a 
:class:`SnowflakeHook`
+so each caller does not need to duplicate the branching or the token caching.
+
+This module intentionally imports nothing from ``common.ai``, ``pydantic-ai``, 
or
+``httpx2`` -- it is plain Snowflake REST auth and must stay usable by callers 
that never
+touch those optional dependencies.
+"""
+
+from __future__ import annotations
+
+import threading
+from collections.abc import Callable
+from dataclasses import dataclass, field
+from datetime import timedelta
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.snowflake.hooks.snowflake import 
_validate_account_component
+from airflow.providers.snowflake.utils.sql_api_generate_jwt import JWTGenerator
+
+if TYPE_CHECKING:
+    from cryptography.hazmat.primitives.asymmetric.types import PrivateKeyTypes
+
+    from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook
+
+LIFETIME = timedelta(minutes=59)  # The tokens will have a 59 minute lifetime
+RENEWAL_DELTA = timedelta(minutes=54)  # Tokens will be renewed after 54 
minutes
+
+
+@dataclass(frozen=True)
+class SnowflakeRestToken:
+    """A REST bearer token and the ``X-Snowflake-Authorization-Token-Type`` 
value it needs."""
+
+    token: str = field(repr=False)
+    token_type: str
+
+
+class SnowflakeRestTokenProvider:
+    """
+    Produce Snowflake REST auth headers from a :class:`SnowflakeHook`, caching 
what it can.
+
+    The branch taken mirrors ``SnowflakeSqlApiHook.get_headers``: 
``authenticator == "oauth"``
+    reads the token ``hook._get_conn_params()`` already resolved (which itself 
refreshes an
+    expiring OAuth or Azure token, so this provider does not cache that branch 
at all);
+    ``authenticator == "programmatic_access_token"`` reads the PAT from the 
connection password;
+    anything else signs a key-pair JWT. The private key is loaded once and 
kept. The
+    ``JWTGenerator`` is also built once and kept -- it renews its own token 
internally, so
+    creating a fresh one on every call (as ``get_headers`` used to) defeated 
``token_renewal_delta``.
+
+    :param hook: The ``SnowflakeHook`` (or subclass) whose connection supplies 
credentials.
+    :param token_life_time: Passed to the ``JWTGenerator`` for the key-pair 
branch.
+    :param token_renewal_delta: Passed to the ``JWTGenerator`` for the 
key-pair branch. When this
+        is greater than or equal to ``token_life_time``, the JWT is renewed on 
every call instead
+        -- otherwise the scheduled renewal would land after the token has 
already expired, and
+        every call in between would serve a token Snowflake rejects.
+    :param private_key_loader: Callable returning the private key for the 
key-pair branch,
+        called at most once. Defaults to ``hook.get_private_key``. A caller 
that already
+        maintains its own ``private_key`` attribute (e.g. 
``SnowflakeSqlApiHook``) can pass a
+        loader that populates it, so that attribute keeps working for existing 
callers.
+    """
+
+    def __init__(
+        self,
+        hook: SnowflakeHook,
+        *,
+        token_life_time: timedelta = LIFETIME,
+        token_renewal_delta: timedelta = RENEWAL_DELTA,
+        private_key_loader: Callable[[], PrivateKeyTypes | None] | None = None,
+    ) -> None:
+        self._hook = hook
+        self._token_life_time = token_life_time
+        # A JWTGenerator only regenerates once `renew_time` (set to `now + 
renewal_delay` at
+        # generation time) has passed. If `renewal_delay >= lifetime`, the 
*next* renewal would
+        # be scheduled for after the token has already expired -- serving a 
rejected token for
+        # the gap between expiry and the (too-late) renewal. Renewing on every 
call (delay=0)
+        # matches what the old per-call `JWTGenerator` construction always did 
for this case.
+        self._token_renewal_delta = (
+            timedelta(0) if token_renewal_delta >= token_life_time else 
token_renewal_delta
+        )
+        self._private_key_loader: Callable[[], PrivateKeyTypes | None] = (
+            private_key_loader or hook.get_private_key
+        )
+        self._private_key: PrivateKeyTypes | None = None
+        self._jwt_generator: JWTGenerator | None = None
+        self._lock = threading.Lock()
+
+    def get_token(self) -> SnowflakeRestToken:
+        """Return the current REST bearer token, refreshing or renewing it as 
needed."""
+        with self._lock:
+            conn_config = self._hook._get_conn_params()
+
+            if conn_config.get("authenticator") == "oauth":
+                token = conn_config.get("token")
+                if not token:
+                    raise ValueError("OAuth authentication did not produce an 
access token.")
+                return SnowflakeRestToken(token=token, token_type="OAUTH")
+
+            if conn_config.get("authenticator") == "programmatic_access_token":
+                pat = conn_config.get("password")
+                if not pat:
+                    raise ValueError(
+                        "Programmatic Access Token (PAT) authentication 
requires the connection "
+                        "password field to contain the PAT token value."
+                    )
+                return SnowflakeRestToken(token=pat, 
token_type="PROGRAMMATIC_ACCESS_TOKEN")
+
+            if conn_config.get("workload_identity_provider"):
+                raise ValueError(
+                    "Workload identity federation is not supported for 
Snowflake REST APIs; use "
+                    "OAuth, PAT, or key-pair authentication instead."
+                )
+
+            if self._private_key is None:
+                self._private_key = self._private_key_loader()
+            if self._private_key is None:
+                raise ValueError(
+                    "Snowflake REST API authentication requires an OAuth 
access token, a "
+                    "Programmatic Access Token (PAT), or a private key for 
key-pair JWT auth; "
+                    "none is configured on this connection."
+                )
+
+            if self._jwt_generator is None:
+                self._jwt_generator = JWTGenerator(
+                    conn_config["account"],  # type: ignore[arg-type]

Review Comment:
   Brought the check back. key-pair auth now raises a `ValueError` naming 
whether the account, the login, or both are missing. Removed `type: ignore`



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to