"""A training loop that saves its progress where Nodus checkpoints it and resumes from it.

It uses only the standard library, so it runs in any image. With the Nodus SDK installed,
`nodus.state_dir()`, `nodus.checkpoint.on_request()` and `nodus.restored()` do the same work.
"""

import json
import os
import socket
import threading
import time
from pathlib import Path

STATE_DIR = Path(os.environ.get("NODUS_STATE_DIR", "/nodus/state"))
STATE = STATE_DIR / "progress.json"
STEPS = int(os.environ.get("STEPS", "120"))
SAVE_EVERY = 20  # like Hugging Face Trainer's save_steps: a regular save even when Nodus does not ask

pending = threading.Event()  # set while Nodus waits for a consistent checkpoint
request_seq = None
events = None


def save(step):
    """Write the state atomically, so a snapshot never sees a half-written file."""
    STATE_DIR.mkdir(parents=True, exist_ok=True)
    tmp = STATE.with_suffix(".tmp")
    tmp.write_text(json.dumps({"step": step}))
    os.replace(tmp, STATE)


def listen():
    """Subscribe to checkpoint requests on the events socket and flag each one for the loop."""
    global events, request_seq
    path = os.environ.get("NODUS_EVENTS_SOCKET", "/run/nodus/events.sock")
    try:
        events = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
        events.connect(path)
    except OSError:
        return  # outside Nodus there is nobody to ask for checkpoints
    events.sendall(b'{"type":"checkpoint.subscribe"}\n')
    for line in events.makefile("r"):
        message = json.loads(line)
        if message.get("type") == "checkpoint.request":
            request_seq = message["seq"]
            pending.set()


def ack():
    """Tell Nodus the files in the state directory are complete."""
    events.sendall((json.dumps({"type": "checkpoint.ack", "seq": request_seq}) + "\n").encode())
    pending.clear()


def main():
    start = 0
    if STATE.exists():
        start = json.loads(STATE.read_text())["step"]
        print(f"resumed from step {start}", flush=True)
    threading.Thread(target=listen, daemon=True).start()
    for step in range(start + 1, STEPS + 1):
        time.sleep(1)  # one step of work
        print(f"step {step}/{STEPS}", flush=True)
        if pending.is_set():
            save(step)
            ack()
        elif step % SAVE_EVERY == 0:
            save(step)
    print("training complete", flush=True)


if __name__ == "__main__":
    main()
