kaxil commented on code in PR #73932: URL: https://github.com/apache/airflow/pull/73932#discussion_r4156674663
########## providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_model.py: ########## @@ -0,0 +1,182 @@ +# 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. +"""Expose Snowflake Cortex's OpenAI-compatible chat endpoint as a pydantic-ai model hook.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.providers.snowflake.utils.rest_auth import SnowflakeRestTokenProvider, get_cortex_base_url + +try: + import httpx2 + from pydantic_ai.providers import snowflake as _pydantic_ai_snowflake_provider # noqa: F401 + + from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook +except ImportError: + raise AirflowOptionalProviderFeatureException( + "This feature requires the 'common.ai' provider, in a version that ships a Snowflake " + "pydantic-ai provider. Install with apache-airflow-providers-snowflake[common.ai]." + ) + +if TYPE_CHECKING: + from httpx2 import Request + +CORTEX_CHAT_COMPLETIONS_PATH = "/api/v2/cortex/v1" + + +class _SnowflakeCortexAuth(httpx2.Auth): + """ + Refresh the ``Authorization`` header on every request from a shared token provider. + + ``build_auth_headers()`` may block: it can call ``requests.post`` with retries for an + expiring OAuth token, or -- for ``azure_conn_id`` -- resolve an Airflow connection and fetch + an Azure token on every call. Resolving a connection synchronously from the event-loop + thread while an async send is in flight raises ``DeadlockImminentError`` (see + ``task-sdk`` ``execution_time/comms.py``), so ``async_auth_flow`` is overridden to run the + refresh in a worker thread instead of the httpx2 default of driving the sync ``auth_flow`` + inline on the loop. ``auth_flow`` itself is kept for sync ``httpx2.Client`` callers, which + have no event loop to block. + """ + + def __init__(self, token_provider: SnowflakeRestTokenProvider) -> None: + self._token_provider = token_provider + + def auth_flow(self, request: Request) -> Any: + request.headers.update(self._token_provider.build_auth_headers()) + yield request + + async def async_auth_flow(self, request: Request) -> Any: + request.headers.update(await asyncio.to_thread(self._token_provider.build_auth_headers)) + yield request + + +class PydanticAISnowflakeHook(PydanticAIHook): + """ + Hook for Snowflake Cortex's OpenAI-compatible chat endpoint via pydantic-ai. + + Unlike the other ``PydanticAI*`` hooks, credentials do not live on this connection: they are + read from an existing ``snowflake`` connection (OAuth, PAT, or key-pair JWT -- whichever that + connection is configured for), refreshed on every request the same way as + ``SnowflakeCortexAgentHook`` and ``SnowflakeSqlApiHook``. See + :class:`~airflow.providers.snowflake.utils.rest_auth.SnowflakeRestTokenProvider`. The + underlying ``httpx2.AsyncClient`` is built once and lives as long as this hook instance; + nothing currently closes it (``SnowflakeProvider`` only owns and closes a client it built + itself, not one passed in). + + Connection fields: + - **extra** JSON: ``{"model": "snowflake:claude-4-sonnet", + "snowflake_conn_id": "snowflake_default"}`` + + Model family support (pydantic-ai-slim's ``SnowflakeProvider.model_profile``): Claude + (``claude*``) and OpenAI (``openai-*``) models support tools and structured output; + other families (``llama*``, ``snowflake-llama*``, ``mistral*``, ``mixtral*``, + ``deepseek*``, and any unlisted family) do not support tools, and structured output + falls back to prompted mode. Use a Claude or OpenAI family model for a tool-using agent. + + :param llm_conn_id: Airflow connection ID for this ``pydanticai_snowflake`` connection. + :param model_id: Model identifier, e.g. ``"snowflake:claude-4-sonnet"``. A bare name (no + recognized platform prefix) is qualified with ``snowflake:``. + :param fallback_conn_ids: See :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`. + :param snowflake_conn_id: Connection ID of an existing Snowflake connection to source + credentials, account, and host from. Takes precedence over the connection extra's + ``snowflake_conn_id``; one of the two is required. + """ + + conn_type = "pydanticai_snowflake" + default_conn_name = "pydanticai_snowflake_default" + hook_name = "Pydantic AI (Snowflake Cortex)" + model_provider = "snowflake" + + def __init__( + self, + llm_conn_id: str | None = None, + model_id: str | None = None, + fallback_conn_ids: list[str] | None = None, + *, + snowflake_conn_id: str | None = None, + **kwargs: Any, + ) -> None: + super().__init__(llm_conn_id, model_id, fallback_conn_ids, **kwargs) + self.snowflake_conn_id = snowflake_conn_id + self._token_provider: SnowflakeRestTokenProvider | None = None + self._cortex_base_url: str | None = None + self._http_client: httpx2.AsyncClient | None = None + + @staticmethod + def get_ui_field_behaviour() -> dict[str, Any]: + """Return custom field behaviour for the Airflow connection form.""" + return { + "hidden_fields": ["schema", "port", "login", "host", "password"], + "relabeling": {}, + "placeholders": { + "extra": '{"model": "snowflake:claude-4-sonnet", "snowflake_conn_id": "snowflake_default"}', + }, + } + + def _get_snowflake_conn_id(self, extra: dict[str, Any]) -> str: + snowflake_conn_id = self.snowflake_conn_id or extra.get("snowflake_conn_id") + if not snowflake_conn_id: + raise ValueError( + f"Connection '{self.llm_conn_id}' has no Snowflake connection to source credentials " + "from. Set snowflake_conn_id on the hook or the connection's extra field, pointing " + "at an existing Snowflake connection." + ) + return snowflake_conn_id + + def _get_token_provider(self, extra: dict[str, Any]) -> SnowflakeRestTokenProvider: + """ + Build the Snowflake hook, token provider, base URL, and HTTP client once. + + Reused for this hook's lifetime -- including the ``httpx2.AsyncClient``, which nothing + else owns or closes (``SnowflakeProvider`` only owns and closes a client it built itself), + so building a fresh one on every call would leak one per call. + """ + if self._token_provider is None: + snowflake_hook = SnowflakeHook(snowflake_conn_id=self._get_snowflake_conn_id(extra)) + self._token_provider = SnowflakeRestTokenProvider(snowflake_hook) + self._cortex_base_url = ( + get_cortex_base_url(snowflake_hook._get_static_conn_params) + CORTEX_CHAT_COMPLETIONS_PATH + ) + self._http_client = httpx2.AsyncClient(auth=_SnowflakeCortexAuth(self._token_provider)) + return self._token_provider + + def _get_provider_kwargs( + self, + api_key: str | None, + base_url: str | None, + extra: dict[str, Any], + ) -> dict[str, Any]: + """ + Return kwargs for ``SnowflakeProvider``. + + .. note:: + ``api_key`` and ``base_url`` (sourced from ``conn.password`` and ``conn.host``) are + intentionally ignored: this connection hides those fields in the UI, and credentials + and host come from the Snowflake connection named by ``snowflake_conn_id`` instead. + """ + token_provider = self._get_token_provider(extra) + return { + "base_url": self._cortex_base_url, + # SnowflakeProvider requires a non-empty token at construction time (and uses it for + # test_connection()); the real per-request token comes from _SnowflakeCortexAuth below. + "token": token_provider.get_token().token, Review Comment: This makes a credential fetch part of building the model, which breaks fallback isolation. `PydanticAIHook.get_conn()` builds every fallback eagerly in `_resolve_fallback_models`, so with this connection as a fallback on OAuth or `azure_conn_id`, a Snowflake token-endpoint outage fails the task before the healthy primary is ever called. As the primary, a failed fetch raises outside `FallbackModel`'s `ModelAPIError` path, so the chain never reaches its fallback either. `_SnowflakeCortexAuth` overwrites `Authorization` on every request, so this token is never sent. Could this pass a placeholder instead, and override `test_connection` to call `get_token()` so it still proves the credentials? ########## providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_model.py: ########## @@ -0,0 +1,182 @@ +# 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. +"""Expose Snowflake Cortex's OpenAI-compatible chat endpoint as a pydantic-ai model hook.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.providers.snowflake.utils.rest_auth import SnowflakeRestTokenProvider, get_cortex_base_url + +try: + import httpx2 + from pydantic_ai.providers import snowflake as _pydantic_ai_snowflake_provider # noqa: F401 + + from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook +except ImportError: + raise AirflowOptionalProviderFeatureException( + "This feature requires the 'common.ai' provider, in a version that ships a Snowflake " + "pydantic-ai provider. Install with apache-airflow-providers-snowflake[common.ai]." + ) + +if TYPE_CHECKING: + from httpx2 import Request + +CORTEX_CHAT_COMPLETIONS_PATH = "/api/v2/cortex/v1" + + +class _SnowflakeCortexAuth(httpx2.Auth): + """ + Refresh the ``Authorization`` header on every request from a shared token provider. + + ``build_auth_headers()`` may block: it can call ``requests.post`` with retries for an + expiring OAuth token, or -- for ``azure_conn_id`` -- resolve an Airflow connection and fetch + an Azure token on every call. Resolving a connection synchronously from the event-loop + thread while an async send is in flight raises ``DeadlockImminentError`` (see + ``task-sdk`` ``execution_time/comms.py``), so ``async_auth_flow`` is overridden to run the + refresh in a worker thread instead of the httpx2 default of driving the sync ``auth_flow`` + inline on the loop. ``auth_flow`` itself is kept for sync ``httpx2.Client`` callers, which + have no event loop to block. + """ + + def __init__(self, token_provider: SnowflakeRestTokenProvider) -> None: + self._token_provider = token_provider + + def auth_flow(self, request: Request) -> Any: + request.headers.update(self._token_provider.build_auth_headers()) + yield request + + async def async_auth_flow(self, request: Request) -> Any: + request.headers.update(await asyncio.to_thread(self._token_provider.build_auth_headers)) + yield request + + +class PydanticAISnowflakeHook(PydanticAIHook): Review Comment: Same point eladkal raised on `snowflake_cortex_agent.py` in #73815: this module repeats the `snowflake_` prefix that AIP-21 removed inside provider packages. Nothing imports it yet, so `hooks/cortex_model.py` (or similar) is free now and needs a deprecation cycle after release. ########## 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: "(as `get_headers` used to)" and the comment at 93-97 describe the code this replaced, which won't mean anything to a reader after merge. Lines 47-48 also add a fifth copy of `LIFETIME` / `RENEWAL_DELTA`; `JWTGenerator.LIFETIME` and `JWTGenerator.RENEWAL_DELTA` are already imported here. ########## 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 Review Comment: Since `utils/rest_auth.py` is a new public module, should `token_type` be `Literal["OAUTH", "PROGRAMMATIC_ACCESS_TOKEN", "KEYPAIR_JWT"]`? If nothing outside this provider is meant to import it, a `_rest_auth.py` name would keep it off the public surface entirely. ########## providers/snowflake/docs/connections/pydantic_ai_snowflake.rst: ########## @@ -0,0 +1,131 @@ + .. 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. + +.. _howto/connection:pydanticai_snowflake: + +Pydantic AI (Snowflake Cortex) connection +========================================== + +The ``pydanticai_snowflake`` connection type configures access to +`Snowflake Cortex <https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api>`__'s +OpenAI-compatible chat endpoint via the pydantic-ai framework. It backs +``PydanticAISnowflakeHook``, the dedicated subclass of ``PydanticAIHook`` for Snowflake Cortex. + +Unlike using a plain ``pydanticai`` connection with a ``snowflake:`` model and pydantic-ai's own +``SNOWFLAKE_ACCOUNT`` / ``SNOWFLAKE_TOKEN`` environment variables, this connection type sources +credentials from an existing :ref:`Snowflake connection <howto/connection:snowflake>` instead: the +credential lives in one place, goes through Airflow's secrets backend, can be rotated in one +place, and -- for key-pair JWT authentication -- is refreshed automatically as the token nears +expiry, which a static environment variable cannot do. + +Install with: + +.. code-block:: bash + + pip install 'apache-airflow-providers-snowflake[common.ai]' + +Default Connection IDs +---------------------- + +The ``PydanticAISnowflakeHook`` uses ``pydanticai_snowflake_default`` by default. + +Configuring the Connection +-------------------------- + +This connection type needs two connections: this one, and an existing ``snowflake`` connection it +points at. All fields below are ``extra`` (JSON) fields on the ``pydanticai_snowflake`` connection; +``Schema``, ``Port``, ``Login``, ``Host``, and ``Password`` are hidden in the connection form +because credentials and host come from the Snowflake connection instead. + +Model + Cortex model identifier (e.g. ``snowflake:claude-4-sonnet``). + + A bare name is automatically resolved to ``snowflake:<name>`` -- Snowflake Cortex is this + connection type's own platform. + +Fallback Connections + Other connection IDs to fail over to, in order, while this provider is unavailable. Stored in + ``extra["fallback_conn_ids"]``. Entries may name any ``pydanticai`` connection type, so one + chain can span vendors. See :doc:`apache-airflow-providers-common-ai:provider_fallback`. + +Snowflake Connection ID + Connection ID of an existing :ref:`Snowflake connection <howto/connection:snowflake>` to + source credentials, account, and host from. Stored in ``extra["snowflake_conn_id"]``. Also + settable via the ``PydanticAISnowflakeHook(snowflake_conn_id=...)`` constructor argument, which + takes precedence over the extra field. One of the two is required. + +Authentication +-------------- + +Credentials are not configured on this connection -- they come from whichever authentication +method the referenced Snowflake connection uses, set the same way as for +:ref:`SnowflakeSqlApiHook and SnowflakeCortexAgentHook <howto/connection:snowflake>`: + +- **OAuth**: set ``authenticator`` to ``oauth`` in the Snowflake connection's extra, and configure + a refresh token, a client credentials grant, or ``azure_conn_id``. +- **PAT (Programmatic Access Token)**: set ``authenticator`` to ``programmatic_access_token`` and + put the PAT value in the Snowflake connection's Password field. +- **Key-pair JWT**: the default when neither of the above is set. Configure Review Comment: Key-pair also needs Login on the Snowflake connection, since it becomes the JWT's `sub`. Worth saying here, as this is the page people will follow. Separately, line 110 cites pydantic-ai-slim 2.44.0 while the floor is 2.33.0. ########## providers/common/ai/docs/model_providers.rst: ########## @@ -89,9 +89,11 @@ type shown, and set the model name with that prefix. * - Snowflake Cortex - ``snowflake:`` - ``pydantic-ai-slim[snowflake]`` - - ``pydanticai`` + - ``pydanticai``, or ``pydanticai_snowflake`` Review Comment: The `pydanticai_snowflake` type comes from `apache-airflow-providers-snowflake[common.ai]`, not `pydantic-ai-slim[snowflake]`. Could the Install cell mention it, so someone picking that connection type knows what to install? ########## providers/snowflake/docs/connections/pydantic_ai_snowflake.rst: ########## @@ -0,0 +1,131 @@ + .. 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. + +.. _howto/connection:pydanticai_snowflake: + +Pydantic AI (Snowflake Cortex) connection +========================================== + +The ``pydanticai_snowflake`` connection type configures access to +`Snowflake Cortex <https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api>`__'s +OpenAI-compatible chat endpoint via the pydantic-ai framework. It backs +``PydanticAISnowflakeHook``, the dedicated subclass of ``PydanticAIHook`` for Snowflake Cortex. + +Unlike using a plain ``pydanticai`` connection with a ``snowflake:`` model and pydantic-ai's own +``SNOWFLAKE_ACCOUNT`` / ``SNOWFLAKE_TOKEN`` environment variables, this connection type sources +credentials from an existing :ref:`Snowflake connection <howto/connection:snowflake>` instead: the +credential lives in one place, goes through Airflow's secrets backend, can be rotated in one +place, and -- for key-pair JWT authentication -- is refreshed automatically as the token nears +expiry, which a static environment variable cannot do. + +Install with: + +.. code-block:: bash + + pip install 'apache-airflow-providers-snowflake[common.ai]' + +Default Connection IDs +---------------------- + +The ``PydanticAISnowflakeHook`` uses ``pydanticai_snowflake_default`` by default. + +Configuring the Connection +-------------------------- + +This connection type needs two connections: this one, and an existing ``snowflake`` connection it +points at. All fields below are ``extra`` (JSON) fields on the ``pydanticai_snowflake`` connection; +``Schema``, ``Port``, ``Login``, ``Host``, and ``Password`` are hidden in the connection form +because credentials and host come from the Snowflake connection instead. + +Model + Cortex model identifier (e.g. ``snowflake:claude-4-sonnet``). + + A bare name is automatically resolved to ``snowflake:<name>`` -- Snowflake Cortex is this + connection type's own platform. + +Fallback Connections + Other connection IDs to fail over to, in order, while this provider is unavailable. Stored in + ``extra["fallback_conn_ids"]``. Entries may name any ``pydanticai`` connection type, so one + chain can span vendors. See :doc:`apache-airflow-providers-common-ai:provider_fallback`. + +Snowflake Connection ID + Connection ID of an existing :ref:`Snowflake connection <howto/connection:snowflake>` to + source credentials, account, and host from. Stored in ``extra["snowflake_conn_id"]``. Also + settable via the ``PydanticAISnowflakeHook(snowflake_conn_id=...)`` constructor argument, which + takes precedence over the extra field. One of the two is required. + +Authentication +-------------- + +Credentials are not configured on this connection -- they come from whichever authentication +method the referenced Snowflake connection uses, set the same way as for +:ref:`SnowflakeSqlApiHook and SnowflakeCortexAgentHook <howto/connection:snowflake>`: + +- **OAuth**: set ``authenticator`` to ``oauth`` in the Snowflake connection's extra, and configure + a refresh token, a client credentials grant, or ``azure_conn_id``. +- **PAT (Programmatic Access Token)**: set ``authenticator`` to ``programmatic_access_token`` and Review Comment: This matches what the code requires, as does the same bullet in `operators/snowflake_cortex_agent.rst`. But the page linked below as the full field reference says the opposite: `connections/snowflake.rst` lines 63-64 and its PAT example say no special authenticator is needed. Someone following that page gets the "none is configured on this connection" error with a PAT sitting in Password. The SQL API hook already had this mismatch, so fixing the `authenticator` bullet and PAT example in `snowflake.rst` here would cover all three hooks. Can one connection with `authenticator: programmatic_access_token` also serve `SnowflakeHook`? If not, the "credential lives in one place" point at the top of this page doesn't hold for PAT. ########## 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"): Review Comment: In 6.18.0 `SnowflakeSqlApiHook.get_headers` had no workload-identity check, so a connection with `workload_identity_provider` and a private key got a key-pair JWT. Since this guard runs before the key is loaded, that connection now raises. Would it work to load the key first and raise this only when there is no key? ########## 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: #73815 raised a `ValueError` here when `account` or `user` was missing, and that check didn't survive the move. Without it, a connection with no Login fails with `AttributeError: 'NoneType' object has no attribute 'upper'` inside `JWTGenerator`, and an empty `account` with `host` set signs a JWT whose `sub` is `.USER`, which Snowflake rejects with a bare 401. Could the check come back? It would also let both `type: ignore`s go. ########## providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_model.py: ########## @@ -0,0 +1,182 @@ +# 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. +"""Expose Snowflake Cortex's OpenAI-compatible chat endpoint as a pydantic-ai model hook.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.providers.snowflake.utils.rest_auth import SnowflakeRestTokenProvider, get_cortex_base_url + +try: + import httpx2 + from pydantic_ai.providers import snowflake as _pydantic_ai_snowflake_provider # noqa: F401 + + from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook +except ImportError: + raise AirflowOptionalProviderFeatureException( + "This feature requires the 'common.ai' provider, in a version that ships a Snowflake " + "pydantic-ai provider. Install with apache-airflow-providers-snowflake[common.ai]." + ) + +if TYPE_CHECKING: + from httpx2 import Request + +CORTEX_CHAT_COMPLETIONS_PATH = "/api/v2/cortex/v1" + + +class _SnowflakeCortexAuth(httpx2.Auth): + """ + Refresh the ``Authorization`` header on every request from a shared token provider. + + ``build_auth_headers()`` may block: it can call ``requests.post`` with retries for an + expiring OAuth token, or -- for ``azure_conn_id`` -- resolve an Airflow connection and fetch + an Azure token on every call. Resolving a connection synchronously from the event-loop + thread while an async send is in flight raises ``DeadlockImminentError`` (see + ``task-sdk`` ``execution_time/comms.py``), so ``async_auth_flow`` is overridden to run the + refresh in a worker thread instead of the httpx2 default of driving the sync ``auth_flow`` + inline on the loop. ``auth_flow`` itself is kept for sync ``httpx2.Client`` callers, which + have no event loop to block. + """ + + def __init__(self, token_provider: SnowflakeRestTokenProvider) -> None: + self._token_provider = token_provider + + def auth_flow(self, request: Request) -> Any: Review Comment: httpx2 declares these as `Generator[Request, Response, None]` and `AsyncGenerator[Request, Response]`. Using those instead of `Any` would let mypy check the overrides. ########## providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_sql_api.py: ########## @@ -561,6 +561,43 @@ def test_get_headers_pat_raises_when_password_missing(self, mock_conn_param): with pytest.raises(ValueError, match="Programmatic Access Token"): hook.get_headers() + @mock.patch(f"{HOOK_PATH}.get_private_key", autospec=True) + @mock.patch(f"{HOOK_PATH}._get_conn_params", autospec=True) + def test_get_headers_reuses_jwt_within_renewal_window( Review Comment: Every test here uses the default lifetimes, so deleting the `token_life_time` / `token_renewal_delta` forwarding in `get_headers` leaves the suite green. A case with a custom `token_renewal_delta`, plus one at or above `token_life_time` (the re-sign-every-call path from the description), would pin both. ########## providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_sql_api.py: ########## @@ -230,54 +231,26 @@ def execute_query( def get_headers(self) -> dict[str, Any]: """Form auth headers based on OAuth token, PAT, or JWT token from private key.""" - conn_config = self._get_conn_params() - - # _get_conn_params() already fetched the OAuth access token for any grant type or azure_conn_id. - if conn_config.get("authenticator") == "oauth": - return { - "Content-Type": "application/json", - "Authorization": f"Bearer {conn_config['token']}", - "Accept": "application/json", - "User-Agent": "snowflakeSQLAPI/1.0", - "X-Snowflake-Authorization-Token-Type": "OAUTH", - } - - # Use PAT (Programmatic Access Token) when authenticator is set to programmatic_access_token - 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 { - "Content-Type": "application/json", - "Authorization": f"Bearer {pat}", - "Accept": "application/json", - "User-Agent": "snowflakeSQLAPI/1.0", - "X-Snowflake-Authorization-Token-Type": "PROGRAMMATIC_ACCESS_TOKEN", - } - - # Fall back to JWT token from the connection details and the private key - if not self.private_key: - self.private_key = self.get_private_key() - - token = JWTGenerator( - conn_config["account"], # type: ignore[arg-type] - conn_config["user"], # type: ignore[arg-type] - private_key=self.private_key, - lifetime=self.token_life_time, - renewal_delay=self.token_renewal_delta, - ).get_token() - + if self._rest_token_provider is None: Review Comment: `get_result_from_successful_sql_api_query` builds one header and reuses it for every partition GET. Before, that header always carried a fresh 59-minute JWT. With the cached generator it can carry as little as `lifetime - renewal_delta` (5 minutes by default), so a large partitioned result fetched late in the window can start getting 401s partway through. Calling `get_headers()` per request in that loop would keep the old guarantee. ########## providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_model.py: ########## @@ -0,0 +1,182 @@ +# 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. +"""Expose Snowflake Cortex's OpenAI-compatible chat endpoint as a pydantic-ai model hook.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.providers.snowflake.utils.rest_auth import SnowflakeRestTokenProvider, get_cortex_base_url + +try: + import httpx2 + from pydantic_ai.providers import snowflake as _pydantic_ai_snowflake_provider # noqa: F401 + + from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook +except ImportError: + raise AirflowOptionalProviderFeatureException( + "This feature requires the 'common.ai' provider, in a version that ships a Snowflake " + "pydantic-ai provider. Install with apache-airflow-providers-snowflake[common.ai]." + ) + +if TYPE_CHECKING: + from httpx2 import Request + +CORTEX_CHAT_COMPLETIONS_PATH = "/api/v2/cortex/v1" + + +class _SnowflakeCortexAuth(httpx2.Auth): + """ + Refresh the ``Authorization`` header on every request from a shared token provider. + + ``build_auth_headers()`` may block: it can call ``requests.post`` with retries for an + expiring OAuth token, or -- for ``azure_conn_id`` -- resolve an Airflow connection and fetch + an Azure token on every call. Resolving a connection synchronously from the event-loop + thread while an async send is in flight raises ``DeadlockImminentError`` (see + ``task-sdk`` ``execution_time/comms.py``), so ``async_auth_flow`` is overridden to run the + refresh in a worker thread instead of the httpx2 default of driving the sync ``auth_flow`` + inline on the loop. ``auth_flow`` itself is kept for sync ``httpx2.Client`` callers, which + have no event loop to block. + """ + + def __init__(self, token_provider: SnowflakeRestTokenProvider) -> None: + self._token_provider = token_provider + + def auth_flow(self, request: Request) -> Any: + request.headers.update(self._token_provider.build_auth_headers()) + yield request + + async def async_auth_flow(self, request: Request) -> Any: + request.headers.update(await asyncio.to_thread(self._token_provider.build_auth_headers)) + yield request + + +class PydanticAISnowflakeHook(PydanticAIHook): + """ + Hook for Snowflake Cortex's OpenAI-compatible chat endpoint via pydantic-ai. + + Unlike the other ``PydanticAI*`` hooks, credentials do not live on this connection: they are + read from an existing ``snowflake`` connection (OAuth, PAT, or key-pair JWT -- whichever that + connection is configured for), refreshed on every request the same way as + ``SnowflakeCortexAgentHook`` and ``SnowflakeSqlApiHook``. See + :class:`~airflow.providers.snowflake.utils.rest_auth.SnowflakeRestTokenProvider`. The + underlying ``httpx2.AsyncClient`` is built once and lives as long as this hook instance; + nothing currently closes it (``SnowflakeProvider`` only owns and closes a client it built + itself, not one passed in). + + Connection fields: + - **extra** JSON: ``{"model": "snowflake:claude-4-sonnet", + "snowflake_conn_id": "snowflake_default"}`` + + Model family support (pydantic-ai-slim's ``SnowflakeProvider.model_profile``): Claude + (``claude*``) and OpenAI (``openai-*``) models support tools and structured output; + other families (``llama*``, ``snowflake-llama*``, ``mistral*``, ``mixtral*``, + ``deepseek*``, and any unlisted family) do not support tools, and structured output + falls back to prompted mode. Use a Claude or OpenAI family model for a tool-using agent. + + :param llm_conn_id: Airflow connection ID for this ``pydanticai_snowflake`` connection. + :param model_id: Model identifier, e.g. ``"snowflake:claude-4-sonnet"``. A bare name (no + recognized platform prefix) is qualified with ``snowflake:``. + :param fallback_conn_ids: See :class:`~airflow.providers.common.ai.hooks.pydantic_ai.PydanticAIHook`. + :param snowflake_conn_id: Connection ID of an existing Snowflake connection to source + credentials, account, and host from. Takes precedence over the connection extra's + ``snowflake_conn_id``; one of the two is required. + """ + + conn_type = "pydanticai_snowflake" + default_conn_name = "pydanticai_snowflake_default" + hook_name = "Pydantic AI (Snowflake Cortex)" + model_provider = "snowflake" + + def __init__( + self, + llm_conn_id: str | None = None, + model_id: str | None = None, + fallback_conn_ids: list[str] | None = None, + *, + snowflake_conn_id: str | None = None, + **kwargs: Any, + ) -> None: + super().__init__(llm_conn_id, model_id, fallback_conn_ids, **kwargs) + self.snowflake_conn_id = snowflake_conn_id + self._token_provider: SnowflakeRestTokenProvider | None = None + self._cortex_base_url: str | None = None + self._http_client: httpx2.AsyncClient | None = None + + @staticmethod + def get_ui_field_behaviour() -> dict[str, Any]: + """Return custom field behaviour for the Airflow connection form.""" + return { + "hidden_fields": ["schema", "port", "login", "host", "password"], + "relabeling": {}, + "placeholders": { + "extra": '{"model": "snowflake:claude-4-sonnet", "snowflake_conn_id": "snowflake_default"}', + }, + } + + def _get_snowflake_conn_id(self, extra: dict[str, Any]) -> str: + snowflake_conn_id = self.snowflake_conn_id or extra.get("snowflake_conn_id") + if not snowflake_conn_id: + raise ValueError( + f"Connection '{self.llm_conn_id}' has no Snowflake connection to source credentials " + "from. Set snowflake_conn_id on the hook or the connection's extra field, pointing " + "at an existing Snowflake connection." + ) + return snowflake_conn_id + + def _get_token_provider(self, extra: dict[str, Any]) -> SnowflakeRestTokenProvider: + """ + Build the Snowflake hook, token provider, base URL, and HTTP client once. + + Reused for this hook's lifetime -- including the ``httpx2.AsyncClient``, which nothing + else owns or closes (``SnowflakeProvider`` only owns and closes a client it built itself), + so building a fresh one on every call would leak one per call. + """ + if self._token_provider is None: + snowflake_hook = SnowflakeHook(snowflake_conn_id=self._get_snowflake_conn_id(extra)) + self._token_provider = SnowflakeRestTokenProvider(snowflake_hook) + self._cortex_base_url = ( Review Comment: If `get_cortex_base_url` raises here (no `host` and no usable `account`), `_token_provider` is already set, so a second `get_conn()` on the same hook skips this block and hands `SnowflakeProvider` `base_url=None, http_client=None`. It then picks the account up from `SNOWFLAKE_ACCOUNT` in the environment and sends the connection's PAT there with no per-request auth. Building all three into locals and assigning them together at the end would close that. ########## providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_model.py: ########## @@ -0,0 +1,182 @@ +# 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. +"""Expose Snowflake Cortex's OpenAI-compatible chat endpoint as a pydantic-ai model hook.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.providers.snowflake.utils.rest_auth import SnowflakeRestTokenProvider, get_cortex_base_url + +try: + import httpx2 + from pydantic_ai.providers import snowflake as _pydantic_ai_snowflake_provider # noqa: F401 Review Comment: `pydantic_ai.providers.snowflake` exists in every pydantic-ai-slim from 2.33 on (2.43 is what the 3.3.2 constraints pin), so this import never fails on a supported version. The case that does break is common.ai 0.9.0, also pinned in those constraints: the imports succeed and the first construction fails with `TypeError: PydanticAIHook.__init__() takes from 1 to 3 positional arguments but 4 were given`, because `fallback_conn_ids` arrived in 0.10.0. Could the guard check for that instead, and the message say common.ai 0.10.0 is needed? ########## 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] + conn_config["user"], # type: ignore[arg-type] + private_key=self._private_key, + lifetime=self._token_life_time, + renewal_delay=self._token_renewal_delta, + ) + token = self._jwt_generator.get_token() + return SnowflakeRestToken(token=token, token_type="KEYPAIR_JWT") + + def build_auth_headers(self) -> dict[str, str]: + """Return the two REST auth headers: ``Authorization`` and the token-type header.""" + token = self.get_token() + return { + "Authorization": f"Bearer {token.token}", + "X-Snowflake-Authorization-Token-Type": token.token_type, + } + + +def get_cortex_base_url(conn_config: dict[str, Any]) -> str: + """ + Return the base URL for a Snowflake account's Cortex REST endpoints. + + The extra field ``host`` wins when set; otherwise the URL is derived from ``account``. + Unlike ``SnowflakeHook.account_identifier`` (used by ``SnowflakeSqlApiHook``, which builds + its own URL and does not call this function), this never appends ``region`` -- Cortex does Review Comment: Is this right for locator-style accounts? Snowflake's account identifier docs give `xy12345.us-east-2.aws.snowflakecomputing.com` for a locator outside AWS us-west-2, and `SnowflakeSqlApiHook` gets there through `account_identifier`. So a connection with `account: xy12345` plus `region` works for the SQL API but sends Cortex requests to the wrong host. The agent hook already did this in 6.18.0, but the new model hook inherits it and this docstring now documents it as intended. ########## providers/snowflake/docs/operators/snowflake_cortex_agent.rst: ########## @@ -64,3 +64,21 @@ An example usage of the ``SnowflakeCortexAgentOperator`` is as follows: Parameters passed to the operator take precedence over the corresponding values configured in the Airflow connection metadata, such as ``database``, ``schema`` and ``role``. + +Authentication +^^^^^^^^^^^^^^ + +``SnowflakeCortexAgentHook`` (and the operator built on it) authenticate the connection's Review Comment: "authenticate the connection's `authenticator` extra" reads oddly; maybe "authenticate according to the connection's `authenticator` extra". The last paragraph also oversells the caching for the operator, which makes one request per `execute`. The reuse matters for a hook instance making many calls. -- 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]
