#!/usr/bin/env python3
"""
Crypto Arbitrage Detection System

Dual-strategy arbitrage model:
1. Spread Arbitrage: Buy low on one exchange, sell high on another
2. Lead-Lag Model: Predict altcoin moves from BTC signals using statistical analysis

Uses Binance as the price leader and Coinbase, Bybit, OKX as lag targets.

Usage:
    python main.py
    python main.py --symbols BTC ETH SOL
    python main.py --threshold 15
    python main.py --no-lead-lag  # Disable lead-lag model
"""

# Install uvloop for high-performance async I/O (critical for HFT)
# This MUST happen before any asyncio imports/usage
try:
    import uvloop
    uvloop.install()
    _USING_UVLOOP = True
except ImportError:
    _USING_UVLOOP = False

import asyncio
import signal
import sys
from argparse import ArgumentParser
from decimal import Decimal

from config import (
    SYMBOLS,
    SPREAD_THRESHOLD_BPS,
    PAPER_TRADE_SIZE_USD,
    ENABLE_LEAD_LAG_MODEL,
    EXECUTION_MODE,
    LIVE_TRADE_SIZE_USD,
    RUST_ORDER_MANAGER_HOST,
    RUST_ORDER_MANAGER_PORT,
)
from connectors import BinanceConnector, CoinbaseConnector, BybitConnector, OKXConnector
from engine import ArbitrageDetector
from engine.paper_trader import PaperTrader
from engine.execution import UDPSignalDispatcher, UDPDispatcherConfig
from utils import setup_logging, get_logger, shutdown_logging


logger = get_logger("main")


async def main(
    symbols: list[str] | None = None,
    threshold_bps: Decimal | None = None,
    trade_size: Decimal | None = None,
    enable_lead_lag: bool | None = None,
    live_mode: bool = False,
) -> None:
    """
    Main entry point for the arbitrage detection system.

    Args:
        symbols: List of canonical symbols to track. Defaults to config.SYMBOLS.
        threshold_bps: Minimum spread threshold in bps. Defaults to config value.
        trade_size: Paper trade size in USD. Defaults to config value.
        enable_lead_lag: Enable lead-lag model. Defaults to config value.
        live_mode: If True, send signals to Rust order manager via UDP.
    """
    symbols = symbols or SYMBOLS
    threshold_bps = threshold_bps or SPREAD_THRESHOLD_BPS
    trade_size = trade_size or (LIVE_TRADE_SIZE_USD if live_mode else PAPER_TRADE_SIZE_USD)
    enable_lead_lag = enable_lead_lag if enable_lead_lag is not None else ENABLE_LEAD_LAG_MODEL

    execution_mode = "LIVE (Rust Order Manager)" if live_mode else "Paper Trading"

    logger.info("=" * 70)
    logger.info("CRYPTO ARBITRAGE DETECTION SYSTEM")
    logger.info("=" * 70)
    if _USING_UVLOOP:
        logger.info("Event Loop:   uvloop (high-perf)")
    else:
        logger.warning("Event Loop:   asyncio (standard) - install uvloop for HFT performance!")
    logger.info(f"Symbols:      {', '.join(symbols)}")
    logger.info(f"Threshold:    {threshold_bps} bps (spread arbitrage)")
    logger.info(f"Trade Size:   ${trade_size}")
    logger.info(f"Mode:         {execution_mode}")
    if live_mode:
        logger.info(f"              -> UDP to {RUST_ORDER_MANAGER_HOST}:{RUST_ORDER_MANAGER_PORT}")
    logger.info(f"Lead-Lag:     {'ENABLED' if enable_lead_lag else 'DISABLED'}")
    if enable_lead_lag:
        logger.info("              (BTC → Altcoin predictive signals)")
    logger.info("=" * 70)

    # Initialize execution handler based on mode
    paper_trader = None
    udp_dispatcher = None

    if live_mode:
        # Live mode: send signals to Rust order manager via UDP
        udp_dispatcher = UDPSignalDispatcher(UDPDispatcherConfig(
            rust_host=RUST_ORDER_MANAGER_HOST,
            rust_port=RUST_ORDER_MANAGER_PORT,
            trade_size_usd=float(trade_size),
        ))

        async def handle_spread_signal(spread):
            await udp_dispatcher.dispatch_spread(spread)

        async def handle_lead_lag_signal(signal):
            await udp_dispatcher.dispatch_lead_lag(signal)

        logger.warning("=" * 70)
        logger.warning("LIVE MODE ACTIVE - Real orders will be sent!")
        logger.warning("Ensure Rust order manager is running.")
        logger.warning("=" * 70)
    else:
        # Paper trading mode
        paper_trader = PaperTrader(trade_size_usd=trade_size)

        async def handle_spread_signal(spread):
            await paper_trader.execute(spread)

        async def handle_lead_lag_signal(signal):
            await paper_trader.execute_lead_lag(
                signal,
                entry_price=Decimal(str(signal.entry_price)),
                exchange=signal.exchange,
            )

    # Initialize arbitrage detector with callbacks for both signal types
    detector = ArbitrageDetector(
        on_signal=handle_spread_signal,
        on_lead_lag_signal=handle_lead_lag_signal,
        threshold_bps=threshold_bps,
        enable_lead_lag=enable_lead_lag,
    )

    # Combined tick handler: feeds detector and paper trader (if in paper mode)
    async def on_tick(tick):
        await detector.on_tick(tick)
        if paper_trader:
            await paper_trader.on_tick(tick)  # For position closing

    # Initialize exchange connectors
    # Note: OKX disabled due to connectivity issues from some regions
    connectors = [
        BinanceConnector(on_tick=on_tick),
        CoinbaseConnector(on_tick=on_tick),
        BybitConnector(on_tick=on_tick),
        # OKXConnector(on_tick=on_tick),  # Re-enable when connectivity stable
    ]

    # Track connector names for logging
    connector_names = [c.exchange_name for c in connectors]
    logger.info(f"Exchanges:    {', '.join(connector_names)}")

    # Set up graceful shutdown
    shutdown_event = asyncio.Event()

    def handle_shutdown(sig):
        logger.info(f"Received {sig.name}, initiating shutdown...")
        shutdown_event.set()

    loop = asyncio.get_running_loop()
    for sig in (signal.SIGINT, signal.SIGTERM):
        loop.add_signal_handler(sig, handle_shutdown, sig)

    try:
        # Connect to all exchanges
        logger.info("Connecting to exchanges...")
        connect_results = await asyncio.gather(
            *[c.connect() for c in connectors],
            return_exceptions=True,
        )

        # Check for connection failures
        for connector, result in zip(connectors, connect_results):
            if isinstance(result, Exception):
                logger.error(f"Failed to connect to {connector.exchange_name}: {result}")
            else:
                logger.info(f"Connected to {connector.exchange_name}")

        # Subscribe to symbols
        logger.info(f"Subscribing to {len(symbols)} symbols...")
        await asyncio.gather(
            *[c.subscribe(symbols) for c in connectors],
            return_exceptions=True,
        )

        # Start paper trader summary loop (paper mode only)
        if paper_trader:
            await paper_trader.start_summary_loop()

        # Run all connectors concurrently
        logger.info("Starting data feeds...")
        logger.info("-" * 70)
        run_tasks = [asyncio.create_task(c.run()) for c in connectors]

        # Also create a task that waits for shutdown
        shutdown_task = asyncio.create_task(shutdown_event.wait())

        # Wait for shutdown signal or connector failure
        done, pending = await asyncio.wait(
            run_tasks + [shutdown_task],
            return_when=asyncio.FIRST_COMPLETED,
        )

        # Cancel pending tasks
        for task in pending:
            task.cancel()
            try:
                await task
            except asyncio.CancelledError:
                pass

    except Exception as e:
        logger.error(f"Fatal error: {e}")
        raise
    finally:
        # Graceful shutdown
        logger.info("Shutting down...")

        # Stop paper trader and print final summary (paper mode only)
        if paper_trader:
            await paper_trader.stop()

        # Close UDP dispatcher (live mode only)
        if udp_dispatcher:
            stats = udp_dispatcher.get_stats()
            logger.info(f"UDP Dispatcher stats: {stats}")
            udp_dispatcher.close()

        # Print detector stats
        stats = detector.get_stats()
        logger.info(f"Detector stats: {stats}")

        # Disconnect all connectors
        for connector in connectors:
            try:
                await connector.disconnect()
            except Exception as e:
                logger.warning(f"Error disconnecting {connector.exchange_name}: {e}")

        logger.info("Shutdown complete.")


def parse_args():
    """Parse command line arguments."""
    parser = ArgumentParser(description="Crypto Arbitrage Detection System")

    parser.add_argument(
        "--symbols",
        nargs="+",
        default=None,
        help=f"Symbols to track (default: {SYMBOLS})",
    )

    parser.add_argument(
        "--threshold",
        type=float,
        default=None,
        help=f"Minimum spread threshold in bps (default: {SPREAD_THRESHOLD_BPS})",
    )

    parser.add_argument(
        "--trade-size",
        type=float,
        default=None,
        help=f"Paper trade size in USD (default: {PAPER_TRADE_SIZE_USD})",
    )

    parser.add_argument(
        "--log-level",
        choices=["DEBUG", "INFO", "WARNING", "ERROR"],
        default="INFO",
        help="Logging level (default: INFO)",
    )

    # Lead-lag model flags
    lead_lag_group = parser.add_mutually_exclusive_group()
    lead_lag_group.add_argument(
        "--lead-lag",
        action="store_true",
        dest="enable_lead_lag",
        default=None,
        help="Enable lead-lag model (default: from config)",
    )
    lead_lag_group.add_argument(
        "--no-lead-lag",
        action="store_false",
        dest="enable_lead_lag",
        help="Disable lead-lag model",
    )

    # Execution mode
    parser.add_argument(
        "--live",
        action="store_true",
        default=False,
        help="Enable live trading via Rust order manager (default: paper trading)",
    )

    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()

    # Set up logging
    import logging
    log_level = getattr(logging, args.log_level)
    setup_logging(level=log_level)

    try:
        asyncio.run(
            main(
                symbols=args.symbols,
                threshold_bps=Decimal(str(args.threshold)) if args.threshold else None,
                trade_size=Decimal(str(args.trade_size)) if args.trade_size else None,
                enable_lead_lag=args.enable_lead_lag,
                live_mode=args.live,
            )
        )
    except KeyboardInterrupt:
        pass
    finally:
        shutdown_logging()
