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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,22 @@ Provides streaming orchestration that delivers LLM responses token-by-token for
- Progressive rendering
- Interruptible generation

### Durable completed-tool checkpoints

A host may register a zero-argument `session.durable_checkpoint` capability on
the coordinator. The Loop invokes it once after a normal tool batch has settled
and all results have been appended in their original order, before a subsequent
provider request (including budget finalization). The host owns storage and
must persist the complete canonical context before returning; count-only
debouncing is not a durable-write guarantee.

The callable may be synchronous or awaitable. Literal `False`, an ordinary
exception, or a non-callable registration fails the turn explicitly without
another provider dispatch. Cancellation propagates with appended results intact;
this does not add a cancellation-recovery policy. A missing capability preserves
the behavior of existing hosts. The Loop itself does not write transcripts.
This checkpoint is not the earlier `tool:post` event, which precedes result append.

### Provider budget preflight

When a provider exposes the optional `request_budget` capability, the loop
Expand Down
23 changes: 23 additions & 0 deletions amplifier_module_loop_streaming/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5391,6 +5391,29 @@ async def count_view(
}
)

# tool:post precedes this ordered append. Only now is the
# settled batch ready for a host-owned durable checkpoint.
# Keep older hosts compatible, but do not dispatch again when
# a host that promises persistence cannot complete its write.
checkpoint = (
coordinator.get_capability("session.durable_checkpoint")
if coordinator
else None
)
if checkpoint is not None:
try:
if not callable(checkpoint):
raise TypeError("session.durable_checkpoint is not callable")
checkpoint_result = checkpoint()
if inspect.isawaitable(checkpoint_result):
checkpoint_result = await checkpoint_result
if checkpoint_result is False:
raise RuntimeError("session.durable_checkpoint returned False")
except Exception as exc:
raise RuntimeError(
"Durable session checkpoint failed; refusing further provider dispatch"
) from exc

# Add exactly one finalization call only when the bounded budget
# prevented a continuation. A normal no-tool break at the cap is a
# natural completion and must return without a duplicate provider call.
Expand Down
204 changes: 204 additions & 0 deletions tests/test_durable_checkpoint.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,204 @@
"""The host checkpoints only a complete, settled, canonically ordered batch."""

from __future__ import annotations

import asyncio
from typing import ClassVar

import pytest
from amplifier_core import ContextLengthError, ToolResult
from amplifier_core.models import ProviderInfo

from amplifier_module_loop_streaming import StreamingOrchestrator
from tests.test_ephemeral_cache_persist_mode import (
MockContext as RequestContext,
)
from tests.test_ephemeral_cache_persist_mode import (
MockCoordinator,
MockResponse,
ScriptedHooks,
ToolCallStub,
)


class MockContext(RequestContext):
async def get_messages(self):
return list(self._messages)


class BatchProvider:
def __init__(self, count=2, *, overflow=False):
self.count = count
self.calls = 0
self.overflow = overflow

def get_info(self):
return ProviderInfo(
id="fixture",
display_name="Fixture",
credential_env_vars=[],
defaults={"model": "fixture"},
)

async def complete(self, request, **kwargs):
self.calls += 1
if self.calls == 2 and self.overflow:
raise ContextLengthError("synthetic overflow after completed tools")
return MockResponse("done")

def parse_tool_calls(self, response):
calls = []
if self.calls == 1:
for i in range(self.count):
call = ToolCallStub(f"call-{i}")
call.arguments = {"index": i}
calls.append(call)
return calls


class OrderedTool:
"""Finish the second call first, without clock-dependent sleeps."""

name = "mock_tool"
description = "deterministic concurrent tool"
input_schema: ClassVar[dict[str, object]] = {
"type": "object",
"properties": {"index": {"type": "integer"}},
}

def __init__(self, count=2):
self.count = count
self.second_finished = asyncio.Event()
self.completed = []

async def execute(self, arguments):
index = arguments["index"]
if index == 0 and self.count == 2:
await self.second_finished.wait()
self.completed.append(index)
if index == 1:
self.second_finished.set()
return ToolResult(success=True, output=f"result-{index}")


def _receipts(context):
return [
(m["tool_call_id"], m["content"])
for m in context._messages
if m["role"] == "tool"
]


async def _execute(provider, context, coordinator, tool, *, max_iterations=3):
return await StreamingOrchestrator({"max_iterations": max_iterations}).execute(
"work",
context,
{"main": provider},
{"mock_tool": tool},
ScriptedHooks({}),
coordinator,
)


@pytest.mark.asyncio
@pytest.mark.parametrize("count", [1, 2])
@pytest.mark.parametrize("max_iterations", [1, 3])
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_checkpoint_precedes_next_dispatch_and_budget_finalization(
count, max_iterations, asynchronous
):
context, coordinator = MockContext(), MockCoordinator()
provider, tool = BatchProvider(count, overflow=True), OrderedTool(count)
checkpoints = []

def checkpoint():
assert provider.calls == 1
assert tool.completed == ([0] if count == 1 else [1, 0])
checkpoints.append(_receipts(context))

async def async_checkpoint():
checkpoint()

coordinator.register_capability(
"session.durable_checkpoint", async_checkpoint if asynchronous else checkpoint
)
with pytest.raises(ContextLengthError, match="synthetic overflow"):
await _execute(
provider, context, coordinator, tool, max_iterations=max_iterations
)

assert checkpoints == [[(f"call-{i}", f"result-{i}") for i in range(count)]]
assert provider.calls == 2


@pytest.mark.asyncio
@pytest.mark.parametrize("max_iterations", [1, 3])
@pytest.mark.parametrize(
"failure", ["malformed", "false", "exception", "async-false", "async-exception"]
)
async def test_checkpoint_failure_prevents_any_next_provider_dispatch(
failure, max_iterations
):
context, coordinator = MockContext(), MockCoordinator()
provider, tool = BatchProvider(), OrderedTool()

def checkpoint():
if "exception" in failure:
raise OSError("disk unavailable")
return False

async def async_checkpoint():
return checkpoint()

callback = (
"invalid"
if failure == "malformed"
else (async_checkpoint if failure.startswith("async") else checkpoint)
)
coordinator.register_capability("session.durable_checkpoint", callback)
with pytest.raises(
RuntimeError, match="Durable session checkpoint failed"
) as caught:
await _execute(
provider, context, coordinator, tool, max_iterations=max_iterations
)
assert caught.value.__cause__ is not None
assert provider.calls == 1
assert _receipts(context) == [("call-0", "result-0"), ("call-1", "result-1")]
assert tool.completed == [1, 0]


@pytest.mark.asyncio
async def test_checkpoint_cancellation_propagates_with_appended_results_intact():
context, coordinator = MockContext(), MockCoordinator()
provider, tool = BatchProvider(), OrderedTool()

async def checkpoint():
raise asyncio.CancelledError()

coordinator.register_capability("session.durable_checkpoint", checkpoint)
with pytest.raises(asyncio.CancelledError):
await _execute(provider, context, coordinator, tool)
assert provider.calls == 1
assert _receipts(context) == [("call-0", "result-0"), ("call-1", "result-1")]


@pytest.mark.asyncio
async def test_absent_checkpoint_preserves_old_host_behavior():
context, coordinator = MockContext(), MockCoordinator()
provider, tool = BatchProvider(), OrderedTool()
assert await _execute(provider, context, coordinator, tool) == "done"
assert provider.calls == 2
assert _receipts(context) == [("call-0", "result-0"), ("call-1", "result-1")]


@pytest.mark.asyncio
async def test_no_tool_turn_does_not_checkpoint():
context, coordinator = MockContext(), MockCoordinator()
coordinator.register_capability(
"session.durable_checkpoint", lambda: pytest.fail("no batch to checkpoint")
)
assert (
await _execute(BatchProvider(count=0), context, coordinator, OrderedTool())
== "done"
)
Loading