diff --git a/providers/common/ai/docs/toolsets/index.rst b/providers/common/ai/docs/toolsets/index.rst index 3d61337e8a37b..eb9c052544850 100644 --- a/providers/common/ai/docs/toolsets/index.rst +++ b/providers/common/ai/docs/toolsets/index.rst @@ -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 diff --git a/providers/common/ai/docs/toolsets/skills.rst b/providers/common/ai/docs/toolsets/skills.rst index bb1cdcc6bf456..1e0b1ed6dd7b7 100644 --- a/providers/common/ai/docs/toolsets/skills.rst +++ b/providers/common/ai/docs/toolsets/skills.rst @@ -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. @@ -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 ---------------------------------------- diff --git a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_skills.py b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_skills.py index f922b40993d30..5618c758e7ad0 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_skills.py +++ b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_skills.py @@ -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, ) ], ) diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/skills.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/skills.py index 8dc1fb68ba955..65d508098e33b 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/skills.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/skills.py @@ -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 @@ -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]"``. """ @@ -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 @@ -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: @@ -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 diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_skills.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_skills.py index a0e768da563e3..7801697948d4d 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_skills.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_skills.py @@ -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 @@ -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: