#!/usr/bin/env python3
"""
Process the o-ring transfer user study CSV and print summary tables.

This script mirrors the paired-comparison output style from
`oring_transfer_process_example.py`, but it reads the transfer CSV directly and
prints results for:
- Overall
- Novices
- Surgeons

Assumptions for `source_data/oring_transfer_data.csv`:
- Surgeons are identified by `Notes == "Surgeon"`
- `Num Trials` is fixed to the cleaned six-trial basis
- Total occurrence counts are stored as `Error <k> Total`
- WeightedErrorPerTrial = sum_k weight_k * error_k_total / Num Trials
- TimePerTrialSec = `Total Time (s)` / `Num Trials`
"""

import argparse
import csv
import math
from pathlib import Path
from typing import Dict, Iterable, List, Sequence, Tuple

import numpy as np
from scipy.stats import shapiro, t as t_dist, ttest_rel, wilcoxon


DEFAULT_WEIGHTS = {"1": 2, "2": 2, "3": 4, "4": 5, "5": 3, "6": 3}
DEFAULT_PLATFORM_ORDER = ["dVRK", "Humanoid", "Manual"]
DEFAULT_SUMMARY_PLATFORM_ORDER = ["Manual", "Humanoid", "dVRK"]
DEFAULT_ERROR_OCCURRENCE_ORDER = ["1", "2", "3", "4", "5", "6"]
DEFAULT_PLATFORM_ALIAS = {
    "dvrk": "dVRK",
    "da vinci": "dVRK",
    "davinci": "dVRK",
    "humanoid": "Humanoid",
    "manual": "Manual",
}


def normalize_platform_name(value: str, alias: Dict[str, str]) -> str:
    if not isinstance(value, str):
        return value
    key = value.strip().lower()
    return alias.get(key, value.strip())


def parse_weights(spec: str) -> Dict[str, float]:
    if not spec:
        return dict(DEFAULT_WEIGHTS)

    parsed: Dict[str, float] = {}
    for part in spec.split(","):
        item = part.strip()
        if not item:
            continue
        key, value = item.split(":")
        parsed[key.strip()] = float(value.strip())
    return parsed


def mean_std(values: Sequence[float]) -> Tuple[float, float]:
    arr = np.asarray(values, dtype=float)
    arr = arr[~np.isnan(arr)]
    if len(arr) == 0:
        return np.nan, np.nan
    if len(arr) == 1:
        return float(arr[0]), 0.0
    return float(np.mean(arr)), float(np.std(arr, ddof=1))


def fmt_mean_std(mean: float, std: float, ndigits: int = 2) -> str:
    if np.isnan(mean):
        return "--"
    return f"{mean:.{ndigits}f} +/- {std:.{ndigits}f}"


def fmt_p(value: float, sci_thresh: float = 1e-3) -> str:
    if value is None or (isinstance(value, float) and np.isnan(value)):
        return "--"
    if value < sci_thresh:
        return f"{value:.2e}"
    return f"{value:.3f}"


def fmt_ci(mean_diff: float, ci_low: float, ci_high: float, ndigits: int = 2) -> str:
    if np.isnan(mean_diff) or np.isnan(ci_low) or np.isnan(ci_high):
        return "--"
    return f"{mean_diff:.{ndigits}f} [{ci_low:.{ndigits}f}, {ci_high:.{ndigits}f}]"


def fmt_d(value: float, ndigits: int = 2) -> str:
    if value is None or (isinstance(value, float) and np.isnan(value)):
        return "--"
    return f"{abs(value):.{ndigits}f}"


def parse_float(value: str) -> float:
    if value is None:
        return np.nan
    text = str(value).strip()
    if not text:
        return np.nan
    return float(text)


def load_and_compute_metrics(
    csv_path: str,
    weights: Dict[str, float],
    platform_order: Sequence[str],
) -> List[Dict[str, object]]:
    with open(csv_path, newline="", encoding="utf-8-sig") as handle:
        reader = csv.DictReader(handle)
        rows = list(reader)

    required = ["Participant", "Notes", "Platform", "Num Trials", "Total Time (s)"]
    required += [f"Error {key} Total" for key in DEFAULT_ERROR_OCCURRENCE_ORDER]
    if not rows:
        raise ValueError("The CSV is empty.")

    missing = [name for name in required if name not in rows[0]]
    if missing:
        raise ValueError(f"Missing required CSV columns: {missing}")

    processed: List[Dict[str, object]] = []
    valid_platforms = set(platform_order)

    for row in rows:
        platform = normalize_platform_name(row["Platform"], DEFAULT_PLATFORM_ALIAS)
        if platform not in valid_platforms:
            continue

        trial_count = parse_float(row["Num Trials"])
        if np.isnan(trial_count) or trial_count <= 0:
            raise ValueError(f"Invalid Num Trials value for participant {row['Participant']!r}: {row['Num Trials']!r}")

        weighted_error = 0.0
        for key, weight in weights.items():
            weighted_error += weight * parse_float(row[f"Error {key} Total"]) / trial_count

        notes = str(row.get("Notes", "") or "").strip().lower()

        processed.append(
            {
                "Participant": str(row["Participant"]).strip(),
                "Platform": platform,
                "IsSurgeon": notes == "surgeon",
                "WeightedErrorPerTrial": weighted_error,
                "TimePerTrialSec": parse_float(row["Total Time (s)"]) / trial_count,
                **{
                    f"Error {key} Total": parse_float(row[f"Error {key} Total"])
                    for key in DEFAULT_ERROR_OCCURRENCE_ORDER
                },
            }
        )

    return processed


def build_participant_platform_map(
    rows: Sequence[Dict[str, object]],
    value_key: str,
) -> Dict[str, Dict[str, float]]:
    participant_platform_values: Dict[str, Dict[str, float]] = {}
    for row in rows:
        participant = str(row["Participant"])
        platform = str(row["Platform"])
        participant_platform_values.setdefault(participant, {})[platform] = float(row[value_key])
    return participant_platform_values


def paired_stats(
    participant_platform_values: Dict[str, Dict[str, float]],
    platform_a: str,
    platform_b: str,
    alpha: float = 0.05,
) -> Dict[str, float]:
    paired_a: List[float] = []
    paired_b: List[float] = []

    for platform_values in participant_platform_values.values():
        if platform_a in platform_values and platform_b in platform_values:
            paired_a.append(float(platform_values[platform_a]))
            paired_b.append(float(platform_values[platform_b]))

    n = len(paired_a)
    if n < 2:
        return {
            "n": n,
            "mean_diff": np.nan,
            "ci_low": np.nan,
            "ci_high": np.nan,
            "p_ttest": np.nan,
            "cohens_dz": np.nan,
            "shapiro_p": np.nan,
            "p_wilcoxon": np.nan,
        }

    a = np.asarray(paired_a, dtype=float)
    b = np.asarray(paired_b, dtype=float)
    diffs = a - b

    mean_diff = float(np.mean(diffs))
    sd_diff = float(np.std(diffs, ddof=1))
    se = sd_diff / math.sqrt(n) if sd_diff > 0 else 0.0

    if sd_diff > 0:
        tcrit = float(t_dist.ppf(1 - alpha / 2, df=n - 1))
        ci_low = mean_diff - tcrit * se
        ci_high = mean_diff + tcrit * se
    else:
        ci_low = mean_diff
        ci_high = mean_diff

    if np.allclose(diffs, 0):
        p_ttest = 1.0
        p_wilcoxon = 1.0
    else:
        try:
            _, p_ttest = ttest_rel(a, b, alternative="two-sided")
            p_ttest = float(p_ttest)
        except Exception:
            p_ttest = np.nan

        try:
            p_wilcoxon = float(wilcoxon(diffs, alternative="two-sided", zero_method="pratt").pvalue)
        except Exception:
            p_wilcoxon = np.nan

    shapiro_p = float(shapiro(diffs).pvalue) if n >= 3 else np.nan
    cohens_dz = (mean_diff / sd_diff) if sd_diff > 0 else np.nan

    return {
        "n": n,
        "mean_diff": mean_diff,
        "ci_low": float(ci_low),
        "ci_high": float(ci_high),
        "p_ttest": p_ttest,
        "cohens_dz": float(cohens_dz) if not np.isnan(cohens_dz) else np.nan,
        "shapiro_p": shapiro_p,
        "p_wilcoxon": p_wilcoxon,
    }


def build_group_tables(
    rows: Sequence[Dict[str, object]],
    group_name: str,
    platform_order: Sequence[str],
) -> Tuple[List[Dict[str, object]], List[Dict[str, object]]]:
    err_by_participant = build_participant_platform_map(rows, "WeightedErrorPerTrial")
    time_by_participant = build_participant_platform_map(rows, "TimePerTrialSec")

    primary_rows: List[Dict[str, object]] = []
    diagnostic_rows: List[Dict[str, object]] = []

    for platform in platform_order:
        group_platform_rows = [row for row in rows if row["Platform"] == platform]
        err_values = [float(row["WeightedErrorPerTrial"]) for row in group_platform_rows]
        time_values = [float(row["TimePerTrialSec"]) for row in group_platform_rows]

        mean_err, std_err = mean_std(err_values)
        mean_time, std_time = mean_std(time_values)

        err_vs_manual = None
        time_vs_manual = None
        err_vs_dvrk = None
        time_vs_dvrk = None

        if platform != "Manual":
            err_vs_manual = paired_stats(err_by_participant, platform, "Manual")
            time_vs_manual = paired_stats(time_by_participant, platform, "Manual")

        if platform != "dVRK":
            err_vs_dvrk = paired_stats(err_by_participant, platform, "dVRK")
            time_vs_dvrk = paired_stats(time_by_participant, platform, "dVRK")

        primary_rows.append(
            {
                "Group": group_name,
                "Platform": platform,
                "Weighted Error (mean +/- std)": fmt_mean_std(mean_err, std_err, 2),
                "Delta Err vs Manual (mean [CI])": "--"
                if platform == "Manual"
                else fmt_ci(err_vs_manual["mean_diff"], err_vs_manual["ci_low"], err_vs_manual["ci_high"], 2),
                "p_t (Err vs Manual)": "--" if platform == "Manual" else fmt_p(err_vs_manual["p_ttest"]),
                "d (Err vs Manual)": "--" if platform == "Manual" else fmt_d(err_vs_manual["cohens_dz"], 2),
                "Time (s) (mean +/- std)": fmt_mean_std(mean_time, std_time, 2),
                "Delta Time vs Manual (mean [CI])": "--"
                if platform == "Manual"
                else fmt_ci(time_vs_manual["mean_diff"], time_vs_manual["ci_low"], time_vs_manual["ci_high"], 2),
                "p_t (Time vs Manual)": "--" if platform == "Manual" else fmt_p(time_vs_manual["p_ttest"]),
                "d (Time vs Manual)": "--" if platform == "Manual" else fmt_d(time_vs_manual["cohens_dz"], 2),
                "Delta Err vs dVRK (mean [CI])": "--"
                if platform == "dVRK"
                else fmt_ci(err_vs_dvrk["mean_diff"], err_vs_dvrk["ci_low"], err_vs_dvrk["ci_high"], 2),
                "p_t (Err vs dVRK)": "--" if platform == "dVRK" else fmt_p(err_vs_dvrk["p_ttest"]),
                "d (Err vs dVRK)": "--" if platform == "dVRK" else fmt_d(err_vs_dvrk["cohens_dz"], 2),
                "Delta Time vs dVRK (mean [CI])": "--"
                if platform == "dVRK"
                else fmt_ci(time_vs_dvrk["mean_diff"], time_vs_dvrk["ci_low"], time_vs_dvrk["ci_high"], 2),
                "p_t (Time vs dVRK)": "--" if platform == "dVRK" else fmt_p(time_vs_dvrk["p_ttest"]),
                "d (Time vs dVRK)": "--" if platform == "dVRK" else fmt_d(time_vs_dvrk["cohens_dz"], 2),
            }
        )

        if platform != "Manual":
            diagnostic_rows.append(
                {
                    "Group": group_name,
                    "Comparison": f"{platform} - Manual",
                    "Metric": "WeightedError",
                    "n": err_vs_manual["n"],
                    "Shapiro p": fmt_p(err_vs_manual["shapiro_p"]),
                    "Wilcoxon p": fmt_p(err_vs_manual["p_wilcoxon"]),
                }
            )
            diagnostic_rows.append(
                {
                    "Group": group_name,
                    "Comparison": f"{platform} - Manual",
                    "Metric": "Time",
                    "n": time_vs_manual["n"],
                    "Shapiro p": fmt_p(time_vs_manual["shapiro_p"]),
                    "Wilcoxon p": fmt_p(time_vs_manual["p_wilcoxon"]),
                }
            )

        if platform != "dVRK":
            diagnostic_rows.append(
                {
                    "Group": group_name,
                    "Comparison": f"{platform} - dVRK",
                    "Metric": "WeightedError",
                    "n": err_vs_dvrk["n"],
                    "Shapiro p": fmt_p(err_vs_dvrk["shapiro_p"]),
                    "Wilcoxon p": fmt_p(err_vs_dvrk["p_wilcoxon"]),
                }
            )
            diagnostic_rows.append(
                {
                    "Group": group_name,
                    "Comparison": f"{platform} - dVRK",
                    "Metric": "Time",
                    "n": time_vs_dvrk["n"],
                    "Shapiro p": fmt_p(time_vs_dvrk["shapiro_p"]),
                    "Wilcoxon p": fmt_p(time_vs_dvrk["p_wilcoxon"]),
                }
            )

    return primary_rows, diagnostic_rows


def ordered_summary_platforms(platform_order: Sequence[str]) -> List[str]:
    ordered = [platform for platform in DEFAULT_SUMMARY_PLATFORM_ORDER if platform in platform_order]
    ordered.extend(platform for platform in platform_order if platform not in ordered)
    return ordered


def build_summary_rows(
    primary_rows: Sequence[Dict[str, object]],
    group_name: str,
    platform_order: Sequence[str],
) -> List[Dict[str, object]]:
    rows_by_platform = {str(row["Platform"]): row for row in primary_rows}
    summary_rows: List[Dict[str, object]] = []

    for platform in ordered_summary_platforms(platform_order):
        if platform not in rows_by_platform:
            continue
        row = rows_by_platform[platform]
        summary_rows.append(
            {
                "Group": group_name,
                "Platform": platform,
                "Weighted Error": row["Weighted Error (mean +/- std)"],
                "Time (s)": row["Time (s) (mean +/- std)"],
            }
        )

    return summary_rows


def build_occurrence_summary_rows(
    rows: Sequence[Dict[str, object]],
    group_name: str,
    platform_order: Sequence[str],
) -> List[Dict[str, object]]:
    occurrence_rows: List[Dict[str, object]] = []
    for platform in ordered_summary_platforms(platform_order):
        platform_rows = [row for row in rows if row["Platform"] == platform]
        summary_row: Dict[str, object] = {"Group": group_name, "Platform": platform}
        for key in DEFAULT_ERROR_OCCURRENCE_ORDER:
            column = f"Error {key} Total"
            values = [float(row[column]) for row in platform_rows]
            mean_value, std_value = mean_std(values)
            summary_row[f"Error {key}"] = fmt_mean_std(mean_value, std_value, 2)
        occurrence_rows.append(summary_row)
    return occurrence_rows


def choose_test(stats: Dict[str, float], normality_alpha: float) -> Tuple[str, str]:
    shapiro_p = stats["shapiro_p"]
    if not np.isnan(shapiro_p) and shapiro_p < normality_alpha:
        return "w", fmt_p(stats["p_wilcoxon"])
    return "t", fmt_p(stats["p_ttest"])


def build_paired_comparison_rows(
    rows: Sequence[Dict[str, object]],
    platform_order: Sequence[str],
    normality_alpha: float,
) -> List[Dict[str, object]]:
    err_by_participant = build_participant_platform_map(rows, "WeightedErrorPerTrial")
    time_by_participant = build_participant_platform_map(rows, "TimePerTrialSec")

    platform_set = set(platform_order)
    comparison_pairs = [
        ("Humanoid", "Manual"),
        ("dVRK", "Manual"),
        ("Humanoid", "dVRK"),
    ]
    comparison_pairs = [
        (first, second)
        for first, second in comparison_pairs
        if first in platform_set and second in platform_set
    ]

    pairwise_rows: List[Dict[str, object]] = []
    for metric_name, participant_values in (
        ("Weighted Error", err_by_participant),
        ("Time (s)", time_by_participant),
    ):
        for first, second in comparison_pairs:
            stats = paired_stats(participant_values, first, second)
            selected_test, selected_p = choose_test(stats, normality_alpha)
            pairwise_rows.append(
                {
                    "Metric": metric_name,
                    "Comparison": f"{first} - {second}",
                    "Delta [95% CI]": fmt_ci(stats["mean_diff"], stats["ci_low"], stats["ci_high"], 2),
                    "p": selected_p,
                    "d_z": fmt_d(stats["cohens_dz"], 2),
                    "Test": selected_test,
                }
            )

    return pairwise_rows


def write_csv(path: str, rows: Sequence[Dict[str, object]]) -> None:
    if not rows:
        return
    output_path = Path(path)
    output_path.parent.mkdir(parents=True, exist_ok=True)
    with output_path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0].keys()))
        writer.writeheader()
        writer.writerows(rows)


def format_table(rows: Sequence[Dict[str, object]], columns: Sequence[str]) -> str:
    if not rows:
        return "(Empty)"

    widths = []
    for column in columns:
        width = len(column)
        for row in rows:
            width = max(width, len(str(row.get(column, ""))))
        widths.append(width)

    def build_line(values: Iterable[str]) -> str:
        return "  ".join(str(value).ljust(width) for value, width in zip(values, widths))

    header = build_line(columns)
    separator = "  ".join("-" * width for width in widths)
    body = [build_line([row.get(column, "") for column in columns]) for row in rows]
    return "\n".join([header, separator, *body])


def print_photo_notes(normality_alpha: float) -> None:
    threshold = f"{normality_alpha:.2f}".rstrip("0").rstrip(".")
    print(
        "\nNotes: Delta is the mean paired difference (first minus second) with a 95% confidence interval (CI). "
        "d_z is Cohen's d for paired samples computed on paired differences. "
        "Test indicates a paired two-sided t-test (t) or Wilcoxon signed-rank test (w), "
        f"selected using the Shapiro-Wilk normality check on paired differences (p >= {threshold} -> t; otherwise w)."
    )


def print_summary_table(rows: Sequence[Dict[str, object]]) -> None:
    print("\n(A) O-ring transfer: weighted error and time (mean +/- std)\n")
    print(format_table(rows, ["Group", "Platform", "Weighted Error", "Time (s)"]))


def print_occurrence_summary_table(rows: Sequence[Dict[str, object]]) -> None:
    print("\n(C) O-ring transfer: error type occurrences (mean +/- std)\n")
    print(
        format_table(
            rows,
            ["Group", "Platform"] + [f"Error {key}" for key in DEFAULT_ERROR_OCCURRENCE_ORDER],
        )
    )


def print_pairwise_tables(group_names: Sequence[str], all_rows: Sequence[Dict[str, object]]) -> None:
    for index, group_name in enumerate(group_names, start=1):
        group_rows = [row for row in all_rows if row["Group"] == group_name]
        print(f"\n(B{index}) O-ring transfer: paired comparison statistics ({group_name})\n")
        print(format_table(group_rows, ["Metric", "Comparison", "Delta [95% CI]", "p", "d_z", "Test"]))


def print_primary_table(rows: Sequence[Dict[str, object]], show_vs_dvrk: bool) -> None:
    columns = [
        "Group",
        "Platform",
        "Weighted Error (mean +/- std)",
        "Delta Err vs Manual (mean [CI])",
        "p_t (Err vs Manual)",
        "d (Err vs Manual)",
        "Time (s) (mean +/- std)",
        "Delta Time vs Manual (mean [CI])",
        "p_t (Time vs Manual)",
        "d (Time vs Manual)",
    ]
    if show_vs_dvrk:
        columns.extend(
            [
                "Delta Err vs dVRK (mean [CI])",
                "p_t (Err vs dVRK)",
                "d (Err vs dVRK)",
                "Delta Time vs dVRK (mean [CI])",
                "p_t (Time vs dVRK)",
                "d (Time vs dVRK)",
            ]
        )

    print("\n=== PRIMARY RESULTS ===\n")
    print(format_table(rows, columns))


def print_diagnostics_table(rows: Sequence[Dict[str, object]]) -> None:
    print("\n=== DIAGNOSTICS ===\n")
    print(format_table(rows, ["Group", "Comparison", "Metric", "n", "Shapiro p", "Wilcoxon p"]))


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--csv", type=str, default="source_data/oring_transfer_data.csv")
    parser.add_argument("--weights", type=str, default="")
    parser.add_argument("--platform_order", type=str, default="dVRK,Humanoid,Manual")
    parser.add_argument("--out_csv", type=str, default="", help="If set, write the primary results to this CSV path.")
    parser.add_argument("--out_diag_csv", type=str, default="", help="If set, write the diagnostics results to this CSV path.")
    parser.add_argument("--out_summary_csv", type=str, default="", help="If set, write the compact summary table to this CSV path.")
    parser.add_argument("--out_occurrence_csv", type=str, default="", help="If set, write the error-occurrence summary table to this CSV path.")
    parser.add_argument("--out_pairwise_csv", type=str, default="", help="If set, write the paired-comparison table to this CSV path.")
    parser.add_argument("--hide_vs_dvrk", action="store_true", help="If set, omit the comparison-vs-dVRK columns.")
    parser.add_argument("--show_diagnostics", action="store_true", help="If set, print the diagnostics table.")
    parser.add_argument("--show_verbose", action="store_true", help="If set, also print the wider comparison table.")
    parser.add_argument("--normality_alpha", type=float, default=0.05, help="Normality threshold used to choose t versus Wilcoxon.")
    args = parser.parse_args()

    weights = parse_weights(args.weights)
    platform_order = [item.strip() for item in args.platform_order.split(",") if item.strip()]

    rows = load_and_compute_metrics(args.csv, weights, platform_order)
    rows_overall = list(rows)
    rows_novices = [row for row in rows if not row["IsSurgeon"]]
    rows_surgeons = [row for row in rows if row["IsSurgeon"]]

    primary_rows: List[Dict[str, object]] = []
    diagnostic_rows: List[Dict[str, object]] = []
    summary_rows: List[Dict[str, object]] = []
    occurrence_rows: List[Dict[str, object]] = []
    pairwise_rows: List[Dict[str, object]] = []
    group_names: List[str] = []

    for group_name, group_rows in (
        ("Overall", rows_overall),
        ("Novices", rows_novices),
        ("Surgeons", rows_surgeons),
    ):
        group_names.append(group_name)
        group_primary, group_diagnostics = build_group_tables(group_rows, group_name, platform_order)
        primary_rows.extend(group_primary)
        diagnostic_rows.extend(group_diagnostics)
        summary_rows.extend(build_summary_rows(group_primary, group_name, platform_order))
        occurrence_rows.extend(build_occurrence_summary_rows(group_rows, group_name, platform_order))
        group_pairwise_rows = build_paired_comparison_rows(group_rows, platform_order, args.normality_alpha)
        pairwise_rows.extend({"Group": group_name, **row} for row in group_pairwise_rows)

    if args.out_csv:
        write_csv(args.out_csv, primary_rows)
        print(f"Saved primary CSV: {args.out_csv}")

    if args.out_diag_csv:
        write_csv(args.out_diag_csv, diagnostic_rows)
        print(f"Saved diagnostics CSV: {args.out_diag_csv}")

    if args.out_summary_csv:
        write_csv(args.out_summary_csv, summary_rows)
        print(f"Saved summary CSV: {args.out_summary_csv}")

    if args.out_occurrence_csv:
        write_csv(args.out_occurrence_csv, occurrence_rows)
        print(f"Saved occurrence CSV: {args.out_occurrence_csv}")

    if args.out_pairwise_csv:
        write_csv(args.out_pairwise_csv, pairwise_rows)
        print(f"Saved paired-comparison CSV: {args.out_pairwise_csv}")

    print_photo_notes(args.normality_alpha)
    print_summary_table(summary_rows)
    print_pairwise_tables(group_names, pairwise_rows)
    print_occurrence_summary_table(occurrence_rows)

    if args.show_verbose:
        print_primary_table(primary_rows, show_vs_dvrk=(not args.hide_vs_dvrk))

    if args.show_diagnostics:
        print_diagnostics_table(diagnostic_rows)


if __name__ == "__main__":
    main()
