from __future__ import annotations import json from datetime import datetime from pathlib import Path from typing import Any from .state_models import ArtifactRecord, RecoveryState, RunError, RunState, StageState class RunStore: def __init__(self, *, path: Path, state: RunState, repo_root: Path | None = None) -> None: self.path = path self.state = state self.repo_root = repo_root @classmethod def create( cls, *, path: Path, run_id: str, workflow: str, run_type: str, started_at: datetime, input_payload: dict[str, Any] | None = None, repo_root: Path | None = None, ) -> "RunStore": state = RunState( run_id=run_id, workflow=workflow, run_type=run_type, status="running", current_stage=None, started_at=started_at, updated_at=started_at, input=input_payload or {}, recovery=RecoveryState(), ) return cls(path=path, state=state, repo_root=repo_root) @classmethod def load(cls, *, path: Path, repo_root: Path | None = None) -> "RunStore": state = RunState.model_validate(json.loads(path.read_text(encoding="utf-8-sig"))) return cls(path=path, state=state, repo_root=repo_root) def save(self) -> None: self.state.updated_at = datetime.now(tz=self.state.started_at.tzinfo) self.path.parent.mkdir(parents=True, exist_ok=True) self.path.write_text( json.dumps(self.state.model_dump(mode="json"), ensure_ascii=False, indent=2), encoding="utf-8", ) def start_stage(self, name: str, *, outputs: dict[str, Any] | None = None) -> None: stage = self._get_or_create_stage(name) now = datetime.now(tz=self.state.started_at.tzinfo) stage.status = "running" stage.started_at = stage.started_at or now stage.finished_at = None if outputs: stage.outputs.update(outputs) stage.error = None self.state.current_stage = name self.state.status = "running" self.state.error = None self._refresh_recovery() self.save() def update_stage(self, name: str, *, outputs: dict[str, Any] | None = None) -> None: stage = self._get_or_create_stage(name) if outputs: stage.outputs.update(outputs) self.state.current_stage = name self._refresh_recovery() self.save() def finish_stage(self, name: str, *, outputs: dict[str, Any] | None = None) -> None: stage = self._get_or_create_stage(name) now = datetime.now(tz=self.state.started_at.tzinfo) stage.status = "success" stage.started_at = stage.started_at or now stage.finished_at = now if outputs: stage.outputs.update(outputs) stage.error = None if self.state.current_stage == name: self.state.current_stage = None self._refresh_recovery() self.save() def fail_stage( self, name: str, *, error: BaseException | RunError, outputs: dict[str, Any] | None = None, ) -> None: stage = self._get_or_create_stage(name) now = datetime.now(tz=self.state.started_at.tzinfo) stage.status = "failed" stage.started_at = stage.started_at or now stage.finished_at = now if outputs: stage.outputs.update(outputs) stage_error = error if isinstance(error, RunError) else self._build_error(error, stage=name) stage.error = stage_error self.state.current_stage = name self.state.status = "failed" self.state.error = stage_error self.state.finished_at = now self._refresh_recovery() self.save() def register_artifact( self, *, name: str, path: Path, kind: str, stage: str, metadata: dict[str, Any] | None = None, ) -> None: artifact = ArtifactRecord( name=name, path=self._normalize_path(path), kind=kind, stage=stage, exists=path.exists(), created_at=datetime.now(tz=self.state.started_at.tzinfo), metadata=metadata or {}, ) existing = next((item for item in self.state.artifacts if item.name == name), None) if existing is None: self.state.artifacts.append(artifact) else: existing.path = artifact.path existing.kind = artifact.kind existing.stage = artifact.stage existing.exists = artifact.exists existing.created_at = artifact.created_at existing.metadata = artifact.metadata self.save() def finish_run(self, *, status: str) -> None: now = datetime.now(tz=self.state.started_at.tzinfo) self.state.status = status self.state.current_stage = None self.state.finished_at = now self.state.error = None self._refresh_recovery() self.save() def _get_or_create_stage(self, name: str) -> StageState: for stage in self.state.stages: if stage.name == name: return stage stage = StageState(name=name) self.state.stages.append(stage) return stage def _refresh_recovery(self) -> None: successful_stages = [stage.name for stage in self.state.stages if stage.status == "success"] last_success_stage = successful_stages[-1] if successful_stages else None resume_from_stage = None if self.state.status == "running": resume_from_stage = self.state.current_stage elif self.state.status == "failed": failed_stage = next((stage.name for stage in reversed(self.state.stages) if stage.status == "failed"), None) resume_from_stage = failed_stage or self.state.current_stage self.state.recovery = RecoveryState( resumable=self.state.status in {"running", "failed"} and resume_from_stage is not None, resume_from_stage=resume_from_stage, last_success_stage=last_success_stage, ) def _build_error(self, error: BaseException, *, stage: str | None = None) -> RunError: return RunError(type=type(error).__name__, message=str(error), stage=stage) def _normalize_path(self, path: Path) -> str: resolved_path = path.resolve() if self.repo_root is not None: try: return str(resolved_path.relative_to(self.repo_root.resolve())) except ValueError: return str(path) return str(path)