"""第1段階向け特徴量定義と学習データ生成。"""
from __future__ import annotations

from collections.abc import Iterable
from datetime import date


def _safe_rate(numerator: int, denominator: int) -> float:
    if denominator <= 0:
        return 0.0
    return float(numerator) / float(denominator)


# 特徴ベクトルに含めてはいけない列（当該レースの結果・識別子・ラベル）
_FORBIDDEN_FEATURE_NAMES = frozenset(
    {
        "finish_position",
        "race_id",
        "horse_id",
        "race_date",
        "label_win",
        "label_top3",
        "label_top5",
    },
)


def get_target_column_names() -> list[str]:
    """学習可能な二値ラベル列（build_training_frame が生成）。"""
    return ["label_win", "label_top3", "label_top5"]


def validate_training_frame_columns(feature_cols: list[str]) -> None:
    """特徴にリーク・ラベル・識別子列が混ざっていないか検査。違反時は ValueError。"""
    for c in feature_cols:
        if c in _FORBIDDEN_FEATURE_NAMES:
            raise ValueError(f"特徴量にリーク疑いの列が含まれます: {c}")


def get_feature_columns() -> list[str]:
    """第1段階（基礎能力足切り）で使う列名（LightGBM 入力と一致させる）。"""
    return [
        "recent_runs",
        "recent_win_rate",
        "recent_top3_rate",
        "recent_avg_finish",
        "recent_best_finish",
        "recent_avg_time",
        "recent_graded_top5_rate",
        "experience_score",
        "form_momentum",
        "same_surface_top3_rate",
        "distance_band_top3_rate",
        "same_track_top3_rate",
        "recent_avg_last_3f",
        "carry_weight_kg",
        "distance_m",
        "surface_turf",
    ]


def _avg_finish_positions(rows: list[dict]) -> float | None:
    if not rows:
        return None
    s = 0.0
    n = 0
    for row in rows:
        p = int(row.get("finish_position") or 0)
        if p > 0:
            s += float(p)
            n += 1
    return s / n if n else None


def _compute_form_momentum(rows: list[dict]) -> float:
    """直近ブロックとその前のブロックの平均着順差から勢いを -1..1 で表す（正＝直近が良化）。"""
    if len(rows) >= 6:
        recent = rows[:3]
        older = rows[3:6]
    elif len(rows) >= 4:
        recent = rows[:2]
        older = rows[2:4]
    else:
        return 0.0
    a_recent = _avg_finish_positions(recent)
    a_older = _avg_finish_positions(older)
    if a_recent is None or a_older is None:
        return 0.0
    raw = a_older - a_recent
    return max(-1.0, min(1.0, raw / 8.0))


def _compute_same_surface_top3(
    rows: list[dict],
    target_surface: str | None,
    neutral: float,
) -> float:
    ts = (target_surface or "").strip()
    if not ts:
        return neutral
    same = [r for r in rows if (r.get("surface") or "").strip() == ts]
    if not same:
        return neutral
    t3 = sum(1 for r in same if int(r.get("finish_position") or 99) <= 3)
    return _safe_rate(t3, len(same))


def _compute_same_track_top3(
    rows: list[dict],
    target_track_code: str | None,
    neutral: float,
) -> float:
    tc = (target_track_code or "").strip()
    if not tc:
        return neutral
    same = [r for r in rows if (r.get("track_code") or "").strip() == tc]
    if not same:
        return neutral
    t3 = sum(1 for r in same if int(r.get("finish_position") or 99) <= 3)
    return _safe_rate(t3, len(same))


def _compute_recent_avg_last_3f(rows: list[dict]) -> float:
    """過去走の上がり3F（秒）の平均。欠損は除外。無ければ 0。"""
    vals: list[float] = []
    for row in rows:
        v = row.get("last_3f_time")
        if v is not None:
            vals.append(float(v))
    return sum(vals) / len(vals) if vals else 0.0


def _compute_distance_band_top3(
    rows: list[dict],
    target_distance_m: float | None,
    neutral: float,
    band_m: float = 200.0,
) -> float:
    if target_distance_m is None or float(target_distance_m) <= 0:
        return neutral
    td = float(target_distance_m)
    near = [
        r
        for r in rows
        if r.get("distance_m") is not None
        and abs(float(r["distance_m"]) - td) <= band_m
    ]
    if not near:
        return neutral
    t3 = sum(1 for r in near if int(r.get("finish_position") or 99) <= 3)
    return _safe_rate(t3, len(near))


def compute_stage1_features(
    history_rows: Iterable[dict],
    *,
    target_surface: str | None = None,
    target_distance_m: float | None = None,
    target_track_code: str | None = None,
    carry_weight_kg: float | None = None,
) -> dict[str, float]:
    rows = list(history_rows)
    runs = len(rows)
    wins = 0
    top3 = 0
    graded_top5 = 0
    finish_sum = 0.0
    best_finish = 99
    time_sum = 0.0
    time_cnt = 0

    for row in rows:
        pos = int(row.get("finish_position") or 0)
        if pos <= 0:
            continue
        finish_sum += pos
        best_finish = min(best_finish, pos)
        if pos == 1:
            wins += 1
        if pos <= 3:
            top3 += 1
        grade = (row.get("grade") or "").upper()
        if grade in {"G1", "G2", "G3"} and pos <= 5:
            graded_top5 += 1
        race_time = row.get("race_time_seconds")
        if race_time is not None:
            time_sum += float(race_time)
            time_cnt += 1

    avg_finish = (finish_sum / runs) if runs else 99.0
    avg_time = (time_sum / time_cnt) if time_cnt else 0.0
    top3_rate = _safe_rate(top3, runs)
    neutral_ctx = max(0.05, min(0.85, top3_rate))

    cw = float(carry_weight_kg) if carry_weight_kg is not None else 0.0

    return {
        "recent_runs": float(runs),
        "recent_win_rate": _safe_rate(wins, runs),
        "recent_top3_rate": top3_rate,
        "recent_avg_finish": avg_finish,
        "recent_best_finish": float(best_finish if best_finish != 99 else 99),
        "recent_avg_time": avg_time,
        "recent_graded_top5_rate": _safe_rate(graded_top5, runs),
        "experience_score": min(runs, 10) / 10.0,
        "race_level_rescue": 1.0 if graded_top5 > 0 else 0.0,
        "form_momentum": _compute_form_momentum(rows),
        "same_surface_top3_rate": _compute_same_surface_top3(rows, target_surface, neutral_ctx),
        "distance_band_top3_rate": _compute_distance_band_top3(rows, target_distance_m, neutral_ctx),
        "same_track_top3_rate": _compute_same_track_top3(rows, target_track_code, neutral_ctx),
        "recent_avg_last_3f": round(_compute_recent_avg_last_3f(rows), 3),
        "carry_weight_kg": cw,
    }


def build_training_frame(conn, lookback_runs: int = 10, limit_rows: int | None = None):
    """race_entries/race_results から第1段階用の学習DataFrameを生成する。"""
    import pandas as pd

    base_sql = """
        SELECT
            re.race_id,
            re.horse_id,
            re.horse_number,
            re.carry_weight,
            r.race_date,
            r.distance_m,
            r.surface,
            t.code AS track_code,
            rr.finish_position
        FROM race_entries re
        INNER JOIN races r ON r.id = re.race_id
        INNER JOIN tracks t ON t.id = r.track_id
        INNER JOIN race_results rr ON rr.race_id = re.race_id AND rr.horse_id = re.horse_id
        WHERE r.circuit = 'JRA'
          AND rr.finish_position > 0
        ORDER BY r.race_date ASC, re.race_id ASC, re.horse_number ASC
    """
    if limit_rows and limit_rows > 0:
        base_sql += f" LIMIT {int(limit_rows)}"

    output_rows: list[dict] = []
    with conn.cursor() as cur:
        cur.execute(base_sql)
        entries = cur.fetchall()

        history_sql = """
            SELECT
                rr.finish_position,
                rr.race_time_seconds,
                rr.last_3f_time,
                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.circuit = 'JRA'
              AND r.race_date < %s
              AND rr.finish_position > 0
            ORDER BY r.race_date DESC, rr.race_id DESC
            LIMIT %s
        """
        for entry in entries:
            cur.execute(history_sql, (entry["horse_id"], entry["race_date"], int(lookback_runs)))
            metrics = compute_stage1_features(
                cur.fetchall(),
                target_surface=str(entry.get("surface") or ""),
                target_distance_m=float(entry.get("distance_m") or 0) or None,
                target_track_code=str(entry.get("track_code") or ""),
                carry_weight_kg=float(entry.get("carry_weight") or 0) or None,
            )
            fp = int(entry.get("finish_position") or 99)
            output_rows.append(
                {
                    **metrics,
                    "distance_m": float(entry.get("distance_m") or 0),
                    "surface_turf": 1.0 if (entry.get("surface") or "") == "芝" else 0.0,
                    "label_win": 1 if fp == 1 else 0,
                    "label_top3": 1 if fp <= 3 else 0,
                    "label_top5": 1 if fp <= 5 else 0,
                    "race_id": int(entry["race_id"]),
                    "horse_id": int(entry["horse_id"]),
                    "race_date": entry["race_date"] if isinstance(entry["race_date"], date) else None,
                }
            )
    return pd.DataFrame(output_rows)
