from __future__ import annotations

from decimal import Decimal, InvalidOperation
from typing import Any


def _decimal(value: object) -> Decimal:
    try:
        return Decimal(str(value))
    except (InvalidOperation, TypeError, ValueError):
        return Decimal("0")


def _text(value: Decimal, places: int = 8) -> str:
    return format(value, f".{places}f").rstrip("0").rstrip(".") or "0"


def _sell_fraction(
    *,
    raw_price: Decimal,
    fraction: Decimal,
    quantity: Decimal,
    effective_entry: Decimal,
    fee_rate: Decimal,
    slippage_rate: Decimal,
) -> tuple[Decimal, Decimal]:
    effective_exit = raw_price * (Decimal("1") - slippage_rate)
    proceeds = effective_exit * quantity * fraction
    gross = (effective_exit - effective_entry) * quantity * fraction
    fee = proceeds * fee_rate
    return gross, fee


def backtest_signal_cases(
    cases: list[dict[str, Any]],
    *,
    notional_usdt: Decimal = Decimal("100"),
    fee_bps: Decimal = Decimal("10"),
    slippage_bps: Decimal = Decimal("5"),
    move_stop_to_break_even: bool = True,
) -> dict[str, Any]:
    """Evaluate historical signal cases without fetching data or reading credentials."""

    safe_notional = max(Decimal("1"), notional_usdt)
    fee_rate = max(Decimal("0"), fee_bps) / Decimal("10000")
    slippage_rate = max(Decimal("0"), slippage_bps) / Decimal("10000")
    results: list[dict[str, Any]] = []

    for case in cases:
        entry = _decimal(case.get("entry_price"))
        stop = _decimal(case.get("stop_loss"))
        target_1 = _decimal(case.get("target_1"))
        target_2 = _decimal(case.get("target_2"))
        bars = case.get("bars") if isinstance(case.get("bars"), list) else []
        if not (Decimal("0") < stop < entry < target_1 < target_2) or not bars:
            results.append(
                {
                    "symbol": str(case.get("symbol", "")),
                    "status": "INVALID_CASE",
                    "error": "Expected stop < entry < target_1 < target_2 and at least one bar.",
                }
            )
            continue

        effective_entry = entry * (Decimal("1") + slippage_rate)
        quantity = safe_notional / effective_entry
        fees = safe_notional * fee_rate
        gross = Decimal("0")
        remaining = Decimal("1")
        current_stop = stop
        target_1_hit = False
        exit_reason = "END_OF_DATA"
        exit_price = _decimal(bars[-1].get("close"))
        exit_time = str(bars[-1].get("timestamp", ""))
        mfe = Decimal("0")
        mae = Decimal("0")

        for bar in bars:
            open_price = _decimal(bar.get("open"))
            high = _decimal(bar.get("high"))
            low = _decimal(bar.get("low"))
            timestamp = str(bar.get("timestamp", ""))
            if high <= 0 or low <= 0 or high < low:
                continue
            mfe = max(mfe, ((high - entry) / entry) * Decimal("100"))
            mae = min(mae, ((low - entry) / entry) * Decimal("100"))

            # Conservative rule: if stop and target are both inside one bar, stop wins.
            if low <= current_stop:
                exit_reason = "STOP_AFTER_TARGET_1" if target_1_hit else "STOP_LOSS"
                exit_price = min(current_stop, open_price) if open_price > 0 else current_stop
                exit_time = timestamp
                break

            if not target_1_hit and high >= target_1:
                part_gross, part_fee = _sell_fraction(
                    raw_price=target_1,
                    fraction=Decimal("0.5"),
                    quantity=quantity,
                    effective_entry=effective_entry,
                    fee_rate=fee_rate,
                    slippage_rate=slippage_rate,
                )
                gross += part_gross
                fees += part_fee
                remaining = Decimal("0.5")
                target_1_hit = True
                if move_stop_to_break_even:
                    current_stop = entry

            if target_1_hit and high >= target_2:
                exit_reason = "TARGET_2"
                exit_price = target_2
                exit_time = timestamp
                break

        close_gross, close_fee = _sell_fraction(
            raw_price=exit_price,
            fraction=remaining,
            quantity=quantity,
            effective_entry=effective_entry,
            fee_rate=fee_rate,
            slippage_rate=slippage_rate,
        )
        gross += close_gross
        fees += close_fee
        net = gross - fees
        return_percent = (net / safe_notional) * Decimal("100")
        results.append(
            {
                "symbol": str(case.get("symbol", "")),
                "signal_time": str(case.get("signal_time", "")),
                "status": "CLOSED",
                "exit_reason": exit_reason,
                "exit_time": exit_time,
                "entry_price": _text(entry),
                "exit_price": _text(exit_price),
                "target_1_hit": target_1_hit,
                "gross_pnl_usdt": _text(gross, 6),
                "fees_usdt": _text(fees, 6),
                "net_pnl_usdt": _text(net, 6),
                "return_percent": _text(return_percent, 4),
                "max_favorable_excursion_percent": _text(mfe, 4),
                "max_adverse_excursion_percent": _text(mae, 4),
            }
        )

    valid = [item for item in results if item.get("status") == "CLOSED"]
    wins = [item for item in valid if _decimal(item.get("net_pnl_usdt")) > 0]
    losses = [item for item in valid if _decimal(item.get("net_pnl_usdt")) < 0]
    total_net = sum((_decimal(item.get("net_pnl_usdt")) for item in valid), Decimal("0"))
    gross_profit = sum((_decimal(item.get("net_pnl_usdt")) for item in wins), Decimal("0"))
    gross_loss = abs(sum((_decimal(item.get("net_pnl_usdt")) for item in losses), Decimal("0")))
    average_return = (
        sum((_decimal(item.get("return_percent")) for item in valid), Decimal("0"))
        / Decimal(len(valid))
        if valid else Decimal("0")
    )
    equity = Decimal("0")
    peak = Decimal("0")
    max_drawdown = Decimal("0")
    for item in valid:
        equity += _decimal(item.get("net_pnl_usdt"))
        peak = max(peak, equity)
        max_drawdown = max(max_drawdown, peak - equity)

    return {
        "summary": {
            "case_count": len(cases),
            "valid_trade_count": len(valid),
            "invalid_case_count": len(cases) - len(valid),
            "wins": len(wins),
            "losses": len(losses),
            "win_rate_percent": format((Decimal(len(wins)) / Decimal(len(valid))) * Decimal("100"), ".2f") if valid else "0.00",
            "total_net_pnl_usdt": _text(total_net, 4),
            "average_return_percent": _text(average_return, 4),
            "profit_factor": _text(gross_profit / gross_loss, 4) if gross_loss > 0 else None,
            "max_drawdown_usdt": _text(max_drawdown, 4),
        },
        "assumptions": {
            "notional_usdt_per_trade": _text(safe_notional, 2),
            "fee_bps_each_side": _text(fee_bps, 2),
            "slippage_bps_each_side": _text(slippage_bps, 2),
            "target_1_fraction": "0.50",
            "move_stop_to_break_even_after_target_1": move_stop_to_break_even,
            "intrabar_conflict_rule": "STOP_FIRST_CONSERVATIVE",
            "data_source": "USER_SUPPLIED_HISTORICAL_CASES",
        },
        "trades": results,
    }
