vatsrahul1001 commented on code in PR #74385: URL: https://github.com/apache/airflow/pull/74385#discussion_r4203668988
########## providers/common/ai/tests/unit/common/ai/durable/test_capability.py: ########## @@ -0,0 +1,699 @@ +# 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. +""" +AirflowDurability through a real pydantic-ai agent loop. + +Each test runs the scenario durable execution exists for: an attempt fails partway, and +the retry runs the agent again from the top with a fresh journal over the same storage. +""" + +from __future__ import annotations + +import asyncio +import dataclasses +from typing import TYPE_CHECKING, Any + +import pytest +from pydantic_ai import Agent, CancellationToken, RunContext +from pydantic_ai.capabilities import AbstractCapability, durable_operation +from pydantic_ai.exceptions import ModelRetry +from pydantic_ai.messages import ModelMessage, ModelResponse, TextPart, ToolCallPart, ToolReturnPart +from pydantic_ai.models.function import AgentInfo, FunctionModel +from pydantic_ai.toolsets import FunctionToolset +from pydantic_ai.usage import RequestUsage, RunUsage, UsageLimits + +from airflow.providers.common.ai.durable import AirflowDurability +from airflow.providers.common.ai.durable.journal import DurableJournal, journal_scope +from airflow.providers.common.ai.exceptions import DurableJournalError +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, ensure_masked + +if TYPE_CHECKING: + from pydantic_ai.toolsets.abstract import ToolsetTool + + +class Calls: + """Counts the live calls one test's model and tools make, across attempts.""" + + def __init__(self) -> None: + self.counts: dict[str, int] = {} + + def bump(self, name: str) -> None: + self.counts[name] = self.counts.get(name, 0) + 1 + + def __getitem__(self, name: str) -> int: + return self.counts.get(name, 0) + + +def responses_so_far(messages: list[ModelMessage]) -> int: + return sum(isinstance(message, ModelResponse) for message in messages) + + +def tool_returns(messages: list[ModelMessage]) -> list[ToolReturnPart]: + return [part for message in messages for part in message.parts if isinstance(part, ToolReturnPart)] + + +async def attempt(storage, agent: Agent[Any, Any], prompt: str = "go", **run_kwargs: Any) -> Any: + """Run one task attempt: a fresh journal over the storage the attempts share.""" + with journal_scope(DurableJournal(storage)): + return await agent.run(prompt, **run_kwargs) + + +class _WarehouseToolset(AirflowToolset): + """An Airflow toolset with one ``query`` tool, standing in for SQLToolset and friends.""" + + def __init__(self, calls: Calls, result: Any = "3 rows", *, replayable: bool = True) -> None: + self._calls = calls + self._result = result + self.replayable = replayable + + def query(sql: str) -> str: + """Run a query.""" + raise AssertionError("served by execute_tool") + + self._inner = FunctionToolset(tools=[query]) + + @property + def id(self) -> str: + return "warehouse" + + async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[Any]]: + tools = await self._inner.get_tools(ctx) + return {name: dataclasses.replace(tool, toolset=self) for name, tool in tools.items()} + + async def execute_tool(self, name, tool_args, *, ctx, tool) -> Any: + self._calls.bump("query") + if isinstance(self._result, BaseException): + raise self._result + return self._result + + +def tool_then_answer(calls: Calls, tool_name: str, *, fail_final: list[bool]): + """A model that calls ``tool_name`` once, then answers with what the tool returned.""" + + def model_fn(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + calls.bump("model") + if responses_so_far(messages) == 0: + return ModelResponse(parts=[ToolCallPart(tool_name, {"sql": "select 1"})]) + if fail_final[0]: + raise RuntimeError("worker died") + return ModelResponse(parts=[TextPart(f"answer: {tool_returns(messages)[-1].content}")]) + + return model_fn + + +class TestReplay: + @pytest.mark.asyncio + async def test_retry_replays_completed_steps_and_runs_only_the_failed_one(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build() -> Agent[None, str]: + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + calls.bump("query") + return "3 rows" + + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError, match="worker died"): + await attempt(memory_storage, build()) + fail_final[0] = False + result = await attempt(memory_storage, build()) + + assert result.output == "answer: 3 rows" + assert calls["query"] == 1 + # The first model step replayed; only the one that failed runs again. + assert calls["model"] == 3 + + @pytest.mark.asyncio + async def test_changed_prompt_runs_everything_again(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build(instructions: str) -> Agent[None, str]: + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + calls.bump("query") + return "3 rows" + + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + instructions=instructions, + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError): + await attempt(memory_storage, build("be terse")) + fail_final[0] = False + await attempt(memory_storage, build("be thorough")) + + assert calls["query"] == 2 + assert calls["model"] == 4 + + @pytest.mark.asyncio + async def test_parallel_tool_calls_replay_by_the_order_they_started(self, memory_storage): + calls = Calls() + delays: dict[int, float] = {} + fail_final = [True] + + def model_fn(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + if responses_so_far(messages) == 0: + return ModelResponse(parts=[ToolCallPart("charge", {"amount": n}) for n in (1, 2, 3)]) + if fail_final[0]: + raise RuntimeError("worker died") + return ModelResponse(parts=[TextPart(",".join(str(p.content) for p in tool_returns(messages)))]) + + def build() -> Agent[None, str]: + toolset = FunctionToolset(id="billing") + + @toolset.tool_plain + async def charge(amount: int) -> str: + await asyncio.sleep(delays.get(amount, 0)) + calls.bump(f"charge_{amount}") + return f"charged {amount}" + + return Agent( + FunctionModel(model_fn), name="biller", toolsets=[toolset], capabilities=[AirflowDurability()] + ) + + # The first attempt finishes the calls in the reverse of the order it started them. + delays.update({1: 0.03, 2: 0.02, 3: 0.0}) + with pytest.raises(RuntimeError): + await attempt(memory_storage, build()) + fail_final[0] = False + delays.clear() + result = await attempt(memory_storage, build()) + + assert result.output == "charged 1,charged 2,charged 3" + assert calls.counts == {"charge_1": 1, "charge_2": 1, "charge_3": 1} + + @pytest.mark.asyncio + async def test_outside_a_journal_the_agent_runs_normally(self, memory_storage): + calls = Calls() + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + calls.bump("query") + return "3 rows" + + agent = Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=[False])), + name="analyst", + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + result = await agent.run("go") + + assert result.output == "answer: 3 rows" + assert memory_storage.entries == {} + + +class TestResultsThatCannotBeRecorded: + @pytest.mark.asyncio + async def test_a_function_tool_result_that_is_not_json_fails_without_retrying(self, memory_storage): + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> object: + return object() + + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(DurableJournalError, match="not JSON-serializable"): + await attempt(memory_storage, agent) + + @pytest.mark.asyncio + async def test_an_airflow_toolset_result_that_is_not_json_fails_without_retrying(self, memory_storage): + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[_WarehouseToolset(Calls(), object())], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(DurableJournalError, match="not JSON-serializable"): + await attempt(memory_storage, agent) + + +class TestStreaming: + @pytest.mark.asyncio + async def test_a_streamed_model_request_replays(self, memory_storage): + calls = Calls() + + async def stream_fn(messages: list[ModelMessage], info: AgentInfo): + calls.bump("model") + yield "streamed answer" + + def build() -> Agent[None, str]: + return Agent( + FunctionModel(stream_function=stream_fn), name="streamer", capabilities=[AirflowDurability()] + ) + + outputs = [] + for _ in range(2): + with journal_scope(DurableJournal(memory_storage)): + async with build().run_stream("go") as result: + outputs.append(await result.get_output()) + + assert outputs == ["streamed answer", "streamed answer"] + assert calls["model"] == 1 + + +class TestCancellation: + @pytest.mark.asyncio + async def test_a_run_inside_the_task_accepts_a_cancellation_token(self, memory_storage): + """AgentOperator.on_kill cancels the run in the task process, which is the durable container.""" + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[_WarehouseToolset(Calls())], + capabilities=[AirflowDurability()], + ) + + result = await attempt(memory_storage, agent, cancellation_token=CancellationToken()) + + assert result.output == "answer: 3 rows" + + +class TestAirflowToolsets: + """Airflow's own toolsets are not durable units of pydantic-ai's backend; the capability journals them.""" + + @pytest.mark.asyncio + async def test_airflow_toolset_calls_replay(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build() -> Agent[None, str]: + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + toolsets=[_WarehouseToolset(calls)], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError): + await attempt(memory_storage, build()) + fail_final[0] = False + result = await attempt(memory_storage, build()) + + assert result.output == "answer: 3 rows" + assert calls["query"] == 1 + + @pytest.mark.asyncio + async def test_a_model_retry_replays_without_running_the_tool(self, memory_storage): + calls = Calls() + fail_final = [True] + + def model_fn(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + if responses_so_far(messages) == 0: + return ModelResponse(parts=[ToolCallPart("query", {"sql": "select nope"})]) + if fail_final[0]: + raise RuntimeError("worker died") + return ModelResponse(parts=[TextPart(f"retry said: {messages[-1].parts[0].content}")]) Review Comment: mypy's red here — this resolves to a union of part types that don't all have the attribute, so it needs an isinstance narrow (or a cast) before the access. ########## providers/common/ai/tests/unit/common/ai/durable/test_capability.py: ########## @@ -0,0 +1,699 @@ +# 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. +""" +AirflowDurability through a real pydantic-ai agent loop. + +Each test runs the scenario durable execution exists for: an attempt fails partway, and +the retry runs the agent again from the top with a fresh journal over the same storage. +""" + +from __future__ import annotations + +import asyncio +import dataclasses +from typing import TYPE_CHECKING, Any + +import pytest +from pydantic_ai import Agent, CancellationToken, RunContext +from pydantic_ai.capabilities import AbstractCapability, durable_operation +from pydantic_ai.exceptions import ModelRetry +from pydantic_ai.messages import ModelMessage, ModelResponse, TextPart, ToolCallPart, ToolReturnPart +from pydantic_ai.models.function import AgentInfo, FunctionModel +from pydantic_ai.toolsets import FunctionToolset +from pydantic_ai.usage import RequestUsage, RunUsage, UsageLimits + +from airflow.providers.common.ai.durable import AirflowDurability +from airflow.providers.common.ai.durable.journal import DurableJournal, journal_scope +from airflow.providers.common.ai.exceptions import DurableJournalError +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, ensure_masked + +if TYPE_CHECKING: + from pydantic_ai.toolsets.abstract import ToolsetTool + + +class Calls: + """Counts the live calls one test's model and tools make, across attempts.""" + + def __init__(self) -> None: + self.counts: dict[str, int] = {} + + def bump(self, name: str) -> None: + self.counts[name] = self.counts.get(name, 0) + 1 + + def __getitem__(self, name: str) -> int: + return self.counts.get(name, 0) + + +def responses_so_far(messages: list[ModelMessage]) -> int: + return sum(isinstance(message, ModelResponse) for message in messages) + + +def tool_returns(messages: list[ModelMessage]) -> list[ToolReturnPart]: + return [part for message in messages for part in message.parts if isinstance(part, ToolReturnPart)] + + +async def attempt(storage, agent: Agent[Any, Any], prompt: str = "go", **run_kwargs: Any) -> Any: + """Run one task attempt: a fresh journal over the storage the attempts share.""" + with journal_scope(DurableJournal(storage)): + return await agent.run(prompt, **run_kwargs) + + +class _WarehouseToolset(AirflowToolset): + """An Airflow toolset with one ``query`` tool, standing in for SQLToolset and friends.""" + + def __init__(self, calls: Calls, result: Any = "3 rows", *, replayable: bool = True) -> None: + self._calls = calls + self._result = result + self.replayable = replayable + + def query(sql: str) -> str: + """Run a query.""" + raise AssertionError("served by execute_tool") + + self._inner = FunctionToolset(tools=[query]) + + @property + def id(self) -> str: + return "warehouse" + + async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[Any]]: + tools = await self._inner.get_tools(ctx) + return {name: dataclasses.replace(tool, toolset=self) for name, tool in tools.items()} + + async def execute_tool(self, name, tool_args, *, ctx, tool) -> Any: + self._calls.bump("query") + if isinstance(self._result, BaseException): + raise self._result + return self._result + + +def tool_then_answer(calls: Calls, tool_name: str, *, fail_final: list[bool]): + """A model that calls ``tool_name`` once, then answers with what the tool returned.""" + + def model_fn(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + calls.bump("model") + if responses_so_far(messages) == 0: + return ModelResponse(parts=[ToolCallPart(tool_name, {"sql": "select 1"})]) + if fail_final[0]: + raise RuntimeError("worker died") + return ModelResponse(parts=[TextPart(f"answer: {tool_returns(messages)[-1].content}")]) + + return model_fn + + +class TestReplay: + @pytest.mark.asyncio + async def test_retry_replays_completed_steps_and_runs_only_the_failed_one(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build() -> Agent[None, str]: + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + calls.bump("query") + return "3 rows" + + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError, match="worker died"): + await attempt(memory_storage, build()) + fail_final[0] = False + result = await attempt(memory_storage, build()) + + assert result.output == "answer: 3 rows" + assert calls["query"] == 1 + # The first model step replayed; only the one that failed runs again. + assert calls["model"] == 3 + + @pytest.mark.asyncio + async def test_changed_prompt_runs_everything_again(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build(instructions: str) -> Agent[None, str]: + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + calls.bump("query") + return "3 rows" + + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + instructions=instructions, + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError): + await attempt(memory_storage, build("be terse")) + fail_final[0] = False + await attempt(memory_storage, build("be thorough")) + + assert calls["query"] == 2 + assert calls["model"] == 4 + + @pytest.mark.asyncio + async def test_parallel_tool_calls_replay_by_the_order_they_started(self, memory_storage): + calls = Calls() + delays: dict[int, float] = {} + fail_final = [True] + + def model_fn(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + if responses_so_far(messages) == 0: + return ModelResponse(parts=[ToolCallPart("charge", {"amount": n}) for n in (1, 2, 3)]) + if fail_final[0]: + raise RuntimeError("worker died") + return ModelResponse(parts=[TextPart(",".join(str(p.content) for p in tool_returns(messages)))]) + + def build() -> Agent[None, str]: + toolset = FunctionToolset(id="billing") + + @toolset.tool_plain + async def charge(amount: int) -> str: + await asyncio.sleep(delays.get(amount, 0)) + calls.bump(f"charge_{amount}") + return f"charged {amount}" + + return Agent( + FunctionModel(model_fn), name="biller", toolsets=[toolset], capabilities=[AirflowDurability()] + ) + + # The first attempt finishes the calls in the reverse of the order it started them. + delays.update({1: 0.03, 2: 0.02, 3: 0.0}) + with pytest.raises(RuntimeError): + await attempt(memory_storage, build()) + fail_final[0] = False + delays.clear() + result = await attempt(memory_storage, build()) + + assert result.output == "charged 1,charged 2,charged 3" + assert calls.counts == {"charge_1": 1, "charge_2": 1, "charge_3": 1} + + @pytest.mark.asyncio + async def test_outside_a_journal_the_agent_runs_normally(self, memory_storage): + calls = Calls() + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + calls.bump("query") + return "3 rows" + + agent = Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=[False])), + name="analyst", + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + result = await agent.run("go") + + assert result.output == "answer: 3 rows" + assert memory_storage.entries == {} + + +class TestResultsThatCannotBeRecorded: + @pytest.mark.asyncio + async def test_a_function_tool_result_that_is_not_json_fails_without_retrying(self, memory_storage): + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> object: + return object() + + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[toolset], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(DurableJournalError, match="not JSON-serializable"): + await attempt(memory_storage, agent) + + @pytest.mark.asyncio + async def test_an_airflow_toolset_result_that_is_not_json_fails_without_retrying(self, memory_storage): + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[_WarehouseToolset(Calls(), object())], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(DurableJournalError, match="not JSON-serializable"): + await attempt(memory_storage, agent) + + +class TestStreaming: + @pytest.mark.asyncio + async def test_a_streamed_model_request_replays(self, memory_storage): + calls = Calls() + + async def stream_fn(messages: list[ModelMessage], info: AgentInfo): + calls.bump("model") + yield "streamed answer" + + def build() -> Agent[None, str]: + return Agent( + FunctionModel(stream_function=stream_fn), name="streamer", capabilities=[AirflowDurability()] + ) + + outputs = [] + for _ in range(2): + with journal_scope(DurableJournal(memory_storage)): + async with build().run_stream("go") as result: + outputs.append(await result.get_output()) + + assert outputs == ["streamed answer", "streamed answer"] + assert calls["model"] == 1 + + +class TestCancellation: + @pytest.mark.asyncio + async def test_a_run_inside_the_task_accepts_a_cancellation_token(self, memory_storage): + """AgentOperator.on_kill cancels the run in the task process, which is the durable container.""" + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[_WarehouseToolset(Calls())], + capabilities=[AirflowDurability()], + ) + + result = await attempt(memory_storage, agent, cancellation_token=CancellationToken()) + + assert result.output == "answer: 3 rows" + + +class TestAirflowToolsets: + """Airflow's own toolsets are not durable units of pydantic-ai's backend; the capability journals them.""" + + @pytest.mark.asyncio + async def test_airflow_toolset_calls_replay(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build() -> Agent[None, str]: + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + toolsets=[_WarehouseToolset(calls)], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError): + await attempt(memory_storage, build()) + fail_final[0] = False + result = await attempt(memory_storage, build()) + + assert result.output == "answer: 3 rows" + assert calls["query"] == 1 + + @pytest.mark.asyncio + async def test_a_model_retry_replays_without_running_the_tool(self, memory_storage): + calls = Calls() + fail_final = [True] + + def model_fn(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + if responses_so_far(messages) == 0: + return ModelResponse(parts=[ToolCallPart("query", {"sql": "select nope"})]) + if fail_final[0]: + raise RuntimeError("worker died") + return ModelResponse(parts=[TextPart(f"retry said: {messages[-1].parts[0].content}")]) + + def build() -> Agent[None, str]: + return Agent( + FunctionModel(model_fn), + name="analyst", + toolsets=[_WarehouseToolset(calls, ModelRetry("no column nope"))], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError): + await attempt(memory_storage, build()) + fail_final[0] = False + result = await attempt(memory_storage, build()) + + assert result.output == "retry said: no column nope" + assert calls["query"] == 1 + + @pytest.mark.asyncio + async def test_a_toolset_that_is_not_replayable_runs_again(self, memory_storage): + calls = Calls() + fail_final = [True] + + def build() -> Agent[None, str]: + return Agent( + FunctionModel(tool_then_answer(calls, "query", fail_final=fail_final)), + name="analyst", + toolsets=[_WarehouseToolset(calls, replayable=False)], + capabilities=[AirflowDurability()], + ) + + with pytest.raises(RuntimeError): + await attempt(memory_storage, build()) + fail_final[0] = False + await attempt(memory_storage, build()) + + assert calls["query"] == 2 + # The model step before it still replayed. + assert calls["model"] == 3 + + [email protected]_redact +class TestMasking: + @pytest.mark.asyncio + async def test_journal_holds_function_tool_results_masked(self, memory_storage, registered_secret): + toolset = FunctionToolset(id="db") + + @toolset.tool_plain + def query(sql: str) -> str: + return f"password={registered_secret}" + + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + # AgentOperator puts the masking wrapper outside the durable unit. + toolsets=[ensure_masked(toolset)], + capabilities=[AirflowDurability()], + ) + + result = await attempt(memory_storage, agent) + + assert result.output == "answer: password=***" + assert registered_secret not in str(memory_storage.entries) + + @pytest.mark.asyncio + async def test_journal_holds_airflow_toolset_results_masked(self, memory_storage, registered_secret): + agent = Agent( + FunctionModel(tool_then_answer(Calls(), "query", fail_final=[False])), + name="analyst", + toolsets=[_WarehouseToolset(Calls(), f"password={registered_secret}")], + capabilities=[AirflowDurability()], + ) + + await attempt(memory_storage, agent) + + assert registered_secret not in str(memory_storage.entries) + + +class _Ledger(AbstractCapability[Any]): + """Stands in for pydantic-ai-harness SpendLimits: it accrues through a durable operation.""" + + def __init__(self, calls: Calls, id: str | None = "ledger") -> None: + self._calls = calls + self._id = id + + @property + def id(self) -> str | None: Review Comment: mypy: this overrides a writeable attribute with a read-only property ([override]). Make it a settable property or just a plain attribute. -- 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]
