#!/usr/bin/env python3
"""
monte_carlo.py — DWPT vs（電池＋MCS）の確率的比較（既定 100,000 ケース）

各ケースで assumptions.csv の LOW〜HIGH の範囲から独立に乱数を引き、
  DWPT 総 LCOS [円/kWh] = インフラ + 車載受電器 + 電力
  代替 LCOS [円/kWh]   = 電池スループット費 + 質量ペナルティ + MCS 充電費 + 時間費用
を計算し、どちらが安いかを分類する。

勝率そのものより「DWPT が勝つケースに共通する条件」を抽出することが目的なので、
勝ちケースと全体のパラメータ分位点を並べて results/monte_carlo_conditions.csv に出す。

    python3 model/monte_carlo.py [n_cases]
"""
from __future__ import annotations

import sys
import csv
import json
import pathlib
import numpy as np

sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent))
from dwpt_model import load_assumptions, RESULTS, HOURS_PER_YEAR, DAYS_PER_YEAR  # noqa: E402

rng = np.random.default_rng(20261005)


def tri(low, base, high, n):
    """三角分布（low/high の順序が逆の変数にも対応）。"""
    lo, hi = min(low, high), max(low, high)
    mode = min(max(base, lo), hi)
    if hi - lo < 1e-12:
        return np.full(n, lo)
    return rng.triangular(lo, mode, hi, n)


def logu(low, high, n):
    """対数一様分布（桁で変わる量：CAPEX、交通量、普及率）。"""
    lo, hi = min(low, high), max(low, high)
    return np.exp(rng.uniform(np.log(lo), np.log(hi), n))


def crf(r, n_years):
    return np.where(r <= 0, 1.0 / n_years, r / (1.0 - (1.0 + r) ** (-n_years)))


def run(n: int = 100_000) -> None:
    L, B, H = load_assumptions("low"), load_assumptions("base"), load_assumptions("high")

    # --- 乱数 ---------------------------------------------------------------
    capex = logu(L["capex_dwpt_yen_per_lane_km"], H["capex_dwpt_yen_per_lane_km"], n)
    om = tri(L["om_fraction"], B["om_fraction"], H["om_fraction"], n)
    life = tri(H["lifetime_years"], B["lifetime_years"], L["lifetime_years"], n)  # 10..30
    r = tri(L["discount_rate"], B["discount_rate"], H["discount_rate"], n)
    aadt = tri(L["aadt_total"], B["aadt_total"], H["aadt_total"], n)
    hv = tri(L["heavy_vehicle_fraction"], B["heavy_vehicle_fraction"], H["heavy_vehicle_fraction"], n)
    # BEV 比率は 2026〜2040 を一様に混ぜる（どの年を評価しているかを明示するため ev_year を記録）
    ev_year = rng.choice([2026, 2030, 2035, 2040], n)
    ev = np.empty(n)
    pack_price = np.empty(n)
    for y in (2026, 2030, 2035, 2040):
        m = ev_year == y
        ev[m] = logu(L[f"ev_fraction_heavy_{y}"], H[f"ev_fraction_heavy_{y}"], m.sum())
        pack_price[m] = tri(L[f"battery_pack_yen_per_kwh_{y}"], B[f"battery_pack_yen_per_kwh_{y}"], H[f"battery_pack_yen_per_kwh_{y}"], m.sum())
    pen = logu(0.01, 1.0, n)
    lane_frac = tri(0.25, 0.5, 1.0, n)
    power = tri(L["peak_power_kw"], B["peak_power_kw"], 300.0, n)  # 将来 300 kW まで許容
    duty = tri(L["coupling_duty"], B["coupling_duty"], H["coupling_duty"], n)
    eff = tri(L["grid_to_battery_eff"], B["grid_to_battery_eff"], H["grid_to_battery_eff"], n)
    speed = tri(H["speed_kmh"], B["speed_kmh"], L["speed_kmh"], n)  # 80 / 90 / 100 km/h（大型貨物）
    elec = tri(L["electricity_yen_per_kwh"], B["electricity_yen_per_kwh"], H["electricity_yen_per_kwh"], n)
    recv_cost = tri(L["receiver_cost_yen"], B["receiver_cost_yen"], H["receiver_cost_yen"], n)
    recv_life = tri(H["receiver_lifetime_years"], B["receiver_lifetime_years"], L["receiver_lifetime_years"], n)
    dwpt_km_per_year = logu(1000.0, 36500.0, n)  # 受電器搭載車が年間に DWPT 区間を走る距離（回廊整備を仮定）

    kg_per_kwh = tri(L["battery_kg_per_kwh"], B["battery_kg_per_kwh"], H["battery_kg_per_kwh"], n)
    cycles = tri(L["battery_cycle_life"], B["battery_cycle_life"], H["battery_cycle_life"], n)
    truck_km = tri(L["truck_km_per_year"], B["truck_km_per_year"], H["truck_km_per_year"], n)
    kwh_km = tri(L["truck_kwh_per_km"], B["truck_kwh_per_km"], H["truck_kwh_per_km"], n)
    truck_life = tri(L["truck_life_years"], B["truck_life_years"], H["truck_life_years"], n)
    payload_val = tri(L["payload_value_yen_per_tonne_km"], B["payload_value_yen_per_tonne_km"], H["payload_value_yen_per_tonne_km"], n)
    payload_bind = tri(0.1, 0.3, 0.6, n)
    delta_kwh = tri(50.0, 200.0, 500.0, n)
    base_pack = tri(300.0, 400.0, 600.0, n)

    mcs_capex = tri(L["mcs_charger_capex_yen"], B["mcs_charger_capex_yen"], H["mcs_charger_capex_yen"], n)
    mcs_grid = tri(L["mcs_site_grid_capex_yen"], B["mcs_site_grid_capex_yen"], H["mcs_site_grid_capex_yen"], n) / tri(4, 6, 8, n)
    mcs_util = tri(L["mcs_utilization"], B["mcs_utilization"], H["mcs_utilization"], n)
    mcs_power = tri(L["mcs_power_kw"], B["mcs_power_kw"], H["mcs_power_kw"], n)
    mcs_life = tri(L["mcs_lifetime_years"], B["mcs_lifetime_years"], H["mcs_lifetime_years"], n)
    mcs_om = tri(L["mcs_om_fraction"], B["mcs_om_fraction"], H["mcs_om_fraction"], n)
    demand = tri(L["demand_charge_yen_per_kw_month"], B["demand_charge_yen_per_kw_month"], H["demand_charge_yen_per_kw_month"], n)
    hour_cost = tri(L["driver_vehicle_hour_yen"], B["driver_vehicle_hour_yen"], H["driver_vehicle_hour_yen"], n)
    overlap = rng.uniform(0.0, 1.0, n)

    # --- DWPT ---------------------------------------------------------------
    flow = aadt * hv * ev * pen * lane_frac                      # 台/日
    e_pass_1km = power * duty * (1000.0 / (speed / 3.6)) / 3600  # kWh/pass/km
    annual_kwh_km = flow * e_pass_1km * DAYS_PER_YEAR
    ann_capex = capex * (crf(r, life) + om)
    dwpt_infra = ann_capex / annual_kwh_km
    recv_kwh_year = e_pass_1km * dwpt_km_per_year
    dwpt_recv = recv_cost * crf(r, recv_life) / recv_kwh_year
    dwpt_elec = elec / eff
    # 契約電力：ピーク時間帯 30 分の平均 kW（下限 50〜150 kW）× 基本料金（レビュー R3-01）
    contract_kw = np.maximum(tri(L["dwpt_contract_kw_min"], B["dwpt_contract_kw_min"], H["dwpt_contract_kw_min"], n), flow * 0.10 * e_pass_1km)
    dwpt_demand = contract_kw * demand * 12 / annual_kwh_km
    dwpt_total = dwpt_infra + dwpt_recv + dwpt_elec + dwpt_demand

    # --- 代替：電池 + MCS -----------------------------------------------------
    annual_use = kwh_km * truck_km
    cyc_year = annual_use / (base_pack + delta_kwh)
    life_b = np.minimum(truck_life, cycles / cyc_year)
    bat_thr = delta_kwh * pack_price * crf(r, life_b) / (delta_kwh * cyc_year)
    mass_t = delta_kwh * kg_per_kwh / 1000.0
    # 分母は ΔE の年間スループット（償却と同じ）に揃える（dwpt_model.battery_throughput_cost と同一の式）
    bat_payload = mass_t * payload_val * payload_bind * truck_km / (delta_kwh * cyc_year)
    bat_energy = mass_t * 0.012 * annual_use * elec / (delta_kwh * cyc_year)
    annual_kwh_stall = mcs_power * HOURS_PER_YEAR * mcs_util * 0.7
    mcs_infra = (mcs_capex + mcs_grid) * (crf(r, mcs_life) + mcs_om) / annual_kwh_stall
    mcs_demand = demand * 12 * mcs_power * 0.6 / annual_kwh_stall
    mcs_elec = elec / 0.93
    mcs_time = hour_cost / (mcs_power * 0.7) * (1.0 - overlap)
    alt_total = bat_thr + bat_payload + bat_energy + mcs_infra + mcs_demand + mcs_elec + mcs_time

    winner = np.where(dwpt_total < alt_total, "DWPT", "BATTERY_MCS")
    margin = alt_total - dwpt_total

    cols = {
        "ev_year": ev_year, "capex_yen_per_lane_km": capex, "om_fraction": om, "lifetime_years": life, "discount_rate": r,
        "aadt": aadt, "heavy_fraction": hv, "ev_fraction": ev, "dwpt_penetration": pen, "lane_fraction": lane_frac,
        "target_flow_per_day": flow, "power_kw": power, "duty": duty, "eff": eff, "speed_kmh": speed,
        "electricity_yen_per_kwh": elec, "receiver_cost_yen": recv_cost, "dwpt_km_per_year": dwpt_km_per_year,
        "pack_price_yen_per_kwh": pack_price, "delta_kwh": delta_kwh, "kg_per_kwh": kg_per_kwh, "cycle_life": cycles,
        "mcs_charger_capex_yen": mcs_capex, "mcs_utilisation": mcs_util, "mcs_power_kw": mcs_power, "logistics_overlap": overlap,
        "hour_cost_yen": hour_cost, "annual_mwh_per_lane_km": annual_kwh_km / 1000, "dwpt_infra_yen_per_kwh": dwpt_infra,
        "dwpt_total_yen_per_kwh": dwpt_total, "alt_total_yen_per_kwh": alt_total, "margin_yen_per_kwh": margin, "winner": winner,
    }
    RESULTS.mkdir(exist_ok=True)
    # 全ケース（サイズ抑制のため 20,000 行を層化せずランダム抽出で保存し、集計は全件で行う）
    keep = rng.choice(n, size=min(n, 20_000), replace=False)
    with open(RESULTS / "monte_carlo.csv", "w", newline="", encoding="utf-8") as f:
        w = csv.writer(f)
        w.writerow(list(cols.keys()))
        for i in keep:
            w.writerow([cols[k][i] if isinstance(cols[k][i], str) else f"{cols[k][i]:.6g}" for k in cols])

    # --- 勝ちケースの条件抽出 ------------------------------------------------
    win = winner == "DWPT"
    summary = {"n_cases": int(n), "dwpt_win_rate": float(win.mean()),
               "dwpt_win_rate_by_year": {int(y): float(win[ev_year == y].mean()) for y in (2026, 2030, 2035, 2040)},
               "dwpt_win_rate_flow_gt_1000": float(win[flow > 1000].mean()) if (flow > 1000).any() else None,
               "dwpt_win_rate_flow_lt_100": float(win[flow < 100].mean()) if (flow < 100).any() else None,
               "median_dwpt_infra_all": float(np.median(dwpt_infra)), "median_alt_all": float(np.median(alt_total))}
    rows = []
    qs = (0.05, 0.25, 0.5, 0.75, 0.95)
    for k in ("target_flow_per_day", "capex_yen_per_lane_km", "lifetime_years", "discount_rate", "dwpt_penetration", "ev_fraction",
              "aadt", "heavy_fraction", "lane_fraction", "power_kw", "duty", "speed_kmh", "annual_mwh_per_lane_km",
              "pack_price_yen_per_kwh", "delta_kwh", "mcs_utilisation", "mcs_charger_capex_yen", "logistics_overlap", "hour_cost_yen", "dwpt_km_per_year"):
        v = cols[k]
        row = {"parameter": k}
        for q in qs:
            row[f"all_p{int(q*100)}"] = float(np.quantile(v, q))
        for q in qs:
            row[f"win_p{int(q*100)}"] = float(np.quantile(v[win], q)) if win.any() else None
        rows.append(row)
    with open(RESULTS / "monte_carlo_conditions.csv", "w", newline="", encoding="utf-8") as f:
        w = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
        w.writeheader()
        for row in rows:
            w.writerow({k: (f"{v:.6g}" if isinstance(v, float) else v) for k, v in row.items()})

    # 必要条件の抽出：勝ちケースの流量の最小値、勝率が 50% を超える流量閾値
    order = np.argsort(flow)
    f_sorted, w_sorted = flow[order], win[order]
    bins = np.logspace(-1, 5, 61)
    idx = np.digitize(f_sorted, bins)
    curve = []
    for b in range(1, len(bins)):
        m = idx == b
        if m.sum() >= 50:
            curve.append({"flow_bin_low": float(bins[b - 1]), "flow_bin_high": float(bins[b]), "n": int(m.sum()), "dwpt_win_rate": float(w_sorted[m].mean())})
    with open(RESULTS / "monte_carlo_winrate_by_flow.csv", "w", newline="", encoding="utf-8") as f:
        w = csv.DictWriter(f, fieldnames=list(curve[0].keys()))
        w.writeheader()
        w.writerows(curve)
    summary["min_flow_in_dwpt_wins"] = float(flow[win].min()) if win.any() else None
    summary["flow_threshold_50pct_win"] = next((c["flow_bin_low"] for c in curve if c["dwpt_win_rate"] >= 0.5), None)
    with open(RESULTS / "monte_carlo_summary.json", "w", encoding="utf-8") as f:
        json.dump(summary, f, ensure_ascii=False, indent=2)
    print(json.dumps(summary, ensure_ascii=False, indent=2))


def run_wide(n: int = 100_000) -> None:
    """館山道に縛らない「汎用回廊」空間。対象車流量 1〜30,000 台/日・車線、CAPEX 0.5〜10 億円を対数一様に引き、
    DWPT 勝率を (流量 × CAPEX) の 2 次元ヒートマップにする。他の変数は run() と同じ分布。"""
    L, B, H = load_assumptions("low"), load_assumptions("base"), load_assumptions("high")
    flow = logu(1.0, 30000.0, n)
    capex = logu(0.5e8, 10e8, n)
    om = tri(L["om_fraction"], B["om_fraction"], H["om_fraction"], n)
    life = tri(H["lifetime_years"], B["lifetime_years"], L["lifetime_years"], n)
    r = tri(L["discount_rate"], B["discount_rate"], H["discount_rate"], n)
    power = tri(100.0, 150.0, 300.0, n)
    duty = tri(L["coupling_duty"], B["coupling_duty"], H["coupling_duty"], n)
    eff = tri(L["grid_to_battery_eff"], B["grid_to_battery_eff"], H["grid_to_battery_eff"], n)
    speed = tri(H["speed_kmh"], B["speed_kmh"], L["speed_kmh"], n)  # 80 / 90 / 100 km/h（大型貨物）
    elec = tri(L["electricity_yen_per_kwh"], B["electricity_yen_per_kwh"], H["electricity_yen_per_kwh"], n)
    recv_cost = tri(L["receiver_cost_yen"], B["receiver_cost_yen"], H["receiver_cost_yen"], n)
    recv_life = tri(H["receiver_lifetime_years"], B["receiver_lifetime_years"], L["receiver_lifetime_years"], n)
    dwpt_km = logu(1000.0, 36500.0, n)
    pack_price = logu(6000.0, 30000.0, n)
    kg_per_kwh = tri(L["battery_kg_per_kwh"], B["battery_kg_per_kwh"], H["battery_kg_per_kwh"], n)
    cycles = tri(L["battery_cycle_life"], B["battery_cycle_life"], H["battery_cycle_life"], n)
    truck_km = tri(L["truck_km_per_year"], B["truck_km_per_year"], H["truck_km_per_year"], n)
    kwh_km = tri(L["truck_kwh_per_km"], B["truck_kwh_per_km"], H["truck_kwh_per_km"], n)
    truck_life = tri(L["truck_life_years"], B["truck_life_years"], H["truck_life_years"], n)
    payload_val = tri(L["payload_value_yen_per_tonne_km"], B["payload_value_yen_per_tonne_km"], H["payload_value_yen_per_tonne_km"], n)
    payload_bind = tri(0.1, 0.3, 0.6, n)
    delta_kwh = tri(50.0, 200.0, 500.0, n)
    base_pack = tri(300.0, 400.0, 600.0, n)
    mcs_capex = tri(L["mcs_charger_capex_yen"], B["mcs_charger_capex_yen"], H["mcs_charger_capex_yen"], n)
    mcs_grid = tri(L["mcs_site_grid_capex_yen"], B["mcs_site_grid_capex_yen"], H["mcs_site_grid_capex_yen"], n) / tri(4, 6, 8, n)
    mcs_util = tri(L["mcs_utilization"], B["mcs_utilization"], H["mcs_utilization"], n)
    mcs_power = tri(L["mcs_power_kw"], B["mcs_power_kw"], H["mcs_power_kw"], n)
    mcs_life = tri(L["mcs_lifetime_years"], B["mcs_lifetime_years"], H["mcs_lifetime_years"], n)
    mcs_om = tri(L["mcs_om_fraction"], B["mcs_om_fraction"], H["mcs_om_fraction"], n)
    demand = tri(L["demand_charge_yen_per_kw_month"], B["demand_charge_yen_per_kw_month"], H["demand_charge_yen_per_kw_month"], n)
    hour_cost = tri(L["driver_vehicle_hour_yen"], B["driver_vehicle_hour_yen"], H["driver_vehicle_hour_yen"], n)
    overlap = rng.uniform(0.0, 1.0, n)

    e_pass_1km = power * duty * (1000.0 / (speed / 3.6)) / 3600
    annual_kwh_km = flow * e_pass_1km * DAYS_PER_YEAR
    contract_kw = np.maximum(tri(L["dwpt_contract_kw_min"], B["dwpt_contract_kw_min"], H["dwpt_contract_kw_min"], n), flow * 0.10 * e_pass_1km)
    dwpt_total = capex * (crf(r, life) + om) / annual_kwh_km + recv_cost * crf(r, recv_life) / (e_pass_1km * dwpt_km) + elec / eff + contract_kw * demand * 12 / annual_kwh_km
    cyc_year = kwh_km * truck_km / (base_pack + delta_kwh)
    life_b = np.minimum(truck_life, cycles / cyc_year)
    mass_t = delta_kwh * kg_per_kwh / 1000.0
    bat = pack_price * crf(r, life_b) / cyc_year + mass_t * (payload_val * payload_bind * truck_km + 0.012 * kwh_km * truck_km * elec) / (delta_kwh * cyc_year)
    annual_kwh_stall = mcs_power * HOURS_PER_YEAR * mcs_util * 0.7
    mcs = (mcs_capex + mcs_grid) * (crf(r, mcs_life) + mcs_om) / annual_kwh_stall + demand * 12 * mcs_power * 0.6 / annual_kwh_stall + elec / 0.93 + hour_cost / (mcs_power * 0.7) * (1 - overlap)
    alt_total = bat + mcs
    win = dwpt_total < alt_total

    fb = np.logspace(0, np.log10(30000), 31)
    cb = np.logspace(np.log10(0.5e8), np.log10(10e8), 21)
    fi = np.clip(np.digitize(flow, fb) - 1, 0, len(fb) - 2)
    ci = np.clip(np.digitize(capex, cb) - 1, 0, len(cb) - 2)
    grid = np.full((len(cb) - 1, len(fb) - 1), np.nan)
    cnt = np.zeros_like(grid)
    for i in range(len(cb) - 1):
        for j in range(len(fb) - 1):
            m = (ci == i) & (fi == j)
            cnt[i, j] = m.sum()
            if m.sum() >= 20:
                grid[i, j] = win[m].mean()
    np.savez(RESULTS / "monte_carlo_wide_grid.npz", grid=grid, flow_bins=fb, capex_bins=cb, count=cnt)
    with open(RESULTS / "monte_carlo_wide_summary.json", "w", encoding="utf-8") as f:
        json.dump({"n_cases": int(n), "dwpt_win_rate": float(win.mean()),
                   "win_rate_flow_gt_2000": float(win[flow > 2000].mean()), "win_rate_flow_gt_5000": float(win[flow > 5000].mean()),
                   "win_rate_flow_gt_5000_capex_lt_2oku": float(win[(flow > 5000) & (capex < 2e8)].mean()),
                   "win_rate_flow_lt_500": float(win[flow < 500].mean()),
                   "p50_flow_in_wins": float(np.median(flow[win])), "p50_capex_in_wins": float(np.median(capex[win])),
                   "p95_alt_total": float(np.quantile(alt_total, 0.95)), "p50_alt_total": float(np.median(alt_total))}, f, ensure_ascii=False, indent=2)
    print(open(RESULTS / "monte_carlo_wide_summary.json", encoding="utf-8").read())


if __name__ == "__main__":
    n = int(sys.argv[1]) if len(sys.argv) > 1 and sys.argv[1].isdigit() else 100_000
    run(n)
    run_wide(n)
