"""
IMDS repository — insert functions for the four IMDS data tables.

Each function:
  • Accepts a typed dataclass / dict of values.
  • Creates the ORM row, adds it to the session, and flushes (does NOT commit).
  • Commits are handled by the caller (typically FastAPI's get_db() dependency
    or the scheduler service wrapper) so multiple inserts can be batched in
    one transaction.

Bulk helpers are also provided for high-throughput scheduler services that
receive a list of records per polling tick.

Usage (inside a FastAPI route or scheduler service):

    from app.db.session import get_db          # FastAPI dep
    from app.db.imds_repository import (
        insert_idx, insert_man, insert_mkistat, insert_trd,
        bulk_insert_idx, bulk_insert_mkistat,
    )

    async with AsyncSession(engine) as db:
        row = await insert_idx(db, idx_index_id="DSEX", ...)
        await db.commit()
"""

from __future__ import annotations

from dataclasses import dataclass, field
from datetime import date, datetime, time
from typing import Optional, Sequence

from sqlalchemy.ext.asyncio import AsyncSession

import time as _time

from sqlalchemy import select

from app.db.imds_models import ImdsIdxData, ImdsManData, ImdsMkistatData, ImdsTrdData, MktSecurityCode


# ===========================================================================
# Data transfer objects  (plain dataclasses — no Pydantic dep needed here)
# ===========================================================================

@dataclass
class IdxRecord:
    """Input DTO for imds_idx_data."""
    idx_index_id:             str
    idx_date_time:            datetime
    idx_capital_value:        Optional[float] = None
    idx_deviation:            Optional[float] = None
    idx_percentage_deviation: Optional[float] = None
    idx_store_date:           Optional[date]  = None
    idx_store_time:           Optional[time]  = None
    # idx_store_timestamp is filled by the DB server — omit here


@dataclass
class ManRecord:
    """Input DTO for imds_man_data."""
    man_announcement_date_time: datetime
    man_announcement_prefix:    Optional[str]  = None
    man_announcement:           Optional[str]  = None
    man_expiry_date:            Optional[date] = None
    man_store_date:             Optional[date] = None
    man_store_time:             Optional[time] = None


@dataclass
class MkistatRecord:
    """Input DTO for imds_mkistat_data."""
    # Required fields (no DB default)
    mkistat_change_yday_close: float
    mkistat_avg_trade_share:   float
    mkistat_adv_shares:        float
    mkistat_adv_value:         float
    mkistat_change_volume:     float
    mkistat_pvr:               float
    mkistat_std_dev_60min:     float
    mkistat_vwap_daily:        float
    mkistat_vwap_1h:           float
    # Enriched / derived columns
    mkistat_sector:               Optional[str]   = None
    mkistat_yesterday_change_pct: Optional[float] = None
    # Optional fields
    mkistat_instrument_code:        Optional[str]      = None
    mkistat_instrument_number:      Optional[str]      = None
    mkistat_quote_bases:            Optional[str]      = None
    mkistat_open_price:             Optional[float]    = None
    mkistat_pub_last_traded_price:  Optional[float]    = None
    mkistat_spot_last_traded_price: Optional[float]    = None
    mkistat_high_price:             Optional[float]    = None
    mkistat_low_price:              Optional[float]    = None
    mkistat_close_price:            Optional[float]    = None
    mkistat_yday_close_price:       Optional[float]    = None
    mkistat_total_trades:           Optional[float]    = None
    mkistat_total_volume:           Optional[float]    = None
    mkistat_total_value:            Optional[float]    = None
    mkistat_public_total_trades:    Optional[float]    = None
    mkistat_public_total_volume:    Optional[float]    = None
    mkistat_public_total_value:     Optional[float]    = None
    mkistat_spot_total_trades:      Optional[float]    = None
    mkistat_spot_total_volume:      Optional[float]    = None
    mkistat_spot_total_value:       Optional[float]    = None
    mkistat_lm_date_time:           Optional[datetime] = None
    mkistat_store_date:             Optional[date]     = None
    mkistat_store_time:             Optional[time]     = None


@dataclass
class TrdRecord:
    """Input DTO for imds_trd_data."""
    trd_sno:          Optional[int]      = None
    trd_total_trades: Optional[float]    = None
    trd_total_volume: Optional[float]    = None
    trd_total_value:  Optional[float]    = None
    trd_mkt_status:   Optional[str]      = None
    trd_lm_date_time: Optional[datetime] = None
    trd_store_date:   Optional[date]     = None
    trd_store_time:   Optional[time]     = None


# ===========================================================================
# Single-row insert functions
# ===========================================================================

async def insert_idx(db: AsyncSession, rec: IdxRecord) -> ImdsIdxData:
    """
    Insert one row into imds_idx_data and return the ORM object.

    The caller must commit the session after this call (or after a batch).
    """
    row = ImdsIdxData(
        IDX_INDEX_ID=rec.idx_index_id,
        IDX_DATE_TIME=rec.idx_date_time,
        IDX_CAPITAL_VALUE=rec.idx_capital_value,
        IDX_DEVIATION=rec.idx_deviation,
        IDX_PERCENTAGE_DEVIATION=rec.idx_percentage_deviation,
        IDX_STORE_DATE=rec.idx_store_date,
        IDX_STORE_TIME=rec.idx_store_time,
    )
    db.add(row)
    await db.flush()   # populate IDX_ID without committing
    return row


async def insert_man(db: AsyncSession, rec: ManRecord) -> ImdsManData:
    """Insert one row into imds_man_data and return the ORM object."""
    row = ImdsManData(
        MAN_ANNOUNCEMENT_DATE_TIME=rec.man_announcement_date_time,
        MAN_ANNOUNCEMENT_PREFIX=rec.man_announcement_prefix,
        MAN_ANNOUNCEMENT=rec.man_announcement,
        MAN_EXPIRY_DATE=rec.man_expiry_date,
        MAN_STORE_DATE=rec.man_store_date,
        MAN_STORE_TIME=rec.man_store_time,
    )
    db.add(row)
    await db.flush()
    return row


async def insert_mkistat(
    db: AsyncSession, rec: MkistatRecord
) -> ImdsMkistatData:
    """Insert one row into imds_mkistat_data and return the ORM object."""
    row = ImdsMkistatData(
        MKISTAT_SECTOR=rec.mkistat_sector,
        MKISTAT_YESTERDAY_CHANGE_PCT=rec.mkistat_yesterday_change_pct,
        MKISTAT_INSTRUMENT_CODE=rec.mkistat_instrument_code,
        MKISTAT_INSTRUMENT_NUMBER=rec.mkistat_instrument_number,
        MKISTAT_QUOTE_BASES=rec.mkistat_quote_bases,
        MKISTAT_OPEN_PRICE=rec.mkistat_open_price,
        MKISTAT_PUB_LAST_TRADED_PRICE=rec.mkistat_pub_last_traded_price,
        MKISTAT_SPOT_LAST_TRADED_PRICE=rec.mkistat_spot_last_traded_price,
        MKISTAT_HIGH_PRICE=rec.mkistat_high_price,
        MKISTAT_LOW_PRICE=rec.mkistat_low_price,
        MKISTAT_CLOSE_PRICE=rec.mkistat_close_price,
        MKISTAT_YDAY_CLOSE_PRICE=rec.mkistat_yday_close_price,
        MKISTAT_TOTAL_TRADES=rec.mkistat_total_trades,
        MKISTAT_TOTAL_VOLUME=rec.mkistat_total_volume,
        MKISTAT_TOTAL_VALUE=rec.mkistat_total_value,
        MKISTAT_PUBLIC_TOTAL_TRADES=rec.mkistat_public_total_trades,
        MKISTAT_PUBLIC_TOTAL_VOLUME=rec.mkistat_public_total_volume,
        MKISTAT_PUBLIC_TOTAL_VALUE=rec.mkistat_public_total_value,
        MKISTAT_SPOT_TOTAL_TRADES=rec.mkistat_spot_total_trades,
        MKISTAT_SPOT_TOTAL_VOLUME=rec.mkistat_spot_total_volume,
        MKISTAT_SPOT_TOTAL_VALUE=rec.mkistat_spot_total_value,
        MKISTAT_LM_DATE_TIME=rec.mkistat_lm_date_time,
        MKISTAT_CHANGE_YDAY_CLOSE=rec.mkistat_change_yday_close,
        MKISTAT_AVG_TRADE_SHARE=rec.mkistat_avg_trade_share,
        MKISTAT_ADV_SHARES=rec.mkistat_adv_shares,
        MKISTAT_ADV_VALUE=rec.mkistat_adv_value,
        MKISTAT_CHANGE_VOLUME=rec.mkistat_change_volume,
        MKISTAT_PVR=rec.mkistat_pvr,
        MKISTAT_STD_DEV_60MIN=rec.mkistat_std_dev_60min,
        MKISTAT_VWAP_DAILY=rec.mkistat_vwap_daily,
        MKISTAT_VWAP_1H=rec.mkistat_vwap_1h,
        MKISTAT_STORE_DATE=rec.mkistat_store_date,
        MKISTAT_STORE_TIME=rec.mkistat_store_time,
    )
    db.add(row)
    await db.flush()
    return row


async def insert_trd(db: AsyncSession, rec: TrdRecord) -> ImdsTrdData:
    """Insert one row into imds_trd_data and return the ORM object."""
    row = ImdsTrdData(
        TRD_SNO=rec.trd_sno,
        TRD_TOTAL_TRADES=rec.trd_total_trades,
        TRD_TOTAL_VOLUME=rec.trd_total_volume,
        TRD_TOTAL_VALUE=rec.trd_total_value,
        TRD_MKT_STATUS=rec.trd_mkt_status,
        TRD_LM_DATE_TIME=rec.trd_lm_date_time,
        TRD_STORE_DATE=rec.trd_store_date,
        TRD_STORE_TIME=rec.trd_store_time,
    )
    db.add(row)
    await db.flush()
    return row


# ===========================================================================
# Bulk insert helpers — one DB round-trip for N rows
# ===========================================================================

async def bulk_insert_idx(
    db: AsyncSession, records: Sequence[IdxRecord]
) -> list[ImdsIdxData]:
    """Insert multiple imds_idx_data rows in a single flush."""
    rows = [
        ImdsIdxData(
            IDX_INDEX_ID=r.idx_index_id,
            IDX_DATE_TIME=r.idx_date_time,
            IDX_CAPITAL_VALUE=r.idx_capital_value,
            IDX_DEVIATION=r.idx_deviation,
            IDX_PERCENTAGE_DEVIATION=r.idx_percentage_deviation,
            IDX_STORE_DATE=r.idx_store_date,
            IDX_STORE_TIME=r.idx_store_time,
        )
        for r in records
    ]
    db.add_all(rows)
    await db.flush()
    return rows


async def bulk_insert_mkistat(
    db: AsyncSession, records: Sequence[MkistatRecord]
) -> list[ImdsMkistatData]:
    """Insert multiple imds_mkistat_data rows in a single flush."""
    rows = [
        ImdsMkistatData(
            MKISTAT_SECTOR=r.mkistat_sector,
            MKISTAT_YESTERDAY_CHANGE_PCT=r.mkistat_yesterday_change_pct,
            MKISTAT_INSTRUMENT_CODE=r.mkistat_instrument_code,
            MKISTAT_INSTRUMENT_NUMBER=r.mkistat_instrument_number,
            MKISTAT_QUOTE_BASES=r.mkistat_quote_bases,
            MKISTAT_OPEN_PRICE=r.mkistat_open_price,
            MKISTAT_PUB_LAST_TRADED_PRICE=r.mkistat_pub_last_traded_price,
            MKISTAT_SPOT_LAST_TRADED_PRICE=r.mkistat_spot_last_traded_price,
            MKISTAT_HIGH_PRICE=r.mkistat_high_price,
            MKISTAT_LOW_PRICE=r.mkistat_low_price,
            MKISTAT_CLOSE_PRICE=r.mkistat_close_price,
            MKISTAT_YDAY_CLOSE_PRICE=r.mkistat_yday_close_price,
            MKISTAT_TOTAL_TRADES=r.mkistat_total_trades,
            MKISTAT_TOTAL_VOLUME=r.mkistat_total_volume,
            MKISTAT_TOTAL_VALUE=r.mkistat_total_value,
            MKISTAT_PUBLIC_TOTAL_TRADES=r.mkistat_public_total_trades,
            MKISTAT_PUBLIC_TOTAL_VOLUME=r.mkistat_public_total_volume,
            MKISTAT_PUBLIC_TOTAL_VALUE=r.mkistat_public_total_value,
            MKISTAT_SPOT_TOTAL_TRADES=r.mkistat_spot_total_trades,
            MKISTAT_SPOT_TOTAL_VOLUME=r.mkistat_spot_total_volume,
            MKISTAT_SPOT_TOTAL_VALUE=r.mkistat_spot_total_value,
            MKISTAT_LM_DATE_TIME=r.mkistat_lm_date_time,
            MKISTAT_CHANGE_YDAY_CLOSE=r.mkistat_change_yday_close,
            MKISTAT_AVG_TRADE_SHARE=r.mkistat_avg_trade_share,
            MKISTAT_ADV_SHARES=r.mkistat_adv_shares,
            MKISTAT_ADV_VALUE=r.mkistat_adv_value,
            MKISTAT_CHANGE_VOLUME=r.mkistat_change_volume,
            MKISTAT_PVR=r.mkistat_pvr,
            MKISTAT_STD_DEV_60MIN=r.mkistat_std_dev_60min,
            MKISTAT_VWAP_DAILY=r.mkistat_vwap_daily,
            MKISTAT_VWAP_1H=r.mkistat_vwap_1h,
            MKISTAT_STORE_DATE=r.mkistat_store_date,
            MKISTAT_STORE_TIME=r.mkistat_store_time,
        )
        for r in records
    ]
    db.add_all(rows)
    await db.flush()
    return rows


async def bulk_insert_trd(
    db: AsyncSession, records: Sequence[TrdRecord]
) -> list[ImdsTrdData]:
    """Insert multiple imds_trd_data rows in a single flush."""
    rows = [
        ImdsTrdData(
            TRD_SNO=r.trd_sno,
            TRD_TOTAL_TRADES=r.trd_total_trades,
            TRD_TOTAL_VOLUME=r.trd_total_volume,
            TRD_TOTAL_VALUE=r.trd_total_value,
            TRD_MKT_STATUS=r.trd_mkt_status,
            TRD_LM_DATE_TIME=r.trd_lm_date_time,
            TRD_STORE_DATE=r.trd_store_date,
            TRD_STORE_TIME=r.trd_store_time,
        )
        for r in records
    ]
    db.add_all(rows)
    await db.flush()
    return rows


# ===========================================================================
# Security → sector lookup  (TTL-cached, 8-hour refresh)
# ===========================================================================

_sector_map_cache: dict[str, str] = {}
_sector_map_loaded_at: float = 0.0
_SECTOR_MAP_TTL: float = 8 * 3600  # 8 hours


async def get_security_sector_map(db: AsyncSession) -> dict[str, str]:
    """
    Return a dict mapping security_code → sector from mkt_security_code.

    Cached in memory for 8 hours (TTL) — matching the original
    LocalStore behaviour. Cache is module-level so it persists across
    scheduler ticks within the same process.

    Returns an empty dict (not an exception) when the table is empty.
    """
    global _sector_map_cache, _sector_map_loaded_at

    now = _time.monotonic()
    if _sector_map_cache and (now - _sector_map_loaded_at) < _SECTOR_MAP_TTL:
        return _sector_map_cache

    result = await db.execute(select(MktSecurityCode))
    rows = result.scalars().all()
    _sector_map_cache = {r.security_code: (r.sector or "Unknown") for r in rows}
    _sector_map_loaded_at = now
    return _sector_map_cache
