#!/usr/bin/env python
# Shared helper: read-level mapping ground truth (mapped_reads / total_reads) per
# (experiment_name, sample_base), sourced from ../consolidated_df_clean.tsv.
#
# WHY: the figures originally used qualimap_mapping_rate = mapped_bases/sequenced_bases,
# which qualimap computes ONLY over reads that mapped — so it saturates at ~1.0 and is
# structurally blind to unmapped content (the appended metagenomic contaminant). The
# correct mapping ground truth is read-level mapped/total: 0.458 for every metagenomic
# sample (matching SNIPE's ~0.46), ~1.0 for the reference-derived assays.
import os
import numpy as np
import pandas as pd

_HERE = os.path.dirname(os.path.abspath(__file__))
_CLEAN = os.path.join(_HERE, "..", "consolidated_df_clean.tsv")


def readlevel_mapping_gt():
    """Return DataFrame[experiment_name, sample_base, mapping_gt_readlevel] (0-1)."""
    c = pd.read_csv(_CLEAN, sep="\t",
                    usecols=lambda x: x in ("experiment_name", "sample_base",
                                            "qualimap_global_no_of_reads",
                                            "qualimap_global_no_of_mapped_reads",
                                            "qualimap_global_percentage_of_mapped_reads"))
    g = c.drop_duplicates(["experiment_name", "sample_base"]).copy()
    tot = g["qualimap_global_no_of_reads"]
    mp = g["qualimap_global_no_of_mapped_reads"]
    frac = mp / tot
    # fallback to the reported percentage where read counts are missing
    pct = g["qualimap_global_percentage_of_mapped_reads"] / 100.0
    g["mapping_gt_readlevel"] = frac.where(tot.notna() & (tot > 0), pct)
    return g[["experiment_name", "sample_base", "mapping_gt_readlevel"]]


def attach_readlevel_mapping(df, base_col="qualimap_mapping_rate"):
    """Merge read-level mapping GT into df, falling back to base_col where unavailable.
    Returns df with a new column 'mapping_gt' to be used as the mapping metric's GT."""
    gt = readlevel_mapping_gt()
    out = df.merge(gt, on=["experiment_name", "sample_base"], how="left")
    out["mapping_gt"] = out["mapping_gt_readlevel"]
    if base_col in out.columns:
        out["mapping_gt"] = out["mapping_gt"].where(out["mapping_gt"].notna(), out[base_col])
    return out
