#!/usr/bin/env python3
"""
結果確定済み（1着あり）の全JRAレースをAI予想し、all_race_ai_compareへ保存する。
"""
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.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 _mk_marks(by_acc: list[dict]) -> list[dict]:
    names = ["◎本命", "〇対抗", "▲単穴", "△連下"]
    out: list[dict] = []
    for i, row in enumerate(by_acc[:4]):
        out.append(
            {
                "mark": names[i],
                "horse_number": int(row["horse_number"]),
                "horse_name": str(row.get("horse_name") or ""),
                "score_accuracy": float(row.get("score_accuracy") or row.get("score") or 0),
                "win_odds": row.get("win_odds"),
                "stage3_prob": row.get("stage3_prob"),
            }
        )
    return out


def _upsert(conn, rec: dict) -> None:
    cols = [
        "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 != "race_id"])
    sql = f"""
        INSERT INTO all_race_ai_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("--lookback-runs", type=int, default=10)
    ap.add_argument("--pass-rate", type=float, default=0.5)
    ap.add_argument("--stage3-min-prob", type=float, default=0.02)
    ap.add_argument("--stage3-odds-cap", type=float, default=80.0)
    ap.add_argument("--stage3-aggressive-min-prob", type=float, default=0.0)
    ap.add_argument("--stage3-aggressive-odds-cap", type=float, default=300.0)
    ap.add_argument("--dry-run", action="store_true")
    args = ap.parse_args()

    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),
            )

        for i, t in enumerate(targets, start=1):
            race_id = int(t["race_id"])
            rec = {
                "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=args.lookback_runs,
                    pass_rate=args.pass_rate,
                    enable_stage1_5=True,
                    enable_stage2=True,
                    enable_stage3=True,
                    stage3_min_prob=args.stage3_min_prob,
                    stage3_odds_cap=args.stage3_odds_cap,
                    stage3_aggressive_min_prob=args.stage3_aggressive_min_prob,
                    stage3_aggressive_odds_cap=args.stage3_aggressive_odds_cap,
                )
                horses = payload.get("horses") or []
                by_acc = sorted(horses, key=lambda x: int(x.get("rank_accuracy") or 999))
                marks = _mk_marks(by_acc)

                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_accuracy"]
                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_accuracy"]
                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_accuracy"]
                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_accuracy"]

                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

                rec["ai_snapshot_json"] = json.dumps(
                    {
                        "race": payload.get("race"),
                        "config": payload.get("config"),
                        "summary": payload.get("summary"),
                        "marks": marks,
                        "horses_top5_accuracy": by_acc[:5],
                    },
                    ensure_ascii=False,
                    default=str,
                )
                if not args.dry_run:
                    _upsert(conn, rec)
                ok += 1
                print(f"[{i}/{len(targets)}] OK 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"[{i}/{len(targets)}] ERR race_id={race_id}: {e}", flush=True)

    finally:
        conn.close()
    print(f"done ok={ok} err={err} total={ok+err}")


if __name__ == "__main__":
    main()
