"""Windway public samples v1.0.0. Python 3.10+, standard library only."""
import hashlib
import io
import json
import time
import urllib.request
import zipfile

RELEASES = {
    "nav": ("windway-gamenav-critical-edits-sample",
            "574a219c1e60ba065eda256b0652fe058fb07c65d1c750c77d288cb61b15b89e"),
    "coord": ("windway-gamecoord-joint-plans-sample",
              "98f29aa51b0674e183bed932a8e04adb4e47f6c882934eb38719c66f57331b56"),
}


def load_sample(kind="nav"):
    slug, expected_hash = RELEASES[kind]
    url = f"https://windwaydata.com/downloads/{slug}/1.0.0/{slug}-1.0.0.zip"
    with urllib.request.urlopen(url, timeout=60) as response:
        payload = response.read()
    actual_hash = hashlib.sha256(payload).hexdigest()
    if actual_hash != expected_hash:
        raise ValueError("Archive SHA-256 mismatch; refusing to read data")
    with zipfile.ZipFile(io.BytesIO(payload)) as archive:
        prefix = slug + "/"
        checksums = archive.read(prefix + "CHECKSUMS.sha256").decode().splitlines()
        for line in checksums:
            digest, filename = line.split(maxsplit=1)
            if hashlib.sha256(archive.read(prefix + filename.strip())).hexdigest() != digest:
                raise ValueError("Payload checksum mismatch: " + filename)
        schema = json.loads(archive.read(prefix + "SCHEMA.json"))
        splits = {}
        for split in ("train", "validation", "test"):
            splits[split] = [json.loads(line) for line in
                             archive.read(prefix + f"data/{split}.jsonl").decode().splitlines()]
    records = [record for split in splits.values() for record in split]
    assert len(records) == len({record["id"] for record in records}) == 100
    return records, schema, splits, actual_hash


if __name__ == "__main__":
    started = time.perf_counter()
    for kind in RELEASES:
        records, schema, splits, digest = load_sample(kind)
        print(kind, "records:", len(records), "SHA-256:", digest)
        print("splits:", {name: len(rows) for name, rows in splits.items()})
        print("first record:", json.dumps(records[0], ensure_ascii=False))
        print("schema fields:", sorted(schema["properties"]))
        families = sorted({record["family"] for record in records})
        print("families:", families)
    print("Seconds to verified data:", round(time.perf_counter() - started, 2))
