Lee-W commented on code in PR #74385:
URL: https://github.com/apache/airflow/pull/74385#discussion_r4203754845


##########
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:
   ```suggestion
       def get_run_id(self, *, default: str) -> str:
   ```



##########
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:
   `cleanup` only walks `self._runs`. A nested run whose parent tool step was 
replayed never starts again, so its keys (`run_<n>.<pos>.<m>_step_*`) are never 
visited, and on the task state store they are written with `NEVER_EXPIRE`.
   
   Should we delete by fix instead? for example,  everything under 
`__commonai_durable__run_` for this task instance. 



##########
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:
   ```suggestion
   def _load_model_response(payload: object, *, streaming: bool) -> 
ModelResponse:
   ```
   
   `JSON_CODEC.load(ModelResponse, ...)` has no guard, so a payload that no 
longer validates raises out of the replay path on every retry.
   
   We should log a warning and call `durable_run._diverge(...)` and run the 
body live, for both the model and tool kinds.



-- 
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]

Reply via email to