"""Virtual trading service for neural network scanner."""

from __future__ import annotations

import json
import time
import threading
from datetime import datetime, timedelta, date
from pathlib import Path
from typing import Dict, List, Optional, Set
from loguru import logger

from core.virtual_trader import VirtualTrader, VirtualTradingMonitor
from core.trade_journal import TradeJournal, DailyReportGenerator


class VirtualTradingService:
    def __init__(
        self,
        initial_capital: float = 1_000_000,
        risk_pct: float = 0.01,
        max_positions: int = 5,
        output_dir: Path = Path("logs/virtual_trading"),
    ):
        self.output_dir = output_dir
        self.monitor = VirtualTradingMonitor(output_dir, initial_capital=initial_capital)
        self.journal = TradeJournal(output_dir)
        self.reporter = DailyReportGenerator(output_dir / "reports")

        self.initial_capital = initial_capital
        self.risk_pct = risk_pct
        self.max_positions = max_positions

        self._last_report_date: Optional[date] = None
        self._closed_by_tp_sl: Set[str] = set()
        self._last_save_ts: float = 0.0
        self._state_lock = threading.Lock()

        self._load_state()

    def _load_state(self):
        service_state_file = self.output_dir / "service_state.json"
        if not service_state_file.exists():
            return
        try:
            with open(service_state_file, "r") as f:
                state = json.load(f)
            self._closed_by_tp_sl = set(state.get("closed_by_tp_sl", []))
            logger.info(f"Service state loaded: {len(self._closed_by_tp_sl)} tickers in filter")
        except Exception as e:
            logger.warning(f"Failed to load service state: {e}")

    def save_state(self):
        if not self._state_lock.acquire(blocking=False):
            return
        try:
            self.monitor.save_all_states()
            service_state_file = self.output_dir / "service_state.json"
            state = {"closed_by_tp_sl": list(self._closed_by_tp_sl)}
            with open(service_state_file, "w") as f:
                json.dump(state, f, indent=2)
        except Exception as e:
            logger.error(f"Failed to save state: {e}")
        finally:
            self._state_lock.release()

    def _auto_save(self):
        now = time.time()
        if now - self._last_save_ts > 300:
            self._last_save_ts = now
            self.save_state()

    def execute_signal(self, signal: Dict, current_price: float) -> Optional[Dict]:
        ticker = signal.get("ticker")
        if not ticker:
            return None

        if ticker in self._closed_by_tp_sl:
            return None

        trader = self.monitor.get_or_create_trader(ticker)

        if len(trader.positions) >= self.max_positions:
            return None

        entry_price = signal.get("entry_price", current_price)
        sl_price = signal.get("sl_price", entry_price * 0.99)
        tp_price = signal.get("tp_price", entry_price * 1.02)

        pos = trader.open_position(
            ticker=ticker,
            signal=signal,
            entry_price=entry_price,
            sl_price=sl_price,
            tp_price=tp_price,
        )

        if pos:
            self._auto_save()
            return {"action": "OPENED", "ticker": ticker, "position": pos}
        return None

    def check_daily_report(self):
        today = date.today()
        if self._last_report_date != today:
            yesterday = today - timedelta(days=1)
            self.reporter.generate_daily_report(yesterday)
            self._last_report_date = today

    def get_status(self) -> Dict:
        stats = self.monitor.get_total_stats()
        positions = self.monitor.get_all_positions()
        # If no traders exist yet, show initial capital
        reported_capital = stats.get("capital", 0)
        if reported_capital == 0 and stats.get("total_trades", 0) == 0:
            reported_capital = self.initial_capital
        return {
            "capital": reported_capital,
            "total_pnl": stats.get("total_pnl", 0),
            "total_trades": stats.get("total_trades", 0),
            "win_rate": stats.get("win_rate", 0),
            "open_positions": len(positions),
        }

    def close_all(self):
        closed = self.monitor.close_all_positions()
        self.save_state()
        logger.info(f"Closed {len(closed)} positions")

    def generate_report(self, report_date: Optional[date] = None) -> str:
        return self.reporter.generate_daily_report(report_date)

    def check_all_exits(self, price_getter, current_ts: int = None) -> List[Dict]:
        """Check all open positions for TP/SL hits, close them, and log to journal.

        Args:
            price_getter: Callable(ticker) -> current_price, or Dict[ticker, price]
            current_ts: Current timestamp (defaults to now)

        Returns:
            List of closed position info dicts
        """
        if current_ts is None:
            current_ts = int(datetime.now().timestamp())

        closed_positions = []

        if callable(price_getter):
            get_price = price_getter
        else:
            get_price = lambda t: price_getter.get(t)

        for ticker, trader in self.monitor.traders.items():
            if not trader.positions:
                continue

            for pos_ticker, pos in list(trader.positions.items()):
                current_price = get_price(pos_ticker)
                if current_price is None:
                    continue

                result = trader.check_exits(pos_ticker, current_price, current_ts)
                if result and result.status in ("TP", "SL"):
                    # Calculate PnL as % of entry
                    if result.side == "LONG":
                        pnl_pct = (result.exit_price - result.entry_price) / result.entry_price * 100
                    else:
                        pnl_pct = (result.entry_price - result.exit_price) / result.entry_price * 100

                    # Write to trade journal
                    self.journal.add_trade(
                        ticker=pos_ticker,
                        side=result.side,
                        entry_time=result.entry_time,
                        exit_time=result.exit_time or current_ts,
                        entry_price=result.entry_price,
                        exit_price=result.exit_price or current_price,
                        sl_price=result.sl_price,
                        tp_price=result.tp_price,
                        size=result.size,
                        pnl=result.pnl,
                        rr=result.rr,
                        status=result.status,
                    )

                    closed_positions.append({
                        "ticker": pos_ticker,
                        "side": result.side,
                        "status": result.status,
                        "pnl": result.pnl,
                        "pnl_pct": round(pnl_pct, 2),
                        "rr": result.rr,
                        "exit_price": result.exit_price,
                        "capital": trader.capital,
                    })
                    self._auto_save()

        return closed_positions


class VirtualTradingScanner:
    """Wraps NeuralScanner and VirtualTradingService."""

    def __init__(self, scanner, vt_service: VirtualTradingService):
        self.scanner = scanner
        self.vt_service = vt_service

    @property
    def live_mode(self) -> bool:
        """Delegate live_mode to underlying scanner."""
        return getattr(self.scanner, 'live_mode', False)

    def get_all_instruments(self):
        return self.scanner.get_all_instruments()

    def get_latest_timestamp(self, ticker: str):
        return self.scanner.get_latest_timestamp(ticker)

    def get_latest_candle_state(self, ticker: str):
        return self.scanner.get_latest_candle_state(ticker)

    def get_latest_price(self, ticker: str):
        return self.scanner.get_latest_price(ticker)

    def scan_instrument(self, ticker: str, lookback_hours: int = 72, **kwargs):
        return self.scanner.scan_instrument(ticker, lookback_hours)

    def scan_all_instruments(self, lookback_hours: int = 72, instruments=None):
        return self.scanner.scan_all_instruments(lookback_hours, instruments)

    def log_signal(self, signal: Dict):
        return self.scanner.log_signal(signal)

    def execute_signal(self, signal: Dict):
        ticker = signal.get("ticker")
        current_price = signal.get("entry_price") or self.get_latest_price(ticker)
        if current_price is None:
            return None
        return self.vt_service.execute_signal(signal, current_price)

    def get_status(self) -> Dict:
        return self.vt_service.get_status()

    def close_all(self):
        return self.vt_service.close_all()

    def check_daily_report(self):
        return self.vt_service.check_daily_report()

    def check_all_exits(self, current_ts: int = None) -> List[Dict]:
        """Check all open positions for TP/SL hits using scanner's price getter."""
        return self.vt_service.check_all_exits(self.get_latest_price, current_ts)
