Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions providers/common/ai/docs/toolsets/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -239,6 +239,8 @@ the call. What counts differs by toolset:
- ``ObjectStorageToolset`` counts invalid arguments only. A path that does not exist or
cannot be read goes back to the model as a failed result without using the budget;
bound repeated failed reads with ``usage_limits``.
- ``AgentSkillsToolset`` counts a failed call to a skills tool, such as a resource name
that does not exist.

These toolsets allow as many corrections as the agent's tool retry budget, pydantic-ai's
``retries`` (one by default), the same way pydantic-ai's own toolsets do. Pass
Expand Down
13 changes: 10 additions & 3 deletions providers/common/ai/docs/toolsets/skills.rst
Original file line number Diff line number Diff line change
Expand Up @@ -109,14 +109,18 @@ reading an excluded file got:

Resource 'warehouse.env' not found in skill 'sql-reporting'. Available resources: ['reference.md']. Use the exact name from load_skill output.

The skills tools allow the model one correction each, and the agent's ``retries``
does not change that. A second refused read in a row failed the run with
``UnexpectedModelBehavior``, even with ``retries`` set to ``{"tools": 3}``:
Without ``max_retries``, each skills tool allows as many corrections as the agent's
``retries``, one by default, and a successful call to that tool resets the count. With
the default, a second refused read in a row failed the run with
``UnexpectedModelBehavior``:

.. code-block:: text

Tool 'read_skill_resource' exceeded max retries count of 1. Consider raising the retry limit, or see the docs on tool retries: https://pydantic.dev/docs/ai/tools-toolsets/tools-advanced/#tool-retries

With the example's ``max_retries=3``, the same two refused reads went back to the model,
which then read ``reference.md`` and finished the run.

``exclude_resources`` hides files from the resource tools only. A skill's scripts can
still read them, which is why the example excludes ``run_skill_script`` as well.

Expand All @@ -137,6 +141,9 @@ Parameters
it does not stop a skill's ``run_skill_script`` from reading them off disk, so
pair it with ``exclude_tools={"run_skill_script"}`` when the files are
genuinely sensitive.
- ``max_retries``: How many times the model may correct failed calls to one skills tool
before the run fails; a successful call to that tool resets the count. Default
``None``, the agent's ``retries``. See :ref:`toolset-retry-budget`.

Using Agent Skills with other frameworks
----------------------------------------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,8 @@ def example_agent_skills_restricted():
exclude_tools={"run_skill_script"},
# Keep matching files out of the resources the model can list and read.
exclude_resources=["*.env", "secrets/*"],
# Allow up to 3 corrections per skills tool; a successful call resets the count.
max_retries=3,
)
],
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,9 +30,11 @@

from __future__ import annotations

import dataclasses
from typing import TYPE_CHECKING, Any

from airflow.providers.common.ai.skills import SkillSource, _materialize_skills
from airflow.providers.common.ai.utils.toolset_base import validate_max_retries

try:
from pydantic_ai.toolsets.abstract import AbstractToolset
Expand Down Expand Up @@ -72,6 +74,10 @@ class AgentSkillsToolset(AbstractToolset):
discovery only -- it does not stop a skill's ``run_skill_script`` from
reading them off disk, so pair it with ``exclude_tools={"run_skill_script"}``
when the files are genuinely sensitive. Requires ``pydantic-ai-skills>=1.2.0``.
:param max_retries: How many times the model may correct failed calls to one skills tool,
such as a resource name that does not exist, before the run fails; a successful call
to that tool resets the count. ``None`` (the default) uses the agent's tool retry
budget, its ``retries``, as the provider's other toolsets do.

Requires the ``skills`` extra: ``pip install "apache-airflow-providers-common-ai[skills]"``.
"""
Expand All @@ -82,10 +88,12 @@ def __init__(
*,
exclude_tools: set[str] | None = None,
exclude_resources: list[str] | None = None,
max_retries: int | None = None,
) -> None:
self._sources = list(sources)
self._exclude_tools = exclude_tools
self._exclude_resources = exclude_resources
self._max_retries = validate_max_retries(max_retries)
self._inner: Any = None
self._cleanup: Callable[[], None] | None = None

Expand All @@ -101,6 +109,7 @@ async def for_run(self, ctx: RunContext) -> AbstractToolset:
self._sources,
exclude_tools=self._exclude_tools,
exclude_resources=self._exclude_resources,
max_retries=self._max_retries,
)

async def __aenter__(self) -> AgentSkillsToolset:
Expand Down Expand Up @@ -150,7 +159,11 @@ def _require_inner(self) -> Any:
return self._inner

async def get_tools(self, ctx: RunContext) -> dict[str, ToolsetTool]:
return await self._require_inner().get_tools(ctx)
# pydantic-ai-skills fixes its tools' budget at one correction, whatever the agent's
# retries say; resolve it the way the provider's other toolsets do instead.
max_retries = ctx.max_retries if self._max_retries is None else self._max_retries
tools = await self._require_inner().get_tools(ctx)
return {name: dataclasses.replace(tool, max_retries=max_retries) for name, tool in tools.items()}

async def call_tool(
self, name: str, tool_args: dict[str, Any], ctx: RunContext, tool: ToolsetTool
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,10 @@
from unittest.mock import MagicMock, patch

import pytest
from pydantic_ai import Agent
from pydantic_ai.exceptions import UnexpectedModelBehavior
from pydantic_ai.messages import ModelResponse, TextPart, ToolCallPart
from pydantic_ai.models.function import FunctionModel

from airflow.providers.common.ai.skills import GitSkills
from airflow.providers.common.ai.toolsets.skills import AgentSkillsToolset
Expand Down Expand Up @@ -145,17 +149,58 @@ def fake_skillstoolset(**kwargs):

assert "exclude_tools" not in captured
assert "exclude_resources" not in captured
assert "max_retries" not in captured

def test_negative_max_retries_is_rejected(self):
with pytest.raises(ValueError, match="max_retries must not be negative"):
AgentSkillsToolset(sources=["/x"], max_retries=-1)

def test_for_run_propagates_optional_kwargs(self):
# for_run hands each run its own instance; dropping a kwarg here would
# silently expose excluded files in concurrent runs.
toolset = AgentSkillsToolset(
sources=["/x"], exclude_tools={"run_skill_script"}, exclude_resources=["*.env"]
sources=["/x"], exclude_tools={"run_skill_script"}, exclude_resources=["*.env"], max_retries=3
)
per_run = asyncio.run(toolset.for_run(MagicMock())) # noqa: spec (for_run ignores ctx)
assert per_run is not toolset
assert per_run._exclude_tools == {"run_skill_script"}
assert per_run._exclude_resources == ["*.env"]
assert per_run._max_retries == 3

@pytest.mark.parametrize(
("max_retries", "agent_retries", "fails"),
[
pytest.param(None, 1, True, id="agent_default_allows_one_correction"),
pytest.param(None, 3, False, id="follows_agent_retries"),
pytest.param(2, 1, False, id="own_budget_wins"),
],
)
def test_max_retries_bounds_corrections_in_a_real_run(self, tmp_path, max_retries, agent_retries, fails):
"""Two unknown resource names in a row need a budget of at least two corrections."""
_write_skill(tmp_path)
calls = iter(
[
{"skill_name": "demo-skill", "resource_name": "missing-1.md"},
{"skill_name": "demo-skill", "resource_name": "missing-2.md"},
]
)

def model(messages, info):
if (args := next(calls, None)) is None:
return ModelResponse(parts=[TextPart("done")])
return ModelResponse(parts=[ToolCallPart("read_skill_resource", args)])

agent = Agent(
FunctionModel(model),
toolsets=[AgentSkillsToolset(sources=[str(tmp_path)], max_retries=max_retries)],
retries={"tools": agent_retries},
)

if fails:
with pytest.raises(UnexpectedModelBehavior, match="exceeded max retries count of 1"):
agent.run_sync("go")
else:
assert agent.run_sync("go").output == "done"


class TestCleanup:
Expand Down