"""Scheduler for neural network scanner."""

from __future__ import annotations

import time
from datetime import datetime
from threading import Thread, Event, Lock
from typing import Dict, Optional, Tuple
from zoneinfo import ZoneInfo

from loguru import logger

_MS = ZoneInfo("Europe/Moscow")

_DEFAULT_FULL_SCAN_INTERVAL = 30 * 60  # 30 minutes


class HourlyScheduler:
    """Periodically scans instruments for trading signals and manages positions."""

    def __init__(
        self,
        scanner,
        interval_seconds: int = 60,
        full_scan_interval: int = _DEFAULT_FULL_SCAN_INTERVAL,
    ):
        self.scanner = scanner
        self.interval_seconds = interval_seconds
        self.full_scan_interval = full_scan_interval
        self.running = False
        self.stop_event = Event()
        self.candle_states: Dict[str, Tuple[int, float]] = {}
        self._last_full_scan_ts: float = 0.0
        self._last_status_ts: float = 0.0
        self.lock = Lock()

    def start(self):
        self.running = True
        self.stop_event.clear()

        logger.info("Scanner scheduler started")
        self._full_scan()

        def run_loop():
            while self.running and not self.stop_event.is_set():
                try:
                    # 1. Check open position exits (TP/SL)
                    self._check_position_exits()

                    # 2. Check for new candles → generate + execute signals
                    self._check_new_candles()

                except Exception as e:
                    logger.error(f"Scan loop error: {e}")

                # 3. Periodic full scan — update all candle states
                now = time.time()
                if now - self._last_full_scan_ts >= self.full_scan_interval:
                    try:
                        self._full_scan()
                    except Exception as e:
                        logger.error(f"Periodic full scan error: {e}")

                self.stop_event.wait(self.interval_seconds)
            logger.info("Scanner loop stopped")

        Thread(target=run_loop, daemon=True).start()

    def stop(self):
        self.running = False
        self.stop_event.set()

    # ── Position exit checking ──────────────────────────────────────────────

    def _check_position_exits(self):
        """Check all open positions for TP/SL hits (delegates to scanner)."""
        if not hasattr(self.scanner, 'check_all_exits'):
            return
        try:
            closed = self.scanner.check_all_exits()
            for pos in closed:
                logger.info(
                    f"[{pos['ticker']}] {pos['side']} {pos['status']} | "
                    f"PnL={pos['pnl']:+.2f} ({pos.get('pnl_pct', 0):+.2f}%) | "
                    f"RR={pos['rr']:.2f} | Capital={pos['capital']:,.2f}"
                )
        except Exception as e:
            logger.error(f"Exit check error: {e}")

    # ── Full scan (periodic) ────────────────────────────────────────────────

    def _full_scan(self):
        """Update candle states for all instruments. Logs open positions only."""
        instruments = self.scanner.get_all_instruments()
        open_tickers = []

        for ticker in instruments:
            state = self.scanner.get_latest_candle_state(ticker)
            cur_price = self.scanner.get_latest_price(ticker) or 0
            if state:
                with self.lock:
                    old_state = self.candle_states.get(ticker)
                    self.candle_states[ticker] = (state["timestamp"], state["close"])

                # Check if this ticker has an open position
                if hasattr(self.scanner, 'vt_service'):
                    trader = self.scanner.vt_service.monitor.traders.get(ticker)
                    if trader and trader.positions:
                        open_tickers.append(ticker)

        self._last_full_scan_ts = time.time()

        if open_tickers:
            logger.info(
                f"Open positions: {len(open_tickers)} — {', '.join(open_tickers)}"
            )

    # ── New candle detection ────────────────────────────────────────────────

    def _check_new_candles(self):
        """Detect new candles, re-scan affected tickers, execute new signals."""
        instruments = self.scanner.get_all_instruments()
        now = datetime.now(_MS).strftime('%H:%M:%S')
        changes = 0

        for ticker in instruments:
            state = self.scanner.get_latest_candle_state(ticker)
            if not state:
                continue

            current_ts = state["timestamp"]
            current_close = state["close"]

            with self.lock:
                last_state = self.candle_states.get(ticker)

            changed = False
            reason = None
            if last_state is None:
                changed = True
                reason = "init"
            elif current_ts != last_state[0]:
                changed = True
                reason = "new candle"
            elif abs(current_close - last_state[1]) > 1e-10:
                changed = True
                reason = "price update"

            if not changed:
                continue

            prev_ts = last_state[0] if last_state else 0

            with self.lock:
                self.candle_states[ticker] = (current_ts, current_close)

            # Log new candles at INFO
            if reason in ("new candle", "init"):
                ts_str = datetime.fromtimestamp(current_ts, _MS).strftime('%H:%M')
                logger.info(f"[{ticker}] new candle {ts_str}")
            else:
                logger.debug(f"[{ticker}] {reason}")

            changes += 1

            # Generate and execute signals from new bars only
            signals = self.scanner.scan_instrument(ticker, lookback_hours=72)
            new_signals = [
                s for s in signals
                if s.get("timestamp", 0) > prev_ts
            ]
            for sig in new_signals:
                sig_ticker = sig.get("ticker", ticker)
                side = sig.get("signal_type", "?")
                entry = sig.get("entry_price", 0)
                sl = sig.get("sl_price", 0)
                tp = sig.get("tp_price", 0)
                proba = sig.get("entry_probability", 0)

                # Attempt execution
                result = None
                if hasattr(self.scanner, 'execute_signal'):
                    try:
                        result = self.scanner.execute_signal(sig)
                    except Exception as e:
                        logger.error(f"[{sig_ticker}] Execution error: {e}")

                opened = "✅" if result else "⏭"
                logger.info(
                    f"[{sig_ticker}] {side} @ {entry:.2f} | "
                    f"SL={sl:.2f} TP={tp:.2f} | "
                    f"P={proba:.0%} | {opened}"
                )

        if changes > 0:
            logger.info(f"[{now}] {changes} changes, scanned {changes} tickers")
