#!/usr/bin/env python3
"""Small sequential Eldrim job runner. Python 3.10+, standard library only.

Run one instance (for example under flock). Results and checkpoints are local.
A checkpoint prevents replay of a saved result, not duplicate inference after
an interrupted request. No automatic retries unless --attempts is increased.
"""
import argparse
import hashlib
import json
import os
from pathlib import Path
import random
import re
import sys
import time
import urllib.error
import urllib.request

API = "https://api.eldrim.net/v1/chat/completions"
RETRYABLE = {429, 500, 502, 503, 504}


class NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, req, fp, code, msg, headers, newurl):
        return None  # Never forward the key or request to a redirect target.


OPENER = urllib.request.build_opener(NoRedirect())


def load_jobs(path):
    jobs, seen = [], set()
    for line in path.read_text().splitlines():
        if not line.strip():
            continue
        job = json.loads(line)
        if (not isinstance(job, dict) or not isinstance(job.get("id"), str)
                or not re.fullmatch(r"[A-Za-z0-9_-]{1,80}", job["id"])
                or not isinstance(job.get("prompt"), str) or not job["prompt"].strip()):
            raise ValueError("Each job needs an id (letters, digits, _ or -) and a prompt")
        if job["id"] in seen:
            raise ValueError("Duplicate job id: " + job["id"])
        seen.add(job["id"])
        jobs.append(job)
    return jobs


def infer(payload, key, timeout):
    request = urllib.request.Request(API, data=json.dumps(payload).encode(), headers={
        "Authorization": "Bearer " + key, "Content-Type": "application/json"})
    with OPENER.open(request, timeout=timeout) as response:
        body = json.load(response)
        if not isinstance(body, dict) or not isinstance(body.get("choices"), list) or not body["choices"]:
            raise ValueError("Unexpected response; no checkpoint saved")
        if body["choices"][0].get("finish_reason") != "stop":
            raise ValueError("Incomplete or tool-call response; review output limits; no checkpoint saved")
        return {"response": body, "route": {h: response.headers.get(h) for h in
            ("x-request-id", "x-eldrim-provider", "x-eldrim-residency")}}


def run_job(job, model, key, output, max_tokens, timeout, attempts):
    payload = {"model": model, "messages": [{"role": "user", "content": job["prompt"]}],
               "max_tokens": max_tokens}
    fingerprint = hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest()
    destination = output / (job["id"] + ".json")
    if destination.exists():
        saved = json.loads(destination.read_text())
        if saved.get("input_sha256") != fingerprint:
            raise ValueError("Job changed since its checkpoint; use a new id: " + job["id"])
        return "already saved"
    for attempt in range(attempts):
        try:
            result = infer(payload, key, timeout)
            break
        except urllib.error.HTTPError as exc:
            status = exc.code
            exc.close()
            if status not in RETRYABLE or attempt + 1 == attempts:
                raise RuntimeError(f"HTTP {status}; no checkpoint saved") from None
        except (urllib.error.URLError, TimeoutError, ConnectionError):
            if attempt + 1 == attempts:
                raise RuntimeError("Transport failure; completion unknown; no checkpoint saved") from None
        time.sleep(min(30, 2 ** attempt) + random.random())
    result.update({"id": job["id"], "input_sha256": fingerprint})
    temporary = destination.with_suffix(".tmp")
    with temporary.open("w") as f:
        json.dump(result, f, indent=2)
        f.write("\n")
        f.flush()
        os.fsync(f.fileno())
    os.replace(temporary, destination)
    return "saved"


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path, help="JSONL file with id and prompt per line")
    parser.add_argument("--output-dir", type=Path, default=Path("results"))
    parser.add_argument("--max-tokens", type=int, default=512)
    parser.add_argument("--timeout", type=int, default=180, help="socket timeout in seconds")
    parser.add_argument("--attempts", type=int, default=1, help="attempts per job; retries may be charged again")
    args = parser.parse_args()
    if not 1 <= args.attempts <= 5 or args.timeout < 1 or args.max_tokens < 1:
        parser.error("Use 1–5 attempts, a positive timeout and a positive token limit")
    key, model = os.environ.get("ELDRIM_API_KEY"), os.environ.get("ELDRIM_MODEL")
    if not key or not model:
        parser.error("Set ELDRIM_API_KEY and ELDRIM_MODEL in the environment")
    os.umask(0o077)
    jobs = load_jobs(args.input)
    args.output_dir.mkdir(parents=True, exist_ok=True)
    for job in jobs:
        status = run_job(job, model, key, args.output_dir, args.max_tokens, args.timeout, args.attempts)
        print(job["id"] + ": " + status, flush=True)


if __name__ == "__main__":
    try:
        main()
    except (OSError, ValueError, RuntimeError) as exc:
        print("Job run stopped: " + str(exc), file=sys.stderr)
        sys.exit(1)
