"""Check a pooled protein embedding archive against an ordered ID manifest."""
import argparse
import json
from pathlib import Path
import sys
from zipfile import BadZipFile

import numpy as np


def check_ids(ids, source):
    if not ids:
        raise ValueError(f"{source}: no sample IDs")
    seen = set()
    for row, value in enumerate(ids):
        if not value or any(char.isspace() for char in value):
            raise ValueError(f"{source}: empty ID or whitespace at row {row}")
        if value in seen:
            raise ValueError(f"{source}: duplicate ID {value!r}")
        seen.add(value)


def audit(path, ids_path, expected_dim):
    if expected_dim < 1:
        raise ValueError("expected dimension must be positive")
    expected = Path(ids_path).read_text(encoding="utf-8-sig").splitlines()
    check_ids(expected, "manifest")
    archive = np.load(path, allow_pickle=False)
    if not isinstance(archive, np.lib.npyio.NpzFile):
        raise ValueError("expected an NPZ archive")
    with archive:
        if len(archive.files) != 2 or set(archive.files) != {"sample_id", "embedding"}:
            raise ValueError("archive must contain only sample_id and embedding")
        ids = archive["sample_id"]
        vectors = archive["embedding"]
    if ids.ndim != 1 or ids.dtype.kind != "U":
        raise ValueError("sample_id must be a one-dimensional Unicode array")
    names = ids.tolist()
    check_ids(names, "archive")
    if vectors.ndim != 2 or vectors.shape != (len(names), expected_dim):
        raise ValueError(f"expected embedding shape ({len(names)}, {expected_dim})")
    if vectors.dtype.kind != "f":
        raise ValueError("embedding must have a floating-point dtype")
    if names != expected:
        if set(names) == set(expected):
            raise ValueError("ID order differs; join by sample_id before export")
        raise ValueError("archive IDs differ from manifest IDs")
    zero_rows = 0
    # Arrays are loaded into RAM; chunking limits temporary boolean arrays.
    for start in range(0, len(names), 1024):
        block = vectors[start:start + 1024]
        finite = np.isfinite(block)
        if not finite.all():
            row, column = np.argwhere(~finite)[0]
            raise ValueError(f"non-finite value at row {start + int(row)}, column {int(column)}")
        zero_rows += int(np.all(block == 0, axis=1).sum())
    return {"status": "ok", "samples": len(names), "dimensions": expected_dim,
            "dtype": str(vectors.dtype), "zero_rows": zero_rows}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("archive", type=Path)
    parser.add_argument("ids", type=Path, help="one ID per line, in downstream label order")
    parser.add_argument("--dim", type=int, required=True, help="expected pooled feature count")
    args = parser.parse_args()
    try:
        result = audit(args.archive, args.ids, args.dim)
    except (OSError, ValueError, EOFError, BadZipFile) as error:
        print(f"ERROR: {error}", file=sys.stderr)
        return 1
    print(json.dumps(result, ensure_ascii=False))
    return 0


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