"""
Tool update listener for dynamic MCP server.

This module provides WebSocket, SSE, and polling support for receiving
real-time tool updates from the Coherence gateway.
"""

import asyncio
import json
import logging
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import Callable, Dict, Any, Optional
from urllib.parse import urljoin

import aiohttp
from aiohttp import ClientSession, ClientWebSocketResponse

logger = logging.getLogger(__name__)


@dataclass
class UpdateListenerConfig:
    """Configuration for the update listener."""
    
    gateway_url: str
    bundle_name: Optional[str] = None
    update_method: str = "websocket"  # websocket, sse, or polling
    polling_interval: int = 30  # seconds
    reconnect_interval: int = 5  # seconds
    max_reconnect_attempts: int = -1  # -1 for infinite
    on_tools_changed: Optional[Callable[[Dict[str, Any]], None]] = None
    on_connection_error: Optional[Callable[[Exception], None]] = None


class UpdateListener:
    """
    Listens for tool updates from the Coherence gateway.
    
    Supports multiple update methods:
    - WebSocket: Real-time bidirectional updates
    - SSE: Server-sent events for one-way updates
    - Polling: Periodic checking for updates
    """
    
    def __init__(self, config: UpdateListenerConfig):
        """
        Initialize the update listener.
        
        Args:
            config: Listener configuration
        """
        self.config = config
        self.session: Optional[ClientSession] = None
        self.websocket: Optional[ClientWebSocketResponse] = None
        self.running = False
        self.reconnect_attempts = 0
        
        # Track last known tool state for polling
        self.last_tool_state: Dict[str, Any] = {}
        self.last_poll_time: Optional[datetime] = None
    
    async def start(self):
        """Start the update listener."""
        self.running = True
        
        try:
            # Create session in the same event loop
            self.session = aiohttp.ClientSession()
            
            if self.config.update_method == "websocket":
                await self._start_websocket()
            elif self.config.update_method == "sse":
                await self._start_sse()
            elif self.config.update_method == "polling":
                await self._start_polling()
            else:
                logger.error(f"Unknown update method: {self.config.update_method}")
                
        except asyncio.CancelledError:
            # Handle cancellation gracefully
            logger.info("Update listener cancelled")
            raise
        except Exception as e:
            logger.error(f"Failed to start update listener: {e}")
            if self.config.on_connection_error:
                self.config.on_connection_error(e)
        finally:
            await self.stop()
    
    async def stop(self):
        """Stop the update listener."""
        self.running = False
        
        if self.websocket:
            try:
                await self.websocket.close()
            except Exception:
                pass
            self.websocket = None
        
        if self.session and not self.session.closed:
            try:
                # Give a small grace period for pending operations
                await asyncio.sleep(0.25)
                await self.session.close()
            except Exception:
                pass
            self.session = None
    
    async def _start_websocket(self):
        """Start WebSocket connection for real-time updates."""
        while self.running:
            try:
                # Build WebSocket URL
                ws_url = self.config.gateway_url.replace("http://", "ws://").replace("https://", "wss://")
                if self.config.bundle_name:
                    ws_url = urljoin(ws_url, f"/ws/tool-updates?bundle_id={self.config.bundle_name}")
                else:
                    ws_url = urljoin(ws_url, "/ws/tool-updates")
                
                logger.info(f"Connecting to WebSocket: {ws_url}")
                
                # Ensure session exists
                if not self.session:
                    logger.error("Session is None, cannot connect to WebSocket")
                    break
                
                # Connect with client metadata
                self.websocket = await self.session.ws_connect(
                    ws_url,
                    headers={
                        "User-Agent": "MCP-DynamicServer/1.0"
                    }
                )
                
                try:
                    self.reconnect_attempts = 0
                    logger.info("WebSocket connected successfully")
                    
                    # Handle messages
                    async for msg in self.websocket:
                        if msg.type == aiohttp.WSMsgType.TEXT:
                            await self._handle_websocket_message(msg.data)
                        elif msg.type == aiohttp.WSMsgType.ERROR:
                            logger.error(f"WebSocket error: {self.websocket.exception()}")
                            break
                        elif msg.type == aiohttp.WSMsgType.CLOSED:
                            logger.info("WebSocket connection closed")
                            break
                            
                finally:
                    # Ensure websocket is closed
                    if self.websocket and not self.websocket.closed:
                        await self.websocket.close()
                    self.websocket = None
                    
            except Exception as e:
                logger.error(f"WebSocket connection error: {e}")
                if self.config.on_connection_error:
                    self.config.on_connection_error(e)
                
                # Reconnect logic
                if self._should_reconnect():
                    await self._wait_before_reconnect()
                else:
                    break
    
    async def _handle_websocket_message(self, data: str):
        """
        Handle a WebSocket message.
        
        Args:
            data: Raw message data
        """
        try:
            message = json.loads(data)
            msg_type = message.get("type")
            msg_data = message.get("data", {})
            
            if msg_type == "initial_state":
                # Initial tool state
                tools = msg_data.get("tools", [])
                logger.info(f"Received initial state with {len(tools)} tools")
                # Store for comparison but don't trigger update
                self._update_tool_state(tools)
                
            elif msg_type == "tool_update":
                # Tool update event
                logger.info(f"Received tool update: {msg_data.get('event_type')}")
                if self.config.on_tools_changed:
                    self.config.on_tools_changed(msg_data)
                    
            elif msg_type == "ping":
                # Respond to ping
                if self.websocket:
                    await self.websocket.send_json({
                        "type": "pong",
                        "timestamp": datetime.utcnow().isoformat()
                    })
                    
            elif msg_type == "error":
                logger.error(f"Received error from server: {msg_data}")
                
        except Exception as e:
            logger.error(f"Error handling WebSocket message: {e}")
    
    async def _start_sse(self):
        """Start SSE connection for server-sent events."""
        while self.running:
            try:
                # Build SSE URL
                sse_url = urljoin(self.config.gateway_url, "/sse/tool-updates")
                if self.config.bundle_name:
                    sse_url += f"?bundle_id={self.config.bundle_name}"
                
                logger.info(f"Connecting to SSE: {sse_url}")
                
                # Connect to SSE endpoint using standard aiohttp
                headers = {
                    "User-Agent": "MCP-DynamicServer/1.0",
                    "Accept": "text/event-stream",
                    "Cache-Control": "no-cache"
                }
                
                async with self.session.get(sse_url, headers=headers) as response:
                    if response.status != 200:
                        raise Exception(f"SSE connection failed with status {response.status}")
                    
                    self.reconnect_attempts = 0
                    logger.info("SSE connected successfully")
                    
                    # Read SSE stream
                    event_type = None
                    event_data = ""
                    
                    async for line in response.content:
                        line = line.decode('utf-8').strip()
                        
                        if line.startswith('event:'):
                            event_type = line[6:].strip()
                        elif line.startswith('data:'):
                            event_data = line[5:].strip()
                        elif line == "" and event_type:
                            # Empty line marks end of event
                            if event_type == "tool_update":
                                await self._handle_sse_event(event_data)
                            elif event_type == "heartbeat":
                                # Heartbeat to keep connection alive
                                logger.debug("Received SSE heartbeat")
                            elif event_type == "initial_state":
                                # Handle initial state
                                await self._handle_sse_event(event_data, is_initial=True)
                            elif event_type == "error":
                                logger.error(f"SSE error event: {event_data}")
                            
                            # Reset for next event
                            event_type = None
                            event_data = ""
                            
            except Exception as e:
                logger.error(f"SSE connection error: {e}")
                if self.config.on_connection_error:
                    self.config.on_connection_error(e)
                
                # Reconnect logic
                if self._should_reconnect():
                    await self._wait_before_reconnect()
                else:
                    break
    
    async def _handle_sse_event(self, data: str, is_initial: bool = False):
        """
        Handle an SSE event.
        
        Args:
            data: Event data
            is_initial: Whether this is an initial state event
        """
        try:
            event_data = json.loads(data)
            
            if is_initial:
                # Initial state - just update our cache
                tools = event_data.get("tools", [])
                logger.info(f"Received SSE initial state with {len(tools)} tools")
                self._update_tool_state(tools)
            else:
                # Tool update event
                logger.info(f"Received SSE event: {event_data.get('event_type')}")
                
                if self.config.on_tools_changed:
                    self.config.on_tools_changed(event_data)
                
        except Exception as e:
            logger.error(f"Error handling SSE event: {e}")
    
    async def _start_polling(self):
        """Start polling for tool updates."""
        logger.info(f"Starting polling with interval: {self.config.polling_interval}s")
        
        while self.running:
            try:
                await self._poll_for_updates()
                
                # Wait for next poll
                await asyncio.sleep(self.config.polling_interval)
                
            except Exception as e:
                logger.error(f"Polling error: {e}")
                if self.config.on_connection_error:
                    self.config.on_connection_error(e)
                
                # Wait before retry
                await asyncio.sleep(self.config.reconnect_interval)
    
    async def _poll_for_updates(self):
        """Poll for tool updates."""
        try:
            # Build polling URL
            if self.config.bundle_name:
                url = urljoin(
                    self.config.gateway_url, 
                    f"/mcp/bundles/{self.config.bundle_name}/tools"
                )
            else:
                url = urljoin(self.config.gateway_url, "/mcp/tools/discovery")
            
            # Add timestamp for incremental updates
            params = {}
            if self.last_poll_time:
                params["since"] = self.last_poll_time.isoformat()
            
            # Make request
            async with self.session.get(url, params=params) as response:
                if response.status == 200:
                    data = await response.json()
                    tools = data.get("tools", [])
                    
                    # Check for changes
                    if self._has_tools_changed(tools):
                        logger.info("Detected tool changes via polling")
                        
                        # Create synthetic event
                        event = {
                            "event_type": "updated",
                            "timestamp": datetime.utcnow().isoformat(),
                            "tool_count": len(tools),
                            "source": "polling"
                        }
                        
                        if self.config.on_tools_changed:
                            self.config.on_tools_changed(event)
                    
                    # Update state
                    self._update_tool_state(tools)
                    self.last_poll_time = datetime.utcnow()
                    
                else:
                    logger.error(f"Polling failed with status: {response.status}")
                    
        except Exception as e:
            logger.error(f"Error during polling: {e}")
            raise
    
    def _update_tool_state(self, tools: list):
        """
        Update the cached tool state.
        
        Args:
            tools: List of tools
        """
        self.last_tool_state = {
            tool.get("name", tool.get("tool_name")): tool
            for tool in tools
        }
    
    def _has_tools_changed(self, tools: list) -> bool:
        """
        Check if tools have changed since last poll.
        
        Args:
            tools: Current list of tools
            
        Returns:
            True if tools have changed
        """
        if not self.last_tool_state:
            return bool(tools)
        
        current_state = {
            tool.get("name", tool.get("tool_name")): tool
            for tool in tools
        }
        
        # Check for differences
        if set(current_state.keys()) != set(self.last_tool_state.keys()):
            return True
        
        # Check for content changes
        for name, tool in current_state.items():
            if tool != self.last_tool_state.get(name):
                return True
        
        return False
    
    def _should_reconnect(self) -> bool:
        """
        Check if we should attempt to reconnect.
        
        Returns:
            True if reconnection should be attempted
        """
        if not self.running:
            return False
        
        if self.config.max_reconnect_attempts == -1:
            return True
        
        self.reconnect_attempts += 1
        return self.reconnect_attempts <= self.config.max_reconnect_attempts
    
    async def _wait_before_reconnect(self):
        """Wait before attempting to reconnect."""
        wait_time = min(
            self.config.reconnect_interval * (2 ** min(self.reconnect_attempts - 1, 5)),
            60  # Max 60 seconds
        )
        logger.info(f"Waiting {wait_time}s before reconnect attempt {self.reconnect_attempts}")
        await asyncio.sleep(wait_time)