# -*- coding: utf-8 -*-
"""DGCUP 2021 A 高铁牵引供电系统 —— 配图脚本（24 张 SVG）。
依赖 _svg.py 基元库；数据全部来自 gen_dgcup2021a.gen_dgcup2021a()。
运行：cd tools/ && python3 fig_dgcup2021a.py
"""
import os
import math
import random
import gen_dgcup2021a as G
from _svg import (_fig, _save, _bar, _grouped_bar, _line, _scatter, _pie,
                  _network, _flow, _txt, _heatmap,
                  C_ACC, C_RED, C_GREEN, C_AMB, C_PUR, C_CYAN, C_TEAL,
                  C_PINK, C_GRID, C_MUT, PALETTE)

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

D = G.gen_dgcup2021a()
s = D["series"]
th = s["time_h"]


def save(name, svg):
    _save(os.path.join(OUT, name), svg)


def clean_pairs(xs, ys):
    return [(x, y) for x, y in zip(xs, ys) if y == y]


def bucket_avg(arr, edges):
    """edges: 区间起点列表（如 [0,6,10,16]），返回各区间均值。"""
    res = []
    for i in range(len(edges) - 1):
        a, b = edges[i], edges[i + 1]
        vals = [arr[t] for t in range(len(arr)) if a <= th[t] < b]
        res.append(sum(vals) / len(vals) if vals else 0.0)
    return res


# ============ Paper 1：数据预处理与运行工况特征 ============
def p1_fig1():  # 清洗前后馈线电压
    raw = clean_pairs(th, [v if v == v else float("nan") for v in s["U_raw"]])
    raw = [(x, y) for x, y in zip(th, s["U_raw"]) if y == y]
    cln = list(zip(th, s["U_clean"]))
    return _line("图1 馈线电压清洗前后对比（线性插值缺失、裁剪离群）",
                 [("原始(含缺失/离群)", C_RED, raw), ("清洗后", C_ACC, cln)],
                 xlabel="时刻 h", ylabel="U (kV)")


def p1_fig2():  # 清洗前后总电流
    raw = [(x, y) for x, y in zip(th, s["Itot_raw"]) if y == y]
    cln = list(zip(th, s["Itot_clean"]))
    return _line("图2 总牵引电流清洗前后对比",
                 [("原始", C_RED, raw), ("清洗后", C_GREEN, cln)],
                 xlabel="时刻 h", ylabel="I (A)")


def p1_fig3():  # 日负荷曲线（按小时平均有功）
    hrs = [str(h) for h in range(24)]
    return _bar("图3 日负荷曲线（各小时平均有功功率 MW）", hrs, s["hourly_P"],
                colors=[C_ACC] * 24, ylabel="P (MW)", fmt="%.1f")


def p1_fig4():  # 各时段区段电流分布（分组柱）
    edges = [0, 6, 10, 16, 24]
    labs = ["夜间", "早高峰", "午间", "晚高峰"]
    gA = bucket_avg(s["I_A"], edges); gB = bucket_avg(s["I_B"], edges); gC = bucket_avg(s["I_C"], edges)
    return _grouped_bar("图4 各时段三相供电臂平均电流（A）", labs, ["A臂", "B臂", "C臂"],
                        [[gA[i], gB[i], gC[i]] for i in range(4)], fmt="%.0f")


def p1_fig5():  # 三相供电侧电流
    return _line("图5 供电侧三相电流 i_a/i_b/i_c 时程",
                 [("i_a", C_ACC, list(zip(th, s["ia"]))),
                  ("i_b", C_RED, list(zip(th, s["ib"]))),
                  ("i_c", C_GREEN, list(zip(th, s["ic"])))],
                 xlabel="时刻 h", ylabel="i (A)")


def p1_fig6():  # 对称分量
    return _line("图6 电流对称分量（正/负/零序幅值）",
                 [("正序", C_ACC, list(zip(th, s["seq1"]))),
                  ("负序", C_RED, list(zip(th, s["seq2"]))),
                  ("零序", C_AMB, list(zip(th, s["seq0"])))],
                 xlabel="时刻 h", ylabel="幅值 (A)")


def p1_fig7():  # 电流不平衡度
    return _line("图7 电流不平衡度（负序/正序×100%）时程",
                 [("不平衡度", C_PUR, list(zip(th, s["unb"])))],
                 xlabel="时刻 h", ylabel="不平衡度 %", ymax=max(s["unb"]) * 1.15)


def p1_fig8():  # 功率因数
    return _line("图8 功率因数 cosφ 时程",
                 [("cosφ", C_TEAL, list(zip(th, s["cos_phi"])))],
                 xlabel="时刻 h", ylabel="cosφ", ymin=0.78, ymax=0.98)


# ============ Paper 2：电能质量（谐波 / 不平衡 / 功率因数） ============
AMP = {
    "空载": [(1, 80), (3, 12), (5, 6), (7, 3), (9, 1.5), (11, 1.0)],
    "牵引": [(1, 1500), (3, 380), (5, 210), (7, 120), (9, 70), (11, 45), (13, 25)],
    "制动": [(1, 1200), (3, 300), (5, 170), (7, 95), (9, 55), (11, 35), (13, 20)],
}


def p2_fig1():  # 牵引电流波形
    rng = random.Random(2021)
    y = G._synth_wave(AMP["牵引"], rng, noise=0.004 * 1500)
    step = max(1, len(y) // 220)
    pts = [(k / len(y), y[k]) for k in range(0, len(y), step)]
    return _line("图1 牵引工况电流波形（一个工频周期采样）",
                 [("i(t)", C_ACC, pts)], xlabel="归一化相位", ylabel="i (A)", ymin=-1800, ymax=1800)


def p2_fig2():  # 牵引电流频谱
    hs = [h for h, _ in AMP["牵引"]]
    amps = [a for _, a in AMP["牵引"]]
    return _bar("图2 牵引工况电流谐波幅值频谱", [("h%d" % h) for h in hs], amps,
                colors=[C_RED] * len(hs), ylabel="幅值 (A)", fmt="%.0f")


def p2_fig3():  # THD 对比
    conds = ["空载", "牵引", "制动"]
    return _bar("图3 三种工况电流总谐波畸变率 THD", conds,
                [D["thd"][c] for c in conds], colors=[C_AMB, C_RED, C_GREEN],
                ylabel="THD %", fmt="%.1f")


def p2_fig4():  # 主导谐波对比（分组柱）
    hs = [3, 5, 7, 9, 11]
    def top5(c):
        d = {h: a for h, a in AMP[c]}
        return [d.get(h, 0) for h in hs]
    cats = ["h3", "h5", "h7", "h9", "h11"]
    return _grouped_bar("图4 各工况前 5 次谐波幅值对比 (A)", ["空载", "牵引", "制动"],
                        cats, [top5(c) for c in ["空载", "牵引", "制动"]], fmt="%.0f")


def p2_fig5():  # 负序 vs 正序（不平衡随风荷载增长）
    pts = [(s["seq1"][t], s["seq2"][t]) for t in range(len(th))]
    return _scatter("图5 负序电流幅值 vs 正序电流幅值（斜率=不平衡度）",
                    [("采样点", C_PUR, pts)], xlabel="正序幅值 (A)", ylabel="负序幅值 (A)")


def p2_fig6():  # 功率因数分布
    bins = [(0.80, 0.84), (0.84, 0.88), (0.88, 0.92), (0.92, 0.96)]
    cnt = [0] * len(bins)
    for v in s["cos_phi"]:
        for i, (a, b) in enumerate(bins):
            if a <= v < b:
                cnt[i] += 1
    return _bar("图6 功率因数分布直方图",
                ["0.80-0.84", "0.84-0.88", "0.88-0.92", "0.92-0.96"], cnt,
                colors=[C_TEAL] * 4, ylabel="频数", fmt="%d")


def p2_fig7():  # 谐波阻抗幅值随阶次
    R, X = D["Zh_R"], D["Zh_X"]
    hs = list(range(1, 14))
    def zmag(RR, XX):
        return [math.sqrt(RR ** 2 + (h * XX) ** 2) for h in hs]
    return _line("图7 谐波阻抗幅值 |Z_h|=√(R²+(hX)²) 随谐波阶次",
                 [("拟合(本数据)", C_ACC, list(zip(hs, zmag(R, X)))),
                  ("理论(1.5+j4.5)", C_RED, list(zip(hs, zmag(1.5, 4.5))))],
                 xlabel="谐波阶次 h", ylabel="|Z_h| (Ω)")


def p2_fig8():  # 各小时最大不平衡度
    edges = list(range(0, 25, 2)) + [24]
    maxu = []
    for i in range(len(edges) - 1):
        a, b = edges[i], edges[i + 1]
        vals = [s["unb"][t] for t in range(len(th)) if a <= th[t] < b]
        maxu.append(max(vals) if vals else 0.0)
    return _line("图8 各时段最大电流不平衡度（夜低昼高）",
                 [("最大不平衡度", C_PUR, list(zip(edges[:-1], maxu)))],
                 xlabel="时刻 h", ylabel="最大不平衡度 %", ymax=max(maxu) * 1.15)


# ============ Paper 3：等值建模与异常检测 ============
def p3_fig1():  # (P,Q) 散点
    pts = [(s["P"][t], s["Q"][t]) for t in range(len(th))]
    return _scatter("图1 有功-无功 (P,Q) 散点（等值建模输入）",
                    [("运行点", C_ACC, pts)], xlabel="P (MW)", ylabel="Q (MVar)")


def p3_fig2():  # (U,I) 散点 + 戴维南拟合线
    pts = [(s["Itot_clean"][t], s["U_clean"][t]) for t in range(len(th))]
    E0 = complex(*D["E0_fit_kV"]); Z = complex(*D["Zeq_fit_Ohm"])
    cphi = math.acos(D["pf_mean"]); ph = -cphi
    Imax = max(s["Itot_clean"]) * 1.05
    line = []
    for I in range(0, int(Imax) + 1, max(1, int(Imax) // 40)):
        Ic = (I / 1000.0) * complex(math.cos(ph), math.sin(ph))
        line.append((I, abs(E0 - Z * Ic)))
    return _scatter("图2 电压-电流 (U,I) 散点与戴维南拟合曲线",
                    [("实测点", C_ACC, pts), ("拟合 U=|E0-Z·I|", C_RED, line)],
                    xlabel="I (A)", ylabel="U (kV)")


def p3_fig3():  # 测量 vs 预测 U
    E0 = complex(*D["E0_fit_kV"]); Z = complex(*D["Zeq_fit_Ohm"])
    cphi = math.acos(D["pf_mean"]); ph = -cphi
    pts = []
    for t in range(len(th)):
        Ic = (s["Itot_clean"][t] / 1000.0) * complex(math.cos(ph), math.sin(ph))
        Up = abs(E0 - Z * Ic)
        pts.append((s["U_clean"][t], Up))
    umax = max(max(p) for p in pts) * 1.05
    ref = [(0, 0), (umax, umax)]
    return _scatter("图3 戴维南模型预测电压 vs 实测电压（贴合 y=x）",
                    [("点", C_GREEN, pts), ("y=x", C_RED, ref)],
                    xlabel="实测 U (kV)", ylabel="预测 U (kV)", xmin=0, xmax=umax, ymin=0, ymax=umax)


def p3_fig4():  # 残差
    return _line("图4 戴维南模型残差时程（突变量=异常）",
                 [("残差", C_AMB, list(zip(th, s["res"])))],
                 xlabel="时刻 h", ylabel="残差 (kV)", ymin=-5, ymax=5)


def p3_fig5():  # U 时程 + 检出异常
    base = list(zip(th, s["U_clean"]))
    det = [(th[t], s["U_clean"][t]) for t in range(len(th)) if s["detected"][t]]
    return _line("图5 馈线电压时程与检出的异常点",
                 [("U(t)", C_ACC, base), ("检出异常", C_RED, det)],
                 xlabel="时刻 h", ylabel="U (kV)")


def p3_fig6():  # dU/dt
    d = [s["U_clean"][t] - s["U_clean"][t - 1] for t in range(1, len(th))]
    xd = th[1:]
    thr = 1.0
    line = [(xd[0], thr), (xd[-1], thr)]
    return _line("图6 电压变化率 dU/dt 时程（阈值 ±1 kV/步）",
                 [("dU/dt", C_PUR, list(zip(xd, d))),
                  ("+阈值", C_RED, line), ("-阈值", C_RED, [(xd[0], -thr), (xd[-1], -thr)])],
                 xlabel="时刻 h", ylabel="dU/dt (kV/步)", ymin=-5, ymax=5)


def p3_fig7():  # 谐波阻抗拟合
    R, X = D["Zh_R"], D["Zh_X"]
    hs = list(range(1, 14))
    def zmag(RR, XX):
        return [math.sqrt(RR ** 2 + (h * XX) ** 2) for h in hs]
    return _line("图7 谐波阻抗模型拟合（|Z_h| 随阶次）",
                 [("拟合 R=%.2f X=%.2f" % (R, X), C_ACC, list(zip(hs, zmag(R, X)))),
                  ("理论", C_RED, list(zip(hs, zmag(1.5, 4.5))))],
                 xlabel="谐波阶次 h", ylabel="|Z_h| (Ω)")


def p3_fig8():  # 方法流程
    return _flow("图8 牵引供电分析-等值建模-异常检测方法流程",
                 [("数据", "采集/合成"), ("清洗", "插值/裁剪"),
                  ("特征", "负荷/序分量"), ("电能质量", "THD/不平衡"),
                  ("等值", "戴维南/谐波Z"), ("异常", "残差检测")])


JOBS = [
    p1_fig1, p1_fig2, p1_fig3, p1_fig4, p1_fig5, p1_fig6, p1_fig7, p1_fig8,
    p2_fig1, p2_fig2, p2_fig3, p2_fig4, p2_fig5, p2_fig6, p2_fig7, p2_fig8,
    p3_fig1, p3_fig2, p3_fig3, p3_fig4, p3_fig5, p3_fig6, p3_fig7, p3_fig8,
]


def main():
    n_bad = 0
    for i, job in enumerate(JOBS):
        paper = (i // 8) + 1
        fig = (i % 8) + 1
        name = "dgcup2021a-%d-fig%d.svg" % (paper, fig)
        try:
            svg = job()
            if svg and "<svg" in svg:
                save(name, svg)
            else:
                n_bad += 1
                print("BAD(empty):", name)
        except Exception as e:
            n_bad += 1
            print("ERR", name, e)
    print("生成 SVG %d 张，失败 %d 张" % (len(JOBS), n_bad))


if __name__ == "__main__":
    main()
