#!/usr/bin/env python3.11
"""
JV-Link（JVLinkServer 経由）から取得したレコードを jv_raw_records に保存する。

環境変数:
  JVLINK_SERVER_HOST  Windows上のJVLinkServerのIP（例: 192.168.1.10）
  JVLINK_SERVER_PORT  既定 8765
  JVLINK_SID           JRA-VAN Data Lab. のソフトウェアID（JVInit 用）。未設定だと API が
                       「Service key is not set」で失敗する。Windows の JVLinkServer.exe も
                       起動時に --sid で同じ値を渡すこと。

例:
  export JVLINK_SERVER_HOST=192.168.x.x
  python3.11 collectors/fetch_jv_raw.py --from-datetime 20210329 --dataspec RACE

  # 蓄積系（option=1）の全 DataSpec を順番に取得し、続けて週次専用（TCOV/RCOV 系）も取得
  python3.11 collectors/fetch_jv_raw.py --from-datetime 20210329 --all-stored

  # オッズを含む RACE の差分（蓄積 option=1）。from_datetime は「更新日時がこの日以降」
  python3.11 collectors/fetch_jv_raw.py --from-datetime 20260501 --weekly-odds

Linux から JVLinkServer に届かない場合は、Windows から ssh -R で 8765 を転送し、本サーバー上で:
  export JVLINK_SERVER_HOST=127.0.0.1
  python3.11 collectors/fetch_jv_raw.py --linux-tunnel-tcp ...
（pyjvlink は Linux+localhost で HTTP /health を要求するため、トンネル時は --linux-tunnel-tcp が必要なことがある）
詳細: docs/ops_jv_link_windows_fetch.md
"""
from __future__ import annotations

import argparse
import asyncio
import json
import os
import platform
import sys
import time
from datetime import datetime

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

import config  # noqa: F401  # loads .env and config/jv_local.env

from collectors.db_util import connect  # noqa: E402


def _normalize_jv_time(s: str) -> str:
    """JVOpen fromtime を 14 桁に揃える（Manus/JRA-VAN SDK: YYYYMMDDhhmmss または範囲）。

    例: 20200101 → 20200101000000
        20200101000000-20210101000000 はそのまま（年単位 setup 推奨）
    """
    s = s.strip()
    if "-" in s:
        start, end = s.split("-", 1)
        return f"{_normalize_jv_time(start)}-{_normalize_jv_time(end)}"
    if len(s) == 8 and s.isdigit():
        return s + "000000"
    return s


def _env_nonempty(key: str, default: str) -> str:
    """環境変数が空文字だけのときも default に落とす（.env に KEY= の行があるケース）。"""
    v = os.environ.get(key)
    if v is None:
        return default
    v2 = v.strip()
    return v2 if v2 else default


def _apply_pyjvlink_linux_ssh_tunnel_compat() -> None:
    """
    Linux 上で SSH -R により 127.0.0.1:8765 が開いている場合、pyjvlink 既定の
    HTTP /health 判定が JV-Link 実装と合わず失敗することがある。TCP 接続のみで
    起動待ちする（JVLINK_LINUX_TUNNEL_TCP=1 または --linux-tunnel-tcp）。
    """
    v = os.environ.get("JVLINK_LINUX_TUNNEL_TCP", "").strip().lower()
    if v not in ("1", "true", "yes"):
        return
    if platform.system() == "Windows":
        return

    import socket

    from pyjvlink._internal.runtime import process_manager as _pm
    from pyjvlink.errors import JVTimeoutError

    orig_probe = _pm.ProcessManager._probe_server
    orig_is_running = _pm.ProcessManager._is_server_running
    orig_wait = _pm.ProcessManager._wait_for_server

    def _tcp_open(host: str, port: int, timeout: float = 3.0) -> bool:
        try:
            with socket.create_connection((host, port), timeout=timeout):
                return True
        except OSError:
            return False

    def _patched_probe_server(self: _pm.ProcessManager):
        if self.config.host in ("127.0.0.1", "localhost") and _tcp_open(
            self.config.host, self.config.port
        ):
            return True, True, {"status": "healthy"}
        return orig_probe(self)

    def _patched_is_server_running(self: _pm.ProcessManager) -> bool:
        if self.config.host in ("127.0.0.1", "localhost"):
            return _tcp_open(self.config.host, self.config.port)
        return orig_is_running(self)

    async def _patched_wait_for_server(self: _pm.ProcessManager) -> None:
        start_time = time.time()
        while time.time() - start_time < self.config.startup_timeout:
            if _tcp_open(self.config.host, self.config.port, timeout=2.0):
                return
            await asyncio.sleep(0.5)
        raise JVTimeoutError(
            f"JVLinkServer TCP wait timeout after {self.config.startup_timeout} seconds (JV_LINUX_TUNNEL_TCP)"
        )

    _pm.ProcessManager._probe_server = _patched_probe_server  # type: ignore[method-assign]
    _pm.ProcessManager._is_server_running = _patched_is_server_running  # type: ignore[method-assign]
    _pm.ProcessManager._wait_for_server = _patched_wait_for_server  # type: ignore[method-assign]


def _serialize(obj) -> str:
    if hasattr(obj, "model_dump"):
        return json.dumps(obj.model_dump(), ensure_ascii=False, default=str)
    if isinstance(obj, dict):
        return json.dumps(obj, ensure_ascii=False, default=str)
    try:
        return json.dumps(vars(obj), ensure_ascii=False, default=str)
    except TypeError:
        return json.dumps(str(obj), ensure_ascii=False)


def _accumulated_dataspecs() -> list[str]:
    from pyjvlink.types import QueryOption, VALID_DATASPECS_BY_OPTION

    specs = sorted(VALID_DATASPECS_BY_OPTION[int(QueryOption.ACCUMULATED)], key=lambda s: s.value)
    return [s.value for s in specs]


def _weekly_only_dataspecs() -> list[str]:
    """option=2 のみで取得できる補完系（蓄積オプション1の対象外）。"""
    from pyjvlink.types import QueryOption, VALID_DATASPECS_BY_OPTION

    w = VALID_DATASPECS_BY_OPTION[int(QueryOption.WEEKLY)]
    a = VALID_DATASPECS_BY_OPTION[int(QueryOption.ACCUMULATED)]
    only = sorted(w - a, key=lambda s: s.value)
    return [s.value for s in only]


async def _stream_to_db(
    client,
    *,
    dataspec: str,
    from_datetime: str,
    option: int,
    max_records: int,
    include_raw: bool,
    cur,
    to_date: str | None = None,
    record_types: list[str] | None = None,
    stream_read_timeout: int | None = None,
) -> int:
    saved = 0
    result = await client.query_stored_raw(
        dataspec=dataspec,
        from_datetime=from_datetime,
        to_date=to_date,
        option=option,
        max_records=max_records,
        include_raw=include_raw,
        record_types=record_types,
        stream_read_timeout=stream_read_timeout,
    )
    async for envelope in result.records:
        rtype = getattr(envelope, "type", None) or "?"
        rec = getattr(envelope, "record", None)
        raw_dict = getattr(envelope, "raw", None) if include_raw else None
        rawb = json.dumps(raw_dict, ensure_ascii=False).encode("utf-8") if raw_dict else None
        key_text = None
        if isinstance(rec, dict):
            for k in ("race_id", "id", "key"):
                if k in rec:
                    key_text = str(rec[k])[:64]
                    break
        pj = _serialize(rec) if rec is not None else "{}"
        cur.execute(
            """INSERT INTO jv_raw_records (dataspec, record_type, key_text, payload_json, payload_raw)
               VALUES (%s,%s,%s,%s,%s)""",
            (dataspec, str(rtype)[:8], key_text, pj, rawb),
        )
        saved += 1
        if saved == 1 or saved % 500 == 0:
            print(f"  ... saved {saved} ({dataspec})", flush=True)
    return saved


def _log_import_start(cur, label: str) -> int:
    cur.execute(
        """INSERT INTO import_logs (source, started_at, status, message)
           VALUES (%s, NOW(), 'running', %s)""",
        ("jv_link_stored", label),
    )
    return int(cur.lastrowid)


def _log_import_finish(cur, log_id: int, status: str, rows: int, message: str | None) -> None:
    cur.execute(
        """UPDATE import_logs SET finished_at = NOW(), status = %s, rows_inserted = %s, message = %s
           WHERE id = %s""",
        (status, rows, message, log_id),
    )


async def fetch_and_store(
    host: str,
    port: int,
    from_datetime: str,
    dataspec: str,
    max_records: int,
    include_raw: bool,
    *,
    sid: str | None = None,
    option: int | None = None,
    to_date: str | None = None,
    busy_retries: int = 12,
    busy_sleep_sec: float = 15.0,
    import_log: bool = True,
    record_types: list[str] | None = None,
) -> int:
    try:
        from pyjvlink import (
            Client,
            JVBusyError,
            JVNoDataError,
            JVServerConfig,
            JVServerError,
            QueryOption,
        )
    except ImportError:
        print("pip install pyjvlink（Python 3.11+）", file=sys.stderr)
        raise SystemExit(1)

    if option is None:
        option = int(QueryOption.ACCUMULATED)

    from_datetime = _normalize_jv_time(from_datetime)
    # to_date は 8 桁 YYYYMMDD のまま pyjvlink に渡す（end_of_day=23:59:59 になる）。
    # 14 桁に 000000 付与すると終了が当日 00:00 になり setup レンジが壊れる。
    if to_date:
        to_date = to_date.strip()

    stream_to = int(_env_nonempty("JVLINK_STREAM_READ_TIMEOUT", "0"))
    if stream_to <= 0:
        # setup(option 3/4) の jv_open は 5 分以上かかることがある（既定 300s だと COM timeout）
        stream_to = 7200 if option in (3, 4) else 600
    cfg_kw: dict = {"host": host, "port": port, "timeout": max(600, stream_to), "stream_read_timeout": stream_to}
    if sid is not None:
        cfg_kw["sid"] = sid
    cfg = JVServerConfig(**cfg_kw)
    conn = connect()
    conn.autocommit(True)
    cur = conn.cursor()
    log_id: int | None = None
    label = f"option={option} dataspec={dataspec}"
    if to_date:
        label += f" to_date={to_date}"
    if record_types:
        label += f" record_types={','.join(record_types)}"
    if import_log:
        log_id = _log_import_start(cur, label)
    saved = 0
    last_err: str | None = None

    try:
        for attempt in range(busy_retries + 1):
            try:
                async with Client(cfg) as client:
                    saved = await _stream_to_db(
                        client,
                        dataspec=dataspec,
                        from_datetime=from_datetime,
                        option=option,
                        max_records=max_records,
                        include_raw=include_raw,
                        cur=cur,
                        to_date=to_date,
                        record_types=record_types,
                        stream_read_timeout=stream_to,
                    )
                last_err = None
                break
            except JVNoDataError:
                saved = 0
                last_err = None
                break
            except JVBusyError as e:
                last_err = str(e)
                ra = getattr(e, "retry_after", None)
                wait = float(ra) if ra else busy_sleep_sec
                if attempt >= busy_retries:
                    raise
                print(f"JVBusyError (attempt {attempt + 1}/{busy_retries}): {e}; sleep {wait}s", file=sys.stderr)
                await asyncio.sleep(wait)
            except JVServerError as e:
                last_err = str(e)
                raise
    except Exception as e:
        if log_id is not None:
            _log_import_finish(cur, log_id, "fail", saved, last_err or str(e))
        raise
    else:
        if log_id is not None:
            _log_import_finish(cur, log_id, "success", saved, None)
    finally:
        cur.close()
        conn.close()

    return saved


async def fetch_all_stored(
    host: str,
    port: int,
    from_datetime: str,
    max_records: int,
    include_raw: bool,
    *,
    sid: str | None = None,
    skip_weekly_extra: bool,
    busy_retries: int,
    busy_sleep_sec: float,
) -> int:
    from pyjvlink.types import QueryOption

    total = 0
    phases: list[tuple[int, list[str]]] = [
        (int(QueryOption.ACCUMULATED), _accumulated_dataspecs()),
    ]
    if not skip_weekly_extra:
        extras = _weekly_only_dataspecs()
        if extras:
            phases.append((int(QueryOption.WEEKLY), extras))

    for opt, specs in phases:
        for ds in specs:
            print(f"--- fetching option={opt} dataspec={ds} ---", flush=True)
            n = await fetch_and_store(
                host,
                port,
                from_datetime,
                ds,
                max_records,
                include_raw,
                sid=sid,
                option=opt,
                busy_retries=busy_retries,
                busy_sleep_sec=busy_sleep_sec,
                import_log=True,
            )
            print(f"saved {n} ({ds})", flush=True)
            total += n
            await asyncio.sleep(0.5)
    return total


async def fetch_weekly_odds(
    host: str,
    port: int,
    from_datetime: str,
    max_records: int,
    include_raw: bool,
    *,
    sid: str | None = None,
    busy_retries: int,
    busy_sleep_sec: float,
) -> int:
    """オッズ(O1-O6)を含む RACE 更新分を取得する。

    JV-Link 実装では「今週 option=2 の RACE」ストリームにオッズレコードが載らないことが多く、
    `record_types=['O1']` のような絞り込みは **0 件**になる。
    オッズ更新は **蓄積 option=1 の RACE** に載るため、ここでは **option=1・dataspec=RACE・全レコード種**
    を **1 回**だけ取り込む（`from_datetime` 以降に更新された RA/SE/O1…がまとめて返る）。

    名前 ``weekly-odds`` は cron 互換のため残す（中身は「オッズを拾うための RACE 差分」）。
    """
    from pyjvlink.types import QueryOption

    opt = int(QueryOption.ACCUMULATED)
    print(
        f"--- fetching option={opt} dataspec=RACE (incremental; includes O1-O6 when JV published) ---",
        flush=True,
    )
    n = await fetch_and_store(
        host,
        port,
        from_datetime,
        "RACE",
        max_records,
        include_raw,
        sid=sid,
        option=opt,
        busy_retries=busy_retries,
        busy_sleep_sec=busy_sleep_sec,
        import_log=True,
        record_types=None,
    )
    print(f"saved {n} (RACE accumulated)", flush=True)
    return n


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--from-datetime", required=True, help="YYYYMMDD")
    ap.add_argument(
        "--to-date",
        default=None,
        help="YYYYMMDD（option=3/4 セットアップ時の開催終了日。未指定時は JV 既定）",
    )
    ap.add_argument("--dataspec", default="RACE")
    ap.add_argument(
        "--option",
        type=int,
        default=None,
        choices=(1, 2, 3, 4),
        help="JVOpen option（1=蓄積 2=今週 3/4=セットアップ）。未指定時は single モードで 1",
    )
    ap.add_argument("--all-stored", action="store_true", help="蓄積系すべて＋週次専用を順に取得")
    ap.add_argument(
        "--weekly-odds",
        action="store_true",
        help="オッズを含む RACE を蓄積(option=1)で1回取得（from_datetime 以降の更新。--all-stored と併用不可）",
    )
    ap.add_argument(
        "--skip-weekly-extra",
        action="store_true",
        help="--all-stored 時、TCOV/RCOV 系（option=2 のみ）を省略",
    )
    ap.add_argument("--host", default=_env_nonempty("JVLINK_SERVER_HOST", "127.0.0.1"))
    ap.add_argument("--port", type=int, default=int(_env_nonempty("JVLINK_SERVER_PORT", "8765")))
    ap.add_argument(
        "--sid",
        default=None,
        help="JVInit 用ソフトウェアID（未指定時は環境変数 JVLINK_SID。Windows の JVLinkServer.exe の --sid と揃える）",
    )
    ap.add_argument("--max-records", type=int, default=-1, help="-1 で無制限")
    ap.add_argument("--no-raw", action="store_true")
    ap.add_argument("--no-import-log", action="store_true", help="import_logs に書かない")
    ap.add_argument("--busy-retries", type=int, default=12)
    ap.add_argument("--busy-sleep", type=float, default=15.0)
    ap.add_argument(
        "--linux-tunnel-tcp",
        action="store_true",
        help="Linux で SSH -R トンネル使用時: HTTP /health ではなく TCP で JVLinkServer を判定",
    )
    args = ap.parse_args()

    if args.linux_tunnel_tcp:
        os.environ["JVLINK_LINUX_TUNNEL_TCP"] = "1"
    _apply_pyjvlink_linux_ssh_tunnel_compat()

    cli_sid = (args.sid or "").strip()
    env_sid = (os.environ.get("JVLINK_SID") or "").strip()
    resolved_sid = cli_sid or env_sid or None

    if args.all_stored and args.weekly_odds:
        raise SystemExit("--all-stored と --weekly-odds は同時に指定できません")

    if args.all_stored:
        n = asyncio.run(
            fetch_all_stored(
                args.host,
                args.port,
                args.from_datetime,
                args.max_records,
                not args.no_raw,
                sid=resolved_sid,
                skip_weekly_extra=args.skip_weekly_extra,
                busy_retries=args.busy_retries,
                busy_sleep_sec=args.busy_sleep,
            )
        )
        print("total_saved", n)
        return

    if args.weekly_odds:
        if args.option is not None and int(args.option) != 1:
            print(
                "WARN: --weekly-odds はオッズ取得のため蓄積 option=1 の RACE を使います（指定の --option は無視されます）",
                file=sys.stderr,
            )
        n = asyncio.run(
            fetch_weekly_odds(
                args.host,
                args.port,
                args.from_datetime,
                args.max_records,
                not args.no_raw,
                sid=resolved_sid,
                busy_retries=args.busy_retries,
                busy_sleep_sec=args.busy_sleep,
            )
        )
        print("weekly_odds_total_saved", n)
        return

    opt = int(args.option) if args.option is not None else None
    n = asyncio.run(
        fetch_and_store(
            args.host,
            args.port,
            args.from_datetime,
            args.dataspec,
            args.max_records,
            not args.no_raw,
            sid=resolved_sid,
            option=opt,
            to_date=args.to_date,
            busy_retries=args.busy_retries,
            busy_sleep_sec=args.busy_sleep,
            import_log=not args.no_import_log,
        )
    )
    to_part = f" to={args.to_date}" if args.to_date else ""
    print(f"saved {n} ({args.dataspec} option={opt or 1} from={args.from_datetime}{to_part})", flush=True)


if __name__ == "__main__":
    main()
