zozo123 commented on code in PR #74198:
URL: https://github.com/apache/airflow/pull/74198#discussion_r4221270450


##########
providers/databricks/src/airflow/providers/databricks/toolsets/unity_mcp.py:
##########
@@ -0,0 +1,339 @@
+# 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.
+"""Toolset that gives a common.ai agent the tools of a Unity Gateway MCP 
Service."""
+
+from __future__ import annotations
+
+import asyncio
+import contextlib
+import re
+import threading
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException
+from airflow.providers.databricks.exceptions import (
+    DatabricksUnityMCPAccessDeniedError,
+    DatabricksUnityMCPError,
+    DatabricksUnityMCPServiceNotFoundError,
+    DatabricksUnityMCPThrottledError,
+    DatabricksUnityMCPTransportError,
+)
+from airflow.providers.databricks.hooks.databricks import DatabricksHook
+
+try:
+    import httpx2
+    from fastmcp.client.transports import StreamableHttpTransport
+    from pydantic_ai.mcp import MCPToolset as PydanticAIMCPToolset
+
+    from airflow.providers.common.ai.toolsets.mcp import MCPToolset
+except ImportError as e:
+    raise AirflowOptionalProviderFeatureException(
+        "DatabricksUnityMCPToolset needs the 'common.ai' extra of the 
databricks provider: "
+        "pip install 'apache-airflow-providers-databricks[common.ai]'"
+    ) from e
+
+if TYPE_CHECKING:
+    from collections.abc import AsyncGenerator, AsyncIterator, Sequence
+
+    from pydantic_ai._run_context import RunContext
+    from pydantic_ai.toolsets.abstract import ToolsetTool
+
+    from airflow.sdk.execution_time.secrets_masker import mask_secret
+else:
+    try:
+        from airflow.sdk.log import mask_secret
+    except ImportError:
+        try:
+            from airflow.sdk.execution_time.secrets_masker import mask_secret
+        except ImportError:
+            from airflow.utils.log.secrets_masker import mask_secret
+
+GATEWAY_MCP_SERVICES_PATH = "ai-gateway/mcp-services"
+
+# Restricting each part to these characters keeps the name from adding a path, 
query or
+# fragment to the URL, so the request can only reach the service path on the 
connection's host.
+_SERVICE_NAME_PART = r"[A-Za-z0-9_-]+"
+_SERVICE_NAME = 
re.compile(rf"{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}")
+
+
+def validate_service_name(service_name: str) -> None:
+    """
+    Raise ``ValueError`` unless ``service_name`` is a three-level 
``catalog.schema.service`` name.
+
+    Each part may contain only ASCII letters, digits, underscores and hyphens.
+    """
+    if not isinstance(service_name, str) or not 
_SERVICE_NAME.fullmatch(service_name):
+        raise ValueError(
+            f"Invalid Unity Gateway MCP Service name {service_name!r}: 
expected "
+            "'catalog.schema.service', each part made of ASCII letters, 
digits, '_' or '-'."
+        )
+
+
+def _parse_retry_after(value: str | None) -> float | None:
+    if value is None:
+        return None
+    try:
+        seconds = float(value)
+    except ValueError:
+        # The HTTP-date form is not worth parsing: the caller only uses this 
as a hint.
+        return None
+    return seconds if seconds >= 0 else None
+
+
+def _find_transport_error(exc: BaseException) -> httpx2.TransportError | None:
+    """Return the network error behind ``exc``, looking through causes and 
exception groups."""
+    seen: set[int] = set()
+    pending: list[BaseException] = [exc]
+    while pending:
+        current = pending.pop()
+        if id(current) in seen:
+            continue
+        seen.add(id(current))
+        if isinstance(current, httpx2.TransportError):
+            return current
+        # Duck-typed: the ExceptionGroup builtin is Python 3.11+, and anyio 
uses the backport on 3.10.
+        if isinstance(grouped := getattr(current, "exceptions", None), (list, 
tuple)):
+            pending.extend(e for e in grouped if isinstance(e, BaseException))
+        pending.extend(e for e in (current.__cause__, current.__context__) if 
e is not None)
+    return None
+
+
+class _DatabricksTokenAuth(httpx2.Auth):
+    """
+    Authenticate each gateway request with a token from the Databricks 
connection.
+
+    Asking the hook on every request, rather than once, lets OAuth tokens 
refresh during a long
+    agent run and on reconnection; the hook caches tokens until they are about 
to expire.
+
+    The MCP client reports every HTTP error from the gateway as the same 
generic error, so the
+    status and ``Retry-After`` of the last error response, and any failure to 
get a token, are
+    kept here for :class:`DatabricksUnityMCPToolset` to report what went 
wrong. With concurrent
+    calls on one toolset, the error reported for a failed call can be another 
call's.
+    """
+
+    def __init__(self, hook: DatabricksHook) -> None:
+        self._hook = hook
+        self.error_status: int | None = None
+        self.retry_after: float | None = None
+        self.token_error: Exception | None = None
+        self._token_lock = threading.Lock()
+
+    def get_token(self) -> str:
+        with self._token_lock:
+            token = self._hook._get_token(raise_error=False)
+        if not token:
+            raise ValueError(
+                f"Connection {self._hook.databricks_conn_id!r} has no 
token-based authentication "
+                "configured. This toolset sends a bearer token: use a personal 
access token, "
+                "service principal OAuth, Azure AD, or workload identity 
federation. Username and "
+                "password authentication is not supported."
+            )
+        # A personal access token is masked when the connection is fetched; 
mask minted
+        # OAuth tokens too, so they never reach task logs.
+        mask_secret(token)
+        return token
+
+    def take_error(self) -> tuple[int | None, float | None, Exception | None]:
+        """Return the last recorded error and forget it."""
+        error = (self.error_status, self.retry_after, self.token_error)
+        self.error_status = self.retry_after = self.token_error = None
+        return error
+
+    def _record(self, response: httpx2.Response) -> None:
+        # The MCP client tolerates some error responses, such as one to a 
notification, so an
+        # error must not outlive the next success, or a later failure would be 
reported as it.
+        if response.status_code >= 400:
+            self.error_status = response.status_code
+            self.retry_after = 
_parse_retry_after(response.headers.get("Retry-After"))
+        else:
+            self.error_status = self.retry_after = None
+
+    async def async_auth_flow(
+        self, request: httpx2.Request
+    ) -> AsyncGenerator[httpx2.Request, httpx2.Response]:
+        try:
+            # Not AirflowToolset.run_blocking: its lock is shared by every 
toolset, so a long SQL
+            # query in another toolset would hold up each request here. The 
connection is already
+            # resolved by _get_server, so fetching a token only calls the 
token endpoint, and the
+            # hook is this toolset's own, so a lock of its own is enough.
+            token = await asyncio.to_thread(self.get_token)
+        except Exception as e:
+            self.token_error = e
+            raise
+        request.headers["Authorization"] = f"Bearer {token}"
+        response = yield request
+        self._record(response)
+
+
+class DatabricksUnityMCPToolset(MCPToolset):
+    """
+    Give an agent the tools of a Unity Gateway MCP Service, authenticated as 
the connection's identity.
+
+    The service is named by its three-level Unity Catalog name, 
``catalog.schema.service``, and
+    reached at ``https://<workspace 
host>/ai-gateway/mcp-services/<catalog.schema.service>``. The
+    workspace host and the credentials both come from the Databricks 
connection, so Dag code holds
+    neither a gateway URL nor a token, and the token is only ever sent to that 
workspace.
+
+    The gateway runs every tool call as the identity of the connection's 
credentials (a user's
+    personal access token, a service principal, or an Azure AD / federated 
identity). That identity
+    needs ``EXECUTE`` on the MCP Service, ``USE CATALOG`` and ``USE SCHEMA`` 
on its catalog and
+    schema, and an assignment to the workspace. It sees only the tools 
selected for the service, and
+    the service's policies apply.
+
+    Tokens are fetched from the connection for each request, so OAuth tokens 
refresh during a long
+    agent run and on reconnection. Gateway errors are raised as
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPAccessDeniedError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPServiceNotFoundError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPThrottledError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPTransportError`,
 or, for any
+    other gateway error, 
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPError`.
+    Tool calls are never retried by this toolset, because a call interrupted 
after it was sent may
+    already have run.
+
+    .. code-block:: python
+
+        from airflow.providers.common.ai.operators.agent import AgentOperator
+        from airflow.providers.databricks.toolsets.unity_mcp import 
DatabricksUnityMCPToolset
+
+        AgentOperator(
+            task_id="ask_mcp_service",
+            prompt="Which tools do you have?",
+            llm_conn_id="pydanticai_default",
+            toolsets=[DatabricksUnityMCPToolset("main.default.my_mcp", 
databricks_conn_id="databricks")],
+        )
+
+    :param service_name: Three-level name of the MCP Service, 
``catalog.schema.service``. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param databricks_conn_id: Databricks connection whose host and 
credentials are used. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param tool_prefix: Optional prefix prepended to tool names.
+    """
+
+    agent_template_fields: Sequence[str] = ("_databricks_conn_id", 
"_service_name")
+
+    def __init__(
+        self,
+        service_name: str,
+        *,
+        databricks_conn_id: str = DatabricksHook.default_conn_name,
+        tool_prefix: str | None = None,
+    ) -> None:
+        super().__init__(databricks_conn_id, tool_prefix=tool_prefix)
+        self._databricks_conn_id = databricks_conn_id
+        self._service_name = service_name
+        self._auth: _DatabricksTokenAuth | None = None
+        # A templated name is checked once it has been rendered, in 
_get_server.
+        if "{{" not in service_name:
+            validate_service_name(service_name)
+
+    @property
+    def id(self) -> str:
+        return 
f"databricks-unity-mcp-{self._databricks_conn_id}-{self._service_name}"
+
+    def get_service_url(self, hook: DatabricksHook) -> str:
+        """Return the gateway URL of the MCP Service on the connection's 
workspace."""
+        validate_service_name(self._service_name)
+        if not hook.host:
+            raise ValueError(f"Connection {self._databricks_conn_id!r} has no 
workspace host.")
+        return 
hook._endpoint_url(f"{GATEWAY_MCP_SERVICES_PATH}/{self._service_name}")
+
+    def _get_server(self) -> Any:
+        if self._server is None:
+            hook = DatabricksHook(self._databricks_conn_id, 
caller=type(self).__name__)
+            url = self.get_service_url(hook)
+            auth = _DatabricksTokenAuth(hook)
+            # Fail here, with a clear message, rather than inside the MCP 
client, which reports
+            # any failure as a generic connection error.
+            auth.get_token()
+            transport = StreamableHttpTransport(url, 
headers=hook.user_agent_header, auth=auth)
+            toolset = PydanticAIMCPToolset(transport)
+            self._auth = auth
+            self._server = toolset.prefixed(self._tool_prefix) if 
self._tool_prefix else toolset
+        return self._server
+
+    def _translate_error(self, error: Exception, *, during_tool_call: bool) -> 
Exception | None:
+        """Return the provider exception for a gateway failure, or ``None`` to 
re-raise ``error`` as is."""
+        status, retry_after, token_error = self._auth.take_error() if 
self._auth else (None, None, None)
+        service = self._service_name
+        if token_error is not None:
+            return DatabricksUnityMCPError(
+                f"Could not get a token from connection 
{self._databricks_conn_id!r} to call MCP Service "
+                f"{service!r}: {token_error}"
+            )
+        if status in (401, 403):
+            return DatabricksUnityMCPAccessDeniedError(
+                f"Unity Gateway denied access to MCP Service {service!r} (HTTP 
{status}). The "
+                f"identity of connection {self._databricks_conn_id!r} needs 
EXECUTE on the service and "
+                "USE CATALOG and USE SCHEMA on its catalog and schema, and its 
credentials must be valid.",
+                http_status_code=status,
+            )
+        if status == 404:

Review Comment:
   On a live workspace a nonexistent service returns 403 with JSON-RPC `-32007 
Not authorized to invoke MCP service.`, not a 404, so this branch doesn't fire 
for a missing service. Maybe fold it into the 403 message ('doesn't exist, or 
lacks EXECUTE...') and fix the docs table.



##########
providers/databricks/provider.yaml:
##########
@@ -143,6 +143,11 @@ integrations:
     how-to-guide:
       - /docs/apache-airflow-providers-databricks/operators/workflow.rst
     tags: [service]
+  - integration-name: Databricks Unity Gateway
+    external-doc-url: https://docs.databricks.com/aws/en/unity-gateway/concepts
+    how-to-guide:

Review Comment:
   `check_doc_files` only accepts how-to guides under 
`docs/operators|sensors|transfer`, which is why Static checks fails. Drop this 
`how-to-guide` (and regenerate `get_provider_info.py`).



##########
providers/databricks/src/airflow/providers/databricks/toolsets/unity_mcp.py:
##########
@@ -0,0 +1,339 @@
+# 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.
+"""Toolset that gives a common.ai agent the tools of a Unity Gateway MCP 
Service."""
+
+from __future__ import annotations
+
+import asyncio
+import contextlib
+import re
+import threading
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException
+from airflow.providers.databricks.exceptions import (
+    DatabricksUnityMCPAccessDeniedError,
+    DatabricksUnityMCPError,
+    DatabricksUnityMCPServiceNotFoundError,
+    DatabricksUnityMCPThrottledError,
+    DatabricksUnityMCPTransportError,
+)
+from airflow.providers.databricks.hooks.databricks import DatabricksHook
+
+try:
+    import httpx2
+    from fastmcp.client.transports import StreamableHttpTransport
+    from pydantic_ai.mcp import MCPToolset as PydanticAIMCPToolset
+
+    from airflow.providers.common.ai.toolsets.mcp import MCPToolset
+except ImportError as e:
+    raise AirflowOptionalProviderFeatureException(
+        "DatabricksUnityMCPToolset needs the 'common.ai' extra of the 
databricks provider: "
+        "pip install 'apache-airflow-providers-databricks[common.ai]'"
+    ) from e
+
+if TYPE_CHECKING:
+    from collections.abc import AsyncGenerator, AsyncIterator, Sequence
+
+    from pydantic_ai._run_context import RunContext
+    from pydantic_ai.toolsets.abstract import ToolsetTool
+
+    from airflow.sdk.execution_time.secrets_masker import mask_secret
+else:
+    try:
+        from airflow.sdk.log import mask_secret
+    except ImportError:
+        try:
+            from airflow.sdk.execution_time.secrets_masker import mask_secret
+        except ImportError:
+            from airflow.utils.log.secrets_masker import mask_secret
+
+GATEWAY_MCP_SERVICES_PATH = "ai-gateway/mcp-services"
+
+# Restricting each part to these characters keeps the name from adding a path, 
query or
+# fragment to the URL, so the request can only reach the service path on the 
connection's host.
+_SERVICE_NAME_PART = r"[A-Za-z0-9_-]+"
+_SERVICE_NAME = 
re.compile(rf"{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}")
+
+
+def validate_service_name(service_name: str) -> None:
+    """
+    Raise ``ValueError`` unless ``service_name`` is a three-level 
``catalog.schema.service`` name.
+
+    Each part may contain only ASCII letters, digits, underscores and hyphens.
+    """
+    if not isinstance(service_name, str) or not 
_SERVICE_NAME.fullmatch(service_name):
+        raise ValueError(
+            f"Invalid Unity Gateway MCP Service name {service_name!r}: 
expected "
+            "'catalog.schema.service', each part made of ASCII letters, 
digits, '_' or '-'."
+        )
+
+
+def _parse_retry_after(value: str | None) -> float | None:
+    if value is None:
+        return None
+    try:
+        seconds = float(value)
+    except ValueError:
+        # The HTTP-date form is not worth parsing: the caller only uses this 
as a hint.
+        return None
+    return seconds if seconds >= 0 else None
+
+
+def _find_transport_error(exc: BaseException) -> httpx2.TransportError | None:
+    """Return the network error behind ``exc``, looking through causes and 
exception groups."""
+    seen: set[int] = set()
+    pending: list[BaseException] = [exc]
+    while pending:
+        current = pending.pop()
+        if id(current) in seen:
+            continue
+        seen.add(id(current))
+        if isinstance(current, httpx2.TransportError):
+            return current
+        # Duck-typed: the ExceptionGroup builtin is Python 3.11+, and anyio 
uses the backport on 3.10.
+        if isinstance(grouped := getattr(current, "exceptions", None), (list, 
tuple)):
+            pending.extend(e for e in grouped if isinstance(e, BaseException))
+        pending.extend(e for e in (current.__cause__, current.__context__) if 
e is not None)
+    return None
+
+
+class _DatabricksTokenAuth(httpx2.Auth):
+    """
+    Authenticate each gateway request with a token from the Databricks 
connection.
+
+    Asking the hook on every request, rather than once, lets OAuth tokens 
refresh during a long
+    agent run and on reconnection; the hook caches tokens until they are about 
to expire.
+
+    The MCP client reports every HTTP error from the gateway as the same 
generic error, so the
+    status and ``Retry-After`` of the last error response, and any failure to 
get a token, are
+    kept here for :class:`DatabricksUnityMCPToolset` to report what went 
wrong. With concurrent
+    calls on one toolset, the error reported for a failed call can be another 
call's.
+    """
+
+    def __init__(self, hook: DatabricksHook) -> None:
+        self._hook = hook
+        self.error_status: int | None = None
+        self.retry_after: float | None = None
+        self.token_error: Exception | None = None
+        self._token_lock = threading.Lock()
+
+    def get_token(self) -> str:
+        with self._token_lock:
+            token = self._hook._get_token(raise_error=False)
+        if not token:
+            raise ValueError(
+                f"Connection {self._hook.databricks_conn_id!r} has no 
token-based authentication "
+                "configured. This toolset sends a bearer token: use a personal 
access token, "
+                "service principal OAuth, Azure AD, or workload identity 
federation. Username and "
+                "password authentication is not supported."
+            )
+        # A personal access token is masked when the connection is fetched; 
mask minted
+        # OAuth tokens too, so they never reach task logs.
+        mask_secret(token)
+        return token
+
+    def take_error(self) -> tuple[int | None, float | None, Exception | None]:
+        """Return the last recorded error and forget it."""
+        error = (self.error_status, self.retry_after, self.token_error)
+        self.error_status = self.retry_after = self.token_error = None
+        return error
+
+    def _record(self, response: httpx2.Response) -> None:

Review Comment:
   This status is shared across concurrent calls. Another call's 2xx can clear 
it before the failed call reads it, and then an ambiguous 5xx on `tools/call` 
surfaces as `ModelRetry` (I reproduced 4 of 20 with two parallel calls). Could 
the fatal-vs-retry decision come from the call's own exception chain (the 
client's synthesized `McpError`), with this kept only as a hint for the message?



##########
providers/databricks/pyproject.toml:
##########
@@ -114,6 +117,7 @@ dev = [
     "apache-airflow-task-sdk",
     "apache-airflow-devel-common",
     "apache-airflow-providers-amazon",
+    "apache-airflow-providers-common-ai",

Review Comment:
   `apache-airflow-providers-common-ai[mcp]`. Otherwise fastmcp isn't installed 
by `uv sync` and `test_unity_mcp.py` skips.



##########
providers/databricks/src/airflow/providers/databricks/toolsets/unity_mcp.py:
##########
@@ -0,0 +1,339 @@
+# 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.
+"""Toolset that gives a common.ai agent the tools of a Unity Gateway MCP 
Service."""
+
+from __future__ import annotations
+
+import asyncio
+import contextlib
+import re
+import threading
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException
+from airflow.providers.databricks.exceptions import (
+    DatabricksUnityMCPAccessDeniedError,
+    DatabricksUnityMCPError,
+    DatabricksUnityMCPServiceNotFoundError,
+    DatabricksUnityMCPThrottledError,
+    DatabricksUnityMCPTransportError,
+)
+from airflow.providers.databricks.hooks.databricks import DatabricksHook
+
+try:
+    import httpx2
+    from fastmcp.client.transports import StreamableHttpTransport
+    from pydantic_ai.mcp import MCPToolset as PydanticAIMCPToolset
+
+    from airflow.providers.common.ai.toolsets.mcp import MCPToolset
+except ImportError as e:
+    raise AirflowOptionalProviderFeatureException(
+        "DatabricksUnityMCPToolset needs the 'common.ai' extra of the 
databricks provider: "
+        "pip install 'apache-airflow-providers-databricks[common.ai]'"
+    ) from e
+
+if TYPE_CHECKING:
+    from collections.abc import AsyncGenerator, AsyncIterator, Sequence
+
+    from pydantic_ai._run_context import RunContext
+    from pydantic_ai.toolsets.abstract import ToolsetTool
+
+    from airflow.sdk.execution_time.secrets_masker import mask_secret
+else:
+    try:
+        from airflow.sdk.log import mask_secret
+    except ImportError:
+        try:
+            from airflow.sdk.execution_time.secrets_masker import mask_secret
+        except ImportError:
+            from airflow.utils.log.secrets_masker import mask_secret
+
+GATEWAY_MCP_SERVICES_PATH = "ai-gateway/mcp-services"
+
+# Restricting each part to these characters keeps the name from adding a path, 
query or
+# fragment to the URL, so the request can only reach the service path on the 
connection's host.
+_SERVICE_NAME_PART = r"[A-Za-z0-9_-]+"
+_SERVICE_NAME = 
re.compile(rf"{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}")
+
+
+def validate_service_name(service_name: str) -> None:
+    """
+    Raise ``ValueError`` unless ``service_name`` is a three-level 
``catalog.schema.service`` name.
+
+    Each part may contain only ASCII letters, digits, underscores and hyphens.
+    """
+    if not isinstance(service_name, str) or not 
_SERVICE_NAME.fullmatch(service_name):
+        raise ValueError(
+            f"Invalid Unity Gateway MCP Service name {service_name!r}: 
expected "
+            "'catalog.schema.service', each part made of ASCII letters, 
digits, '_' or '-'."
+        )
+
+
+def _parse_retry_after(value: str | None) -> float | None:
+    if value is None:
+        return None
+    try:
+        seconds = float(value)
+    except ValueError:
+        # The HTTP-date form is not worth parsing: the caller only uses this 
as a hint.
+        return None
+    return seconds if seconds >= 0 else None
+
+
+def _find_transport_error(exc: BaseException) -> httpx2.TransportError | None:
+    """Return the network error behind ``exc``, looking through causes and 
exception groups."""
+    seen: set[int] = set()
+    pending: list[BaseException] = [exc]
+    while pending:
+        current = pending.pop()
+        if id(current) in seen:
+            continue
+        seen.add(id(current))
+        if isinstance(current, httpx2.TransportError):
+            return current
+        # Duck-typed: the ExceptionGroup builtin is Python 3.11+, and anyio 
uses the backport on 3.10.
+        if isinstance(grouped := getattr(current, "exceptions", None), (list, 
tuple)):
+            pending.extend(e for e in grouped if isinstance(e, BaseException))
+        pending.extend(e for e in (current.__cause__, current.__context__) if 
e is not None)
+    return None
+
+
+class _DatabricksTokenAuth(httpx2.Auth):
+    """
+    Authenticate each gateway request with a token from the Databricks 
connection.
+
+    Asking the hook on every request, rather than once, lets OAuth tokens 
refresh during a long
+    agent run and on reconnection; the hook caches tokens until they are about 
to expire.
+
+    The MCP client reports every HTTP error from the gateway as the same 
generic error, so the
+    status and ``Retry-After`` of the last error response, and any failure to 
get a token, are
+    kept here for :class:`DatabricksUnityMCPToolset` to report what went 
wrong. With concurrent
+    calls on one toolset, the error reported for a failed call can be another 
call's.
+    """
+
+    def __init__(self, hook: DatabricksHook) -> None:
+        self._hook = hook
+        self.error_status: int | None = None
+        self.retry_after: float | None = None
+        self.token_error: Exception | None = None
+        self._token_lock = threading.Lock()
+
+    def get_token(self) -> str:
+        with self._token_lock:
+            token = self._hook._get_token(raise_error=False)
+        if not token:
+            raise ValueError(
+                f"Connection {self._hook.databricks_conn_id!r} has no 
token-based authentication "
+                "configured. This toolset sends a bearer token: use a personal 
access token, "
+                "service principal OAuth, Azure AD, or workload identity 
federation. Username and "
+                "password authentication is not supported."
+            )
+        # A personal access token is masked when the connection is fetched; 
mask minted
+        # OAuth tokens too, so they never reach task logs.
+        mask_secret(token)
+        return token
+
+    def take_error(self) -> tuple[int | None, float | None, Exception | None]:
+        """Return the last recorded error and forget it."""
+        error = (self.error_status, self.retry_after, self.token_error)
+        self.error_status = self.retry_after = self.token_error = None
+        return error
+
+    def _record(self, response: httpx2.Response) -> None:
+        # The MCP client tolerates some error responses, such as one to a 
notification, so an
+        # error must not outlive the next success, or a later failure would be 
reported as it.
+        if response.status_code >= 400:
+            self.error_status = response.status_code
+            self.retry_after = 
_parse_retry_after(response.headers.get("Retry-After"))
+        else:
+            self.error_status = self.retry_after = None
+
+    async def async_auth_flow(
+        self, request: httpx2.Request
+    ) -> AsyncGenerator[httpx2.Request, httpx2.Response]:
+        try:
+            # Not AirflowToolset.run_blocking: its lock is shared by every 
toolset, so a long SQL
+            # query in another toolset would hold up each request here. The 
connection is already
+            # resolved by _get_server, so fetching a token only calls the 
token endpoint, and the
+            # hook is this toolset's own, so a lock of its own is enough.
+            token = await asyncio.to_thread(self.get_token)
+        except Exception as e:
+            self.token_error = e
+            raise
+        request.headers["Authorization"] = f"Bearer {token}"
+        response = yield request
+        self._record(response)
+
+
+class DatabricksUnityMCPToolset(MCPToolset):
+    """
+    Give an agent the tools of a Unity Gateway MCP Service, authenticated as 
the connection's identity.
+
+    The service is named by its three-level Unity Catalog name, 
``catalog.schema.service``, and
+    reached at ``https://<workspace 
host>/ai-gateway/mcp-services/<catalog.schema.service>``. The
+    workspace host and the credentials both come from the Databricks 
connection, so Dag code holds
+    neither a gateway URL nor a token, and the token is only ever sent to that 
workspace.
+
+    The gateway runs every tool call as the identity of the connection's 
credentials (a user's
+    personal access token, a service principal, or an Azure AD / federated 
identity). That identity
+    needs ``EXECUTE`` on the MCP Service, ``USE CATALOG`` and ``USE SCHEMA`` 
on its catalog and
+    schema, and an assignment to the workspace. It sees only the tools 
selected for the service, and
+    the service's policies apply.
+
+    Tokens are fetched from the connection for each request, so OAuth tokens 
refresh during a long
+    agent run and on reconnection. Gateway errors are raised as
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPAccessDeniedError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPServiceNotFoundError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPThrottledError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPTransportError`,
 or, for any
+    other gateway error, 
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPError`.
+    Tool calls are never retried by this toolset, because a call interrupted 
after it was sent may
+    already have run.
+
+    .. code-block:: python
+
+        from airflow.providers.common.ai.operators.agent import AgentOperator
+        from airflow.providers.databricks.toolsets.unity_mcp import 
DatabricksUnityMCPToolset
+
+        AgentOperator(
+            task_id="ask_mcp_service",
+            prompt="Which tools do you have?",
+            llm_conn_id="pydanticai_default",
+            toolsets=[DatabricksUnityMCPToolset("main.default.my_mcp", 
databricks_conn_id="databricks")],
+        )
+
+    :param service_name: Three-level name of the MCP Service, 
``catalog.schema.service``. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param databricks_conn_id: Databricks connection whose host and 
credentials are used. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param tool_prefix: Optional prefix prepended to tool names.
+    """
+
+    agent_template_fields: Sequence[str] = ("_databricks_conn_id", 
"_service_name")
+
+    def __init__(
+        self,
+        service_name: str,
+        *,
+        databricks_conn_id: str = DatabricksHook.default_conn_name,
+        tool_prefix: str | None = None,
+    ) -> None:
+        super().__init__(databricks_conn_id, tool_prefix=tool_prefix)
+        self._databricks_conn_id = databricks_conn_id
+        self._service_name = service_name
+        self._auth: _DatabricksTokenAuth | None = None
+        # A templated name is checked once it has been rendered, in 
_get_server.
+        if "{{" not in service_name:
+            validate_service_name(service_name)
+
+    @property
+    def id(self) -> str:
+        return 
f"databricks-unity-mcp-{self._databricks_conn_id}-{self._service_name}"
+
+    def get_service_url(self, hook: DatabricksHook) -> str:
+        """Return the gateway URL of the MCP Service on the connection's 
workspace."""
+        validate_service_name(self._service_name)
+        if not hook.host:
+            raise ValueError(f"Connection {self._databricks_conn_id!r} has no 
workspace host.")
+        return 
hook._endpoint_url(f"{GATEWAY_MCP_SERVICES_PATH}/{self._service_name}")
+
+    def _get_server(self) -> Any:
+        if self._server is None:
+            hook = DatabricksHook(self._databricks_conn_id, 
caller=type(self).__name__)
+            url = self.get_service_url(hook)
+            auth = _DatabricksTokenAuth(hook)
+            # Fail here, with a clear message, rather than inside the MCP 
client, which reports
+            # any failure as a generic connection error.
+            auth.get_token()
+            transport = StreamableHttpTransport(url, 
headers=hook.user_agent_header, auth=auth)
+            toolset = PydanticAIMCPToolset(transport)
+            self._auth = auth
+            self._server = toolset.prefixed(self._tool_prefix) if 
self._tool_prefix else toolset
+        return self._server
+
+    def _translate_error(self, error: Exception, *, during_tool_call: bool) -> 
Exception | None:
+        """Return the provider exception for a gateway failure, or ``None`` to 
re-raise ``error`` as is."""
+        status, retry_after, token_error = self._auth.take_error() if 
self._auth else (None, None, None)
+        service = self._service_name
+        if token_error is not None:
+            return DatabricksUnityMCPError(
+                f"Could not get a token from connection 
{self._databricks_conn_id!r} to call MCP Service "
+                f"{service!r}: {token_error}"
+            )
+        if status in (401, 403):
+            return DatabricksUnityMCPAccessDeniedError(
+                f"Unity Gateway denied access to MCP Service {service!r} (HTTP 
{status}). The "
+                f"identity of connection {self._databricks_conn_id!r} needs 
EXECUTE on the service and "
+                "USE CATALOG and USE SCHEMA on its catalog and schema, and its 
credentials must be valid.",
+                http_status_code=status,
+            )
+        if status == 404:
+            return DatabricksUnityMCPServiceNotFoundError(
+                f"MCP Service {service!r} was not found on the workspace of 
connection "
+                f"{self._databricks_conn_id!r}, or is not visible to its 
identity (HTTP 404).",
+                http_status_code=status,
+            )
+        if status == 429:
+            hint = f" Retry after {retry_after:g} seconds." if retry_after is 
not None else ""
+            return DatabricksUnityMCPThrottledError(
+                f"Unity Gateway rate-limited calls to MCP Service {service!r} 
(HTTP 429).{hint}",
+                http_status_code=status,
+                retry_after=retry_after,
+            )
+        ambiguous = (
+            " The tool call may or may not have run, so it was not retried." 
if during_tool_call else ""
+        )
+        if status is not None:

Review Comment:
   A `400` that carries a JSON-RPC error body (e.g. `-32602 Invalid params`) 
ends up here as a fatal error, so the model can't correct its arguments. The 
MCP client already surfaces those bodies as a real `McpError`. Let them pass 
through as `ModelRetry`.



##########
providers/databricks/src/airflow/providers/databricks/toolsets/unity_mcp.py:
##########
@@ -0,0 +1,339 @@
+# 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.
+"""Toolset that gives a common.ai agent the tools of a Unity Gateway MCP 
Service."""
+
+from __future__ import annotations
+
+import asyncio
+import contextlib
+import re
+import threading
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException
+from airflow.providers.databricks.exceptions import (
+    DatabricksUnityMCPAccessDeniedError,
+    DatabricksUnityMCPError,
+    DatabricksUnityMCPServiceNotFoundError,
+    DatabricksUnityMCPThrottledError,
+    DatabricksUnityMCPTransportError,
+)
+from airflow.providers.databricks.hooks.databricks import DatabricksHook
+
+try:
+    import httpx2
+    from fastmcp.client.transports import StreamableHttpTransport
+    from pydantic_ai.mcp import MCPToolset as PydanticAIMCPToolset
+
+    from airflow.providers.common.ai.toolsets.mcp import MCPToolset
+except ImportError as e:
+    raise AirflowOptionalProviderFeatureException(
+        "DatabricksUnityMCPToolset needs the 'common.ai' extra of the 
databricks provider: "
+        "pip install 'apache-airflow-providers-databricks[common.ai]'"
+    ) from e
+
+if TYPE_CHECKING:
+    from collections.abc import AsyncGenerator, AsyncIterator, Sequence
+
+    from pydantic_ai._run_context import RunContext
+    from pydantic_ai.toolsets.abstract import ToolsetTool
+
+    from airflow.sdk.execution_time.secrets_masker import mask_secret
+else:
+    try:
+        from airflow.sdk.log import mask_secret
+    except ImportError:
+        try:
+            from airflow.sdk.execution_time.secrets_masker import mask_secret
+        except ImportError:
+            from airflow.utils.log.secrets_masker import mask_secret
+
+GATEWAY_MCP_SERVICES_PATH = "ai-gateway/mcp-services"
+
+# Restricting each part to these characters keeps the name from adding a path, 
query or
+# fragment to the URL, so the request can only reach the service path on the 
connection's host.
+_SERVICE_NAME_PART = r"[A-Za-z0-9_-]+"
+_SERVICE_NAME = 
re.compile(rf"{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}")
+
+
+def validate_service_name(service_name: str) -> None:
+    """
+    Raise ``ValueError`` unless ``service_name`` is a three-level 
``catalog.schema.service`` name.
+
+    Each part may contain only ASCII letters, digits, underscores and hyphens.
+    """
+    if not isinstance(service_name, str) or not 
_SERVICE_NAME.fullmatch(service_name):
+        raise ValueError(
+            f"Invalid Unity Gateway MCP Service name {service_name!r}: 
expected "
+            "'catalog.schema.service', each part made of ASCII letters, 
digits, '_' or '-'."
+        )
+
+
+def _parse_retry_after(value: str | None) -> float | None:
+    if value is None:
+        return None
+    try:
+        seconds = float(value)
+    except ValueError:
+        # The HTTP-date form is not worth parsing: the caller only uses this 
as a hint.
+        return None
+    return seconds if seconds >= 0 else None
+
+
+def _find_transport_error(exc: BaseException) -> httpx2.TransportError | None:
+    """Return the network error behind ``exc``, looking through causes and 
exception groups."""
+    seen: set[int] = set()
+    pending: list[BaseException] = [exc]
+    while pending:
+        current = pending.pop()
+        if id(current) in seen:
+            continue
+        seen.add(id(current))
+        if isinstance(current, httpx2.TransportError):
+            return current
+        # Duck-typed: the ExceptionGroup builtin is Python 3.11+, and anyio 
uses the backport on 3.10.
+        if isinstance(grouped := getattr(current, "exceptions", None), (list, 
tuple)):
+            pending.extend(e for e in grouped if isinstance(e, BaseException))
+        pending.extend(e for e in (current.__cause__, current.__context__) if 
e is not None)
+    return None
+
+
+class _DatabricksTokenAuth(httpx2.Auth):
+    """
+    Authenticate each gateway request with a token from the Databricks 
connection.
+
+    Asking the hook on every request, rather than once, lets OAuth tokens 
refresh during a long
+    agent run and on reconnection; the hook caches tokens until they are about 
to expire.
+
+    The MCP client reports every HTTP error from the gateway as the same 
generic error, so the
+    status and ``Retry-After`` of the last error response, and any failure to 
get a token, are
+    kept here for :class:`DatabricksUnityMCPToolset` to report what went 
wrong. With concurrent
+    calls on one toolset, the error reported for a failed call can be another 
call's.
+    """
+
+    def __init__(self, hook: DatabricksHook) -> None:
+        self._hook = hook
+        self.error_status: int | None = None
+        self.retry_after: float | None = None
+        self.token_error: Exception | None = None
+        self._token_lock = threading.Lock()
+
+    def get_token(self) -> str:
+        with self._token_lock:
+            token = self._hook._get_token(raise_error=False)
+        if not token:
+            raise ValueError(
+                f"Connection {self._hook.databricks_conn_id!r} has no 
token-based authentication "
+                "configured. This toolset sends a bearer token: use a personal 
access token, "
+                "service principal OAuth, Azure AD, or workload identity 
federation. Username and "
+                "password authentication is not supported."
+            )
+        # A personal access token is masked when the connection is fetched; 
mask minted
+        # OAuth tokens too, so they never reach task logs.
+        mask_secret(token)
+        return token
+
+    def take_error(self) -> tuple[int | None, float | None, Exception | None]:
+        """Return the last recorded error and forget it."""
+        error = (self.error_status, self.retry_after, self.token_error)
+        self.error_status = self.retry_after = self.token_error = None
+        return error
+
+    def _record(self, response: httpx2.Response) -> None:
+        # The MCP client tolerates some error responses, such as one to a 
notification, so an
+        # error must not outlive the next success, or a later failure would be 
reported as it.
+        if response.status_code >= 400:
+            self.error_status = response.status_code
+            self.retry_after = 
_parse_retry_after(response.headers.get("Retry-After"))
+        else:
+            self.error_status = self.retry_after = None
+
+    async def async_auth_flow(
+        self, request: httpx2.Request
+    ) -> AsyncGenerator[httpx2.Request, httpx2.Response]:
+        try:
+            # Not AirflowToolset.run_blocking: its lock is shared by every 
toolset, so a long SQL
+            # query in another toolset would hold up each request here. The 
connection is already
+            # resolved by _get_server, so fetching a token only calls the 
token endpoint, and the
+            # hook is this toolset's own, so a lock of its own is enough.
+            token = await asyncio.to_thread(self.get_token)
+        except Exception as e:
+            self.token_error = e
+            raise
+        request.headers["Authorization"] = f"Bearer {token}"
+        response = yield request
+        self._record(response)
+
+
+class DatabricksUnityMCPToolset(MCPToolset):
+    """
+    Give an agent the tools of a Unity Gateway MCP Service, authenticated as 
the connection's identity.
+
+    The service is named by its three-level Unity Catalog name, 
``catalog.schema.service``, and
+    reached at ``https://<workspace 
host>/ai-gateway/mcp-services/<catalog.schema.service>``. The
+    workspace host and the credentials both come from the Databricks 
connection, so Dag code holds
+    neither a gateway URL nor a token, and the token is only ever sent to that 
workspace.
+
+    The gateway runs every tool call as the identity of the connection's 
credentials (a user's
+    personal access token, a service principal, or an Azure AD / federated 
identity). That identity
+    needs ``EXECUTE`` on the MCP Service, ``USE CATALOG`` and ``USE SCHEMA`` 
on its catalog and
+    schema, and an assignment to the workspace. It sees only the tools 
selected for the service, and
+    the service's policies apply.
+
+    Tokens are fetched from the connection for each request, so OAuth tokens 
refresh during a long
+    agent run and on reconnection. Gateway errors are raised as
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPAccessDeniedError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPServiceNotFoundError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPThrottledError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPTransportError`,
 or, for any
+    other gateway error, 
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPError`.
+    Tool calls are never retried by this toolset, because a call interrupted 
after it was sent may
+    already have run.
+
+    .. code-block:: python
+
+        from airflow.providers.common.ai.operators.agent import AgentOperator
+        from airflow.providers.databricks.toolsets.unity_mcp import 
DatabricksUnityMCPToolset
+
+        AgentOperator(
+            task_id="ask_mcp_service",
+            prompt="Which tools do you have?",
+            llm_conn_id="pydanticai_default",
+            toolsets=[DatabricksUnityMCPToolset("main.default.my_mcp", 
databricks_conn_id="databricks")],
+        )
+
+    :param service_name: Three-level name of the MCP Service, 
``catalog.schema.service``. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param databricks_conn_id: Databricks connection whose host and 
credentials are used. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param tool_prefix: Optional prefix prepended to tool names.
+    """
+
+    agent_template_fields: Sequence[str] = ("_databricks_conn_id", 
"_service_name")
+
+    def __init__(
+        self,
+        service_name: str,
+        *,
+        databricks_conn_id: str = DatabricksHook.default_conn_name,
+        tool_prefix: str | None = None,
+    ) -> None:
+        super().__init__(databricks_conn_id, tool_prefix=tool_prefix)
+        self._databricks_conn_id = databricks_conn_id
+        self._service_name = service_name
+        self._auth: _DatabricksTokenAuth | None = None
+        # A templated name is checked once it has been rendered, in 
_get_server.
+        if "{{" not in service_name:
+            validate_service_name(service_name)
+
+    @property
+    def id(self) -> str:
+        return 
f"databricks-unity-mcp-{self._databricks_conn_id}-{self._service_name}"
+
+    def get_service_url(self, hook: DatabricksHook) -> str:
+        """Return the gateway URL of the MCP Service on the connection's 
workspace."""
+        validate_service_name(self._service_name)
+        if not hook.host:
+            raise ValueError(f"Connection {self._databricks_conn_id!r} has no 
workspace host.")
+        return 
hook._endpoint_url(f"{GATEWAY_MCP_SERVICES_PATH}/{self._service_name}")
+
+    def _get_server(self) -> Any:
+        if self._server is None:
+            hook = DatabricksHook(self._databricks_conn_id, 
caller=type(self).__name__)
+            url = self.get_service_url(hook)
+            auth = _DatabricksTokenAuth(hook)
+            # Fail here, with a clear message, rather than inside the MCP 
client, which reports
+            # any failure as a generic connection error.
+            auth.get_token()
+            transport = StreamableHttpTransport(url, 
headers=hook.user_agent_header, auth=auth)

Review Comment:
   `hook.proxies` isn't applied here, so connections that need the `proxies` 
extra won't reach the gateway. `StreamableHttpTransport` takes 
`httpx_client_factory`.



##########
providers/databricks/src/airflow/providers/databricks/toolsets/unity_mcp.py:
##########
@@ -0,0 +1,339 @@
+# 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.
+"""Toolset that gives a common.ai agent the tools of a Unity Gateway MCP 
Service."""
+
+from __future__ import annotations
+
+import asyncio
+import contextlib
+import re
+import threading
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException
+from airflow.providers.databricks.exceptions import (
+    DatabricksUnityMCPAccessDeniedError,
+    DatabricksUnityMCPError,
+    DatabricksUnityMCPServiceNotFoundError,
+    DatabricksUnityMCPThrottledError,
+    DatabricksUnityMCPTransportError,
+)
+from airflow.providers.databricks.hooks.databricks import DatabricksHook
+
+try:
+    import httpx2
+    from fastmcp.client.transports import StreamableHttpTransport
+    from pydantic_ai.mcp import MCPToolset as PydanticAIMCPToolset
+
+    from airflow.providers.common.ai.toolsets.mcp import MCPToolset
+except ImportError as e:
+    raise AirflowOptionalProviderFeatureException(
+        "DatabricksUnityMCPToolset needs the 'common.ai' extra of the 
databricks provider: "
+        "pip install 'apache-airflow-providers-databricks[common.ai]'"
+    ) from e
+
+if TYPE_CHECKING:
+    from collections.abc import AsyncGenerator, AsyncIterator, Sequence
+
+    from pydantic_ai._run_context import RunContext
+    from pydantic_ai.toolsets.abstract import ToolsetTool
+
+    from airflow.sdk.execution_time.secrets_masker import mask_secret
+else:
+    try:
+        from airflow.sdk.log import mask_secret
+    except ImportError:
+        try:
+            from airflow.sdk.execution_time.secrets_masker import mask_secret
+        except ImportError:
+            from airflow.utils.log.secrets_masker import mask_secret
+
+GATEWAY_MCP_SERVICES_PATH = "ai-gateway/mcp-services"
+
+# Restricting each part to these characters keeps the name from adding a path, 
query or
+# fragment to the URL, so the request can only reach the service path on the 
connection's host.
+_SERVICE_NAME_PART = r"[A-Za-z0-9_-]+"
+_SERVICE_NAME = 
re.compile(rf"{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}\.{_SERVICE_NAME_PART}")
+
+
+def validate_service_name(service_name: str) -> None:
+    """
+    Raise ``ValueError`` unless ``service_name`` is a three-level 
``catalog.schema.service`` name.
+
+    Each part may contain only ASCII letters, digits, underscores and hyphens.
+    """
+    if not isinstance(service_name, str) or not 
_SERVICE_NAME.fullmatch(service_name):
+        raise ValueError(
+            f"Invalid Unity Gateway MCP Service name {service_name!r}: 
expected "
+            "'catalog.schema.service', each part made of ASCII letters, 
digits, '_' or '-'."
+        )
+
+
+def _parse_retry_after(value: str | None) -> float | None:
+    if value is None:
+        return None
+    try:
+        seconds = float(value)
+    except ValueError:
+        # The HTTP-date form is not worth parsing: the caller only uses this 
as a hint.
+        return None
+    return seconds if seconds >= 0 else None
+
+
+def _find_transport_error(exc: BaseException) -> httpx2.TransportError | None:
+    """Return the network error behind ``exc``, looking through causes and 
exception groups."""
+    seen: set[int] = set()
+    pending: list[BaseException] = [exc]
+    while pending:
+        current = pending.pop()
+        if id(current) in seen:
+            continue
+        seen.add(id(current))
+        if isinstance(current, httpx2.TransportError):
+            return current
+        # Duck-typed: the ExceptionGroup builtin is Python 3.11+, and anyio 
uses the backport on 3.10.
+        if isinstance(grouped := getattr(current, "exceptions", None), (list, 
tuple)):
+            pending.extend(e for e in grouped if isinstance(e, BaseException))
+        pending.extend(e for e in (current.__cause__, current.__context__) if 
e is not None)
+    return None
+
+
+class _DatabricksTokenAuth(httpx2.Auth):
+    """
+    Authenticate each gateway request with a token from the Databricks 
connection.
+
+    Asking the hook on every request, rather than once, lets OAuth tokens 
refresh during a long
+    agent run and on reconnection; the hook caches tokens until they are about 
to expire.
+
+    The MCP client reports every HTTP error from the gateway as the same 
generic error, so the
+    status and ``Retry-After`` of the last error response, and any failure to 
get a token, are
+    kept here for :class:`DatabricksUnityMCPToolset` to report what went 
wrong. With concurrent
+    calls on one toolset, the error reported for a failed call can be another 
call's.
+    """
+
+    def __init__(self, hook: DatabricksHook) -> None:
+        self._hook = hook
+        self.error_status: int | None = None
+        self.retry_after: float | None = None
+        self.token_error: Exception | None = None
+        self._token_lock = threading.Lock()
+
+    def get_token(self) -> str:
+        with self._token_lock:
+            token = self._hook._get_token(raise_error=False)
+        if not token:
+            raise ValueError(
+                f"Connection {self._hook.databricks_conn_id!r} has no 
token-based authentication "
+                "configured. This toolset sends a bearer token: use a personal 
access token, "
+                "service principal OAuth, Azure AD, or workload identity 
federation. Username and "
+                "password authentication is not supported."
+            )
+        # A personal access token is masked when the connection is fetched; 
mask minted
+        # OAuth tokens too, so they never reach task logs.
+        mask_secret(token)
+        return token
+
+    def take_error(self) -> tuple[int | None, float | None, Exception | None]:
+        """Return the last recorded error and forget it."""
+        error = (self.error_status, self.retry_after, self.token_error)
+        self.error_status = self.retry_after = self.token_error = None
+        return error
+
+    def _record(self, response: httpx2.Response) -> None:
+        # The MCP client tolerates some error responses, such as one to a 
notification, so an
+        # error must not outlive the next success, or a later failure would be 
reported as it.
+        if response.status_code >= 400:
+            self.error_status = response.status_code
+            self.retry_after = 
_parse_retry_after(response.headers.get("Retry-After"))
+        else:
+            self.error_status = self.retry_after = None
+
+    async def async_auth_flow(
+        self, request: httpx2.Request
+    ) -> AsyncGenerator[httpx2.Request, httpx2.Response]:
+        try:
+            # Not AirflowToolset.run_blocking: its lock is shared by every 
toolset, so a long SQL
+            # query in another toolset would hold up each request here. The 
connection is already
+            # resolved by _get_server, so fetching a token only calls the 
token endpoint, and the
+            # hook is this toolset's own, so a lock of its own is enough.
+            token = await asyncio.to_thread(self.get_token)
+        except Exception as e:
+            self.token_error = e
+            raise
+        request.headers["Authorization"] = f"Bearer {token}"
+        response = yield request
+        self._record(response)
+
+
+class DatabricksUnityMCPToolset(MCPToolset):
+    """
+    Give an agent the tools of a Unity Gateway MCP Service, authenticated as 
the connection's identity.
+
+    The service is named by its three-level Unity Catalog name, 
``catalog.schema.service``, and
+    reached at ``https://<workspace 
host>/ai-gateway/mcp-services/<catalog.schema.service>``. The
+    workspace host and the credentials both come from the Databricks 
connection, so Dag code holds
+    neither a gateway URL nor a token, and the token is only ever sent to that 
workspace.
+
+    The gateway runs every tool call as the identity of the connection's 
credentials (a user's
+    personal access token, a service principal, or an Azure AD / federated 
identity). That identity
+    needs ``EXECUTE`` on the MCP Service, ``USE CATALOG`` and ``USE SCHEMA`` 
on its catalog and
+    schema, and an assignment to the workspace. It sees only the tools 
selected for the service, and
+    the service's policies apply.
+
+    Tokens are fetched from the connection for each request, so OAuth tokens 
refresh during a long
+    agent run and on reconnection. Gateway errors are raised as
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPAccessDeniedError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPServiceNotFoundError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPThrottledError`,
+    
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPTransportError`,
 or, for any
+    other gateway error, 
:class:`~airflow.providers.databricks.exceptions.DatabricksUnityMCPError`.
+    Tool calls are never retried by this toolset, because a call interrupted 
after it was sent may
+    already have run.
+
+    .. code-block:: python
+
+        from airflow.providers.common.ai.operators.agent import AgentOperator
+        from airflow.providers.databricks.toolsets.unity_mcp import 
DatabricksUnityMCPToolset
+
+        AgentOperator(
+            task_id="ask_mcp_service",
+            prompt="Which tools do you have?",
+            llm_conn_id="pydanticai_default",
+            toolsets=[DatabricksUnityMCPToolset("main.default.my_mcp", 
databricks_conn_id="databricks")],
+        )
+
+    :param service_name: Three-level name of the MCP Service, 
``catalog.schema.service``. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param databricks_conn_id: Databricks connection whose host and 
credentials are used. Templated
+        when the toolset is passed to ``AgentOperator`` / ``@task.agent``.
+    :param tool_prefix: Optional prefix prepended to tool names.
+    """
+
+    agent_template_fields: Sequence[str] = ("_databricks_conn_id", 
"_service_name")
+
+    def __init__(
+        self,
+        service_name: str,
+        *,
+        databricks_conn_id: str = DatabricksHook.default_conn_name,
+        tool_prefix: str | None = None,
+    ) -> None:
+        super().__init__(databricks_conn_id, tool_prefix=tool_prefix)
+        self._databricks_conn_id = databricks_conn_id
+        self._service_name = service_name
+        self._auth: _DatabricksTokenAuth | None = None
+        # A templated name is checked once it has been rendered, in 
_get_server.
+        if "{{" not in service_name:
+            validate_service_name(service_name)
+
+    @property
+    def id(self) -> str:
+        return 
f"databricks-unity-mcp-{self._databricks_conn_id}-{self._service_name}"
+
+    def get_service_url(self, hook: DatabricksHook) -> str:
+        """Return the gateway URL of the MCP Service on the connection's 
workspace."""
+        validate_service_name(self._service_name)
+        if not hook.host:
+            raise ValueError(f"Connection {self._databricks_conn_id!r} has no 
workspace host.")
+        return 
hook._endpoint_url(f"{GATEWAY_MCP_SERVICES_PATH}/{self._service_name}")
+
+    def _get_server(self) -> Any:

Review Comment:
   This overrides private `MCPToolset` internals (`_get_server`, `_server`, 
`_tool_prefix`) from another provider. Could we either add a public per-request 
auth hook to common.ai, or subclass `AirflowToolset` directly?



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