#!/usr/bin/env python3
"""Validate every published Role 2 preprocessing shard against competition metadata.

This is a gate for Role 3. It verifies exact study/series coverage against the
competition's train_series.csv, validates manifest schema, reconciles every HDF5
path with its manifest, and fully reads deterministic HDF5 samples to exercise
compression/decompression and validate array metadata.
"""

import ast
import json
import os
import sys
from pathlib import Path

import h5py
import numpy as np
import pandas as pd

VALID_PLANES = {"sagittal", "coronal", "axial", "unknown"}
VALID_SEQUENCES = {"T1", "T2", "PD", "other"}
EXPECTED_SHARDS = set(range(24))
HDF5_SAMPLE_FRACTION = 0.10
HDF5_MIN_SAMPLE = 5
HDF5_MAX_SAMPLE = 50


def find_shard_outputs():
    """Find Kaggle notebook-output mounts for all Role 2 shards."""
    mounted = {}
    input_root = Path("/kaggle/input")
    for shard_dir in input_root.glob("**/rsna-role2-shard-*/shards"):
        if not shard_dir.is_dir():
            continue
        try:
            index = int(shard_dir.parent.name.rsplit("-", 1)[1])
        except (IndexError, ValueError):
            continue
        if index in mounted:
            raise RuntimeError(f"Shard {index:02d} was mounted more than once")
        mounted[index] = shard_dir
        print(f"Found shard {index:02d}: {shard_dir}", flush=True)
    return mounted


def load_train_series():
    paths = [
        Path("/kaggle/input/competitions/rsna-knee-abnormality-detection/train_series.csv"),
        Path("/kaggle/input/rsna-knee-abnormality-detection/train_series.csv"),
    ]
    csv_path = next((p for p in paths if p.exists()), None)
    if csv_path is None:
        raise FileNotFoundError("Could not locate competition train_series.csv")
    frame = pd.read_csv(csv_path, dtype={"StudyInstanceUID": str, "SeriesInstanceUID": str})
    required = {"StudyInstanceUID", "SeriesInstanceUID"}
    missing = required - set(frame.columns)
    if missing:
        raise ValueError(f"train_series.csv lacks required columns: {sorted(missing)}")
    if frame.SeriesInstanceUID.duplicated().any():
        raise ValueError("train_series.csv contains duplicate SeriesInstanceUID values")
    print(f"Loaded authoritative train_series.csv: {len(frame)} series, "
          f"{frame.StudyInstanceUID.nunique()} studies", flush=True)
    return frame


def read_csv_or_empty(path, columns):
    if not path.exists() or path.stat().st_size == 0:
        return pd.DataFrame(columns=columns)
    return pd.read_csv(path, dtype={"StudyInstanceUID": str, "SeriesInstanceUID": str})


def load_sidecars(shards):
    manifests, failures, done_by_shard = [], [], {}
    for index, shard_dir in sorted(shards.items()):
        manifest_path = shard_dir / f"shard_{index:02d}_manifest.csv"
        done_path = shard_dir / f"shard_{index:02d}_done_studies.txt"
        failure_path = shard_dir / "failures.csv"

        manifest = read_csv_or_empty(manifest_path, [])
        if manifest.empty:
            raise FileNotFoundError(f"Shard {index:02d}: missing or empty manifest: {manifest_path}")
        manifest["shard"] = index
        manifests.append(manifest)

        failure = read_csv_or_empty(
            failure_path,
            ["timestamp", "StudyInstanceUID", "SeriesInstanceUID", "expected_by_csv", "error"],
        )
        if not failure.empty:
            failure["shard"] = index
            failures.append(failure)

        if not done_path.exists():
            raise FileNotFoundError(f"Shard {index:02d}: missing done-study file: {done_path}")
        study_ids = [line.strip() for line in done_path.read_text().splitlines() if line.strip()]
        duplicates = len(study_ids) - len(set(study_ids))
        if duplicates:
            raise ValueError(f"Shard {index:02d}: {duplicates} duplicated done-study entries")
        done_by_shard[index] = set(study_ids)
        print(f"Shard {index:02d}: manifest={len(manifest)}, failures={len(failure)}, "
              f"done_studies={len(study_ids)}", flush=True)

    return (
        pd.concat(manifests, ignore_index=True),
        pd.concat(failures, ignore_index=True) if failures else pd.DataFrame(),
        done_by_shard,
    )


def schema_validation(manifest):
    required = {
        "StudyInstanceUID", "SeriesInstanceUID", "plane", "sequence_type",
        "n_slices", "array_shape", "dtype", "slice_order_method",
    }
    missing_columns = sorted(required - set(manifest.columns))
    invalid_planes = sorted(set(manifest.get("plane", [])) - VALID_PLANES)
    invalid_sequences = sorted(set(manifest.get("sequence_type", [])) - VALID_SEQUENCES)
    duplicate_series = int(manifest.SeriesInstanceUID.duplicated().sum())
    blank_studies = int(manifest.StudyInstanceUID.isna().sum() + (manifest.StudyInstanceUID == "").sum())
    blank_series = int(manifest.SeriesInstanceUID.isna().sum() + (manifest.SeriesInstanceUID == "").sum())
    return {
        "missing_columns": missing_columns,
        "invalid_planes": invalid_planes,
        "invalid_sequences": invalid_sequences,
        "duplicate_series": duplicate_series,
        "blank_studies": blank_studies,
        "blank_series": blank_series,
    }


def coverage_validation(train, manifest, failures, done_by_shard):
    expected_studies = set(train.StudyInstanceUID)
    expected_series = set(train.SeriesInstanceUID)
    written_studies = set(manifest.StudyInstanceUID)
    written_series = set(manifest.SeriesInstanceUID)
    done_studies = set().union(*done_by_shard.values())

    failed_expected_series = set()
    if not failures.empty and "expected_by_csv" in failures.columns:
        expected_rows = failures[failures.expected_by_csv.astype(str).isin({"1", "1.0", "True", "true"})]
        failed_expected_series = set(expected_rows.SeriesInstanceUID.dropna().astype(str))

    dispositioned_series = written_series | failed_expected_series
    return {
        "expected_studies": len(expected_studies),
        "expected_series": len(expected_series),
        "written_studies": len(written_studies),
        "written_series": len(written_series),
        "done_studies": len(done_studies),
        "failed_expected_series": len(failed_expected_series),
        "unaccounted_studies": sorted(expected_studies - done_studies),
        "extra_done_studies": sorted(done_studies - expected_studies),
        "unaccounted_series": sorted(expected_series - dispositioned_series),
        "extra_written_series": sorted(written_series - expected_series),
        "failed_series_preview": sorted(failed_expected_series)[:20],
    }


def hdf5_validation(shards, manifest):
    result = {
        "h5_files": 0,
        "h5_series": 0,
        "manifest_series_missing_in_h5": [],
        "h5_series_missing_in_manifest": [],
        "datasets_sampled_and_read": 0,
        "dtype_counts": {},
        "shape_errors": [],
        "n_slices_errors": [],
        "read_errors": [],
    }

    for index, shard_dir in sorted(shards.items()):
        h5_path = shard_dir / f"shard_{index:02d}.h5"
        if not h5_path.exists():
            result["read_errors"].append(f"Shard {index:02d}: missing HDF5 file")
            continue

        shard_manifest = manifest[manifest.shard == index].copy().reset_index(drop=True)
        manifest_keys = set(
            shard_manifest.StudyInstanceUID.astype(str) + "/" + shard_manifest.SeriesInstanceUID.astype(str)
        )
        h5_keys = set()
        result["h5_files"] += 1

        print(f"Shard {index:02d}: enumerate and reconcile HDF5 paths...", flush=True)
        try:
            with h5py.File(h5_path, "r") as file:
                def visit(name, obj):
                    if isinstance(obj, h5py.Dataset):
                        h5_keys.add(name)
                file.visititems(visit)

                result["h5_series"] += len(h5_keys)
                result["manifest_series_missing_in_h5"].extend(
                    f"{index:02d}:{key}" for key in sorted(manifest_keys - h5_keys)
                )
                result["h5_series_missing_in_manifest"].extend(
                    f"{index:02d}:{key}" for key in sorted(h5_keys - manifest_keys)
                )

                sample_count = max(HDF5_MIN_SAMPLE, int(np.ceil(len(shard_manifest) * HDF5_SAMPLE_FRACTION)))
                sample_count = min(HDF5_MAX_SAMPLE, sample_count, len(shard_manifest))
                sample = shard_manifest.iloc[np.linspace(0, len(shard_manifest) - 1, sample_count, dtype=int)]
                print(f"Shard {index:02d}: fully reading {len(sample)}/{len(shard_manifest)} deterministic samples...", flush=True)
                for _, row in sample.iterrows():
                    key = f"{row.StudyInstanceUID}/{row.SeriesInstanceUID}"
                    if key not in file:
                        continue
                    dataset = file[key]
                    try:
                        values = dataset[()]  # forces gzip decompression / checksum validation
                    except Exception as exc:
                        result["read_errors"].append(f"Shard {index:02d}:{key}: {type(exc).__name__}: {exc}")
                        continue
                    result["datasets_sampled_and_read"] += 1
                    dtype = str(values.dtype)
                    result["dtype_counts"][dtype] = result["dtype_counts"].get(dtype, 0) + 1
                    try:
                        expected_shape = tuple(ast.literal_eval(row.array_shape))
                    except (ValueError, SyntaxError, TypeError):
                        result["shape_errors"].append(f"Shard {index:02d}:{key}: invalid manifest array_shape={row.array_shape!r}")
                        continue
                    if values.shape != expected_shape or values.ndim != 3:
                        result["shape_errors"].append(
                            f"Shard {index:02d}:{key}: manifest={expected_shape}, h5={values.shape}, ndim={values.ndim}"
                        )
                    if int(row.n_slices) != values.shape[0]:
                        result["n_slices_errors"].append(
                            f"Shard {index:02d}:{key}: manifest={row.n_slices}, h5={values.shape[0]}"
                        )
        except Exception as exc:
            result["read_errors"].append(f"Shard {index:02d}: {type(exc).__name__}: {exc}")
    return result


def main():
    print("=" * 78)
    print("RSNA ROLE 2 — ALL-SHARD VALIDATION GATE")
    print("=" * 78)

    shards = find_shard_outputs()
    found = set(shards)
    if found != EXPECTED_SHARDS:
        print(f"FATAL: expected shards {sorted(EXPECTED_SHARDS)}, found {sorted(found)}")
        return 2

    train = load_train_series()
    manifest, failures, done_by_shard = load_sidecars(shards)
    schema = schema_validation(manifest)
    coverage = coverage_validation(train, manifest, failures, done_by_shard)
    hdf5 = hdf5_validation(shards, manifest)

    blockers = []
    warnings = []
    if schema["missing_columns"]:
        blockers.append(f"missing manifest columns: {schema['missing_columns']}")
    if schema["invalid_planes"]:
        blockers.append(f"invalid plane labels: {schema['invalid_planes']}")
    if schema["invalid_sequences"]:
        blockers.append(f"invalid sequence labels: {schema['invalid_sequences']}")
    for field in ("duplicate_series", "blank_studies", "blank_series"):
        if schema[field]:
            blockers.append(f"{field}={schema[field]}")
    if coverage["unaccounted_studies"]:
        blockers.append(f"unaccounted studies={len(coverage['unaccounted_studies'])}")
    if coverage["extra_done_studies"]:
        blockers.append(f"extra done studies={len(coverage['extra_done_studies'])}")
    if coverage["unaccounted_series"]:
        blockers.append(f"unaccounted expected series={len(coverage['unaccounted_series'])}")
    if coverage["extra_written_series"]:
        blockers.append(f"extra written series={len(coverage['extra_written_series'])}")
    for field in ("manifest_series_missing_in_h5", "h5_series_missing_in_manifest", "shape_errors", "n_slices_errors", "read_errors"):
        if hdf5[field]:
            blockers.append(f"{field}={len(hdf5[field])}")
    if set(hdf5["dtype_counts"]) != {"uint16"}:
        blockers.append(f"sampled dtype distribution is not exclusively uint16: {hdf5['dtype_counts']}")
    if coverage["failed_expected_series"]:
        warnings.append(f"Expected series recorded as failed={coverage['failed_expected_series']}")

    report = {
        "ready_for_role3": not blockers and not warnings,
        "coverage": coverage,
        "schema": schema,
        "hdf5": hdf5,
        "failure_rows": int(len(failures)),
        "blockers": blockers,
        "warnings": warnings,
    }
    os.makedirs("/kaggle/working", exist_ok=True)
    with open("/kaggle/working/validation_summary.json", "w") as out:
        json.dump(report, out, indent=2)
    manifest.to_csv("/kaggle/working/series_manifest_full.csv", index=False)

    print("\n" + "=" * 78)
    print("VALIDATION RESULT")
    print("=" * 78)
    print(f"Expected: {coverage['expected_studies']} studies / {coverage['expected_series']} series")
    print(f"Written:  {coverage['written_studies']} studies / {coverage['written_series']} series")
    print(f"Done:     {coverage['done_studies']} studies")
    print(f"HDF5:     {hdf5['h5_files']} files / {hdf5['h5_series']} datasets")
    print(f"Sampled full reads: {hdf5['datasets_sampled_and_read']} datasets; dtypes={hdf5['dtype_counts']}")
    print(f"Failure rows: {len(failures)}")
    if blockers:
        print("ROLE 3 BLOCKED")
        for item in blockers:
            print(f"  BLOCKER: {item}")
    elif warnings:
        print("ROLE 3 NOT YET APPROVED — review warnings")
        for item in warnings:
            print(f"  WARNING: {item}")
    else:
        print("ROLE 3 APPROVED — all output-validation gates passed")
    print("Summary: /kaggle/working/validation_summary.json")
    print("Merged manifest: /kaggle/working/series_manifest_full.csv")
    return 0 if report["ready_for_role3"] else 1


if __name__ == "__main__":
    sys.exit(main())
