191 lines
6.6 KiB
Python
191 lines
6.6 KiB
Python
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)
|