"""Command-line interface for the control-anchored BRCA LGA screen."""

from __future__ import annotations

import argparse
import csv
import hashlib
import json
import sys
from datetime import datetime, timezone
from pathlib import Path

import numpy as np

from .clinical_depth import (
    CLAIMS_BOUNDARY,
    AnchoredDepthModel,
    DepthTarget,
)
from .provenance import atomic_write_json


def _sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for block in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(block)
    return digest.hexdigest()


def read_targets(path: str | Path) -> tuple[DepthTarget, ...]:
    source = Path(path)
    with source.open(newline="", encoding="utf-8-sig") as handle:
        reader = csv.DictReader(handle, delimiter="\t")
        required = {"target_id", "gene", "exon", "order", "reportable"}
        missing = required - set(reader.fieldnames or ())
        if missing:
            raise ValueError(
                "target manifest is missing column(s): "
                + ", ".join(sorted(missing))
            )
        rows = []
        for line_number, row in enumerate(reader, start=2):
            raw_reportable = row["reportable"].strip().lower()
            if raw_reportable not in {"0", "1", "false", "true"}:
                raise ValueError(
                    f"invalid reportable value on line {line_number}"
                )
            rows.append(
                DepthTarget(
                    target_id=row["target_id"].strip(),
                    gene=row["gene"].strip(),
                    exon=row["exon"].strip(),
                    order=int(row["order"]),
                    reportable=raw_reportable in {"1", "true"},
                )
            )
    if not rows:
        raise ValueError("target manifest is empty")
    return tuple(rows)


def read_counts(
    path: str | Path,
    targets: tuple[DepthTarget, ...],
) -> tuple[tuple[str, ...], np.ndarray]:
    source = Path(path)
    with source.open(newline="", encoding="utf-8-sig") as handle:
        reader = csv.reader(handle, delimiter="\t")
        try:
            header = next(reader)
        except StopIteration as exc:
            raise ValueError("count matrix is empty") from exc
        if not header or header[0].strip().lower() != "sample":
            raise ValueError("first count-matrix column must be 'sample'")
        observed = [value.strip() for value in header[1:]]
        expected = [target.target_id for target in targets]
        if len(observed) != len(set(observed)):
            raise ValueError("count-matrix target columns must be unique")
        missing = sorted(set(expected) - set(observed))
        extra = sorted(set(observed) - set(expected))
        if missing or extra:
            raise ValueError(
                f"count-matrix targets differ; missing={missing}, extra={extra}"
            )
        column_order = [observed.index(target_id) + 1 for target_id in expected]
        samples: list[str] = []
        values: list[list[float]] = []
        for line_number, row in enumerate(reader, start=2):
            if len(row) != len(header):
                raise ValueError(
                    f"count-matrix line {line_number} has {len(row)} "
                    f"columns; expected {len(header)}"
                )
            sample = row[0].strip()
            if not sample or sample in samples:
                raise ValueError(
                    f"sample IDs must be non-empty and unique "
                    f"(line {line_number})"
                )
            samples.append(sample)
            values.append([float(row[index]) for index in column_order])
    matrix = np.asarray(values, dtype=float)
    if matrix.ndim != 2 or matrix.shape[0] == 0:
        raise ValueError("count matrix has no sample rows")
    if not np.all(np.isfinite(matrix)) or np.any(matrix < 0):
        raise ValueError("counts must be finite and nonnegative")
    return tuple(samples), matrix


def cmd_fit(args: argparse.Namespace) -> int:
    targets = read_targets(args.targets)
    fit_samples, fit_counts = read_counts(args.fit_counts, targets)
    calibration_samples, calibration_counts = read_counts(
        args.calibration_counts, targets
    )
    overlap = sorted(set(fit_samples).intersection(calibration_samples))
    if overlap:
        preview = ", ".join(overlap[:5])
        raise ValueError(
            "fit and calibration sample IDs overlap"
            + (f": {preview}" if preview else "")
        )
    model = AnchoredDepthModel.fit(
        fit_counts,
        calibration_counts,
        targets,
        alpha=args.alpha,
        n_components=args.components,
        min_median_control_depth=args.min_control_depth,
        min_target_median_depth=args.min_target_depth,
        max_target_sigma=args.max_target_sigma,
        max_window_targets=args.max_window_targets,
        fit_sample_ids=fit_samples,
        calibration_sample_ids=calibration_samples,
    )
    model.save(args.out)
    print(
        json.dumps(
            {
                "model": str(Path(args.out).resolve()),
                "n_fit_normals": model.n_fit,
                "n_calibration_normals": model.n_calibration,
                "n_components": model.n_components,
                "callable_reportable_targets": int(
                    model.callable_mask.sum()
                ),
                "fit_sample_ids_sha256":
                    model.fit_sample_ids_sha256,
                "calibration_sample_ids_sha256":
                    model.calibration_sample_ids_sha256,
                "claims_boundary": CLAIMS_BOUNDARY,
            },
            indent=2,
        )
    )
    return 0


def cmd_screen(args: argparse.Namespace) -> int:
    model_path = Path(args.model)
    model = AnchoredDepthModel.load(model_path)
    samples, counts = read_counts(args.counts, model.targets)
    results = [
        model.screen(row, sample=sample).to_json()
        for sample, row in zip(samples, counts)
    ]
    payload = {
        "schema": "lgasieve.anchored-depth-screen-results.v1",
        "created_utc": datetime.now(timezone.utc)
        .isoformat()
        .replace("+00:00", "Z"),
        "model": {
            "path": str(model_path.resolve()),
            "sha256": _sha256(model_path),
        },
        "input_counts": {
            "path": str(Path(args.counts).resolve()),
            "sha256": _sha256(Path(args.counts)),
        },
        "claims_boundary": CLAIMS_BOUNDARY,
        "results": results,
    }
    atomic_write_json(str(args.out), payload)
    counts_by_status: dict[str, int] = {}
    for result in results:
        status = result["status"]
        counts_by_status[status] = counts_by_status.get(status, 0) + 1
    print(
        json.dumps(
            {
                "output": str(Path(args.out).resolve()),
                "samples": len(results),
                "status_counts": counts_by_status,
            },
            indent=2,
        )
    )
    return 0


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        prog="lgasieve-depth",
        description=(
            "Fit or run the broad-panel, control-anchored BRCA LGA "
            "screening model."
        ),
    )
    subparsers = parser.add_subparsers(dest="command", required=True)

    fit = subparsers.add_parser(
        "fit", help="fit and independently calibrate a site-specific model"
    )
    fit.add_argument("--targets", required=True)
    fit.add_argument("--fit-counts", required=True)
    fit.add_argument("--calibration-counts", required=True)
    fit.add_argument("--out", required=True)
    fit.add_argument("--alpha", type=float, default=0.01)
    fit.add_argument("--components", type=int, default=12)
    fit.add_argument("--min-control-depth", type=float, default=100.0)
    fit.add_argument("--min-target-depth", type=float, default=100.0)
    fit.add_argument("--max-target-sigma", type=float, default=0.13)
    fit.add_argument("--max-window-targets", type=int, default=60)
    fit.set_defaults(func=cmd_fit)

    screen = subparsers.add_parser(
        "screen", help="screen samples and emit reflex/no-call results"
    )
    screen.add_argument("--model", required=True)
    screen.add_argument("--counts", required=True)
    screen.add_argument("--out", required=True)
    screen.set_defaults(func=cmd_screen)
    return parser


def main(argv: list[str] | None = None) -> int:
    parser = build_parser()
    args = parser.parse_args(argv)
    try:
        return int(args.func(args))
    except (OSError, ValueError, json.JSONDecodeError) as exc:
        print(f"[lgasieve-depth] {exc}", file=sys.stderr)
        return 2


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