"""GM! Occupancy Trader policy, as pure functions (no dependencies).

Reference implementation of ``SOUL.md`` in this directory, for Hyperliquid
perpetuals. Constants in ``PROVISIONAL`` are carried from a reference book on
another venue and are not yet validated on Hyperliquid; paper confirmation
is the validation. Closes are Hyperliquid 1h closes, oldest first; ``signed``
is one GM! ``signed_conviction`` per hour, oldest first, excluding the
current hour unless stated.
"""

from __future__ import annotations

import math
import random
from typing import Iterable, Literal, Sequence

Side = Literal["long", "short"]

HOLD_HOURS = 4
WINDOW = 720
BASIS_HORIZON = "4H"
OPPOSITE_SIGN = "opposite_sign_oscillator"
MIN_HISTORY = 168
RANK_Q = 0.125
PI_MULT_PROFIT = 1.0
MIN_TAKEN = 20
# observation_ts is the open of the last closed 1h bar, so a healthy cell is 1-2 h old.
STALE_MS = 3 * 3600 * 1000

K = 6.0
L_MULT = 12.0
MOVE_AGAINST_MULT = 12.0
LM_C = 2.5
CLAMP_LO = 0.3
CLAMP_HI = 3.0

MASS_MIN = WINDOW // 2
LATTICE = (3, 5, 8)
BASE_LOTS = 5

# Kill in R: each lot's net P&L over its own initial floor risk, 12·τ_b of
# notional, so the kill scales with each instrument's loss mass.
KILL_R = 10.0
LIVE_KILL_DD_MULT = 2.0
LOOKS = (60, 120, 180)
CONFIRM_P = 0.80
REVERT_P = 0.10
N_BOOT = 1000

PROVISIONAL = (
    "RANK_Q", "WINDOW", "MIN_HISTORY", "PI_MULT_PROFIT", "HOLD_HOURS", "K", "L_MULT",
    "MOVE_AGAINST_MULT", "LM_C", "CLAMP_LO", "CLAMP_HI", "LATTICE", "KILL_R",
)

# Time value: a risk-free rate on the full notional over the hold. Declared by
# the trader with LOT_USD; replace the placeholder with the USDC rate you use.
RF_APR = 0.045
HOURS_PER_YEAR = 8760.0

# Hyperliquid perps, 14-day volume tiers (userFees.feeSchedule). Use your own
# userAddRate / userCrossRate: staking and referral discounts lower them.
FEES = {
    "hyperliquid_base": {"maker": 0.00015, "taker": 0.00045},
    "hyperliquid_vip1": {"maker": 0.00012, "taker": 0.00040},
    "hyperliquid_vip2": {"maker": 0.00008, "taker": 0.00035},
}


def tau(*, maker: float, taker: float, builder: float = 0.0) -> float:
    """Round-trip hurdle of one lot, as a return: post-only (maker) entry, taker exit.

    ``builder`` is the builder fee per fill, charged on both legs when the
    agent trades through a builder code.
    """
    if maker < -0.001 or taker < 0 or builder < 0:
        raise ValueError("fee rates out of range")
    return maker + taker + 2.0 * builder


def rho(rf_apr: float = RF_APR) -> float:
    """Time value per hour, as a return on notional."""
    if not 0.0 <= rf_apr < 1.0:
        raise ValueError("rf_apr out of range")
    return rf_apr / HOURS_PER_YEAR


def tau_hold(tau_rt: float, rho_h: float, hours: float = HOLD_HOURS) -> float:
    """τ_H = τ + ρ·H: fees plus the time value of the notional over the hold."""
    return tau_rt + rho_h * max(hours, 0.0)


def funding_cost(side: Side, hourly_rate: float, hours: float = HOLD_HOURS) -> float:
    """Funding the side pays over the hold at the current hourly rate (negative = received)."""
    return (1.0 if side == "long" else -1.0) * hourly_rate * hours


def _quantile(values: Sequence[float], q: float) -> float:
    xs = sorted(values)
    pos = (len(xs) - 1) * q
    lo, hi = math.floor(pos), math.ceil(pos)
    return xs[lo] + (xs[hi] - xs[lo]) * (pos - lo)


def band(signed_history: Sequence[float]) -> float | None:
    """87.5th percentile of |signed| over the trailing 720 hours; None under 168."""
    hist = [abs(float(x)) for x in signed_history[-WINDOW:] if math.isfinite(float(x))]
    if len(hist) < MIN_HISTORY:
        return None
    return _quantile(hist, 1.0 - RANK_Q)


def admit(signed: float, band_level: float | None) -> Side | None:
    if band_level is None or not math.isfinite(signed):
        return None
    if signed >= band_level:
        return "long"
    if signed <= -band_level:
        return "short"
    return None


def book_support(metadata: dict, horizon: str = BASIS_HORIZON, series: str = "GM!/HYPERLIQUID") -> str | None:
    """The ``support`` a production book declares in ``/v1/feed/metadata`` ``venue_series``."""
    for row in metadata.get("venue_series") or []:
        if row.get("series") != series:
            continue
        for book in row.get("books") or []:
            if book.get("horizon") == horizon and book.get("status") == "production":
                return book.get("support")
    return None


def hurdle_net(side: Side, cell: dict, *, opposite_sign: bool = False) -> bool:
    """The cell's omega, netted at its own ``omega_hurdle``, points to this side.

    On an opposite-sign book the conviction fades the move that the cell's omega
    reads, so omega points the other way by construction and is not same-side
    support. The cell part of the gate does not apply there.
    """
    if opposite_sign:
        return True
    if cell.get("value_mode") == "omega":
        direction = float(cell.get("signed_conviction") or 0.0)
    else:
        direction = float(cell.get("omega_direction") or 0.0)
    return direction > 0 if side == "long" else direction < 0


def omega_at(returns: Sequence[float], hurdle: float) -> float | None:
    """Ω at a hurdle: Σ(x − h)+ / Σ(h − x)+. Ω > 1 exactly when mean(x) > h."""
    if not returns:
        return None
    gain = sum(max(x - hurdle, 0.0) for x in returns)
    loss = sum(max(hurdle - x, 0.0) for x in returns)
    if loss == 0.0:
        return math.inf if gain > 0 else None
    return gain / loss


def profit_bar(tau_rt: float, rho_h: float = 0.0, hours: float = HOLD_HOURS) -> float:
    """h* = τ_H + π with π = τ_H, so h* = 2·τ_H."""
    return tau_hold(tau_rt, rho_h, hours) * (1.0 + PI_MULT_PROFIT)


def profit_net(
    expected: float | None,
    tau_rt: float,
    side: Side,
    hourly_funding: float,
    rho_h: float = 0.0,
) -> bool:
    """E[R | taken] − funding over the hold clears h* = 2·τ_H."""
    if expected is None:
        return False
    return expected - funding_cost(side, hourly_funding) > profit_bar(tau_rt, rho_h)


def expected_taken_return(
    closes: Sequence[float], signed: Sequence[float], side: Side
) -> tuple[float | None, int]:
    """Mean of ``taken_returns``; None under 20 taken hours."""
    taken = taken_returns(closes, signed, side)
    if len(taken) < MIN_TAKEN:
        return None, len(taken)
    return sum(taken) / len(taken), len(taken)


def taken_returns(closes: Sequence[float], signed: Sequence[float], side: Side) -> list[float]:
    """Side-signed log returns over the hold, at trailing hours where the band admitted ``side``.

    ``closes[i]`` and ``signed[i]`` are the same hour. Each hour's band uses only
    the hours before it. Only hours whose forward hold has closed count.
    """
    n = min(len(closes), len(signed))
    closes, signed = list(closes[-n:]), list(signed[-n:])
    start = max(0, n - WINDOW)
    sign = 1.0 if side == "long" else -1.0
    taken: list[float] = []
    for i in range(start, n - HOLD_HOURS):
        lvl = band(signed[:i])
        if admit(signed[i], lvl) != side:
            continue
        c0, c1 = closes[i], closes[i + HOLD_HOURS]
        if c0 > 0 and c1 > 0:
            taken.append(sign * math.log(c1 / c0))
    return taken


def _r_hold(closes: Sequence[float]) -> list[float]:
    return [
        math.log(closes[i] / closes[i - HOLD_HOURS])
        for i in range(HOLD_HOURS, len(closes))
        if closes[i] > 0 and closes[i - HOLD_HOURS] > 0
    ]


def _mass_series(closes: Sequence[float], side: Side) -> list[float | None]:
    rh = [None] * HOLD_HOURS + [
        math.log(closes[i] / closes[i - HOLD_HOURS]) for i in range(HOLD_HOURS, len(closes))
    ]
    out: list[float | None] = []
    for i in range(len(closes)):
        win = [x for x in rh[max(0, i - WINDOW + 1) : i + 1] if x is not None]
        if len(win) < MASS_MIN:
            out.append(None)
            continue
        part = [max(x, 0.0) if side == "long" else max(-x, 0.0) for x in win]
        out.append(sum(part) / len(part))
    return out


def gain_mass_rank(closes: Sequence[float], side: Side) -> float | None:
    """Rank of the latest trade-side gain mass among the prior 720 hours."""
    mass = [m for m in _mass_series(closes, side) if m is not None]
    if len(mass) < MIN_HISTORY + 1:
        return None
    prior = mass[-(WINDOW + 1) : -1]
    if len(prior) < MIN_HISTORY:
        return None
    now = mass[-1]
    return sum(1 for m in prior if m < now) / len(prior)


def lots(rank: float | None) -> int:
    if rank is None or not math.isfinite(rank):
        return BASE_LOTS
    return LATTICE[0] if rank < 1 / 3 else (LATTICE[1] if rank < 2 / 3 else LATTICE[2])


def loss_mass(closes: Sequence[float]) -> float | None:
    rh = _r_hold(closes)[-WINDOW:]
    if len(rh) < WINDOW // 2:
        return None
    g0 = sum(max(x, 0.0) for x in rh) / len(rh)
    l0 = sum(max(-x, 0.0) for x in rh) / len(rh)
    lm = (g0 + l0) / 2.0
    return lm if math.isfinite(lm) and lm > 0 else None


def tau_b(closes: Sequence[float], tau_rt: float) -> float:
    """Floor unit stamped at entry: clamp(2.5·LM/12, 0.3τ, 3τ); τ when LM is missing."""
    lm = loss_mass(closes)
    if lm is None:
        return tau_rt
    return min(max(LM_C * lm / L_MULT, CLAMP_LO * tau_rt), CLAMP_HI * tau_rt)


def floor_return(age_hours: float, tau_floor: float) -> float:
    """floor(t) = κ·min(t, H) − L·τ_b with κ = k·τ_b/H."""
    kappa = K * tau_floor / HOLD_HOURS
    return kappa * min(max(age_hours, 0.0), HOLD_HOURS) - L_MULT * tau_floor


def exit_reason(
    side: Side, entry: float, mid: float, age_hours: float, tau_floor: float
) -> str | None:
    """First exit that holds now: 'floor', 'move_against', 'time', or None."""
    f = floor_return(age_hours, tau_floor)
    pi = MOVE_AGAINST_MULT * tau_floor
    if side == "long":
        if mid <= entry * (1.0 + f):
            return "floor"
        if mid <= entry * (1.0 - pi):
            return "move_against"
    else:
        if mid >= entry * (1.0 - f):
            return "floor"
        if mid >= entry * (1.0 + pi):
            return "move_against"
    if age_hours >= HOLD_HOURS:
        return "time"
    return None


def stay_out(side: Side, mid: float, failed_side: Side | None, failed_entry: float | None) -> bool:
    """Same side stays out after a stop until mid trades back through that entry."""
    if failed_side != side or not failed_entry:
        return False
    return mid <= failed_entry if side == "long" else mid >= failed_entry


def reentry_hold(mid: float, flatten_px: float | None, tau_rt: float) -> bool:
    if not flatten_px or flatten_px <= 0:
        return False
    return abs(mid - flatten_px) / flatten_px < tau_rt


def cell_usable(cell: dict, *, basis_address: str, venue: str = "hyperliquid") -> str | None:
    """Reason to refuse a cell for new entries, or None."""
    if cell.get("source_venue") != venue:
        return "venue"
    if cell.get("index_address") != basis_address:
        return "address"
    if cell.get("methodology_status") != "production":
        return "not_production"
    stale = cell.get("staleness_ms")
    if stale is None or float(stale) > STALE_MS:
        return "stale"
    return None


def lot_r(net_usd: float, notional: float, tau_floor: float) -> float:
    """A closed lot's net P&L in units of its initial floor risk, L·τ_b of notional."""
    if not notional > 0 or not tau_floor > 0:
        raise ValueError("notional and tau_floor must be positive")
    return net_usd / (L_MULT * tau_floor * notional)


def max_drawdown(closed_r: Iterable[float]) -> float:
    cum = peak = dd = 0.0
    for x in closed_r:
        cum += x
        peak = max(peak, cum)
        dd = min(dd, cum - peak)
    return dd


def live_kill_r(paper_r: Iterable[float]) -> float:
    """Kill for the live record: twice the confirmed paper record's drawdown in R, never under KILL_R."""
    return max(KILL_R, LIVE_KILL_DD_MULT * -max_drawdown(paper_r))


def history_usable(
    history: dict, *, basis_address: str, venue: str = "hyperliquid", live_params_hash: str | None = None
) -> str | None:
    """Reason to refuse a served conviction history for the band, or None.

    The history is only the same series as the live cell when it carries the
    same venue, address, methodology and ``params_hash``.
    """
    if history.get("source_venue") != venue:
        return "venue"
    if history.get("index_address") != basis_address:
        return "address"
    if history.get("methodology_status") != "production":
        return "not_production"
    if live_params_hash is not None and history.get("params_hash") != live_params_hash:
        return "params_hash"
    if not history.get("records"):
        return "empty"
    return None


def seed_signed(records: Sequence[dict], bar_open_ms: Sequence[int]) -> list[float]:
    """Signed conviction for each Hyperliquid bar, by open time in ms, oldest first.

    A record with timestamp ``t`` is the cell at the bar that opens at ``t``,
    so it pairs with that bar's close. NaN where the cell has no record.
    """
    by_ts = {
        int(r["ts"]): float(r["signed_conviction"])
        for r in records
        if r.get("signed_conviction") is not None
    }
    return [by_ts.get(int(t), math.nan) for t in bar_open_ms]


def killed(closed_r: Iterable[float], kill_r: float = KILL_R) -> bool:
    return sum(closed_r) < -kill_r


def day_clustered_p(net_bps: Sequence[float], days: Sequence[str], *, seed: int = 20260926) -> float | None:
    if len(net_bps) < 2:
        return None
    by: dict[str, list[float]] = {}
    for x, d in zip(net_bps, days):
        by.setdefault(d, []).append(float(x))
    keys = sorted(by)
    rng = random.Random(seed)
    wins = 0
    for _ in range(N_BOOT):
        pick = [x for k in (rng.choice(keys) for _ in keys) for x in by[k]]
        wins += (sum(pick) / len(pick)) > 0.0
    return wins / N_BOOT


def confirmation(net_bps: Sequence[float], days: Sequence[str]) -> str:
    """Confirmation law over closed lots in close order: 'running', 'confirmed' or 'revert'."""
    for look in LOOKS:
        if len(net_bps) < look:
            return "running"
        sub = list(net_bps[:look])
        p = day_clustered_p(sub, list(days[:look])) or 0.0
        if sum(sub) / look > 0 and p >= CONFIRM_P:
            return "confirmed"
        if p <= REVERT_P or look == LOOKS[-1]:
            return "revert"
    return "running"
