#!/usr/bin/env python3
"""第3段階（妙味）パラメータのグリッド探索。全券種対応。"""
from __future__ import annotations

import argparse
import itertools
import json
import os
import re
import sys
from collections import defaultdict
from dataclasses import dataclass, field
from datetime import datetime

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from ai.stage1_screening import run_stage1  # noqa: E402
from collectors.db_util import connect  # noqa: E402

BET_TYPES_ALL = ["tansho", "fukusho", "wakuren", "wide", "umaren", "umatan", "sanrenpuku", "sanrentan"]
ORDERED_BET_TYPES = {"umatan", "sanrentan"}
UNORDERED_BET_TYPES = {"wakuren", "wide", "umaren", "sanrenpuku"}


@dataclass
class EvalResult:
    params: dict
    races: int
    bets: int
    hits: int
    staked_yen: int
    returned_yen: int
    per_type: dict[str, dict] = field(default_factory=dict)

    @property
    def hit_rate(self) -> float:
        return (self.hits / self.bets) if self.bets else 0.0

    @property
    def roi(self) -> float:
        return (self.returned_yen / self.staked_yen) if self.staked_yen else 0.0


def _parse_date(s: str) -> str:
    datetime.strptime(s, "%Y-%m-%d")
    return s


def _parse_bet_types(s: str) -> list[str]:
    if s.strip().lower() == "all":
        return BET_TYPES_ALL[:]
    items = [x.strip().lower() for x in s.split(",") if x.strip()]
    for it in items:
        if it not in BET_TYPES_ALL:
            raise ValueError(f"未対応 bet_type: {it}")
    return items


def _norm_combo(bet_type: str, combination: str) -> str | None:
    nums = [int(x) for x in re.findall(r"\d+", combination or "")]
    if not nums:
        return None
    if bet_type in {"tansho", "fukusho"}:
        return str(nums[0])
    if bet_type in ORDERED_BET_TYPES:
        need = 2 if bet_type == "umatan" else 3
        if len(nums) < need:
            return None
        return "-".join(str(x) for x in nums[:need])
    if bet_type in UNORDERED_BET_TYPES:
        need = 2 if bet_type in {"wakuren", "wide", "umaren"} else 3
        if len(nums) < need:
            return None
        use = sorted(nums[:need])
        return "-".join(str(x) for x in use)
    return None


def _fetch_candidate_race_ids(from_date: str, to_date: str, limit_races: int | None) -> list[int]:
    conn = connect()
    try:
        with conn.cursor() as cur:
            sql = """
                SELECT r.id AS race_id
                FROM races r
                INNER JOIN payouts p ON p.race_id = r.id AND p.bet_type = 'tansho'
                WHERE r.circuit = 'JRA'
                  AND r.race_date BETWEEN %s AND %s
                GROUP BY r.id
                ORDER BY r.race_date DESC, r.id DESC
            """
            if limit_races and limit_races > 0:
                sql += f" LIMIT {int(limit_races)}"
            cur.execute(sql, (from_date, to_date))
            return [int(r["race_id"]) for r in cur.fetchall()]
    finally:
        conn.close()


def _fetch_payout_maps(race_id: int, bet_types: list[str]) -> dict[str, dict[str, int]]:
    conn = connect()
    try:
        with conn.cursor() as cur:
            placeholders = ",".join(["%s"] * len(bet_types))
            cur.execute(
                f"""
                SELECT bet_type, combination, payout_yen
                FROM payouts
                WHERE race_id = %s
                  AND bet_type IN ({placeholders})
                """,
                (race_id, *bet_types),
            )
            out: dict[str, dict[str, int]] = {b: {} for b in bet_types}
            for row in cur.fetchall():
                bt = row["bet_type"]
                key = _norm_combo(bt, row.get("combination") or "")
                if key is None:
                    continue
                out[bt][key] = int(row.get("payout_yen") or 0)
            return out
    finally:
        conn.close()


def _fetch_bracket_map(race_id: int) -> dict[int, int]:
    conn = connect()
    try:
        with conn.cursor() as cur:
            cur.execute(
                """
                SELECT horse_number, bracket_number
                FROM race_entries
                WHERE race_id = %s
                  AND is_scratched = 0
                """,
                (race_id,),
            )
            out: dict[int, int] = {}
            for row in cur.fetchall():
                if row.get("horse_number") is None or row.get("bracket_number") is None:
                    continue
                out[int(row["horse_number"])] = int(row["bracket_number"])
            return out
    finally:
        conn.close()


def _top_horses(rows: list[dict], top_horses: int) -> list[dict]:
    base = sorted(rows, key=lambda x: int(x.get("rank_value") or 999999))
    return base[:top_horses]


def _generate_tickets_for_type(
    bet_type: str,
    rows: list[dict],
    bracket_map: dict[int, int],
    top_horses: int,
    tickets_per_type: int,
) -> list[tuple[str, float]]:
    picks = _top_horses(rows, top_horses)
    out: list[tuple[str, float]] = []
    by_combo: dict[str, float] = {}

    def add(combo: str, score: float) -> None:
        prev = by_combo.get(combo)
        if prev is None or score > prev:
            by_combo[combo] = score

    if bet_type in {"tansho", "fukusho"}:
        for h in picks:
            combo = str(int(h["horse_number"]))
            score = float(h.get("stage3_prob") or 0.0)
            add(combo, score)
    elif bet_type in {"umaren", "wide", "wakuren"}:
        for a, b in itertools.combinations(picks, 2):
            n1 = int(a["horse_number"])
            n2 = int(b["horse_number"])
            if bet_type == "wakuren":
                if n1 not in bracket_map or n2 not in bracket_map:
                    continue
                c1, c2 = sorted([bracket_map[n1], bracket_map[n2]])
            else:
                c1, c2 = sorted([n1, n2])
            combo = f"{c1}-{c2}"
            score = float(a.get("stage3_prob") or 0.0) * float(b.get("stage3_prob") or 0.0)
            add(combo, score)
    elif bet_type == "umatan":
        for a, b in itertools.permutations(picks, 2):
            n1 = int(a["horse_number"])
            n2 = int(b["horse_number"])
            combo = f"{n1}-{n2}"
            score = (
                float(a.get("stage3_prob") or 0.0)
                * float(b.get("stage3_prob") or 0.0)
                * (1.0 + max(0.0, (float(a.get("score_accuracy") or 0.0) - float(b.get("score_accuracy") or 0.0)) / 100.0))
            )
            add(combo, score)
    elif bet_type == "sanrenpuku":
        for a, b, c in itertools.combinations(picks, 3):
            nums = sorted([int(a["horse_number"]), int(b["horse_number"]), int(c["horse_number"])])
            combo = f"{nums[0]}-{nums[1]}-{nums[2]}"
            score = (
                float(a.get("stage3_prob") or 0.0)
                * float(b.get("stage3_prob") or 0.0)
                * float(c.get("stage3_prob") or 0.0)
            )
            add(combo, score)
    elif bet_type == "sanrentan":
        for a, b, c in itertools.permutations(picks, 3):
            combo = f"{int(a['horse_number'])}-{int(b['horse_number'])}-{int(c['horse_number'])}"
            score = (
                float(a.get("stage3_prob") or 0.0)
                * float(b.get("stage3_prob") or 0.0)
                * float(c.get("stage3_prob") or 0.0)
                * (1.0 + float(a.get("score_accuracy") or 0.0) / 120.0)
            )
            add(combo, score)

    out = sorted(by_combo.items(), key=lambda x: x[1], reverse=True)
    return out[:tickets_per_type]


def _evaluate_params(
    race_ids: list[int],
    *,
    min_prob: float,
    odds_cap: float,
    pass_rate: float,
    lookback_runs: int,
    stake_yen: int,
    disable_stage1_5: bool,
    disable_stage2: bool,
    model_path: str | None,
    model_meta_path: str | None,
    bet_types: list[str],
    top_horses: int,
    tickets_per_type: int,
) -> EvalResult:
    bets = 0
    hits = 0
    staked = 0
    returned = 0
    evaluated_races = 0
    per_type: dict[str, dict[str, int]] = defaultdict(lambda: {"bets": 0, "hits": 0, "staked_yen": 0, "returned_yen": 0})

    for rid in race_ids:
        payload = run_stage1(
            race_id=rid,
            lookback_runs=lookback_runs,
            pass_rate=pass_rate,
            enable_stage1_5=not disable_stage1_5,
            enable_stage2=not disable_stage2,
            enable_stage3=True,
            stage3_min_prob=min_prob,
            stage3_odds_cap=odds_cap,
            model_path=model_path,
            model_meta_path=model_meta_path,
        )
        horses = payload["horses"]
        if not horses:
            continue
        payout_maps = _fetch_payout_maps(rid, bet_types)
        if not any(payout_maps.values()):
            continue
        bracket_map = _fetch_bracket_map(rid)
        evaluated_races += 1

        for bt in bet_types:
            tickets = _generate_tickets_for_type(
                bt,
                horses,
                bracket_map=bracket_map,
                top_horses=top_horses,
                tickets_per_type=tickets_per_type,
            )
            paymap = payout_maps.get(bt, {})
            if not tickets:
                continue
            for combo, _score in tickets:
                bets += 1
                staked += stake_yen
                per_type[bt]["bets"] += 1
                per_type[bt]["staked_yen"] += stake_yen
                payoff = int(paymap.get(combo, 0))
                if payoff > 0:
                    hits += 1
                    won = int(round(payoff * (stake_yen / 100.0)))
                    returned += won
                    per_type[bt]["hits"] += 1
                    per_type[bt]["returned_yen"] += won

    per_type_out: dict[str, dict] = {}
    for bt, v in per_type.items():
        b = int(v["bets"])
        s = int(v["staked_yen"])
        r = int(v["returned_yen"])
        h = int(v["hits"])
        per_type_out[bt] = {
            "bets": b,
            "hits": h,
            "hit_rate": round((h / b) if b else 0.0, 4),
            "roi": round((r / s) if s else 0.0, 4),
            "staked_yen": s,
            "returned_yen": r,
        }

    return EvalResult(
        params={
            "stage3_min_prob": min_prob,
            "stage3_odds_cap": odds_cap,
            "pass_rate": pass_rate,
            "lookback_runs": lookback_runs,
            "stake_yen": stake_yen,
            "disable_stage1_5": disable_stage1_5,
            "disable_stage2": disable_stage2,
            "model_path": model_path,
            "model_meta_path": model_meta_path,
            "bet_types": bet_types,
            "top_horses": top_horses,
            "tickets_per_type": tickets_per_type,
        },
        races=evaluated_races,
        bets=bets,
        hits=hits,
        staked_yen=staked,
        returned_yen=returned,
        per_type=per_type_out,
    )


def main() -> None:
    p = argparse.ArgumentParser(description="第3段階（妙味）パラメータのグリッド探索（全券種対応）")
    p.add_argument("--from-date", type=_parse_date, required=True, help="YYYY-MM-DD")
    p.add_argument("--to-date", type=_parse_date, required=True, help="YYYY-MM-DD")
    p.add_argument("--limit-races", type=int, default=200, help="評価対象レース上限")
    p.add_argument("--lookback-runs", type=int, default=10)
    p.add_argument("--pass-rate", type=float, default=0.5)
    p.add_argument("--stake-yen", type=int, default=100)
    p.add_argument("--bet-types", type=str, default="all", help="all または tansho,umaren,...")
    p.add_argument("--top-horses", type=int, default=5, help="組合せ生成に使う上位頭数")
    p.add_argument("--tickets-per-type", type=int, default=3, help="券種ごとの購入点数")
    p.add_argument("--disable-stage1-5", action="store_true")
    p.add_argument("--disable-stage2", action="store_true")
    p.add_argument(
        "--model-path",
        type=str,
        default=None,
        help="Stage1 の LightGBM モデル（本番 p6 と揃える場合は stage1_lgbm_enhanced.txt 等）",
    )
    p.add_argument(
        "--model-meta-path",
        type=str,
        default=None,
        help="モデルメタ JSON（--model-path 指定時に推奨）",
    )
    p.add_argument("--min-prob-grid", type=str, default="0.00,0.01,0.02,0.03,0.05")
    p.add_argument("--odds-cap-grid", type=str, default="30,50,80,120")
    p.add_argument("--json", action="store_true")
    args = p.parse_args()

    if args.lookback_runs <= 0:
        raise SystemExit("--lookback-runs は 1 以上で指定してください")
    if not (0.0 < args.pass_rate <= 1.0):
        raise SystemExit("--pass-rate は 0.0 より大きく 1.0 以下で指定してください")
    if args.stake_yen <= 0:
        raise SystemExit("--stake-yen は 1 以上で指定してください")
    if args.top_horses < 3:
        raise SystemExit("--top-horses は 3 以上を推奨（3連系のため）")
    if args.tickets_per_type <= 0:
        raise SystemExit("--tickets-per-type は 1 以上で指定してください")

    bet_types = _parse_bet_types(args.bet_types)
    min_probs = [float(x.strip()) for x in args.min_prob_grid.split(",") if x.strip() != ""]
    odds_caps = [float(x.strip()) for x in args.odds_cap_grid.split(",") if x.strip() != ""]

    race_ids = _fetch_candidate_race_ids(args.from_date, args.to_date, args.limit_races)
    if not race_ids:
        raise SystemExit("評価対象レースが0件です（期間・データ取得状況を確認してください）")

    results: list[EvalResult] = []
    for min_prob, cap in itertools.product(min_probs, odds_caps):
        r = _evaluate_params(
            race_ids,
            min_prob=min_prob,
            odds_cap=cap,
            pass_rate=args.pass_rate,
            lookback_runs=args.lookback_runs,
            stake_yen=args.stake_yen,
            disable_stage1_5=args.disable_stage1_5,
            disable_stage2=args.disable_stage2,
            model_path=args.model_path,
            model_meta_path=args.model_meta_path,
            bet_types=bet_types,
            top_horses=args.top_horses,
            tickets_per_type=args.tickets_per_type,
        )
        results.append(r)

    results.sort(key=lambda x: (x.roi, x.hit_rate, x.races), reverse=True)
    best = results[0]

    if args.json:
        out = {
            "best": {
                **best.params,
                "races": best.races,
                "bets": best.bets,
                "hits": best.hits,
                "hit_rate": round(best.hit_rate, 4),
                "roi": round(best.roi, 4),
                "staked_yen": best.staked_yen,
                "returned_yen": best.returned_yen,
                "per_type": best.per_type,
            },
            "all": [
                {
                    **r.params,
                    "races": r.races,
                    "bets": r.bets,
                    "hits": r.hits,
                    "hit_rate": round(r.hit_rate, 4),
                    "roi": round(r.roi, 4),
                    "staked_yen": r.staked_yen,
                    "returned_yen": r.returned_yen,
                    "per_type": r.per_type,
                }
                for r in results
            ],
        }
        print(json.dumps(out, ensure_ascii=False, indent=2))
        return

    print(f"races: {len(race_ids)} (evaluated up to limit={args.limit_races})")
    print("best params:")
    print(json.dumps(best.params, ensure_ascii=False, indent=2))
    print(
        f"score: roi={best.roi:.4f} hit_rate={best.hit_rate:.4f} "
        f"bets={best.bets} hits={best.hits} staked={best.staked_yen} returned={best.returned_yen}"
    )
    if best.per_type:
        print("\nBest per type:")
        for bt in BET_TYPES_ALL:
            if bt in best.per_type:
                v = best.per_type[bt]
                print(f"{bt:10s} roi={v['roi']:.4f} hit={v['hit_rate']:.4f} bets={v['bets']}")
    print("\nTop 10:")
    for r in results[:10]:
        print(
            f"roi={r.roi:.4f} hit={r.hit_rate:.4f} races={r.races} "
            f"min_prob={r.params['stage3_min_prob']} cap={r.params['stage3_odds_cap']}"
        )


if __name__ == "__main__":
    main()

