"""Shared sizing helpers for overlay / IV research backtests (contracts vs broker-risk budget)."""

from __future__ import annotations

import math


def resolve_overlay_contracts(
    *,
    contracts: int | None,
    target_broker_risk_usd: float | None,
    per_contract_broker_risk: float,
) -> tuple[int, float, float]:
    """
    Return ``(contracts, per_contract_risk, total_broker_risk_usd)``.

    Precedence: explicit ``contracts`` (>= 1); else ``target_broker_risk_usd`` uses
    ``max(1, floor(target / per_contract_risk))``; else ``contracts = 1``.

    ``per_contract_broker_risk`` must be the model's max loss or debit for **one** contract
    (or one 1x1 spread where each leg has the same qty).
    """
    per = max(abs(float(per_contract_broker_risk)), 1e-9)
    if contracts is not None and int(contracts) >= 1:
        n = int(contracts)
    elif target_broker_risk_usd is not None and math.isfinite(float(target_broker_risk_usd)):
        t = float(target_broker_risk_usd)
        if t > 0:
            n = max(1, int(t // per))
        else:
            n = 1
    else:
        n = 1
    return n, per, per * float(n)


def nav_pct_target_and_applied(
    *,
    initial_portfolio_capital: float,
    realized_pnl_to_date: float,
    broker_risk_pct_of_portfolio: float | None,
) -> tuple[float, float | None, float]:
    """
    NAV before a new entry (for logging) and optional broker-risk budget from a fraction of NAV.

    Returns ``(nav_usd, target_broker_risk_usd_or_none, risk_pct_applied)``.
    When ``broker_risk_pct_of_portfolio`` is set, ``target_broker_risk_usd`` is ``max(1, nav * pct)``;
    otherwise the second value is ``None`` and sizing should use fixed ``contracts`` / ``target_broker_risk_usd``.
    """
    nav = max(float(initial_portfolio_capital) + float(realized_pnl_to_date), 1.0)
    if broker_risk_pct_of_portfolio is None:
        return nav, None, 0.0
    p = float(broker_risk_pct_of_portfolio)
    if not (0.0 < p <= 1.0):
        raise ValueError("broker_risk_pct_of_portfolio must be in (0, 1]")
    return nav, max(nav * p, 1.0), p
