"""Pure causal sizing math for the research-only D48 Turtle pyramid."""

from dataclasses import dataclass


@dataclass(frozen=True)
class PyramidAddition:
    """One eligible position increase and the frozen trigger that caused it."""

    addition_number: int
    trigger_rate: float
    stake: float


def planned_risk_fraction(
    *,
    initial_risk_fraction: float,
    add_stake_fraction: float,
    max_additions: int,
) -> float:
    """Upper-bound planned risk when every unit shares the initial loss fraction."""
    if initial_risk_fraction <= 0 or add_stake_fraction <= 0 or max_additions < 0:
        raise ValueError("risk inputs must be positive")
    return initial_risk_fraction * (1 + add_stake_fraction * max_additions)


def pyramid_addition(
    *,
    initial_stake: float,
    initial_rate: float,
    initial_atr: float,
    current_rate: float,
    successful_entries: int,
    is_short: bool,
    add_stake_fraction: float,
    spacing_atr: float,
    max_additions: int,
    min_stake: float | None,
    max_stake: float,
) -> PyramidAddition | None:
    """Return the next add using only state frozen at the initial fill.

    ``successful_entries`` includes the initial entry.  A value of one therefore
    evaluates the first add at 0.5 ATR, while two evaluates the second at 1 ATR.
    Only one add is emitted per call; the strategy waits for it to fill before the
    successful-entry counter permits the next level.
    """
    values = (
        initial_stake,
        initial_rate,
        initial_atr,
        current_rate,
        add_stake_fraction,
        spacing_atr,
        max_stake,
    )
    if any(value <= 0 for value in values):
        raise ValueError("pyramid inputs must be positive")
    if successful_entries < 1 or max_additions < 0:
        raise ValueError("entry counts must be valid")

    addition_number = successful_entries
    if addition_number > max_additions:
        return None

    offset = spacing_atr * initial_atr * addition_number
    trigger_rate = initial_rate - offset if is_short else initial_rate + offset
    reached = current_rate <= trigger_rate if is_short else current_rate >= trigger_rate
    if not reached:
        return None

    stake = min(initial_stake * add_stake_fraction, max_stake)
    if min_stake is not None and stake < min_stake:
        return None
    return PyramidAddition(
        addition_number=addition_number,
        trigger_rate=trigger_rate,
        stake=stake,
    )
