"""
Service 1 — DSE full data sync (MKISTAT + MAN + IDX + TRD) with error logging.

Every cycle writes a row to `scheduler_job_log` so ops can audit any
sync problem without scraping terminal logs.

Stages per cycle:
  1  — read last_sync checkpoints from PostgreSQL
  2  — fetch all four tables from DSE MySQL
  3  — load sector map (TTL-cached 8 h)
  4a — enrich + bulk insert MKISTAT
  4b — bulk insert MAN
  4c — bulk insert IDX
  4d — bulk insert TRD  (full snapshot, no WHERE)
  5  — update last_sync for each table that had data
  6  — commit

Each stage is independent — a failure in one table does NOT prevent the
others from running. Only the failed table's last_sync stays unchanged.

Query errors in PostgreSQL:
  SELECT * FROM scheduler_job_log
  WHERE job_id = 'service_1' AND status = 'error'
  ORDER BY started_at DESC;
"""

from __future__ import annotations

import time
import traceback
from datetime import datetime, timezone
from typing import Any

import structlog
from sqlalchemy.ext.asyncio import AsyncSession

from app.db.imds_repository import (
    IdxRecord,
    ManRecord,
    MkistatRecord,
    TrdRecord,
    bulk_insert_idx,
    bulk_insert_mkistat,
    bulk_insert_trd,
    get_security_sector_map,
    insert_man,
)
from app.db.models import SchedulerJobLog
from app.db.session import AsyncSessionLocal
from app.services.dse_retrieve import DseRetrieve

log = structlog.get_logger(__name__)

_JOB_ID = "service_1"


# ===========================================================================
# Public entry point
# ===========================================================================

async def run() -> None:
    """Scheduler entry point — called every 5 s within the active window."""
    started_at = datetime.now(timezone.utc)
    tick_start  = time.perf_counter()

    status:    str       = "success"
    error_msg: str | None = None
    meta:      dict[str, Any] = {}

    async with AsyncSessionLocal() as db:
        try:
            meta = await _sync_cycle(db, tick_start)
            await db.commit()

        except ConnectionError as exc:
            status    = "error"
            error_msg = f"DSE connection unavailable: {exc}"
            meta["stage_failed"] = "dse_fetch"
            meta["error"]        = str(exc)
            log.error(
                "service_1_dse_unavailable",
                error=str(exc),
                hint="Check dse_host/dse_user/dse_password in .env",
            )

        except _StageError as exc:
            # A known stage raised a labelled error — already logged in-situ
            status    = "error"
            error_msg = str(exc)
            meta.update(exc.meta)

        except Exception:
            status    = "error"
            error_msg = traceback.format_exc()
            meta["stage_failed"] = "unknown"
            log.exception("service_1_unhandled_error")
            await db.rollback()

        finally:
            duration_ms = int((time.perf_counter() - tick_start) * 1000)
            await _write_job_log(
                started_at  = started_at,
                status      = status,
                duration_ms = duration_ms,
                error_msg   = error_msg,
                metadata    = meta,
            )


# ===========================================================================
# Sync cycle — each stage has its own try/except
# ===========================================================================

async def _sync_cycle(
    db: AsyncSession,
    tick_start: float,
) -> dict[str, Any]:
    """
    Execute all sync stages and return a metadata dict for the job log.
    Raises _StageError on any unrecoverable stage failure.
    """
    meta: dict[str, Any] = {}

    # ------------------------------------------------------------------
    # Stage 1: read last_sync checkpoints
    # ------------------------------------------------------------------
    try:
        last_sync = await DseRetrieve.get_last_sync_date_time(db)
        meta["checkpoints"] = {k: str(v) for k, v in last_sync.items()}
        # log.debug("service_1_stage1_ok", checkpoints=meta["checkpoints"])

    except Exception as exc:
        _raise_stage("checkpoint_read", exc, meta,
                     hint="Check PostgreSQL connection and last_sync table.")

    # ------------------------------------------------------------------
    # Stage 2: fetch all four tables from DSE MySQL
    # ------------------------------------------------------------------
    try:
        new_data_other   = await DseRetrieve.get_new_data_other(last_sync)
        new_data_mkistat = await DseRetrieve.get_new_data_mkistat(last_sync)

        fetch_elapsed = time.perf_counter() - tick_start
        mkistat_rows = new_data_mkistat["MKISTAT"]
        man_rows     = new_data_other["MAN"]
        idx_rows     = new_data_other["IDX"]
        trd_rows     = new_data_other["TRD"]

        meta.update({
            "fetched_mkistat_rows": len(mkistat_rows),
            "fetched_man_rows":     len(man_rows),
            "fetched_idx_rows":     len(idx_rows),
            "fetched_trd_rows":     len(trd_rows),
            "fetch_elapsed_s":      round(fetch_elapsed, 3),
        })
        # log.info(
        #     "service_1_stage2_ok",
        #     mkistat=len(mkistat_rows), man=len(man_rows),
        #     idx=len(idx_rows), trd=len(trd_rows),
        #     elapsed_s=meta["fetch_elapsed_s"],
        # )

    except ConnectionError:
        raise
    except Exception as exc:
        _raise_stage("dse_fetch", exc, meta,
                     hint="DSE MySQL query failed. Check DSE server availability.")

    now_dt     = datetime.now()
    store_date = now_dt.date()
    store_time = now_dt.time()

    # Tracks which tables were actually inserted so we update last_sync correctly
    checkpoint_updates: dict[str, str | None] = {}
    checkpoint_errors:  list[str] = []

    # ------------------------------------------------------------------
    # Stage 3: load sector map (TTL-cached, non-fatal)
    # ------------------------------------------------------------------
    try:
        sector_map = await get_security_sector_map(db)
        meta["sector_map_size"] = len(sector_map)
    except Exception as exc:
        sector_map = {}
        meta["sector_map_warning"] = str(exc)
        # log.warning("service_1_sector_map_failed", error=str(exc),
        #             hint="Rows will store sector='Unknown'.")

    # ------------------------------------------------------------------
    # Stage 4a: enrich + bulk insert MKISTAT
    # ------------------------------------------------------------------
    if mkistat_rows:
        try:
            records:      list[MkistatRecord] = []
            skipped_rows: list[dict]          = []

            for i, row in enumerate(mkistat_rows):
                try:
                    instrument_code = _safe_str(row.get("MKISTAT_INSTRUMENT_CODE"))
                    records.append(MkistatRecord(
                        mkistat_change_yday_close = _safe_float(row.get("MKISTAT_CHANGE_YDAY_CLOSE")),
                        mkistat_avg_trade_share   = _safe_float(row.get("MKISTAT_AVG_TRADE_SHARE")),
                        mkistat_adv_shares        = _safe_float(row.get("MKISTAT_ADV_SHARES")),
                        mkistat_adv_value         = _safe_float(row.get("MKISTAT_ADV_VALUE")),
                        mkistat_change_volume     = _safe_float(row.get("MKISTAT_CHANGE_VOLUME")),
                        mkistat_pvr               = _safe_float(row.get("MKISTAT_PVR")),
                        mkistat_std_dev_60min     = _safe_float(row.get("MKISTAT_STD_DEV_60MIN")),
                        mkistat_vwap_daily        = _safe_float(row.get("MKISTAT_VWAP_DAILY")),
                        mkistat_vwap_1h           = _safe_float(row.get("MKISTAT_VWAP_1H")),
                        mkistat_sector               = sector_map.get(instrument_code or "", "Unknown"),
                        mkistat_yesterday_change_pct = DseRetrieve.calc_yesterday_change_pct(row),
                        mkistat_instrument_code        = instrument_code,
                        mkistat_instrument_number      = _safe_str(row.get("MKISTAT_INSTRUMENT_NUMBER")),
                        mkistat_quote_bases            = _safe_str(row.get("MKISTAT_QUOTE_BASES")),
                        mkistat_open_price             = row.get("MKISTAT_OPEN_PRICE"),
                        mkistat_pub_last_traded_price  = row.get("MKISTAT_PUB_LAST_TRADED_PRICE"),
                        mkistat_spot_last_traded_price = row.get("MKISTAT_SPOT_LAST_TRADED_PRICE"),
                        mkistat_high_price             = row.get("MKISTAT_HIGH_PRICE"),
                        mkistat_low_price              = row.get("MKISTAT_LOW_PRICE"),
                        mkistat_close_price            = row.get("MKISTAT_CLOSE_PRICE"),
                        mkistat_yday_close_price       = row.get("MKISTAT_YDAY_CLOSE_PRICE"),
                        mkistat_total_trades           = row.get("MKISTAT_TOTAL_TRADES"),
                        mkistat_total_volume           = row.get("MKISTAT_TOTAL_VOLUME"),
                        mkistat_total_value            = row.get("MKISTAT_TOTAL_VALUE"),
                        mkistat_public_total_trades    = row.get("MKISTAT_PUBLIC_TOTAL_TRADES"),
                        mkistat_public_total_volume    = row.get("MKISTAT_PUBLIC_TOTAL_VOLUME"),
                        mkistat_public_total_value     = row.get("MKISTAT_PUBLIC_TOTAL_VALUE"),
                        mkistat_spot_total_trades      = row.get("MKISTAT_SPOT_TOTAL_TRADES"),
                        mkistat_spot_total_volume      = row.get("MKISTAT_SPOT_TOTAL_VOLUME"),
                        mkistat_spot_total_value       = row.get("MKISTAT_SPOT_TOTAL_VALUE"),
                        mkistat_lm_date_time           = _safe_datetime(row.get("MKISTAT_LM_DATE_TIME")),
                        mkistat_store_date             = store_date,
                        mkistat_store_time             = store_time,
                    ))
                except Exception as row_exc:
                    skipped_rows.append({"row_index": i, "error": str(row_exc)})
                    # log.warning("service_1_mkistat_row_skipped",
                    #             row_index=i, error=str(row_exc))

            meta["mkistat_enriched"] = len(records)
            meta["mkistat_skipped"]  = len(skipped_rows)

            if records:
                inserted = await bulk_insert_mkistat(db, records)
                meta["mkistat_inserted"] = len(inserted)
                checkpoint_updates["MKISTAT"] = new_data_mkistat.get("mkistat_sync_date_time")
                # log.info("service_1_mkistat_ok", inserted=len(inserted), skipped=len(skipped_rows))
            else:
                pass
                # log.warning("service_1_mkistat_all_skipped", total=len(mkistat_rows))

        except Exception as exc:
            meta["mkistat_error"] = str(exc)
            log.error("service_1_mkistat_insert_failed", error=str(exc),
                      traceback=traceback.format_exc())
    else:
        pass
        # log.info("service_1_mkistat_no_new_data")


    # ------------------------------------------------------------------
    # Stage 4b: insert MAN (announcements)
    # ------------------------------------------------------------------
    if man_rows:
        try:
            man_inserted = 0
            man_skipped  = 0
            for i, row in enumerate(man_rows):
                try:
                    await insert_man(db, ManRecord(
                        man_announcement_date_time = _safe_datetime(row.get("MAN_ANNOUNCEMENT_DATE_TIME")),
                        man_announcement_prefix    = _safe_str(row.get("MAN_ANNOUNCEMENT_PREFIX")),
                        man_announcement           = row.get("MAN_ANNOUNCEMENT"),
                        man_expiry_date            = row.get("MAN_EXPIRY_DATE"),
                        man_store_date             = store_date,
                        man_store_time             = store_time,
                    ))
                    man_inserted += 1
                except Exception as row_exc:
                    man_skipped += 1
                    # log.warning("service_1_man_row_skipped",
                    #             row_index=i, error=str(row_exc))

            meta["man_inserted"] = man_inserted
            meta["man_skipped"]  = man_skipped
            if man_inserted:
                checkpoint_updates["MAN"] = new_data_other.get("man_sync_date_time")
            # log.info("service_1_man_ok", inserted=man_inserted, skipped=man_skipped)

        except Exception as exc:
            meta["man_error"] = str(exc)
            log.error("service_1_man_insert_failed", error=str(exc),
                      traceback=traceback.format_exc())
    else:
        pass
        # log.info("service_1_man_no_new_data")

    # ------------------------------------------------------------------
    # Stage 4c: insert IDX (index snapshots)
    # ------------------------------------------------------------------
    if idx_rows:
        try:
            idx_records = []
            idx_skipped = 0
            for i, row in enumerate(idx_rows):
                try:
                    idx_records.append(IdxRecord(
                        idx_index_id             = _safe_str(row.get("IDX_INDEX_ID")),
                        idx_date_time            = _safe_datetime(row.get("IDX_DATE_TIME")),
                        idx_capital_value        = row.get("IDX_CAPITAL_VALUE"),
                        idx_deviation            = row.get("IDX_DEVIATION"),
                        # idx_percentage_deviation = row.get("IDX_PERCENTAGE_DEVIATION"),
                        idx_percentage_deviation = row.get("lDX_PERCENTAGE_DEVIATION"),
                        idx_store_date           = store_date,
                        idx_store_time           = store_time,
                    ))
                except Exception as row_exc:
                    idx_skipped += 1
                    # log.warning("service_1_idx_row_skipped",
                    #             row_index=i, error=str(row_exc))

            if idx_records:
                inserted_idx = await bulk_insert_idx(db, idx_records)
                meta["idx_inserted"] = len(inserted_idx)
                meta["idx_skipped"]  = idx_skipped
                checkpoint_updates["IDX"] = new_data_other.get("idx_sync_date_time")
                # log.info("service_1_idx_ok", inserted=len(inserted_idx), skipped=idx_skipped)

        except Exception as exc:
            meta["idx_error"] = str(exc)
            log.error("service_1_idx_insert_failed", error=str(exc),
                      traceback=traceback.format_exc())
    else:
        pass
        # log.info("service_1_idx_no_new_data")

    # ------------------------------------------------------------------
    # Stage 4d: insert TRD (trade summary — full snapshot, one row)
    # ------------------------------------------------------------------
    if trd_rows:
        try:
            trd_records = []
            for i, row in enumerate(trd_rows):
                try:
                    trd_records.append(TrdRecord(
                        trd_sno          = row.get("TRD_SNO"),
                        trd_total_trades = row.get("TRD_TOTAL_TRADES"),
                        trd_total_volume = row.get("TRD_TOTAL_VOLUME"),
                        trd_total_value  = row.get("TRD_TOTAL_VALUE"),
                        trd_mkt_status   = _safe_str(row.get("TRD_MKT_STATUS")),
                        trd_lm_date_time = _safe_datetime(row.get("TRD_LM_DATE_TIME")),
                        trd_store_date   = store_date,
                        trd_store_time   = store_time,
                    ))
                except Exception as row_exc:
                    pass
                    # log.warning("service_1_trd_row_skipped",
                    #             row_index=i, error=str(row_exc))

            if trd_records:
                inserted_trd = await bulk_insert_trd(db, trd_records)
                meta["trd_inserted"] = len(inserted_trd)
                checkpoint_updates["TRD"] = new_data_other.get("trd_sync_date_time")
                # log.info("service_1_trd_ok", inserted=len(inserted_trd))

        except Exception as exc:
            meta["trd_error"] = str(exc)
            log.error("service_1_trd_insert_failed", error=str(exc),
                      traceback=traceback.format_exc())
    else:
        pass
        # log.info("service_1_trd_no_new_data")

    # ------------------------------------------------------------------
    # Stage 5: update last_sync for every table that was inserted
    # ------------------------------------------------------------------
    for table, ts in checkpoint_updates.items():
        if ts is None:
            continue
        try:
            await DseRetrieve.update_last_sync_date_time(db, table, ts)
            # log.debug("service_1_checkpoint_updated", table=table, ts=ts)
        except Exception as exc:
            checkpoint_errors.append(f"{table}: {exc}")
            log.error("service_1_checkpoint_update_failed",
                      table=table, ts=ts, error=str(exc))

    if checkpoint_errors:
        meta["checkpoint_errors"] = checkpoint_errors

    # ------------------------------------------------------------------
    # Final summary
    # ------------------------------------------------------------------
    total_elapsed           = time.perf_counter() - tick_start
    meta["total_elapsed_s"] = round(total_elapsed, 3)
    meta["result"]          = "success" if not checkpoint_errors else "success_with_checkpoint_errors"

    log.info(
        "service_1_cycle_complete",
        mkistat  = meta.get("mkistat_inserted", 0),
        man      = meta.get("man_inserted", 0),
        idx      = meta.get("idx_inserted", 0),
        trd      = meta.get("trd_inserted", 0),
        elapsed  = meta["total_elapsed_s"],
        result   = meta["result"],
    )
    return meta


# ===========================================================================
# Persistent job log writer (uses a SEPARATE session so rollback on main
# session never suppresses the error record)
# ===========================================================================

async def _write_job_log(
    started_at:  datetime,
    status:      str,
    duration_ms: int,
    error_msg:   str | None,
    metadata:    dict,
) -> None:
    try:
        async with AsyncSessionLocal() as log_db:
            log_db.add(
                SchedulerJobLog(
                    job_id      = _JOB_ID,
                    status      = status,
                    duration_ms = duration_ms,
                    error_msg   = error_msg,
                    job_metadata= metadata,
                    started_at  = started_at,
                )
            )
            await log_db.commit()
    except Exception:
        # Log writer must NEVER crash the scheduler
        log.exception("service_1_job_log_write_failed")


# ===========================================================================
# Internal helpers
# ===========================================================================

class _StageError(Exception):
    """Wraps a stage failure with structured metadata."""
    def __init__(self, stage: str, cause: Exception, meta: dict, hint: str = ""):
        self.stage = stage
        self.meta  = {
            "stage_failed": stage,
            "error":        str(cause),
            "traceback":    traceback.format_exc(),
            **({"hint": hint} if hint else {}),
        }
        super().__init__(f"[{stage}] {cause}")


def _raise_stage(
    stage: str,
    exc:   Exception,
    meta:  dict,
    hint:  str = "",
) -> None:
    """Log the stage failure then raise _StageError."""
    log.error(
        f"service_1_stage_{stage}_failed",
        stage     = stage,
        error     = str(exc),
        traceback = traceback.format_exc(),
        **({"hint": hint} if hint else {}),
    )
    raise _StageError(stage, exc, meta, hint)


def _safe_float(value: Any, default: float = 0.0) -> float:
    """Cast value to float; return default when value is None or invalid."""
    if value is None:
        return default
    try:
        return float(value)
    except (TypeError, ValueError):
        return default


def _safe_datetime(value: Any) -> datetime | None:
    """
    Ensure a datetime column value is a datetime object.

    aiomysql occasionally returns DATETIME columns as strings
    (e.g. '2026-06-03 13:55:00') instead of datetime instances.
    asyncpg refuses str for TIMESTAMP columns — parse it here.
    """
    if value is None:
        return None
    if isinstance(value, datetime):
        return value
    if isinstance(value, str):
        try:
            return datetime.strptime(value, "%Y-%m-%d %H:%M:%S")
        except ValueError:
            try:
                return datetime.fromisoformat(value)
            except ValueError:
                return None
    return None


def _safe_str(value: Any) -> str | None:
    """
    Cast value to str, return None when value is None.

    asyncpg is strict about VARCHAR columns — it rejects int/Decimal values
    even when the DB column type is text-compatible. DSE MySQL occasionally
    returns numeric types for columns like MKISTAT_INSTRUMENT_NUMBER.
    """
    if value is None:
        return None
    return str(value)
