"""
Data delivery routes — pull (REST) and push (WebSocket / SSE).

GET  /data/pull/{dataset}      Pull on request (Caller Types 1 + 2)
WS   /data/push/subscribe      WebSocket push feed (Caller Types 1 + 2)

RBAC matrix:
  Caller Type 1 (user JWT)   — viewer/analyst/admin → pull; analyst/admin → push WS
  Caller Type 2 (HMAC key)   — scope 'data:pull'   → pull; scope 'feed:subscribe' → WS
  Caller Type 3 (cron)       — rejected from all data endpoints

WebSocket auth:
  Headers X-Service-Key / X-Timestamp / X-Signature (or Bearer JWT)
  sent in the WebSocket upgrade request.
  HMAC validated once on connect — not per-message (stateful connection).
  Key revocation checked on each heartbeat ping (every WS_HEARTBEAT_INTERVAL s).
"""

from __future__ import annotations

import asyncio
import uuid
from datetime import datetime, timezone

import structlog
from fastapi import (
    APIRouter,
    Depends,
    HTTPException,
    Request,
    WebSocket,
    WebSocketDisconnect,
)
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession

from app.config import get_settings
from app.core.rate_limit import data_limiter
from app.db.market_data_repository import (
    TABLE_REGISTRY,
    fetch_all,
    fetch_day_end_data,
    fetch_new_data,
)
from app.db.models import ApiKey
from app.db.session import get_db
from app.dependencies.auth import (
    CallerContext,
    get_client_ip,
    require_user_or_scope,
    verify_api_key,
    verify_user_jwt,
)
from app.schemas.market_data import (
    DayEndDataResponse,
    MarketDataRequest,
    MarketDataResponse,
)

settings = get_settings()
log = structlog.get_logger(__name__)
router = APIRouter(prefix="/data", tags=["data"])

# WebSocket close codes (4000–4999 are application-defined)
_WS_UNAUTHORIZED = 4001
_WS_FORBIDDEN = 4003
_WS_KEY_REVOKED = 4004


# ===========================================================================
# GET /data/pull/{dataset}
# ===========================================================================

@router.get("/pull/{dataset}", status_code=200)
@data_limiter.limit("1000/minute")   # per user_id or key_id via _get_user_id_or_ip
async def pull_dataset(
    dataset: str,
    request: Request,
    caller: CallerContext = Depends(
        require_user_or_scope(
            roles=("viewer", "analyst", "admin"),
            scope="data:pull",
        )
    ),
    db: AsyncSession = Depends(get_db),
) -> dict:
    """
    Pull dataset on request.

    Accepted callers:
      User JWT   — any role (viewer / analyst / admin)
      HMAC key   — must have scope 'data:pull'
      Cron key   — rejected (cron accesses data internally, not via this endpoint)

    Rate limit: 1000 req/min per user_id / key_id (Redis counter).
    Per-key ceiling for API key callers is stored in api_keys.rate_limit;
    enforce that separately if the per-key ceiling < 1000.
    """
    # Enforce per-key rate ceiling for API key callers
    if caller.caller_type == "api_key" and caller.rate_limit < 1000:
        # slowapi's @data_limiter.limit decorator uses 1000/min as default;
        # for keys with lower ceilings, the Redis counter approach below
        # is a placeholder. In production, use a dynamic limit string:
        # @data_limiter.limit(lambda req: f"{caller.rate_limit}/minute")
        pass

    # --- Placeholder: replace with real data fetch logic ---
    data = await _fetch_dataset(dataset, caller)
    # -------------------------------------------------------

    return {
        "dataset": dataset,
        "records": data,
        "fetched_at": datetime.now(timezone.utc).isoformat(),
        "caller": caller.caller_type,
    }


# ===========================================================================
# POST /data/market-data
# ===========================================================================

@router.post("/market-data", status_code=200)
@data_limiter.limit("1000/minute")
async def market_data(
    payload: MarketDataRequest,
    request: Request,
    caller: CallerContext = Depends(
        require_user_or_scope(
            roles=("viewer", "analyst", "admin"),
            scope="data:pull",
        )
    ),
    db: AsyncSession = Depends(get_db),
) -> MarketDataResponse | DayEndDataResponse:
    """
    PostgreSQL port of the legacy `market-data.php` pull endpoint.

    Authentication (S2S + user JWT — same gateway as every data route):
      • User JWT   — role viewer / analyst / admin
      • HMAC key   — scope 'data:pull' (X-Service-Key / X-Timestamp / X-Signature)
      • Cron key   — rejected

    Actions:
      fetch_new_data      Incremental pull (ID-based preferred, timestamp fallback).
      fetch_all           Full backfill (ID-based preferred, offset fallback,
                          optional `until_date` ceiling).
      fetch_day_end_data  Latest non-zero-trade row per instrument for a date
                          (imds_mkistat_data only).

    Error bodies preserve the PHP `{error, code}` shape for client compatibility.
    """
    cfg = TABLE_REGISTRY[payload.table]

    try:
        if payload.action == "fetch_day_end_data":
            return await _handle_day_end(db, payload)

        if payload.action == "fetch_new_data":
            data, has_more = await fetch_new_data(
                db,
                cfg,
                last_id=payload.last_id,
                last_timestamp=payload.last_timestamp,
                primary_key_column=payload.primary_key_column,
                batch_size=payload.batch_size,
            )
        else:  # fetch_all
            data, has_more = await fetch_all(
                db,
                cfg,
                last_id=payload.last_id,
                primary_key_column=payload.primary_key_column,
                batch_size=payload.batch_size,
                offset=payload.offset,
                until_date=payload.until_date,
                timestamp_column=payload.timestamp_column,
            )

    except ValueError as exc:
        # Bad column name / malformed date — client error, not a 500.
        raise HTTPException(
            status_code=400,
            detail={"error": str(exc), "code": "INVALID_PARAMETER"},
        ) from exc
    except Exception:
        log.exception("market_data_query_failed", table=payload.table, action=payload.action)
        raise HTTPException(
            status_code=500,
            detail={"error": "Query failed", "code": "QUERY_FAILED"},
        )

    return MarketDataResponse(data=data, count=len(data), has_more=has_more)


async def _handle_day_end(
    db: AsyncSession,
    payload: MarketDataRequest,
) -> DayEndDataResponse:
    """Validate and run the day-end snapshot query (mkistat only)."""
    if payload.table != "imds_mkistat_data":
        raise HTTPException(
            status_code=400,
            detail={
                "error": "Day-end data only supported for imds_mkistat_data table",
                "code": "INVALID_TABLE",
            },
        )
    if not payload.date:
        raise HTTPException(
            status_code=400,
            detail={"error": "Date parameter required", "code": "MISSING_DATE"},
        )

    try:
        data = await fetch_day_end_data(db, payload.date)
    except ValueError:
        raise HTTPException(
            status_code=400,
            detail={
                "error": "Invalid date format. Use YYYY-MM-DD",
                "code": "INVALID_DATE_FORMAT",
            },
        )
    except Exception:
        log.exception("market_data_day_end_failed", date=payload.date)
        raise HTTPException(
            status_code=500,
            detail={"error": "Query failed", "code": "QUERY_FAILED"},
        )

    return DayEndDataResponse(data=data, count=len(data), date=payload.date)


# ===========================================================================
# WebSocket /data/push/subscribe
# ===========================================================================

@router.websocket("/push/subscribe")
async def ws_push_subscribe(
    websocket: WebSocket,
    db: AsyncSession = Depends(get_db),
) -> None:
    """
    WebSocket push-feed endpoint.

    Auth flow:
      1. Extract auth from upgrade-request headers (same headers as HTTP)
      2. Verify on connect — not per-message (WS is stateful)
      3. Heartbeat ping every WS_HEARTBEAT_INTERVAL seconds
      4. Re-check key/user status on each heartbeat (revocation detection)
      5. Close with code 4004 if key is revoked mid-session
    """
    caller: CallerContext | None = None
    ip = get_client_ip(websocket)

    # --- Authenticate from upgrade headers ---
    auth = websocket.headers.get("Authorization", "")
    key_id_header = websocket.headers.get("X-Service-Key", "")

    try:
        if key_id_header:
            # Caller Type 2: HMAC-signed WebSocket connect
            # For WS, the "body" is empty bytes — sign timestamp + "" per protocol
            caller = await verify_api_key(websocket, db)  # type: ignore[arg-type]
            if "feed:subscribe" not in caller.scopes:
                await websocket.close(code=_WS_FORBIDDEN)
                return

        elif auth.startswith("Bearer "):
            # Caller Type 1: JWT user
            caller = await verify_user_jwt(auth[len("Bearer "):])
            if caller.role not in ("analyst", "admin"):
                # viewer role cannot subscribe to the push feed
                await websocket.close(code=_WS_FORBIDDEN)
                return

        else:
            await websocket.close(code=_WS_UNAUTHORIZED)
            return

    except HTTPException:
        await websocket.close(code=_WS_UNAUTHORIZED)
        return

    await websocket.accept()
    log.info(
        "ws_connected",
        caller_type=caller.caller_type,
        caller_id=getattr(caller, "sub", getattr(caller, "key_id", "?")),
        ip=ip,
    )

    # --- Message loop with heartbeat-based revocation check ---
    try:
        while True:
            await asyncio.sleep(settings.ws_heartbeat_interval)

            # Re-check key revocation on every heartbeat (30-s window)
            if caller.caller_type == "api_key":
                still_active = await _check_key_active(db, caller.key_id)
                if not still_active:
                    log.warning(
                        "ws_key_revoked",
                        key_id=caller.key_id,
                        ip=ip,
                    )
                    await websocket.close(code=_WS_KEY_REVOKED)
                    return

            # Heartbeat ping to keep the connection alive and detect dead clients
            await websocket.send_json(
                {
                    "type": "ping",
                    "timestamp": datetime.now(timezone.utc).isoformat(),
                }
            )

            # --- Placeholder: push real data events here ---
            # events = await _get_pending_events(caller)
            # for event in events:
            #     await websocket.send_json(event)
            # ------------------------------------------------

    except WebSocketDisconnect:
        log.info(
            "ws_disconnected",
            caller_type=caller.caller_type,
            ip=ip,
        )
    except Exception:
        log.exception("ws_error", ip=ip)
        await websocket.close(code=1011)  # 1011 = Internal Error


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

async def _check_key_active(db: AsyncSession, key_id_str: str) -> bool:
    """Return True if the API key is still active. Used for heartbeat checks."""
    try:
        key_id = uuid.UUID(key_id_str)
    except ValueError:
        return False
    result = await db.execute(
        select(ApiKey.is_active).where(ApiKey.key_id == key_id)
    )
    is_active = result.scalar_one_or_none()
    return bool(is_active)


async def _fetch_dataset(dataset: str, caller: CallerContext) -> list:
    """
    Placeholder data fetch. Replace with real data-source queries.
    Validate `dataset` name against an allowlist before querying.
    """
    _ALLOWED_DATASETS = {"prices", "volumes", "signals", "alerts"}
    if dataset not in _ALLOWED_DATASETS:
        raise HTTPException(
            status_code=404,
            detail=f"Dataset '{dataset}' not found. Available: {sorted(_ALLOWED_DATASETS)}",
        )
    return []
