import argparse
import pandas as pd
from gp_common import Run, checked_run, digest

def compare(folder, input_path="outputs/google_play_reviews.csv", root="outputs/gp_runs"):
    folder, meta = checked_run(folder)
    if digest(input_path) != meta["input_sha256"]:
        raise ValueError("入力CSVが学習時と異なります。元ファイルを確認してください。")
    required = ["lda_assignments.csv", "bertopic_assignments.csv"]
    if meta.get("statuses", {}).get("lda") != "OK" or meta.get("statuses", {}).get("bertopic") not in {"OK", "ALL_OUTLIERS"}:
        raise ValueError("同じ実行で両手法が完了していません。手法別statusを確認してください。")
    if not all(name in meta["files"] for name in required):
        raise ValueError("片方のモデルが省略・失敗しています。別実行の成功ファイルを混ぜないでください。")
    tables = []
    for name in required:
        d = pd.read_csv(folder / name, dtype=str, keep_default_na=False)
        if d["row_id"].duplicated().any() or d["row_id"].eq("").any():
            raise ValueError("row_id重複・欠損")
        if not d["run_id"].eq(meta["run_id"]).all() or not d["input_sha256"].eq(meta["input_sha256"]).all():
            raise ValueError("別実行の結果が混ざっています。")
        tables.append(d.set_index("row_id").sort_index())
    a, b = tables
    columns = ["content_raw", "score_raw", "at_raw", "reviewId"]
    if not a[columns].equals(b[columns]) or set(a.index) != {str(i) for i in meta["selected_row_ids"]}:
        raise ValueError("ID・本文・評価・日付または分析対象が一致しません。")
    joined = pd.DataFrame({"lda": a["topic_id"], "bertopic": b["topic_id"]})
    outliers = int(joined["bertopic"].eq("-1").sum())
    valid = joined[joined["bertopic"].ne("-1")]
    count = valid.groupby(["bertopic", "lda"]).size().reset_index(name="count")
    count["denominator"] = count.groupby("bertopic")["count"].transform("sum")
    count["ratio"] = count["count"] / count["denominator"]
    summary = count.sort_values("count", ascending=False).drop_duplicates("bertopic").rename(columns={"lda": "largest_lda", "ratio": "purity"})
    run = Run("topic-comparison", folder / "metadata.json", {"source_run": meta["run_id"], "source_input_sha256": meta["input_sha256"]}, root)
    run.csv("correspondence.csv", count)
    run.csv("group_largest_share.csv", summary)
    run.csv("joined_assignments.csv", joined.reset_index())
    for _, row in summary.iterrows():
        print(f"BERTopic {row.bertopic}: {int(row.denominator)}件中{int(row['count'])}件がLDA {row.largest_lda} ({row.purity:.1%})。意味の一致は上位語と原文で確認。")
    return run.finish(paired_rows=len(joined), comparison_rows=len(valid), outlier_rows=outliers,
                      note="最大割当比率。意味の一致・正しさ・再学習時の安定性ではない。")

if __name__ == "__main__":
    p = argparse.ArgumentParser(); p.add_argument("run_folder"); p.add_argument("--input", default="outputs/google_play_reviews.csv")
    a = p.parse_args(); compare(a.run_folder, a.input)
