fix(dream): advance cursor when tool errors were recovered, and report why a run did not complete

This commit is contained in:
flobo3
2026-08-22 02:05:40 +08:00
committed by chengyongru
parent dd993b4f70
commit d853ac239f
4 changed files with 162 additions and 11 deletions
+52 -8
View File
@@ -54,10 +54,21 @@ if TYPE_CHECKING:
class DreamRunProgress:
"""Track tool failures that make a nominally completed Dream run unsafe to advance."""
"""Track tool failures that make a nominally completed Dream run unsafe to advance.
A failure in an earlier tool round is tolerated when the model observed
the error, retried, and the final tool round ran clean. Only failures the
model never got to correct invalidate the run.
"""
def __init__(self) -> None:
self.had_tool_errors = False
self.last_tool_round_had_errors: bool | None = None
@property
def recovered_from_tool_errors(self) -> bool:
"""True when errors occurred but the final tool round finished clean."""
return self.had_tool_errors and self.last_tool_round_had_errors is False
async def __call__(
self,
@@ -65,11 +76,17 @@ class DreamRunProgress:
tool_events: list[dict[str, Any]] | None = None,
**_kwargs: Any,
) -> None:
if any(
isinstance(cast(object, event), dict) and event.get("phase") == "error"
for event in tool_events or ()
):
events = [
event for event in tool_events or ()
if isinstance(cast(object, event), dict)
]
round_had_errors = any(event.get("phase") == "error" for event in events)
if round_had_errors:
self.had_tool_errors = True
# Terminal payloads ("end"/"error") close a tool round; the most
# recent closed round decides whether earlier failures were recovered.
if any(event.get("phase") in ("end", "error") for event in events):
self.last_tool_round_had_errors = round_had_errors
class MemoryStore:
@@ -689,12 +706,39 @@ class MemoryStore:
resp: object | None,
*,
had_tool_errors: bool = False,
recovered_tool_errors: bool = False,
) -> bool:
"""Return True only when a Dream turn completed without tool failures."""
"""Return True only when a Dream turn finished cleanly enough to advance.
Tool failures from earlier rounds are acceptable when the final tool
round ran clean (``recovered_tool_errors``): the model observed the
failure, corrected it, and produced a consistent final state. Failures
in the final round still invalidate the run.
"""
metadata = getattr(resp, "metadata", None)
if had_tool_errors or not isinstance(metadata, dict):
if not isinstance(metadata, dict):
return False
return cast(dict[str, Any], metadata).get("_stop_reason") == "completed"
if cast(dict[str, Any], metadata).get("_stop_reason") != "completed":
return False
return not had_tool_errors or recovered_tool_errors
@staticmethod
def dream_incompletion_reason(
resp: object | None,
*,
had_tool_errors: bool = False,
recovered_tool_errors: bool = False,
) -> str:
"""Human-readable explanation of why a Dream run cannot advance."""
metadata = getattr(resp, "metadata", None)
if isinstance(metadata, dict):
stop_reason = cast(dict[str, Any], metadata).get("_stop_reason", "unknown")
else:
stop_reason = "missing response metadata"
parts = [f"stop_reason: {stop_reason}"]
if had_tool_errors and not recovered_tool_errors:
parts.append("unrecovered tool errors")
return ", ".join(parts)
# -- message formatting utility ------------------------------------------
+7 -1
View File
@@ -536,6 +536,7 @@ def _run_gateway(
completed = MemoryStore.dream_run_completed(
resp,
had_tool_errors=progress.had_tool_errors,
recovered_tool_errors=progress.recovered_from_tool_errors,
)
if completed:
store.set_last_dream_cursor(last_cursor)
@@ -552,7 +553,12 @@ def _run_gateway(
)
else:
logger.warning(
"Dream cron job did not complete; cursor remains at {}",
"Dream cron job did not complete ({}); cursor remains at {}",
MemoryStore.dream_incompletion_reason(
resp,
had_tool_errors=progress.had_tool_errors,
recovered_tool_errors=progress.recovered_from_tool_errors,
),
store.get_last_dream_cursor(),
)
except Exception:
+7 -1
View File
@@ -462,6 +462,7 @@ async def cmd_dream(ctx: CommandContext) -> OutboundMessage:
completed = MemoryStore.dream_run_completed(
resp,
had_tool_errors=progress.had_tool_errors,
recovered_tool_errors=progress.recovered_from_tool_errors,
)
if completed:
store.set_last_dream_cursor(last_cursor)
@@ -470,8 +471,13 @@ async def cmd_dream(ctx: CommandContext) -> OutboundMessage:
else:
content = f"Dream completed in {elapsed:.1f}s; no memory changes."
else:
reason = MemoryStore.dream_incompletion_reason(
resp,
had_tool_errors=progress.had_tool_errors,
recovered_tool_errors=progress.recovered_from_tool_errors,
)
content = (
f"Dream did not complete after {elapsed:.1f}s; "
f"Dream did not complete after {elapsed:.1f}s ({reason}); "
"memory cursor was not advanced."
)
except Exception as e:
+96 -1
View File
@@ -2,7 +2,7 @@
import pytest
from nanobot.agent.memory import MemoryStore
from nanobot.agent.memory import DreamRunProgress, MemoryStore
from nanobot.config.schema import ModelPresetConfig
from nanobot.providers.base import LLMResponse
from nanobot.security.workspace_access import (
@@ -186,6 +186,101 @@ class TestBuildDreamPrompt:
assert "Always strip these bracketed tags from saved memory content" in prompt
class TestDreamRunCompletion:
"""DreamRunProgress + dream_run_completed gate cursor advancement."""
class _Resp:
def __init__(self, stop_reason: str = "completed") -> None:
self.metadata = {"_stop_reason": stop_reason}
@staticmethod
async def _feed(progress: DreamRunProgress, *batches: list[dict]) -> None:
for batch in batches:
await progress("", tool_events=batch)
@pytest.mark.asyncio
async def test_clean_run_completes(self):
progress = DreamRunProgress()
await self._feed(
progress,
[{"phase": "start"}, {"phase": "end"}],
)
assert not progress.had_tool_errors
assert not progress.recovered_from_tool_errors
assert MemoryStore.dream_run_completed(
self._Resp(), had_tool_errors=progress.had_tool_errors,
)
@pytest.mark.asyncio
async def test_recovered_tool_error_completes(self):
"""A failed edit_file retry that later succeeds must not block the cursor."""
progress = DreamRunProgress()
await self._feed(
progress,
[{"phase": "start"}, {"phase": "error"}],
[{"phase": "start"}, {"phase": "end"}],
)
assert progress.had_tool_errors
assert progress.recovered_from_tool_errors
assert MemoryStore.dream_run_completed(
self._Resp(),
had_tool_errors=progress.had_tool_errors,
recovered_tool_errors=progress.recovered_from_tool_errors,
)
@pytest.mark.asyncio
async def test_error_in_final_tool_round_blocks(self):
"""An error the model never corrected still invalidates the run."""
progress = DreamRunProgress()
await self._feed(
progress,
[{"phase": "start"}, {"phase": "end"}],
[{"phase": "start"}, {"phase": "error"}],
)
assert progress.had_tool_errors
assert not progress.recovered_from_tool_errors
assert not MemoryStore.dream_run_completed(
self._Resp(),
had_tool_errors=progress.had_tool_errors,
recovered_tool_errors=progress.recovered_from_tool_errors,
)
@pytest.mark.asyncio
async def test_start_only_events_do_not_close_a_round(self):
progress = DreamRunProgress()
await self._feed(progress, [{"phase": "start"}])
assert not progress.had_tool_errors
assert progress.last_tool_round_had_errors is None
assert not progress.recovered_from_tool_errors
@pytest.mark.asyncio
async def test_thought_progress_calls_are_ignored(self):
progress = DreamRunProgress()
await progress("thinking...", tool_hint=True)
await progress("", file_edit_events=[{"phase": "edit"}])
assert not progress.had_tool_errors
assert progress.last_tool_round_had_errors is None
def test_non_completed_stop_reason_blocks_despite_clean_tools(self):
assert not MemoryStore.dream_run_completed(self._Resp("max_iterations"))
assert not MemoryStore.dream_run_completed(None)
def test_incompletion_reason_names_the_cause(self):
reason = MemoryStore.dream_incompletion_reason(
self._Resp("max_iterations"),
had_tool_errors=True,
recovered_tool_errors=False,
)
assert "stop_reason: max_iterations" in reason
assert "unrecovered tool errors" in reason
assert MemoryStore.dream_incompletion_reason(
self._Resp(), had_tool_errors=True, recovered_tool_errors=True,
) == "stop_reason: completed"
assert MemoryStore.dream_incompletion_reason(None) == (
"stop_reason: missing response metadata"
)
class TestDreamTools:
def test_dream_tools_are_restricted_to_file_edits(self, store):
tools = store.build_dream_tools()