"""Risk management calculations: stop loss, take profit, position sizing."""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Optional, Tuple

from .trend_analysis import Candle, detect_trend, TrendDirection
from .entry_detection import EntrySignal, SignalType, detect_entry


@dataclass
class RiskParameters:
    """Risk calculation parameters."""
    account_balance: float = 10000.0
    risk_percent: float = 1.0  # Risk per trade as % of account
    max_open_trades: int = 5
    slippage: float = 0.1  # Slippage in points
    spread: float = 0.0  # Spread in price units
    max_risk_per_trade_pct: float = 5.0  # Maximum risk as % of entry price


@dataclass
class RiskResult:
    """Result of risk calculation."""
    entry_price: float
    stop_loss: float
    take_profit: float
    position_size: float  # Number of units
    risk_amount: float  # Dollar amount at risk
    risk_pct: float  # Risk as % of account
    reward_risk_ratio: float
    invalid: bool = False
    reason: str = ""


def calculate_stop_loss_take_profit(entry_price: float,
                                     signal_type: SignalType,
                                     second_to_last_low: float,
                                     second_to_last_high: float,
                                     atr: Optional[float] = None,
                                     multiplier: float = 3.0,
                                     atr_multiplier: float = 2.0) -> Tuple[float, float]:
    """
    Calculate stop loss and take profit based on strategy rules.

    Uses ATR-based stops when ATR is available (volatility-scaled).
    Falls back to candle structure when ATR is unavailable.

    Args:
        entry_price: The entry price
        signal_type: BUY or SELL signal
        second_to_last_low: Low of second-to-last candle (fallback)
        second_to_last_high: High of second-to-last candle (fallback)
        atr: Optional ATR value for volatility-based stop distance
        multiplier: Take profit multiplier (default 3x)
        atr_multiplier: ATR multiplier for stop distance (default 2x)

    Returns:
        Tuple of (stop_loss_price, take_profit_price)
    """
    if atr and atr > 0:
        stop_distance = atr * atr_multiplier
    elif signal_type == SignalType.BUY:
        stop_distance = entry_price - second_to_last_low
        if stop_distance <= 0:
            stop_distance = entry_price * 0.02  # Fallback: 2%
    else:
        stop_distance = second_to_last_high - entry_price
        if stop_distance <= 0:
            stop_distance = entry_price * 0.02  # Fallback: 2%

    if signal_type == SignalType.BUY:
        stop_loss = entry_price - stop_distance
        take_profit = entry_price + (stop_distance * multiplier)
    else:
        stop_loss = entry_price + stop_distance
        take_profit = entry_price - (stop_distance * multiplier)

    return stop_loss, take_profit


def calculate_position_size(entry_price: float,
                             stop_loss: float,
                             risk_params: RiskParameters) -> Tuple[float, float, float]:
    """
    Calculate position size based on risk management rules.

    Args:
        entry_price: Entry price per unit
        stop_loss: Stop loss price
        risk_params: Risk parameters

    Returns:
        Tuple of (position_size_units, risk_amount, risk_pct)
    """
    risk_amount = risk_params.account_balance * (risk_params.risk_percent / 100.0)

    price_risk = abs(entry_price - stop_loss)
    if price_risk == 0:
        return 0.0, 0.0, 0.0

    # Account for slippage
    effective_risk_per_unit = price_risk + risk_params.slippage

    position_size = risk_amount / effective_risk_per_unit

    # Cap position size to max risk per trade
    max_position = (risk_params.account_balance * risk_params.max_risk_per_trade_pct / 100.0) / entry_price
    position_size = min(position_size, max_position)

    risk_pct = (risk_amount / risk_params.account_balance) * 100.0

    return round(position_size, 4), round(risk_amount, 2), round(risk_pct, 4)


def calculate_atr(candles: List[Candle], period: int = 14) -> float:
    """
    Calculate Average True Range for volatility-based calculations.

    Args:
        candles: List of candles (oldest to newest)
        period: ATR period

    Returns:
        ATR value
    """
    if len(candles) < period + 1:
        return 0.0

    true_ranges = []
    for i in range(1, len(candles)):
        high = candles[i].high
        low = candles[i].low
        prev_close = candles[i - 1].close

        tr = max(
            high - low,
            abs(high - prev_close),
            abs(low - prev_close)
        )
        true_ranges.append(tr)

    if len(true_ranges) < period:
        return 0.0

    return sum(true_ranges[-period:]) / period


def validate_trade(entry_signal: EntrySignal,
                    risk_params: RiskParameters,
                    current_position_size: float = 0.0) -> RiskResult:
    """
    Validate a trade against risk management rules.

    Args:
        entry_signal: The detected entry signal
        risk_params: Risk parameters
        current_position_size: Currently open position size

    Returns:
        RiskResult with validation outcome
    """
    if entry_signal.stop_loss is None:
        return RiskResult(
            entry_price=entry_signal.price,
            stop_loss=0.0,
            take_profit=0.0,
            position_size=0.0,
            risk_amount=0.0,
            risk_pct=0.0,
            reward_risk_ratio=0.0,
            invalid=True,
            reason="No stop loss defined for entry signal"
        )

    # Check max open trades
    if current_position_size >= risk_params.max_open_trades:
        return RiskResult(
            entry_price=entry_signal.price,
            stop_loss=entry_signal.stop_loss,
            take_profit=entry_signal.take_profit or 0.0,
            position_size=0.0,
            risk_amount=0.0,
            risk_pct=0.0,
            reward_risk_ratio=0.0,
            invalid=True,
            reason=f"Maximum open trades ({risk_params.max_open_trades}) reached"
        )

    # Calculate position size
    if entry_signal.take_profit:
        reward_risk = abs(entry_signal.take_profit - entry_signal.price) / abs(entry_signal.price - entry_signal.stop_loss)
    else:
        reward_risk = 0.0

    position_size, risk_amount, risk_pct = calculate_position_size(
        entry_signal.price,
        entry_signal.stop_loss,
        risk_params
    )

    # Validate risk per trade limit
    price_risk_pct = abs(entry_signal.price - entry_signal.stop_loss) / entry_signal.price * 100
    if price_risk_pct > risk_params.max_risk_per_trade_pct:
        return RiskResult(
            entry_price=entry_signal.price,
            stop_loss=entry_signal.stop_loss,
            take_profit=entry_signal.take_profit or 0.0,
            position_size=0.0,
            risk_amount=0.0,
            risk_pct=0.0,
            reward_risk_ratio=reward_risk,
            invalid=True,
            reason=f"Price risk ({price_risk_pct:.2f}%) exceeds max ({risk_params.max_risk_per_trade_pct}%)"
        )

    return RiskResult(
        entry_price=entry_signal.price,
        stop_loss=entry_signal.stop_loss,
        take_profit=entry_signal.take_profit or 0.0,
        position_size=position_size,
        risk_amount=risk_amount,
        risk_pct=risk_pct,
        reward_risk_ratio=round(reward_risk, 4),
        invalid=False,
        reason="Trade validated successfully"
    )


def full_risk_analysis(candles: List[Candle],
                        risk_params: Optional[RiskParameters] = None) -> Optional[RiskResult]:
    """
    Perform complete risk analysis: detect entry, calculate SL/TP, size position.

    Args:
        candles: List of candles (oldest to newest)
        risk_params: Risk parameters (uses defaults if None)

    Returns:
        RiskResult or None if no valid entry found
    """
    if risk_params is None:
        risk_params = RiskParameters()

    if len(candles) < 3:
        return None

    # Detect entry signal
    entry = detect_entry(candles)
    if entry is None:
        return None

    # Calculate SL/TP if not already set
    if entry.stop_loss is None and len(candles) >= 2:
        second_to_last = candles[-2]
        entry.stop_loss, entry.take_profit = calculate_stop_loss_take_profit(
            entry_price=entry.price,
            signal_type=entry.signal_type,
            second_to_last_low=second_to_last.low,
            second_to_last_high=second_to_last.high,
            atr=calculate_atr(candles),
        )

    # Validate and size
    return validate_trade(entry, risk_params)