From 4a273b61e463ca24641ac29565c0d4f9fdd58620 Mon Sep 17 00:00:00 2001 From: divo12 Date: Fri, 7 Aug 2026 15:36:16 +0530 Subject: [PATCH 1/6] feat(state): Hermes-style human rewind for shadow checkpoints MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Expose Session.list_checkpoints / restore_checkpoint so operators can restore the worktree and truncate the matching transcript turns — the gap vs Hermes /rollback after auto pre-mutate snaps landed. Co-authored-by: Cursor --- src/dream/_factory.py | 21 +++++ src/dream/config/paths.py | 5 + src/dream/engine/_engine.py | 6 ++ src/dream/harness.py | 3 + src/dream/session.py | 72 ++++++++++++++ src/dream/state/__init__.py | 8 ++ src/dream/state/shadow/__init__.py | 10 ++ src/dream/state/shadow/_manager.py | 28 ++++++ src/dream/state/shadow/_rewind.py | 72 ++++++++++++++ src/dream/state/shadow/_types.py | 12 +++ tests/test_session/test_checkpoint_rewind.py | 92 ++++++++++++++++++ tests/test_state/test_combined_restore.py | 99 ++++++++++++++++++++ tests/test_state/test_transcript_rewind.py | 99 ++++++++++++++++++++ 13 files changed, 527 insertions(+) create mode 100644 src/dream/state/shadow/_rewind.py create mode 100644 tests/test_session/test_checkpoint_rewind.py create mode 100644 tests/test_state/test_combined_restore.py create mode 100644 tests/test_state/test_transcript_rewind.py diff --git a/src/dream/_factory.py b/src/dream/_factory.py index 187db0bd..5f517697 100644 --- a/src/dream/_factory.py +++ b/src/dream/_factory.py @@ -62,6 +62,11 @@ build_session_skill_registry, render_skill_catalogue, ) +from dream.state.shadow import ( + ShadowCheckpointHook, + ShadowCheckpointManager, + ShadowCheckpointStore, +) from dream.subagents._async_delegation import AsyncDelegationManager from dream.subagents._declaration import SubagentSet from dream.tasks import ( @@ -106,6 +111,7 @@ def build_harness( policy_warning_sink: PolicyWarningSink | None = None, env: Mapping[str, str] | None = None, wake_model: str | None = None, + shadow_checkpoints: bool = True, ) -> Harness: """Build a Harness whose engine factory produces a real, tool-wired engine. @@ -123,6 +129,10 @@ def build_harness( no caller effort. Pass ``skill_registry`` to supply your own (it wins); pass ``skills=False`` to disable discovery entirely. + ``shadow_checkpoints`` (default True) registers Hermes-style pre-mutate + filesystem snapshots and exposes :meth:`~dream.session.Session.restore_checkpoint` + for operator rewind (FS + transcript). + Workspace memory (the durable per-project record store under :func:`~dream.memory.project_memory_dir`) is wired by default: its catalogue lands in the system prompt and the ``memory_search`` / @@ -224,12 +234,18 @@ def build_harness( # so a scheduler tick loop knows where to poll, and `paths` carries the # env-resolved roots. del wake_model # ponytail: compat no-op — the wake runtime is gone + checkpoint_manager: ShadowCheckpointManager | None = None + if shadow_checkpoints: + checkpoint_manager = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=paths.checkpoints_dir), + ) config = HarnessConfig( working_dir=working_dir, task_manager=task_manager, delegations=AsyncDelegationManager(), cron_registry_path=task_context.cron_registry_path, paths=paths, + checkpoint_manager=checkpoint_manager, # MCP connect + plugin import are async/IO, so they hang off the # async-open chokepoint (``Harness._ensure_open``) rather than running # in this sync factory. ``None`` when both surfaces are disabled so the @@ -247,6 +263,10 @@ def build_harness( ), ) harness = Harness(config) + if checkpoint_manager is not None: + harness.register_hook( + ShadowCheckpointHook(manager=checkpoint_manager, working_dir=working_dir) + ) # The engine factory closes over ``harness`` (not a hooks snapshot) so the # spec-13 lifecycle executor is assembled lazily at session construction @@ -718,4 +738,5 @@ def _build_session_engine( model=options.model or model, hook_executor=hook_executor, delegations=harness.config.delegations, + checkpoint_manager=harness.config.checkpoint_manager, ) diff --git a/src/dream/config/paths.py b/src/dream/config/paths.py index dfd3c32d..fdee76a5 100644 --- a/src/dream/config/paths.py +++ b/src/dream/config/paths.py @@ -203,6 +203,11 @@ def memory_dir(self) -> Path: def skills_dir(self) -> Path: return self.home / "skills" + @property + def checkpoints_dir(self) -> Path: + """Shared shadow-git store root (Hermes ``~/.hermes/checkpoints/store``).""" + return self.home / "checkpoints" + # --- the one explicit side effect --- def ensure(self) -> DreamPaths: diff --git a/src/dream/engine/_engine.py b/src/dream/engine/_engine.py index 4b1c6f68..68bb0e22 100644 --- a/src/dream/engine/_engine.py +++ b/src/dream/engine/_engine.py @@ -40,6 +40,7 @@ from dream.tools._registry import ToolRegistry if TYPE_CHECKING: + from dream.state.shadow import ShadowCheckpointManager from dream.subagents._async_delegation import AsyncDelegationManager @@ -86,6 +87,9 @@ class QueryEngine: # Role name from session metadata (planner/generator/evaluator) for STOP hooks. role: str | None = None delegations: AsyncDelegationManager | None = None + # Hermes shadow checkpoints — shared across sessions on the harness; enables + # :meth:`Session.list_checkpoints` / :meth:`Session.restore_checkpoint`. + checkpoint_manager: ShadowCheckpointManager | None = None def make_session_config( self, @@ -146,6 +150,7 @@ def build_query_engine( hook_executor: HookExecutor | None = None, orientation: OrientationConfig | None = None, delegations: AsyncDelegationManager | None = None, + checkpoint_manager: ShadowCheckpointManager | None = None, ) -> QueryEngine: """Wrap a ``ToolRegistry`` in the canonical dispatcher and bind a streamer. @@ -186,6 +191,7 @@ def build_query_engine( orientation=orientation, role=role_raw if isinstance(role_raw, str) else None, delegations=delegations, + checkpoint_manager=checkpoint_manager, ) diff --git a/src/dream/harness.py b/src/dream/harness.py index ac11e770..80aba7cf 100644 --- a/src/dream/harness.py +++ b/src/dream/harness.py @@ -36,6 +36,7 @@ SprintGoalProvider, ) from dream.sprint import EvaluatorPropose, GeneratorRespond + from dream.state.shadow import ShadowCheckpointManager from dream.subagents._async_delegation import AsyncDelegationManager from dream.tasks import BackgroundTaskManager @@ -80,6 +81,8 @@ class HarnessConfig: # so the runtime reuses the exact same roots (DREAM_HOME honoured) rather # than re-resolving and risking divergence. paths: DreamPaths | None = None + # Shared Hermes-style shadow checkpoint manager (FS snaps + human rewind). + checkpoint_manager: ShadowCheckpointManager | None = None extra: dict[str, Any] = field(default_factory=dict) _engine_factory: EngineFactory | None = None # Async setup run once before the first session — MCP connect + plugin diff --git a/src/dream/session.py b/src/dream/session.py index 2b875320..d5a7f4f0 100644 --- a/src/dream/session.py +++ b/src/dream/session.py @@ -57,6 +57,11 @@ if TYPE_CHECKING: from dream.engine._engine import QueryEngine + from dream.state.shadow import ( + CheckpointSnapshot, + CombinedRestoreResult, + ShadowCheckpointManager, + ) def _tool_failed_marker(tool_name: str) -> str: @@ -177,6 +182,73 @@ def transcript(self) -> list[ConversationMessage]: """ return self._transcript + @property + def checkpoint_manager(self) -> ShadowCheckpointManager | None: + """Shadow FS checkpoint manager when the bound engine was built with one.""" + engine = self._engine + if engine is None: + return None + return engine.checkpoint_manager + + def list_checkpoints(self) -> list[CheckpointSnapshot]: + """List shadow checkpoints for this session's working directory (newest first).""" + engine = self._engine + manager = self.checkpoint_manager + if engine is None or manager is None: + return [] + return manager.list_for(engine.working_dir) + + def restore_checkpoint( + self, + commit_sha: str | None = None, + *, + rewind_turns: int = 1, + ) -> CombinedRestoreResult: + """Hermes-style human rewind: restore FS and truncate the transcript. + + ``commit_sha=None`` selects the newest checkpoint. Refuses while a + ``send`` is in flight so restore cannot race the live turn loop. + """ + from dream.state.shadow import CombinedRestoreResult, RestoreOutcome, RestoreResult + + if self._active: + raise RuntimeError("cannot restore checkpoint while a send is in flight") + engine = self._engine + manager = self.checkpoint_manager + if engine is None or manager is None: + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.DISABLED, + detail="no checkpoint manager bound on this session", + ), + messages=tuple(self._transcript), + transcript_removed=0, + ) + + sha = commit_sha + if sha is None: + listed = manager.list_for(engine.working_dir) + if not listed: + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.NOT_FOUND, + detail="no checkpoints for working directory", + ), + messages=tuple(self._transcript), + transcript_removed=0, + ) + sha = listed[0].commit_sha + + result = manager.restore_and_rewind( + engine.working_dir, + commit_sha=sha, + messages=self._transcript, + rewind_turns=rewind_turns, + ) + if result.fs.outcome is RestoreOutcome.RESTORED: + self._transcript = list(result.messages) + return result + async def send(self, prompt: str) -> AsyncIterator[Event]: """Submit a user prompt and stream typed events back. diff --git a/src/dream/state/__init__.py b/src/dream/state/__init__.py index 8751028e..4b03d856 100644 --- a/src/dream/state/__init__.py +++ b/src/dream/state/__init__.py @@ -5,21 +5,29 @@ from dream.state.shadow import ( CheckpointOutcome, CheckpointReason, + CheckpointSnapshot, + CombinedRestoreResult, MutatingToolName, RestoreOutcome, + RestoreResult, ShadowCheckpointConfig, ShadowCheckpointHook, ShadowCheckpointManager, ShadowCheckpointStore, + rewind_transcript, ) __all__ = [ "CheckpointOutcome", "CheckpointReason", + "CheckpointSnapshot", + "CombinedRestoreResult", "MutatingToolName", "RestoreOutcome", + "RestoreResult", "ShadowCheckpointConfig", "ShadowCheckpointHook", "ShadowCheckpointManager", "ShadowCheckpointStore", + "rewind_transcript", ] diff --git a/src/dream/state/shadow/__init__.py b/src/dream/state/shadow/__init__.py index 4b519d5c..0d58a5ba 100644 --- a/src/dream/state/shadow/__init__.py +++ b/src/dream/state/shadow/__init__.py @@ -4,11 +4,17 @@ from dream.state.shadow._hook import ShadowCheckpointHook from dream.state.shadow._manager import ShadowCheckpointManager +from dream.state.shadow._rewind import ( + is_user_prompt_message, + rewind_transcript, + user_prompt_indices, +) from dream.state.shadow._store import ShadowCheckpointStore from dream.state.shadow._types import ( CheckpointOutcome, CheckpointReason, CheckpointSnapshot, + CombinedRestoreResult, EnsureResult, MutatingToolName, RestoreOutcome, @@ -20,6 +26,7 @@ "CheckpointOutcome", "CheckpointReason", "CheckpointSnapshot", + "CombinedRestoreResult", "EnsureResult", "MutatingToolName", "RestoreOutcome", @@ -28,4 +35,7 @@ "ShadowCheckpointHook", "ShadowCheckpointManager", "ShadowCheckpointStore", + "is_user_prompt_message", + "rewind_transcript", + "user_prompt_indices", ] diff --git a/src/dream/state/shadow/_manager.py b/src/dream/state/shadow/_manager.py index 58100a0a..a931a137 100644 --- a/src/dream/state/shadow/_manager.py +++ b/src/dream/state/shadow/_manager.py @@ -4,13 +4,17 @@ import contextlib import shutil +from collections.abc import Sequence from pathlib import Path +from dream.engine._messages import ConversationMessage +from dream.state.shadow._rewind import rewind_transcript from dream.state.shadow._store import ShadowCheckpointStore from dream.state.shadow._types import ( CheckpointOutcome, CheckpointReason, CheckpointSnapshot, + CombinedRestoreResult, EnsureResult, RestoreOutcome, RestoreResult, @@ -158,6 +162,30 @@ def restore(self, working_dir: Path, *, commit_sha: str) -> RestoreResult: return RestoreResult(outcome=RestoreOutcome.RESTORED, restored_to=commit_sha[:8]) + def restore_and_rewind( + self, + working_dir: Path, + *, + commit_sha: str, + messages: Sequence[ConversationMessage], + rewind_turns: int = 1, + ) -> CombinedRestoreResult: + """Restore the worktree and truncate the conversation (Hermes ``/rollback``). + + Transcript rewind runs only after a successful filesystem restore so a + failed snap never desyncs chat from disk. ``rewind_turns=0`` keeps the + transcript unchanged (FS-only restore). + """ + fs = self.restore(working_dir, commit_sha=commit_sha) + if fs.outcome is not RestoreOutcome.RESTORED: + return CombinedRestoreResult(fs=fs, messages=tuple(messages), transcript_removed=0) + kept, removed = rewind_transcript(messages, turns=rewind_turns) + return CombinedRestoreResult( + fs=fs, + messages=tuple(kept), + transcript_removed=removed, + ) + def _take(self, working_dir: Path, reason: CheckpointReason) -> EnsureResult: err = self._store.ensure_initialized(working_dir) if err: diff --git a/src/dream/state/shadow/_rewind.py b/src/dream/state/shadow/_rewind.py new file mode 100644 index 00000000..1b7cf526 --- /dev/null +++ b/src/dream/state/shadow/_rewind.py @@ -0,0 +1,72 @@ +"""Transcript rewind helpers — Hermes ``/rollback`` conversation half. + +Filesystem restore lives on :class:`ShadowCheckpointManager`. These pure +functions truncate the in-memory conversation so FS and chat stay aligned +after an operator rewind. +""" + +from __future__ import annotations + +from collections.abc import Sequence + +from dream.engine._messages import ConversationMessage, TextBlock, ToolResultBlock + + +def is_user_prompt_message(message: ConversationMessage) -> bool: + """True when ``message`` starts a human/agent turn (not a tool-result batch). + + Tool-result user messages only carry :class:`ToolResultBlock`s. A real + prompt has at least one :class:`TextBlock` (and may also mix other blocks). + """ + if message.role != "user": + return False + has_text = False + for block in message.content: + if isinstance(block, ToolResultBlock): + continue + if isinstance(block, TextBlock): + if block.text.strip(): + has_text = True + continue + # Image / other non-tool content still counts as a prompt. + return True + return has_text + + +def user_prompt_indices(messages: Sequence[ConversationMessage]) -> tuple[int, ...]: + """Indices of user-prompt messages in order (oldest first).""" + return tuple(i for i, msg in enumerate(messages) if is_user_prompt_message(msg)) + + +def rewind_transcript( + messages: Sequence[ConversationMessage], + *, + turns: int, +) -> tuple[list[ConversationMessage], int]: + """Drop the last ``turns`` user-prompt cycles from ``messages``. + + A cycle starts at a user-prompt message and includes everything until + (but not including) the next user-prompt, or EOF. ``turns=0`` is a no-op. + + Returns ``(kept_messages, removed_count)``. + """ + if turns < 0: + raise ValueError(f"turns must be >= 0; got {turns}") + if turns == 0: + return list(messages), 0 + + indices = user_prompt_indices(messages) + if not indices: + return list(messages), 0 + + drop_from = indices[-turns] if turns <= len(indices) else indices[0] + kept = list(messages[:drop_from]) + removed = len(messages) - len(kept) + return kept, removed + + +__all__ = [ + "is_user_prompt_message", + "rewind_transcript", + "user_prompt_indices", +] diff --git a/src/dream/state/shadow/_types.py b/src/dream/state/shadow/_types.py index 5f3b3eb4..300c1b6e 100644 --- a/src/dream/state/shadow/_types.py +++ b/src/dream/state/shadow/_types.py @@ -6,6 +6,8 @@ from enum import Enum, StrEnum from pathlib import Path +from dream.engine._messages import ConversationMessage + class MutatingToolName(StrEnum): """Built-in tools that may mutate the working tree.""" @@ -87,6 +89,15 @@ class RestoreResult: detail: str = "" +@dataclass(frozen=True, slots=True) +class CombinedRestoreResult: + """Filesystem restore plus optional transcript rewind (Hermes ``/rollback``).""" + + fs: RestoreResult + messages: tuple[ConversationMessage, ...] = () + transcript_removed: int = 0 + + _TOOL_TO_REASON: dict[MutatingToolName, CheckpointReason] = { MutatingToolName.WRITE_FILE: CheckpointReason.BEFORE_WRITE_FILE, MutatingToolName.EDIT_FILE: CheckpointReason.BEFORE_EDIT_FILE, @@ -109,6 +120,7 @@ def reason_for_tool(tool_name: str) -> CheckpointReason | None: "CheckpointOutcome", "CheckpointReason", "CheckpointSnapshot", + "CombinedRestoreResult", "EnsureResult", "MutatingToolName", "RestoreOutcome", diff --git a/tests/test_session/test_checkpoint_rewind.py b/tests/test_session/test_checkpoint_rewind.py new file mode 100644 index 00000000..76fe07f0 --- /dev/null +++ b/tests/test_session/test_checkpoint_rewind.py @@ -0,0 +1,92 @@ +"""Eval: Session.list_checkpoints / restore_checkpoint (Hermes human rewind).""" + +from __future__ import annotations + +from pathlib import Path +from typing import Never + +import pytest + +from dream.engine._engine import QueryEngine +from dream.engine._messages import ConversationMessage, TextBlock +from dream.session import Session +from dream.state.shadow import ( + CheckpointReason, + RestoreOutcome, + ShadowCheckpointConfig, + ShadowCheckpointManager, + ShadowCheckpointStore, +) + + +class _NoopStreamer: + async def stream_turn(self, *_args: object, **_kwargs: object) -> Never: + raise RuntimeError("streamer unused") + yield # pragma: no cover — makes this an async generator type + + +class _NoopDispatcher: + async def dispatch(self, *_args: object, **_kwargs: object) -> tuple[str, bool]: + return ("", False) + + +def _engine( + working_dir: Path, manager: ShadowCheckpointManager | None +) -> QueryEngine: + return QueryEngine( + streamer=_NoopStreamer(), # type: ignore[arg-type] + dispatcher=_NoopDispatcher(), # type: ignore[arg-type] + session_id="s1", + working_dir=working_dir, + checkpoint_manager=manager, + ) + + +def test_session_restore_checkpoint_rewinds_fs_and_transcript(tmp_path: Path) -> None: + work = tmp_path / "proj" + work.mkdir() + (work / "f.txt").write_text("clean\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "ck"), + config=ShadowCheckpointConfig(enabled=True), + ) + snap = mgr.ensure(work, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert snap.snapshot is not None + + (work / "f.txt").write_text("dirty\n", encoding="utf-8") + session = Session(id="s1", _engine=_engine(work, mgr)) + session.transcript.extend( + [ + ConversationMessage(role="user", content=[TextBlock(text="edit it")]), + ConversationMessage(role="assistant", content=[TextBlock(text="done")]), + ConversationMessage(role="user", content=[TextBlock(text="more")]), + ConversationMessage(role="assistant", content=[TextBlock(text="more done")]), + ] + ) + + result = session.restore_checkpoint(rewind_turns=1) + assert result.fs.outcome is RestoreOutcome.RESTORED + assert (work / "f.txt").read_text(encoding="utf-8") == "clean\n" + assert len(session.transcript) == 2 + assert session.transcript[0].text == "edit it" + assert result.transcript_removed == 2 + + +def test_session_list_checkpoints_empty_without_manager(tmp_path: Path) -> None: + session = Session(id="s1", _engine=_engine(tmp_path, None)) + assert session.list_checkpoints() == [] + result = session.restore_checkpoint() + assert result.fs.outcome is RestoreOutcome.DISABLED + + +def test_session_restore_refuses_during_active_send(tmp_path: Path) -> None: + work = tmp_path / "proj" + work.mkdir() + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "ck"), + config=ShadowCheckpointConfig(enabled=True), + ) + session = Session(id="s1", _engine=_engine(work, mgr)) + session._active = True + with pytest.raises(RuntimeError, match="in flight"): + session.restore_checkpoint() diff --git a/tests/test_state/test_combined_restore.py b/tests/test_state/test_combined_restore.py new file mode 100644 index 00000000..f3297f69 --- /dev/null +++ b/tests/test_state/test_combined_restore.py @@ -0,0 +1,99 @@ +"""Eval: combined FS + transcript restore (Hermes /rollback).""" + +from __future__ import annotations + +from pathlib import Path + +import pytest + +from dream.engine._messages import ConversationMessage, TextBlock +from dream.state.shadow import ( + CheckpointReason, + CombinedRestoreResult, + RestoreOutcome, + ShadowCheckpointConfig, + ShadowCheckpointManager, + ShadowCheckpointStore, +) + + +@pytest.fixture() +def work_dir(tmp_path: Path) -> Path: + d = tmp_path / "project" + d.mkdir() + (d / "README.md").write_text("v1\n", encoding="utf-8") + return d + + +@pytest.fixture() +def mgr(tmp_path: Path) -> ShadowCheckpointManager: + return ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "checkpoints"), + config=ShadowCheckpointConfig(enabled=True, max_snapshots=10), + ) + + +def _msgs(*texts: str) -> list[ConversationMessage]: + out: list[ConversationMessage] = [] + for i, text in enumerate(texts): + role = "user" if i % 2 == 0 else "assistant" + out.append(ConversationMessage(role=role, content=[TextBlock(text=text)])) + return out + + +def test_restore_and_rewind_aligns_fs_and_transcript( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + taken = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert taken.outcome.value == "taken" + assert taken.snapshot is not None + sha = taken.snapshot.commit_sha + + (work_dir / "README.md").write_text("v2-broken\n", encoding="utf-8") + messages = _msgs("fix the bug", "I edited README", "looks good?", "ship it") + + result = mgr.restore_and_rewind( + work_dir, + commit_sha=sha, + messages=messages, + rewind_turns=1, + ) + assert isinstance(result, CombinedRestoreResult) + assert result.fs.outcome is RestoreOutcome.RESTORED + assert (work_dir / "README.md").read_text(encoding="utf-8") == "v1\n" + assert result.messages == tuple(messages[:2]) + assert result.transcript_removed == 2 + + +def test_restore_and_rewind_fs_only_when_rewind_zero( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + taken = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_EDIT_FILE) + assert taken.snapshot is not None + (work_dir / "README.md").write_text("changed\n", encoding="utf-8") + messages = _msgs("a", "b") + + result = mgr.restore_and_rewind( + work_dir, + commit_sha=taken.snapshot.commit_sha, + messages=messages, + rewind_turns=0, + ) + assert result.fs.outcome is RestoreOutcome.RESTORED + assert result.messages == tuple(messages) + assert result.transcript_removed == 0 + + +def test_restore_and_rewind_propagates_fs_failure( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + messages = _msgs("a", "b") + result = mgr.restore_and_rewind( + work_dir, + commit_sha="deadbeef" * 5, + messages=messages, + rewind_turns=1, + ) + assert result.fs.outcome is RestoreOutcome.NOT_FOUND + assert result.messages == tuple(messages) + assert result.transcript_removed == 0 diff --git a/tests/test_state/test_transcript_rewind.py b/tests/test_state/test_transcript_rewind.py new file mode 100644 index 00000000..f22e4315 --- /dev/null +++ b/tests/test_state/test_transcript_rewind.py @@ -0,0 +1,99 @@ +"""Transcript rewind — Hermes /rollback conversation half (pure).""" + +from __future__ import annotations + +from dream.engine._messages import ConversationMessage, TextBlock, ToolResultBlock, ToolUseBlock +from dream.state.shadow._rewind import ( + is_user_prompt_message, + rewind_transcript, + user_prompt_indices, +) + + +def _user(text: str) -> ConversationMessage: + return ConversationMessage(role="user", content=[TextBlock(text=text)]) + + +def _assistant(text: str) -> ConversationMessage: + return ConversationMessage(role="assistant", content=[TextBlock(text=text)]) + + +def _assistant_tools(tool_id: str, name: str) -> ConversationMessage: + return ConversationMessage( + role="assistant", + content=[ToolUseBlock(id=tool_id, name=name, input={"path": "a.py"})], + ) + + +def _tool_results(tool_id: str) -> ConversationMessage: + return ConversationMessage( + role="user", + content=[ToolResultBlock(tool_use_id=tool_id, content="ok")], + ) + + +def test_is_user_prompt_ignores_tool_result_only_messages() -> None: + assert is_user_prompt_message(_user("do the thing")) is True + assert is_user_prompt_message(_tool_results("t1")) is False + assert is_user_prompt_message(_assistant("hi")) is False + + +def test_user_prompt_indices_finds_prompt_boundaries() -> None: + messages = [ + _user("task A"), + _assistant_tools("t1", "write_file"), + _tool_results("t1"), + _assistant("done A"), + _user("task B"), + _assistant("done B"), + ] + assert user_prompt_indices(messages) == (0, 4) + + +def test_rewind_zero_turns_is_noop() -> None: + messages = [_user("a"), _assistant("b")] + kept, removed = rewind_transcript(messages, turns=0) + assert kept == messages + assert removed == 0 + + +def test_rewind_one_turn_drops_last_prompt_cycle() -> None: + messages = [ + _user("task A"), + _assistant("done A"), + _user("task B"), + _assistant_tools("t1", "edit_file"), + _tool_results("t1"), + _assistant("done B"), + ] + kept, removed = rewind_transcript(messages, turns=1) + assert kept == messages[:2] + assert removed == 4 + + +def test_rewind_two_turns_clears_both_cycles() -> None: + messages = [ + _user("A"), + _assistant("a"), + _user("B"), + _assistant("b"), + ] + kept, removed = rewind_transcript(messages, turns=2) + assert kept == [] + assert removed == 4 + + +def test_rewind_more_than_available_clears_all() -> None: + messages = [_user("only"), _assistant("ok")] + kept, removed = rewind_transcript(messages, turns=9) + assert kept == [] + assert removed == 2 + + +def test_rewind_rejects_negative_turns() -> None: + try: + rewind_transcript([_user("x")], turns=-1) + except ValueError as exc: + assert "turns" in str(exc) + else: + raise AssertionError("expected ValueError") From 94489055b1f0cceded5de1ac540c4fd70c79e8fd Mon Sep 17 00:00:00 2001 From: divyansh_b Date: Fri, 7 Aug 2026 11:05:29 +0000 Subject: [PATCH 2/6] fix(shadow): scope checkpoints by session Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- src/dream/config/paths.py | 2 +- src/dream/state/shadow/_hook.py | 8 ++++-- src/dream/state/shadow/_manager.py | 28 ++++++++++++++----- tests/test_state/test_shadow_checkpoint.py | 31 +++++++++++++++++++++- 4 files changed, 59 insertions(+), 10 deletions(-) diff --git a/src/dream/config/paths.py b/src/dream/config/paths.py index fdee76a5..6404b919 100644 --- a/src/dream/config/paths.py +++ b/src/dream/config/paths.py @@ -205,7 +205,7 @@ def skills_dir(self) -> Path: @property def checkpoints_dir(self) -> Path: - """Shared shadow-git store root (Hermes ``~/.hermes/checkpoints/store``).""" + """Shared shadow-git store root under ``$DREAM_HOME/checkpoints``.""" return self.home / "checkpoints" # --- the one explicit side effect --- diff --git a/src/dream/state/shadow/_hook.py b/src/dream/state/shadow/_hook.py index 2c49e442..855654c9 100644 --- a/src/dream/state/shadow/_hook.py +++ b/src/dream/state/shadow/_hook.py @@ -34,7 +34,9 @@ def __init__( async def __call__(self, event: HookEvent, payload: Mapping[str, Any]) -> HookResult: if event is HookEvent.USER_PROMPT_SUBMIT: - self._manager.begin_turn() + raw_session_id = payload.get("session_id") + session_id = str(raw_session_id) if raw_session_id is not None else None + self._manager.begin_turn(session_id) return HookResult() if event is not HookEvent.PRE_TOOL_USE: @@ -45,7 +47,9 @@ async def __call__(self, event: HookEvent, payload: Mapping[str, Any]) -> HookRe if reason is None: return HookResult() - self._manager.ensure(self._working_dir, reason=reason) + raw_session_id = payload.get("session_id") + session_id = str(raw_session_id) if raw_session_id is not None else None + self._manager.ensure(self._working_dir, reason=reason, session_id=session_id) return HookResult() diff --git a/src/dream/state/shadow/_manager.py b/src/dream/state/shadow/_manager.py index a931a137..b11c83d7 100644 --- a/src/dream/state/shadow/_manager.py +++ b/src/dream/state/shadow/_manager.py @@ -41,18 +41,24 @@ def __init__( ) -> None: self._store = store self._config = config or ShadowCheckpointConfig() - self._checkpointed_dirs: set[Path] = set() + self._checkpointed_dirs: dict[str | None, set[Path]] = {} self._git_available: bool | None = None @property def config(self) -> ShadowCheckpointConfig: return self._config - def begin_turn(self) -> None: + def begin_turn(self, session_id: str | None = None) -> None: """Reset per-turn dedup (call on each USER_PROMPT_SUBMIT / agent turn).""" - self._checkpointed_dirs.clear() + self._checkpointed_dirs.pop(session_id, None) - def ensure(self, working_dir: Path, *, reason: CheckpointReason) -> EnsureResult: + def ensure( + self, + working_dir: Path, + *, + reason: CheckpointReason, + session_id: str | None = None, + ) -> EnsureResult: """Take a checkpoint if enabled and not already done this turn.""" if not self._config.enabled: return EnsureResult(outcome=CheckpointOutcome.DISABLED) @@ -70,7 +76,8 @@ def ensure(self, working_dir: Path, *, reason: CheckpointReason) -> EnsureResult if abs_dir == Path("/").resolve() or abs_dir == Path.home().resolve(): return EnsureResult(outcome=CheckpointOutcome.DIRECTORY_TOO_BROAD) - if abs_dir in self._checkpointed_dirs: + checkpointed_dirs = self._checkpointed_dirs.setdefault(session_id, set()) + if abs_dir in checkpointed_dirs: return EnsureResult(outcome=CheckpointOutcome.ALREADY_THIS_TURN) try: @@ -81,7 +88,7 @@ def ensure(self, working_dir: Path, *, reason: CheckpointReason) -> EnsureResult # Only suppress retries after a conclusive snap (or no-op). Transient # FAILED must leave the turn open so a later mutate can try again. if result.outcome is not CheckpointOutcome.FAILED: - self._checkpointed_dirs.add(abs_dir) + checkpointed_dirs.add(abs_dir) return result def list_for(self, working_dir: Path) -> list[CheckpointSnapshot]: @@ -176,6 +183,15 @@ def restore_and_rewind( failed snap never desyncs chat from disk. ``rewind_turns=0`` keeps the transcript unchanged (FS-only restore). """ + if rewind_turns < 0: + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.FAILED, + detail=f"turns must be >= 0; got {rewind_turns}", + ), + messages=tuple(messages), + transcript_removed=0, + ) fs = self.restore(working_dir, commit_sha=commit_sha) if fs.outcome is not RestoreOutcome.RESTORED: return CombinedRestoreResult(fs=fs, messages=tuple(messages), transcript_removed=0) diff --git a/tests/test_state/test_shadow_checkpoint.py b/tests/test_state/test_shadow_checkpoint.py index ea4a0a4e..42002851 100644 --- a/tests/test_state/test_shadow_checkpoint.py +++ b/tests/test_state/test_shadow_checkpoint.py @@ -68,6 +68,30 @@ def test_dedup_same_turn(mgr: ShadowCheckpointManager, work_dir: Path) -> None: assert second.outcome is CheckpointOutcome.ALREADY_THIS_TURN +def test_dedup_is_scoped_to_session(mgr: ShadowCheckpointManager, work_dir: Path) -> None: + first = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE, session_id="a") + (work_dir / "README.md").write_text("changed\n", encoding="utf-8") + second = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_BASH, session_id="b") + assert first.outcome is CheckpointOutcome.TAKEN + assert second.outcome is CheckpointOutcome.TAKEN + + +def test_negative_rewind_does_not_restore( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + taken = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert taken.snapshot is not None + (work_dir / "README.md").write_text("mutated\n", encoding="utf-8") + result = mgr.restore_and_rewind( + work_dir, + commit_sha=taken.snapshot.commit_sha, + messages=[], + rewind_turns=-1, + ) + assert result.fs.outcome is RestoreOutcome.FAILED + assert (work_dir / "README.md").read_text(encoding="utf-8") == "mutated\n" + + def test_new_turn_allows_another_when_changed(mgr: ShadowCheckpointManager, work_dir: Path) -> None: assert ( mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE).outcome @@ -229,5 +253,10 @@ async def test_hook_begin_turn_on_user_prompt(mgr: ShadowCheckpointManager, work await hook(HookEvent.USER_PROMPT_SUBMIT, {"session_id": "s1", "prompt": "go"}) (work_dir / "README.md").write_text("next\n", encoding="utf-8") assert ( - mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_BASH).outcome is CheckpointOutcome.TAKEN + mgr.ensure( + work_dir, + reason=CheckpointReason.BEFORE_BASH, + session_id="s1", + ).outcome + is CheckpointOutcome.TAKEN ) From 221f890205208875be60956141927e03228f6695 Mon Sep 17 00:00:00 2001 From: divyansh_b Date: Fri, 7 Aug 2026 11:11:32 +0000 Subject: [PATCH 3/6] fix: bound rewind session state and track prompt boundaries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- src/dream/session.py | 8 ++ src/dream/state/shadow/__init__.py | 8 +- src/dream/state/shadow/_hook.py | 17 ++-- src/dream/state/shadow/_manager.py | 25 +++++- src/dream/state/shadow/_rewind.py | 43 ++-------- tests/test_session.py | 2 + tests/test_state/test_combined_restore.py | 3 + tests/test_state/test_shadow_checkpoint.py | 1 + tests/test_state/test_transcript_rewind.py | 93 +++++++--------------- 9 files changed, 83 insertions(+), 117 deletions(-) diff --git a/src/dream/session.py b/src/dream/session.py index d5a7f4f0..da604978 100644 --- a/src/dream/session.py +++ b/src/dream/session.py @@ -141,6 +141,7 @@ def __init__( self.cost = SessionCost() self._engine = _engine self._transcript: list[ConversationMessage] = [] + self._prompt_indices: list[int] = [] self._cancel_event: asyncio.Event | None = None self._closed = False # Single-flight guard: ``Session`` keeps per-call cancel state on the @@ -243,10 +244,14 @@ def restore_checkpoint( engine.working_dir, commit_sha=sha, messages=self._transcript, + prompt_indices=self._prompt_indices, rewind_turns=rewind_turns, ) if result.fs.outcome is RestoreOutcome.RESTORED: self._transcript = list(result.messages) + self._prompt_indices = [ + index for index in self._prompt_indices if index < len(self._transcript) + ] return result async def send(self, prompt: str) -> AsyncIterator[Event]: @@ -274,6 +279,7 @@ async def send(self, prompt: str) -> AsyncIterator[Event]: resume = list(self._transcript) if self._transcript else None user_msg = ConversationMessage(role="user", content=[TextBlock(text=prompt)]) self._transcript.append(user_msg) + self._prompt_indices.append(len(self._transcript) - 1) config = self._engine.make_session_config() inner: AsyncGenerator[Any, None] = run_session( # type: ignore[assignment] @@ -449,6 +455,7 @@ def _apply_compaction(self) -> None: carryover = engine.carryover_metadata if carryover is not None and carryover.last_compacted_transcript is not None: self._transcript[:] = list(carryover.last_compacted_transcript) + self._prompt_indices.clear() carryover.last_compacted_transcript = None return compactor = engine.compactor @@ -470,6 +477,7 @@ def _apply_compaction(self) -> None: ) if result is not None: self._transcript[:] = new_transcript + self._prompt_indices.clear() async def cancel(self) -> None: """Cancel the in-flight ``send``, if any. diff --git a/src/dream/state/shadow/__init__.py b/src/dream/state/shadow/__init__.py index 0d58a5ba..ba3f4809 100644 --- a/src/dream/state/shadow/__init__.py +++ b/src/dream/state/shadow/__init__.py @@ -4,11 +4,7 @@ from dream.state.shadow._hook import ShadowCheckpointHook from dream.state.shadow._manager import ShadowCheckpointManager -from dream.state.shadow._rewind import ( - is_user_prompt_message, - rewind_transcript, - user_prompt_indices, -) +from dream.state.shadow._rewind import rewind_transcript from dream.state.shadow._store import ShadowCheckpointStore from dream.state.shadow._types import ( CheckpointOutcome, @@ -35,7 +31,5 @@ "ShadowCheckpointHook", "ShadowCheckpointManager", "ShadowCheckpointStore", - "is_user_prompt_message", "rewind_transcript", - "user_prompt_indices", ] diff --git a/src/dream/state/shadow/_hook.py b/src/dream/state/shadow/_hook.py index 855654c9..a78f9cc2 100644 --- a/src/dream/state/shadow/_hook.py +++ b/src/dream/state/shadow/_hook.py @@ -32,11 +32,14 @@ def __init__( self._manager = manager self._working_dir = working_dir + @staticmethod + def _session_id(payload: Mapping[str, Any]) -> str | None: + raw_session_id = payload.get("session_id") + return str(raw_session_id) if raw_session_id is not None else None + async def __call__(self, event: HookEvent, payload: Mapping[str, Any]) -> HookResult: if event is HookEvent.USER_PROMPT_SUBMIT: - raw_session_id = payload.get("session_id") - session_id = str(raw_session_id) if raw_session_id is not None else None - self._manager.begin_turn(session_id) + self._manager.begin_turn(self._session_id(payload)) return HookResult() if event is not HookEvent.PRE_TOOL_USE: @@ -47,9 +50,11 @@ async def __call__(self, event: HookEvent, payload: Mapping[str, Any]) -> HookRe if reason is None: return HookResult() - raw_session_id = payload.get("session_id") - session_id = str(raw_session_id) if raw_session_id is not None else None - self._manager.ensure(self._working_dir, reason=reason, session_id=session_id) + self._manager.ensure( + self._working_dir, + reason=reason, + session_id=self._session_id(payload), + ) return HookResult() diff --git a/src/dream/state/shadow/_manager.py b/src/dream/state/shadow/_manager.py index b11c83d7..2be6e92b 100644 --- a/src/dream/state/shadow/_manager.py +++ b/src/dream/state/shadow/_manager.py @@ -4,6 +4,7 @@ import contextlib import shutil +from collections import OrderedDict from collections.abc import Sequence from pathlib import Path @@ -41,7 +42,7 @@ def __init__( ) -> None: self._store = store self._config = config or ShadowCheckpointConfig() - self._checkpointed_dirs: dict[str | None, set[Path]] = {} + self._checkpointed_dirs: OrderedDict[str | None, set[Path]] = OrderedDict() self._git_available: bool | None = None @property @@ -77,6 +78,9 @@ def ensure( return EnsureResult(outcome=CheckpointOutcome.DIRECTORY_TOO_BROAD) checkpointed_dirs = self._checkpointed_dirs.setdefault(session_id, set()) + self._checkpointed_dirs.move_to_end(session_id) + while len(self._checkpointed_dirs) > 64: + self._checkpointed_dirs.popitem(last=False) if abs_dir in checkpointed_dirs: return EnsureResult(outcome=CheckpointOutcome.ALREADY_THIS_TURN) @@ -175,6 +179,7 @@ def restore_and_rewind( *, commit_sha: str, messages: Sequence[ConversationMessage], + prompt_indices: Sequence[int], rewind_turns: int = 1, ) -> CombinedRestoreResult: """Restore the worktree and truncate the conversation (Hermes ``/rollback``). @@ -192,10 +197,26 @@ def restore_and_rewind( messages=tuple(messages), transcript_removed=0, ) + if rewind_turns > 0 and rewind_turns > len(prompt_indices): + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.FAILED, + detail=( + "requested rewind boundary is unavailable " + f"for {rewind_turns} turn(s)" + ), + ), + messages=tuple(messages), + transcript_removed=0, + ) fs = self.restore(working_dir, commit_sha=commit_sha) if fs.outcome is not RestoreOutcome.RESTORED: return CombinedRestoreResult(fs=fs, messages=tuple(messages), transcript_removed=0) - kept, removed = rewind_transcript(messages, turns=rewind_turns) + kept, removed = rewind_transcript( + messages, + prompt_indices=prompt_indices, + turns=rewind_turns, + ) return CombinedRestoreResult( fs=fs, messages=tuple(kept), diff --git a/src/dream/state/shadow/_rewind.py b/src/dream/state/shadow/_rewind.py index 1b7cf526..5bde7dec 100644 --- a/src/dream/state/shadow/_rewind.py +++ b/src/dream/state/shadow/_rewind.py @@ -9,44 +9,16 @@ from collections.abc import Sequence -from dream.engine._messages import ConversationMessage, TextBlock, ToolResultBlock - - -def is_user_prompt_message(message: ConversationMessage) -> bool: - """True when ``message`` starts a human/agent turn (not a tool-result batch). - - Tool-result user messages only carry :class:`ToolResultBlock`s. A real - prompt has at least one :class:`TextBlock` (and may also mix other blocks). - """ - if message.role != "user": - return False - has_text = False - for block in message.content: - if isinstance(block, ToolResultBlock): - continue - if isinstance(block, TextBlock): - if block.text.strip(): - has_text = True - continue - # Image / other non-tool content still counts as a prompt. - return True - return has_text - - -def user_prompt_indices(messages: Sequence[ConversationMessage]) -> tuple[int, ...]: - """Indices of user-prompt messages in order (oldest first).""" - return tuple(i for i, msg in enumerate(messages) if is_user_prompt_message(msg)) +from dream.engine._messages import ConversationMessage def rewind_transcript( messages: Sequence[ConversationMessage], *, + prompt_indices: Sequence[int], turns: int, ) -> tuple[list[ConversationMessage], int]: - """Drop the last ``turns`` user-prompt cycles from ``messages``. - - A cycle starts at a user-prompt message and includes everything until - (but not including) the next user-prompt, or EOF. ``turns=0`` is a no-op. + """Drop the last ``turns`` prompt cycles using Session-owned boundaries. Returns ``(kept_messages, removed_count)``. """ @@ -55,18 +27,17 @@ def rewind_transcript( if turns == 0: return list(messages), 0 - indices = user_prompt_indices(messages) - if not indices: + if turns > len(prompt_indices): + raise ValueError("requested rewind boundary is unavailable") + if not prompt_indices: return list(messages), 0 - drop_from = indices[-turns] if turns <= len(indices) else indices[0] + drop_from = prompt_indices[-turns] kept = list(messages[:drop_from]) removed = len(messages) - len(kept) return kept, removed __all__ = [ - "is_user_prompt_message", "rewind_transcript", - "user_prompt_indices", ] diff --git a/tests/test_session.py b/tests/test_session.py index 6cce1236..a5791836 100644 --- a/tests/test_session.py +++ b/tests/test_session.py @@ -358,6 +358,7 @@ async def work() -> tuple[SubagentResult, ...]: engine = _engine(FakeStreamer([]), FakeDispatcher()) engine.delegations = manager session = Session(id="s1", _engine=engine) + session._prompt_indices.append(0) await session.cancel() @@ -688,3 +689,4 @@ def test_translate_compaction_applies_compacted_shape_to_transcript() -> None: # now the compacted shape, not the stale full history. assert post_blobs < pre_blobs assert cleared > 0 + assert session._prompt_indices == [] diff --git a/tests/test_state/test_combined_restore.py b/tests/test_state/test_combined_restore.py index f3297f69..c79792b4 100644 --- a/tests/test_state/test_combined_restore.py +++ b/tests/test_state/test_combined_restore.py @@ -56,6 +56,7 @@ def test_restore_and_rewind_aligns_fs_and_transcript( work_dir, commit_sha=sha, messages=messages, + prompt_indices=(0, 2), rewind_turns=1, ) assert isinstance(result, CombinedRestoreResult) @@ -77,6 +78,7 @@ def test_restore_and_rewind_fs_only_when_rewind_zero( work_dir, commit_sha=taken.snapshot.commit_sha, messages=messages, + prompt_indices=(0,), rewind_turns=0, ) assert result.fs.outcome is RestoreOutcome.RESTORED @@ -92,6 +94,7 @@ def test_restore_and_rewind_propagates_fs_failure( work_dir, commit_sha="deadbeef" * 5, messages=messages, + prompt_indices=(0,), rewind_turns=1, ) assert result.fs.outcome is RestoreOutcome.NOT_FOUND diff --git a/tests/test_state/test_shadow_checkpoint.py b/tests/test_state/test_shadow_checkpoint.py index 42002851..30b0c1ca 100644 --- a/tests/test_state/test_shadow_checkpoint.py +++ b/tests/test_state/test_shadow_checkpoint.py @@ -86,6 +86,7 @@ def test_negative_rewind_does_not_restore( work_dir, commit_sha=taken.snapshot.commit_sha, messages=[], + prompt_indices=(), rewind_turns=-1, ) assert result.fs.outcome is RestoreOutcome.FAILED diff --git a/tests/test_state/test_transcript_rewind.py b/tests/test_state/test_transcript_rewind.py index f22e4315..7dfbe7e7 100644 --- a/tests/test_state/test_transcript_rewind.py +++ b/tests/test_state/test_transcript_rewind.py @@ -1,98 +1,59 @@ -"""Transcript rewind — Hermes /rollback conversation half (pure).""" +"""Transcript rewind — Session-owned prompt boundaries.""" from __future__ import annotations -from dream.engine._messages import ConversationMessage, TextBlock, ToolResultBlock, ToolUseBlock -from dream.state.shadow._rewind import ( - is_user_prompt_message, - rewind_transcript, - user_prompt_indices, -) +from dream.engine._messages import ConversationMessage, TextBlock +from dream.state.shadow._rewind import rewind_transcript -def _user(text: str) -> ConversationMessage: - return ConversationMessage(role="user", content=[TextBlock(text=text)]) - - -def _assistant(text: str) -> ConversationMessage: - return ConversationMessage(role="assistant", content=[TextBlock(text=text)]) - - -def _assistant_tools(tool_id: str, name: str) -> ConversationMessage: - return ConversationMessage( - role="assistant", - content=[ToolUseBlock(id=tool_id, name=name, input={"path": "a.py"})], - ) - - -def _tool_results(tool_id: str) -> ConversationMessage: - return ConversationMessage( - role="user", - content=[ToolResultBlock(tool_use_id=tool_id, content="ok")], - ) - - -def test_is_user_prompt_ignores_tool_result_only_messages() -> None: - assert is_user_prompt_message(_user("do the thing")) is True - assert is_user_prompt_message(_tool_results("t1")) is False - assert is_user_prompt_message(_assistant("hi")) is False - - -def test_user_prompt_indices_finds_prompt_boundaries() -> None: - messages = [ - _user("task A"), - _assistant_tools("t1", "write_file"), - _tool_results("t1"), - _assistant("done A"), - _user("task B"), - _assistant("done B"), - ] - assert user_prompt_indices(messages) == (0, 4) +def _message(role: str, text: str) -> ConversationMessage: + return ConversationMessage(role=role, content=[TextBlock(text=text)]) def test_rewind_zero_turns_is_noop() -> None: - messages = [_user("a"), _assistant("b")] - kept, removed = rewind_transcript(messages, turns=0) + messages = [_message("user", "a"), _message("assistant", "b")] + kept, removed = rewind_transcript(messages, prompt_indices=(0,), turns=0) assert kept == messages assert removed == 0 def test_rewind_one_turn_drops_last_prompt_cycle() -> None: messages = [ - _user("task A"), - _assistant("done A"), - _user("task B"), - _assistant_tools("t1", "edit_file"), - _tool_results("t1"), - _assistant("done B"), + _message("user", "task A"), + _message("assistant", "done A"), + _message("user", "task B"), + _message("assistant", "done B"), ] - kept, removed = rewind_transcript(messages, turns=1) + kept, removed = rewind_transcript(messages, prompt_indices=(0, 2), turns=1) assert kept == messages[:2] - assert removed == 4 + assert removed == 2 def test_rewind_two_turns_clears_both_cycles() -> None: messages = [ - _user("A"), - _assistant("a"), - _user("B"), - _assistant("b"), + _message("user", "A"), + _message("assistant", "a"), + _message("user", "B"), + _message("assistant", "b"), ] - kept, removed = rewind_transcript(messages, turns=2) + kept, removed = rewind_transcript(messages, prompt_indices=(0, 2), turns=2) assert kept == [] assert removed == 4 -def test_rewind_more_than_available_clears_all() -> None: - messages = [_user("only"), _assistant("ok")] - kept, removed = rewind_transcript(messages, turns=9) - assert kept == [] - assert removed == 2 +def test_rewind_rejects_unavailable_boundary() -> None: + messages = [_message("user", "compacted")] + try: + rewind_transcript(messages, prompt_indices=(), turns=1) + except ValueError as exc: + assert "unavailable" in str(exc) + else: + raise AssertionError("expected ValueError") def test_rewind_rejects_negative_turns() -> None: try: - rewind_transcript([_user("x")], turns=-1) + rewind_transcript([_message("user", "x")], prompt_indices=(0,), turns=-1) except ValueError as exc: assert "turns" in str(exc) else: From b83d3b3da20569da98f234f884ba8ed6d05fd37e Mon Sep 17 00:00:00 2001 From: divyansh_b Date: Fri, 7 Aug 2026 11:18:35 +0000 Subject: [PATCH 4/6] test: seed explicit rewind boundaries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_session/test_checkpoint_rewind.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_session/test_checkpoint_rewind.py b/tests/test_session/test_checkpoint_rewind.py index 76fe07f0..d452e63d 100644 --- a/tests/test_session/test_checkpoint_rewind.py +++ b/tests/test_session/test_checkpoint_rewind.py @@ -63,6 +63,7 @@ def test_session_restore_checkpoint_rewinds_fs_and_transcript(tmp_path: Path) -> ConversationMessage(role="assistant", content=[TextBlock(text="more done")]), ] ) + session._prompt_indices.extend([0, 2]) result = session.restore_checkpoint(rewind_turns=1) assert result.fs.outcome is RestoreOutcome.RESTORED From 586b335441527137d4a64a219b475a2160ccb729 Mon Sep 17 00:00:00 2001 From: divyansh_b Date: Fri, 7 Aug 2026 12:07:09 +0000 Subject: [PATCH 5/6] Skip checkpoints for oversized worktrees Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- src/dream/_factory.py | 8 +++- src/dream/state/shadow/_manager.py | 34 ++++++++++++++- src/dream/state/shadow/_rewind.py | 2 - src/dream/state/shadow/_types.py | 2 + tests/test_state/test_shadow_checkpoint.py | 51 ++++++++++++++++++++++ 5 files changed, 93 insertions(+), 4 deletions(-) diff --git a/src/dream/_factory.py b/src/dream/_factory.py index 5f517697..30647511 100644 --- a/src/dream/_factory.py +++ b/src/dream/_factory.py @@ -63,6 +63,7 @@ render_skill_catalogue, ) from dream.state.shadow import ( + ShadowCheckpointConfig, ShadowCheckpointHook, ShadowCheckpointManager, ShadowCheckpointStore, @@ -112,6 +113,7 @@ def build_harness( env: Mapping[str, str] | None = None, wake_model: str | None = None, shadow_checkpoints: bool = True, + shadow_checkpoint_config: ShadowCheckpointConfig | None = None, ) -> Harness: """Build a Harness whose engine factory produces a real, tool-wired engine. @@ -131,7 +133,10 @@ def build_harness( ``shadow_checkpoints`` (default True) registers Hermes-style pre-mutate filesystem snapshots and exposes :meth:`~dream.session.Session.restore_checkpoint` - for operator rewind (FS + transcript). + for operator rewind (FS + transcript). Worktrees over the checkpoint + manager's ``max_files`` threshold (10,000 by default) are skipped to keep + per-turn overhead bounded; pass a custom ``ShadowCheckpointConfig`` to + override that threshold. Workspace memory (the durable per-project record store under :func:`~dream.memory.project_memory_dir`) is wired by default: its @@ -238,6 +243,7 @@ def build_harness( if shadow_checkpoints: checkpoint_manager = ShadowCheckpointManager( store=ShadowCheckpointStore(base_dir=paths.checkpoints_dir), + config=shadow_checkpoint_config, ) config = HarnessConfig( working_dir=working_dir, diff --git a/src/dream/state/shadow/_manager.py b/src/dream/state/shadow/_manager.py index 2be6e92b..a3f29198 100644 --- a/src/dream/state/shadow/_manager.py +++ b/src/dream/state/shadow/_manager.py @@ -3,6 +3,7 @@ from __future__ import annotations import contextlib +import os import shutil from collections import OrderedDict from collections.abc import Sequence @@ -21,6 +22,7 @@ RestoreResult, ShadowCheckpointConfig, ) +from dream.utils.git import run_git # Safety snap outcomes that still allow restore to proceed. _SAFETY_OK = frozenset( @@ -29,6 +31,7 @@ CheckpointOutcome.NO_CHANGES, } ) +_MAX_SESSION_STATES = 64 class ShadowCheckpointManager: @@ -44,6 +47,7 @@ def __init__( self._config = config or ShadowCheckpointConfig() self._checkpointed_dirs: OrderedDict[str | None, set[Path]] = OrderedDict() self._git_available: bool | None = None + self._worktree_size_ok: dict[Path, bool] = {} @property def config(self) -> ShadowCheckpointConfig: @@ -77,9 +81,16 @@ def ensure( if abs_dir == Path("/").resolve() or abs_dir == Path.home().resolve(): return EnsureResult(outcome=CheckpointOutcome.DIRECTORY_TOO_BROAD) + size_ok = self._worktree_size_ok.get(abs_dir) + if size_ok is None: + size_ok = self._probe_worktree_size(abs_dir) + self._worktree_size_ok[abs_dir] = size_ok + if not size_ok: + return EnsureResult(outcome=CheckpointOutcome.DIRECTORY_TOO_LARGE) + checkpointed_dirs = self._checkpointed_dirs.setdefault(session_id, set()) self._checkpointed_dirs.move_to_end(session_id) - while len(self._checkpointed_dirs) > 64: + while len(self._checkpointed_dirs) > _MAX_SESSION_STATES: self._checkpointed_dirs.popitem(last=False) if abs_dir in checkpointed_dirs: return EnsureResult(outcome=CheckpointOutcome.ALREADY_THIS_TURN) @@ -95,6 +106,27 @@ def ensure( checkpointed_dirs.add(abs_dir) return result + def _probe_worktree_size(self, working_dir: Path) -> bool: + """Return whether the worktree is cheap enough to snapshot.""" + limit = self._config.max_files + if limit < 1: + return True + rc, stdout, _err = run_git( + ["ls-files", "-co", "--exclude-standard", "-z"], + cwd=working_dir, + timeout=self._config.timeout_seconds, + ) + if rc == 0: + return len([path for path in stdout.split("\0") if path]) <= limit + + count = 0 + for _root, dirs, files in os.walk(working_dir): + dirs[:] = [name for name in dirs if name != ".git"] + count += len(files) + if count > limit: + return False + return True + def list_for(self, working_dir: Path) -> list[CheckpointSnapshot]: """List retained checkpoints for ``working_dir`` (newest first).""" abs_dir = working_dir.resolve() diff --git a/src/dream/state/shadow/_rewind.py b/src/dream/state/shadow/_rewind.py index 5bde7dec..54153f3d 100644 --- a/src/dream/state/shadow/_rewind.py +++ b/src/dream/state/shadow/_rewind.py @@ -29,8 +29,6 @@ def rewind_transcript( if turns > len(prompt_indices): raise ValueError("requested rewind boundary is unavailable") - if not prompt_indices: - return list(messages), 0 drop_from = prompt_indices[-turns] kept = list(messages[:drop_from]) diff --git a/src/dream/state/shadow/_types.py b/src/dream/state/shadow/_types.py index 300c1b6e..b4c5b5d0 100644 --- a/src/dream/state/shadow/_types.py +++ b/src/dream/state/shadow/_types.py @@ -39,6 +39,7 @@ class CheckpointOutcome(Enum): ALREADY_THIS_TURN = "already_this_turn" NO_CHANGES = "no_changes" DIRECTORY_TOO_BROAD = "directory_too_broad" + DIRECTORY_TOO_LARGE = "directory_too_large" GIT_UNAVAILABLE = "git_unavailable" FAILED = "failed" @@ -59,6 +60,7 @@ class ShadowCheckpointConfig: enabled: bool = True max_snapshots: int = 20 timeout_seconds: float = 30.0 + max_files: int = 10_000 @dataclass(frozen=True, slots=True) diff --git a/tests/test_state/test_shadow_checkpoint.py b/tests/test_state/test_shadow_checkpoint.py index 30b0c1ca..ad84aa85 100644 --- a/tests/test_state/test_shadow_checkpoint.py +++ b/tests/test_state/test_shadow_checkpoint.py @@ -115,6 +115,57 @@ def test_no_changes_skips(mgr: ShadowCheckpointManager, work_dir: Path) -> None: assert skipped.outcome is CheckpointOutcome.NO_CHANGES +def test_small_worktree_checkpoints_normally(tmp_path: Path) -> None: + work_dir = tmp_path / "small" + work_dir.mkdir() + (work_dir / "README.md").write_text("hello\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "checkpoints"), + config=ShadowCheckpointConfig(max_files=1), + ) + result = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert result.outcome is CheckpointOutcome.TAKEN + + +def test_large_worktree_skips_without_creating_store(tmp_path: Path) -> None: + work_dir = tmp_path / "large" + work_dir.mkdir() + for index in range(3): + (work_dir / f"file-{index}.txt").write_text("x\n", encoding="utf-8") + store_root = tmp_path / "checkpoints" + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=store_root), + config=ShadowCheckpointConfig(max_files=2), + ) + result = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert result.outcome is CheckpointOutcome.DIRECTORY_TOO_LARGE + assert not store_root.exists() + + +def test_worktree_size_probe_runs_once_per_directory( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + work_dir = tmp_path / "large" + work_dir.mkdir() + (work_dir / "file.txt").write_text("x\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "checkpoints"), + config=ShadowCheckpointConfig(max_files=0), + ) + calls = 0 + original = mgr._probe_worktree_size + + def probe_once(path: Path) -> bool: + nonlocal calls + calls += 1 + return original(path) + + monkeypatch.setattr(mgr, "_probe_worktree_size", probe_once) + mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_BASH) + assert calls == 1 + + def test_skips_home_and_root(mgr: ShadowCheckpointManager) -> None: assert ( mgr.ensure(Path("/"), reason=CheckpointReason.BEFORE_BASH).outcome From 878ef1bf5712ea108e3b84c346ae361756a0d3e3 Mon Sep 17 00:00:00 2001 From: divo12 Date: Sun, 16 Aug 2026 23:29:53 +0530 Subject: [PATCH 6/6] feat(state): Hermes-style human rewind for shadow checkpoints Restack operator rewind onto current main so Session.list_checkpoints / restore_checkpoint restore the worktree and truncate Session-owned prompt turns, with build_harness auto-wiring the shared shadow store. Co-authored-by: Cursor --- CHANGELOG.md | 6 + src/dream/_factory.py | 27 +++++ src/dream/config/paths.py | 5 + src/dream/engine/_engine.py | 6 + src/dream/harness.py | 3 + src/dream/session.py | 84 +++++++++++++ src/dream/state/__init__.py | 8 ++ src/dream/state/shadow/__init__.py | 4 + src/dream/state/shadow/_hook.py | 13 +- src/dream/state/shadow/_manager.py | 106 ++++++++++++++++- src/dream/state/shadow/_rewind.py | 41 +++++++ src/dream/state/shadow/_types.py | 14 +++ tests/test_config/test_paths.py | 2 + tests/test_session.py | 1 + tests/test_session/test_checkpoint_rewind.py | 118 +++++++++++++++++++ tests/test_state/test_combined_restore.py | 102 ++++++++++++++++ tests/test_state/test_shadow_checkpoint.py | 81 ++++++++++++- tests/test_state/test_transcript_rewind.py | 60 ++++++++++ 18 files changed, 672 insertions(+), 9 deletions(-) create mode 100644 src/dream/state/shadow/_rewind.py create mode 100644 tests/test_session/test_checkpoint_rewind.py create mode 100644 tests/test_state/test_combined_restore.py create mode 100644 tests/test_state/test_transcript_rewind.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 867fc592..48b786bb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] ### Added +- Hermes-style human rewind: `Session.list_checkpoints` / `Session.restore_checkpoint` + restore the worktree and truncate matching transcript turns. Turn boundaries + come from the session (prompt-submit indices), not message inspection. + `build_harness` auto-wires `ShadowCheckpointHook` and a shared store under + `DreamPaths.checkpoints_dir` (default-on; oversized trees skip via + `ShadowCheckpointConfig.max_files`). - OpenTelemetry is **default-on**: core deps ship the OTLP SDK; sessions fan JSONL traces to OTel (`CompositeTracer`). Endpoint defaults to `http://localhost:4318`; override with `OTEL_EXPORTER_OTLP_ENDPOINT`. Opt out diff --git a/src/dream/_factory.py b/src/dream/_factory.py index 233e204a..d38665af 100644 --- a/src/dream/_factory.py +++ b/src/dream/_factory.py @@ -72,6 +72,12 @@ build_session_skill_registry, render_skill_catalogue, ) +from dream.state.shadow import ( + ShadowCheckpointConfig, + ShadowCheckpointHook, + ShadowCheckpointManager, + ShadowCheckpointStore, +) from dream.subagents._async_delegation import AsyncDelegationManager from dream.subagents._catalogue import SubagentCatalogue from dream.subagents._declaration import SubagentSet @@ -181,6 +187,8 @@ def build_harness( env: Mapping[str, str] | None = None, wake_model: str | None = None, verify_on_stop: bool = True, + shadow_checkpoints: bool = True, + shadow_checkpoint_config: ShadowCheckpointConfig | None = None, ) -> Harness: """Build a Harness whose engine factory produces a real, tool-wired engine. @@ -233,6 +241,13 @@ def build_harness( mutating file tools without a subsequent evidence tool (read/grep/glob) nudge another turn before seal (capped by ``max_verify_nudges``). + ``shadow_checkpoints`` (default True) registers Hermes-style pre-mutate + filesystem snapshots and exposes :meth:`~dream.session.Session.restore_checkpoint` + for operator rewind (FS + transcript). Worktrees over the checkpoint + manager's ``max_files`` threshold (10,000 by default) are skipped to keep + per-turn overhead bounded; pass a custom ``ShadowCheckpointConfig`` to + override that threshold. + ``env`` is consulted only for host resolution — ``DREAM_HOME`` path overrides and shell detection for the runtime-info prompt block — and defaults to ``os.environ``. Credentials never come from it. @@ -346,12 +361,19 @@ def build_harness( # so a scheduler tick loop knows where to poll, and `paths` carries the # env-resolved roots. del wake_model # ponytail: compat no-op — the wake runtime is gone + checkpoint_manager: ShadowCheckpointManager | None = None + if shadow_checkpoints: + checkpoint_manager = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=paths.checkpoints_dir), + config=shadow_checkpoint_config, + ) config = HarnessConfig( working_dir=working_dir, task_manager=task_manager, delegations=AsyncDelegationManager(), cron_registry_path=task_context.cron_registry_path, paths=paths, + checkpoint_manager=checkpoint_manager, # MCP connect + plugin import are async/IO, so they hang off the # async-open chokepoint (``Harness._ensure_open``) rather than running # in this sync factory. ``None`` when both surfaces are disabled so the @@ -369,6 +391,10 @@ def build_harness( ), ) harness = Harness(config) + if checkpoint_manager is not None: + harness.register_hook( + ShadowCheckpointHook(manager=checkpoint_manager, working_dir=working_dir) + ) if verify_on_stop: harness.register_hook(VerifyOnStopHook()) @@ -877,4 +903,5 @@ def _build_session_engine( initial_context=render_runtime_context(runtime_info), delegations=harness.config.delegations, prompt_surfaces=prompt_surfaces, + checkpoint_manager=harness.config.checkpoint_manager, ) diff --git a/src/dream/config/paths.py b/src/dream/config/paths.py index 2f25c7bf..de33d177 100644 --- a/src/dream/config/paths.py +++ b/src/dream/config/paths.py @@ -207,6 +207,11 @@ def memory_dir(self) -> Path: def skills_dir(self) -> Path: return self.home / "skills" + @property + def checkpoints_dir(self) -> Path: + """Shared shadow-git store root under ``$DREAM_HOME/checkpoints``.""" + return self.home / "checkpoints" + # --- the one explicit side effect --- def ensure(self) -> DreamPaths: diff --git a/src/dream/engine/_engine.py b/src/dream/engine/_engine.py index 682b20c3..ff958e8e 100644 --- a/src/dream/engine/_engine.py +++ b/src/dream/engine/_engine.py @@ -41,6 +41,7 @@ from dream.tools._registry import ToolRegistry if TYPE_CHECKING: + from dream.state.shadow import ShadowCheckpointManager from dream.subagents._async_delegation import AsyncDelegationManager @@ -88,6 +89,9 @@ class QueryEngine: # Role name from session metadata (planner/generator/evaluator) for STOP hooks. role: str | None = None delegations: AsyncDelegationManager | None = None + # Hermes shadow checkpoints — shared across sessions on the harness; enables + # :meth:`Session.list_checkpoints` / :meth:`Session.restore_checkpoint`. + checkpoint_manager: ShadowCheckpointManager | None = None # Typed request surfaces for `/context` (FailoverStreamer hides streamer attrs). prompt_surfaces: PromptSurfaces | None = None @@ -153,6 +157,7 @@ def build_query_engine( orientation: OrientationConfig | None = None, delegations: AsyncDelegationManager | None = None, prompt_surfaces: PromptSurfaces | None = None, + checkpoint_manager: ShadowCheckpointManager | None = None, ) -> QueryEngine: """Wrap a ``ToolRegistry`` in the canonical dispatcher and bind a streamer. @@ -195,6 +200,7 @@ def build_query_engine( role=role_raw if isinstance(role_raw, str) else None, delegations=delegations, prompt_surfaces=prompt_surfaces, + checkpoint_manager=checkpoint_manager, ) diff --git a/src/dream/harness.py b/src/dream/harness.py index d5c428d1..f5f29517 100644 --- a/src/dream/harness.py +++ b/src/dream/harness.py @@ -44,6 +44,7 @@ RunTaskResult, SprintGoalProvider, ) + from dream.state.shadow import ShadowCheckpointManager from dream.subagents._async_delegation import AsyncDelegationManager from dream.tasks import BackgroundTaskManager @@ -89,6 +90,8 @@ class HarnessConfig: # than re-resolving and risking divergence. paths: DreamPaths | None = None session_store: FileSessionStore | None = None + # Shared Hermes-style shadow checkpoint manager (FS snaps + human rewind). + checkpoint_manager: ShadowCheckpointManager | None = None extra: dict[str, Any] = field(default_factory=dict) _engine_factory: EngineFactory | None = None # Async setup run once before the first session — MCP connect + plugin diff --git a/src/dream/session.py b/src/dream/session.py index 5c36e2e9..a70f1a09 100644 --- a/src/dream/session.py +++ b/src/dream/session.py @@ -72,6 +72,11 @@ if TYPE_CHECKING: from dream.context import ContextBreakdown from dream.engine._engine import QueryEngine + from dream.state.shadow import ( + CheckpointSnapshot, + CombinedRestoreResult, + ShadowCheckpointManager, + ) def _zero_cost() -> SessionCostSnapshot: @@ -165,6 +170,7 @@ def __init__( # path is missing. Resumed sessions receive the exact loaded revision. self._snapshot_revision = _snapshot_revision self._transcript: list[ConversationMessage] = [] + self._prompt_indices: list[int] = [] self._cancel_event: asyncio.Event | None = None self._closed = False # Single-flight guard: ``Session`` keeps per-call cancel state on the @@ -231,6 +237,80 @@ def transcript(self) -> list[ConversationMessage]: """ return self._transcript + @property + def checkpoint_manager(self) -> ShadowCheckpointManager | None: + """Shadow FS checkpoint manager when the bound engine was built with one.""" + engine = self._engine + if engine is None: + return None + return engine.checkpoint_manager + + def list_checkpoints(self) -> list[CheckpointSnapshot]: + """List shadow checkpoints for this session's working directory (newest first).""" + engine = self._engine + manager = self.checkpoint_manager + if engine is None or manager is None: + return [] + return manager.list_for(engine.working_dir) + + def restore_checkpoint( + self, + commit_sha: str | None = None, + *, + rewind_turns: int = 1, + ) -> CombinedRestoreResult: + """Hermes-style human rewind: restore FS and truncate the transcript. + + ``commit_sha=None`` selects the newest checkpoint. Refuses while a + ``send`` is in flight so restore cannot race the live turn loop. + Turn boundaries come from ``Session`` (prompt submit indices), not + from inspecting message roles — engine-generated user messages are + not human turns. + """ + from dream.state.shadow import CombinedRestoreResult, RestoreOutcome, RestoreResult + + if self._active: + raise RuntimeError("cannot restore checkpoint while a send is in flight") + engine = self._engine + manager = self.checkpoint_manager + if engine is None or manager is None: + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.DISABLED, + detail="no checkpoint manager bound on this session", + ), + messages=tuple(self._transcript), + transcript_removed=0, + ) + + sha = commit_sha + if sha is None: + listed = manager.list_for(engine.working_dir) + if not listed: + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.NOT_FOUND, + detail="no checkpoints for working directory", + ), + messages=tuple(self._transcript), + transcript_removed=0, + ) + sha = listed[0].commit_sha + + result = manager.restore_and_rewind( + engine.working_dir, + commit_sha=sha, + messages=self._transcript, + prompt_indices=self._prompt_indices, + rewind_turns=rewind_turns, + ) + if result.fs.outcome is RestoreOutcome.RESTORED: + self._transcript[:] = list(result.messages) + self._prompt_indices = [ + index for index in self._prompt_indices if index < len(self._transcript) + ] + return result + def _current_cost(self) -> SessionCostSnapshot: return cost_snapshot_from_fields( SessionCostFields( @@ -324,6 +404,7 @@ def restore_from_snapshot(self, snapshot: SessionSnapshot) -> None: """Replace transcript and cost counters from a saved snapshot.""" restored = messages_from_records(snapshot.messages) self._transcript[:] = sanitize_conversation_messages(restored) + self._prompt_indices.clear() self.cost.input_tokens = snapshot.cost.input_tokens self.cost.output_tokens = snapshot.cost.output_tokens self.cost.cache_read_tokens = snapshot.cost.cache_read_tokens @@ -359,6 +440,7 @@ async def send(self, prompt: str) -> AsyncIterator[Event]: resume = list(self._transcript) if self._transcript else None user_msg = ConversationMessage(role="user", content=[TextBlock(text=prompt)]) self._transcript.append(user_msg) + self._prompt_indices.append(len(self._transcript) - 1) config = self._engine.make_session_config() if self._has_sent: @@ -537,6 +619,7 @@ def _apply_compaction(self) -> None: carryover = engine.carryover_metadata if carryover is not None and carryover.last_compacted_transcript is not None: self._transcript[:] = list(carryover.last_compacted_transcript) + self._prompt_indices.clear() carryover.last_compacted_transcript = None return compactor = engine.compactor @@ -558,6 +641,7 @@ def _apply_compaction(self) -> None: ) if result is not None: self._transcript[:] = new_transcript + self._prompt_indices.clear() async def cancel(self) -> None: """Cancel the in-flight ``send``, if any. diff --git a/src/dream/state/__init__.py b/src/dream/state/__init__.py index 8751028e..4b03d856 100644 --- a/src/dream/state/__init__.py +++ b/src/dream/state/__init__.py @@ -5,21 +5,29 @@ from dream.state.shadow import ( CheckpointOutcome, CheckpointReason, + CheckpointSnapshot, + CombinedRestoreResult, MutatingToolName, RestoreOutcome, + RestoreResult, ShadowCheckpointConfig, ShadowCheckpointHook, ShadowCheckpointManager, ShadowCheckpointStore, + rewind_transcript, ) __all__ = [ "CheckpointOutcome", "CheckpointReason", + "CheckpointSnapshot", + "CombinedRestoreResult", "MutatingToolName", "RestoreOutcome", + "RestoreResult", "ShadowCheckpointConfig", "ShadowCheckpointHook", "ShadowCheckpointManager", "ShadowCheckpointStore", + "rewind_transcript", ] diff --git a/src/dream/state/shadow/__init__.py b/src/dream/state/shadow/__init__.py index 4b519d5c..ba3f4809 100644 --- a/src/dream/state/shadow/__init__.py +++ b/src/dream/state/shadow/__init__.py @@ -4,11 +4,13 @@ from dream.state.shadow._hook import ShadowCheckpointHook from dream.state.shadow._manager import ShadowCheckpointManager +from dream.state.shadow._rewind import rewind_transcript from dream.state.shadow._store import ShadowCheckpointStore from dream.state.shadow._types import ( CheckpointOutcome, CheckpointReason, CheckpointSnapshot, + CombinedRestoreResult, EnsureResult, MutatingToolName, RestoreOutcome, @@ -20,6 +22,7 @@ "CheckpointOutcome", "CheckpointReason", "CheckpointSnapshot", + "CombinedRestoreResult", "EnsureResult", "MutatingToolName", "RestoreOutcome", @@ -28,4 +31,5 @@ "ShadowCheckpointHook", "ShadowCheckpointManager", "ShadowCheckpointStore", + "rewind_transcript", ] diff --git a/src/dream/state/shadow/_hook.py b/src/dream/state/shadow/_hook.py index 2c49e442..a78f9cc2 100644 --- a/src/dream/state/shadow/_hook.py +++ b/src/dream/state/shadow/_hook.py @@ -32,9 +32,14 @@ def __init__( self._manager = manager self._working_dir = working_dir + @staticmethod + def _session_id(payload: Mapping[str, Any]) -> str | None: + raw_session_id = payload.get("session_id") + return str(raw_session_id) if raw_session_id is not None else None + async def __call__(self, event: HookEvent, payload: Mapping[str, Any]) -> HookResult: if event is HookEvent.USER_PROMPT_SUBMIT: - self._manager.begin_turn() + self._manager.begin_turn(self._session_id(payload)) return HookResult() if event is not HookEvent.PRE_TOOL_USE: @@ -45,7 +50,11 @@ async def __call__(self, event: HookEvent, payload: Mapping[str, Any]) -> HookRe if reason is None: return HookResult() - self._manager.ensure(self._working_dir, reason=reason) + self._manager.ensure( + self._working_dir, + reason=reason, + session_id=self._session_id(payload), + ) return HookResult() diff --git a/src/dream/state/shadow/_manager.py b/src/dream/state/shadow/_manager.py index 58100a0a..06844e03 100644 --- a/src/dream/state/shadow/_manager.py +++ b/src/dream/state/shadow/_manager.py @@ -3,19 +3,26 @@ from __future__ import annotations import contextlib +import os import shutil +from collections import OrderedDict +from collections.abc import Sequence from pathlib import Path +from dream.engine._messages import ConversationMessage +from dream.state.shadow._rewind import rewind_transcript from dream.state.shadow._store import ShadowCheckpointStore from dream.state.shadow._types import ( CheckpointOutcome, CheckpointReason, CheckpointSnapshot, + CombinedRestoreResult, EnsureResult, RestoreOutcome, RestoreResult, ShadowCheckpointConfig, ) +from dream.utils.git import run_git # Safety snap outcomes that still allow restore to proceed. _SAFETY_OK = frozenset( @@ -24,6 +31,7 @@ CheckpointOutcome.NO_CHANGES, } ) +_MAX_SESSION_STATES = 64 class ShadowCheckpointManager: @@ -37,18 +45,25 @@ def __init__( ) -> None: self._store = store self._config = config or ShadowCheckpointConfig() - self._checkpointed_dirs: set[Path] = set() + self._checkpointed_dirs: OrderedDict[str | None, set[Path]] = OrderedDict() self._git_available: bool | None = None + self._worktree_size_ok: dict[Path, bool] = {} @property def config(self) -> ShadowCheckpointConfig: return self._config - def begin_turn(self) -> None: + def begin_turn(self, session_id: str | None = None) -> None: """Reset per-turn dedup (call on each USER_PROMPT_SUBMIT / agent turn).""" - self._checkpointed_dirs.clear() + self._checkpointed_dirs.pop(session_id, None) - def ensure(self, working_dir: Path, *, reason: CheckpointReason) -> EnsureResult: + def ensure( + self, + working_dir: Path, + *, + reason: CheckpointReason, + session_id: str | None = None, + ) -> EnsureResult: """Take a checkpoint if enabled and not already done this turn.""" if not self._config.enabled: return EnsureResult(outcome=CheckpointOutcome.DISABLED) @@ -66,7 +81,18 @@ def ensure(self, working_dir: Path, *, reason: CheckpointReason) -> EnsureResult if abs_dir == Path("/").resolve() or abs_dir == Path.home().resolve(): return EnsureResult(outcome=CheckpointOutcome.DIRECTORY_TOO_BROAD) - if abs_dir in self._checkpointed_dirs: + size_ok = self._worktree_size_ok.get(abs_dir) + if size_ok is None: + size_ok = self._probe_worktree_size(abs_dir) + self._worktree_size_ok[abs_dir] = size_ok + if not size_ok: + return EnsureResult(outcome=CheckpointOutcome.DIRECTORY_TOO_LARGE) + + checkpointed_dirs = self._checkpointed_dirs.setdefault(session_id, set()) + self._checkpointed_dirs.move_to_end(session_id) + while len(self._checkpointed_dirs) > _MAX_SESSION_STATES: + self._checkpointed_dirs.popitem(last=False) + if abs_dir in checkpointed_dirs: return EnsureResult(outcome=CheckpointOutcome.ALREADY_THIS_TURN) try: @@ -77,9 +103,30 @@ def ensure(self, working_dir: Path, *, reason: CheckpointReason) -> EnsureResult # Only suppress retries after a conclusive snap (or no-op). Transient # FAILED must leave the turn open so a later mutate can try again. if result.outcome is not CheckpointOutcome.FAILED: - self._checkpointed_dirs.add(abs_dir) + checkpointed_dirs.add(abs_dir) return result + def _probe_worktree_size(self, working_dir: Path) -> bool: + """Return whether the worktree is cheap enough to snapshot.""" + limit = self._config.max_files + if limit < 1: + return True + rc, stdout, _err = run_git( + ["ls-files", "-co", "--exclude-standard", "-z"], + cwd=working_dir, + timeout=self._config.timeout_seconds, + ) + if rc == 0: + return len([path for path in stdout.split("\0") if path]) <= limit + + count = 0 + for _root, dirs, files in os.walk(working_dir): + dirs[:] = [name for name in dirs if name != ".git"] + count += len(files) + if count > limit: + return False + return True + def list_for(self, working_dir: Path) -> list[CheckpointSnapshot]: """List retained checkpoints for ``working_dir`` (newest first).""" abs_dir = working_dir.resolve() @@ -158,6 +205,53 @@ def restore(self, working_dir: Path, *, commit_sha: str) -> RestoreResult: return RestoreResult(outcome=RestoreOutcome.RESTORED, restored_to=commit_sha[:8]) + def restore_and_rewind( + self, + working_dir: Path, + *, + commit_sha: str, + messages: Sequence[ConversationMessage], + prompt_indices: Sequence[int], + rewind_turns: int = 1, + ) -> CombinedRestoreResult: + """Restore the worktree and truncate the conversation (Hermes ``/rollback``). + + Transcript rewind runs only after a successful filesystem restore so a + failed snap never desyncs chat from disk. ``rewind_turns=0`` keeps the + transcript unchanged (FS-only restore). + """ + if rewind_turns < 0: + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.FAILED, + detail=f"turns must be >= 0; got {rewind_turns}", + ), + messages=tuple(messages), + transcript_removed=0, + ) + if rewind_turns > 0 and rewind_turns > len(prompt_indices): + return CombinedRestoreResult( + fs=RestoreResult( + outcome=RestoreOutcome.FAILED, + detail=f"requested rewind boundary is unavailable for {rewind_turns} turn(s)", + ), + messages=tuple(messages), + transcript_removed=0, + ) + fs = self.restore(working_dir, commit_sha=commit_sha) + if fs.outcome is not RestoreOutcome.RESTORED: + return CombinedRestoreResult(fs=fs, messages=tuple(messages), transcript_removed=0) + kept, removed = rewind_transcript( + messages, + prompt_indices=prompt_indices, + turns=rewind_turns, + ) + return CombinedRestoreResult( + fs=fs, + messages=tuple(kept), + transcript_removed=removed, + ) + def _take(self, working_dir: Path, reason: CheckpointReason) -> EnsureResult: err = self._store.ensure_initialized(working_dir) if err: diff --git a/src/dream/state/shadow/_rewind.py b/src/dream/state/shadow/_rewind.py new file mode 100644 index 00000000..54153f3d --- /dev/null +++ b/src/dream/state/shadow/_rewind.py @@ -0,0 +1,41 @@ +"""Transcript rewind helpers — Hermes ``/rollback`` conversation half. + +Filesystem restore lives on :class:`ShadowCheckpointManager`. These pure +functions truncate the in-memory conversation so FS and chat stay aligned +after an operator rewind. +""" + +from __future__ import annotations + +from collections.abc import Sequence + +from dream.engine._messages import ConversationMessage + + +def rewind_transcript( + messages: Sequence[ConversationMessage], + *, + prompt_indices: Sequence[int], + turns: int, +) -> tuple[list[ConversationMessage], int]: + """Drop the last ``turns`` prompt cycles using Session-owned boundaries. + + Returns ``(kept_messages, removed_count)``. + """ + if turns < 0: + raise ValueError(f"turns must be >= 0; got {turns}") + if turns == 0: + return list(messages), 0 + + if turns > len(prompt_indices): + raise ValueError("requested rewind boundary is unavailable") + + drop_from = prompt_indices[-turns] + kept = list(messages[:drop_from]) + removed = len(messages) - len(kept) + return kept, removed + + +__all__ = [ + "rewind_transcript", +] diff --git a/src/dream/state/shadow/_types.py b/src/dream/state/shadow/_types.py index cf4c9f55..92eb5771 100644 --- a/src/dream/state/shadow/_types.py +++ b/src/dream/state/shadow/_types.py @@ -6,6 +6,8 @@ from enum import Enum, StrEnum from pathlib import Path +from dream.engine._messages import ConversationMessage + class MutatingToolName(StrEnum): """Built-in tools that may mutate the working tree.""" @@ -37,6 +39,7 @@ class CheckpointOutcome(Enum): ALREADY_THIS_TURN = "already_this_turn" NO_CHANGES = "no_changes" DIRECTORY_TOO_BROAD = "directory_too_broad" + DIRECTORY_TOO_LARGE = "directory_too_large" GIT_UNAVAILABLE = "git_unavailable" FAILED = "failed" @@ -57,6 +60,7 @@ class ShadowCheckpointConfig: enabled: bool = True max_snapshots: int = 20 timeout_seconds: float = 30.0 + max_files: int = 10_000 @dataclass(frozen=True, slots=True) @@ -87,6 +91,15 @@ class RestoreResult: detail: str = "" +@dataclass(frozen=True, slots=True) +class CombinedRestoreResult: + """Filesystem restore plus optional transcript rewind (Hermes ``/rollback``).""" + + fs: RestoreResult + messages: tuple[ConversationMessage, ...] = () + transcript_removed: int = 0 + + _TOOL_TO_REASON: dict[MutatingToolName, CheckpointReason] = { MutatingToolName.WRITE_FILE: CheckpointReason.BEFORE_WRITE_FILE, MutatingToolName.APPLY_PATCH: CheckpointReason.BEFORE_APPLY_PATCH, @@ -109,6 +122,7 @@ def reason_for_tool(tool_name: str) -> CheckpointReason | None: "CheckpointOutcome", "CheckpointReason", "CheckpointSnapshot", + "CombinedRestoreResult", "EnsureResult", "MutatingToolName", "RestoreOutcome", diff --git a/tests/test_config/test_paths.py b/tests/test_config/test_paths.py index 2d33b33a..fc728288 100644 --- a/tests/test_config/test_paths.py +++ b/tests/test_config/test_paths.py @@ -102,6 +102,7 @@ def test_home_side_paths(tmp_path: Path) -> None: assert p.tasks_dir == p.home / "data" / "tasks" assert p.memory_dir == p.home / "memory" assert p.skills_dir == p.home / "skills" + assert p.checkpoints_dir == p.home / "checkpoints" def test_reading_properties_creates_nothing(tmp_path: Path) -> None: @@ -121,6 +122,7 @@ def test_reading_properties_creates_nothing(tmp_path: Path) -> None: p.tasks_dir, p.memory_dir, p.skills_dir, + p.checkpoints_dir, p.worktree("T1"), p.sidecar("T1"), ) diff --git a/tests/test_session.py b/tests/test_session.py index d22acecd..bf9046b5 100644 --- a/tests/test_session.py +++ b/tests/test_session.py @@ -755,3 +755,4 @@ def test_translate_compaction_applies_compacted_shape_to_transcript() -> None: # now the compacted shape, not the stale full history. assert post_blobs < pre_blobs assert cleared > 0 + assert session._prompt_indices == [] diff --git a/tests/test_session/test_checkpoint_rewind.py b/tests/test_session/test_checkpoint_rewind.py new file mode 100644 index 00000000..3b3136f3 --- /dev/null +++ b/tests/test_session/test_checkpoint_rewind.py @@ -0,0 +1,118 @@ +"""Eval: Session.list_checkpoints / restore_checkpoint (Hermes human rewind).""" + +from __future__ import annotations + +from pathlib import Path +from typing import Never + +import pytest + +from dream.engine._engine import QueryEngine +from dream.engine._messages import ConversationMessage, TextBlock +from dream.session import Session +from dream.state.shadow import ( + CheckpointReason, + RestoreOutcome, + ShadowCheckpointConfig, + ShadowCheckpointManager, + ShadowCheckpointStore, +) + + +class _NoopStreamer: + async def stream_turn(self, *_args: object, **_kwargs: object) -> Never: + raise RuntimeError("streamer unused") + yield # pragma: no cover — makes this an async generator type + + +class _NoopDispatcher: + async def dispatch(self, *_args: object, **_kwargs: object) -> tuple[str, bool]: + return ("", False) + + +def _engine(working_dir: Path, manager: ShadowCheckpointManager | None) -> QueryEngine: + return QueryEngine( + streamer=_NoopStreamer(), # type: ignore[arg-type] + dispatcher=_NoopDispatcher(), # type: ignore[arg-type] + session_id="s1", + working_dir=working_dir, + checkpoint_manager=manager, + ) + + +def test_session_restore_checkpoint_rewinds_fs_and_transcript(tmp_path: Path) -> None: + work = tmp_path / "proj" + work.mkdir() + (work / "f.txt").write_text("clean\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "ck"), + config=ShadowCheckpointConfig(enabled=True), + ) + snap = mgr.ensure(work, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert snap.snapshot is not None + + (work / "f.txt").write_text("dirty\n", encoding="utf-8") + session = Session(id="s1", _engine=_engine(work, mgr)) + session.transcript.extend( + [ + ConversationMessage(role="user", content=[TextBlock(text="edit it")]), + ConversationMessage(role="assistant", content=[TextBlock(text="done")]), + ConversationMessage(role="user", content=[TextBlock(text="more")]), + ConversationMessage(role="assistant", content=[TextBlock(text="more done")]), + ] + ) + session._prompt_indices.extend([0, 2]) + + result = session.restore_checkpoint(rewind_turns=1) + assert result.fs.outcome is RestoreOutcome.RESTORED + assert (work / "f.txt").read_text(encoding="utf-8") == "clean\n" + assert len(session.transcript) == 2 + assert session.transcript[0].text == "edit it" + assert result.transcript_removed == 2 + + +def test_session_list_checkpoints_empty_without_manager(tmp_path: Path) -> None: + session = Session(id="s1", _engine=_engine(tmp_path, None)) + assert session.list_checkpoints() == [] + result = session.restore_checkpoint() + assert result.fs.outcome is RestoreOutcome.DISABLED + + +def test_session_restore_refuses_during_active_send(tmp_path: Path) -> None: + work = tmp_path / "proj" + work.mkdir() + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "ck"), + config=ShadowCheckpointConfig(enabled=True), + ) + session = Session(id="s1", _engine=_engine(work, mgr)) + session._active = True + with pytest.raises(RuntimeError, match="in flight"): + session.restore_checkpoint() + + +def test_session_restore_unavailable_after_compaction(tmp_path: Path) -> None: + work = tmp_path / "proj" + work.mkdir() + (work / "f.txt").write_text("clean\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "ck"), + config=ShadowCheckpointConfig(enabled=True), + ) + snap = mgr.ensure(work, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert snap.snapshot is not None + (work / "f.txt").write_text("dirty\n", encoding="utf-8") + + session = Session(id="s1", _engine=_engine(work, mgr)) + session.transcript.extend( + [ + ConversationMessage(role="user", content=[TextBlock(text="compacted")]), + ] + ) + session._prompt_indices.clear() + + result = session.restore_checkpoint(rewind_turns=1) + assert result.fs.outcome is RestoreOutcome.FAILED + assert "unavailable" in result.fs.detail + assert (work / "f.txt").read_text(encoding="utf-8") == "dirty\n" + assert len(session.transcript) == 1 diff --git a/tests/test_state/test_combined_restore.py b/tests/test_state/test_combined_restore.py new file mode 100644 index 00000000..a7aafb8e --- /dev/null +++ b/tests/test_state/test_combined_restore.py @@ -0,0 +1,102 @@ +"""Eval: combined FS + transcript restore (Hermes /rollback).""" + +from __future__ import annotations + +from pathlib import Path + +import pytest + +from dream.engine._messages import ConversationMessage, TextBlock +from dream.state.shadow import ( + CheckpointReason, + CombinedRestoreResult, + RestoreOutcome, + ShadowCheckpointConfig, + ShadowCheckpointManager, + ShadowCheckpointStore, +) + + +@pytest.fixture() +def work_dir(tmp_path: Path) -> Path: + d = tmp_path / "project" + d.mkdir() + (d / "README.md").write_text("v1\n", encoding="utf-8") + return d + + +@pytest.fixture() +def mgr(tmp_path: Path) -> ShadowCheckpointManager: + return ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "checkpoints"), + config=ShadowCheckpointConfig(enabled=True, max_snapshots=10), + ) + + +def _msgs(*texts: str) -> list[ConversationMessage]: + out: list[ConversationMessage] = [] + for i, text in enumerate(texts): + role = "user" if i % 2 == 0 else "assistant" + out.append(ConversationMessage(role=role, content=[TextBlock(text=text)])) + return out + + +def test_restore_and_rewind_aligns_fs_and_transcript( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + taken = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert taken.outcome.value == "taken" + assert taken.snapshot is not None + sha = taken.snapshot.commit_sha + + (work_dir / "README.md").write_text("v2-broken\n", encoding="utf-8") + messages = _msgs("fix the bug", "I edited README", "looks good?", "ship it") + + result = mgr.restore_and_rewind( + work_dir, + commit_sha=sha, + messages=messages, + prompt_indices=(0, 2), + rewind_turns=1, + ) + assert isinstance(result, CombinedRestoreResult) + assert result.fs.outcome is RestoreOutcome.RESTORED + assert (work_dir / "README.md").read_text(encoding="utf-8") == "v1\n" + assert result.messages == tuple(messages[:2]) + assert result.transcript_removed == 2 + + +def test_restore_and_rewind_fs_only_when_rewind_zero( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + taken = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert taken.snapshot is not None + (work_dir / "README.md").write_text("changed\n", encoding="utf-8") + messages = _msgs("a", "b") + + result = mgr.restore_and_rewind( + work_dir, + commit_sha=taken.snapshot.commit_sha, + messages=messages, + prompt_indices=(0,), + rewind_turns=0, + ) + assert result.fs.outcome is RestoreOutcome.RESTORED + assert result.messages == tuple(messages) + assert result.transcript_removed == 0 + + +def test_restore_and_rewind_propagates_fs_failure( + mgr: ShadowCheckpointManager, work_dir: Path +) -> None: + messages = _msgs("a", "b") + result = mgr.restore_and_rewind( + work_dir, + commit_sha="deadbeef" * 5, + messages=messages, + prompt_indices=(0,), + rewind_turns=1, + ) + assert result.fs.outcome is RestoreOutcome.NOT_FOUND + assert result.messages == tuple(messages) + assert result.transcript_removed == 0 diff --git a/tests/test_state/test_shadow_checkpoint.py b/tests/test_state/test_shadow_checkpoint.py index d00f0729..26631f75 100644 --- a/tests/test_state/test_shadow_checkpoint.py +++ b/tests/test_state/test_shadow_checkpoint.py @@ -68,6 +68,29 @@ def test_dedup_same_turn(mgr: ShadowCheckpointManager, work_dir: Path) -> None: assert second.outcome is CheckpointOutcome.ALREADY_THIS_TURN +def test_dedup_is_scoped_to_session(mgr: ShadowCheckpointManager, work_dir: Path) -> None: + first = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE, session_id="a") + (work_dir / "README.md").write_text("changed\n", encoding="utf-8") + second = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_BASH, session_id="b") + assert first.outcome is CheckpointOutcome.TAKEN + assert second.outcome is CheckpointOutcome.TAKEN + + +def test_negative_rewind_does_not_restore(mgr: ShadowCheckpointManager, work_dir: Path) -> None: + taken = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert taken.snapshot is not None + (work_dir / "README.md").write_text("mutated\n", encoding="utf-8") + result = mgr.restore_and_rewind( + work_dir, + commit_sha=taken.snapshot.commit_sha, + messages=[], + prompt_indices=(), + rewind_turns=-1, + ) + assert result.fs.outcome is RestoreOutcome.FAILED + assert (work_dir / "README.md").read_text(encoding="utf-8") == "mutated\n" + + def test_new_turn_allows_another_when_changed(mgr: ShadowCheckpointManager, work_dir: Path) -> None: assert ( mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE).outcome @@ -90,6 +113,57 @@ def test_no_changes_skips(mgr: ShadowCheckpointManager, work_dir: Path) -> None: assert skipped.outcome is CheckpointOutcome.NO_CHANGES +def test_small_worktree_checkpoints_normally(tmp_path: Path) -> None: + work_dir = tmp_path / "small" + work_dir.mkdir() + (work_dir / "README.md").write_text("hello\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "checkpoints"), + config=ShadowCheckpointConfig(max_files=1), + ) + result = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert result.outcome is CheckpointOutcome.TAKEN + + +def test_large_worktree_skips_without_creating_store(tmp_path: Path) -> None: + work_dir = tmp_path / "large" + work_dir.mkdir() + for index in range(3): + (work_dir / f"file-{index}.txt").write_text("x\n", encoding="utf-8") + store_root = tmp_path / "checkpoints" + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=store_root), + config=ShadowCheckpointConfig(max_files=2), + ) + result = mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + assert result.outcome is CheckpointOutcome.DIRECTORY_TOO_LARGE + assert not store_root.exists() + + +def test_worktree_size_probe_runs_once_per_directory( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + work_dir = tmp_path / "large" + work_dir.mkdir() + (work_dir / "file.txt").write_text("x\n", encoding="utf-8") + mgr = ShadowCheckpointManager( + store=ShadowCheckpointStore(base_dir=tmp_path / "checkpoints"), + config=ShadowCheckpointConfig(max_files=0), + ) + calls = 0 + original = mgr._probe_worktree_size + + def probe_once(path: Path) -> bool: + nonlocal calls + calls += 1 + return original(path) + + monkeypatch.setattr(mgr, "_probe_worktree_size", probe_once) + mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_WRITE_FILE) + mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_BASH) + assert calls == 1 + + def test_skips_home_and_root(mgr: ShadowCheckpointManager) -> None: assert ( mgr.ensure(Path("/"), reason=CheckpointReason.BEFORE_BASH).outcome @@ -229,5 +303,10 @@ async def test_hook_begin_turn_on_user_prompt(mgr: ShadowCheckpointManager, work await hook(HookEvent.USER_PROMPT_SUBMIT, {"session_id": "s1", "prompt": "go"}) (work_dir / "README.md").write_text("next\n", encoding="utf-8") assert ( - mgr.ensure(work_dir, reason=CheckpointReason.BEFORE_BASH).outcome is CheckpointOutcome.TAKEN + mgr.ensure( + work_dir, + reason=CheckpointReason.BEFORE_BASH, + session_id="s1", + ).outcome + is CheckpointOutcome.TAKEN ) diff --git a/tests/test_state/test_transcript_rewind.py b/tests/test_state/test_transcript_rewind.py new file mode 100644 index 00000000..7dfbe7e7 --- /dev/null +++ b/tests/test_state/test_transcript_rewind.py @@ -0,0 +1,60 @@ +"""Transcript rewind — Session-owned prompt boundaries.""" + +from __future__ import annotations + +from dream.engine._messages import ConversationMessage, TextBlock +from dream.state.shadow._rewind import rewind_transcript + + +def _message(role: str, text: str) -> ConversationMessage: + return ConversationMessage(role=role, content=[TextBlock(text=text)]) + + +def test_rewind_zero_turns_is_noop() -> None: + messages = [_message("user", "a"), _message("assistant", "b")] + kept, removed = rewind_transcript(messages, prompt_indices=(0,), turns=0) + assert kept == messages + assert removed == 0 + + +def test_rewind_one_turn_drops_last_prompt_cycle() -> None: + messages = [ + _message("user", "task A"), + _message("assistant", "done A"), + _message("user", "task B"), + _message("assistant", "done B"), + ] + kept, removed = rewind_transcript(messages, prompt_indices=(0, 2), turns=1) + assert kept == messages[:2] + assert removed == 2 + + +def test_rewind_two_turns_clears_both_cycles() -> None: + messages = [ + _message("user", "A"), + _message("assistant", "a"), + _message("user", "B"), + _message("assistant", "b"), + ] + kept, removed = rewind_transcript(messages, prompt_indices=(0, 2), turns=2) + assert kept == [] + assert removed == 4 + + +def test_rewind_rejects_unavailable_boundary() -> None: + messages = [_message("user", "compacted")] + try: + rewind_transcript(messages, prompt_indices=(), turns=1) + except ValueError as exc: + assert "unavailable" in str(exc) + else: + raise AssertionError("expected ValueError") + + +def test_rewind_rejects_negative_turns() -> None: + try: + rewind_transcript([_message("user", "x")], prompt_indices=(0,), turns=-1) + except ValueError as exc: + assert "turns" in str(exc) + else: + raise AssertionError("expected ValueError")