from fastapi import APIRouter, WebSocket, WebSocketDisconnect, Query
from typing import List, Dict
import json
import logging
import anyio
from app.services.integrations.firestore import save_chat_message
from app.services.infra.database import insert_feedback

logger = logging.getLogger(__name__)

router = APIRouter()

@router.get("/ws-test/{tenant_id}/{thread_id}")
async def ws_diagnostic(tenant_id: str, thread_id: str):
    return {
        "status": "ok", 
        "message": f"WebSocket router is mounted. Testing for {tenant_id}/{thread_id}",
        "path_verified": True
    }

import asyncio
from app.services.infra.redis import get_redis_client, publish_event

class ConnectionManager:
    def __init__(self):
        # Dictionary mapping channel_id to a list of active WebSocket connections
        self.active_connections: Dict[str, List[WebSocket]] = {}
        # Keep track of pubsub tasks per websocket
        self.pubsub_tasks: Dict[WebSocket, asyncio.Task] = {}

    async def connect(self, websocket: WebSocket, channel_id: str):
        await websocket.accept()
        if channel_id not in self.active_connections:
            self.active_connections[channel_id] = []
        self.active_connections[channel_id].append(websocket)
        logger.info(f"New connection to channel {channel_id}. Total connections: {len(self.active_connections[channel_id])}")
        
        # Start a background task to listen to Redis for this websocket
        task = asyncio.create_task(self._listen_to_redis(websocket, channel_id))
        self.pubsub_tasks[websocket] = task

    async def _listen_to_redis(self, websocket: WebSocket, channel_id: str):
        retry_delay = 1
        max_delay = 60
        
        while True:
            try:
                redis_client = get_redis_client()
                async with redis_client.pubsub() as pubsub:
                    await pubsub.subscribe(channel_id)
                    logger.info(f"Subscribed to Redis channel: {channel_id}")
                    retry_delay = 1 # Reset delay on successful connection
                    
                    async for message in pubsub.listen():
                        if message["type"] == "message":
                            data = message["data"]
                            try:
                                await websocket.send_text(data)
                                logger.info(f"Successfully sent message to client on channel {channel_id}")
                                logger.debug(f"Forwarded Redis message to {channel_id}")
                            except Exception as e:
                                logger.error(f"Error sending message to websocket in {channel_id}: {e}")
                                return # Exit if websocket is dead
            except asyncio.CancelledError:
                logger.info(f"Redis listen task cancelled for channel {channel_id}")
                return
            except Exception as e:
                logger.error(f"Redis subscription error for {channel_id}: {e}. Retrying in {retry_delay}s...")
                await asyncio.sleep(retry_delay)
                retry_delay = min(retry_delay * 2, max_delay)

    def disconnect(self, websocket: WebSocket, channel_id: str):
        if channel_id in self.active_connections:
            if websocket in self.active_connections[channel_id]:
                self.active_connections[channel_id].remove(websocket)
            if not self.active_connections[channel_id]:
                del self.active_connections[channel_id]
                
        # Cancel the redis listen task
        if websocket in self.pubsub_tasks:
            self.pubsub_tasks[websocket].cancel()
            del self.pubsub_tasks[websocket]
            
        logger.info(f"Disconnected from channel {channel_id}")

    async def broadcast(self, message: str, channel_id: str, sender_socket: WebSocket = None):
        """
        Broadcasts a message to all participants in a channel via Redis.
        """
        try:
            payload = json.loads(message)
        except json.JSONDecodeError:
            payload = {"text": message}
            
        await publish_event(channel_id, payload)

    async def broadcast_system_event(self, channel_id: str, event_type: str, message: str, data: dict = None):
        """
        Broadcasts a system-level event to all participants in a channel.
        """
        payload = {
            "type": "system",
            "event": event_type,
            "message": message,
            "data": data or {}
        }
        await publish_event(channel_id, payload)

manager = ConnectionManager()

@router.websocket("/ws/{tenant_id}/dashboard")
async def dashboard_websocket_endpoint(
    websocket: WebSocket,
    tenant_id: str,
    role: str = Query("admin")
):
    """
    Dedicated WebSocket for the Admin Dashboard to receive real-time organization-wide updates.
    """
    if role != "admin":
        await websocket.close(code=1008) # Policy Violation
        return

    # Diagnostic: Check if schema exists
    from app.services.infra.database import get_db_connection
    try:
        conn = get_db_connection()
        with conn.cursor() as cur:
            cur.execute("SELECT schema_name FROM information_schema.schemata WHERE schema_name = %s", (tenant_id,))
            if not cur.fetchone():
                logger.warning(f"Dashboard WebSocket connected to tenant '{tenant_id}' which has NO database schema.")
        conn.close()
    except Exception as e:
        logger.warning(f"Failed to verify schema for tenant {tenant_id}: {e}")

    # Use namespacing for multi-tenant isolation
    channel_id = f"tenant_{tenant_id}:dashboard"
    await manager.connect(websocket, channel_id)
    logger.info(f"Admin Dashboard connected to real-time stream for tenant {tenant_id}")
    
    try:
        while True:
            # Dashboard is primarily a listener, but we keep the connection alive
            await websocket.receive_text() 
    except WebSocketDisconnect:
        manager.disconnect(websocket, channel_id)
        logger.info(f"Admin Dashboard disconnected for {tenant_id}")
    except Exception as e:
        logger.error(f"Unexpected error in Dashboard WebSocket for {tenant_id}: {e}", exc_info=True)
        manager.disconnect(websocket, channel_id)

@router.websocket("/ws/{tenant_id}/conversation_list")
async def conversation_list_websocket_endpoint(websocket: WebSocket, tenant_id: str):
    # Diagnostic: Check if schema exists
    from app.services.infra.database import get_db_connection
    try:
        conn = get_db_connection()
        with conn.cursor() as cur:
            cur.execute("SELECT schema_name FROM information_schema.schemata WHERE schema_name = %s", (tenant_id,))
            if not cur.fetchone():
                logger.warning(f"Conversation list WebSocket connected to tenant '{tenant_id}' which has NO database schema.")
        conn.close()
    except Exception as e:
        logger.warning(f"Failed to verify schema for tenant {tenant_id}: {e}")

    channel_id = f"tenant_{tenant_id}:conversation_list"
    await manager.connect(websocket, channel_id)
    logger.info(f"Connected to conversation_list stream for tenant {tenant_id}")
    try:
        while True:
            await websocket.receive_text()
    except WebSocketDisconnect:
        manager.disconnect(websocket, channel_id)
    except Exception as e:
        logger.error(f"Error in conversation_list WebSocket: {e}")
        manager.disconnect(websocket, channel_id)

@router.websocket("/ws/{tenant_id}/tickets")
async def tickets_websocket_endpoint(websocket: WebSocket, tenant_id: str):
    # Diagnostic: Check if schema exists to help debug prefix issues
    from app.services.infra.database import get_db_connection
    try:
        conn = get_db_connection()
        with conn.cursor() as cur:
            cur.execute("SELECT schema_name FROM information_schema.schemata WHERE schema_name = %s", (tenant_id,))
            if not cur.fetchone():
                logger.warning(f"WebSocket connected to tenant '{tenant_id}' which has NO corresponding database schema. Updates may not be received.")
        conn.close()
    except Exception as e:
        logger.warning(f"Failed to verify schema for tenant {tenant_id}: {e}")

    channel_id = f"tenant_{tenant_id}:tickets"
    await manager.connect(websocket, channel_id)
    logger.info(f"Connected to tickets stream for tenant {tenant_id}")
    try:
        while True:
            await websocket.receive_text()
    except WebSocketDisconnect:
        manager.disconnect(websocket, channel_id)
    except Exception as e:
        logger.error(f"Error in tickets WebSocket: {e}")
        manager.disconnect(websocket, channel_id)

@router.websocket("/ws/{tenant_id}/{thread_id}")
async def websocket_endpoint(
    websocket: WebSocket, 
    tenant_id: str, 
    thread_id: str,
    role: str = Query("user"),
    user_name: str = Query("Anonymous"),
    user_id: str = Query(None)
):
    # Use strict namespacing
    channel_id = f"tenant_{tenant_id}:chat_{thread_id}"
    await manager.connect(websocket, channel_id)
    logger.info(f"WebSocket session started for {role} ({user_name}) on channel {channel_id}, user_id: {user_id}")
    
    try:
        while True:
            # Receive message from the client
            data = await websocket.receive_text()
            try:
                payload = json.loads(data)
                message_text = payload.get("text", "")
            except json.JSONDecodeError:
                message_text = data
            
            if not message_text.strip():
                continue

            logger.info(f"[{role}] {user_name}: {message_text}")

            # 1. Persist to Firestore (Ongoing conversation history)
            # We run this in a thread to avoid blocking the async loop
            await anyio.to_thread.run_sync(
                save_chat_message,
                tenant_id,
                role,
                message_text,
                thread_id,
                user_name,
                f"Human Takeover Interaction ({role})",
                user_id
            )

            # 2. Persist to Postgres (Analytics/Feedback)
            try:
                await anyio.to_thread.run_sync(
                    insert_feedback,
                    tenant_id,
                    {
                        "thread_id": thread_id,
                        "user_name": user_name,
                        "question": message_text if role == "user" else f"[Agent Response]: {message_text}",
                        "answer": "Human-to-Human Handled" if role == "user" else "Response Delivered",
                        "metadata": {
                            "interaction_type": "human_takeover",
                            "role": role,
                            "real_time": True,
                            "user_id": user_id
                        }
                    }
                )
            except Exception as pg_err:
                logger.error(f"PG Persistence failed in WS: {pg_err}")

            # 3. Broadcast via Redis
            broadcast_payload = json.dumps({
                "role": role,
                "user_name": user_name,
                "text": message_text,
                "thread_id": thread_id,
                "user_id": user_id,
                "type": "chat_message"
            })
            await manager.broadcast(broadcast_payload, channel_id, sender_socket=websocket)

    except WebSocketDisconnect:
        manager.disconnect(websocket, channel_id)
        logger.info(f"WebSocket disconnected for {user_name} in {channel_id}")
    except Exception as e:
        logger.error(f"Unexpected error in websocket for {channel_id}: {e}", exc_info=True)
        manager.disconnect(websocket, channel_id)
