This is an automated email from the ASF dual-hosted git repository.
kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 389b4e1751f Allow templated connection IDs in agent toolsets (#73578)
389b4e1751f is described below
commit 389b4e1751f32e32c6a40356c7d3e41893b7798a
Author: Kaxil Naik <[email protected]>
AuthorDate: Wed Sep 23 22:47:33 2026 +0100
Allow templated connection IDs in agent toolsets (#73578)
The connection IDs of Common AI agent toolsets are now Jinja templates,
rendered for each task instance just before it runs. One toolset definition can
then reach a different system depending on the run:
- **Per environment:** the same Dag reads the staging warehouse in staging
and the production one in production, with the environment in a Variable:
`SQLToolset(db_conn_id="warehouse_{{ var.value.environment }}")`.
- **Per unit of work:** a mapped `@task.agent` gives each map index its own
connection, e.g. one per customer for customer-facing analytics, where each
customer's rows sit behind their own database role:
```python
@task.agent(
llm_conn_id="pydanticai_default",
toolsets=[SQLToolset(db_conn_id="analytics_{{ task.op_kwargs.customer
}}")],
)
def report(customer: str) -> str:
return f"Summarize this month's orders for {customer}."
report.expand(customer=customers())
```
Until now toolsets took their connection when the Dag was parsed.
`@task.llm_sql` already templates `db_conn_id`, but an agent's toolsets could
not, so every mapped agent shared one connection.
What is templated: `SQLToolset.db_conn_id`, `MCPToolset.mcp_conn_id`, and a
`HookToolset`'s hook connection ID. `HookToolset` reads the attribute the
hook's `conn_name_attr` names (`postgres_conn_id`, ...), and falls back to
`conn_id` for hooks such as `WasbHook` that keep it there. Only connection IDs
are templated.
## Design rationale
**Why not add `toolsets` to `AgentOperator.template_fields`?** Template
fields are serialized into the Dag. A toolset's repr is what gets serialized,
and pydantic-ai's wrappers (`.prefixed()`, `.filtered()`) are dataclasses whose
repr can embed function addresses, so the Dag hash would change on every parse.
`toolsets` stays out of `template_fields`, and the operator renders the
toolsets itself.
**Each task instance renders a copy.** `MappedOperator.unmap` passes the
partial's toolset objects straight to every unmapped task, and `dag.test()`
runs every task in one process. Rendering in place would hand map index 0's
connection to map index 1. Each opt-in leaf toolset is copied before rendering
(for `HookToolset`, the hook too). Wrappers and `Toolset` capabilities are
walked with pydantic-ai's `visit_and_replace`, and only when something in them
is templated, so an untemplated [...]
**The opt-in attribute is `agent_template_fields`, not `template_fields`.**
Airflow's templater renders any object that carries `template_fields` in place,
wherever it is nested inside another template field. `agent_params` is a
template field, and `agent_params["toolsets"]` is a supported way to pass
toolsets, so the familiar name would bring the leak back through that path.
Third-party toolsets opt in by declaring the same attribute.
**Rendering hangs off `_do_render_template_fields`.** A mapped task never
calls `render_template_fields`: `MappedOperator` renders the unmapped task
through `_do_render_template_fields`. `KubernetesPodOperator` hooks the same
method for the same reason.
The rendered connection is not recorded anywhere else, so each task
instance logs it once, e.g. `Rendered toolset sql-analytics_acme`.
## Screenshots
A two-customer demo on a real scheduler and API server, each customer with
its own SQLite connection, and a `test` model that calls every tool. Map index
0 renders the `acme` connection for both the SQL toolset and the hook toolset;
map index 1 renders `globex`:


Each agent's tool results come from its own customer's database:


## Gotchas
- **Build the connection ID from values the Dag controls (a Variable,
upstream task output), not `params` or `dag_run.conf`.** A task can read any
connection it names, so a template driven by trigger input lets whoever
triggers the Dag pick the database (for an MCP `stdio` connection, the command
that runs on the worker). The docs say so next to each example.
- `{{ customer }}` does not work: the task's arguments are not template
variables. Use `{{ task.op_kwargs.customer }}`. With
`AgentOperator.partial(...).expand(prompt=...)`, the connection has to come
from the map index; the docs show that form and its ordering caveat.
- `HookToolset.id` now includes the connection ID
(`hook-PostgresHook-analytics_acme`, previously `hook-PostgresHook`). The
toolset id is part of the durable-execution step fingerprint, so a
`durable=True` task that retries across the upgrade misses its cache once.
- Not templated: `allowed_tables` (validated when the toolset is created,
so a template stays a literal), `DataFusionToolset`, a `Toolset` capability
built from a callable, and hooks that keep their connection ID under some other
attribute. A hook that looks its connection up in `__init__` fails at Dag parse
time, because the template is not a connection ID yet.
---
providers/common/ai/docs/toolsets/hook.rst | 41 +++-
providers/common/ai/docs/toolsets/mcp.rst | 8 +-
providers/common/ai/docs/toolsets/sql.rst | 92 +++++++-
.../airflow/providers/common/ai/operators/agent.py | 111 ++++++++-
.../airflow/providers/common/ai/toolsets/hook.py | 41 +++-
.../airflow/providers/common/ai/toolsets/mcp.py | 9 +-
.../airflow/providers/common/ai/toolsets/sql.py | 10 +-
.../tests/unit/common/ai/decorators/test_agent.py | 7 +-
.../tests/unit/common/ai/operators/test_agent.py | 248 ++++++++++++++++++++-
.../ai/tests/unit/common/ai/toolsets/test_hook.py | 55 +++++
10 files changed, 597 insertions(+), 25 deletions(-)
diff --git a/providers/common/ai/docs/toolsets/hook.rst
b/providers/common/ai/docs/toolsets/hook.rst
index 0a4fd657e61..5f387b872ee 100644
--- a/providers/common/ai/docs/toolsets/hook.rst
+++ b/providers/common/ai/docs/toolsets/hook.rst
@@ -43,10 +43,49 @@ For each listed method, the introspection engine:
3. Enriches parameter descriptions from Sphinx ``:param:`` or Google
``Args:`` blocks.
+.. _hook-toolset-templated-connection:
+
+Templated connection IDs
+------------------------
+
+The hook's connection ID is a Jinja template, rendered for each task instance
just
+before it runs, so one toolset can reach a different system depending on the
run:
+``PostgresHook(postgres_conn_id="warehouse_{{ var.value.environment }}")``
switches
+between staging and production, and a mapped task can give each map index its
own
+connection:
+
+.. code-block:: python
+
+ @task.agent(
+ llm_conn_id="pydanticai_default",
+ toolsets=[
+ HookToolset(
+ PostgresHook(postgres_conn_id="analytics_{{
task.op_kwargs.customer }}"),
+ allowed_methods=["get_records"],
+ )
+ ],
+ )
+ def report(customer: str) -> str:
+ return f"Summarize this month's orders for {customer}."
+
+
+ report.expand(customer=customers())
+
+Each task instance runs against a copy of the hook with its rendered connection
+ID; the hook in the Dag file keeps the template. ``HookToolset`` reads the ID
+from the attribute ``conn_name_attr`` names, or from ``conn_id`` for hooks
such as
+``WasbHook`` that keep it there; a hook that stores it anywhere else is not
+templated. This works for hooks that read their connection when a method is
+called, which is what Airflow hooks are expected to do: a hook that looks the
+connection up in its constructor fails at Dag parse time, because the template
+is not a connection ID yet. The same warning as for ``SQLToolset`` applies:
build
+the ID from values the Dag controls, not from ``params`` or ``dag_run.conf``
(see
+:ref:`sql-toolset-templated-connection`).
+
Parameters
----------
-- ``hook``: An instantiated Airflow Hook.
+- ``hook``: An instantiated Airflow Hook. Its connection ID is templated.
- ``allowed_methods``: Method names to expose as tools. Required. Methods
are validated with ``hasattr`` + ``callable`` at instantiation time.
- ``tool_name_prefix``: Optional prefix prepended to each tool name
diff --git a/providers/common/ai/docs/toolsets/mcp.rst
b/providers/common/ai/docs/toolsets/mcp.rst
index 05ded1632fd..639b5a665e1 100644
--- a/providers/common/ai/docs/toolsets/mcp.rst
+++ b/providers/common/ai/docs/toolsets/mcp.rst
@@ -48,7 +48,13 @@ Requires the ``mcp`` extra: ``pip install
"apache-airflow-providers-common-ai[mc
Parameters
----------
-- ``mcp_conn_id``: Airflow connection ID for the MCP server.
+- ``mcp_conn_id``: Airflow connection ID for the MCP server. Templated when the
+ toolset is passed to ``AgentOperator`` / ``@task.agent``, like
+ ``SQLToolset.db_conn_id`` (see :ref:`sql-toolset-templated-connection`).
Build it
+ from values the Dag controls, never from ``params`` or ``dag_run.conf``: a
``stdio``
+ connection runs its ``Extra.command`` on the worker, so whoever picks the
+ connection picks the command. A ``token_provider`` or ``env_provider`` is
+ shared by every connection the template renders to.
- ``tool_prefix``: Optional prefix prepended to tool names to avoid
collisions when using multiple MCP servers (e.g. ``"weather"`` produces
``"weather_get_forecast"``).
diff --git a/providers/common/ai/docs/toolsets/sql.rst
b/providers/common/ai/docs/toolsets/sql.rst
index 1dc78217a87..fa82aed4559 100644
--- a/providers/common/ai/docs/toolsets/sql.rst
+++ b/providers/common/ai/docs/toolsets/sql.rst
@@ -86,10 +86,100 @@ fall back to ``schema``, and table-name matching is
case-insensitive (databases
reflect identifiers in their own case). For tables in a different *database*,
use
a separate toolset whose connection points at that database.
+.. _sql-toolset-templated-connection:
+
+Templated connection IDs
+------------------------
+
+``db_conn_id`` is a Jinja template, rendered for each task instance just
before it
+runs, so one toolset definition can reach a different database depending on
where
+and for what the task runs:
+
+- **Per environment.** The same Dag reads the staging warehouse in staging and
the
+ production one in production, with the environment name kept in a Variable:
+ ``SQLToolset(db_conn_id="warehouse_{{ var.value.environment }}")``.
+- **Per unit of work.** A mapped task gives each map index its own connection
--
+ one per customer, region, or shard -- as in the example below.
+
+Each task instance renders its own copy of the toolset, so the object in the
Dag
+file keeps its template and no rendered connection carries over to another task
+instance. The task log records which connection each instance got, as a
+``Rendered toolset sql-warehouse_prod`` line. A toolset wrapped with
+``.prefixed()`` or ``.filtered()``, passed as a ``Toolset`` capability, or
passed
+in ``agent_params["toolsets"]`` is rendered the same way. A ``Toolset``
+capability built from a callable is resolved when the run starts and is not
+rendered. ``MCPToolset.mcp_conn_id`` and ``HookToolset``'s hook connection ID
are
+templated the same way (see :ref:`hook-toolset-templated-connection`).
+
+.. warning::
+
+ Build the connection ID from values the Dag controls -- a Variable,
upstream
+ task output -- not from ``params`` or ``dag_run.conf``. Whoever triggers
the Dag
+ controls those, and a task can read any connection it names, so a templated
+ ``db_conn_id`` taken from trigger input lets the trigger pick the database.
+
+Only the connection ID is templated. ``allowed_tables`` is validated when the
+toolset is created, so a template in it stays a literal table name.
+``DataFusionToolset`` takes data source configs and is not templated.
+
+One connection per customer
+^^^^^^^^^^^^^^^^^^^^^^^^^^^
+
+Customer-facing analytics must only ever read one customer's rows. The boundary
+that holds is the database's own: a role or database per customer, reached
+through its own Airflow connection. A mapped agent task can give each
customer's
+task instance that customer's connection:
+
+.. code-block:: python
+
+ from airflow.providers.common.ai.toolsets.sql import SQLToolset
+ from airflow.sdk import dag, task
+
+
+ @dag
+ def customer_reports():
+ @task
+ def customers() -> list[str]:
+ return ["acme", "globex"]
+
+ @task.agent(
+ llm_conn_id="pydanticai_default",
+ toolsets=[SQLToolset(db_conn_id="analytics_{{
task.op_kwargs.customer }}")],
+ )
+ def report(customer: str) -> str:
+ return f"Summarize this month's orders for {customer}."
+
+ report.expand(customer=customers())
+
+
+ customer_reports()
+
+The ``acme`` task instance queries through ``analytics_acme`` and the
``globex``
+one through ``analytics_globex``. Write ``{{ task.op_kwargs.customer }}``, not
+``{{ customer }}``: the task's arguments are not template variables, and the
+undefined name fails the task.
+
+``AgentOperator`` mapped over prompts has no customer argument to read, so the
+connection has to come from the map index. That works when the prompts are
built
+from the same list, in the same order, as the one the template indexes:
+
+.. code-block:: python
+
+ names = customers()
+ AgentOperator.partial(
+ task_id="report",
+ llm_conn_id="pydanticai_default",
+ toolsets=[SQLToolset(db_conn_id="analytics_{{
ti.xcom_pull(task_ids='customers')[ti.map_index] }}")],
+ ).expand(prompt=names.map(lambda name: f"Summarize this month's orders for
{name}."))
+
+Prefer ``@task.agent`` where you can: ``task.op_kwargs`` names the customer
+directly instead of relying on the two lists lining up.
+
Parameters
----------
-- ``db_conn_id``: Airflow connection ID for the database.
+- ``db_conn_id``: Airflow connection ID for the database. Templated (see
+ :ref:`sql-toolset-templated-connection`).
- ``allowed_tables``: Restrict the agent to a fixed set of tables. Omit the
argument (the default) to expose all tables in ``schema``. No value means
allow-all: ``None`` and an empty list both raise ``ValueError``, so an
allow-list
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
index 03be4e1af8f..b7b64d1a42a 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
@@ -18,8 +18,9 @@
from __future__ import annotations
+import copy
import json
-from collections.abc import Sequence
+from collections.abc import Iterable, Sequence
from dataclasses import replace
from datetime import timedelta
from functools import cached_property
@@ -27,6 +28,7 @@ from typing import TYPE_CHECKING, Any, ClassVar
from pydantic import BaseModel
from pydantic_ai.capabilities import Toolset
+from pydantic_ai.toolsets.abstract import AbstractToolset
from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook
from airflow.providers.common.ai.mixins.cancellable_run import
CancellableAgentRunMixin
@@ -38,7 +40,7 @@ from airflow.providers.common.ai.observability import (
from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset
from airflow.providers.common.ai.utils.logging import log_run_summary,
wrap_toolsets_for_logging
from airflow.providers.common.ai.utils.output_type import
rehydrate_pydantic_output
-from airflow.providers.common.ai.utils.toolsets import find_toolset
+from airflow.providers.common.ai.utils.toolsets import find_toolset,
iter_toolsets
from airflow.providers.common.ai.utils.usage import coerce_usage_limits
from airflow.providers.common.compat.sdk import (
AirflowOptionalProviderFeatureException,
@@ -57,9 +59,9 @@ except ImportError: # pragma: no cover - cores before the
worker-side registrat
_CORE_WALKER = False
if TYPE_CHECKING:
+ import jinja2
from pydantic_ai import Agent
from pydantic_ai.messages import ModelMessage
- from pydantic_ai.toolsets.abstract import AbstractToolset
from pydantic_ai.usage import UsageLimits
from airflow.providers.common.ai.durable.base import DurableStorageProtocol
@@ -101,6 +103,18 @@ class HITLReviewLink(BaseOperatorLink):
)
+def _is_concrete_toolset_capability(capability: Any) -> bool:
+ """Whether *capability* is a ``Toolset`` holding a toolset, not a callable
factory resolved per run."""
+ return isinstance(capability, Toolset) and isinstance(capability.toolset,
AbstractToolset)
+
+
+def _declares_agent_template_fields(toolset: Any) -> bool:
+ """Whether *toolset*, or a toolset it wraps or combines, has connection
IDs to render."""
+ return isinstance(toolset, AbstractToolset) and any(
+ getattr(leaf, "agent_template_fields", None) for leaf in
iter_toolsets(toolset)
+ )
+
+
def _build_code_mode() -> Any:
"""
Return a pydantic-ai-harness ``CodeMode`` capability, or raise if not
installed.
@@ -162,7 +176,15 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
directly. The class must be defined at module scope -- nested classes
cannot be deserialized from XCom.
:param toolsets: List of pydantic-ai toolsets the agent can use
- (e.g. ``SQLToolset``, ``HookToolset``).
+ (e.g. ``SQLToolset``, ``HookToolset``). The connection IDs of
+ ``SQLToolset``, ``MCPToolset`` and ``HookToolset`` (its hook's
+ ``conn_name_attr``) are templated, e.g.
+ ``SQLToolset(db_conn_id="warehouse_{{ var.value.environment }}")`` per
+ environment, or ``"tenant_{{ task.op_kwargs.customer }}"`` per map
index of
+ a mapped ``@task.agent``. Each task instance renders its own copy and
logs
+ the rendered toolset id; the toolset object in the Dag file is not
+ modified. Derive the connection ID from values the Dag controls rather
than
+ ``params`` or ``dag_run.conf``, which whoever triggers the Dag
controls.
:param enable_tool_logging: When ``True`` (default), wraps each toolset in
a
``LoggingToolset`` that logs tool calls with timing at INFO level and
arguments at DEBUG level. Set to ``False`` to disable.
@@ -393,7 +415,7 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator,
HITLReviewMixin):
"""
candidates = list(self.toolsets or [])
for capability in self.agent_params.get("capabilities") or ():
- if isinstance(capability, Toolset) and not
callable(capability.toolset):
+ if _is_concrete_toolset_capability(capability):
candidates.append(capability.toolset)
if find_toolset(candidates, SandboxToolset) is None:
return
@@ -409,6 +431,78 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
f"Drop {flag}, or move the sandbox work into its own task."
)
+ def _do_render_template_fields(
+ self,
+ parent: Any,
+ template_fields: Iterable[str],
+ context: Context,
+ jinja_env: jinja2.Environment,
+ seen_oids: set[int],
+ ) -> None:
+ super()._do_render_template_fields(parent, template_fields, context,
jinja_env, seen_oids)
+ # Hooked here rather than in render_template_fields because a mapped
task never calls
+ # that one -- MappedOperator renders through
_do_render_template_fields on the unmapped task.
+ if parent is self:
+ self._render_toolsets(context, jinja_env, seen_oids)
+
+ def _render_toolsets(self, context: Context, jinja_env:
jinja2.Environment, seen_oids: set[int]) -> None:
+ """
+ Render the connection IDs of toolsets that declare
``agent_template_fields``.
+
+ ``toolsets`` is not itself a template field: serializing it would put
each
+ toolset's repr -- which for pydantic-ai's dataclass toolsets embeds
function
+ addresses -- into the Dag hash and the rendered-fields view. Instead,
each leaf
+ toolset that opts in (``SQLToolset``, ``MCPToolset``, ``HookToolset``)
is
+ rendered here, found with pydantic-ai's ``visit_and_replace`` inside
+ ``.prefixed()`` / ``.filtered()`` wrappers, ``Toolset`` capabilities,
and a
+ ``toolsets`` list passed through ``agent_params``. A ``Toolset``
capability
+ backed by a callable factory is resolved per run and is not rendered.
+
+ A rendered *copy* replaces the original, which is left untouched:
mapped task
+ instances and ``dag.test()`` share one toolset object across runs in
the same
+ process, and rendering it in place would hand one map index's
connection to
+ the next. That is also why the opt-in is ``agent_template_fields`` and
not
+ ``template_fields``: Airflow's templater renders any object carrying
+ ``template_fields`` in place wherever it sits inside another template
field,
+ such as ``agent_params``.
+ """
+
+ def render(toolset: AbstractToolset[Any]) -> AbstractToolset[Any]:
+ fields = getattr(toolset, "agent_template_fields", None)
+ if not fields:
+ return toolset
+ rendered = copy.copy(toolset)
+ self._do_render_template_fields(rendered, fields, context,
jinja_env, seen_oids)
+ # The rendered connection is recorded nowhere else, so this line
is the audit trail
+ # of which connection this task instance's agent was given.
@task.agent renders a
+ # second time, when the id no longer changes, so this logs once
per task instance.
+ if rendered.id != toolset.id:
+ self.log.info("Rendered toolset %s", rendered.id)
+ return rendered
+
+ def render_all(toolsets: list[Any]) -> list[Any]:
+ # Leave anything without a templated leaf alone: rebuilding a
wrapper via
+ # visit_and_replace breaks wrapper subclasses with their own
__init__.
+ return [
+ toolset.visit_and_replace(render) if
_declares_agent_template_fields(toolset) else toolset
+ for toolset in toolsets
+ ]
+
+ if self.toolsets:
+ self.toolsets = render_all(self.toolsets)
+ agent_params = dict(self.agent_params)
+ if agent_params.get("toolsets"):
+ agent_params["toolsets"] = render_all(agent_params["toolsets"])
+ if agent_params.get("capabilities"):
+ agent_params["capabilities"] = [
+ replace(capability,
toolset=capability.toolset.visit_and_replace(render))
+ if _is_concrete_toolset_capability(capability)
+ and _declares_agent_template_fields(capability.toolset)
+ else capability
+ for capability in agent_params["capabilities"]
+ ]
+ self.agent_params = agent_params
+
@cached_property
def llm_hook(self) -> PydanticAIHook:
"""Return PydanticAIHook for the configured LLM connection."""
@@ -469,18 +563,13 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
unchanged, as does a ``Toolset`` holding a callable factory rather
than a
concrete toolset (only a concrete toolset can be wrapped here).
"""
- # pydantic-ai (and the pydantic-ai-importing CachingToolset) are
imported
- # lazily to keep them out of DAG-parse-time imports, matching
- # ``_build_durable_toolsets`` and the rest of this module.
- from pydantic_ai.toolsets.abstract import AbstractToolset
-
from airflow.providers.common.ai.durable.caching_toolset import
CachingToolset
rewrapped: list[Any] = []
for capability in capabilities:
# ``Toolset.toolset`` can be a concrete toolset or a callable
factory
# resolved per run; only a concrete toolset can be wrapped here.
- if isinstance(capability, Toolset) and
isinstance(capability.toolset, AbstractToolset):
+ if _is_concrete_toolset_capability(capability):
cached = CachingToolset(wrapped=capability.toolset,
storage=storage, counter=counter)
rewrapped.append(replace(capability, toolset=cached))
continue
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
index 13412b82c35..609deba8e9a 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
@@ -18,6 +18,7 @@
from __future__ import annotations
+import copy
import inspect
import re
import types
@@ -33,7 +34,7 @@ from airflow.providers.common.ai.utils.tool_definition import
(
)
if TYPE_CHECKING:
- from collections.abc import Callable
+ from collections.abc import Callable, Sequence
from pydantic_ai._run_context import RunContext
@@ -59,13 +60,22 @@ class HookToolset(AbstractToolset[Any]):
hook to build :class:`~pydantic_ai.tools.ToolDefinition` objects that an
LLM
agent can call.
- :param hook: An instantiated Airflow Hook.
+ :param hook: An instantiated Airflow Hook. Its connection ID -- the
attribute
+ the hook's ``conn_name_attr`` names, such as ``postgres_conn_id`` -- is
+ templated when the toolset is passed to ``AgentOperator`` /
``@task.agent``,
+ so ``HookToolset(PostgresHook(postgres_conn_id="tenant_{{ ... }}"),
...)``
+ reaches a different database per task instance. The hook in the Dag
file
+ is not modified; each task instance gets a copy.
:param allowed_methods: Method names to expose as tools. Required —
auto-discovery is intentionally not supported for safety.
:param tool_name_prefix: Optional prefix prepended to each tool name
(e.g. ``"s3_"`` → ``"s3_list_keys"``).
"""
+ # Rendered, on a copy, by AgentOperator. Deliberately not
``template_fields``, which
+ # Airflow's templater would render in place wherever the toolset is nested.
+ agent_template_fields: Sequence[str] = ("conn_id",)
+
def __init__(
self,
hook: BaseHook,
@@ -88,11 +98,34 @@ class HookToolset(AbstractToolset[Any]):
self._hook = hook
self._allowed_methods = allowed_methods
self._tool_name_prefix = tool_name_prefix
- self._id = f"hook-{type(hook).__name__}"
+ # The attribute holding the hook's connection ID, e.g.
``postgres_conn_id``. Some hooks
+ # name one attribute in conn_name_attr but keep the ID in ``conn_id``
(WasbHook,
+ # KubernetesHook), so fall back to that.
+ conn_attr: str | None = getattr(hook, "conn_name_attr", None)
+ if conn_attr is None or not hasattr(hook, conn_attr):
+ conn_attr = "conn_id" if hasattr(hook, "conn_id") else None
+ self._conn_attr = conn_attr
+
+ @property
+ def conn_id(self) -> str | None:
+ """The hook's connection ID, or ``None`` when the hook keeps it under
neither attribute."""
+ return getattr(self._hook, self._conn_attr, None) if self._conn_attr
else None
+
+ @conn_id.setter
+ def conn_id(self, value: str) -> None:
+ if self._conn_attr is None:
+ raise AttributeError(f"{type(self._hook).__name__} keeps no
connection ID to set.")
+ # Set on a copy: the hook in the Dag file backs every task instance
that shares this
+ # toolset, so writing the rendered ID onto it would carry one
instance's connection
+ # into the next.
+ hook = copy.copy(self._hook)
+ setattr(hook, self._conn_attr, value)
+ self._hook = hook
@property
def id(self) -> str:
- return self._id
+ name = type(self._hook).__name__
+ return f"hook-{name}-{self.conn_id}" if self.conn_id else
f"hook-{name}"
async def get_tools(self, ctx: RunContext[Any]) -> dict[str,
ToolsetTool[Any]]:
tools: dict[str, ToolsetTool[Any]] = {}
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
index 158d3bd7c35..5a1d7281477 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
@@ -24,7 +24,7 @@ from pydantic_ai.toolsets.abstract import AbstractToolset,
ToolsetTool
from typing_extensions import Self
if TYPE_CHECKING:
- from collections.abc import Callable
+ from collections.abc import Callable, Sequence
from pydantic_ai._run_context import RunContext
@@ -61,7 +61,8 @@ class MCPToolset(AbstractToolset[Any]):
merged over the connection's static ``Extra.env`` (``env_provider`` wins on
key conflicts).
- :param mcp_conn_id: Airflow connection ID for the MCP server.
+ :param mcp_conn_id: Airflow connection ID for the MCP server. Templated
when
+ the toolset is passed to ``AgentOperator`` / ``@task.agent``.
:param tool_prefix: Optional prefix prepended to tool names
(e.g. ``"weather"`` → ``"weather_get_forecast"``).
:param token_provider: Optional zero-argument callable returning a bearer
@@ -73,6 +74,10 @@ class MCPToolset(AbstractToolset[Any]):
first time this toolset establishes a connection.
"""
+ # Rendered, on a copy, by AgentOperator. Deliberately not
``template_fields``, which
+ # Airflow's templater would render in place wherever the toolset is nested.
+ agent_template_fields: Sequence[str] = ("_mcp_conn_id",)
+
def __init__(
self,
mcp_conn_id: str,
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
index bbfd32108b6..6d312847384 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
@@ -50,6 +50,8 @@ from airflow.providers.common.ai.utils.tool_definition import
build_args_validat
from airflow.providers.common.compat.sdk import BaseHook
if TYPE_CHECKING:
+ from collections.abc import Sequence
+
from pydantic_ai._run_context import RunContext
# Sentinel distinguishing "caller did not pass ``allowed_tables``" (expose
every
@@ -176,7 +178,9 @@ class SQLToolset(AbstractToolset[Any]):
failure -- exhausts the retries and fails the task for Airflow to retry.
The
toolset does not inspect the error type or message.
- :param db_conn_id: Airflow connection ID for the database.
+ :param db_conn_id: Airflow connection ID for the database. Templated when
the
+ toolset is passed to ``AgentOperator`` / ``@task.agent``, so each task
+ instance can reach its own database, e.g. one connection per customer.
:param allowed_tables: Restrict the agent to a fixed set of tables. Omit
the
argument (the default) to expose every table in ``schema``. No *value*
means
allow-all: ``None`` and an empty list both raise ``ValueError``, so an
allow-list
@@ -246,6 +250,10 @@ class SQLToolset(AbstractToolset[Any]):
its projection rather than page through the table.
"""
+ # Rendered, on a copy, by AgentOperator. Deliberately not
``template_fields``, which
+ # Airflow's templater would render in place wherever the toolset is nested.
+ agent_template_fields: Sequence[str] = ("_db_conn_id",)
+
def __init__(
self,
db_conn_id: str,
diff --git a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
index 2be4009e078..eae3b6bf393 100644
--- a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
+++ b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py
@@ -21,6 +21,7 @@ from unittest.mock import ANY, MagicMock, patch
import pytest
from pydantic import BaseModel
from pydantic_ai.messages import ImageUrl
+from pydantic_ai.toolsets.function import FunctionToolset
from airflow.providers.common.ai.decorators.agent import
_AgentDecoratedOperator
from airflow.providers.common.ai.toolsets.logging import LoggingToolset
@@ -161,13 +162,13 @@ class TestAgentDecoratedOperator:
mock_agent.run_sync.return_value = make_mock_run_result("result")
mock_hook_cls.get_hook.return_value.create_agent.return_value =
mock_agent
- mock_toolset = MagicMock()
+ toolset = FunctionToolset()
op = _AgentDecoratedOperator(
task_id="test",
python_callable=lambda: "Do something",
llm_conn_id="my_llm",
- toolsets=[mock_toolset],
+ toolsets=[toolset],
)
op.execute(context=_make_context())
@@ -175,7 +176,7 @@ class TestAgentDecoratedOperator:
passed_toolsets = create_call[1]["toolsets"]
assert len(passed_toolsets) == 1
assert isinstance(passed_toolsets[0], LoggingToolset)
- assert passed_toolsets[0].wrapped is mock_toolset
+ assert passed_toolsets[0].wrapped is toolset
@requires_typed_xcom
@patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
index 129749b1498..f6d65279ef2 100644
--- a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
+++ b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
@@ -19,6 +19,7 @@ from __future__ import annotations
import sys
from datetime import timedelta
from decimal import Decimal
+from types import SimpleNamespace
from unittest.mock import ANY, MagicMock, patch
import pytest
@@ -39,6 +40,7 @@ from pydantic_ai.messages import (
from pydantic_ai.models.function import AgentInfo, FunctionModel
from pydantic_ai.toolsets.combined import CombinedToolset
from pydantic_ai.toolsets.function import FunctionToolset
+from pydantic_ai.toolsets.wrapper import WrapperToolset
from pydantic_ai.usage import RequestUsage, UsageLimits
from airflow.providers.common.ai.durable.base import DurableStorageProtocol
@@ -47,9 +49,14 @@ from airflow.providers.common.ai.durable.step_counter import
DurableStepCounter
from airflow.providers.common.ai.durable.storage import DurableStorage
from airflow.providers.common.ai.operators.agent import AgentOperator,
HITLReviewLink, _build_code_mode
from airflow.providers.common.ai.sandbox.base import SandboxBackend
+from airflow.providers.common.ai.toolsets.hook import HookToolset
from airflow.providers.common.ai.toolsets.logging import LoggingToolset
+from airflow.providers.common.ai.toolsets.mcp import MCPToolset
from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset
-from airflow.providers.common.compat.sdk import
AirflowOptionalProviderFeatureException
+from airflow.providers.common.ai.toolsets.sql import SQLToolset
+from airflow.providers.common.ai.utils.toolsets import find_toolset
+from airflow.providers.common.compat.sdk import
AirflowOptionalProviderFeatureException, BaseHook
+from airflow.sdk import DAG, task
from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS,
AIRFLOW_V_3_3_PLUS
@@ -188,6 +195,245 @@ class TestAgentOperatorTemplateFields:
assert set(AgentOperator.template_fields) == expected
+class _TenantHook(BaseHook):
+ conn_name_attr = "tenant_conn_id"
+
+ def __init__(self, tenant_conn_id: str):
+ super().__init__()
+ self.tenant_conn_id = tenant_conn_id
+
+ def get_records(self, sql: str) -> list:
+ """Run a query."""
+ return []
+
+
+class TestAgentOperatorToolsetTemplating:
+ """Connection IDs on SQLToolset / MCPToolset / HookToolset render per task
instance, on a copy."""
+
+ CONTEXT = {"params": {"customer": "acme"}}
+
+ def test_sql_toolset_conn_id_is_rendered_on_a_copy(self):
+ toolset = SQLToolset(db_conn_id="tenant_{{ params.customer }}")
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
toolsets=[toolset])
+
+ op.render_template_fields(self.CONTEXT)
+
+ (rendered,) = op.toolsets
+ assert rendered._db_conn_id == "tenant_acme"
+ assert rendered.id == "sql-tenant_acme"
+ assert rendered is not toolset
+ assert toolset._db_conn_id == "tenant_{{ params.customer }}"
+
+ def test_shared_toolset_renders_independently_per_task(self):
+ """Mapped task instances and dag.test() share one toolset object in a
process;
+ rendering it in place would hand the first customer's connection to
the next."""
+ shared = SQLToolset(db_conn_id="tenant_{{ params.customer }}")
+ first = AgentOperator(task_id="a", prompt="p", llm_conn_id="llm",
toolsets=[shared])
+ second = AgentOperator(task_id="b", prompt="p", llm_conn_id="llm",
toolsets=[shared])
+
+ first.render_template_fields({"params": {"customer": "acme"}})
+ second.render_template_fields({"params": {"customer": "globex"}})
+
+ assert first.toolsets[0]._db_conn_id == "tenant_acme"
+ assert second.toolsets[0]._db_conn_id == "tenant_globex"
+
+ def test_hook_toolset_conn_id_is_rendered_on_a_copy(self):
+ hook = _TenantHook(tenant_conn_id="tenant_{{ params.customer }}")
+ op = AgentOperator(
+ task_id="t",
+ prompt="p",
+ llm_conn_id="llm",
+ toolsets=[HookToolset(hook, allowed_methods=["get_records"])],
+ )
+
+ op.render_template_fields(self.CONTEXT)
+
+ assert op.toolsets[0].id == "hook-_TenantHook-tenant_acme"
+ assert hook.tenant_conn_id == "tenant_{{ params.customer }}"
+
+ def test_mcp_toolset_conn_id_is_rendered(self):
+ op = AgentOperator(
+ task_id="t",
+ prompt="p",
+ llm_conn_id="llm",
+ toolsets=[MCPToolset(mcp_conn_id="mcp_{{ params.customer }}")],
+ )
+
+ op.render_template_fields(self.CONTEXT)
+
+ assert op.toolsets[0].id == "mcp-mcp_acme"
+
+ @pytest.mark.parametrize(
+ "wrap",
+ [
+ pytest.param(lambda ts: ts.prefixed("crm"), id="prefixed"),
+ pytest.param(lambda ts: ts.filtered(lambda ctx, tool: True),
id="filtered"),
+ pytest.param(lambda ts: CombinedToolset([FunctionToolset(), ts]),
id="combined"),
+ ],
+ )
+ def test_toolset_inside_wrapper_is_rendered(self, wrap):
+ inner = SQLToolset(db_conn_id="tenant_{{ params.customer }}")
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
toolsets=[wrap(inner)])
+
+ op.render_template_fields(self.CONTEXT)
+
+ found = find_toolset(op.toolsets, SQLToolset)
+ assert found is not None
+ assert found._db_conn_id == "tenant_acme"
+ assert inner._db_conn_id == "tenant_{{ params.customer }}"
+
+ def test_toolset_capability_is_rendered(self):
+ capability = Toolset(SQLToolset(db_conn_id="tenant_{{ params.customer
}}"))
+ op = AgentOperator(
+ task_id="t", prompt="p", llm_conn_id="llm",
agent_params={"capabilities": [capability]}
+ )
+
+ op.render_template_fields(self.CONTEXT)
+
+ (rendered,) = op.agent_params["capabilities"]
+ assert rendered.toolset._db_conn_id == "tenant_acme"
+ assert capability.toolset._db_conn_id == "tenant_{{ params.customer }}"
+
+ def test_callable_toolset_capability_is_left_as_is(self):
+ """A factory resolved per run has no toolset to render until the run
starts."""
+ capability = Toolset(lambda ctx: SQLToolset(db_conn_id="tenant_{{
params.customer }}"))
+ op = AgentOperator(
+ task_id="t", prompt="p", llm_conn_id="llm",
agent_params={"capabilities": [capability]}
+ )
+
+ op.render_template_fields(self.CONTEXT)
+
+ assert op.agent_params["capabilities"][0] is capability
+
+ def test_only_connection_ids_are_templated(self):
+ """allowed_tables is validated and canonicalised in __init__, so
rendering it later
+ would bypass the fail-closed empty-list check."""
+ assert SQLToolset.agent_template_fields == ("_db_conn_id",)
+ assert MCPToolset.agent_template_fields == ("_mcp_conn_id",)
+ assert HookToolset.agent_template_fields == ("conn_id",)
+
+ @pytest.mark.parametrize("toolset_cls", [SQLToolset, MCPToolset,
HookToolset])
+ def test_toolsets_do_not_opt_in_through_template_fields(self, toolset_cls):
+ """Airflow's templater renders any object with ``template_fields`` in
place wherever it is
+ nested in a template field, which would leak one task instance's
connection to the next."""
+ assert not hasattr(toolset_cls, "template_fields")
+
+ def test_toolsets_in_agent_params_render_on_a_copy_per_task(self):
+ shared = SQLToolset(db_conn_id="tenant_{{ params.customer }}")
+ first = AgentOperator(task_id="a", prompt="p", llm_conn_id="llm",
agent_params={"toolsets": [shared]})
+ second = AgentOperator(
+ task_id="b", prompt="p", llm_conn_id="llm",
agent_params={"toolsets": [shared]}
+ )
+
+ first.render_template_fields({"params": {"customer": "acme"}})
+ second.render_template_fields({"params": {"customer": "globex"}})
+
+ assert first.agent_params["toolsets"][0].id == "sql-tenant_acme"
+ assert second.agent_params["toolsets"][0].id == "sql-tenant_globex"
+ assert shared._db_conn_id == "tenant_{{ params.customer }}"
+
+ def test_wrapped_toolset_inside_a_capability_is_rendered(self):
+ capability = Toolset(SQLToolset(db_conn_id="tenant_{{ params.customer
}}").prefixed("crm"))
+ op = AgentOperator(
+ task_id="t", prompt="p", llm_conn_id="llm",
agent_params={"capabilities": [capability]}
+ )
+
+ op.render_template_fields(self.CONTEXT)
+
+ found = find_toolset([op.agent_params["capabilities"][0].toolset],
SQLToolset)
+ assert found is not None
+ assert found.id == "sql-tenant_acme"
+
+ def test_rendered_toolset_id_is_logged(self, caplog):
+ op = AgentOperator(
+ task_id="t",
+ prompt="p",
+ llm_conn_id="llm",
+ toolsets=[SQLToolset(db_conn_id="tenant_{{ params.customer }}")],
+ )
+
+ with caplog.at_level("INFO"):
+ op.render_template_fields(self.CONTEXT)
+
+ assert "Rendered toolset sql-tenant_acme" in caplog.text
+
+ def test_rendering_twice_logs_once(self, caplog):
+ """@task.agent renders a second time; by then the id no longer
changes."""
+ op = AgentOperator(
+ task_id="t",
+ prompt="p",
+ llm_conn_id="llm",
+ toolsets=[SQLToolset(db_conn_id="tenant_{{ params.customer }}")],
+ )
+
+ with caplog.at_level("INFO"):
+ op.render_template_fields(self.CONTEXT)
+ op.render_template_fields(self.CONTEXT)
+
+ assert caplog.text.count("Rendered toolset sql-tenant_acme") == 1
+
+ def test_toolset_without_template_fields_is_left_as_is(self):
+ toolset = FunctionToolset()
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
toolsets=[toolset])
+
+ op.render_template_fields(self.CONTEXT)
+
+ assert op.toolsets[0] is toolset
+
+ def test_untemplated_wrapper_subclass_is_not_rebuilt(self):
+ """visit_and_replace rebuilds wrappers with dataclasses.replace, which
a subclass with its
+ own __init__ does not survive; nothing to render means nothing to
rebuild."""
+
+ class Audited(WrapperToolset):
+ def __init__(self, wrapped, *, audit_name):
+ super().__init__(wrapped=wrapped)
+ self.audit_name = audit_name
+
+ toolset = Audited(FunctionToolset(), audit_name="x")
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
toolsets=[toolset])
+
+ op.render_template_fields(self.CONTEXT)
+
+ assert op.toolsets[0] is toolset
+
+ @pytest.mark.parametrize(
+ ("form", "template"),
+ [
+ pytest.param("operator", "tenant_{{ task.prompt }}",
id="operator"),
+ pytest.param("decorator", "tenant_{{ task.op_kwargs.customer }}",
id="decorator"),
+ ],
+ )
+ def test_each_map_index_gets_its_own_connection(self, form, template):
+ """Through the real MappedOperator render path, for both authoring
forms."""
+ shared = SQLToolset(db_conn_id=template)
+ with DAG("d", schedule=None) as dag:
+ if form == "operator":
+ mapped = AgentOperator.partial(task_id="m", llm_conn_id="llm",
toolsets=[shared]).expand(
+ prompt=["acme", "globex"]
+ )
+ else:
+
+ @task.agent(llm_conn_id="llm", toolsets=[shared])
+ def report(customer: str) -> str:
+ return customer
+
+ mapped = report.expand(customer=["acme", "globex"]).operator
+
+ ids = []
+ for map_index in (0, 1):
+ context: dict = {
+ "ti": SimpleNamespace(map_index=map_index),
+ "params": {},
+ "dag": dag,
+ "dag_run": SimpleNamespace(conf={}),
+ }
+ mapped.render_template_fields(context, dag.get_template_env())
+ ids.append(context["task"].toolsets[0].id)
+
+ assert ids == ["sql-tenant_acme", "sql-tenant_globex"]
+ assert shared._db_conn_id == template
+
+
class TestAgentOperatorExecute:
@pytest.mark.parametrize(
"bad",
diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py
b/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py
index b8ce959db02..8ddc911b144 100644
--- a/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py
+++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py
@@ -85,6 +85,61 @@ class TestHookToolsetInit:
assert "FakeHook" in ts.id
+class _FakeConnHook(_FakeHook):
+ """A hook that names its connection attribute, like every provider hook
does."""
+
+ conn_name_attr = "fake_conn_id"
+
+ def __init__(self, fake_conn_id: str = "fake_default"):
+ self.fake_conn_id = fake_conn_id
+
+
+class TestHookToolsetConnId:
+ def test_conn_id_is_read_from_the_hooks_conn_name_attr(self):
+ ts = HookToolset(_FakeConnHook("warehouse"),
allowed_methods=["list_keys"])
+
+ assert ts.conn_id == "warehouse"
+ assert ts.id == "hook-_FakeConnHook-warehouse"
+
+ def test_a_hook_without_conn_name_attr_has_no_conn_id(self):
+ ts = HookToolset(_FakeHook(), allowed_methods=["list_keys"])
+
+ assert ts.conn_id is None
+ assert ts.id == "hook-_FakeHook"
+
+ def test_setting_conn_id_copies_the_hook(self):
+ """The hook in the Dag file is shared by every task instance that uses
the toolset."""
+ hook = _FakeConnHook("tenant_{{ customer }}")
+ ts = HookToolset(hook, allowed_methods=["list_keys"])
+
+ ts.conn_id = "tenant_acme"
+
+ assert ts.conn_id == "tenant_acme"
+ assert ts._hook is not hook
+ assert hook.fake_conn_id == "tenant_{{ customer }}"
+
+ def test_setting_conn_id_on_a_hook_without_one_raises(self):
+ ts = HookToolset(_FakeHook(), allowed_methods=["list_keys"])
+
+ with pytest.raises(AttributeError, match="keeps no connection ID"):
+ ts.conn_id = "x"
+
+ def test_falls_back_to_conn_id_when_conn_name_attr_is_not_set(self):
+ """WasbHook and KubernetesHook declare one attribute and keep the ID
in ``conn_id``."""
+
+ class _WasbShapedHook(_FakeHook):
+ conn_name_attr = "wasb_conn_id"
+
+ def __init__(self, wasb_conn_id: str):
+ self.conn_id = wasb_conn_id
+
+ ts = HookToolset(_WasbShapedHook("blob_{{ customer }}"),
allowed_methods=["list_keys"])
+ ts.conn_id = "blob_acme"
+
+ assert ts.conn_id == "blob_acme"
+ assert ts.id == "hook-_WasbShapedHook-blob_acme"
+
+
class TestHookToolsetGetTools:
def test_returns_tools_for_allowed_methods(self):
hook = _FakeHook()