"""
risk_manager.py

Two responsibilities:
1. Position sizing (how many shares, given account equity and risk %)
2. Stop calculation — initial stop and the dynamic trailing stop that
   only ever moves up (never back down), per the project's explicit
   spec and worked examples.

Supports four configurable stop methods: fixed_cents, percentage, atr,
volatility_adjusted. The method is chosen entirely through config.json.
"""

import math
from config_loader import get_config
from logger_setup import get_logger
from indicators import atr as calc_atr

log = get_logger("risk_manager")


def _round_to_cent(x: float) -> float:
    return round(x + 1e-9, 2)


def _min_stop_distance(cfg: dict, price: float) -> float:
    """
    [FEATURE 2026-09-03] min_stop_distance_cents alone is a flat dollar
    floor -- fine at one price level, negligible at another. Several
    real 2026-09-03 entries (ACHR, CRML, HAFN, JOBY) landed exactly on
    or barely above that $0.03 floor because the ATR-derived distance
    was tiny at entry time, leaving only $0.03-$0.05 (0.4-0.7% of price)
    of actual room -- not enough to survive ordinary 1-minute noise on
    this project's $5-15 cohort. Takes the larger of the flat-cents
    floor and min_stop_distance_pct% of the current price, so the
    percentage floor is what actually governs at this project's normal
    price range while min_stop_distance_cents remains a floor-of-the-
    floor for any very low-priced name where 1% would be a fraction of
    a cent.
    """
    return max(cfg["min_stop_distance_cents"], price * (cfg.get("min_stop_distance_pct", 0) / 100.0))


def compute_stop_distance(current_price: float, bars: list = None) -> float:
    """
    Returns the trailing distance in dollars, based on the configured
    method. `bars` (recent 1-min bars) is required for atr/volatility
    methods; ignored otherwise.
    """
    cfg = get_config()["stop"]
    method = cfg["method"]

    if method == "fixed_cents":
        distance = cfg["trailing_distance"]

    elif method == "percentage":
        distance = current_price * (cfg["percentage_trailing_pct"] / 100.0)

    elif method == "atr":
        if not bars:
            log.warning("atr stop method requested but no bars supplied; falling back to fixed_cents")
            distance = cfg["trailing_distance"]
        else:
            a = calc_atr(bars, period=cfg["atr_period"])
            distance = a * cfg["atr_multiplier_trailing"]

    elif method == "volatility_adjusted":
        # Blend of ATR and percentage, whichever is tighter, to avoid
        # absurdly wide stops on spiky low-liquidity names while still
        # respecting the stock's real volatility.
        pct_distance = current_price * (cfg["percentage_trailing_pct"] / 100.0)
        if bars:
            a = calc_atr(bars, period=cfg["atr_period"])
            atr_distance = a * cfg["atr_multiplier_trailing"]
            distance = min(pct_distance, atr_distance) if atr_distance > 0 else pct_distance
        else:
            distance = pct_distance

    else:
        log.warning(f"Unknown stop method '{method}', defaulting to fixed_cents")
        distance = cfg["trailing_distance"]

    # Enforce sane bounds regardless of method
    min_dist = _min_stop_distance(cfg, current_price)
    max_dist = current_price * (cfg["max_stop_distance_pct"] / 100.0)
    distance = max(min_dist, min(distance, max_dist))
    return _round_to_cent(distance)


def compute_initial_stop(entry_price: float, bars: list = None) -> float:
    cfg = get_config()["stop"]
    method = cfg["method"]

    if method == "fixed_cents":
        distance = cfg["initial_distance"]
    elif method == "percentage":
        distance = entry_price * (cfg["percentage_initial_pct"] / 100.0)
    elif method == "atr":
        a = calc_atr(bars, period=cfg["atr_period"]) if bars else 0
        distance = a * cfg["atr_multiplier_initial"] if a else cfg["initial_distance"]
    elif method == "volatility_adjusted":
        pct_distance = entry_price * (cfg["percentage_initial_pct"] / 100.0)
        a = calc_atr(bars, period=cfg["atr_period"]) if bars else 0
        atr_distance = a * cfg["atr_multiplier_initial"] if a else pct_distance
        distance = min(pct_distance, atr_distance)
    else:
        distance = cfg["initial_distance"]

    min_dist = _min_stop_distance(cfg, entry_price)
    max_dist = entry_price * (cfg["max_stop_distance_pct"] / 100.0)
    distance = max(min_dist, min(distance, max_dist))
    stop = _round_to_cent(entry_price - distance)
    return stop


def update_trailing_stop(highest_price: float, current_stop: float, bars: list = None) -> float:
    """
    Core trailing-stop rule from the spec:
        new_stop = highest_price - trailing_distance
        only update if new_stop > current_stop  (stop NEVER moves down)
    """
    distance = compute_stop_distance(highest_price, bars)
    candidate_stop = _round_to_cent(highest_price - distance)
    if candidate_stop > current_stop:
        return candidate_stop
    return current_stop


def compute_position_size(account_equity: float, entry_price: float, initial_stop: float) -> int:
    """
    Risk-based sizing: risk `account_risk_pct_per_trade`% of equity on
    the distance between entry and initial stop, capped by a maximum
    notional percentage of equity to avoid oversized positions on
    very tight stops.
    """
    cfg = get_config()["risk"]
    trading_cfg = get_config()["trading"]

    risk_dollars = account_equity * (cfg["account_risk_pct_per_trade"] / 100.0)
    per_share_risk = max(entry_price - initial_stop, 0.01)
    shares_by_risk = math.floor(risk_dollars / per_share_risk)

    max_notional = account_equity * (cfg["max_position_notional_pct_of_equity"] / 100.0)
    shares_by_notional = math.floor(max_notional / entry_price) if entry_price > 0 else 0

    shares = min(shares_by_risk, shares_by_notional)
    shares = max(shares, cfg["min_shares"] if shares > 0 else 0)
    return int(shares)


def check_daily_loss_limit(realized_pl_today: float, account_equity: float) -> bool:
    """Returns True if the daily loss limit has been breached (halt new entries)."""
    cfg = get_config()["risk"]
    max_loss = -account_equity * (cfg["max_daily_loss_pct"] / 100.0)
    return realized_pl_today <= max_loss


def check_consecutive_loss_pause(recent_trade_results: list) -> bool:
    """
    recent_trade_results: chronological list of booleans, True = win.
    Returns True if bot should pause new entries due to a losing streak.
    """
    cfg = get_config()["risk"]
    limit = cfg["max_consecutive_losses_pause"]
    if len(recent_trade_results) < limit:
        return False
    tail = recent_trade_results[-limit:]
    return all(result is False for result in tail)
