kaxil commented on code in PR #74385: URL: https://github.com/apache/airflow/pull/74385#discussion_r4206743991
########## 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: Thanks, added an `isinstance(..., RetryPromptPart)` narrow before reading `content` (8f5152314f3). ########## 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: Made `id` a plain attribute in `__init__` (8f5152314f3). I'd only run mypy on the source locally; ran it on the tests too this time and it's clean. ########## providers/common/ai/src/airflow/providers/common/ai/durable/journal.py: ########## @@ -0,0 +1,510 @@ +# 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. +""" +The durable journal: replay an agent's completed steps when Airflow retries its task. + +An Airflow task retry is a fresh process that runs the agent again from the top. The +journal records each step an agent run completes (a model response, a tool result, a +capability operation) under the step's position in that run, and on the retry hands +back the recorded result instead of running the step again. Steps the previous +attempt never reached run live and are recorded in turn. + +A step is replayed only when the step at the same position in the previous attempt +had the same name and the same fingerprint. The name says what kind of step it was +(``agent__model.request``, ``agent__function_toolset__x.call_tool:query``); the +fingerprint, computed by the framework adapter, says what it was asked (the message +history, the tool arguments). The first mismatch means the run took a different path +from the previous attempt, so that step and every step after it run live. So does +everything the previous attempt did after a step that raised. + +A :class:`DurableJournal` belongs to one task attempt. Each agent run in it is a +:class:`DurableRun`, numbered in the order the runs start, with positions of its own, +so a task that runs several agents replays each of them independently. + +Nothing here depends on an agent framework. Payloads are JSON-compatible values and +fingerprints are opaque strings. +:class:`~airflow.providers.common.ai.durable.capability.AirflowDurability` drives the +journal for pydantic-ai. Another framework's adapter starts a run with +:meth:`DurableJournal.start_run` and calls one of two shapes of API on it: + +* :meth:`DurableRun.run` wraps a step: + ``await run.run(name, kind=..., fingerprint=..., body=...)``. +* :meth:`DurableRun.claim` and :meth:`JournalStep.record` split it in two, for a + framework whose hooks see a step's input and its output in separate callbacks. +""" + +from __future__ import annotations + +import collections +from collections.abc import Awaitable, Callable, Iterator +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Literal + +from airflow.providers.common.ai.durable.base import RUN_ID_KEY, build_run_meta_key, build_step_key +from airflow.providers.common.ai.durable.storage import DurableStorage +from airflow.providers.common.ai.observability import make_task_instance_run_key +from airflow.providers.common.ai.utils.task_logger import get_task_logger +from airflow.providers.common.compat.sdk import AirflowException, get_current_context +from airflow.providers.common.compat.version_compat import AIRFLOW_V_3_3_PLUS + +if TYPE_CHECKING: + import logging + + from airflow.providers.common.ai.durable.base import DurableStorageProtocol + from airflow.sdk import Context + +log = get_task_logger() + +StepKind = Literal["model", "tool", "other"] +"""What a step is, for the end-of-run summary and for adapters that treat model steps specially.""" + + +@dataclass +class JournalStep: + """ + One position in a run, claimed by :meth:`DurableRun.claim`. + + When :attr:`replayed` is true, :attr:`payload` holds what the previous attempt + recorded and the caller returns it instead of running the step. Otherwise the + caller runs the step inside :meth:`executing` and passes its result to + :meth:`record`, or the exception it raised to :meth:`fail`; :meth:`run` does all three. + """ + + durable_run: DurableRun = field(repr=False) + position: int + name: str + kind: StepKind + fingerprint: str | None + replayable: bool = True + replayed: bool = False + payload: Any = None + # Agent runs started while this step executes; they are numbered under it. + nested_runs: int = field(default=0, repr=False) + + async def run( + self, body: Callable[[], Awaitable[Any]], *, to_record: Callable[[Any], Any] | None = None + ) -> Any: + """ + Run a step that was not replayed and record what ``body`` returns, or that it raised. + + :param to_record: Transforms the live result before it is recorded, such as + masking secrets in a tool result. The live result is returned unchanged. + """ + try: + with self.executing(): + payload = await body() + except Exception as error: + self.fail(error) + raise + self.record(payload if to_record is None else to_record(payload)) + return payload + + @contextmanager + def executing(self) -> Iterator[None]: + """ + Mark the code in the ``with`` block as this step's work. + + An agent run that starts inside the block, such as an agent a tool calls, is + numbered under this step, so the runs of later steps keep their numbers on a + retry that replays this step instead of running it. + """ + token = _EXECUTING_STEP.set(self) + try: + yield + finally: + _EXECUTING_STEP.reset(token) + + def record(self, payload: Any) -> bool: + """ + Record the step's result so a retry can replay it. + + :param payload: A JSON-compatible value. + :return: Whether it was recorded. A step that was not recorded runs again on retry, + and one claimed with ``replayable=False`` is never recorded. + """ + if not self.replayable: + return False + entry = {"name": self.name, "kind": self.kind, "fingerprint": self.fingerprint, "payload": payload} + if self.durable_run.save(self.position, entry): + self.durable_run.journal.stats.recorded[self.kind] += 1 + log.debug("Durable: recorded step", position=self.position, step=self.name) + return True + self.durable_run.journal.stats.not_recorded.append((self.kind, self.name)) + # Named here, not only in the summary: this line is logged on every path, + # including the failed attempt that Airflow retries. + log.warning( + "Durable: step not recorded; a retry runs it again, and may run the steps after it again", + position=self.position, + step=self.name, + ) + return False + + def fail(self, error: BaseException) -> None: + """ + Note that the step raised ``error``. + + Nothing is recorded for the step, so it runs again on retry. If ``error`` goes on to + fail the run, :meth:`DurableRun.fail` uses this to tell which steps the run had + already started when it failed; a run that recovers from the error is unaffected. + """ + self.durable_run._raised.append((error, self.durable_run.position)) + + +@dataclass +class JournalStats: + """What one task attempt replayed and recorded, by step kind.""" + + replayed: collections.Counter[StepKind] = field(default_factory=collections.Counter) + recorded: collections.Counter[StepKind] = field(default_factory=collections.Counter) + # Steps that ran live and could not be recorded, so a retry runs them again. + not_recorded: list[tuple[StepKind, str]] = field(default_factory=list) + + +class DurableRun: + """ + One agent run's steps in the journal. + + Positions are handed out in the order steps are claimed. A step must be claimed + before its first ``await``, so concurrent steps (parallel tool calls) take their + positions in the order they were started rather than the order they finish. + + Besides its steps, a run keeps one entry about itself: where the last attempt that + failed stopped being trustworthy, and the furthest position any attempt reached, so + a successful run deletes every step an earlier attempt left behind. + """ + + def __init__(self, journal: DurableJournal, key: str) -> None: + self.journal = journal + self.key = key + self.diverged = False + self._position = 0 + # Read on first use: where the last failed attempt's entries stop being trusted, + # and how far any attempt got. + self._diverge_from: int | None = None + self._high_water = 0 + self._meta_loaded = False + # Whether the step claimed last replayed; a step without a fingerprint replays only + # if it did, so it is never matched against a different preceding conversation. + self._previous_replayed = True + # Entries read ahead by ``peek`` and not yet claimed. + self._peeked: dict[int, dict[str, Any] | None] = {} + # Exceptions steps raised this attempt, with the next position when each was raised. + self._raised: list[tuple[BaseException, int]] = [] + # The furthest position this attempt recorded a step at. + self._last_saved = -1 + + @property + def position(self) -> int: + """The position the next claimed step will take.""" + return self._position + + def _load_meta(self) -> None: + if self._meta_loaded: + return + self._meta_loaded = True + meta = self.journal.storage.load_step(build_run_meta_key(self.key)) or {} + diverge_from, high_water = meta.get("diverge_from"), meta.get("high_water") + self._diverge_from = diverge_from if isinstance(diverge_from, int) else None + self._high_water = high_water if isinstance(high_water, int) else 0 + + def peek(self, position: int) -> dict[str, Any] | None: + """ + Return the entry the previous attempt recorded at ``position``, without claiming it. + + ``None`` when nothing usable was recorded there, or the run has already diverged. + The entry is a dict with ``name``, ``kind``, ``fingerprint`` and ``payload``. + """ + self._load_meta() + if self.diverged or (self._diverge_from is not None and position >= self._diverge_from): + return None + if position not in self._peeked: + self._peeked[position] = self.journal.storage.load_step(build_step_key(self.key, position)) + return self._peeked[position] + + def claim( + self, name: str, *, kind: StepKind, fingerprint: str | None, replayable: bool = True + ) -> JournalStep: + """ + Take the next position and look up what the previous attempt recorded there. + + :param name: What kind of step this is. Must be the same on every attempt for the + same step, and should differ between steps that are not interchangeable. + :param kind: ``"model"``, ``"tool"`` or ``"other"``. + :param fingerprint: What the step was asked, or ``None`` when it cannot be + fingerprinted. A step without one replays on a matching name, and only right + after a step that replayed. + :param replayable: ``False`` for a step that must always run live, such as a tool + whose effect Airflow cannot observe. It still takes a position, so the steps + after it keep theirs, but nothing is looked up or recorded for it. + """ + position = self._position + self._position += 1 + step = JournalStep(self, position, name, kind, fingerprint, replayable=replayable) + self._load_meta() + if not self.diverged and self._diverge_from is not None and position >= self._diverge_from: + self._diverge(position, name, reason="the previous attempt failed before reaching this step") + entry = self.peek(position) if replayable else None + self._peeked.pop(position, None) + previous_replayed, self._previous_replayed = self._previous_replayed, False + if entry is None: + return step + if entry.get("name") == name and entry.get("fingerprint") == fingerprint: + if fingerprint is None and not previous_replayed: + return step + step.replayed = self._previous_replayed = True + step.payload = entry.get("payload") + self.journal.stats.replayed[kind] += 1 + log.debug("Durable: replayed step", position=position, step=name) + return step + self._diverge( + position, + name, + reason=( + f"the previous attempt ran {entry.get('name')!r} here" + if entry.get("name") != name + else "the request differs from the previous attempt's (prompt, model, settings, " + "tools, arguments or conversation so far)" + ), + ) + return step + + async def run( + self, + name: str, + *, + kind: StepKind, + fingerprint: str | None, + body: Callable[[], Awaitable[Any]], + replayable: bool = True, + to_record: Callable[[Any], Any] | None = None, + ) -> Any: + """Replay the step at the next position, or run ``body`` and record what it returns.""" + step = self.claim(name, kind=kind, fingerprint=fingerprint, replayable=replayable) + if step.replayed: + return step.payload + return await step.run(body, to_record=to_record) + + def save(self, position: int, entry: dict[str, Any]) -> bool: + """Store ``entry`` at ``position``; ``False`` when the storage skipped it.""" + self._last_saved = max(self._last_saved, position) + return self.journal.storage.save_step(build_step_key(self.key, position), entry) + + def fail(self, error: BaseException) -> None: + """ + Record that the run failed with ``error``, so a retry does not replay its error path. + + Steps the run started before the step that raised ``error`` (tool calls running + alongside it, say) still replay on retry; everything from then on, including what + error handling recorded, runs again. + """ + failed_at = self._position + cause: BaseException | None = error + while cause is not None: + failed_at = min([failed_at, *(position for raised, position in self._raised if raised is cause)]) + cause = cause.__cause__ or cause.__context__ + self._load_meta() + # An attempt that was still replaying an earlier one when it failed, and recorded + # nothing from the failure on, leaves that earlier attempt's steps as they were. + diverge_from = failed_at if self.diverged or self._last_saved >= failed_at else self._diverge_from + meta = {"diverge_from": diverge_from, "high_water": max(self._high_water, self._position)} + self.journal.storage.save_step(build_run_meta_key(self.key), meta) + + def cleanup(self) -> None: + """Delete this run's steps. Call only once the run, and anything that consumes it, has succeeded.""" + self._load_meta() + end = max(self._high_water, self._position) + # An attempt killed before it could note how far it got may have gone further. + while self.journal.storage.load_step(build_step_key(self.key, end)) is not None: + end += 1 + keys = [build_step_key(self.key, position) for position in range(end)] + self.journal.storage.delete_steps([*keys, build_run_meta_key(self.key)]) + + def _diverge(self, position: int, name: str, *, reason: str) -> None: + self.diverged = True + self._peeked.clear() + log.warning( + "Durable: the run took a different path from the previous attempt; this step and " + "every step after it run again", + position=position, + step=name, + reason=reason, + ) + + +class DurableJournal: + """ + One task attempt's view of the durable journal. + + :param storage: Where steps are stored; see :func:`build_task_storage`. + :param clean_up_after_run: Delete a run's steps as soon as it succeeds. A journal the + operator manages leaves this off and calls :meth:`cleanup` itself after the + task's own post-run work has succeeded. + """ + + def __init__(self, storage: DurableStorageProtocol, *, clean_up_after_run: bool = False) -> None: + self.storage = storage + self.clean_up_after_run = clean_up_after_run + self.stats = JournalStats() + self._runs: list[DurableRun] = [] + self._top_level_runs = 0 + self._run_id: str | None = None + + def run_id(self, *, default: str) -> str: Review Comment: Renamed in 8f5152314f3. ########## providers/common/ai/src/airflow/providers/common/ai/durable/journal.py: ########## @@ -0,0 +1,510 @@ +# 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. +""" +The durable journal: replay an agent's completed steps when Airflow retries its task. + +An Airflow task retry is a fresh process that runs the agent again from the top. The +journal records each step an agent run completes (a model response, a tool result, a +capability operation) under the step's position in that run, and on the retry hands +back the recorded result instead of running the step again. Steps the previous +attempt never reached run live and are recorded in turn. + +A step is replayed only when the step at the same position in the previous attempt +had the same name and the same fingerprint. The name says what kind of step it was +(``agent__model.request``, ``agent__function_toolset__x.call_tool:query``); the +fingerprint, computed by the framework adapter, says what it was asked (the message +history, the tool arguments). The first mismatch means the run took a different path +from the previous attempt, so that step and every step after it run live. So does +everything the previous attempt did after a step that raised. + +A :class:`DurableJournal` belongs to one task attempt. Each agent run in it is a +:class:`DurableRun`, numbered in the order the runs start, with positions of its own, +so a task that runs several agents replays each of them independently. + +Nothing here depends on an agent framework. Payloads are JSON-compatible values and +fingerprints are opaque strings. +:class:`~airflow.providers.common.ai.durable.capability.AirflowDurability` drives the +journal for pydantic-ai. Another framework's adapter starts a run with +:meth:`DurableJournal.start_run` and calls one of two shapes of API on it: + +* :meth:`DurableRun.run` wraps a step: + ``await run.run(name, kind=..., fingerprint=..., body=...)``. +* :meth:`DurableRun.claim` and :meth:`JournalStep.record` split it in two, for a + framework whose hooks see a step's input and its output in separate callbacks. +""" + +from __future__ import annotations + +import collections +from collections.abc import Awaitable, Callable, Iterator +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Literal + +from airflow.providers.common.ai.durable.base import RUN_ID_KEY, build_run_meta_key, build_step_key +from airflow.providers.common.ai.durable.storage import DurableStorage +from airflow.providers.common.ai.observability import make_task_instance_run_key +from airflow.providers.common.ai.utils.task_logger import get_task_logger +from airflow.providers.common.compat.sdk import AirflowException, get_current_context +from airflow.providers.common.compat.version_compat import AIRFLOW_V_3_3_PLUS + +if TYPE_CHECKING: + import logging + + from airflow.providers.common.ai.durable.base import DurableStorageProtocol + from airflow.sdk import Context + +log = get_task_logger() + +StepKind = Literal["model", "tool", "other"] +"""What a step is, for the end-of-run summary and for adapters that treat model steps specially.""" + + +@dataclass +class JournalStep: + """ + One position in a run, claimed by :meth:`DurableRun.claim`. + + When :attr:`replayed` is true, :attr:`payload` holds what the previous attempt + recorded and the caller returns it instead of running the step. Otherwise the + caller runs the step inside :meth:`executing` and passes its result to + :meth:`record`, or the exception it raised to :meth:`fail`; :meth:`run` does all three. + """ + + durable_run: DurableRun = field(repr=False) + position: int + name: str + kind: StepKind + fingerprint: str | None + replayable: bool = True + replayed: bool = False + payload: Any = None + # Agent runs started while this step executes; they are numbered under it. + nested_runs: int = field(default=0, repr=False) + + async def run( + self, body: Callable[[], Awaitable[Any]], *, to_record: Callable[[Any], Any] | None = None + ) -> Any: + """ + Run a step that was not replayed and record what ``body`` returns, or that it raised. + + :param to_record: Transforms the live result before it is recorded, such as + masking secrets in a tool result. The live result is returned unchanged. + """ + try: + with self.executing(): + payload = await body() + except Exception as error: + self.fail(error) + raise + self.record(payload if to_record is None else to_record(payload)) + return payload + + @contextmanager + def executing(self) -> Iterator[None]: + """ + Mark the code in the ``with`` block as this step's work. + + An agent run that starts inside the block, such as an agent a tool calls, is + numbered under this step, so the runs of later steps keep their numbers on a + retry that replays this step instead of running it. + """ + token = _EXECUTING_STEP.set(self) + try: + yield + finally: + _EXECUTING_STEP.reset(token) + + def record(self, payload: Any) -> bool: + """ + Record the step's result so a retry can replay it. + + :param payload: A JSON-compatible value. + :return: Whether it was recorded. A step that was not recorded runs again on retry, + and one claimed with ``replayable=False`` is never recorded. + """ + if not self.replayable: + return False + entry = {"name": self.name, "kind": self.kind, "fingerprint": self.fingerprint, "payload": payload} + if self.durable_run.save(self.position, entry): + self.durable_run.journal.stats.recorded[self.kind] += 1 + log.debug("Durable: recorded step", position=self.position, step=self.name) + return True + self.durable_run.journal.stats.not_recorded.append((self.kind, self.name)) + # Named here, not only in the summary: this line is logged on every path, + # including the failed attempt that Airflow retries. + log.warning( + "Durable: step not recorded; a retry runs it again, and may run the steps after it again", + position=self.position, + step=self.name, + ) + return False + + def fail(self, error: BaseException) -> None: + """ + Note that the step raised ``error``. + + Nothing is recorded for the step, so it runs again on retry. If ``error`` goes on to + fail the run, :meth:`DurableRun.fail` uses this to tell which steps the run had + already started when it failed; a run that recovers from the error is unaffected. + """ + self.durable_run._raised.append((error, self.durable_run.position)) + + +@dataclass +class JournalStats: + """What one task attempt replayed and recorded, by step kind.""" + + replayed: collections.Counter[StepKind] = field(default_factory=collections.Counter) + recorded: collections.Counter[StepKind] = field(default_factory=collections.Counter) + # Steps that ran live and could not be recorded, so a retry runs them again. + not_recorded: list[tuple[StepKind, str]] = field(default_factory=list) + + +class DurableRun: + """ + One agent run's steps in the journal. + + Positions are handed out in the order steps are claimed. A step must be claimed + before its first ``await``, so concurrent steps (parallel tool calls) take their + positions in the order they were started rather than the order they finish. + + Besides its steps, a run keeps one entry about itself: where the last attempt that + failed stopped being trustworthy, and the furthest position any attempt reached, so + a successful run deletes every step an earlier attempt left behind. + """ + + def __init__(self, journal: DurableJournal, key: str) -> None: + self.journal = journal + self.key = key + self.diverged = False + self._position = 0 + # Read on first use: where the last failed attempt's entries stop being trusted, + # and how far any attempt got. + self._diverge_from: int | None = None + self._high_water = 0 + self._meta_loaded = False + # Whether the step claimed last replayed; a step without a fingerprint replays only + # if it did, so it is never matched against a different preceding conversation. + self._previous_replayed = True + # Entries read ahead by ``peek`` and not yet claimed. + self._peeked: dict[int, dict[str, Any] | None] = {} + # Exceptions steps raised this attempt, with the next position when each was raised. + self._raised: list[tuple[BaseException, int]] = [] + # The furthest position this attempt recorded a step at. + self._last_saved = -1 + + @property + def position(self) -> int: + """The position the next claimed step will take.""" + return self._position + + def _load_meta(self) -> None: + if self._meta_loaded: + return + self._meta_loaded = True + meta = self.journal.storage.load_step(build_run_meta_key(self.key)) or {} + diverge_from, high_water = meta.get("diverge_from"), meta.get("high_water") + self._diverge_from = diverge_from if isinstance(diverge_from, int) else None + self._high_water = high_water if isinstance(high_water, int) else 0 + + def peek(self, position: int) -> dict[str, Any] | None: + """ + Return the entry the previous attempt recorded at ``position``, without claiming it. + + ``None`` when nothing usable was recorded there, or the run has already diverged. + The entry is a dict with ``name``, ``kind``, ``fingerprint`` and ``payload``. + """ + self._load_meta() + if self.diverged or (self._diverge_from is not None and position >= self._diverge_from): + return None + if position not in self._peeked: + self._peeked[position] = self.journal.storage.load_step(build_step_key(self.key, position)) + return self._peeked[position] + + def claim( + self, name: str, *, kind: StepKind, fingerprint: str | None, replayable: bool = True + ) -> JournalStep: + """ + Take the next position and look up what the previous attempt recorded there. + + :param name: What kind of step this is. Must be the same on every attempt for the + same step, and should differ between steps that are not interchangeable. + :param kind: ``"model"``, ``"tool"`` or ``"other"``. + :param fingerprint: What the step was asked, or ``None`` when it cannot be + fingerprinted. A step without one replays on a matching name, and only right + after a step that replayed. + :param replayable: ``False`` for a step that must always run live, such as a tool + whose effect Airflow cannot observe. It still takes a position, so the steps + after it keep theirs, but nothing is looked up or recorded for it. + """ + position = self._position + self._position += 1 + step = JournalStep(self, position, name, kind, fingerprint, replayable=replayable) + self._load_meta() + if not self.diverged and self._diverge_from is not None and position >= self._diverge_from: + self._diverge(position, name, reason="the previous attempt failed before reaching this step") + entry = self.peek(position) if replayable else None + self._peeked.pop(position, None) + previous_replayed, self._previous_replayed = self._previous_replayed, False + if entry is None: + return step + if entry.get("name") == name and entry.get("fingerprint") == fingerprint: + if fingerprint is None and not previous_replayed: + return step + step.replayed = self._previous_replayed = True + step.payload = entry.get("payload") + self.journal.stats.replayed[kind] += 1 + log.debug("Durable: replayed step", position=position, step=name) + return step + self._diverge( + position, + name, + reason=( + f"the previous attempt ran {entry.get('name')!r} here" + if entry.get("name") != name + else "the request differs from the previous attempt's (prompt, model, settings, " + "tools, arguments or conversation so far)" + ), + ) + return step + + async def run( + self, + name: str, + *, + kind: StepKind, + fingerprint: str | None, + body: Callable[[], Awaitable[Any]], + replayable: bool = True, + to_record: Callable[[Any], Any] | None = None, + ) -> Any: + """Replay the step at the next position, or run ``body`` and record what it returns.""" + step = self.claim(name, kind=kind, fingerprint=fingerprint, replayable=replayable) + if step.replayed: + return step.payload + return await step.run(body, to_record=to_record) + + def save(self, position: int, entry: dict[str, Any]) -> bool: + """Store ``entry`` at ``position``; ``False`` when the storage skipped it.""" + self._last_saved = max(self._last_saved, position) + return self.journal.storage.save_step(build_step_key(self.key, position), entry) + + def fail(self, error: BaseException) -> None: + """ + Record that the run failed with ``error``, so a retry does not replay its error path. + + Steps the run started before the step that raised ``error`` (tool calls running + alongside it, say) still replay on retry; everything from then on, including what + error handling recorded, runs again. + """ + failed_at = self._position + cause: BaseException | None = error + while cause is not None: + failed_at = min([failed_at, *(position for raised, position in self._raised if raised is cause)]) + cause = cause.__cause__ or cause.__context__ + self._load_meta() + # An attempt that was still replaying an earlier one when it failed, and recorded + # nothing from the failure on, leaves that earlier attempt's steps as they were. + diverge_from = failed_at if self.diverged or self._last_saved >= failed_at else self._diverge_from + meta = {"diverge_from": diverge_from, "high_water": max(self._high_water, self._position)} + self.journal.storage.save_step(build_run_meta_key(self.key), meta) + + def cleanup(self) -> None: + """Delete this run's steps. Call only once the run, and anything that consumes it, has succeeded.""" + self._load_meta() + end = max(self._high_water, self._position) + # An attempt killed before it could note how far it got may have gone further. + while self.journal.storage.load_step(build_step_key(self.key, end)) is not None: + end += 1 + keys = [build_step_key(self.key, position) for position in range(end)] + self.journal.storage.delete_steps([*keys, build_run_meta_key(self.key)]) + + def _diverge(self, position: int, name: str, *, reason: str) -> None: + self.diverged = True + self._peeked.clear() + log.warning( + "Durable: the run took a different path from the previous attempt; this step and " + "every step after it run again", + position=position, + step=name, + reason=reason, + ) + + +class DurableJournal: + """ + One task attempt's view of the durable journal. + + :param storage: Where steps are stored; see :func:`build_task_storage`. + :param clean_up_after_run: Delete a run's steps as soon as it succeeds. A journal the + operator manages leaves this off and calls :meth:`cleanup` itself after the + task's own post-run work has succeeded. + """ + + def __init__(self, storage: DurableStorageProtocol, *, clean_up_after_run: bool = False) -> None: + self.storage = storage + self.clean_up_after_run = clean_up_after_run + self.stats = JournalStats() + self._runs: list[DurableRun] = [] + self._top_level_runs = 0 + self._run_id: str | None = None + + def run_id(self, *, default: str) -> str: + """ + Return the id that names the task's agent run on every attempt. + + The first attempt stores ``default`` beside the journal; a retry reads it back. It + is deleted with the journal when the task succeeds, so a later clear of the task + starts over with a new id. Capabilities that key their own state on the agent's + ``run_id``, such as pydantic-ai-harness ``SpendLimits`` and ``StepPersistence``, + need it to stay the same across a replay. + + :param default: The id to use when no earlier attempt stored one, such as the + first attempt's task instance id. + """ + if self._run_id is None: + entry = self.storage.load_step(RUN_ID_KEY) + stored = entry.get("run_id") if entry is not None else None + if isinstance(stored, str): + self._run_id = stored + else: + self._run_id = default + self.storage.save_step(RUN_ID_KEY, {"run_id": default}) + return self._run_id + + def start_run(self) -> DurableRun: + """ + Start an agent run. + + A run started while a step of this journal executes (an agent a tool calls) is + numbered under that step; any other run takes the next top-level number. + """ + parent = _EXECUTING_STEP.get() + if parent is not None and parent.durable_run.journal is self: + parent.nested_runs += 1 + key = f"{parent.durable_run.key}.{parent.position}.{parent.nested_runs}" + else: + key = str(self._top_level_runs) + self._top_level_runs += 1 + durable_run = DurableRun(self, key) + self._runs.append(durable_run) + return durable_run + + def cleanup(self) -> None: Review Comment: Yes, that leaked. A prefix delete would be nicer, but the task state store accessor can't list keys (it only has `get`/`set`/`delete`/`clear`, and `clear` would wipe the task's own keys too). So the journal now keeps a registry entry of every run any attempt started, and cleanup walks that instead of `self._runs`. `test_cleanup_reaches_a_run_started_by_a_step_that_now_replays` reproduces your case. 8f5152314f3 ########## providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py: ########## @@ -1,141 +0,0 @@ -# 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. -"""Caching toolset wrapper for durable execution.""" - -from __future__ import annotations - -from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any - -from pydantic_ai.toolsets.wrapper import WrapperToolset - -from airflow.providers.common.ai.durable.base import build_tool_step_key -from airflow.providers.common.ai.durable.fingerprint import fingerprint_tool_call -from airflow.providers.common.ai.utils.task_logger import get_task_logger -from airflow.providers.common.ai.utils.tool_metrics import record_tool_call -from airflow.providers.common.ai.utils.toolset_base import AirflowToolset - -if TYPE_CHECKING: - from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool - - from airflow.providers.common.ai.durable.base import DurableStorageProtocol - from airflow.providers.common.ai.durable.replay_usage import ReplayUsageLedger - from airflow.providers.common.ai.durable.step_counter import DurableStepCounter - -log = get_task_logger() - - -@dataclass -class CachingToolset(WrapperToolset[Any]): Review Comment: Dropped it (8f5152314f3). The replacement wrapper is private (leading `_`), and the registration check skips private classes, so it doesn't need an entry. ########## providers/common/ai/src/airflow/providers/common/ai/durable/capability.py: ########## @@ -0,0 +1,453 @@ +# 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. +"""Durable execution for pydantic-ai agents in Airflow tasks, on pydantic-ai's durable backend API.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Iterator, Mapping +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, ClassVar, TypeAlias, cast + +from pydantic_ai.durable_exec import ( + JSON_CODEC, + BaseDurabilityCapability, + CapabilityOperationId, + DurabilityEngineSpec, + DurableOperationId, + JournalCallableOperationBackend, + ModelRequestId, + RoleBasedOperationConfig, + ToolsetCallToolId, + ToolsetValidateToolArgumentsId, +) +from pydantic_ai.exceptions import ApprovalRequired, CallDeferred, ModelRetry, ToolFailed +from pydantic_ai.messages import ModelResponse, ToolReturn +from pydantic_ai.tools import AgentDepsT +from pydantic_ai.toolsets import DynamicToolset, FunctionToolset +from pydantic_ai.toolsets.wrapper import WrapperToolset +from pydantic_core import PydanticSerializationError + +from airflow.providers.common.ai.durable.fingerprint import fingerprint_model_request, fingerprint_tool_call +from airflow.providers.common.ai.durable.journal import current_journal +from airflow.providers.common.ai.durable.replay_usage import ReplayUsageLedger, is_successful_tool_payload +from airflow.providers.common.ai.exceptions import DurableJournalError +from airflow.providers.common.ai.utils.masking import mask_secrets +from airflow.providers.common.ai.utils.tool_metrics import record_tool_call +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset + +if TYPE_CHECKING: + from pydantic_ai.agent import EventStreamHandler + from pydantic_ai.capabilities import WrapRunHandler + from pydantic_ai.messages import ModelMessage + from pydantic_ai.models import Model, ModelRequestParameters + from pydantic_ai.run import AgentRunResult + from pydantic_ai.settings import ModelSettings + from pydantic_ai.tools import RunContext + from pydantic_ai.toolsets import AbstractToolset, ToolsetTool + + from airflow.providers.common.ai.durable.journal import DurableRun, JournalStep + +__all__ = ["AirflowDurability"] + +# What pydantic-ai's durable backend passes as ``cache_key``: a projection of the +# operation's parameters that it types as an opaque tuple. These spell out the shapes +# of the two projections fingerprinted here. +_ModelRequestKey: TypeAlias = ( + "tuple[str | None, list[ModelMessage], ModelSettings | None, ModelRequestParameters, RunContext[Any]]" +) +_ToolCallKey: TypeAlias = "tuple[str, dict[str, Any], RunContext[Any]]" + +try: + from pydantic_ai.mcp import MCPToolset +except ImportError: # the ``mcp`` extra is not installed + DURABLE_UNIT_TOOLSETS: tuple[type, ...] = (FunctionToolset, DynamicToolset) +else: + DURABLE_UNIT_TOOLSETS = (FunctionToolset, DynamicToolset, MCPToolset) +"""Leaf toolsets pydantic-ai's durable backend runs as durable units; each needs a unique ``id``.""" + +_NO_CONFIG: RoleBasedOperationConfig[None] = RoleBasedOperationConfig( + model=None, event=None, capability=None, tool=None +) + + +@dataclass(frozen=True) +class _ActiveRun: + """The journal run an agent's steps go to, and the ledger that keeps their replays out of its usage.""" + + durable_run: DurableRun + ledger: ReplayUsageLedger + + +_ACTIVE_RUN: ContextVar[_ActiveRun | None] = ContextVar("commonai_durability_run", default=None) + + +@contextmanager +def _active_run(active: _ActiveRun) -> Iterator[None]: + token = _ACTIVE_RUN.set(active) + try: + yield + finally: + _ACTIVE_RUN.reset(token) + + +def _not_journalable(error: Exception) -> BaseException: + # The same value fails the same way on every retry, so retrying the task cannot help. + return DurableJournalError( + f"durable execution could not record a step's result, because it is not JSON-serializable: {error}" + ) + + +class AirflowDurability(BaseDurabilityCapability[AgentDepsT]): + """ + Replay an agent's completed steps when Airflow retries the task it runs in. + + Attach it to a pydantic-ai ``Agent`` running inside an Airflow task. Each model + request, tool call (function, MCP and dynamic toolsets, and Airflow's own toolsets), + tool discovery and ``@durable_operation`` of another capability + is recorded in the task instance's durable journal as it completes. When the task + fails and Airflow retries it, the agent runs again from the start and every step the + previous attempt completed is replayed from the journal instead of being run again: + no second model call, no second tool side effect, no second charge against a spend + limit. See :ref:`durable-execution` for how replay is verified. + + ``AgentOperator(durable=True)`` and ``@task.agent(durable=True)`` attach it for you. + In a plain ``@task``, attach it yourself: + + .. code-block:: python + + @task(retries=3) + def research(question: str) -> str: + agent = Agent( + "anthropic:claude-sonnet-4-5", + name="researcher", + toolsets=[SQLToolset("warehouse")], + capabilities=[AirflowDurability()], + ) + return agent.run_sync(question).output + + Outside a running Airflow task the capability does nothing, so the same agent runs + normally in a test or a notebook. In a plain ``@task``, a run's steps are deleted as + soon as the run succeeds, so that clearing the task later starts it fresh: a task that + runs several agents replays only the run that failed, and runs the ones before it again. + + Toolsets the agent calls through pydantic-ai's durable backend need a unique ``id`` + (``FunctionToolset(id=...)``, ``MCPToolset(..., id=...)``), as do capabilities that + contribute ``@durable_operation`` methods; the ids name the steps in the journal. + + :param models: Extra models the run may switch to, keyed by id; see pydantic-ai's + ``BaseDurabilityCapability``. + :param event_stream_handler: Optional handler for the run's events. + :param name: Prefix for the names of the agent's steps. Defaults to the agent's + ``name``; one of the two is required. + """ + + engine_spec: ClassVar[DurabilityEngineSpec] = DurabilityEngineSpec( + engine_name="Airflow", + durable_unit_noun="step", + durable_container_noun="task", + codec=JSON_CODEC, + serialization_failure=_not_journalable, + ) + + def __init__( + self, + *, + models: Mapping[str, Model] | None = None, + event_stream_handler: EventStreamHandler[AgentDepsT] | None = None, + name: str | None = None, + ) -> None: + super().__init__(models=models, event_stream_handler=event_stream_handler, name=name) + # Keyed by agent too: pydantic-ai binds a shallow copy of this capability to each + # agent it is attached to, and the copies share this dict. + self._fingerprint_models: dict[tuple[int, str | None], Model] = {} + + @property + def in_durable_context(self) -> bool: + return current_journal() is not None + + def get_durable_operation_backend(self) -> _AirflowOperationBackend: + return _AirflowOperationBackend(self) + + def get_wrapper_toolset(self, toolset: AbstractToolset[AgentDepsT]) -> AbstractToolset[AgentDepsT] | None: + """Journal the leaf toolsets pydantic-ai's backend does not, then let the base wrap the rest.""" + journaled = toolset.visit_and_replace(self._journal_other_leaf) + return super().get_wrapper_toolset(journaled) or journaled + + def _journal_other_leaf(self, toolset: AbstractToolset[AgentDepsT]) -> AbstractToolset[AgentDepsT]: + # Function, dynamic and MCP toolsets are durable units of pydantic-ai's backend. + # Anything else, such as Airflow's own SQL, hook and managed-agent toolsets, is + # journaled here instead, or its calls would run again on every retry. + if isinstance(toolset, DURABLE_UNIT_TOOLSETS): + return toolset + return _JournaledToolset(wrapped=toolset, durability=self) + + async def before_run(self, ctx: RunContext[AgentDepsT]) -> None: + # The base refuses a run with a ``cancellation_token`` inside the durable container, + # because for its engines the container runs elsewhere and cancelling it out of band + # would break replay. Airflow's container is the task process itself: AgentOperator's + # on_kill cancels the run there, the attempt fails, and the retry replays what the + # journal recorded. That check is all the base hook does, so it is not called. + return + + async def wrap_run(self, ctx: RunContext[AgentDepsT], *, handler: WrapRunHandler) -> AgentRunResult[Any]: + journal = current_journal() + if journal is None: + return await super().wrap_run(ctx, handler=handler) + durable_run = journal.start_run() + # Keeps replays out of the usage the run counts and is limited by, which is the + # cross-attempt total when AgentOperator passes one as ``usage=``. + ledger = ReplayUsageLedger(run_usage=ctx.usage, usage_limits=ctx.usage_limits) + ledger.credit_first_replay(durable_run) + try: + with _active_run(_ActiveRun(durable_run, ledger)): + result = await super().wrap_run(ctx, handler=handler) + except BaseException as error: + durable_run.fail(error) + raise + finally: + ledger.settle() + if journal.clean_up_after_run: + durable_run.cleanup() + return result + + async def _model_for_fingerprint( + self, model_id: str | None, run_context: RunContext[AgentDepsT] + ) -> Model: + """Return the model a request with ``model_id`` goes to, resolved once per agent and id.""" + key = (id(self.agent), model_id) + if (model := self._fingerprint_models.get(key)) is None: + model = self._fingerprint_models[key] = await self._resolve_model_for_request( + model_id, run_context + ) + return model + + +class _AirflowOperationBackend(JournalCallableOperationBackend[None]): + """Runs each durable operation through the task's durable journal.""" + + def __init__(self, durability: AirflowDurability[Any]) -> None: + super().__init__( + agent_name=durability.name, default_model_id=durability.default_model_id, config=_NO_CONFIG + ) + self._durability = durability + + async def execute( + self, + *, + operation_id: DurableOperationId, + name: str, + body: Callable[[], Awaitable[object]], + cache_key: tuple[object, ...], + config: None, + ) -> object: + active = _ACTIVE_RUN.get() + if active is None: + return await body() + match operation_id: + case ModelRequestId(streaming=streaming): + return await self._model_request(active, name, body, cache_key, streaming=streaming) + case ToolsetCallToolId(): + step = _claim_tool_call(active, name, _tool_fingerprint(cache_key)) + if step.replayed: + return step.payload + # The masking wrapper AgentOperator adds sits outside this durable unit, so mask + # before recording: secrets must not reach the journal any more than the model. + return await step.run(body, to_record=mask_secrets) + case ToolsetValidateToolArgumentsId(): + # Local validation of the arguments the recorded call was made with: nothing to + # save by replaying it, and it is the same on every attempt, so it takes no position. + return await body() + case CapabilityOperationId(): + step = active.durable_run.claim(name, kind="other", fingerprint=None) + if step.replayed: + active.ledger.record_capability_replay(step.payload) + return step.payload + return await step.run(body) + case _: + # Tool discovery, compaction, event handling, and the operations later + # pydantic-ai versions add: replayed by name and position. + return await active.durable_run.run(name, kind="other", fingerprint=None, body=body) + + async def _model_request( + self, + active: _ActiveRun, + name: str, + body: Callable[[], Awaitable[object]], + cache_key: tuple[object, ...], + *, + streaming: bool, + ) -> object: + model_id, messages, model_settings, parameters, run_context = cast("_ModelRequestKey", cache_key) + fingerprint = await self._fingerprint_model_request( + model_id, messages, model_settings, parameters, run_context + ) + ledger = active.ledger + continuation = ledger.begin_model_request(messages) + had_request_credit = ledger.settle() + step = active.durable_run.claim(name, kind="model", fingerprint=fingerprint) + if step.replayed: + response = _model_response(step.payload, streaming=streaming) + ledger.record_model_replay(response, continuation=continuation) + ledger.track_chain(response, parameters) + if response.state != "suspended": + ledger.credit_successors(active.durable_run, step.position) + return step.payload + ledger.record_live_model_request(had_request_credit=had_request_credit) + payload = await step.run(body) + ledger.track_chain(_model_response(payload, streaming=streaming), parameters) + return payload + + async def _fingerprint_model_request( + self, + model_id: str | None, + messages: list[ModelMessage], + model_settings: ModelSettings | None, + parameters: ModelRequestParameters, + run_context: RunContext[Any], + ) -> str | None: + # Fingerprint the request as the model will prepare it, not the raw arguments. + # ``prepare_request`` merges the model's own settings and applies profile + # transforms (thinking, native tools, output mode) before the provider sees the + # request, so a change that lives only on the connection, such as a different + # temperature, still invalidates the recorded response. It is pure, so calling it + # here as well as in the request itself is safe. + model = await self._durability._model_for_fingerprint(model_id, run_context) + prepared_settings, prepared_parameters = model.prepare_request(model_settings, parameters) + return fingerprint_model_request( + f"{model.system}:{model.model_name}", messages, prepared_settings, prepared_parameters + ) + + +@dataclass +class _JournaledToolset(WrapperToolset[Any]): + """ + Journals the calls of a leaf toolset that pydantic-ai's durable backend does not wrap. + + Results, and the control-flow exceptions a tool raises for the model (``ModelRetry``, + ``ToolFailed``, ``ApprovalRequired``, ``CallDeferred``), are recorded as values, so a + retry replays them exactly. Any other exception is recorded as a failure and the call + runs again on retry. A result that is not JSON-serializable fails the task, as it does + for the toolsets pydantic-ai runs as durable steps. + """ + + durability: AirflowDurability[Any] = field(repr=False, kw_only=True) + + def visit_and_replace( + self, visitor: Callable[[AbstractToolset[Any]], AbstractToolset[Any]] + ) -> AbstractToolset[Any]: + # A durable unit, like pydantic-ai's own durable toolsets: a later visit, such as the + # journaling pass of the next run, must not reach the leaf and wrap it again. + return self + + async def call_tool( + self, name: str, tool_args: dict[str, Any], ctx: RunContext[Any], tool: ToolsetTool[Any] + ) -> Any: + active = _ACTIVE_RUN.get() + if active is None: + return await self.wrapped.call_tool(name, tool_args, ctx, tool) + leaf = self.wrapped + step = _claim_tool_call( + active, + f"{self.durability.name}__airflow_toolset__{leaf.id or type(leaf).__name__}.call_tool:{name}", + fingerprint_tool_call(name, tool_args, ctx.tool_call_id), + # A toolset whose calls act on a system Airflow cannot observe, such as a managed + # agent, runs them again on every attempt. + replayable=leaf.replayable if isinstance(leaf, AirflowToolset) else True, + ) + if step.replayed: + if isinstance(leaf, AirflowToolset): + record_tool_call(type(leaf).__name__, "replayed") + return _decode_tool_payload(step.payload) + return await _run_and_record_tool(step, lambda: leaf.call_tool(name, tool_args, ctx, tool)) + + +def _claim_tool_call( + active: _ActiveRun, name: str, fingerprint: str | None, *, replayable: bool = True +) -> JournalStep: + """Claim a tool call's step and keep the ledger's count of tool calls in step with it.""" + step = active.durable_run.claim(name, kind="tool", fingerprint=fingerprint, replayable=replayable) + if not step.replayed: + active.ledger.record_live_tool_call(step.position) + elif is_successful_tool_payload(step.payload): + active.ledger.record_tool_replay(step.position) + return step + + +async def _run_and_record_tool(step: JournalStep, call: Callable[[], Awaitable[Any]]) -> Any: + try: + with step.executing(): + result = await call() + except ModelRetry as e: + _record_tool_payload(step, {"kind": "model_retry", "message": e.message}) + raise + except ToolFailed as e: + _record_tool_payload(step, {"kind": "tool_failed", "message": e.message}) + raise + except ApprovalRequired as e: + _record_tool_payload(step, {"kind": "approval_required", "metadata": e.metadata}) + raise + except CallDeferred as e: + _record_tool_payload(step, {"kind": "call_deferred", "metadata": e.metadata}) + raise + except Exception as error: + step.fail(error) + raise + key = "tool_return" if isinstance(result, ToolReturn) else "result" + _record_tool_payload(step, {"kind": "tool_return", key: result}) + return result + + +def _record_tool_payload(step: JournalStep, payload: dict[str, Any]) -> None: + try: + encoded = JSON_CODEC.dump(Any, payload) + except (PydanticSerializationError, TypeError, ValueError) as error: + # As for the toolsets pydantic-ai runs as durable steps: the same value fails the same + # way on every attempt, and the model could not have read it either. + raise _not_journalable(error) from error + step.record(mask_secrets(encoded)) + + +def _decode_tool_payload(payload: Any) -> Any: + kind = payload.get("kind") if isinstance(payload, dict) else None + match kind: + case "tool_return" if "tool_return" in payload: + return JSON_CODEC.load(ToolReturn, payload["tool_return"]) + case "tool_return": + return payload["result"] + case "model_retry": + raise ModelRetry(payload["message"]) + case "tool_failed": + raise ToolFailed(payload["message"]) + case "approval_required": + raise ApprovalRequired(metadata=payload.get("metadata")) + case "call_deferred": + raise CallDeferred(metadata=payload.get("metadata")) + raise DurableJournalError(f"durable execution found a tool result it cannot replay: {kind!r}") + + +def _model_response(payload: object, *, streaming: bool) -> ModelResponse: Review Comment: Agreed, and renamed. In 8f5152314f3 a record that no longer loads gets rejected: the journal warns, marks the run diverged, and that step plus everything after it runs live. Model responses get validated as `ModelResponse` and Airflow-toolset results against their full shape. For tool results that pydantic-ai records itself I can only check the `kind` field, since the type it decodes them with is private. Covered in `TestUnloadableRecords`. -- 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]
