#!/usr/bin/env python3
"""
STEP 3: 1レース・馬単位で stage1 の入力データ（過去走・特徴量）をダンプする。
旧サーバ / 新本番で同じコマンド → 出力を diff。

例:
  python3 scripts/diagnose_stage1_horse_data_one_race.py \\
    --race-date 2026-01-04 --track-code 06 --race-number 1 \\
    --out tmp/stage1_horse_163.txt
"""
from __future__ import annotations

import argparse
import os
import socket
import sys

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

from ai.feature_engineering import compute_stage1_features, get_feature_columns
from ai.prediction_profiles import get_profiles
from ai.stage1_screening import _load_model_cached, _model_win_prob, run_stage1
from collectors.db_util import connect

DEFAULT_MODEL = "ai/models/stage1_lgbm_enhanced.txt"
DEFAULT_META = "ai/models/stage1_lgbm_enhanced.meta.json"


def _resolve_race(cur, race_date: str, track_code: str, race_number: int) -> dict | None:
    cur.execute(
        """
        SELECT r.id, r.race_date, r.distance_m, r.surface, t.code AS track_code, t.name AS track_name
        FROM races r
        JOIN tracks t ON t.id = r.track_id
        WHERE r.race_date = %s AND t.code = %s AND r.race_number = %s AND r.circuit = 'JRA'
        LIMIT 1
        """,
        (race_date, track_code, race_number),
    )
    return cur.fetchone()


def main() -> None:
    ap = argparse.ArgumentParser(description="1レース・馬単位 stage1 入力データ診断")
    ap.add_argument("--race-date", required=True)
    ap.add_argument("--track-code", required=True)
    ap.add_argument("--race-number", type=int, required=True)
    ap.add_argument("--model-path", default=DEFAULT_MODEL)
    ap.add_argument("--model-meta-path", default=DEFAULT_META)
    ap.add_argument("--out", default=None)
    args = ap.parse_args()

    lines: list[str] = []

    def out(s: str = "") -> None:
        lines.append(s)

    prof = get_profiles("p6_latest")[0]
    feature_cols = get_feature_columns()

    out("=== stage1 馬単位データ診断 ===")
    out(f"host={socket.gethostname()}")
    out(f"race={args.race_date} track={args.track_code} R={args.race_number}")
    out(f"profile=p6_latest -> {prof.key} lookback={prof.lookback_runs} pipeline={prof.forecast_pipeline}")
    out()

    conn = connect()
    try:
        with conn.cursor() as cur:
            race = _resolve_race(cur, args.race_date, args.track_code, args.race_number)
            if not race:
                out("ERROR: レースなし")
                text = "\n".join(lines) + "\n"
                if args.out:
                    with open(args.out, "w", encoding="utf-8") as f:
                        f.write(text)
                else:
                    sys.stdout.write(text)
                sys.exit(1)

            race_id = int(race["id"])
            out(f"race_id={race_id} surface={race.get('surface')} distance_m={race.get('distance_m')}")
            out()

            cur.execute(
                """
                SELECT re.horse_id, re.horse_number, re.carry_weight, h.name AS horse_name
                FROM race_entries re
                JOIN horses h ON h.id = re.horse_id
                WHERE re.race_id = %s AND re.is_scratched = 0
                ORDER BY re.horse_number ASC
                """,
                (race_id,),
            )
            entries = list(cur.fetchall())

            history_sql = """
                SELECT
                    rr.finish_position,
                    rr.race_time_seconds,
                    rr.last_3f_time,
                    rr.passing_order,
                    r.grade,
                    r.race_date,
                    r.distance_m,
                    r.surface,
                    t.code AS track_code
                FROM race_results rr
                INNER JOIN races r ON r.id = rr.race_id
                INNER JOIN tracks t ON t.id = r.track_id
                WHERE rr.horse_id = %s
                  AND r.race_date < %s
                  AND r.circuit = 'JRA'
                  AND rr.finish_position > 0
                ORDER BY r.race_date DESC, rr.race_id DESC
                LIMIT %s
            """

            models_weights, model_features = _load_model_cached(args.model_path, args.model_meta_path)
            surface_turf = 1.0 if (race.get("surface") or "") == "芝" else 0.0
            dist_m = float(race.get("distance_m") or 0.0)

            for entry in entries:
                hn = entry.get("horse_number")
                hid = int(entry["horse_id"])
                out(f"--- horse_number={hn} horse_id={hid} name={entry.get('horse_name')} ---")
                cur.execute(
                    history_sql,
                    (hid, race["race_date"], int(prof.lookback_runs)),
                )
                history = list(cur.fetchall())
                out(f"history_count={len(history)} (limit={prof.lookback_runs})")
                out("idx\trace_date\ttrack\tsurface\tdist\tfinish\tlast_3f")
                for i, h in enumerate(history, 1):
                    rd = h.get("race_date")
                    rd_s = rd.isoformat() if hasattr(rd, "isoformat") else str(rd)
                    out(
                        f"{i}\t{rd_s}\t{h.get('track_code')}\t{h.get('surface')}\t"
                        f"{h.get('distance_m')}\t{h.get('finish_position')}\t{h.get('last_3f_time')}"
                    )
                metrics = compute_stage1_features(
                    history,
                    target_surface=str(race.get("surface") or ""),
                    target_distance_m=float(race.get("distance_m") or 0) or None,
                    target_track_code=str(race.get("track_code") or ""),
                    carry_weight_kg=float(entry.get("carry_weight") or 0) or None,
                )
                mwp = _model_win_prob(
                    models_weights, model_features, metrics, dist_m, surface_turf
                )
                out("features:")
                for col in feature_cols:
                    out(f"  {col}={metrics.get(col)}")
                out(f"model_win_prob={mwp}")
                out()

        payload = run_stage1(
            race_id=race_id,
            lookback_runs=prof.lookback_runs,
            pass_rate=prof.pass_rate,
            enable_stage1_5=prof.enable_stage1_5,
            enable_stage2=prof.enable_stage2,
            enable_stage3=prof.enable_stage3,
            stage3_min_prob=prof.stage3_min_prob,
            stage3_odds_cap=prof.stage3_odds_cap,
            stage3_aggressive_min_prob=prof.stage3_aggressive_min_prob,
            stage3_aggressive_odds_cap=prof.stage3_aggressive_odds_cap,
            model_path=args.model_path,
            model_meta_path=args.model_meta_path,
            forecast_pipeline=prof.forecast_pipeline,
            rank_mode=prof.rank_mode,
        )
        out("--- stage1 出力（score_accuracy 降順）---")
        out("horse_no\tscore_accuracy\tmodel_win_prob\tstage2_bonus")
        for row in sorted(
            payload.get("horses") or [],
            key=lambda x: float(x.get("score_accuracy") or 0),
            reverse=True,
        ):
            out(
                f"{row.get('horse_number')}\t{row.get('score_accuracy')}\t"
                f"{row.get('model_win_prob')}\t{row.get('stage2_bonus')}"
            )
    finally:
        conn.close()

    text = "\n".join(lines) + "\n"
    if args.out:
        with open(args.out, "w", encoding="utf-8") as f:
            f.write(text)
        print(f"wrote {args.out}", flush=True)
    else:
        sys.stdout.write(text)


if __name__ == "__main__":
    main()
