"""Resume a single-worker text batch using SQLite. Python 3.12+."""
import argparse
import hashlib
import json
from pathlib import Path
import sqlite3
import sys

WORKER_VERSION = "text-length-v1"


def load_jobs(path):
    jobs = json.loads(Path(path).read_text(encoding="utf-8"))
    if not isinstance(jobs, list):
        raise ValueError("input must be a JSON list")
    seen = set()
    for job in jobs:
        if not isinstance(job, dict) or set(job) != {"id", "text"}:
            raise ValueError("each job needs exactly id and text")
        job_id, text = job["id"], job["text"]
        if not isinstance(job_id, str) or not job_id.strip():
            raise ValueError("job id must be a nonempty string")
        if not isinstance(text, str):
            raise ValueError("job text must be a string")
        if job_id in seen:
            raise ValueError(f"duplicate job id: {job_id}")
        seen.add(job_id)
    return jobs


def fingerprint(text):
    payload = json.dumps([WORKER_VERSION, text], ensure_ascii=True)
    return hashlib.sha256(payload.encode("utf-8")).hexdigest()


def calculate(text):
    return {"length": len(text)}


def run_batch(jobs, database, *, stop_after=None):
    completed = skipped = 0
    con = sqlite3.connect(database, autocommit=False)
    try:
        with con:
            con.execute("""CREATE TABLE IF NOT EXISTS results (
                job_id TEXT PRIMARY KEY,
                fingerprint TEXT NOT NULL,
                result_json TEXT NOT NULL
            )""")
        # Check every current input before calculating any new result.
        for job in jobs:
            old = con.execute(
                "SELECT fingerprint FROM results WHERE job_id = ?",
                (job["id"],),
            ).fetchone()
            if old is not None and old[0] != fingerprint(job["text"]):
                raise ValueError(f"input or worker changed for {job['id']}; use a new database")
        for job in jobs:
            old = con.execute(
                "SELECT 1 FROM results WHERE job_id = ?", (job["id"],)
            ).fetchone()
            if old is not None:
                skipped += 1
                print(f"SKIP {job['id']}", flush=True)
                continue
            result = calculate(job["text"])
            encoded = json.dumps(result, allow_nan=False, sort_keys=True)
            # The result row itself is the completion record.
            with con:
                con.execute(
                    "INSERT INTO results VALUES (?, ?, ?)",
                    (job["id"], fingerprint(job["text"]), encoded),
                )
            completed += 1
            print(f"DONE {job['id']}", flush=True)
            if stop_after is not None and completed >= stop_after:
                print(f"STOPPED: new={completed}; skipped={skipped}", flush=True)
                return
        print(f"FINISHED: new={completed}; skipped={skipped}", flush=True)
    finally:
        con.close()


def positive_int(value):
    number = int(value)
    if number < 1:
        raise argparse.ArgumentTypeError("must be at least 1")
    return number


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path, help="JSON list of id/text jobs")
    parser.add_argument("database", type=Path, help="SQLite checkpoint; parent must exist")
    parser.add_argument("--stop-after", type=positive_int, help="pause after N new results")
    args = parser.parse_args()
    try:
        jobs = load_jobs(args.input)
        run_batch(jobs, args.database, stop_after=args.stop_after)
    except (OSError, ValueError, sqlite3.Error) as exc:
        print(f"ERROR: {exc}", file=sys.stderr)
        return 1
    except KeyboardInterrupt:
        print("INTERRUPTED: rerun with the same input and database", file=sys.stderr)
        return 130
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
