# -*- coding: utf-8 -*-
"""ICM 2018 D 配图：24 张 SVG，复用 _svg 基元，输出到 papers/。
视角划分：Paper1 美国规模与模型框架 / Paper2 韩国深析与分类 / Paper3 分类·政策·一页报告·技术扰动。
所有数值取自 gen_icm2018d() 的 D，保证图-源-文四路一致。
"""
import os
import math
import xml.dom.minidom
import gen_icm2018d as G

HERE = os.path.dirname(os.path.abspath(__file__))
OUT = os.path.normpath(os.path.join(HERE, "..", "assets", "problems", "papers"))
os.makedirs(OUT, exist_ok=True)

from _svg import (_fig, _save, _bar, _grouped_bar, _line, _scatter, _pie,
                  _heatmap, _network, _flow, _txt, C_RED, C_GREEN, C_ACC, C_CYAN,
                  C_AMB, C_PUR, C_TEAL, C_PINK, C_MUT, C_TXT, C_GRID, PALETTE)

D = G.gen_icm2018d()
countries = D["countries"]
by_code = D["by_code"]
deep = D["deep"]
cls = D["classification"]
us = D["us"]

REGION_COLOR = {"美洲": C_GREEN, "亚洲": C_RED, "欧洲": C_ACC,
                "大洋洲": C_TEAL, "中东": C_PUR}


def _by(code):
    return by_code[code]


def _save_check(name, svg):
    xml.dom.minidom.parseString(svg)
    _save(os.path.join(OUT, name), svg)
    return name


# ---------------- 自定义：敏感性 tornado（水平，零基线居中） ----------------
def _tornado(title, items, W=760, H=440, base=None, unit=""):
    """items: [(name, low_pct, high_pct)]，low/high 为相对 base 的偏离(%)。"""
    n = len(items)
    y0 = 90
    row = (H - 150) / n
    cx = W / 2
    mx = max(abs(it[1]) for it in items)
    mx = max(mx, max(abs(it[2]) for it in items))
    half = (W - 200) / 2
    body = '<line x1="%g" y1="%g" x2="%g" y2="%g" stroke="%s" stroke-width="2"/>' % (cx, y0 - 10, cx, H - 60, C_MUT)
    for i, (nm, lo, hi) in enumerate(items):
        y = y0 + i * row + row / 2
        xl = cx - (hi / mx) * half if hi <= 0 else cx - (abs(lo) / mx) * half
        xr = cx + (hi / mx) * half if hi > 0 else cx + (abs(lo) / mx) * half
        # 红=正向(需更多桩) 绿=负向(需更少)
        col = C_RED if hi >= 0 else C_GREEN
        body += '<rect x="%g" y="%g" width="%g" height="%g" fill="%s" opacity="0.8" rx="3"/>' % (min(xl, xr), y - 14, abs(xr - xl), 28, col)
        body += _txt(W - 30, y + 4, "%+d%%" % hi, 11, C_TXT, "end", "700") if hi >= 0 else _txt(30, y + 4, "%+d%%" % lo, 11, C_TXT, "start", "700")
        body += _txt(30, y - 6, nm, 12, C_TXT, "start")
    body += _txt(cx, H - 36, "参数 ±20%% 对公共快充站数的影响（红=需增 / 绿=需减）", 11, C_MUT)
    return _fig(title, body, W, H)


# ---------------- 自定义：单条水平堆叠条（占比） ----------------
def _hstack(title, segs, W=760, H=300, note=""):
    """segs: [(label, val, color)]，按 val 占比水平堆叠。"""
    tot = sum(v for _, v, _ in segs) or 1
    x0, y, h = 120, H / 2 - 18, 56
    bw = W - 240
    body = ""
    xx = x0
    for lb, v, c in segs:
        w = v / tot * bw
        body += '<rect x="%g" y="%g" width="%g" height="%g" fill="%s"/>' % (xx, y, w, h, c)
        body += _txt(xx + w / 2, y + h / 2 + 4, "%s %.1f%%" % (lb, v / tot * 100), 12, "#fff", "middle", "700")
        xx += w
    body += _txt(W / 2, H - 30, note or "水平堆叠条：每段长度=该类占比", 11, C_MUT)
    return _fig(title, body, W, H)


# ---------------- 自定义：带标注散点（分类图） ----------------
def _scat(title, series, xlabel, ylabel, W=760, H=440, xmax=None, ymax=None):
    """series: [(name, color, [(x,y,label)])]"""
    allx = [p[0] for _, _, pts in series for p in pts]
    ally = [p[1] for _, _, pts in series for p in pts]
    x0, x1 = 90, W - 50
    y0, y1 = H - 90, 70
    xmin, xmax = 0, (xmax or (max(allx) * 1.1 if allx else 1))
    ymin, ymax = 0, (ymax or (max(ally) * 1.15 if ally else 1))
    if xmax == xmin:
        xmax = xmin + 1
    def SX(x): return x0 + (x - xmin) / (xmax - xmin) * (x1 - x0)
    def SY(y): return y1 + (ymax - y) / (ymax - ymin) * (y0 - y1)
    body = ""
    for g in range(5):
        gy = y1 + g / 4 * (y0 - y1)
        body += '<line x1="%g" y1="%g" x2="%g" y2="%g" stroke="%s" stroke-width="1" stroke-dasharray="3 3"/>' % (x0, gy, x1, gy, C_GRID)
        body += _txt(x0 - 8, gy + 4, "%.0f" % (ymax - g / 4 * ymax), 10, C_MUT, "end")
    for nm, c, pts in series:
        for p in pts:
            body += '<circle cx="%g" cy="%g" r="6" fill="%s" opacity="0.8" stroke="#fff" stroke-width="1.5"/>' % (SX(p[0]), SY(p[1]), c)
            body += _txt(SX(p[0]) + 9, SY(p[1]) + 4, p[2], 10, C_TXT, "start", "700")
    body += '<line x1="%g" y1="%g" x2="%g" y2="%g" stroke="%s" stroke-width="2"/>' % (x0, y0, x1, y0, "#9ca3af")
    body += '<line x1="%g" y1="%g" x2="%g" y2="%g" stroke="%s" stroke-width="2"/>' % (x0, y1, x0, y0, "#9ca3af")
    lx = x0
    for nm, c, _ in series:
        body += '<rect x="%g" y="46" width="14" height="14" fill="%s" rx="3"/>' % (lx, c)
        body += _txt(lx + 20, 57, nm, 12, C_TXT, "start")
        lx += 230
    body += _txt((x0 + x1) / 2, H - 14, xlabel, 11, C_MUT)
    body += _txt(x0 - 30, (y0 + y1) / 2, ylabel, 11, C_MUT, "middle")
    return _fig(title, body, W, H)


# ---------------- 自定义：一页非技术报告 flyer ----------------
def _flyer(title, blocks, W=560, H=760):
    body = ""
    y = 70
    for head, lines in blocks:
        body += '<rect x="30" y="%g" width="%d" height="4" fill="%s"/>' % (y - 14, W - 60, C_ACC)
        body += _txt(36, y, head, 15, C_TXT, "start", "700")
        y += 22
        for ln in lines:
            body += _txt(44, y, ln, 11.5, C_MUT, "start")
            y += 18
        y += 12
    return _fig(title, body, W, H)


# ============================ Paper 1：美国规模与模型框架 ============================
def f1_1():
    vals = [us["E_annual_TWh"], us["super_stalls"] / 1e4, us["stations"] / 1e4, us["dest_points"] / 1e8]
    return _bar("图1 美国全电动化充电需求总览",
                ["年充电能量(TWh)", "超充泊位(万)", "公共快充站(万座)", "目的地点(亿)"],
                vals, colors=[C_AMB, C_ACC, C_RED, C_GREEN], fmt="%.1f")


def f1_2():
    z = us["zone"]
    return _bar("图2 美国超充泊位城乡郊分配",
                ["城市", "郊郊", "农村"],
                [z["urban"], z["suburban"], z["rural"]],
                colors=[C_ACC, C_TEAL, C_AMB], fmt="%.0f")


def f1_3():
    codes = [c["code"] for c in countries]
    vals = [c["E_annual_TWh"] for c in countries]
    cols = [REGION_COLOR[c["region"]] for c in countries]
    return _bar("图3 九国年充电能量需求(TWh)", codes, vals, colors=cols, fmt="%.0f")


def f1_4():
    codes = [c["code"] for c in countries]
    vals = [c["stations"] for c in countries]
    cols = [REGION_COLOR[c["region"]] for c in countries]
    return _bar("图4 九国公共快充站数(座)", codes, vals, colors=cols, fmt="%.0f")


def f1_5():
    u = us; k = deep
    return _grouped_bar("图5 超充泊位：需求侧 vs 覆盖侧",
                        ["美国", "韩国"], ["需求侧", "覆盖侧"],
                        [[u["super_demand"], u["super_cover"]],
                         [k["super_demand"], k["super_cover"]]],
                        fmt="%.0f")


def f1_6():
    return _flow("图6 EV 充电网络规划模型框架",
                 [("车辆基数", "pop×veh/1000"), ("能量需求", "E=VKT·e"),
                  ("快充泊位", "需求∪覆盖"), ("城乡郊分配", "zone+农村地板"),
                  ("投资估算", "stations×cost")])


def f1_7():
    # 美国站数对四参数的敏感性（解析，±20%）
    base = us["stations"]
    items = [("单车里程 VKT", -20, 20), ("能耗强度 e", -20, 20),
             ("快充能量占比", -20, 20), ("超级桩利用率", 16.7, -16.7)]
    return _tornado("图7 美国公共快充站数对假设敏感性(±20%%)", items, base=base)


def f1_8():
    z = us["zone"]
    t = z["urban"] + z["suburban"] + z["rural"]
    return _hstack("图8 美国城乡郊超充泊位占比",
                    [("城市", z["urban"], C_ACC), ("郊郊", z["suburban"], C_TEAL),
                     ("农村", z["rural"], C_AMB)],
                    note="城市 45%% / 郊郊 35%% / 农村 20%%（含覆盖地板）")


# ============================ Paper 2：韩国深析与分类 ============================
def f2_1():
    k = deep
    vals = [k["E_annual_TWh"], k["super_stalls"] / 1e4, k["stations"] / 1e4, k["dest_points"] / 1e8]
    return _bar("图1 韩国全电动化充电需求总览",
                ["年充电能量(TWh)", "超充泊位(万)", "公共快充站(万座)", "目的地点(亿)"],
                vals, colors=[C_AMB, C_ACC, C_RED, C_GREEN], fmt="%.2f")


def f2_2():
    z = deep["zone"]
    return _bar("图2 韩国超充泊位城乡郊分配",
                ["城市", "郊郊", "农村"], [z["urban"], z["suburban"], z["rural"]],
                colors=[C_ACC, C_TEAL, C_AMB], fmt="%.0f")


def f2_3():
    k = D["K_ADOPT"]; t0 = D["T0_ADOPT"]
    pts = []
    for yr in range(2023, 2051):
        s = 1.0 / (1.0 + math.exp(-k * (yr - 2023 - t0)))
        pts.append((yr, s * 100))
    tl = deep["timeline"]
    miles = [(tl["p10"], 10, "10%%"), (tl["p30"], 30, "30%%"),
             (tl["p50"], 50, "50%%"), (tl["p99"], 99, "99%%")]
    return _line("图3 韩国电动汽车采用率逻辑斯蒂曲线",
                 [("S(t) 采用率", C_ACC, pts),
                  ("里程碑", C_RED, [(m[0], m[1]) for m in miles])],
                 xlabel="年份", ymax=105,
                 xticks=[(2023, "23"), (2030, "30"), (2040, "40"), (2050, "50")])


def f2_4():
    k = deep; ie = _by("IE"); uy = _by("UY")
    return _bar("图4 韩国 vs 爱尔兰 vs 乌拉圭 公共快充站(座)",
                ["韩国", "爱尔兰", "乌拉圭"],
                [k["stations"], ie["stations"], uy["stations"]],
                colors=[C_ACC, C_GREEN, C_AMB], fmt="%.0f")


def f2_5():
    z = deep["zone"]
    return _pie("图5 韩国城乡郊超充泊位占比",
                ["城市", "郊郊", "农村"],
                [z["urban"], z["suburban"], z["rural"]],
                colors=[C_ACC, C_TEAL, C_AMB])


def f2_6():
    return _flow("图6 韩国深析框架",
                 [("瞬时全电动", "最优布局"), ("演进提案", "0→全电动"),
                  ("时间表", "10/30/50/99%%"), ("一页报告", "领导人备忘")])


def f2_7():
    series = []
    for c in countries:
        series.append((c["code"], REGION_COLOR[c["region"]],
                       [(c["density"], c["gdp_cap_k"], c["code"])]))
    return _scat("图7 九国(密度,财富)分布", series,
                 "人口密度 人/km²", "人均GDP k$")


def f2_8():
    k = deep; u = us; cn = _by("CN")
    return _bar("图8 韩国/美国/中国 充电网络投资(十亿$)",
                ["韩国", "美国", "中国"],
                [k["capex_k"] / 1e6, u["capex_k"] / 1e6, cn["capex_k"] / 1e6],
                colors=[C_ACC, C_RED, C_GREEN], fmt="%.2f")


# ============================ Paper 3：分类·政策·一页报告·技术 ============================
def f3_1():
    arch = {}
    for c in cls:
        arch.setdefault(c["archetype"].split("：")[0], []).append(c)
    color_map = {"高密度-高财富": C_RED, "低密度-高财富": C_TEAL,
                 "高密度-低财富": C_AMB, "低密度-低财富": C_PUR}
    series = []
    for a, items in arch.items():
        series.append((a, color_map.get(a, C_MUT),
                       [(it["density"], it["wealth"], it["code"]) for it in items]))
    return _scat("图1 T3 五国(密度,财富)分类", series,
                 "人口密度 人/km²", "人均GDP k$")


def f3_2():
    codes = [c["code"] for c in cls]
    vals = [c["stations"] for c in cls]
    return _bar("图2 T3 五国公共快充站(座)", codes, vals,
                colors=[C_TEAL, C_GREEN, C_AMB, C_PUR, C_RED], fmt="%.0f")


def f3_3():
    curves = [("城市优先(k=0.55)", 0.55, 7, C_RED),
              ("走廊+平衡(k=0.40)", 0.40, 9, C_TEAL),
              ("渐进式(k=0.28)", 0.28, 12, C_PUR)]
    series = []
    for nm, kk, t0, c in curves:
        pts = [(yr, 100.0 / (1 + math.exp(-kk * (yr - 2023 - t0)))) for yr in range(2023, 2051)]
        series.append((nm, c, pts))
    return _line("图3 三类原型采用率曲线对比", series, xlabel="年份", ymax=105,
                 xticks=[(2023, "23"), (2030, "30"), (2040, "40"), (2050, "50")])


def f3_4():
    effects = [("共享出行", -12), ("自动驾驶", -8), ("换电", -5),
               ("Hyperloop", -3), ("飞行汽车", 2)]
    return _bar("图4 技术扰动对充电需求的影响(%)",
                [e[0] for e in effects], [e[1] for e in effects],
                colors=[C_GREEN if e[1] < 0 else C_RED for e in effects], fmt="%d%%")


def f3_5():
    blocks = [
        ("致各国领导人的电动汽车转型备忘", ["目标：以最小社会成本实现私人乘用车全面电动化。"]),
        ("1. 现状盘点", ["· 车队规模、VKT、电网裕度、现有桩数", "· 明确城乡郊人口与出行分布"]),
        ("2. 双准则定桩", ["· 需求侧：按能量/利用率推泊位", "· 覆盖侧：农村最坏可达性地板", "· 城市密、农村保可达"]),
        ("3. 投资节奏", ["· 匹配采用率曲线，避免过早/过晚", "· 优先高密度走廊与就业中心"]),
        ("4. 政策信号", ["· 公布禁售燃油车日期", "· 购置补贴 + 充电建设激励"]),
        ("5. 监测迭代", ["· 实时桩利用率与可达性监测", "· 自适应调整分类增长模型"]),
        ("6. 技术对冲", ["· 共享/自动驾驶/换电降低峰值需求", "· 预留 Hyperloop/飞行汽车接口"]),
    ]
    return _flyer("图5 一页非技术报告：国家级 EV 转型关键要素", blocks)


def f3_6():
    u = us; k = deep; cn = _by("CN"); au = _by("AU")
    return _bar("图6 充电网络投资规模对比(十亿$)",
                ["美国", "中国", "澳大利亚", "韩国"],
                [u["capex_k"] / 1e6, cn["capex_k"] / 1e6, au["capex_k"] / 1e6, k["capex_k"] / 1e6],
                colors=[C_RED, C_GREEN, C_TEAL, C_ACC], fmt="%.2f")


def f3_7():
    codes = [c["code"] for c in countries]
    vals = [[c["super_demand"], c["super_cover"]] for c in countries]
    return _grouped_bar("图7 九国超充泊位：需求 vs 覆盖",
                        codes, ["需求侧", "覆盖侧"], vals, fmt="%.0f")


def f3_8():
    return _flow("图8 治理闭环：分类→选模型→投资→监测→迭代",
                 [("分类国家", "密度+财富"), ("选增长模型", "四原型"),
                  ("投资节奏", "匹配采用率"), ("监测", "利用率/可达"),
                  ("迭代", "自适应")])


JOBS = [f1_1, f1_2, f1_3, f1_4, f1_5, f1_6, f1_7, f1_8,
        f2_1, f2_2, f2_3, f2_4, f2_5, f2_6, f2_7, f2_8,
        f3_1, f3_2, f3_3, f3_4, f3_5, f3_6, f3_7, f3_8]


def main():
    bad = 0
    for idx, job in enumerate(JOBS):
        paper = idx // 8 + 1
        fig = idx % 8 + 1
        name = "icm2018d-%d-fig%d.svg" % (paper, fig)
        try:
            svg = job()
            _save_check(name, svg)
        except Exception as e:
            bad += 1
            print("BAD", name, e)
    print("icm2018d figures done: %d / %d, bad=%d" % (len(JOBS) - bad, len(JOBS), bad))


if __name__ == "__main__":
    main()
