Files
reader/src/summary_mcp/runtime/run_store.py
T

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)