File size: 3,722 Bytes
e317359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8c1bce7
 
 
e317359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8c1bce7
 
 
 
 
e317359
 
 
8c1bce7
e317359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
"""Supervise independent UI and evaluation processes on the Space."""
from __future__ import annotations

import json
import hashlib
import os
from pathlib import Path
import signal
import shutil
import subprocess
import sys
import time

ROOT = Path(__file__).resolve().parents[1]


def ensure_bootstrap(root=ROOT, *, downloader=None):
    archive = root / "bootstrap/seed.tar.gz"
    if archive.exists():
        return
    from huggingface_hub import hf_hub_download
    expected = json.loads((root / "bootstrap/manifest.json").read_text())["sha256"]
    download = downloader or hf_hub_download
    source = Path(download(os.environ["HF_STATE_REPO"], f"bootstrap/{expected}.tar.gz",
                           repo_type="dataset", token=os.environ["HF_TOKEN"]))
    if hashlib.sha256(source.read_bytes()).hexdigest() != expected:
        raise ValueError("Remote bootstrap checksum mismatch")
    shutil.copy2(source, archive)


def main():
    os.chdir(ROOT)
    ensure_bootstrap()
    subprocess.run([sys.executable, "scripts/bootstrap_data.py"], check=True)
    (ROOT / ".cloud-state").mkdir(exist_ok=True)
    fallback = ROOT / ".cloud-state/ui-results"
    shutil.rmtree(fallback, ignore_errors=True)
    shutil.copytree(ROOT / "space/results", fallback)
    os.environ["TSFM_RESULTS_PATH"] = str(fallback)
    revision = json.loads((ROOT / "deployment.json").read_text())["source_revision"]
    os.environ["LIVEHOUSE_SOURCE_REVISION"] = revision
    processes = {}
    enabled = os.getenv("TSFM_EVALUATOR_ENABLED", "0") == "1"
    commands = {"web": [sys.executable, "cloud/web.py"]}
    if enabled:
        commands["worker"] = [sys.executable, "cloud/worker.py"]
    if os.getenv("LIVEHOUSE_ACCEPTANCE_ID"):
        commands["acceptance"] = [sys.executable, "cloud/acceptance.py"]
    completed = set()
    stopping = False

    def shutdown(*_):
        nonlocal stopping
        stopping = True
        for proc in processes.values():
            if proc.poll() is None:
                proc.terminate()

    signal.signal(signal.SIGTERM, shutdown)
    signal.signal(signal.SIGINT, shutdown)
    restart_at = {}
    try:
        while not stopping:
            for name, command in commands.items():
                proc = processes.get(name)
                if proc is not None and proc.poll() is not None:
                    if name == "acceptance":
                        print(f"Acceptance process finished with {proc.returncode}; no automatic rerun", flush=True)
                        completed.add(name)
                        del processes[name]
                        continue
                    print(f"{name} exited with {proc.returncode}; restart after 60 seconds", flush=True)
                    restart_at[name] = time.monotonic() + 60
                    del processes[name]
                if name not in processes and name not in completed and time.monotonic() >= restart_at.get(name, 0):
                    processes[name] = subprocess.Popen(command, cwd=ROOT)
            status = {"updated_at": time.time(), "source_revision": revision,
                      "evaluator_enabled": enabled,
                      "worker_running": "worker" in processes and processes["worker"].poll() is None}
            path = ROOT / ".cloud-state/supervisor.json"
            temporary = path.with_suffix(".tmp")
            temporary.write_text(json.dumps(status))
            temporary.replace(path)
            time.sleep(5)
    finally:
        shutdown()
        for proc in processes.values():
            try:
                proc.wait(timeout=35)
            except subprocess.TimeoutExpired:
                proc.kill()
                proc.wait()


if __name__ == "__main__":
    main()