"""
DseRetrieve — async service for fetching incremental data from the DSE remote
MySQL database and tracking sync checkpoints in local PostgreSQL.

Mirrors the structure of the original synchronous DseRetrive class but:
  • Uses aiomysql for async MySQL queries (DSE remote source)
  • Uses SQLAlchemy AsyncSession for local PostgreSQL reads/writes
  • Returns the same data shapes so downstream insert logic is unchanged

Tables read from DSE MySQL:
  TRD      — trade summary         (full refresh every tick)
  IDX      — index snapshots       (incremental by IDX_DATE_TIME)
  MAN      — announcements         (incremental by MAN_ANNOUNCEMENT_DATE_TIME)
  MKISTAT  — market statistics     (incremental by MKISTAT_LM_DATE_TIME)

Local PostgreSQL:
  last_sync — checkpoint table (table_name → last_synced_timestamp)

Typical usage inside a scheduler service:

    from app.db.session import AsyncSessionLocal
    from app.services.dse_retrieve import DseRetrieve

    async def run():
        async with AsyncSessionLocal() as db:
            last_sync = await DseRetrieve.get_last_sync_date_time(db)

            other = await DseRetrieve.get_new_data_other(last_sync)
            mkistat = await DseRetrieve.get_new_data_mkistat(last_sync)

            # ... insert rows into imds_* tables ...

            if other["trd_sync_date_time"]:
                await DseRetrieve.update_last_sync_date_time(
                    db, "TRD", other["trd_sync_date_time"]
                )
            # ... update remaining checkpoints ...
            await db.commit()
"""

from __future__ import annotations

from datetime import datetime
from typing import Any

import aiomysql
import structlog
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession

from app.core.dse_client import dse_connection
from app.db.imds_models import LastSync

log = structlog.get_logger(__name__)

# ---------------------------------------------------------------------------
# Datetime coercion helpers
# ---------------------------------------------------------------------------
# Datetime columns that aiomysql may return as strings instead of datetime objects.
# Normalise them immediately after fetchall() so every downstream consumer
# always receives proper datetime instances — asyncpg rejects str for TIMESTAMP.

_DATETIME_FIELDS: dict[str, tuple[str, ...]] = {
    "TRD":     ("TRD_LM_DATE_TIME",),
    "IDX":     ("IDX_DATE_TIME",),
    "MAN":     ("MAN_ANNOUNCEMENT_DATE_TIME", "MAN_EXPIRY_DATE"),
    "MKISTAT": ("MKISTAT_LM_DATE_TIME",),
}

_DT_FMT = "%Y-%m-%d %H:%M:%S"


def _coerce_datetimes(rows: list[dict], table: str) -> list[dict]:
    """
    Parse string-typed datetime values in DSE rows to datetime objects.
    Mutates rows in place and returns the same list for zero-copy.
    """
    fields = _DATETIME_FIELDS.get(table, ())
    if not fields:
        return rows
    for row in rows:
        for field in fields:
            val = row.get(field)
            if isinstance(val, str) and val:
                try:
                    row[field] = datetime.strptime(val, _DT_FMT)
                except ValueError:
                    try:
                        row[field] = datetime.fromisoformat(val)
                    except ValueError:
                        row[field] = None
    return rows


class DseRetrieve:
    """
    Async re-implementation of the original DseRetrive class.

    All methods are static — no instance state needed.
    Pass an open AsyncSession for any method that touches local PostgreSQL.
    """

    # =======================================================================
    # Local PostgreSQL — sync checkpoint helpers
    # =======================================================================

    @staticmethod
    async def get_last_sync_date_time(db: AsyncSession) -> dict[str, datetime | None]:
        """
        Read last_sync checkpoints for all four tables.

        Returns:
            {
                "TRD":     <datetime | None>,
                "IDX":     <datetime | None>,
                "MAN":     <datetime | None>,
                "MKISTAT": <datetime | None>,
            }
        """
        result = await db.execute(
            select(LastSync).where(
                LastSync.table_name.in_(["TRD", "IDX", "MAN", "MKISTAT"])
            )
        )
        rows = result.scalars().all()
        return {row.table_name: row.last_synced_timestamp for row in rows}

    @staticmethod
    async def update_last_sync_date_time(
        db: AsyncSession,
        table_name: str,
        sync_date_time: datetime | str,
    ) -> bool:
        """
        Update the last_synced_timestamp checkpoint for one table.

        Args:
            db:             open AsyncSession (caller must commit)
            table_name:     "TRD" | "IDX" | "MAN" | "MKISTAT"
            sync_date_time: latest data timestamp from the fetched batch

        Returns True on success.
        """
        if isinstance(sync_date_time, str):
            sync_date_time = datetime.strptime(sync_date_time, "%Y-%m-%d %H:%M:%S")

        now = datetime.now()
        await db.execute(
            update(LastSync)
            .where(LastSync.table_name == table_name)
            .values(
                last_synced_timestamp=sync_date_time,
                last_synced_at=now,
            )
        )
        log.debug(
            "last_sync_updated",
            table=table_name,
            synced_to=str(sync_date_time),
        )
        return True

    # =======================================================================
    # DSE MySQL — data fetch helpers
    # =======================================================================

    @staticmethod
    async def get_new_data_other(
        last_sync_date_time: dict[str, datetime | None],
    ) -> dict[str, Any]:
        """
        Fetch TRD (full), IDX (incremental), MAN (incremental) from DSE MySQL.

        Mirrors get_new_data_other() from the original code.

        Returns:
            {
                "TRD":                  [dict, ...],
                "IDX":                  [dict, ...],
                "MAN":                  [dict, ...],
                "trd_sync_date_time":   str | None,
                "idx_sync_date_time":   str | None,
                "man_sync_date_time":   str | None,
            }
        """
        trd_data: list[dict] = []
        idx_data: list[dict] = []
        man_data: list[dict] = []
        sync_date_time_trd: str | None = None
        sync_date_time_idx: str | None = None
        sync_date_time_man: str | None = None

        async with dse_connection() as conn:
            async with conn.cursor(aiomysql.DictCursor) as cur:

                # --- TRD — full refresh (no WHERE clause; matches original) ---
                try:
                    await cur.execute("SELECT * FROM TRD WHERE `TRD_LM_DATE_TIME` > %s ORDER BY `TRD_LM_DATE_TIME` ASC", (last_sync_date_time.get("TRD"),))
                    trd_data = _coerce_datetimes(list(await cur.fetchall()), "TRD")
                    if trd_data:
                        lm = trd_data[-1].get("TRD_LM_DATE_TIME")
                        sync_date_time_trd = (
                            lm.strftime("%Y-%m-%d %H:%M:%S")
                            if hasattr(lm, "strftime")
                            else str(lm)
                        )
                    # log.info("dse_fetch_trd", rows=len(trd_data))
                except aiomysql.Error as exc:
                    log.error("dse_fetch_trd_error", error=str(exc))
                    trd_data = []
                    sync_date_time_trd = None

                # --- IDX — incremental by IDX_DATE_TIME ----------------------
                try:
                    sync_dt = last_sync_date_time.get("IDX")
                    await cur.execute(
                        "SELECT * FROM IDX "
                        "WHERE `IDX_DATE_TIME` > %s "
                        "ORDER BY `IDX_DATE_TIME` ASC",
                        (sync_dt,),
                    )
                    idx_data = _coerce_datetimes(list(await cur.fetchall()), "IDX")
                    if idx_data:
                        lm = idx_data[-1].get("IDX_DATE_TIME")
                        sync_date_time_idx = (
                            lm.strftime("%Y-%m-%d %H:%M:%S")
                            if hasattr(lm, "strftime")
                            else str(lm)
                        )
                    # log.info("dse_fetch_idx", rows=len(idx_data), since=str(sync_dt))
                except aiomysql.Error as exc:
                    log.error("dse_fetch_idx_error", error=str(exc))
                    idx_data = []
                    sync_date_time_idx = None

                # --- MAN — incremental by MAN_ANNOUNCEMENT_DATE_TIME ---------
                try:
                    sync_dt = last_sync_date_time.get("MAN")
                    await cur.execute(
                        "SELECT * FROM MAN "
                        "WHERE `MAN_ANNOUNCEMENT_DATE_TIME` > %s "
                        "ORDER BY `MAN_ANNOUNCEMENT_DATE_TIME` ASC",
                        (sync_dt,),
                    )
                    man_data = _coerce_datetimes(list(await cur.fetchall()), "MAN")
                    if man_data:
                        lm = man_data[-1].get("MAN_ANNOUNCEMENT_DATE_TIME")
                        sync_date_time_man = (
                            lm.strftime("%Y-%m-%d %H:%M:%S")
                            if hasattr(lm, "strftime")
                            else str(lm)
                        )
                    # log.info("dse_fetch_man", rows=len(man_data), since=str(sync_dt))
                except aiomysql.Error as exc:
                    log.error("dse_fetch_man_error", error=str(exc))
                    man_data = []
                    sync_date_time_man = None

        return {
            "TRD": trd_data,
            "IDX": idx_data,
            "MAN": man_data,
            "trd_sync_date_time": sync_date_time_trd,
            "idx_sync_date_time": sync_date_time_idx,
            "man_sync_date_time": sync_date_time_man,
        }

    @staticmethod
    async def get_new_data_mkistat(
        last_sync_date_time: dict[str, datetime | None],
    ) -> dict[str, Any]:
        """
        Fetch MKISTAT (incremental by MKISTAT_LM_DATE_TIME) from DSE MySQL.

        Mirrors get_new_data_mkistat() from the original code.

        Returns:
            {
                "MKISTAT":               [dict, ...],
                "mkistat_sync_date_time": str,
            }
        """
        mkistat_data: list[dict] = []
        sync_date_time_mkistat: str = datetime.now().strftime("%Y-%m-%d %H:%M:%S")

        async with dse_connection() as conn:
            async with conn.cursor(aiomysql.DictCursor) as cur:
                try:
                    sync_dt = last_sync_date_time.get("MKISTAT")
                    await cur.execute(
                        "SELECT * FROM MKISTAT "
                        "WHERE `MKISTAT_LM_DATE_TIME` > %s "
                        "ORDER BY `MKISTAT_LM_DATE_TIME` ASC",
                        (sync_dt,),
                    )
                    mkistat_data = _coerce_datetimes(list(await cur.fetchall()), "MKISTAT")
                    if mkistat_data:
                        lm = mkistat_data[-1].get("MKISTAT_LM_DATE_TIME")
                        sync_date_time_mkistat = (
                            lm.strftime("%Y-%m-%d %H:%M:%S")
                            if hasattr(lm, "strftime")
                            else str(lm)
                        )
                    else:
                        # No new data — keep checkpoint at now so next run
                        # doesn't re-scan old data unnecessarily.
                        sync_date_time_mkistat = datetime.now().strftime(
                            "%Y-%m-%d %H:%M:%S"
                        )
                    # log.info(
                    #     "dse_fetch_mkistat",
                    #     rows=len(mkistat_data),
                    #     since=str(sync_dt),
                    # )
                except aiomysql.Error as exc:
                    log.error("dse_fetch_mkistat_error", error=str(exc))
                    mkistat_data = []

        return {
            "MKISTAT": mkistat_data,
            "mkistat_sync_date_time": sync_date_time_mkistat,
        }

    # =======================================================================
    # Pure calculation helper (no I/O)
    # =======================================================================

    @staticmethod
    def calc_yesterday_change_pct(row: dict) -> float | None:
        """
        Calculate percentage change vs. yesterday's close.

        Uses MKISTAT_PUB_LAST_TRADED_PRICE and MKISTAT_YDAY_CLOSE_PRICE.
        Returns 0 when either value is None/zero (matches original behaviour).
        Returns None on unexpected type errors.
        """
        close = row.get("MKISTAT_PUB_LAST_TRADED_PRICE")
        yday  = row.get("MKISTAT_YDAY_CLOSE_PRICE")

        if close is None or yday is None or yday == 0 or close == 0:
            return 0.0
        try:
            return round((float(close) - float(yday)) / float(yday) * 100, 2)
        except (TypeError, ZeroDivisionError):
            return None
