diff --git a/CHANGELOG.md b/CHANGELOG.md index 9367ca21..f56355dd 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")