"""Step definitions for checkpoint_rollback.feature. Exercises checkpoint domain models, CheckpointService operations (create, list, rollback, prune, delete), guard validations, and correction-service integration. """ from __future__ import annotations import os import shutil import subprocess import tempfile from pathlib import Path from behave import given, then, when from pydantic import ValidationError as PydanticValidationError from cleveragents.application.services.checkpoint_service import CheckpointService from cleveragents.application.services.correction_service import CorrectionService from cleveragents.core.exceptions import ( BusinessRuleViolation, ResourceNotFoundError, ValidationError, ) from cleveragents.domain.models.core.checkpoint import ( Checkpoint, CheckpointMetadata, CheckpointRetentionPolicy, RollbackResult, ) from cleveragents.infrastructure.events.models import DomainEvent _VALID_ULID = "01ARZ3NDEKTSV4RRFFQ69G5FAV" def _create_git_workspace(context: object) -> str: """Create a temporary git workspace and register cleanup on context.""" tmpdir = tempfile.mkdtemp(prefix="checkpoint_test_") subprocess.run( ["git", "init", "--initial-branch=main"], cwd=tmpdir, capture_output=True, text=True, check=True, ) subprocess.run( ["git", "config", "user.email", "test@example.com"], cwd=tmpdir, capture_output=True, text=True, check=True, ) subprocess.run( ["git", "config", "user.name", "Test"], cwd=tmpdir, capture_output=True, text=True, check=True, ) # Create an initial commit so the repo is not empty placeholder = os.path.join(tmpdir, ".gitkeep") Path(placeholder).write_text("") subprocess.run( ["git", "add", ".gitkeep"], cwd=tmpdir, capture_output=True, text=True, check=True, ) subprocess.run( ["git", "commit", "-m", "Initial commit"], cwd=tmpdir, capture_output=True, text=True, check=True, ) def _cleanup() -> None: shutil.rmtree(tmpdir, ignore_errors=True) if hasattr(context, "add_cleanup"): context.add_cleanup(_cleanup) return tmpdir def _get_head_sha(cwd: str) -> str: """Return the HEAD commit SHA of a git repository.""" result = subprocess.run( ["git", "rev-parse", "HEAD"], cwd=cwd, capture_output=True, text=True, check=True, ) return result.stdout.strip() # ------------------------------------------------------------------- # Domain model steps # ------------------------------------------------------------------- @given('a valid checkpoint with plan_id "{plan_id}" and sandbox_ref "{sandbox_ref}"') def step_create_valid_checkpoint(context, plan_id, sandbox_ref): from ulid import ULID context.checkpoint = Checkpoint( checkpoint_id=str(ULID()), plan_id=plan_id, sandbox_ref=sandbox_ref, ) @then('the checkpoint plan_id should be "{expected}"') def step_check_plan_id(context, expected): assert context.checkpoint.plan_id == expected @then('the checkpoint sandbox_ref should be "{expected}"') def step_check_sandbox_ref(context, expected): assert context.checkpoint.sandbox_ref == expected @given( 'a checkpoint with reason "{reason}" and source_tool "{tool}" and phase "{phase}"' ) def step_create_checkpoint_with_metadata(context, reason, tool, phase): from ulid import ULID context.checkpoint = Checkpoint( checkpoint_id=str(ULID()), plan_id=_VALID_ULID, sandbox_ref="ref123", metadata=CheckpointMetadata(reason=reason, source_tool=tool, phase=phase), ) @then('the checkpoint metadata reason should be "{expected}"') def step_check_metadata_reason(context, expected): assert context.checkpoint.metadata.reason == expected @then('the checkpoint metadata source_tool should be "{expected}"') def step_check_metadata_source_tool(context, expected): assert context.checkpoint.metadata.source_tool == expected @then('the checkpoint metadata phase should be "{expected}"') def step_check_metadata_phase(context, expected): assert context.checkpoint.metadata.phase == expected @given("a default retention policy") def step_default_retention(context): context.policy = CheckpointRetentionPolicy() @then("the max_checkpoints should be {count:d}") def step_check_max_checkpoints(context, count): assert context.policy.max_checkpoints == count @then("auto_prune should be true") def step_check_auto_prune_true(context): assert context.policy.auto_prune is True @then("auto_prune should be false") def step_check_auto_prune_false(context): assert context.policy.auto_prune is False @given("a retention policy with max_checkpoints {max_cp:d} and auto_prune false") def step_custom_retention(context, max_cp): context.policy = CheckpointRetentionPolicy(max_checkpoints=max_cp, auto_prune=False) @given('a rollback result with {count:d} restored files and checkpoint "{cp_id}"') def step_create_rollback_result(context, count, cp_id): context.rollback_result = RollbackResult( restored_files_count=count, changed_paths=[f"file_{i}.py" for i in range(count)], from_checkpoint_id=cp_id, ) @then("the restored_files_count should be {count:d}") def step_check_restored_count(context, count): assert context.rollback_result.restored_files_count == count @then('the from_checkpoint_id should be "{expected}"') def step_check_from_checkpoint(context, expected): assert context.rollback_result.from_checkpoint_id == expected # ------------------------------------------------------------------- # CheckpointService steps # ------------------------------------------------------------------- @given("a checkpoint service") def step_create_service(context): context.svc = CheckpointService() context.checkpoint = None context.checkpoints = None context.rollback_result = None context.error = None context.pruned_ids = None @given('a checkpoint service with sandbox for plan "{plan_id}"') def step_create_service_with_sandbox(context, plan_id): workspace = _create_git_workspace(context) context.svc = CheckpointService() context.svc.register_sandbox(plan_id, workspace) context.workspace_dir = workspace context.checkpoint = None context.checkpoints = None context.rollback_result = None context.error = None context.pruned_ids = None @given('plan "{plan_id}" is marked as applied') def step_mark_applied(context, plan_id): context.svc.mark_plan_applied(plan_id) @given('I create a checkpoint for plan "{plan_id}" with sandbox_ref "{ref}"') @when('I create a checkpoint for plan "{plan_id}" with sandbox_ref "{ref}"') def step_create_checkpoint(context, plan_id, ref): context.checkpoint = context.svc.create_checkpoint(plan_id=plan_id, sandbox_ref=ref) @then("the checkpoint should be created successfully") def step_checkpoint_created(context): assert context.checkpoint is not None assert context.checkpoint.checkpoint_id @then('the checkpoint plan_id should match "{expected}"') def step_checkpoint_plan_match(context, expected): assert context.checkpoint.plan_id == expected @when('I list checkpoints for plan "{plan_id}"') def step_list_checkpoints(context, plan_id): context.checkpoints = context.svc.list_checkpoints(plan_id) @then("I should see {count:d} checkpoints") def step_check_checkpoint_count(context, count): assert len(context.checkpoints) == count @when('I rollback plan "{plan_id}" to the created checkpoint') def step_rollback_to_checkpoint(context, plan_id): context.rollback_result = context.svc.rollback_to_checkpoint( plan_id, context.checkpoint.checkpoint_id ) @then("the rollback should succeed") def step_rollback_success(context): assert context.rollback_result is not None @then("the rollback result should show restored files") def step_rollback_shows_files(context): assert context.rollback_result.restored_files_count > 0 @when('I attempt rollback plan "{plan_id}" to the created checkpoint') def step_attempt_rollback(context, plan_id): try: context.rollback_result = context.svc.rollback_to_checkpoint( plan_id, context.checkpoint.checkpoint_id ) context.error = None except (BusinessRuleViolation, ResourceNotFoundError, ValidationError) as e: context.error = e @when('I attempt rollback plan "{plan_id}" to checkpoint "{cp_id}"') def step_attempt_rollback_by_id(context, plan_id, cp_id): try: context.rollback_result = context.svc.rollback_to_checkpoint(plan_id, cp_id) context.error = None except (BusinessRuleViolation, ResourceNotFoundError, ValidationError) as e: context.error = e @then('the rollback should be rejected with "{msg}"') def step_rollback_rejected(context, msg): assert context.error is not None assert msg in str(context.error) @then("the rollback should fail with not found error") def step_rollback_not_found(context): assert context.error is not None assert isinstance(context.error, ResourceNotFoundError) @when("I delete the created checkpoint") def step_delete_checkpoint(context): context.svc.delete_checkpoint(context.checkpoint.checkpoint_id) @then("the checkpoint should be deleted") def step_checkpoint_deleted(context): try: context.svc.get_checkpoint(context.checkpoint.checkpoint_id) msg = "Checkpoint should have been deleted" raise AssertionError(msg) except ResourceNotFoundError: pass @given('I create {count:d} checkpoints for plan "{plan_id}"') def step_create_multiple_checkpoints(context, count, plan_id): for i in range(count): context.checkpoint = context.svc.create_checkpoint( plan_id=plan_id, sandbox_ref=f"commit-{i}" ) @when("I prune checkpoints with max {max_cp:d}") def step_prune_checkpoints(context, max_cp): policy = CheckpointRetentionPolicy(max_checkpoints=max_cp, auto_prune=True) context.pruned_ids = context.svc.prune_checkpoints(_VALID_ULID, policy) @then("{count:d} checkpoints should be pruned") def step_check_pruned_count(context, count): assert len(context.pruned_ids) == count @then("{count:d} checkpoint snapshots should remain") def step_check_remaining_count(context, count): remaining = context.svc.list_checkpoints(_VALID_ULID) assert len(remaining) == count @when("I create a checkpoint with retention policy max {max_cp:d}") def step_create_with_retention(context, max_cp): policy = CheckpointRetentionPolicy(max_checkpoints=max_cp, auto_prune=True) before_count = len(context.svc.list_checkpoints(_VALID_ULID)) context.checkpoint = context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="commit-retention", retention_policy=policy, ) after_count = len(context.svc.list_checkpoints(_VALID_ULID)) context.pruned_on_creation = (before_count + 1) - after_count @then("{count:d} checkpoints should have been pruned on creation") def step_check_pruned_on_creation(context, count): assert context.pruned_on_creation == count @when("I attempt to create a checkpoint with empty plan_id") def step_create_empty_plan(context): try: context.svc.create_checkpoint(plan_id="", sandbox_ref="ref") context.error = None except (ValidationError, PydanticValidationError) as e: context.error = e @when("I attempt to create a checkpoint with empty sandbox_ref") def step_create_empty_ref(context): try: context.svc.create_checkpoint(plan_id=_VALID_ULID, sandbox_ref="") context.error = None except (ValidationError, PydanticValidationError) as e: context.error = e @then("a checkpoint validation error should be raised") def step_checkpoint_validation_error(context): assert context.error is not None assert isinstance(context.error, (ValidationError, PydanticValidationError)) @when("I get the checkpoint by its ID") def step_get_checkpoint_by_id(context): context.result_checkpoint = context.svc.get_checkpoint( context.checkpoint.checkpoint_id ) @then("the checkpoint should be returned successfully") def step_checkpoint_returned(context): assert context.result_checkpoint is not None assert context.result_checkpoint.checkpoint_id == context.checkpoint.checkpoint_id @when('I attempt to get checkpoint "{cp_id}"') def step_get_nonexistent(context, cp_id): try: context.svc.get_checkpoint(cp_id) context.error = None except ResourceNotFoundError as e: context.error = e @then("a checkpoint not found error should be raised") def step_checkpoint_not_found_error(context): assert context.error is not None assert isinstance(context.error, ResourceNotFoundError) # ------------------------------------------------------------------- # Correction service integration steps # ------------------------------------------------------------------- @given("a correction service with checkpoint support") def step_correction_with_checkpoint(context): cp_svc = CheckpointService() context.correction_svc = CorrectionService(checkpoint_service=cp_svc) @then("the correction service should have checkpoint support") def step_check_checkpoint_support(context): assert context.correction_svc._checkpoint_service is not None @when( 'I create a checkpoint with reason "{reason}" and source_tool "{tool}" and phase "{phase}"' ) def step_create_with_metadata(context, reason, tool, phase): context.checkpoint = context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="ref-meta", reason=reason, source_tool=tool, phase=phase, ) @then('the stored checkpoint metadata should have reason "{expected}"') def step_stored_reason(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.metadata.reason == expected @then('the stored checkpoint metadata should have source_tool "{expected}"') def step_stored_source_tool(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.metadata.source_tool == expected @then('the stored checkpoint metadata should have phase "{expected}"') def step_stored_phase(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.metadata.phase == expected # ------------------------------------------------------------------- # Spec-aligned field steps (decision_id, checkpoint_type, etc.) # ------------------------------------------------------------------- @when('I create a checkpoint aligned to decision "{decision_id}"') def step_create_with_decision(context, decision_id): context.checkpoint = context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="ref-decision", decision_id=decision_id, ) @then('the stored checkpoint decision_id should be "{expected}"') def step_check_decision_id(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.decision_id == expected @when('I create a pre_write checkpoint for plan "{plan_id}"') def step_create_prewrite(context, plan_id): context.checkpoint = context.svc.create_checkpoint( plan_id=plan_id, sandbox_ref="ref-prewrite", checkpoint_type="pre_write", ) @then('the stored checkpoint_type should be "{expected}"') def step_check_checkpoint_type(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.checkpoint_type == expected @when('I attempt to create a checkpoint with type "{cp_type}"') def step_create_invalid_type(context, cp_type): try: context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="ref-invalid", checkpoint_type=cp_type, ) context.error = None except (ValidationError, PydanticValidationError) as e: context.error = e @when('I create a checkpoint with resource "{resource_id}"') def step_create_with_resource(context, resource_id): context.checkpoint = context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="ref-resource", resource_id=resource_id, ) @then('the stored checkpoint resource_id should be "{expected}"') def step_check_resource_id(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.resource_id == expected @when('I create a checkpoint with path "{path}" and size {size:d}') def step_create_with_path_and_size(context, path, size): context.checkpoint = context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="ref-path", filesystem_path=path, size_bytes=size, ) @then('the stored checkpoint filesystem_path should be "{expected}"') def step_check_filesystem_path(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.filesystem_path == expected @then("the stored checkpoint size_bytes should be {expected:d}") def step_check_size_bytes(context, expected): cp = context.svc.get_checkpoint(context.checkpoint.checkpoint_id) assert cp.size_bytes == expected @then("the first checkpoint should survive pruning") def step_first_survives_pruning(context): """Verify the first (oldest) checkpoint survived pruning.""" remaining = context.svc.list_checkpoints(_VALID_ULID) # The first checkpoint created has sandbox_ref "commit-0" assert any(cp.sandbox_ref == "commit-0" for cp in remaining), ( "First checkpoint (commit-0) should be preserved after pruning" ) @then("the most recent checkpoint should survive pruning") def step_most_recent_survives_pruning(context): """Verify the most recent checkpoint survived pruning.""" remaining = context.svc.list_checkpoints(_VALID_ULID) # The last checkpoint created has the highest commit-N last_ref = context.checkpoint.sandbox_ref assert any(cp.sandbox_ref == last_ref for cp in remaining), ( f"Most recent checkpoint ({last_ref}) should be preserved after pruning" ) # ------------------------------------------------------------------- # Persistent guard steps (lifecycle-service backed) # ------------------------------------------------------------------- class _FakePlan: """Minimal plan stub for lifecycle-service guard tests.""" def __init__(self, plan_id, processing_state, sandbox_refs=None): self.plan_id = plan_id self.processing_state = processing_state self.sandbox_refs = sandbox_refs or [] class _FakeLifecycleService: """Stub PlanLifecycleService that returns a canned plan.""" def __init__(self, plan): self._plan = plan def get_plan(self, plan_id): return self._plan @given( 'a checkpoint service backed by a lifecycle service with an applied plan "{plan_id}"' ) def step_service_lifecycle_applied(context, plan_id): from cleveragents.domain.models.core.plan import ProcessingState fake_plan = _FakePlan(plan_id, ProcessingState.APPLIED, sandbox_refs=["sb"]) fake_ls = _FakeLifecycleService(fake_plan) context.svc = CheckpointService(plan_lifecycle_service=fake_ls) # type: ignore[arg-type] context.checkpoint = None context.error = None context.rollback_result = None @given( 'a checkpoint service backed by a lifecycle service with no sandbox for plan "{plan_id}"' ) def step_service_lifecycle_no_sandbox(context, plan_id): from cleveragents.domain.models.core.plan import ProcessingState fake_plan = _FakePlan(plan_id, ProcessingState.PROCESSING, sandbox_refs=[]) fake_ls = _FakeLifecycleService(fake_plan) context.svc = CheckpointService(plan_lifecycle_service=fake_ls) # type: ignore[arg-type] context.checkpoint = None context.error = None context.rollback_result = None @given( 'a checkpoint service backed by a lifecycle service with sandbox for plan "{plan_id}"' ) def step_service_lifecycle_with_sandbox(context, plan_id): from cleveragents.domain.models.core.plan import ProcessingState fake_plan = _FakePlan( plan_id, ProcessingState.PROCESSING, sandbox_refs=["sandbox-root"] ) fake_ls = _FakeLifecycleService(fake_plan) context.svc = CheckpointService(plan_lifecycle_service=fake_ls) # type: ignore[arg-type] context.checkpoint = None context.error = None context.rollback_result = None # ------------------------------------------------------------------- # Default retention policy step # ------------------------------------------------------------------- @then('at most {max_cp:d} checkpoint snapshots should remain for plan "{plan_id}"') def step_at_most_checkpoints_remain(context, max_cp, plan_id): remaining = context.svc.list_checkpoints(plan_id) assert len(remaining) <= max_cp, ( f"Expected at most {max_cp} checkpoints, but found {len(remaining)}" ) @given('the sandbox for plan "{plan_id}" is unregistered') def step_unregister_sandbox(context, plan_id): context.svc.unregister_sandbox(plan_id) @when("I prune checkpoints with auto_prune disabled") def step_prune_no_auto(context): policy = CheckpointRetentionPolicy(max_checkpoints=3, auto_prune=False) context.pruned_ids = context.svc.prune_checkpoints(_VALID_ULID, policy) # ------------------------------------------------------------------- # Stub repository for repository-backed tests # ------------------------------------------------------------------- class _StubCheckpointRepository: """In-memory repository stub that implements the CheckpointRepository interface. Used to exercise the repository-backed code paths in CheckpointService without requiring a real database connection. """ def __init__(self) -> None: self._store: dict[str, Checkpoint] = {} self._plan_index: dict[str, list[str]] = {} def create(self, checkpoint: Checkpoint) -> Checkpoint: self._store[checkpoint.checkpoint_id] = checkpoint self._plan_index.setdefault(checkpoint.plan_id, []).append( checkpoint.checkpoint_id ) return checkpoint def get_by_id(self, checkpoint_id: str) -> Checkpoint: from cleveragents.infrastructure.database.repositories import ( CheckpointNotFoundError, ) cp = self._store.get(checkpoint_id) if cp is None: raise CheckpointNotFoundError(checkpoint_id) return cp def list_by_plan(self, plan_id: str) -> list[Checkpoint]: ids = self._plan_index.get(plan_id, []) return [self._store[cid] for cid in ids if cid in self._store] def delete(self, checkpoint_id: str) -> bool: cp = self._store.pop(checkpoint_id, None) if cp is None: return False plan_list = self._plan_index.get(cp.plan_id, []) if checkpoint_id in plan_list: plan_list.remove(checkpoint_id) return True def prune(self, plan_id: str, max_checkpoints: int) -> list[str]: ids = self._plan_index.get(plan_id, []) if len(ids) <= max_checkpoints or len(ids) < 3: return [] interior = ids[1:-1] excess = len(ids) - max_checkpoints to_remove = interior[:excess] for cid in to_remove: self._store.pop(cid, None) self._plan_index[plan_id] = [cid for cid in ids if cid not in set(to_remove)] return to_remove @given("a checkpoint service backed by a stub repository") def step_create_repo_backed_service(context): repo = _StubCheckpointRepository() context.stub_repo = repo context.svc = CheckpointService(repository=repo) # type: ignore[arg-type] context.checkpoint = None context.checkpoints = None context.rollback_result = None context.error = None context.pruned_ids = None @then("the stub repository should have {count:d} checkpoint stored") def step_check_stub_repo_count(context, count): assert len(context.stub_repo._store) == count @when('I attempt to delete checkpoint "{cp_id}"') def step_attempt_delete_checkpoint(context, cp_id): try: context.svc.delete_checkpoint(cp_id) context.error = None except ResourceNotFoundError as e: context.error = e # ------------------------------------------------------------------- # In-memory edge case steps # ------------------------------------------------------------------- @given("a checkpoint service with an orphaned plan index entry") def step_service_with_orphaned_index(context): """Create a service where _plan_index has an ID not in _checkpoints.""" context.svc = CheckpointService() # Create a real checkpoint first real_cp = context.svc.create_checkpoint( plan_id=_VALID_ULID, sandbox_ref="commit-real" ) # Add a ghost entry to plan_index that has no corresponding checkpoint context.svc._plan_index[_VALID_ULID].append("01GHOST0000000000000000000") context.checkpoint = real_cp context.checkpoints = None context.rollback_result = None context.error = None context.pruned_ids = None @given("a checkpoint service with a checkpoint missing from plan index") def step_service_with_missing_plan_index(context): """Create a service where a checkpoint exists in _checkpoints but not in _plan_index.""" context.svc = CheckpointService() cp = context.svc.create_checkpoint(plan_id=_VALID_ULID, sandbox_ref="commit-orphan") # Remove the checkpoint from plan_index but leave it in _checkpoints context.svc._plan_index[_VALID_ULID].clear() context.checkpoint = cp context.checkpoints = None context.rollback_result = None context.error = None context.pruned_ids = None # ------------------------------------------------------------------- # Domain model validator edge case steps # ------------------------------------------------------------------- @when("I invoke the checkpoint ULID validator with an empty string") def step_invoke_ulid_validator_empty(context): try: Checkpoint.validate_ulid("") context.error = None except ValueError as e: context.error = e @then('the checkpoint validator should raise ValueError "{msg}"') def step_check_checkpoint_value_error(context, msg): assert context.error is not None, "Expected a ValueError" assert isinstance(context.error, ValueError) assert msg in str(context.error) # ------------------------------------------------------------------- # Real git rollback steps (bug #822 fix) # ------------------------------------------------------------------- _SANDBOX_INITIAL_CONTENT = "initial sandbox content\n" _SANDBOX_MODIFIED_CONTENT = "modified after checkpoint\n" _SANDBOX_TRACKED_FILE = "sandbox_tracked.txt" _SANDBOX_NEW_FILE = "sandbox_extra.txt" @given('a checkpoint is created from the sandbox HEAD for plan "{plan_id}"') def step_create_checkpoint_from_sandbox_head(context, plan_id): """Create a checkpoint using the real HEAD SHA of the sandbox workspace.""" head_sha = _get_head_sha(context.workspace_dir) context.checkpoint = context.svc.create_checkpoint( plan_id=plan_id, sandbox_ref=head_sha, reason="checkpoint from sandbox HEAD", checkpoint_type="manual", ) context.checkpoint_head_sha = head_sha @given("a file is modified in the sandbox after the checkpoint") def step_modify_sandbox_file(context): """Modify a file in the sandbox and commit to advance HEAD.""" tracked_path = os.path.join(context.workspace_dir, _SANDBOX_TRACKED_FILE) Path(tracked_path).write_text(_SANDBOX_MODIFIED_CONTENT) subprocess.run( ["git", "add", _SANDBOX_TRACKED_FILE], cwd=context.workspace_dir, capture_output=True, text=True, check=True, ) subprocess.run( ["git", "commit", "-m", "Modify sandbox file"], cwd=context.workspace_dir, capture_output=True, text=True, check=True, ) @given("a tracked sandbox file is modified and committed after the checkpoint") def step_modify_tracked_sandbox_file(context): """Create, commit, then modify a tracked sandbox file after checkpoint.""" tracked_path = os.path.join(context.workspace_dir, _SANDBOX_TRACKED_FILE) # Write initial content and commit (this is the pre-checkpoint state # since .gitkeep was the only file at checkpoint time) Path(tracked_path).write_text(_SANDBOX_INITIAL_CONTENT) subprocess.run( ["git", "add", _SANDBOX_TRACKED_FILE], cwd=context.workspace_dir, capture_output=True, text=True, check=True, ) subprocess.run( ["git", "commit", "-m", "Add tracked file"], cwd=context.workspace_dir, capture_output=True, text=True, check=True, ) context.sandbox_tracked_path = tracked_path @given("a new file is added and committed in the sandbox after the checkpoint") def step_add_new_sandbox_file(context): """Add a new file and commit it in the sandbox after the checkpoint.""" new_path = os.path.join(context.workspace_dir, _SANDBOX_NEW_FILE) Path(new_path).write_text("new file content\n") subprocess.run( ["git", "add", _SANDBOX_NEW_FILE], cwd=context.workspace_dir, capture_output=True, text=True, check=True, ) subprocess.run( ["git", "commit", "-m", "Add new sandbox file"], cwd=context.workspace_dir, capture_output=True, text=True, check=True, ) context.sandbox_new_file_path = new_path @then("the sandbox file content should match the pre-checkpoint state") def step_assert_sandbox_file_reverted(context): """Assert the tracked sandbox file no longer exists (it was added after checkpoint).""" assert not os.path.exists(context.sandbox_tracked_path), ( f"Tracked file should not exist after rollback to pre-checkpoint state. " f"File still exists: {context.sandbox_tracked_path}" ) @then("the file added after the checkpoint should not exist") def step_assert_sandbox_new_file_removed(context): """Assert the file added after checkpoint is removed by rollback.""" assert not os.path.exists(context.sandbox_new_file_path), ( f"New file should not exist after rollback. " f"File still exists: {context.sandbox_new_file_path}" ) class _RecordingEventBus: """In-memory event bus that records emitted events for assertions.""" def __init__(self) -> None: self.events: list[DomainEvent] = [] def emit(self, event: DomainEvent) -> None: self.events.append(event) def subscribe(self, event_type: object, handler: object) -> None: pass @given('a checkpoint service with sandbox and event bus for plan "{plan_id}"') def step_create_service_with_sandbox_and_event_bus(context, plan_id): workspace = _create_git_workspace(context) event_bus = _RecordingEventBus() context.svc = CheckpointService(event_bus=event_bus) # type: ignore[arg-type] context.svc.register_sandbox(plan_id, workspace) context.workspace_dir = workspace context.event_bus = event_bus context.checkpoint = None context.checkpoints = None context.rollback_result = None context.error = None context.pruned_ids = None @then("a CHECKPOINT_RESTORED domain event should have been emitted") def step_assert_checkpoint_restored_event(context): from cleveragents.infrastructure.events.types import EventType events = context.event_bus.events restored_events = [ e for e in events if e.event_type == EventType.CHECKPOINT_RESTORED ] assert len(restored_events) == 1, ( f"Expected exactly 1 CHECKPOINT_RESTORED event, got {len(restored_events)}" ) @given('a checkpoint service with nonexistent sandbox path for plan "{plan_id}"') def step_create_service_with_nonexistent_sandbox(context, plan_id): context.svc = CheckpointService() context.svc.register_sandbox(plan_id, "/tmp/nonexistent_sandbox_path_822") context.checkpoint = None context.error = None context.rollback_result = None @given('a checkpoint service with non-git sandbox path for plan "{plan_id}"') def step_create_service_with_non_git_sandbox(context, plan_id): tmpdir = tempfile.mkdtemp(prefix="checkpoint_non_git_") context.svc = CheckpointService() context.svc.register_sandbox(plan_id, tmpdir) context.checkpoint = None context.error = None context.rollback_result = None def _cleanup() -> None: shutil.rmtree(tmpdir, ignore_errors=True) context.add_cleanup(_cleanup) @given( 'a checkpoint service backed by a lifecycle service with real sandbox for plan "{plan_id}"' ) def step_service_lifecycle_with_real_sandbox(context, plan_id): from cleveragents.domain.models.core.plan import ProcessingState workspace = _create_git_workspace(context) fake_plan = _FakePlan(plan_id, ProcessingState.PROCESSING, sandbox_refs=[workspace]) fake_ls = _FakeLifecycleService(fake_plan) context.svc = CheckpointService(plan_lifecycle_service=fake_ls) # type: ignore[arg-type] context.workspace_dir = workspace context.checkpoint = None context.error = None context.rollback_result = None