import json
import random
import csv
import math
from pathlib import Path
from collections import defaultdict


# =========================
# 設定
# =========================

BASE_DIR = Path("all/5")

MODEL_DIRS = {
    "gpt-5": BASE_DIR / "gpt-5",
    "gpt-4o": BASE_DIR / "gpt-4o",
}

TRIALS = 1000
STEP = 5
RANDOM_SEED = 42

# 敷居判定の条件
SPEARMAN_THRESHOLD = 0.95
TOP3_THRESHOLD = 0.95
TOP1_THRESHOLD = 0.95


# =========================
# JSON読み込み
# =========================

def load_json_files(model_name, folder_path):
    """
    指定モデルのフォルダからjsonファイルをすべて読み込む。
    1ファイル = 1試合分の評価ログとして扱う。
    """
    files = sorted(folder_path.rglob("*.json"))

    if not files:
        raise FileNotFoundError(f"{folder_path} に .json ファイルが見つかりません。")

    games = []

    for file_path in files:
        with open(file_path, "r", encoding="utf-8") as f:
            data = json.load(f)

        game_id = data.get("game_id", file_path.stem)

        games.append({
            "game_id": game_id,
            "file": file_path.name,
            "model": model_name,
            "judgements": [data],
        })

    return games


def load_all_patterns():
    """
    gpt-5単独、gpt-4o単独、gpt-5+gpt-4統合の3パターンを作る。
    """
    loaded_by_model = {}

    for model_name, folder_path in MODEL_DIRS.items():
        loaded_by_model[model_name] = load_json_files(model_name, folder_path)

    patterns = {}

    # モデル別
    patterns["gpt-5"] = loaded_by_model["gpt-5"]
    patterns["gpt-4o"] = loaded_by_model["gpt-4o"]

    # 統合版
    patterns["combined"] = combine_games_by_game_id(loaded_by_model)

    return patterns


def combine_games_by_game_id(loaded_by_model):
    """
    gpt-5とgpt-4の同じgame_idをまとめる。
    統合版では、1つのgame_idに複数モデルの評価を入れる。

    注意：
    240試合として扱うのではなく、
    120試合それぞれに2モデル分の評価があるものとして扱う。
    """
    combined = {}

    for model_name, games in loaded_by_model.items():
        for game in games:
            game_id = game["game_id"]

            if game_id not in combined:
                combined[game_id] = {
                    "game_id": game_id,
                    "file": game["file"],
                    "model": "combined",
                    "judgements": [],
                }

            combined[game_id]["judgements"].extend(game["judgements"])

    return list(combined.values())


# =========================
# チーム一覧の取得
# =========================

def extract_all_teams(games):
    """
    teamフィールドから全チーム名を集める。
    """
    teams = set()

    for game in games:
        for data in game["judgements"]:
            evaluations = data.get("evaluations", {})

            for eval_name, eval_data in evaluations.items():
                rankings = eval_data.get("rankings", [])

                for row in rankings:
                    team = row.get("team")
                    if team:
                        teams.add(team)

    return sorted(teams)


# =========================
# スコア集計
# =========================

def score_games(games_subset, all_teams):
    """
    指定された試合群からランキングを作る。

    各評価項目ごとに、
    1位=5点、2位=4点、3位=3点、4位=2点、5位=1点
    として集計する。

    combinedの場合は、同じ試合に対してgpt-5とgpt-4oの評価を両方使う。
    """
    total_points = defaultdict(float)
    eval_count = defaultdict(int)
    appearance_count = defaultdict(int)

    for game in games_subset:
        teams_in_this_game = set()

        for data in game["judgements"]:
            evaluations = data.get("evaluations", {})

            for eval_name, eval_data in evaluations.items():
                rankings = eval_data.get("rankings", [])
                player_count = len(rankings)

                for row in rankings:
                    team = row.get("team")
                    rank = row.get("ranking")

                    if team is None or rank is None:
                        continue

                    teams_in_this_game.add(team)

                    # 5人村なら 1位=5点、2位=4点、...、5位=1点
                    points = player_count - rank + 1

                    total_points[team] += points
                    eval_count[team] += 1

        # この試合に出たチームの出場回数を1増やす
        # combinedでも、同じ試合を2回出場扱いにはしない
        for team in teams_in_this_game:
            appearance_count[team] += 1

    results = {}

    for team in all_teams:
        if eval_count[team] > 0:
            mean_score = total_points[team] / eval_count[team]
        else:
            mean_score = -1.0

        results[team] = {
            "team": team,
            "total_points": total_points[team],
            "eval_count": eval_count[team],
            "appearance_count": appearance_count[team],
            "mean_score": mean_score,
        }

    ranking = sorted(
        results.values(),
        key=lambda x: (-x["mean_score"], -x["total_points"], x["team"])
    )

    return ranking, results


# =========================
# ランキング比較
# =========================

def rank_positions(ranking):
    return {
        row["team"]: i + 1
        for i, row in enumerate(ranking)
    }


def spearman_rank_correlation(ranking_a, ranking_b):
    """
    Spearman順位相関を計算する。
    scipy不要版。
    """
    pos_a = rank_positions(ranking_a)
    pos_b = rank_positions(ranking_b)

    teams = sorted(set(pos_a.keys()) & set(pos_b.keys()))
    n = len(teams)

    if n <= 1:
        return 1.0

    d2_sum = 0

    for team in teams:
        d = pos_a[team] - pos_b[team]
        d2_sum += d * d

    return 1 - (6 * d2_sum) / (n * (n * n - 1))


def compare_rankings(sample_ranking, final_ranking):
    sample_order = [row["team"] for row in sample_ranking]
    final_order = [row["team"] for row in final_ranking]

    top1_match = sample_order[0] == final_order[0]
    top3_match = set(sample_order[:3]) == set(final_order[:3])
    full_match = sample_order == final_order

    spearman = spearman_rank_correlation(sample_ranking, final_ranking)

    sample_pos = rank_positions(sample_ranking)
    final_pos = rank_positions(final_ranking)

    max_rank_diff = max(
        abs(sample_pos[team] - final_pos[team])
        for team in final_pos
    )

    return {
        "top1_match": top1_match,
        "top3_match": top3_match,
        "full_match": full_match,
        "spearman": spearman,
        "max_rank_diff": max_rank_diff,
    }


# =========================
# 敷居分析
# =========================

def analyze_threshold(games, all_teams, final_ranking):
    random.seed(RANDOM_SEED)

    total_games = len(games)
    summary_rows = []

    sample_sizes = list(range(STEP, total_games + 1, STEP))

    if total_games not in sample_sizes:
        sample_sizes.append(total_games)

    for sample_size in sample_sizes:
        metrics_list = []

        if sample_size == total_games:
            sample_ranking, _ = score_games(games, all_teams)
            metrics_list.append(compare_rankings(sample_ranking, final_ranking))
        else:
            for _ in range(TRIALS):
                sampled_games = random.sample(games, sample_size)
                sample_ranking, _ = score_games(sampled_games, all_teams)
                metrics_list.append(compare_rankings(sample_ranking, final_ranking))

        top1_rate = sum(m["top1_match"] for m in metrics_list) / len(metrics_list)
        top3_rate = sum(m["top3_match"] for m in metrics_list) / len(metrics_list)
        full_match_rate = sum(m["full_match"] for m in metrics_list) / len(metrics_list)
        avg_spearman = sum(m["spearman"] for m in metrics_list) / len(metrics_list)
        avg_max_rank_diff = sum(m["max_rank_diff"] for m in metrics_list) / len(metrics_list)

        # 5人村なので、平均出場回数 = 試合数 * 5 / チーム数
        avg_appearances_per_team = sample_size * 5 / len(all_teams)

        summary_rows.append({
            "sample_games": sample_size,
            "avg_appearances_per_team": avg_appearances_per_team,
            "top1_match_rate": top1_rate,
            "top3_match_rate": top3_rate,
            "full_ranking_match_rate": full_match_rate,
            "avg_spearman": avg_spearman,
            "avg_max_rank_diff": avg_max_rank_diff,
        })

    return summary_rows


def find_threshold(summary_rows):
    for row in summary_rows:
        if (
            row["avg_spearman"] >= SPEARMAN_THRESHOLD
            and row["top3_match_rate"] >= TOP3_THRESHOLD
            and row["top1_match_rate"] >= TOP1_THRESHOLD
        ):
            return row

    return None


# =========================
# 今回大会の推奨試合数
# =========================

def recommend_game_counts(threshold_appearances, min_teams=15, max_teams=19):
    """
    前回大会から推定した1チームあたり必要出場数をもとに、
    15〜19チームの場合の推奨試合数を計算する。

    役職公平性のため、総試合数はチーム数の倍数にする。

    総試合数 = チーム数 * k
    とすると、
    各チームの平均出場数 = 5k
    各チームの人狼回数 = k
    各チームの狂人回数 = k
    各チームの占い師回数 = k
    各チームの村人回数 = 2k
    """
    rows = []

    required_k = math.ceil(threshold_appearances / 5)

    for team_count in range(min_teams, max_teams + 1):
        games = team_count * required_k
        avg_appearances = games * 5 / team_count

        rows.append({
            "team_count": team_count,
            "recommended_games": games,
            "avg_appearances_per_team": avg_appearances,
            "werewolf_count_per_team": required_k,
            "madman_count_per_team": required_k,
            "seer_count_per_team": required_k,
            "villager_count_per_team": required_k * 2,
        })

    return rows


# =========================
# CSV出力
# =========================

def write_csv(path, rows):
    if not rows:
        return

    with open(path, "w", encoding="utf-8-sig", newline="") as f:
        writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
        writer.writeheader()
        writer.writerows(rows)


# =========================
# 表示
# =========================

def print_final_ranking(pattern_name, final_ranking):
    print()
    print("================================")
    print(f"{pattern_name}: 全試合を使った最終ランキング")
    print("================================")

    for i, row in enumerate(final_ranking, start=1):
        print(
            f"{i}位: {row['team']} "
            f"平均点={row['mean_score']:.4f} "
            f"出場={row['appearance_count']}試合 "
            f"評価数={row['eval_count']}"
        )


def print_summary(pattern_name, summary_rows):
    print()
    print("================================")
    print(f"{pattern_name}: 試合数ごとのランキング安定度")
    print("================================")
    print(
        "試合数 | 1チーム平均出場 | 1位一致率 | 上位3一致率 | 全順位一致率 | 平均順位相関 | 平均最大順位差"
    )

    for row in summary_rows:
        print(
            f"{row['sample_games']:>4} | "
            f"{row['avg_appearances_per_team']:>15.2f} | "
            f"{row['top1_match_rate']:>8.3f} | "
            f"{row['top3_match_rate']:>10.3f} | "
            f"{row['full_ranking_match_rate']:>10.3f} | "
            f"{row['avg_spearman']:>10.3f} | "
            f"{row['avg_max_rank_diff']:>10.3f}"
        )


def print_threshold_and_recommendations(pattern_name, threshold_row):
    print()
    print("================================")
    print(f"{pattern_name}: 推定された敷居")
    print("================================")

    if threshold_row is None:
        print("設定した条件を満たす試合数は見つかりませんでした。")
        return None

    threshold_appearances = threshold_row["avg_appearances_per_team"]

    print(f"敷居となる全体試合数: {threshold_row['sample_games']}試合")
    print(f"敷居となる1チーム平均出場数: {threshold_appearances:.2f}試合")
    print(f"1位一致率: {threshold_row['top1_match_rate']:.3f}")
    print(f"上位3一致率: {threshold_row['top3_match_rate']:.3f}")
    print(f"平均順位相関: {threshold_row['avg_spearman']:.3f}")

    recommended_rows = recommend_game_counts(threshold_appearances, 15, 19)

    print()
    print(f"{pattern_name}: 15〜19チームの場合の推奨試合数")
    print("チーム数 | 推奨試合数 | 1チーム平均出場 | 人狼 | 狂人 | 占い師 | 村人")

    for row in recommended_rows:
        print(
            f"{row['team_count']:>7} | "
            f"{row['recommended_games']:>10} | "
            f"{row['avg_appearances_per_team']:>15.2f} | "
            f"{row['werewolf_count_per_team']:>4} | "
            f"{row['madman_count_per_team']:>4} | "
            f"{row['seer_count_per_team']:>6} | "
            f"{row['villager_count_per_team']:>4}"
        )

    return recommended_rows


# =========================
# メイン処理
# =========================

def main():
    patterns = load_all_patterns()

    threshold_overview = []

    for pattern_name, games in patterns.items():
        all_teams = extract_all_teams(games)

        print()
        print("################################################")
        print(f"分析対象: {pattern_name}")
        print("################################################")
        print(f"試合数: {len(games)}")
        print(f"チーム数: {len(all_teams)}")
        print("チーム一覧:")
        for team in all_teams:
            print(f"  - {team}")

        final_ranking, _ = score_games(games, all_teams)
        print_final_ranking(pattern_name, final_ranking)

        summary_rows = analyze_threshold(games, all_teams, final_ranking)
        print_summary(pattern_name, summary_rows)

        threshold_row = find_threshold(summary_rows)
        recommended_rows = print_threshold_and_recommendations(pattern_name, threshold_row)

        # CSV出力
        write_csv(f"threshold_summary_{pattern_name}.csv", summary_rows)

        if recommended_rows is not None:
            write_csv(f"recommended_games_{pattern_name}.csv", recommended_rows)

        if threshold_row is not None:
            threshold_overview.append({
                "pattern": pattern_name,
                "threshold_games": threshold_row["sample_games"],
                "threshold_avg_appearances_per_team": threshold_row["avg_appearances_per_team"],
                "top1_match_rate": threshold_row["top1_match_rate"],
                "top3_match_rate": threshold_row["top3_match_rate"],
                "avg_spearman": threshold_row["avg_spearman"],
            })
        else:
            threshold_overview.append({
                "pattern": pattern_name,
                "threshold_games": None,
                "threshold_avg_appearances_per_team": None,
                "top1_match_rate": None,
                "top3_match_rate": None,
                "avg_spearman": None,
            })

    print()
    print("################################################")
    print("3パターンの敷居まとめ")
    print("################################################")

    for row in threshold_overview:
        print(
            f"{row['pattern']}: "
            f"敷居試合数={row['threshold_games']}, "
            f"1チーム平均出場={row['threshold_avg_appearances_per_team']}, "
            f"1位一致率={row['top1_match_rate']}, "
            f"上位3一致率={row['top3_match_rate']}, "
            f"順位相関={row['avg_spearman']}"
        )

    write_csv("threshold_overview.csv", threshold_overview)

    print()
    print("CSVファイルを出力しました。")
    print("  - threshold_summary_gpt-5.csv")
    print("  - threshold_summary_gpt-4o.csv")
    print("  - threshold_summary_combined.csv")
    print("  - recommended_games_gpt-5.csv")
    print("  - recommended_games_gpt-4o.csv")
    print("  - recommended_games_combined.csv")
    print("  - threshold_overview.csv")


if __name__ == "__main__":
    main()
