#!/usr/bin/env python3
"""create_tables.sql と seed_tracks.sql を pymysql で順に実行。環境変数 MYSQL_PWD 必須。"""
from __future__ import annotations

import os
import sys

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

try:
    import pymysql
except ImportError:
    print("pip install pymysql", file=sys.stderr)
    raise SystemExit(1)

PW = os.environ.get("MYSQL_PWD") or os.environ.get("KEIBA_MYSQL_ROOT_PASSWORD", "")


def split_statements(sql: str) -> list[str]:
    parts: list[str] = []
    buf: list[str] = []
    for line in sql.splitlines():
        s = line.strip()
        if not s or s.startswith("--"):
            continue
        buf.append(line)
        if s.endswith(";"):
            stmt = "\n".join(buf).rstrip()[:-1].strip()
            buf = []
            if stmt:
                parts.append(stmt)
    return parts


def main() -> None:
    if not PW:
        sys.exit("export MYSQL_PWD=... を設定して実行してください")
    conn = pymysql.connect(
        host=os.environ.get("KEIBA_DB_HOST", "localhost"),
        user="root",
        password=PW,
        charset="utf8mb4",
    )
    try:
        conn.autocommit(True)
        cur = conn.cursor()
        for rel in ("db/create_tables.sql", "db/seed_tracks.sql"):
            path = os.path.join(ROOT, rel)
            with open(path, encoding="utf-8") as f:
                body = f.read()
            for stmt in split_statements(body):
                cur.execute(stmt)
            print("OK", rel)
        cur.close()
    finally:
        conn.close()


if __name__ == "__main__":
    main()
