"""
Read repository for the market-data pull endpoint.

Re-implements the query logic of the legacy `market-data.php` against
PostgreSQL using SQLAlchemy 2.0 async ORM. The IMDS column names are
preserved verbatim (UPPERCASE) and SQLAlchemy quotes them automatically,
so no raw SQL string interpolation is needed — this removes the SQL
injection surface the PHP version exposed via dynamic column names.

All public functions return JSON-serialisable dicts (Decimal → float,
datetime/date/time → ISO-8601 strings).
"""

from __future__ import annotations

from datetime import date as date_cls
from datetime import datetime, time
from decimal import Decimal
from typing import Any

from sqlalchemy import and_, func, select
from sqlalchemy.ext.asyncio import AsyncSession

from app.db.imds_models import (
    ImdsBase,
    ImdsIdxData,
    ImdsManData,
    ImdsMkistatData,
    ImdsTrdData,
)

# ---------------------------------------------------------------------------
# Table registry — maps table name → ORM model, primary key, timestamp column
# ---------------------------------------------------------------------------


class TableConfig:
    """Static descriptor for one readable IMDS table."""

    def __init__(self, model: type[ImdsBase], primary_key: str, timestamp: str):
        self.model = model
        self.primary_key = primary_key
        self.timestamp = timestamp

    @property
    def column_names(self) -> set[str]:
        return {c.name for c in self.model.__table__.columns}


TABLE_REGISTRY: dict[str, TableConfig] = {
    "imds_idx_data": TableConfig(ImdsIdxData, "IDX_ID", "IDX_STORE_TIMESTAMP"),
    "imds_man_data": TableConfig(ImdsManData, "MAN_ID", "MAN_STORE_TIMESTAMP"),
    "imds_mkistat_data": TableConfig(
        ImdsMkistatData, "MKISTAT_ID", "MKISTAT_STORE_TIMESTAMP"
    ),
    "imds_trd_data": TableConfig(ImdsTrdData, "TRD_ID", "TRD_STORE_TIMESTAMP"),
}


# ---------------------------------------------------------------------------
# Serialisation helpers
# ---------------------------------------------------------------------------


def _coerce_value(value: Any) -> Any:
    """Convert a single column value into a JSON-serialisable primitive."""
    if value is None:
        return None
    if isinstance(value, Decimal):
        return float(value)
    if isinstance(value, datetime):
        return value.strftime("%Y-%m-%d %H:%M:%S")
    if isinstance(value, date_cls):
        return value.isoformat()
    if isinstance(value, time):
        return value.strftime("%H:%M:%S")
    return value


def _row_to_dict(row: ImdsBase) -> dict[str, Any]:
    """Serialise an ORM row into a plain JSON-friendly dict."""
    return {
        col.name: _coerce_value(getattr(row, col.name))
        for col in row.__table__.columns
    }


def _resolve_column(cfg: TableConfig, name: str | None, default: str):
    """
    Return the ORM column attribute for `name`, falling back to `default`.

    Validates against the model's real columns so a caller cannot inject an
    arbitrary identifier (the PHP version trusted client-supplied names).
    """
    chosen = name or default
    if chosen not in cfg.column_names:
        raise ValueError(f"Unknown column '{chosen}' for table {cfg.model.__tablename__}")
    return getattr(cfg.model, chosen)


def _parse_date(value: str) -> date_cls:
    """Parse a YYYY-MM-DD string into a date, raising ValueError if malformed."""
    return datetime.strptime(value, "%Y-%m-%d").date()


# ---------------------------------------------------------------------------
# fetch_new_data
# ---------------------------------------------------------------------------


async def fetch_new_data(
    db: AsyncSession,
    cfg: TableConfig,
    *,
    last_id: int | None,
    last_timestamp: str | None,
    primary_key_column: str | None,
    batch_size: int,
) -> tuple[list[dict[str, Any]], bool]:
    """
    Incremental pull.

    Pagination precedence (matches the PHP contract):
      1. last_id + primary key  → ID-based (preferred, no row drift)
      2. last_timestamp         → timestamp-based fallback
      3. neither                → latest batch (newest rows first)
    """
    model = cfg.model
    pk_col = _resolve_column(cfg, primary_key_column, cfg.primary_key)
    ts_col = getattr(model, cfg.timestamp)

    if last_id is not None:
        stmt = (
            select(model)
            .where(pk_col > last_id)
            .order_by(pk_col.asc())
            .limit(batch_size)
        )
    elif last_timestamp:
        stmt = (
            select(model)
            .where(ts_col > last_timestamp)
            .order_by(ts_col.asc())
            .limit(batch_size)
        )
    else:
        # First call — return the newest batch.
        stmt = select(model).order_by(pk_col.desc()).limit(batch_size)

    rows = list((await db.execute(stmt)).scalars().all())
    data = [_row_to_dict(r) for r in rows]

    has_more = False
    if data:
        if last_id is not None:
            last_pk = data[-1][cfg.primary_key]
            has_more = await _count_where(db, model, pk_col > last_pk) > 0
        elif last_timestamp:
            last_ts = data[-1][cfg.timestamp]
            has_more = await _count_where(db, model, ts_col > last_ts) > 0

    return data, has_more


# ---------------------------------------------------------------------------
# fetch_all
# ---------------------------------------------------------------------------


async def fetch_all(
    db: AsyncSession,
    cfg: TableConfig,
    *,
    last_id: int | None,
    primary_key_column: str | None,
    batch_size: int,
    offset: int,
    until_date: str | None,
    timestamp_column: str | None,
) -> tuple[list[dict[str, Any]], bool]:
    """
    Full backfill pull.

    ID-based pagination is preferred. Offset-based is the fallback when no
    `last_id` cursor is supplied. An optional `until_date` (inclusive, end of
    day) caps the result against `timestamp_column`.
    """
    model = cfg.model
    pk_col = _resolve_column(cfg, primary_key_column, cfg.primary_key)

    until_dt: datetime | None = None
    filter_col = None
    if until_date and timestamp_column:
        filter_col = _resolve_column(cfg, timestamp_column, cfg.timestamp)
        until_dt = datetime.combine(_parse_date(until_date), time(23, 59, 59))

    if last_id is not None:
        conditions = [pk_col > last_id]
        if until_dt is not None and filter_col is not None:
            conditions.append(filter_col <= until_dt)
        stmt = (
            select(model)
            .where(and_(*conditions))
            .order_by(pk_col.asc())
            .limit(batch_size)
        )
        rows = list((await db.execute(stmt)).scalars().all())
        data = [_row_to_dict(r) for r in rows]

        has_more = False
        if data:
            last_pk = data[-1][cfg.primary_key]
            check = [pk_col > last_pk]
            if until_dt is not None and filter_col is not None:
                check.append(filter_col <= until_dt)
            has_more = await _count_where(db, model, and_(*check)) > 0
        return data, has_more

    # Offset-based fallback
    ts_col = getattr(model, cfg.timestamp)
    if until_dt is not None and filter_col is not None:
        stmt = (
            select(model)
            .where(filter_col <= until_dt)
            .order_by(ts_col.asc())
            .limit(batch_size)
            .offset(offset)
        )
    else:
        stmt = select(model).order_by(ts_col.asc()).limit(batch_size).offset(offset)

    rows = list((await db.execute(stmt)).scalars().all())
    data = [_row_to_dict(r) for r in rows]
    # Offset paging can only estimate: a full page implies more may remain.
    has_more = len(data) >= batch_size
    return data, has_more


# ---------------------------------------------------------------------------
# fetch_day_end_data  (imds_mkistat_data only)
# ---------------------------------------------------------------------------


async def fetch_day_end_data(
    db: AsyncSession,
    trade_date: str,
) -> list[dict[str, Any]]:
    """
    Latest non-zero-trade row per instrument for a given trade date.

    Reproduces the PHP self-join: for each MKISTAT_INSTRUMENT_CODE take the
    row with MAX(MKISTAT_ID) where MKISTAT_STORE_DATE = date and
    MKISTAT_TOTAL_TRADES != 0.
    """
    model = ImdsMkistatData
    date_obj = _parse_date(trade_date)

    latest = (
        select(
            model.MKISTAT_INSTRUMENT_CODE.label("code"),
            func.max(model.MKISTAT_ID).label("max_id"),
        )
        .where(
            model.MKISTAT_STORE_DATE == date_obj,
            model.MKISTAT_TOTAL_TRADES != 0,
        )
        .group_by(model.MKISTAT_INSTRUMENT_CODE)
        .subquery()
    )

    stmt = (
        select(model)
        .join(
            latest,
            and_(
                model.MKISTAT_INSTRUMENT_CODE == latest.c.code,
                model.MKISTAT_ID == latest.c.max_id,
            ),
        )
        .where(
            model.MKISTAT_STORE_DATE == date_obj,
            model.MKISTAT_TOTAL_TRADES != 0,
        )
    )

    rows = list((await db.execute(stmt)).scalars().all())
    return [_row_to_dict(r) for r in rows]


# ---------------------------------------------------------------------------
# Internal count helpers
# ---------------------------------------------------------------------------


async def _count_where(db: AsyncSession, model: type[ImdsBase], condition) -> int:
    """Return COUNT(*) for `model` filtered by `condition`."""
    stmt = select(func.count()).select_from(model).where(condition)
    return (await db.execute(stmt)).scalar_one()
