#!/usr/bin/env python3
"""
結果確定済みJRA全レースを、予想プロファイルで計算して保存する。

公開用は期待値予想（ev_forecast）。ベース予想（base_forecast）はオッズ無しの順位・勝率見立て。
"""
from __future__ import annotations

import argparse
import json
import os
import sys

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from ai.prediction_profiles import get_profiles
from ai.stage1_screening import run_stage1
from collectors.db_util import connect


def _winner(cur, race_id: int) -> tuple[int | None, str | None]:
    cur.execute(
        """
        SELECT re.horse_number, h.name
        FROM race_results rr
        JOIN race_entries re ON re.race_id = rr.race_id AND re.horse_id = rr.horse_id
        JOIN horses h ON h.id = re.horse_id
        WHERE rr.race_id = %s AND rr.finish_position = 1
        LIMIT 1
        """,
        (race_id,),
    )
    row = cur.fetchone()
    if not row:
        return None, None
    return int(row["horse_number"]), str(row["name"] or "")


def _select_targets(cur, from_date: str | None, to_date: str | None, limit: int | None) -> list[dict]:
    where = ["r.circuit = 'JRA'"]
    params: list = []
    if from_date:
        where.append("r.race_date >= %s")
        params.append(from_date)
    if to_date:
        where.append("r.race_date <= %s")
        params.append(to_date)
    sql = f"""
        SELECT r.id AS race_id, r.race_date, t.code AS track_code, r.race_number, r.race_name
        FROM races r
        JOIN tracks t ON t.id = r.track_id
        WHERE {" AND ".join(where)}
          AND EXISTS (
            SELECT 1 FROM race_results rr
            WHERE rr.race_id = r.id AND rr.finish_position = 1
          )
        ORDER BY r.race_date DESC, t.code ASC, r.race_number ASC
    """
    if limit and limit > 0:
        sql += f" LIMIT {int(limit)}"
    cur.execute(sql, params)
    return list(cur.fetchall())


def _ranked(horses: list[dict], rank_mode: str) -> tuple[list[dict], str]:
    if rank_mode == "prob_odds":
        return sorted(horses, key=lambda x: int(x.get("rank_prob_odds") or 999)), "prob_times_odds"
    if rank_mode == "value":
        return sorted(horses, key=lambda x: int(x.get("rank_value") or 999)), "score_value"
    if rank_mode == "aggressive":
        return sorted(horses, key=lambda x: int(x.get("rank_value_aggressive") or 999)), "expected_roi_aggressive"
    return sorted(horses, key=lambda x: int(x.get("rank_accuracy") or 999)), "score_accuracy"


def _mk_marks(rows: list[dict], score_key: str) -> list[dict]:
    names = ["◎本命", "〇対抗", "▲単穴", "△連下"]
    out: list[dict] = []
    for i, row in enumerate(rows[:4]):
        out.append(
            {
                "mark": names[i],
                "horse_number": int(row["horse_number"]),
                "horse_name": str(row.get("horse_name") or ""),
                "score": float(row.get(score_key) or row.get("score_accuracy") or row.get("score") or 0),
                "win_odds": row.get("win_odds"),
                "win_probability": row.get("win_probability"),
                "win_rate": row.get("win_rate"),
                "model_win_prob": row.get("model_win_prob"),
                "cut": row.get("cut"),
                "cut_reason": row.get("cut_reason"),
                "expected_roi": row.get("expected_roi"),
            }
        )
    return out


def _upsert(conn, rec: dict) -> None:
    cols = [
        "profile_key",
        "profile_label",
        "race_id",
        "race_date",
        "track_code",
        "race_number",
        "race_name",
        "ai_honmei_umaban",
        "ai_honmei_name",
        "ai_honmei_score",
        "ai_taikou_umaban",
        "ai_taikou_name",
        "ai_taikou_score",
        "ai_tanana_umaban",
        "ai_tanana_name",
        "ai_tanana_score",
        "ai_renka_umaban",
        "ai_renka_name",
        "ai_renka_score",
        "winner_umaban",
        "winner_name",
        "match_honmei_winner",
        "match_taikou_winner",
        "match_honmei_or_taikou_winner",
        "status",
        "error_message",
        "ai_snapshot_json",
    ]
    placeholders = ", ".join(["%s"] * len(cols))
    updates = ", ".join([f"{c}=VALUES({c})" for c in cols if c not in ("profile_key", "race_id")])
    sql = f"""
        INSERT INTO all_race_ai_profile_compare ({", ".join(cols)})
        VALUES ({placeholders})
        ON DUPLICATE KEY UPDATE {updates}
    """
    with conn.cursor() as cur:
        cur.execute(sql, [rec.get(c) for c in cols])
    conn.commit()


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--from-date", default=None, help="YYYY-MM-DD")
    ap.add_argument("--to-date", default=None, help="YYYY-MM-DD")
    ap.add_argument("--limit", type=int, default=0, help="0=全件")
    ap.add_argument("--profile", default="all", help="all または profile key")
    ap.add_argument("--dry-run", action="store_true")
    ap.add_argument("--model-path", default=None, help="LightGBM モデル（例: ai/models/stage1_lgbm_enhanced.txt）")
    ap.add_argument("--model-meta-path", default=None, help="メタJSON（例: ai/models/stage1_lgbm_enhanced.meta.json）")
    args = ap.parse_args()

    profiles = get_profiles(args.profile)
    conn = connect()
    ok = 0
    err = 0
    try:
        with conn.cursor() as cur:
            targets = _select_targets(
                cur=cur,
                from_date=args.from_date,
                to_date=args.to_date,
                limit=(args.limit if args.limit > 0 else None),
            )

        total = len(targets) * len(profiles)
        done = 0
        # レースは DB から race_date DESC（新しい順）。外側をレースにし、各レースで全プロファイルを連続計算する。
        for t in targets:
            for profile in profiles:
                done += 1
                race_id = int(t["race_id"])
                rec = {
                    "profile_key": profile.key,
                    "profile_label": profile.label,
                    "race_id": race_id,
                    "race_date": t["race_date"],
                    "track_code": t["track_code"],
                    "race_number": int(t["race_number"]),
                    "race_name": t.get("race_name"),
                    "ai_honmei_umaban": None,
                    "ai_honmei_name": None,
                    "ai_honmei_score": None,
                    "ai_taikou_umaban": None,
                    "ai_taikou_name": None,
                    "ai_taikou_score": None,
                    "ai_tanana_umaban": None,
                    "ai_tanana_name": None,
                    "ai_tanana_score": None,
                    "ai_renka_umaban": None,
                    "ai_renka_name": None,
                    "ai_renka_score": None,
                    "winner_umaban": None,
                    "winner_name": None,
                    "match_honmei_winner": None,
                    "match_taikou_winner": None,
                    "match_honmei_or_taikou_winner": None,
                    "status": "ok",
                    "error_message": None,
                    "ai_snapshot_json": None,
                }
                try:
                    payload = run_stage1(
                        race_id=race_id,
                        lookback_runs=profile.lookback_runs,
                        pass_rate=profile.pass_rate,
                        enable_stage1_5=profile.enable_stage1_5,
                        enable_stage2=profile.enable_stage2,
                        enable_stage3=profile.enable_stage3,
                        stage3_min_prob=profile.stage3_min_prob,
                        stage3_odds_cap=profile.stage3_odds_cap,
                        stage3_aggressive_min_prob=profile.stage3_aggressive_min_prob,
                        stage3_aggressive_odds_cap=profile.stage3_aggressive_odds_cap,
                        model_path=args.model_path,
                        model_meta_path=args.model_meta_path,
                        forecast_pipeline=profile.forecast_pipeline,
                        rank_mode=profile.rank_mode,
                    )
                    horses = payload.get("horses") or []
                    ranked_rows, score_key = _ranked(horses, profile.rank_mode)
                    marks = _mk_marks(ranked_rows, score_key)

                    if len(marks) >= 1:
                        rec["ai_honmei_umaban"] = marks[0]["horse_number"]
                        rec["ai_honmei_name"] = marks[0]["horse_name"]
                        rec["ai_honmei_score"] = marks[0]["score"]
                    if len(marks) >= 2:
                        rec["ai_taikou_umaban"] = marks[1]["horse_number"]
                        rec["ai_taikou_name"] = marks[1]["horse_name"]
                        rec["ai_taikou_score"] = marks[1]["score"]
                    if len(marks) >= 3:
                        rec["ai_tanana_umaban"] = marks[2]["horse_number"]
                        rec["ai_tanana_name"] = marks[2]["horse_name"]
                        rec["ai_tanana_score"] = marks[2]["score"]
                    if len(marks) >= 4:
                        rec["ai_renka_umaban"] = marks[3]["horse_number"]
                        rec["ai_renka_name"] = marks[3]["horse_name"]
                        rec["ai_renka_score"] = marks[3]["score"]

                    with conn.cursor() as cur:
                        w_umaban, w_name = _winner(cur, race_id)
                    rec["winner_umaban"] = w_umaban
                    rec["winner_name"] = w_name

                    h = rec["ai_honmei_umaban"]
                    o = rec["ai_taikou_umaban"]
                    w = rec["winner_umaban"]
                    if w is not None:
                        rec["match_honmei_winner"] = 1 if (h is not None and h == w) else 0
                        rec["match_taikou_winner"] = 1 if (o is not None and o == w) else 0
                        rec["match_honmei_or_taikou_winner"] = 1 if ((h is not None and h == w) or (o is not None and o == w)) else 0

                    snap_body = {
                        "profile": {
                            "key": profile.key,
                            "label": profile.label,
                            "description": profile.description,
                            "rank_mode": profile.rank_mode,
                            "forecast_pipeline": profile.forecast_pipeline,
                            "enable_stage1_5": profile.enable_stage1_5,
                            "enable_stage2": profile.enable_stage2,
                            "enable_stage3": profile.enable_stage3,
                            "lookback_runs": profile.lookback_runs,
                            "pass_rate": profile.pass_rate,
                        },
                        "race": payload.get("race"),
                        "config": payload.get("config"),
                        "summary": payload.get("summary"),
                        "marks": marks,
                        "horses_top5_ranked": ranked_rows[:5],
                        "horses_all_ranked": sorted(
                            horses,
                            key=lambda x: int(x.get("rank_accuracy") or 999),
                        ),
                    }
                    hc = payload.get("honmei_confidence")
                    if hc is not None:
                        snap_body["honmei_confidence"] = hc
                    emp = payload.get("empirical_hit_rates_ev")
                    if emp is not None:
                        snap_body["empirical_hit_rates_ev"] = emp
                    rec["ai_snapshot_json"] = json.dumps(
                        snap_body,
                        ensure_ascii=False,
                        default=str,
                    )
                    if not args.dry_run:
                        _upsert(conn, rec)
                    ok += 1
                    print(f"[{done}/{total}] OK {profile.key} race_id={race_id}", flush=True)
                except Exception as e:
                    rec["status"] = "error"
                    rec["error_message"] = str(e)[:500]
                    if not args.dry_run:
                        _upsert(conn, rec)
                    err += 1
                    print(f"[{done}/{total}] ERR {profile.key} race_id={race_id}: {e}", flush=True)
    finally:
        conn.close()

    print(f"done ok={ok} err={err} total={ok+err}")


if __name__ == "__main__":
    main()
