from __future__ import annotations

import json
from datetime import datetime
from pathlib import Path
from typing import Dict, List, Optional
from dataclasses import dataclass, field
from loguru import logger

from core.risk_manager import calculate_position_size


@dataclass
class VirtualPosition:
    ticker: str
    entry_time: int
    entry_price: float
    sl_price: float
    tp_price: float
    size: int
    side: str
    status: str = "OPEN"
    exit_time: Optional[int] = None
    exit_price: Optional[float] = None
    pnl: float = 0.0
    rr: float = 0.0
    signal_data: Dict = field(default_factory=dict)

    @property
    def duration_bars(self) -> int:
        if self.exit_time:
            return (self.exit_time - self.entry_time) // 3600
        return (int(datetime.now().timestamp()) - self.entry_time) // 3600


class VirtualTrader:
    def __init__(
        self,
        initial_capital: float = 1_000_000,
        risk_pct: float = 0.01,
        rr_ratio: float = 2.0,
        max_positions: int = 5,
        state_dir: Optional[Path] = None,
    ):
        self.initial_capital = initial_capital
        self.capital = initial_capital
        self.risk_pct = risk_pct
        self.rr_ratio = rr_ratio
        self.max_positions = max_positions
        self.positions: Dict[str, VirtualPosition] = {}
        self.journal: List[VirtualPosition] = []
        self.state_dir = Path(state_dir) if state_dir else Path("logs/virtual_trading")
        self.state_dir.mkdir(parents=True, exist_ok=True)

    def open_position(
        self,
        ticker: str,
        signal: Dict,
        entry_price: float,
        sl_price: float,
        tp_price: Optional[float] = None,
        size: Optional[int] = None,
    ) -> Optional[VirtualPosition]:
        if len(self.positions) >= self.max_positions:
            logger.warning(f"Max positions reached ({self.max_positions})")
            return None

        if ticker in self.positions:
            logger.warning(f"Position already exists for {ticker}")
            return None

        if tp_price is None:
            risk = abs(entry_price - sl_price)
            if signal.get("signal_type") == "LONG":
                tp_price = entry_price + risk * self.rr_ratio
            else:
                tp_price = entry_price - risk * self.rr_ratio

        side = signal.get("signal_type", "LONG")

        if size is None:
            size = self._calculate_size(entry_price, sl_price, side)

        if size <= 0:
            logger.warning(f"Invalid position size: {size}")
            return None

        pos = VirtualPosition(
            ticker=ticker,
            entry_time=int(signal.get("timestamp", int(datetime.now().timestamp()))),
            entry_price=entry_price,
            sl_price=sl_price,
            tp_price=tp_price,
            size=size,
            side=signal.get("signal_type", "LONG"),
            signal_data=signal,
        )

        self.positions[ticker] = pos
        return pos

    def _calculate_size(self, entry: float, sl: float, side: str = "long") -> int:
        return calculate_position_size(self.capital, self.risk_pct, entry, sl, side=side)

    def check_exits(self, ticker: str, current_price: float, current_ts: int) -> Optional[VirtualPosition]:
        if ticker not in self.positions:
            return None

        pos = self.positions[ticker]
        if pos.status != "OPEN":
            return None

        entry = pos.entry_price
        sl = pos.sl_price
        tp = pos.tp_price
        side = pos.side
        size = pos.size
        pnl = 0.0
        status = "OPEN"
        exit_price = None

        if side == "LONG":
            if current_price <= sl:
                pnl = (sl - entry) * size
                status = "SL"
                exit_price = sl
            elif current_price >= tp:
                pnl = (tp - entry) * size
                status = "TP"
                exit_price = tp
        else:
            if current_price >= sl:
                pnl = (entry - sl) * size
                status = "SL"
                exit_price = sl
            elif current_price <= tp:
                pnl = (entry - tp) * size
                status = "TP"
                exit_price = tp
        if status != "OPEN":
            pos.exit_time = current_ts
            pos.exit_price = exit_price or current_price
            pos.pnl = pnl
            pos.status = status

            risk = abs(entry - sl)
            if risk > 0:
                pos.rr = pnl / (size * risk)

            self.capital += pnl
            self.journal.append(pos)
            del self.positions[ticker]

        return pos

    def close_position(
        self, ticker: str, exit_price: float, status: str = "MANUAL"
    ) -> Optional[VirtualPosition]:
        if ticker not in self.positions:
            return None

        pos = self.positions[ticker]
        size = pos.size
        entry = pos.entry_price
        side = pos.side

        if side == "LONG":
            pnl = (exit_price - entry) * size
        else:
            pnl = (entry - exit_price) * size

        pos.exit_time = int(datetime.now().timestamp())
        pos.exit_price = exit_price
        pos.pnl = pnl
        pos.status = status

        risk = abs(entry - pos.sl_price)
        if risk > 0:
            pos.rr = pnl / (size * risk)

        self.capital += pnl
        self.journal.append(pos)
        del self.positions[ticker]

        return pos

    def close_all(self, prices: Dict[str, float] = None) -> List[VirtualPosition]:
        closed = []
        for ticker in list(self.positions.keys()):
            if prices and ticker in prices:
                self.close_position(ticker, prices[ticker], "EOD_CLOSE")
            else:
                pos = self.positions[ticker]
                self.close_position(ticker, pos.entry_price, "EOD_CLOSE")
            closed.append(self.journal[-1])
        return closed

    def get_stats(self) -> Dict:
        if not self.journal:
            return {
                "total_trades": 0,
                "winning_trades": 0,
                "losing_trades": 0,
                "win_rate": 0.0,
                "total_pnl": 0.0,
                "capital": self.capital,
                "profit_factor": 0.0,
                "avg_rr": 0.0,
            }

        winning = [t for t in self.journal if t.pnl > 0]
        losing = [t for t in self.journal if t.pnl <= 0]

        total_pnl = sum(t.pnl for t in self.journal)
        wins_pnl = sum(t.pnl for t in winning)
        loss_pnl = sum(t.pnl for t in losing)

        avg_rr = sum(t.rr for t in self.journal) / len(self.journal) if self.journal else 0.0

        return {
            "total_trades": len(self.journal),
            "winning_trades": len(winning),
            "losing_trades": len(losing),
            "win_rate": len(winning) / len(self.journal) if self.journal else 0.0,
            "total_pnl": total_pnl,
            "capital": self.capital,
            "profit_factor": abs(wins_pnl / loss_pnl) if loss_pnl != 0 else float("inf"),
            "avg_rr": avg_rr,
        }

    def get_open_positions(self) -> List[Dict]:
        result = []
        for ticker, pos in self.positions.items():
            result.append({
                "ticker": ticker,
                "side": pos.side,
                "entry_price": pos.entry_price,
                "current_price": pos.signal_data.get("current_price", pos.entry_price),
                "sl_price": pos.sl_price,
                "tp_price": pos.tp_price,
                "size": pos.size,
                "unrealized_pnl": self._calc_unrealized_pnl(pos),
                "entry_time": pos.entry_time,
            })
        return result

    def _calc_unrealized_pnl(self, pos: VirtualPosition) -> float:
        current_price = pos.signal_data.get("current_price", pos.entry_price)
        if pos.side == "LONG":
            return (current_price - pos.entry_price) * pos.size
        else:
            return (pos.entry_price - current_price) * pos.size

    def save_state(self, ticker: str = None) -> Path:
        state_file = self.state_dir / (f"{ticker}_state.json" if ticker else "trader_state.json")
        state = {
            "initial_capital": self.initial_capital,
            "capital": self.capital,
            "risk_pct": self.risk_pct,
            "rr_ratio": self.rr_ratio,
            "max_positions": self.max_positions,
            "journal_count": len(self.journal),
            "positions": {},
        }
        for tk, pos in self.positions.items():
            state["positions"][tk] = {
                "ticker": pos.ticker,
                "entry_time": int(pos.entry_time),
                "entry_price": float(pos.entry_price),
                "sl_price": float(pos.sl_price),
                "tp_price": float(pos.tp_price),
                "size": int(pos.size),
                "side": pos.side,
                "status": pos.status,
                "signal_data": {k: v for k, v in pos.signal_data.items()
                                if not callable(getattr(v, "items", None))},
            }
        with open(state_file, "w") as f:
            json.dump(state, f, indent=2, default=str)
        logger.debug(f"State saved to {state_file}")
        return state_file

    def load_state(self, ticker: str = None) -> bool:
        state_file = self.state_dir / (f"{ticker}_state.json" if ticker else "trader_state.json")
        if not state_file.exists():
            return False
        try:
            with open(state_file, "r") as f:
                state = json.load(f)
            self.initial_capital = state.get("initial_capital", self.initial_capital)
            self.capital = state.get("capital", self.initial_capital)
            self.risk_pct = state.get("risk_pct", self.risk_pct)
            self.rr_ratio = state.get("rr_ratio", self.rr_ratio)
            self.max_positions = state.get("max_positions", self.max_positions)
            self.positions.clear()
            for tk, pos_data in state.get("positions", {}).items():
                sd = pos_data.get("signal_data", {})
                self.positions[tk] = VirtualPosition(
                    ticker=pos_data["ticker"],
                    entry_time=int(pos_data["entry_time"]),
                    entry_price=float(pos_data["entry_price"]),
                    sl_price=float(pos_data["sl_price"]),
                    tp_price=float(pos_data["tp_price"]),
                    size=int(pos_data["size"]),
                    side=pos_data["side"],
                    status=pos_data.get("status", "OPEN"),
                    signal_data=sd,
                )
                self.positions[tk].signal_data["current_price"] = sd.get("current_price", pos_data["entry_price"])
            logger.debug(f"State loaded from {state_file}: {len(self.positions)} open positions, capital={self.capital:,.2f}")
            return True
        except Exception as e:
            logger.warning(f"Failed to load state from {state_file}: {e}")
            return False


class VirtualTradingMonitor:
    def __init__(self, output_dir: Path = Path("logs"), initial_capital: float = 1_000_000):
        self.output_dir = output_dir
        self.output_dir.mkdir(parents=True, exist_ok=True)
        self.state_dir = output_dir / "states"
        self.state_dir.mkdir(parents=True, exist_ok=True)
        self.traders: Dict[str, VirtualTrader] = {}
        self.initial_capital = initial_capital

    def get_or_create_trader(self, ticker: str) -> VirtualTrader:
        if ticker not in self.traders:
            trader = VirtualTrader(state_dir=self.state_dir)
            trader.load_state(ticker)
            self.traders[ticker] = trader
        return self.traders[ticker]

    def save_all_states(self) -> List[Path]:
        saved = []
        for ticker, trader in self.traders.items():
            path = trader.save_state(ticker)
            saved.append(path)
        return saved

    def get_all_positions(self) -> List[Dict]:
        all_positions = []
        for ticker, trader in self.traders.items():
            all_positions.extend(trader.get_open_positions())
        return all_positions

    def close_all_positions(self) -> List[VirtualPosition]:
        all_closed = []
        for trader in self.traders.values():
            all_closed.extend(trader.close_all())
        return all_closed

    def get_total_stats(self) -> Dict:
        total_trades = sum(len(t.journal) for t in self.traders.values())
        total_pnl = sum(sum(tr.pnl for tr in t.journal) for t in self.traders.values())
        # Use shared initial capital + total PnL, not sum of per-ticker capitals
        capital = self.initial_capital + total_pnl

        winning = sum(1 for trader in self.traders.values() for t in trader.journal if t.pnl > 0)
        win_rate = winning / total_trades if total_trades > 0 else 0.0

        return {
            "total_trades": total_trades,
            "winning_trades": winning,
            "win_rate": win_rate,
            "total_pnl": total_pnl,
            "capital": capital,
        }