Lee-W commented on code in PR #73932: URL: https://github.com/apache/airflow/pull/73932#discussion_r4177997883
########## 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``. Review Comment: Removed both comments and the module-level copies. Defaults are now changed to `JWTGenerator.LIFETIME` / `JWTGenerator.RENEWAL_DELTA`. -- 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]
