"""sim-4firms-ev.py の検証用。numba を使わず、口座のルールを素直に書き直した別実装で同じ条件を回し、
(1) 1年間の期待値が誤差の範囲で一致するか
(2) ルール違反（1日2.5%超の損失・1回3%超のリスク・失格ラインを割った口座の生存）が起きていないか
(3) 1トレードの期待値が式 (2p − 1.05)R と一致するか
を確かめる。乱数は numba 版と別。

使い方: python scripts/verify-4firms-ev.py [回数=3000] [numba版の結果JSON]
"""
import json
import math
import random
import sys

DAYS, TRADES_PER_DAY, CYCLE_DAYS = 250, 3, 10
DAILY_STOP = 0.025
COST = 0.05

# 各社のルール（sim-4firms-ev.py と同じ値を、名前つきで書き直したもの）
PLANS = {
    "FTMO 2-Step": dict(price=619, targets=[0.10, 0.05], max_loss=0.10, daily=0.05, daily_from_initial=True,
                        min_days=("trading", 4), payout_share=0.8, refund_after_payouts=1, full_payout=False),
    "FTMO 1-Step": dict(price=572, targets=[0.10], max_loss=0.10, daily=0.03, daily_from_initial=True,
                        trailing_eod=True, reset_on_payout=True, best_day=0.5, payout_share=0.9, full_payout=True),
    "The5ers Classic": dict(price=455, targets=[0.08, 0.05], max_loss=0.08, daily=0.04, min_days=("profit", 3),
                            payout_share=0.8, credits=[0.1, 0.2], funded_cash=0.7, payout_cap=0.04, payout_min=0.005,
                            payout_profit_days=3),
    "The5ers New": dict(price=405, targets=[0.10, 0.05], max_loss=0.08, daily=0.04, min_days=("profit", 3),
                        payout_share=0.8, credits=[0.1, 0.2], funded_cash=0.7, payout_cap=0.04, payout_min=0.005,
                        payout_profit_days=3),
    "Fintokei": dict(price=549, targets=[0.08, 0.06], max_loss=0.10, daily=0.05, min_days=("trading", 3),
                     payout_share=0.8, refund_after_payouts=2, refund_min_funded_days=20),
    "Hantec Enhanced": dict(price=599, targets=[0.10, 0.05], max_loss=0.10, daily=0.05, min_days=("profit", 3),
                            payout_share=0.8, lock_on_payout=True, payout_profit_days=3),
    "Hantec Endurance": dict(price=299, targets=[0.06, 0.06, 0.06], max_loss=0.08, daily=0.04, min_days=("trading", 3),
                             payout_share=0.8),
}


class Violation(Exception):
    pass


def one_year(plan, p, r, keep, rng, stats):
    """1年ぶんを回して (報酬の合計, 払った参加費の合計) を返す（口座比）。"""
    price = plan["price"] / 100000
    paid = 0.0
    spent = price
    credit = 0.0

    def new_account():
        return dict(stage=0, bal=1.0, line=1.0 - plan["max_loss"], eod_high=1.0, days=0, profit_days=0,
                    best=0.0, pos_sum=0.0, funded=False, since=0, payouts=0, funded_days=0, refunded=False, locked=False)

    acc = new_account()
    for _ in range(DAYS):
        start = acc["bal"]
        firm_floor = start - plan["daily"] if plan.get("daily_from_initial") else start * (1 - plan["daily"])
        my_floor = max(firm_floor, start - DAILY_STOP)
        failed = False
        reached = False
        for _ in range(TRADES_PER_DAY):
            if reached:
                break  # その日のトレードで目標に届いたら、その日はやめる（最低日数の待ちの日も1回は建てる）
            if acc["bal"] - r * (1 + COST) < my_floor - 1e-12:
                break  # 次に負けたら自分の上限（2.5%）か会社の日次を超えるならやめる
            if r > 0.03:
                raise Violation("1回のリスクが3%超")
            win = rng.random() < p
            pnl = r * (1 - COST) if win else -r * (1 + COST)
            stats["trades"] += 1
            stats["pnl_R"] += pnl / r
            acc["bal"] += pnl
            if acc["bal"] < acc["line"] - 1e-12 or acc["bal"] < firm_floor - 1e-12:
                failed = True
                break
            if not acc["funded"] and acc["bal"] >= 1 + plan["targets"][acc["stage"]] - 1e-12:
                reached = True
        day_loss = start - acc["bal"]
        if not failed and day_loss > DAILY_STOP + 1e-9:
            raise Violation(f"1日の損失が2.5%超: {day_loss:.4f}")
        stats["max_day_loss"] = max(stats["max_day_loss"], day_loss)
        if failed:
            stats["fails"] += 1
            use = min(credit, price)
            spent += price - use
            credit -= use
            acc = new_account()
            continue
        dp = acc["bal"] - start
        acc["days"] += 1
        if dp >= 0.005 - 1e-12:
            acc["profit_days"] += 1
        acc["best"] = max(acc["best"], dp)
        if dp > 0:
            acc["pos_sum"] += dp
        if plan.get("trailing_eod"):
            acc["eod_high"] = max(acc["eod_high"], acc["bal"])
            acc["line"] = max(acc["line"], acc["eod_high"] - plan["max_loss"])
        if not acc["funded"]:
            tgt = plan["targets"][acc["stage"]]
            if acc["bal"] >= 1 + tgt - 1e-12:
                kind, n = plan.get("min_days", (None, 0))
                ok = not (kind == "trading" and acc["days"] < n) and not (kind == "profit" and acc["profit_days"] < n)
                if plan.get("best_day") and acc["best"] > plan["best_day"] * acc["pos_sum"] + 1e-12:
                    ok = False
                if ok:
                    credits = plan.get("credits")
                    if credits and acc["stage"] < len(credits):
                        credit += credits[acc["stage"]] * price
                    stage = acc["stage"] + 1
                    acc = new_account()
                    acc["stage"] = stage
                    if stage == len(plan["targets"]):
                        acc["funded"] = True
                        acc["bal"] += plan.get("funded_cash", 0.0) * price
                        stats["funded"] += 1
        else:
            acc["since"] += 1
            acc["funded_days"] += 1
            if acc["since"] >= CYCLE_DAYS:
                profit = acc["bal"] - 1.0
                ok = True
                if plan.get("payout_profit_days") and acc["profit_days"] < plan["payout_profit_days"]:
                    ok = False
                if plan.get("best_day") and profit > 0 and acc["best"] > plan["best_day"] * acc["pos_sum"] + 1e-12:
                    ok = False
                amount = profit if plan.get("full_payout") else acc["bal"] - (1 + keep)
                amount = min(amount, plan.get("payout_cap", 9.0))
                if ok and amount >= plan.get("payout_min", 0.0002):
                    if plan.get("payout_cap") and amount > plan["payout_cap"] + 1e-12:
                        raise Violation("出金上限超")
                    paid += amount * plan["payout_share"]
                    acc["payouts"] += 1
                    n = plan.get("refund_after_payouts", 0)
                    if (n and not acc["refunded"] and acc["payouts"] >= n
                            and acc["funded_days"] >= plan.get("refund_min_funded_days", 0)):
                        paid += price
                        acc["refunded"] = True
                    acc["bal"] -= amount
                    acc.update(since=0, profit_days=0, best=0.0, pos_sum=0.0)
                    if plan.get("lock_on_payout"):
                        acc["line"] = 1.0
                    if plan.get("reset_on_payout"):
                        acc.update(bal=1.0, line=1.0 - plan["max_loss"], eod_high=1.0)
        # 生き残っている口座が失格ラインより下にいないか
        if acc["bal"] < acc["line"] - 1e-9:
            raise Violation("失格ラインを割った口座が生きている")
    return paid, spent


def main():
    n = int(sys.argv[1]) if len(sys.argv) > 1 else 3000
    ref_path = sys.argv[2] if len(sys.argv) > 2 else None
    ref = json.load(open(ref_path, encoding="utf-8")) if ref_path else []
    edges = [("ゼロ", 0.50), ("小", 0.52), ("中", 0.55), ("強", 0.58)]
    rng = random.Random(20260927)
    worst_z = 0.0
    print(f"{'プラン':16s}{'実力':4s}{'r':>6s}{'B':>5s}{'検証の平均':>12s}{'±誤差':>9s}{'numba版':>11s}{'z':>7s}{'1回の期待値R':>12s}{'式':>8s}")
    for r in (0.01, 0.015):
        for name, plan in PLANS.items():
            for ename, p in edges:
                # numba 版で一番よかった「残す利益」を使う（無ければ0）
                cand = [x for x in ref if x["plan"] == name and x["edge"] == ename and abs(x["r"] - r) < 1e-9]
                best = max(cand, key=lambda x: x["net"]) if cand else None
                keep = best["buf"] if best else 0.0
                stats = dict(trades=0, pnl_R=0.0, max_day_loss=0.0, fails=0, funded=0)
                nets = []
                for _ in range(n):
                    paid, spent = one_year(plan, p, r, keep, rng, stats)
                    nets.append((paid - spent) * 100000)
                mean = sum(nets) / n
                sd = math.sqrt(sum((x - mean) ** 2 for x in nets) / (n - 1))
                se = sd / math.sqrt(n)
                ref_mean = best["net"] if best else float("nan")
                z = (mean - ref_mean) / se if best and se > 0 else float("nan")
                if best:
                    worst_z = max(worst_z, abs(z))
                per_trade = stats["pnl_R"] / max(stats["trades"], 1)
                print(f"{name:16s}{ename:4s}{r*100:5.1f}%{keep*100:4.0f}%{mean:+12.0f}{1.96*se:9.0f}{ref_mean:+11.0f}{z:+7.2f}"
                      f"{per_trade:+12.4f}{2*p-1.05:+8.3f}   最大の1日損失 {stats['max_day_loss']*100:.2f}%")
    print(f"\nルール違反: なし（違反があれば例外で止まる）／ numba 版との差の最大 |z| = {worst_z:.2f}")


if __name__ == "__main__":
    main()
