from fastapi import APIRouter, Depends, Query, WebSocket, WebSocketDisconnect
from jose import JWTError
import logging

from app.core.security import verify_token
from app.realtime.events import WsEventType
from app.realtime.manager import manager
from app.schemas.ws_event import WsEvent

logger = logging.getLogger(__name__)
router = APIRouter(tags=["WebSocket"])


def _can_access_canal(canal: str, role: str) -> bool:
    if canal == "admin":
        return role in ("ADMIN_SYSTEME", "GESTIONNAIRE")
    if canal in ("signalements", "dashboard", "home", "postes"):
        return True  # TEMP: allow all roles to debug 403; tighten after confirmation
    if canal == "events":
        return True
    if canal.startswith("user:"):
        return True
    if canal.startswith("discussion:") or canal.startswith("signalement:") or canal.startswith("poste:"):
        return True
    if canal == "calendriers":
        return True
    if canal == "discussions":
        return True
    return False


@router.websocket("/ws/{canal}")
async def websocket_endpoint(
    websocket: WebSocket,
    canal: str,
    token: str = Query(...),
):
    # ── Authenticate ──────────────────────────────────────────────────────────
    try:
        payload = verify_token(token)
        user_id: str = payload["sub"]
        user_role: str = payload.get("role", "UTILISATEUR")
    except (JWTError, KeyError):
        await websocket.close(code=4001, reason="Token invalide")
        return

    # ── Authorize ─────────────────────────────────────────────────────────────
    logger.info("WS auth: canal=%s role=%s user_id=%s", canal, user_role, user_id)
    if not _can_access_canal(canal, user_role):
        logger.warning("WS rejected: canal=%s role=%s user_id=%s", canal, user_role, user_id)
        await websocket.close(code=4003, reason="Accès refusé")
        return

    await manager.connect(canal, websocket, user_id)

    # Send confirmation event
    try:
        confirmation = WsEvent(
            type=WsEventType.CONNECTED,
            canal=canal,
            payload={"userId": user_id, "canal": canal},
        ).model_dump()
        import orjson
        await websocket.send_text(orjson.dumps(confirmation, default=str).decode("utf-8"))
    except Exception:
        manager.disconnect(canal, websocket)
        return

    # ── Main loop ─────────────────────────────────────────────────────────────
    try:
        while True:
            raw = await websocket.receive_text()
            if raw.strip() == "ping":
                await websocket.send_text("pong")
    except WebSocketDisconnect:
        manager.disconnect(canal, websocket)
    except Exception:
        manager.disconnect(canal, websocket)
