"""
Base WebSocket connector with reconnection and heartbeat logic.
"""
import asyncio
import ssl
from abc import ABC, abstractmethod
from enum import Enum
from typing import Callable, Awaitable
import time

import websockets
from websockets.client import WebSocketClientProtocol

from models import Tick
from utils import get_logger
from config import RECONNECT_BASE_DELAY_S, RECONNECT_MAX_DELAY_S


class ConnectionState(Enum):
    DISCONNECTED = "disconnected"
    CONNECTING = "connecting"
    CONNECTED = "connected"
    RECONNECTING = "reconnecting"


TickCallback = Callable[[Tick], Awaitable[None]]


class BaseConnector(ABC):
    """
    Abstract base class for exchange WebSocket connectors.

    Subclasses must implement:
    - _get_endpoint() -> str
    - _get_subscribe_message(symbols) -> dict | list[dict]
    - _parse_message(data) -> Tick | None
    - _get_heartbeat_interval() -> float | None
    - _get_heartbeat_payload() -> str | dict | None
    """

    def __init__(
        self,
        exchange_name: str,
        on_tick: TickCallback | None = None,
    ):
        self.exchange_name = exchange_name
        self.on_tick = on_tick
        self.state = ConnectionState.DISCONNECTED
        self.ws: WebSocketClientProtocol | None = None
        self.symbols: list[str] = []
        self._reconnect_delay = RECONNECT_BASE_DELAY_S
        self._last_reconnect_time: float = 0
        self._running = False
        self._tasks: list[asyncio.Task] = []
        self.logger = get_logger(f"connector.{exchange_name}")

    @abstractmethod
    def _get_endpoint(self) -> str:
        """Return the WebSocket endpoint URL."""
        pass

    @abstractmethod
    def _get_subscribe_message(self, symbols: list[str]) -> dict | list[dict]:
        """Return the subscription message(s) for the given symbols."""
        pass

    @abstractmethod
    def _parse_message(self, data: dict) -> Tick | None:
        """Parse a WebSocket message and return a Tick, or None if not a tick."""
        pass

    @abstractmethod
    def _get_heartbeat_interval(self) -> float | None:
        """Return heartbeat interval in seconds, or None if not needed."""
        pass

    @abstractmethod
    def _get_heartbeat_payload(self) -> str | dict | None:
        """Return the heartbeat payload to send."""
        pass

    def _get_min_reconnect_interval(self) -> float:
        """Return minimum time between reconnection attempts. Override for rate limits."""
        return 0.0

    def _get_ssl_context(self) -> ssl.SSLContext | None:
        """Return custom SSL context for connection. Override if exchange needs special SSL handling."""
        return None

    async def connect(self) -> None:
        """Establish WebSocket connection."""
        self.state = ConnectionState.CONNECTING
        endpoint = self._get_endpoint()
        self.logger.info(f"Connecting to {endpoint}")

        try:
            # Get optional custom SSL context (some exchanges need special SSL handling)
            ssl_context = self._get_ssl_context()

            connect_kwargs = {
                "ping_interval": 20,
                "ping_timeout": 10,
                "close_timeout": 5,
                "open_timeout": 15,  # Increased timeout for slow connections
            }
            if ssl_context:
                connect_kwargs["ssl"] = ssl_context

            self.ws = await websockets.connect(endpoint, **connect_kwargs)
            self.state = ConnectionState.CONNECTED
            self._reconnect_delay = RECONNECT_BASE_DELAY_S
            self.logger.info("Connected successfully")
        except Exception as e:
            self.state = ConnectionState.DISCONNECTED
            self.logger.error(f"Connection failed: {e}")
            raise

    async def disconnect(self) -> None:
        """Gracefully close the connection."""
        self._running = False
        for task in self._tasks:
            task.cancel()
        self._tasks.clear()

        if self.ws:
            await self.ws.close()
            self.ws = None
        self.state = ConnectionState.DISCONNECTED
        self.logger.info("Disconnected")

    async def subscribe(self, symbols: list[str]) -> None:
        """Subscribe to ticker updates for the given symbols."""
        if not self.ws or self.state != ConnectionState.CONNECTED:
            raise RuntimeError("Not connected")

        self.symbols = symbols
        messages = self._get_subscribe_message(symbols)

        if not isinstance(messages, list):
            messages = [messages]

        for msg in messages:
            import orjson
            await self.ws.send(orjson.dumps(msg))
            self.logger.debug(f"Sent subscribe: {msg}")

    async def run(self) -> None:
        """Main run loop: receive messages and handle reconnection."""
        self._running = True

        # Start heartbeat task if needed
        heartbeat_interval = self._get_heartbeat_interval()
        if heartbeat_interval:
            task = asyncio.create_task(self._heartbeat_loop(heartbeat_interval))
            self._tasks.append(task)

        while self._running:
            # If not connected, attempt to reconnect first
            if not self.ws or self.state != ConnectionState.CONNECTED:
                if self._running:
                    await self._reconnect()
                continue

            try:
                await self._receive_loop()
            except websockets.ConnectionClosed as e:
                self.logger.warning(f"Connection closed: {e}")
                if self._running:
                    await self._reconnect()
            except Exception as e:
                self.logger.error(f"Error in receive loop: {e}")
                if self._running:
                    await self._reconnect()

    async def _receive_loop(self) -> None:
        """Receive and process messages."""
        import orjson

        if not self.ws:
            return

        async for message in self.ws:
            if not self._running:
                break

            try:
                # Handle string pong responses
                if isinstance(message, str) and message.lower() == "pong":
                    continue

                data = orjson.loads(message)
                tick = self._parse_message(data)
                if tick and self.on_tick:
                    await self.on_tick(tick)
            except Exception as e:
                self.logger.debug(f"Failed to parse message: {e}")

    async def _heartbeat_loop(self, interval: float) -> None:
        """Send periodic heartbeat messages."""
        import orjson

        while self._running:
            await asyncio.sleep(interval)
            if self.ws and self.state == ConnectionState.CONNECTED:
                try:
                    payload = self._get_heartbeat_payload()
                    if payload:
                        if isinstance(payload, str):
                            await self.ws.send(payload)
                        else:
                            await self.ws.send(orjson.dumps(payload))
                        self.logger.debug("Sent heartbeat")
                except Exception as e:
                    self.logger.warning(f"Heartbeat failed: {e}")

    async def _reconnect(self) -> None:
        """Reconnect with exponential backoff."""
        self.state = ConnectionState.RECONNECTING

        # Respect minimum reconnect interval (for exchanges like Coinbase)
        min_interval = self._get_min_reconnect_interval()
        if min_interval > 0:
            elapsed = time.time() - self._last_reconnect_time
            if elapsed < min_interval:
                wait_time = min_interval - elapsed
                self.logger.info(f"Rate limit: waiting {wait_time:.1f}s before reconnect")
                await asyncio.sleep(wait_time)

        while self._running:
            self.logger.info(f"Reconnecting in {self._reconnect_delay:.1f}s...")
            await asyncio.sleep(self._reconnect_delay)

            try:
                self._last_reconnect_time = time.time()
                await self.connect()
                if self.symbols:
                    await self.subscribe(self.symbols)
                return
            except Exception as e:
                self.logger.error(f"Reconnection failed: {e}")
                self._reconnect_delay = min(
                    self._reconnect_delay * 2,
                    RECONNECT_MAX_DELAY_S,
                )
