"""Join tidy sequence and label CSVs by sample_id, with a strict 1:1 policy."""
import argparse
from pathlib import Path
import sys
import warnings

import pandas as pd


def read_table(path, value_column):
    with warnings.catch_warnings():
        warnings.simplefilter("error", pd.errors.ParserWarning)
        table = pd.read_csv(path, dtype="string", keep_default_na=False,
                            skip_blank_lines=False, index_col=False)
    columns = ["sample_id", value_column]
    if len(table.columns) != 2 or set(table.columns) != set(columns):
        raise ValueError(f"{path}: expected exactly {columns}")
    table = table[columns]
    if table.empty:
        raise ValueError(f"{path}: no data rows")
    for column in columns:
        values = table[column]
        if values.isna().any() or values.str.strip().eq("").any():
            raise ValueError(f"{path}: blank {column}")
    ids = table["sample_id"]
    if ids.ne(ids.str.strip()).any():
        raise ValueError(f"{path}: sample_id has surrounding whitespace")
    duplicates = ids[ids.duplicated(keep=False)].unique().tolist()
    if duplicates:
        raise ValueError(f"{path}: duplicate sample_id: {duplicates[:5]}")
    return table


def join_tables(sequences, labels):
    audit = sequences.merge(labels, on="sample_id", how="outer",
                            validate="one_to_one", indicator=True)
    missing = audit.loc[audit["_merge"].eq("left_only"), "sample_id"].tolist()
    extra = audit.loc[audit["_merge"].eq("right_only"), "sample_id"].tolist()
    if missing or extra:
        raise ValueError(f"ID mismatch: missing labels={missing[:5]}; "
                         f"labels without sequences={extra[:5]}")
    # Use the left join only after the full coverage audit has passed.
    return sequences.merge(labels, on="sample_id", how="left", sort=False,
                           validate="one_to_one")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("sequences", type=Path)
    parser.add_argument("labels", type=Path)
    parser.add_argument("output", type=Path)
    args = parser.parse_args()
    try:
        sequences = read_table(args.sequences, "sequence")
        labels = read_table(args.labels, "label")
        joined = join_tables(sequences, labels)
        # Exclusive creation preserves earlier outputs and the input files.
        with args.output.open("x", encoding="utf-8", newline="") as handle:
            joined.to_csv(handle, index=False, lineterminator="\n")
        print(f"OK: {len(joined)} sequences, {len(labels)} labels, "
              f"{len(joined)} matched rows")
        return 0
    except (OSError, ValueError, pd.errors.ParserError,
            pd.errors.ParserWarning, pd.errors.EmptyDataError,
            pd.errors.MergeError) as error:
        print(f"ERROR: {error}", file=sys.stderr)
        return 2


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