"""Download a paid XpertSystems.ai order. Needs Python 3.8+ and nothing else.

    python xsai_download.py --key xsk_... --order ord_...
    python xsai_download.py --key xsk_... --list          (show your orders)

Files already downloaded are skipped, so running it again finishes an
interrupted download. You can re-download an order as often as you like
until its download window closes.
"""
from __future__ import annotations

import argparse
import json
import os
import sys
import time
import urllib.error
import urllib.request
from concurrent.futures import ThreadPoolExecutor, as_completed

DEFAULT_BASE = "https://data.xpertsystems.ai"


def api(base: str, key: str, path: str) -> dict:
    req = urllib.request.Request(base.rstrip("/") + path,
                                 headers={"Authorization": f"Bearer {key}",
                                          "User-Agent": "xsai-download/1.0"})
    try:
        with urllib.request.urlopen(req, timeout=60) as r:
            return json.loads(r.read())
    except urllib.error.HTTPError as e:
        try:
            detail = json.loads(e.read()).get("detail")
        except Exception:
            detail = e.reason
        sys.exit(f"Error {e.code}: {detail}")


def fetch(url: str, dest: str, size: int) -> str:
    if os.path.exists(dest) and (size == 0 or os.path.getsize(dest) == size):
        return "skipped"
    tmp = dest + ".part"
    for attempt in range(4):
        try:
            with urllib.request.urlopen(url, timeout=120) as r, open(tmp, "wb") as f:
                while True:
                    chunk = r.read(1 << 16)
                    if not chunk:
                        break
                    f.write(chunk)
            os.replace(tmp, dest)
            return "ok"
        except Exception as e:                      # noqa: BLE001
            if attempt == 3:
                return f"failed: {e}"
            time.sleep(2 ** attempt)
    return "failed"


def main() -> int:
    ap = argparse.ArgumentParser(description="Download an XpertSystems.ai order")
    ap.add_argument("--key", default=os.environ.get("XSAI_API_KEY"),
                    help="your API key (or set XSAI_API_KEY)")
    ap.add_argument("--order", help="order id, e.g. ord_1a2b3c4d")
    ap.add_argument("--list", action="store_true", help="list your orders")
    ap.add_argument("--out", help="folder to save into (default: the order id)")
    ap.add_argument("--base", default=os.environ.get("XSAI_BASE_URL", DEFAULT_BASE))
    ap.add_argument("--workers", type=int, default=8)
    a = ap.parse_args()

    if not a.key:
        sys.exit("Give your API key with --key or set XSAI_API_KEY.")

    if a.list or not a.order:
        orders = api(a.base, a.key, "/v1/orders")["orders"]
        if not orders:
            print("No orders yet.")
        for o in orders:
            print(f"{o['order_id']}  {o['dataset_name']:<20} {o['quantity']:>6,} files  "
                  f"{o['status']:<8} until {o['download_until'] or '-'}")
        return 0

    info = api(a.base, a.key, f"/v1/orders/{a.order}")
    if info["status"] != "paid":
        sys.exit(f"Order {a.order} is {info['status']}, nothing to download.")
    out = a.out or a.order
    os.makedirs(out, exist_ok=True)
    print(f"{info['quantity']:,} {info['dataset_name']} files -> {os.path.abspath(out)}")

    done = skipped = 0
    failed: list[str] = []
    offset = 0
    while offset is not None:
        page = api(a.base, a.key, f"/v1/orders/{a.order}/files?offset={offset}&limit=500")
        with ThreadPoolExecutor(max_workers=max(1, a.workers)) as pool:
            jobs = {pool.submit(fetch, f["url"], os.path.join(out, f["name"]), f["size"]):
                    f["name"] for f in page["files"]}
            for j in as_completed(jobs):
                r = j.result()
                if r == "ok":
                    done += 1
                elif r == "skipped":
                    skipped += 1
                else:
                    failed.append(f"{jobs[j]}: {r}")
        print(f"  {offset + page['count']:,} / {page['total']:,}")
        offset = page["next_offset"]

    print(f"Downloaded {done:,}, already had {skipped:,}, failed {len(failed):,}.")
    for f in failed[:20]:
        print("  " + f)
    if failed:
        print("Run the same command again to retry the failed files.")
    return 1 if failed else 0


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