"""
UDP Signal Dispatcher for Rust Order Manager.

Sends trading signals to the Rust order manager via UDP using a fixed-size
binary protocol for minimum latency.

Protocol:
- Header (8 bytes): message_type (1), version (1), sequence (2), checksum (4)
- LeadLag body (64 bytes): see _pack_lead_lag_body
- Spread body (56 bytes): see _pack_spread_body
"""
import socket
import struct
import time
import zlib
from dataclasses import dataclass
from typing import Callable, Awaitable

from engine.lead_lag_brain import LeadLagSignal, SignalDirection
from engine.spread_calculator import SpreadResult
from config import (
    RUST_ORDER_MANAGER_HOST,
    RUST_ORDER_MANAGER_PORT,
    RUST_ACK_PORT,
    EXCHANGE_IDS,
    LIVE_TRADE_SIZE_USD,
)
from utils import get_logger

logger = get_logger(__name__)

# Message types
MSG_LEAD_LAG = 0x01
MSG_SPREAD = 0x02
MSG_HEARTBEAT = 0x03
MSG_ACK = 0x10
MSG_FILL = 0x20

# Protocol version
PROTOCOL_VERSION = 0x01


@dataclass
class UDPDispatcherConfig:
    """Configuration for the UDP signal dispatcher."""
    rust_host: str = RUST_ORDER_MANAGER_HOST
    rust_port: int = RUST_ORDER_MANAGER_PORT
    ack_port: int = RUST_ACK_PORT
    enable_ack_listener: bool = False
    trade_size_usd: float = float(LIVE_TRADE_SIZE_USD)


class UDPSignalDispatcher:
    """
    Dispatches trading signals to Rust order manager via UDP.

    Uses a fixed-size binary protocol for minimum latency:
    - No JSON parsing overhead
    - No string allocations
    - Fixed packet sizes for predictable memory

    Protocol Design:
    - All multi-byte integers are big-endian (network byte order)
    - Floats are IEEE 754 double precision (8 bytes)
    - Symbols are 4-byte null-padded ASCII strings
    - Exchange IDs are single bytes (0=binance, 1=coinbase, 2=bybit)
    """

    # Struct formats (big-endian)
    # Header: type(1) + version(1) + sequence(2) + checksum(4) = 8 bytes
    HEADER_FORMAT = "!BBHI"

    # LeadLag body: timestamp(8) + symbol(4) + direction(1) + exchange(1) + pad(2)
    #              + entry_price(8) + expected_return(8) + return_gap(8)
    #              + beta(8) + confidence(4) + max_hold_ms(4) + quantity(8)
    #              = 64 bytes
    LEAD_LAG_BODY_FORMAT = "!q4sBBxxddddfidd"

    # Spread body: timestamp(8) + symbol(4) + buy_exch(1) + sell_exch(1) + pad(2)
    #             + buy_price(8) + sell_price(8) + spread_bps(8) + quantity(8)
    #             = 48 bytes + 8 reserved = 56 bytes
    SPREAD_BODY_FORMAT = "!q4sBBxxdddd8x"

    def __init__(self, config: UDPDispatcherConfig | None = None):
        self.config = config or UDPDispatcherConfig()
        self._socket = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
        self._socket.setblocking(False)
        self._sequence = 0
        self._pending_signals: dict[int, float] = {}  # seq -> timestamp
        self._signals_sent = 0
        self._signals_acked = 0

        logger.info(
            f"UDPSignalDispatcher initialized: "
            f"{self.config.rust_host}:{self.config.rust_port}"
        )

    def _next_seq(self) -> int:
        """Get next sequence number (wraps at 65535)."""
        self._sequence = (self._sequence + 1) % 65536
        return self._sequence

    def _compute_checksum(self, data: bytes) -> int:
        """Compute CRC32 checksum for packet integrity."""
        return zlib.crc32(data) & 0xFFFFFFFF

    def _pack_header(self, msg_type: int, body: bytes) -> bytes:
        """Pack message header with checksum of body."""
        seq = self._next_seq()
        checksum = self._compute_checksum(body)
        return struct.pack(self.HEADER_FORMAT, msg_type, PROTOCOL_VERSION, seq, checksum)

    def _symbol_to_bytes(self, symbol: str) -> bytes:
        """Convert symbol to 4-byte null-padded ASCII."""
        return symbol.encode("ascii")[:4].ljust(4, b"\x00")

    def _pack_lead_lag_body(self, signal: LeadLagSignal, quantity: float) -> bytes:
        """Pack LeadLag signal into binary body."""
        symbol_bytes = self._symbol_to_bytes(signal.lagger_symbol)
        direction = 0 if signal.direction == SignalDirection.LONG else 1
        exchange_id = EXCHANGE_IDS.get(signal.exchange, 0)

        return struct.pack(
            self.LEAD_LAG_BODY_FORMAT,
            signal.timestamp_ms,
            symbol_bytes,
            direction,
            exchange_id,
            # After 2-byte padding (handled by xx in format)
            signal.entry_price,
            signal.expected_return_pct / 100,  # Convert % to decimal
            signal.return_gap_pct / 100,
            signal.rolling_beta,
            signal.confidence,
            signal.max_hold_ms,
            quantity,
            signal.rolling_correlation,
        )

    def _pack_spread_body(self, spread: SpreadResult, quantity: float) -> bytes:
        """Pack Spread signal into binary body."""
        symbol_bytes = self._symbol_to_bytes(spread.symbol)
        buy_exch_id = EXCHANGE_IDS.get(spread.buy_exchange, 0)
        sell_exch_id = EXCHANGE_IDS.get(spread.sell_exchange, 0)

        return struct.pack(
            self.SPREAD_BODY_FORMAT,
            int(time.time() * 1000),
            symbol_bytes,
            buy_exch_id,
            sell_exch_id,
            # After 2-byte padding
            float(spread.buy_price),
            float(spread.sell_price),
            float(spread.spread_bps),
            quantity,
            # 8 bytes reserved (handled by 8x)
        )

    async def dispatch_lead_lag(self, signal: LeadLagSignal) -> None:
        """
        Send lead-lag signal to Rust order manager.

        Args:
            signal: LeadLagSignal from the lead-lag brain
        """
        # Calculate quantity based on trade size and entry price
        if signal.entry_price > 0:
            quantity = self.config.trade_size_usd / signal.entry_price
        else:
            quantity = 0.0

        body = self._pack_lead_lag_body(signal, quantity)
        header = self._pack_header(MSG_LEAD_LAG, body)
        packet = header + body

        try:
            self._socket.sendto(
                packet,
                (self.config.rust_host, self.config.rust_port)
            )
            self._pending_signals[self._sequence] = time.time()
            self._signals_sent += 1

            logger.info(
                f"Dispatched LeadLag signal: {signal.lagger_symbol} "
                f"{signal.direction.value} @{signal.exchange} "
                f"seq={self._sequence}"
            )
        except Exception as e:
            logger.error(f"Failed to dispatch LeadLag signal: {e}")

    async def dispatch_spread(self, spread: SpreadResult) -> None:
        """
        Send spread signal to Rust order manager.

        Args:
            spread: SpreadResult from the spread calculator
        """
        # Calculate quantity based on trade size and buy price
        if spread.buy_price > 0:
            quantity = self.config.trade_size_usd / float(spread.buy_price)
        else:
            quantity = 0.0

        body = self._pack_spread_body(spread, quantity)
        header = self._pack_header(MSG_SPREAD, body)
        packet = header + body

        try:
            self._socket.sendto(
                packet,
                (self.config.rust_host, self.config.rust_port)
            )
            self._pending_signals[self._sequence] = time.time()
            self._signals_sent += 1

            logger.info(
                f"Dispatched Spread signal: {spread.symbol} "
                f"BUY@{spread.buy_exchange} SELL@{spread.sell_exchange} "
                f"{spread.spread_bps:.1f}bps seq={self._sequence}"
            )
        except Exception as e:
            logger.error(f"Failed to dispatch Spread signal: {e}")

    async def send_heartbeat(self) -> None:
        """Send heartbeat to Rust order manager."""
        body = struct.pack("!q", int(time.time() * 1000))
        header = self._pack_header(MSG_HEARTBEAT, body)
        packet = header + body

        try:
            self._socket.sendto(
                packet,
                (self.config.rust_host, self.config.rust_port)
            )
        except Exception as e:
            logger.warning(f"Failed to send heartbeat: {e}")

    def get_stats(self) -> dict:
        """Get dispatcher statistics."""
        return {
            "signals_sent": self._signals_sent,
            "signals_acked": self._signals_acked,
            "pending_signals": len(self._pending_signals),
            "current_sequence": self._sequence,
        }

    def close(self) -> None:
        """Close the UDP socket."""
        self._socket.close()
        logger.info("UDPSignalDispatcher closed")
