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


##########
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:
   Now passes `_UNUSED_TOKEN_PLACEHOLDER` into `_get_provider_kwargs`, so 
building the model never fetches credentials, and `test_connection` is 
overridden to call `get_token()` after resolving the model.



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