#!/usr/bin/env python3
"""
JRA公式 accessS.html から中央競馬の過去レース結果を取得し DB に保存。

例（5年分・終了日指定）:
  python3 collectors/fetch_jra_history.py --to-date 2026-03-29 --years 5

短い試験:
  python3 collectors/fetch_jra_history.py --from-date 2021-08-01 --to-date 2021-08-02 --place 4

事前に db/alter_jra_fetch.sql を実行すること（source_cname 列）。
"""
from __future__ import annotations

import argparse
import logging
import os
import re
import sys
from datetime import date, datetime, timedelta

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

from collectors.db_util import connect
from collectors.jra_http import fetch_access_s, get_session
from collectors.jra_parse import parse_race_result_html

logger = logging.getLogger(__name__)

TRACK_NAMES = {
    1: "札幌",
    2: "函館",
    3: "福島",
    4: "新潟",
    5: "東京",
    6: "中山",
    7: "中京",
    8: "京都",
    9: "阪神",
    10: "小倉",
}


def daterange(start: date, end: date):
    d = start
    while d <= end:
        yield d
        d += timedelta(days=1)


def build_mid_body(place_code: int, race_date: date, kai: int, nichi: int) -> str:
    y = race_date.year
    ds = race_date.strftime("%Y%m%d")
    return f"10{place_code:02d}{y}{kai:02d}{nichi:02d}01{ds}"


def race_hex(race_number: int) -> str:
    return f"{0xDF + race_number - 1:X}"


def cname_full(place_code: int, race_date: date, kai: int, nichi: int, race_number: int) -> str:
    mid = build_mid_body(place_code, race_date, kai, nichi)
    return f"pw01sde{mid}/{race_hex(race_number)}"


def parse_time_to_sec(raw: str):
    if not raw:
        return None
    raw = raw.strip()
    m = re.match(r"(\d):(\d{2})\.(\d)", raw)
    if m:
        return int(m.group(1)) * 60 + int(m.group(2)) + int(m.group(3)) / 10.0
    m = re.match(r"(\d):(\d{2})$", raw)
    if m:
        return int(m.group(1)) * 60 + int(m.group(2))
    return None


def load_track_ids(cur) -> dict[int, int]:
    cur.execute("SELECT id, code FROM tracks WHERE circuit='JRA'")
    return {int(r["code"]): r["id"] for r in cur.fetchall()}


def ensure_person(cur, table: str, name: str) -> int | None:
    name = (name or "").strip()
    if not name:
        return None
    cur.execute(f"SELECT id FROM {table} WHERE name=%s LIMIT 1", (name,))
    row = cur.fetchone()
    if row:
        return row["id"]
    cur.execute(f"INSERT INTO {table} (name) VALUES (%s)", (name,))
    return cur.lastrowid


def ensure_horse(cur, name: str) -> int:
    name = name.strip()
    cur.execute("SELECT id FROM horses WHERE name=%s LIMIT 1", (name,))
    row = cur.fetchone()
    if row:
        return row["id"]
    cur.execute("INSERT INTO horses (name) VALUES (%s)", (name,))
    return cur.lastrowid


def guess_grade(race_name: str | None) -> str | None:
    if not race_name:
        return None
    if "Ｇ１" in race_name or "G1" in race_name:
        return "G1"
    if "Ｇ２" in race_name or "G2" in race_name:
        return "G2"
    if "Ｇ３" in race_name or "G3" in race_name:
        return "G3"
    return None


def save_parsed(conn, cur, track_ids: dict[int, int], source_cname: str, pr) -> None:
    tn = pr.track_name
    place_code = None
    for code, nm in TRACK_NAMES.items():
        if nm == tn:
            place_code = code
            break
    if place_code is None:
        raise ValueError(f"unknown track name: {tn}")
    track_id = track_ids[place_code]

    cur.execute(
        "SELECT id FROM races WHERE source_cname=%s",
        (source_cname,),
    )
    ex = cur.fetchone()

    grade = guess_grade(pr.race_name)
    rname = pr.race_name[:128] if pr.race_name else None

    if ex:
        race_id = ex["id"]
        cur.execute(
            """UPDATE races SET weather=%s, going_turf=%s, going_dirt=%s, surface=%s,
               direction=%s, distance_m=%s, start_time=%s, head_count=%s, race_name=%s,
               grade=%s, result_fetched=1, data_source='jra_accessS', source_cname=%s,
               age_condition=%s, kai=%s, nichi=%s, race_number=%s
               WHERE id=%s""",
            (
                pr.weather,
                pr.going_turf,
                pr.going_dirt,
                pr.surface,
                pr.direction,
                pr.distance_m,
                pr.start_time,
                pr.head_count,
                rname,
                grade,
                source_cname,
                pr.age_condition,
                pr.kai,
                pr.nichi,
                pr.race_number,
                race_id,
            ),
        )
    else:
        cur.execute(
            """INSERT INTO races (
              race_date, track_id, circuit, kai, nichi, race_number, race_name, grade,
              age_condition, surface, direction, distance_m, weather, going_turf, going_dirt,
              head_count, start_time, data_source, result_fetched, source_cname
            ) VALUES (%s,%s,'JRA',%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,'jra_accessS',1,%s)""",
            (
                pr.race_date,
                track_id,
                pr.kai,
                pr.nichi,
                pr.race_number,
                rname,
                grade,
                pr.age_condition,
                pr.surface,
                pr.direction,
                pr.distance_m,
                pr.weather,
                pr.going_turf,
                pr.going_dirt,
                pr.head_count,
                pr.start_time,
                source_cname,
            ),
        )
        race_id = cur.lastrowid

    cur.execute("DELETE FROM race_results WHERE race_id=%s", (race_id,))
    cur.execute("DELETE FROM race_entries WHERE race_id=%s", (race_id,))

    for row in pr.result_rows:
        hid = ensure_horse(cur, row["horse_name"])
        jid = ensure_person(cur, "jockeys", row.get("jockey_name") or "")
        tid = ensure_person(cur, "trainers", row.get("trainer_name") or "")
        um = row["horse_number"]
        if not um:
            continue
        cur.execute(
            """INSERT INTO race_entries (
              race_id, horse_id, bracket_number, horse_number, jockey_id, trainer_id,
              popularity
            ) VALUES (%s,%s,%s,%s,%s,%s,%s)""",
            (
                race_id,
                hid,
                row.get("bracket"),
                um,
                jid,
                tid,
                row.get("popularity"),
            ),
        )
        rt = parse_time_to_sec(row.get("race_time_raw") or "")
        lf = row.get("last_3f")
        cur.execute(
            """INSERT INTO race_results (
              race_id, horse_id, finish_position, last_3f_time, race_time_seconds,
              margin, passing_order, win_time_raw
            ) VALUES (%s,%s,%s,%s,%s,%s,%s,%s)""",
            (
                race_id,
                hid,
                row["finish_position"],
                float(lf) if lf is not None else None,
                rt,
                row.get("margin"),
                row.get("passing_order"),
                row.get("race_time_raw"),
            ),
        )
    conn.commit()


def discover_and_fetch_day(
    session,
    conn,
    cur,
    track_ids: dict[int, int],
    race_date: date,
    place_code: int,
    max_kai: int,
    max_nichi: int,
    max_race: int,
    dry_run: bool,
    refresh_existing: bool,
) -> int:
    saved = 0
    for kai in range(1, max_kai + 1):
        for nichi in range(1, max_nichi + 1):
            c1 = cname_full(place_code, race_date, kai, nichi, 1)
            html, _ = fetch_access_s(session, c1)
            if not html or "パラメータエラー" in html:
                continue
            pr = parse_race_result_html(html)
            if not pr or pr.race_date != race_date.isoformat():
                continue

            for rnum in range(1, max_race + 1):
                sc = cname_full(place_code, race_date, kai, nichi, rnum)
                if rnum > 1:
                    html_r, _ = fetch_access_s(session, sc)
                    if not html_r or "パラメータエラー" in html_r:
                        break
                    pr = parse_race_result_html(html_r)
                    if not pr:
                        break

                cur.execute("SELECT id FROM races WHERE source_cname=%s", (sc,))
                if cur.fetchone() and not refresh_existing:
                    saved += 1
                    continue
                if dry_run:
                    logger.info("dry-run %s", sc)
                    saved += 1
                    continue
                save_parsed(conn, cur, track_ids, sc, pr)
                logger.info("saved %s %s", sc, pr.race_name)
                saved += 1
            return saved
    return saved


def main() -> None:
    logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
    ap = argparse.ArgumentParser()
    ap.add_argument("--to-date", default="2026-03-29")
    ap.add_argument("--years", type=int, default=5)
    ap.add_argument("--from-date", default=None)
    ap.add_argument("--place", type=int, default=None)
    ap.add_argument("--max-kai", type=int, default=12)
    ap.add_argument("--max-nichi", type=int, default=12)
    ap.add_argument("--max-race", type=int, default=12, help="通常は12R")
    ap.add_argument("--dry-run", action="store_true")
    ap.add_argument("--weekends-only", action="store_true")
    ap.add_argument("--refresh-existing", action="store_true", help="既存source_cnameも上書き更新")
    args = ap.parse_args()

    end_d = datetime.strptime(args.to_date, "%Y-%m-%d").date()
    if args.from_date:
        start_d = datetime.strptime(args.from_date, "%Y-%m-%d").date()
    else:
        start_d = end_d - timedelta(days=365 * args.years - 1)

    conn = connect()
    conn.autocommit(False)
    cur = conn.cursor()

    cur.execute(
        """SELECT COUNT(*) AS c FROM information_schema.COLUMNS
           WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'races' AND COLUMN_NAME = 'source_cname'"""
    )
    if cur.fetchone()["c"] == 0:
        logger.error("mysql < db/alter_jra_fetch.sql を実行してください")
        sys.exit(1)

    track_ids = load_track_ids(cur)
    session = get_session()
    total = 0
    places = [args.place] if args.place else list(range(1, 11))

    for d in daterange(start_d, end_d):
        if args.weekends_only and d.weekday() not in (5, 6):
            continue
        for p in places:
            try:
                total += discover_and_fetch_day(
                    session, conn, cur, track_ids, d, p,
                    args.max_kai, args.max_nichi, args.max_race, args.dry_run,
                    args.refresh_existing,
                )
            except Exception as e:
                logger.exception("error %s place %s: %s", d, p, e)
                conn.rollback()

    cur.close()
    conn.close()
    logger.info("done count=%s", total)


if __name__ == "__main__":
    main()
