"""◎〇▲△の4頭選出: レース内勝率（合計100%）×単勝オッズの順位。"""
from __future__ import annotations

MARK_LABELS = ("◎本命", "〇対抗", "▲単穴", "△連下")


def normalize_win_probabilities(
    rows: list[dict],
    *,
    score_key: str = "score_accuracy",
) -> None:
    """
    各馬の勝率を score_key の正の値で按分し、レース内で合計 1.0（100%）にする。
    """
    weights: list[float] = []
    for row in rows:
        w = max(float(row.get(score_key) or row.get("score") or 0.0), 0.01)
        weights.append(w)
    total = sum(weights)
    if total <= 0:
        total = float(len(rows)) or 1.0
        weights = [1.0] * len(rows)
    for row, w in zip(rows, weights):
        prob = w / total
        row["win_probability"] = round(prob, 6)
        row["stage3_prob"] = row["win_probability"]


def apply_prob_odds_ranking(rows: list[dict]) -> None:
    """
    勝率×オッズ（prob_times_odds）の降順で rank_prob_odds を付与する。
    オッズが無い馬は勝率のみで後方に並べ、上位4頭の印候補に使う。
    """
    if not rows:
        return
    if rows[0].get("win_probability") is None:
        normalize_win_probabilities(rows)

    with_odds: list[dict] = []
    without_odds: list[dict] = []
    for row in rows:
        prob = float(row.get("win_probability") or row.get("stage3_prob") or 0.0)
        odds = row.get("win_odds")
        if odds is not None and float(odds) > 0:
            row["prob_times_odds"] = round(prob * float(odds), 6)
            with_odds.append(row)
        else:
            row["prob_times_odds"] = None
            without_odds.append(row)

    with_odds.sort(
        key=lambda x: (
            float(x.get("prob_times_odds") or 0.0),
            float(x.get("win_probability") or 0.0),
        ),
        reverse=True,
    )
    without_odds.sort(
        key=lambda x: float(x.get("win_probability") or 0.0),
        reverse=True,
    )
    ordered = with_odds + without_odds
    for idx, row in enumerate(ordered, start=1):
        row["rank_prob_odds"] = idx


def pick_four_marks(
    rows: list[dict],
    *,
    score_key: str = "prob_times_odds",
) -> list[dict]:
    """上位4頭に印ラベルを付けて返す（import / 画面用）。"""
    apply_prob_odds_ranking(rows)
    ranked = sorted(rows, key=lambda x: int(x.get("rank_prob_odds") or 999))
    out: list[dict] = []
    for i, row in enumerate(ranked[:4]):
        out.append(
            {
                "mark": MARK_LABELS[i],
                "horse_number": int(row["horse_number"]),
                "horse_name": str(row.get("horse_name") or ""),
                "score": float(row.get(score_key) or row.get("prob_times_odds") or 0.0),
                "win_odds": row.get("win_odds"),
                "win_probability": row.get("win_probability"),
                "stage3_prob": row.get("stage3_prob"),
                "prob_times_odds": row.get("prob_times_odds"),
            }
        )
    return out


def verify_probability_sum(rows: list[dict], tol: float = 1e-4) -> float:
    """検証用: 勝率の合計（1.0 に近いか）。"""
    return sum(float(r.get("win_probability") or r.get("stage3_prob") or 0.0) for r in rows)
