"""
Export imds_mkistat_data month-by-month to SQL INSERT files.

Fetches only selected columns with a trading-session time filter,
one calendar month at a time, so the DB is never hit with one huge query.
"""

from __future__ import annotations

from datetime import datetime
from pathlib import Path

import pymysql
from pymysql.converters import escape_string

# =========================
# DATABASE CONFIG
# =========================
DB_CONFIG = {
    "host": "127.0.0.1",
    "port": 3306,
    "user": "root",
    "password": "",
    "database": "local_data",
    "charset": "utf8mb4",
    "cursorclass": pymysql.cursors.SSCursor,  # streaming cursor — low memory
}

TABLE_NAME = "imds_mkistat_data"
TS_COLUMN = "MKISTAT_LM_DATE_TIME"

COLUMNS = [
    "MKISTAT_LM_DATE_TIME",
    "MKISTAT_INSTRUMENT_CODE",
    "MKISTAT_HIGH_PRICE",
    "MKISTAT_LOW_PRICE",
    "MKISTAT_TOTAL_VOLUME",
    "MKISTAT_PUB_LAST_TRADED_PRICE",
    "MKISTAT_OPEN_PRICE",
]

# Overall range (inclusive start, exclusive end on the date bound)
RANGE_START = datetime(2026, 1, 1, 10, 0, 0)
RANGE_END = datetime(2026, 8, 7, 0, 0, 0)

# Intraday session filter applied every month
SESSION_TIME_START = "10:00:00"
SESSION_TIME_END = "14:30:00"

# Rows per INSERT statement (keeps each SQL statement manageable)
BATCH_SIZE = 1000

OUTPUT_DIR = Path(__file__).resolve().parent / "exports_mkistat"


def month_windows(start: datetime, end: datetime) -> list[tuple[datetime, datetime]]:
    """Split [start, end) into calendar-month windows clipped to the range."""
    windows: list[tuple[datetime, datetime]] = []
    year, month = start.year, start.month

    while True:
        month_start = datetime(year, month, 1)
        # First instant of next month
        if month == 12:
            month_end = datetime(year + 1, 1, 1)
        else:
            month_end = datetime(year, month + 1, 1)

        win_start = max(month_start, start)
        win_end = min(month_end, end)

        if win_start < win_end:
            windows.append((win_start, win_end))

        if month_end >= end:
            break

        if month == 12:
            year += 1
            month = 1
        else:
            month += 1

    return windows


def sql_literal(value) -> str:
    if value is None:
        return "NULL"
    if isinstance(value, datetime):
        return f"'{value.strftime('%Y-%m-%d %H:%M:%S')}'"
    if isinstance(value, (int, float)):
        return str(value)
    if isinstance(value, bytes):
        value = value.decode("utf-8", errors="replace")
    return f"'{escape_string(str(value))}'"


def write_insert_batch(fh, rows: list[tuple]) -> None:
    cols = ", ".join(f"`{c}`" for c in COLUMNS)
    values_sql = []
    for row in rows:
        literals = ", ".join(sql_literal(v) for v in row)
        values_sql.append(f"({literals})")
    fh.write(
        f"INSERT INTO `{TABLE_NAME}` ({cols}) VALUES\n"
        + ",\n".join(values_sql)
        + ";\n\n"
    )


def export_month(conn, win_start: datetime, win_end: datetime, out_path: Path) -> int:
    cols = ", ".join(f"`{c}`" for c in COLUMNS)
    sql = f"""
        SELECT {cols}
        FROM `{TABLE_NAME}`
        WHERE `{TS_COLUMN}` >= %s
          AND `{TS_COLUMN}` < %s
          AND TIME(`{TS_COLUMN}`) BETWEEN %s AND %s
    """

    params = (
        win_start.strftime("%Y-%m-%d %H:%M:%S"),
        win_end.strftime("%Y-%m-%d %H:%M:%S"),
        SESSION_TIME_START,
        SESSION_TIME_END,
    )

    print(f"  Query: {params[0]} <= {TS_COLUMN} < {params[1]}  (session {SESSION_TIME_START}-{SESSION_TIME_END})")

    total = 0

    with conn.cursor() as cursor, out_path.open("w", encoding="utf-8", newline="\n") as fh:
        fh.write(f"-- Export of `{TABLE_NAME}`\n")
        fh.write(f"-- Window: {params[0]} <= {TS_COLUMN} < {params[1]}\n")
        fh.write(f"-- Session TIME between {SESSION_TIME_START} and {SESSION_TIME_END}\n")
        fh.write(f"-- Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n\n")
        fh.write("SET NAMES utf8mb4;\n")
        fh.write("SET FOREIGN_KEY_CHECKS = 0;\n\n")

        cursor.execute(sql, params)

        while True:
            rows = cursor.fetchmany(BATCH_SIZE)
            if not rows:
                break
            write_insert_batch(fh, rows)
            total += len(rows)
            print(f"    ... {total:,} rows", end="\r")

        fh.write("SET FOREIGN_KEY_CHECKS = 1;\n")

    print(f"  Wrote {total:,} rows -> {out_path.name}")
    return total


def main() -> None:
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    windows = month_windows(RANGE_START, RANGE_END)

    print(f"Database: {DB_CONFIG['database']} @ {DB_CONFIG['host']}")
    print(f"Table:    {TABLE_NAME}")
    print(f"Range:    {RANGE_START} .. {RANGE_END}")
    print(f"Months:   {len(windows)}")
    print(f"Output:   {OUTPUT_DIR}\n")

    conn = pymysql.connect(**DB_CONFIG)
    try:
        grand_total = 0
        for win_start, win_end in windows:
            label = win_start.strftime("%Y_%m")
            out_path = OUTPUT_DIR / f"{TABLE_NAME}_{label}.sql"
            print(f"[{label}] {win_start} -> {win_end}")
            grand_total += export_month(conn, win_start, win_end, out_path)
        print(f"\nDone. Total rows exported: {grand_total:,}")
        print(f"SQL files in: {OUTPUT_DIR}")
    finally:
        conn.close()


if __name__ == "__main__":
    main()
