mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
2.65 kB
import hashlib
import json
import os
import re
import time
import uuid
from pathlib import Path
from .contracts import StageEvent
OUTPUTS = Path(os.environ.get("STUDIO_OUTPUTS", "outputs")).resolve()
def digest(path: Path) -> str:
with path.open("rb") as stream:
return hashlib.file_digest(stream, "sha256").hexdigest()
class Run:
def __init__(self, request):
self.path = OUTPUTS / uuid.uuid4().hex
self.path.mkdir(parents=True)
self.data = {"id": self.path.name, "request": request.model_dump(mode="json"),
"created": time.time(), "status": "running", "events": [],
"completed_stages": [], "skipped_stages": [], "next_stage": "reference"}
# Avoid retaining the upload cache's private path in downloadable metadata.
self.data["request"]["image"] = "input.png" if request.image else None
self.save()
@classmethod
def load(cls, run_id):
if not isinstance(run_id, str) or not re.fullmatch(r"[0-9a-f]{32}", run_id):
raise ValueError("Invalid run ID; expected 32 lowercase hexadecimal characters.")
path = OUTPUTS / run_id
if path.is_symlink() or path.resolve().parent != OUTPUTS.resolve():
raise ValueError("Run path is outside the output directory.")
manifest = path / "manifest.json"
if manifest.is_symlink():
raise ValueError("Run manifest must not be a symlink.")
data = json.loads(manifest.read_text())
if not isinstance(data, dict) or data.get("id") != run_id:
raise ValueError("Run manifest ID does not match its directory.")
if data.get("request", {}).get("image") not in {None, "input.png"}:
raise ValueError("Run request contains an invalid input path.")
run = cls.__new__(cls)
run.path, run.data = path, data
return run
def save(self):
temporary = self.path / "manifest.tmp"
temporary.write_text(json.dumps(self.data, indent=2))
temporary.replace(self.path / "manifest.json")
def event(self, stage, status, detail):
self.data["events"].append(StageEvent(stage=stage, status=status, detail=detail).model_dump())
self.save()
def finish(self, status):
self.data["status"] = status
self.data["elapsed_seconds"] = round(time.time() - self.data["created"], 2)
self.data["files"] = {str(p.relative_to(self.path)): digest(p)
for p in self.path.rglob("*") if p.is_file()
and p.name not in {"manifest.json", "asset.zip"}}
self.save()