#!/usr/bin/env python3
"""Concurrency sweep to find the clone throughput knee and a usable worker range.

Controls:
  - Only L+XL leftover repos (bandwidth-bound, not handshake-bound, not XXL stragglers)
  - Size-matched round-robin so every worker level sees a similar size mix
  - Unique repos per cell (no GitHub/CDN repeat-cache)
  - Randomized level order
  - Two tracks: clone-only (download ceiling) and clone+in-memory-tar.gz (pipeline)

Writes speed_test/reports/concurrency_sweep_report.md
"""
from __future__ import annotations

import json
import os
import random
import statistics
import sys
import threading
import time
from collections import defaultdict
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))
import clone_speed_test as cst

SEED = 20260911
ROOT = Path("/www/wwwroot/dataset/github_clone/speed_test")
REPORTS = ROOT / "reports"
LOGS = ROOT / "logs"
PREV_METRICS = REPORTS / "metrics.jsonl"
PREV_SAMPLE = REPORTS / "sample.json"
MANIFEST = Path("/www/wwwroot/dataset/github_clone/star30_repos_not_cloned_part013.jsonl")

E2E_LEVELS = [1, 2, 4, 6, 8, 12, 16, 20]
CLONE_ONLY_LEVELS = [4, 8, 12, 16]
N_PER = 8
N_CLONE_ONLY = 6
WARMUP_N = 2

cst.ARCHIVES = ROOT / "concurrency_archives"
cst.METRICS_PATH = REPORTS / "concurrency_metrics.jsonl"
cst.LOG_PATH = LOGS / "concurrency_run.log"
cst.WORK = Path("/dev/shm/clone_conc_work")
cst.DISK_WORK = ROOT / "concurrency_work"
SUMMARY_PATH = REPORTS / "concurrency_summary.json"
REPORT_PATH = REPORTS / "concurrency_sweep_report.md"
PLAN_PATH = REPORTS / "concurrency_plan.json"


def default_iface() -> str:
    cp = cst.run(["ip", "-4", "route", "show", "default"])
    for tok in (cp.stdout or "").split():
        pass
    parts = (cp.stdout or "").split()
    if "dev" in parts:
        return parts[parts.index("dev") + 1]
    return "eth0"


def read_cpu() -> tuple[int, int]:
    with open("/proc/stat", encoding="utf-8") as fh:
        vals = [int(x) for x in fh.readline().split()[1:]]
    idle = vals[3] + (vals[4] if len(vals) > 4 else 0)
    return idle, sum(vals)


def read_mem_avail() -> int:
    with open("/proc/meminfo", encoding="utf-8") as fh:
        for line in fh:
            if line.startswith("MemAvailable:"):
                return int(line.split()[1]) * 1024
    return 0


def read_net(iface: str) -> tuple[int, int]:
    with open("/proc/net/dev", encoding="utf-8") as fh:
        for line in fh:
            if line.strip().startswith(iface + ":"):
                cols = line.split(":", 1)[1].split()
                return int(cols[0]), int(cols[8])
    return 0, 0


class SysSampler:
    def __init__(self, iface: str, interval: float = 0.5):
        self.iface = iface
        self.interval = interval
        self._stop = threading.Event()
        self._thr = None
        self.samples: list[dict] = []

    def start(self) -> None:
        self._stop.clear()
        self.samples = []
        self._thr = threading.Thread(target=self._run, daemon=True)
        self._thr.start()

    def _run(self) -> None:
        prev_idle, prev_total = read_cpu()
        prev_rx, prev_tx = read_net(self.iface)
        prev_t = time.perf_counter()
        time.sleep(self.interval)
        while not self._stop.is_set():
            idle, total = read_cpu()
            rx, tx = read_net(self.iface)
            now = time.perf_counter()
            dt = max(now - prev_t, 1e-6)
            d_total = max(total - prev_total, 1)
            cpu = 100.0 * (1.0 - (idle - prev_idle) / d_total)
            self.samples.append(
                {
                    "t": now,
                    "cpu_pct": cpu,
                    "rx_bps": (rx - prev_rx) / dt,
                    "tx_bps": (tx - prev_tx) / dt,
                    "mem_avail": read_mem_avail(),
                }
            )
            prev_idle, prev_total, prev_rx, prev_tx, prev_t = idle, total, rx, tx, now
            self._stop.wait(self.interval)

    def stop(self) -> dict:
        self._stop.set()
        if self._thr:
            self._thr.join(timeout=2)
        xs = self.samples
        if not xs:
            return {"n": 0}
        cpu = [s["cpu_pct"] for s in xs]
        rx = [s["rx_bps"] for s in xs]
        return {
            "n": len(xs),
            "iface": self.iface,
            "cpu_mean": statistics.fmean(cpu),
            "cpu_max": max(cpu),
            "cpu_p95": cst.percentile(cpu, 95),
            "rx_mean_bps": statistics.fmean(rx),
            "rx_max_bps": max(rx),
            "rx_p95_bps": cst.percentile(rx, 95),
            "mem_avail_min": min(s["mem_avail"] for s in xs),
        }


def used_names() -> set[str]:
    used: set[str] = set()
    if PREV_METRICS.exists():
        for line in PREV_METRICS.read_text(encoding="utf-8").splitlines():
            if line.strip():
                used.add(json.loads(line).get("full_name"))
    if PREV_SAMPLE.exists():
        sample = json.loads(PREV_SAMPLE.read_text(encoding="utf-8"))
        for names in (sample.get("sample_meta") or {}).get("phases", {}).values():
            used.update(names)
    return {u for u in used if u}


def leftover_lx() -> list[dict]:
    sample = json.loads(PREV_SAMPLE.read_text(encoding="utf-8"))
    used = used_names()
    out = []
    for r in sample["api_rows"]:
        fn = r.get("full_name")
        if not fn or fn in used:
            continue
        st = cst.stratum_of(r.get("size_kib"))
        if st not in ("L", "XL"):
            continue
        owner, repo = fn.split("/", 1)
        out.append(
            {
                "owner": owner,
                "repo": repo,
                "full_name": fn,
                "stratum": st,
                "size_kib": r.get("size_kib"),
                "stars": r.get("stargazers_count"),
                "language": r.get("language"),
                "default_branch": r.get("default_branch"),
                "mode": "shallow",
            }
        )
    out.sort(key=lambda x: x["size_kib"] or 0)
    return out


def leftover_m_warmup(n: int, already: set[str]) -> list[dict]:
    sample = json.loads(PREV_SAMPLE.read_text(encoding="utf-8"))
    used = used_names() | already
    pool = []
    for r in sample["api_rows"]:
        fn = r.get("full_name")
        if not fn or fn in used:
            continue
        if cst.stratum_of(r.get("size_kib")) != "M":
            continue
        owner, repo = fn.split("/", 1)
        pool.append(
            {
                "owner": owner,
                "repo": repo,
                "full_name": fn,
                "stratum": "M",
                "size_kib": r.get("size_kib"),
                "stars": r.get("stargazers_count"),
                "language": r.get("language"),
                "default_branch": r.get("default_branch"),
                "mode": "shallow",
                "phase": "warmup",
                "pack": True,
            }
        )
    rng = random.Random(SEED)
    rng.shuffle(pool)
    return pool[:n]


def assign_round_robin(sorted_repos: list[dict], n_groups: int, n_each: int) -> tuple[list[list[dict]], list[dict]]:
    groups = [[] for _ in range(n_groups)]
    for i, rec in enumerate(sorted_repos):
        groups[i % n_groups].append(rec)
    chosen = []
    rest = []
    for g in groups:
        chosen.append(g[:n_each])
        rest.extend(g[n_each:])
    return chosen, rest


def phase_stats(results: list[dict], wall_s: float, workers: int, sysinfo: dict, pack: bool) -> dict:
    st = cst.stats_block(results, f"w{workers}")
    payload = st.get("sum_payload_bytes") or 0
    sum_clone = st.get("sum_clone_s") or 0
    sum_total = st.get("sum_total_s") or 0
    wall_bps = payload / wall_s if wall_s and payload else None
    efficiency = (sum_total / wall_s / workers) if wall_s and workers and sum_total else None
    clone_eff = (sum_clone / wall_s / workers) if wall_s and workers and sum_clone else None
    sizes = [r.get("size_kib_api") or 0 for r in results]
    return {
        "workers": workers,
        "pack": pack,
        "wall_s": wall_s,
        "wall_payload_bps": wall_bps,
        "parallel_efficiency_e2e": efficiency,
        "parallel_efficiency_clone": clone_eff,
        "speedup_e2e": (sum_total / wall_s) if wall_s and sum_total else None,
        "mean_size_kib": statistics.fmean(sizes) if sizes else None,
        "median_size_kib": statistics.median(sizes) if sizes else None,
        "sys": sysinfo,
        "stats": st,
        "repos": [r.get("full_name") for r in results],
    }


def run_level(workers: int, tasks: list[dict], pack: bool, iface: str) -> dict:
    kind = "e2e" if pack else "clone_only"
    phase = f"{kind}_w{workers}"
    for t in tasks:
        t["phase"] = phase
        t["pack"] = pack
        t["mode"] = "shallow"
    cst.log(f"=== {phase} n={len(tasks)} workers={workers} pack={pack} ===")
    sampler = SysSampler(iface)
    sampler.start()
    t0 = time.perf_counter()
    results = []
    if workers <= 1:
        for t in tasks:
            results.append(cst.clone_and_pack(t))
    else:
        import concurrent.futures

        with concurrent.futures.ThreadPoolExecutor(max_workers=workers) as ex:
            futs = [ex.submit(cst.clone_and_pack, t) for t in tasks]
            for fut in concurrent.futures.as_completed(futs):
                results.append(fut.result())
    wall = time.perf_counter() - t0
    sysinfo = sampler.stop()
    rec = phase_stats(results, wall, workers, sysinfo, pack)
    rec["results"] = results
    ok = rec["stats"]["n_ok"]
    cst.log(
        f"=== {phase} done wall={cst.fmt_s(wall)} ok={ok}/{len(results)} "
        f"wall_bps={cst.fmt_bps(rec['wall_payload_bps'])} "
        f"agg_clone={cst.fmt_bps(rec['stats'].get('aggregate_clone_bps'))} "
        f"cpu={rec['sys'].get('cpu_mean')} rx={cst.fmt_bps(rec['sys'].get('rx_mean_bps'))} ==="
    )
    return rec


def pick_range(rows: list[dict]) -> dict:
    ok = [r for r in rows if (r["stats"].get("success_rate") or 0) >= 0.999]
    if not ok:
        ok = rows[:]
    peak = max(ok, key=lambda r: r.get("wall_payload_bps") or 0)
    peak_bps = peak.get("wall_payload_bps") or 0
    seq = next((r for r in rows if r["workers"] == 1), None)
    seq_med = None
    if seq and seq["stats"].get("clone_bps"):
        seq_med = seq["stats"]["clone_bps"].get("p50")
    knee = None
    for r in sorted(ok, key=lambda x: x["workers"]):
        bps = r.get("wall_payload_bps") or 0
        if peak_bps and bps >= 0.90 * peak_bps:
            knee = r
            break
    upper = peak
    for r in sorted(ok, key=lambda x: x["workers"]):
        if r["workers"] < peak["workers"]:
            continue
        bps = r.get("wall_payload_bps") or 0
        drop = (peak_bps - bps) / peak_bps if peak_bps else 0
        med = (r["stats"].get("clone_bps") or {}).get("p50")
        conn_drop = 0
        if seq_med and med:
            conn_drop = 1 - (med / seq_med)
        eff = r.get("parallel_efficiency_clone")
        fails = r["stats"].get("n_fail") or 0
        if fails or drop > 0.12 or (eff is not None and eff < 0.35) or conn_drop > 0.35:
            if r["workers"] != peak["workers"]:
                break
        upper = r
    # recommended: knee through upper, prefer smallest that is within 8% of peak if efficiency better
    band = [r for r in ok if knee and upper and knee["workers"] <= r["workers"] <= upper["workers"]]
    rec = None
    if band:
        # score: high wall bps, high efficiency, fewer workers
        def score(r):
            bps = (r.get("wall_payload_bps") or 0) / (peak_bps or 1)
            eff = r.get("parallel_efficiency_clone") or 0
            return bps * 0.6 + eff * 0.3 - r["workers"] / 200.0

        rec = max(band, key=score)
    return {
        "peak_workers": peak["workers"],
        "peak_wall_bps": peak_bps,
        "knee_workers": knee["workers"] if knee else None,
        "upper_workers": upper["workers"] if upper else None,
        "recommended_workers": rec["workers"] if rec else None,
        "band": [r["workers"] for r in band] if band else [],
    }


def write_report(plan: dict, warmup: dict, e2e: list[dict], clone_only: list[dict], iface: str) -> None:
    lines = []
    a = lines.append
    a("# Git Clone 并发极限与调优区间")
    a("")
    a(f"- 生成时间（UTC）：`{cst.utc_now()}`")
    a(f"- 种子：`{SEED}`")
    a(f"- 网卡采样：`{iface}`")
    a("- 样本：上一轮 API 抽样中尚未 clone 的 **L + XL**（10–200 MiB GitHub size）")
    a("- 分配：按 size 排序后 round-robin，各并发档体积分布接近")
    a("- 仓库互斥，随机打乱档位顺序")
    a("")
    a("## 1. 结论")
    a("")
    e2e_pick = pick_range(e2e)
    co_pick = pick_range(clone_only) if clone_only else {}
    a("这轮专门回答两件事：**下载能被并发推到多快**，以及 **生产流水线该用几路**。")
    a("")
    if co_pick:
        a(
            f"- **纯 clone 墙钟峰值**：并发 {co_pick.get('peak_workers')} → "
            f"{cst.fmt_bps(co_pick.get('peak_wall_bps'))}"
        )
        a(
            f"- **纯 clone 调优区间**：{co_pick.get('knee_workers')}–{co_pick.get('upper_workers')} "
            f"（推荐 {co_pick.get('recommended_workers')}）"
        )
    a(
        f"- **clone + 内存 tar.gz 落盘峰值**：并发 {e2e_pick.get('peak_workers')} → "
        f"{cst.fmt_bps(e2e_pick.get('peak_wall_bps'))}"
    )
    a(
        f"- **生产流水线调优区间**：{e2e_pick.get('knee_workers')}–{e2e_pick.get('upper_workers')} "
        f"（推荐 {e2e_pick.get('recommended_workers')}）"
    )
    a("")
    a("区间含义：左端是达到峰值 90% 的最小并发（knee），右端是吞吐掉 12%、单连接变慢 35%、并行效率 <0.35 或开始失败之前的最后一档。推荐值在区间内综合吞吐和效率。")
    a("")
    a("## 2. 方法")
    a("")
    a("上一轮 1/4/8 不能当极限：8 路被 1.97 GiB 的 XXL 拖死，且小仓握手开销会压低平均值。本轮：")
    a("")
    a("1. 只用 L/XL，让测量进入带宽区而不是 RTT 区。")
    a("2. 每档固定 8 个仓（clone-only 每档 6 个），体积 round-robin 对齐。")
    a("3. 两套轨道：只 clone（下载天花板）vs clone+内存打包落盘（真实流水线）。")
    a("4. 记录 CPU%、网卡 RX、并行效率 = Σ任务时间 / (墙钟 × 并发)。")
    a("")
    a("## 3. 各档体积对齐")
    a("")
    a("| 轨道 | 并发 | n | size 中位 KiB | size 均 KiB | L/XL |")
    a("|---|---:|---:|---:|---:|---|")
    for rec in e2e + clone_only:
        rows = rec.get("results") or []
        lx = f"{sum(1 for r in rows if r.get('stratum')=='L')}/{sum(1 for r in rows if r.get('stratum')=='XL')}"
        a(
            f"| {'e2e' if rec['pack'] else 'clone-only'} | {rec['workers']} | {rec['stats']['n']} | "
            f"{rec.get('median_size_kib'):.0f} | {rec.get('mean_size_kib'):.0f} | {lx} |"
        )
    a("")
    a("## 4. 纯 clone（下载极限）")
    a("")
    if clone_only:
        a("| 并发 | 成功 | 墙钟 | 墙钟吞吐 | 聚合 clone 吞吐 | 单仓吞吐中位 | clone 并行效率 | CPU均/峰值 | 网卡 RX 均 |")
        a("|---:|---:|---:|---:|---:|---:|---:|---:|---:|")
        for rec in sorted(clone_only, key=lambda r: r["workers"]):
            st = rec["stats"]
            sysinfo = rec.get("sys") or {}
            a(
                f"| {rec['workers']} | {st['n_ok']}/{st['n']} | {cst.fmt_s(rec['wall_s'])} | "
                f"{cst.fmt_bps(rec.get('wall_payload_bps'))} | {cst.fmt_bps(st.get('aggregate_clone_bps'))} | "
                f"{cst.fmt_bps((st.get('clone_bps') or {}).get('p50'))} | "
                f"{rec.get('parallel_efficiency_clone'):.2f} | "
                f"{sysinfo.get('cpu_mean', 0):.0f}% / {sysinfo.get('cpu_max', 0):.0f}% | "
                f"{cst.fmt_bps(sysinfo.get('rx_mean_bps'))} |"
            )
        a("")
        fails = [f for rec in clone_only for f in (rec["stats"].get("failures") or [])]
        if fails:
            a("失败：")
            for f in fails:
                a(f"- `{f.get('full_name')}` · {f.get('error_class')} · {(f.get('error') or '')[:180]}")
            a("")
    else:
        a("（本轮未跑 clone-only）")
        a("")
    a("## 5. clone + 内存 tar.gz 落盘（生产流水线）")
    a("")
    a("| 并发 | 成功 | 墙钟 | 墙钟吞吐 | 聚合 clone 吞吐 | 单仓吞吐中位 | e2e 并行效率 | 打包中位 | CPU均/峰值 | 网卡 RX 均 |")
    a("|---:|---:|---:|---:|---:|---:|---:|---:|---:|---:|")
    for rec in sorted(e2e, key=lambda r: r["workers"]):
        st = rec["stats"]
        sysinfo = rec.get("sys") or {}
        a(
            f"| {rec['workers']} | {st['n_ok']}/{st['n']} | {cst.fmt_s(rec['wall_s'])} | "
            f"{cst.fmt_bps(rec.get('wall_payload_bps'))} | {cst.fmt_bps(st.get('aggregate_clone_bps'))} | "
            f"{cst.fmt_bps((st.get('clone_bps') or {}).get('p50'))} | "
            f"{(rec.get('parallel_efficiency_e2e') or 0):.2f} | "
            f"{cst.fmt_s((st.get('pack_s') or {}).get('p50'))} | "
            f"{sysinfo.get('cpu_mean', 0):.0f}% / {sysinfo.get('cpu_max', 0):.0f}% | "
            f"{cst.fmt_bps(sysinfo.get('rx_mean_bps'))} |"
        )
    a("")
    a("墙钟吞吐 = Σ clone_bytes / 阶段墙钟，这才是「开 N 路时机器实际搬数据的速度」。")
    a("单仓吞吐中位下降 = 连接开始互相抢带宽或 CPU。并行效率 1.0 表示 N 路没有互相等待。")
    a("")
    a("## 6. 怎么读调优区间")
    a("")
    a("- **knee**：再加并发几乎不再涨吞吐，再少则明显变慢。")
    a("- **推荐值**：区间内吞吐高、效率不差、CPU/网卡都还没打满到掉速。")
    a("- **不要超过右端**：再加线程只会让 gzip 和 git 解包抢 CPU，墙钟变差，还更容易 429。")
    a("- 全量 5.4 万仓仍应按 size 分队列：本区间适用于 L/XL；XS/S 可以更高并发（它们受握手限制），XXL 必须更低并发。")
    a("")
    a("## 7. 产物")
    a("")
    a(f"- 报告：`{REPORT_PATH}`")
    a(f"- JSON：`{SUMMARY_PATH}`")
    a(f"- 每仓指标：`{cst.METRICS_PATH}`")
    a(f"- 归档（仅 e2e）：`{cst.ARCHIVES}`")
    a("")
    REPORT_PATH.write_text("\n".join(lines) + "\n", encoding="utf-8")


def compact(rec: dict) -> dict:
    out = {k: v for k, v in rec.items() if k != "results"}
    st = out.get("stats") or {}
    if "failures" in st:
        pass
    rows = rec.get("results") or []
    keys = [
        "ok", "phase", "full_name", "stratum", "size_kib_api", "clone_s", "pack_s",
        "write_s", "total_s", "clone_bytes", "archive_bytes", "clone_bps",
        "error", "error_class",
    ]
    out["results"] = [{k: r.get(k) for k in keys} for r in rows]
    return out


def main() -> int:
    for p in (cst.ARCHIVES, cst.WORK, cst.DISK_WORK, REPORTS, LOGS):
        p.mkdir(parents=True, exist_ok=True)
    if cst.METRICS_PATH.exists():
        cst.METRICS_PATH.unlink()
    tok = cst.token()
    iface = default_iface()
    cst.log(f"iface={iface} login probe")
    code, body, _hdrs = cst.gh_api("https://api.github.com/user", tok)
    cst.log(f"github user http={code} login={(body or {}).get('login') if isinstance(body, dict) else None}")

    pool = leftover_lx()
    cst.log(f"leftover L+XL n={len(pool)}")
    n_e2e = len(E2E_LEVELS) * N_PER
    n_co = len(CLONE_ONLY_LEVELS) * N_CLONE_ONLY
    if len(pool) < n_e2e + n_co:
        raise SystemExit(f"not enough L+XL leftovers: have {len(pool)} need {n_e2e + n_co}")

    e2e_groups, rest = assign_round_robin(pool, len(E2E_LEVELS), N_PER)
    co_groups, leftover = assign_round_robin(rest, len(CLONE_ONLY_LEVELS), N_CLONE_ONLY)
    already = {t["full_name"] for g in e2e_groups + co_groups for t in g}
    warmup_tasks = leftover_m_warmup(WARMUP_N, already)

    rng = random.Random(SEED + 7)
    e2e_plan = list(zip(E2E_LEVELS, e2e_groups))
    co_plan = list(zip(CLONE_ONLY_LEVELS, co_groups))
    rng.shuffle(e2e_plan)
    rng.shuffle(co_plan)

    plan = {
        "iface": iface,
        "e2e_order": [w for w, _ in e2e_plan],
        "clone_only_order": [w for w, _ in co_plan],
        "e2e": {str(w): [t["full_name"] for t in g] for w, g in zip(E2E_LEVELS, e2e_groups)},
        "clone_only": {str(w): [t["full_name"] for t in g] for w, g in zip(CLONE_ONLY_LEVELS, co_groups)},
        "warmup": [t["full_name"] for t in warmup_tasks],
        "n_pool": len(pool),
        "leftover_unused": [t["full_name"] for t in leftover],
    }
    PLAN_PATH.write_text(json.dumps(plan, ensure_ascii=False, indent=2), encoding="utf-8")
    cst.log(f"e2e order {plan['e2e_order']} clone-only order {plan['clone_only_order']}")

    warmup_rec = run_level(1, warmup_tasks, pack=True, iface=iface) if warmup_tasks else {}

    # clone-only first: isolate download ceiling before gzip CPU contention
    clone_only_recs = []
    for workers, tasks in co_plan:
        clone_only_recs.append(run_level(workers, tasks, pack=False, iface=iface))

    e2e_recs = []
    for workers, tasks in e2e_plan:
        e2e_recs.append(run_level(workers, tasks, pack=True, iface=iface))

    summary = {
        "status": "complete",
        "completed_at": cst.utc_now(),
        "plan": plan,
        "clone_only_pick": pick_range(clone_only_recs) if clone_only_recs else None,
        "e2e_pick": pick_range(e2e_recs),
        "warmup": compact(warmup_rec) if warmup_rec else None,
        "clone_only": [compact(r) for r in clone_only_recs],
        "e2e": [compact(r) for r in e2e_recs],
    }
    SUMMARY_PATH.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
    write_report(plan, warmup_rec, e2e_recs, clone_only_recs, iface)
    cst.log(f"report {REPORT_PATH}")
    cst.log("STATUS complete")
    return 0


if __name__ == "__main__":
    try:
        raise SystemExit(main())
    except Exception:
        cst.log("STATUS failed")
        cst.log(cst.redact(__import__("traceback").format_exc()))
        SUMMARY_PATH.write_text(
            json.dumps({"status": "failed", "completed_at": cst.utc_now(), "error": "see log"}, indent=2),
            encoding="utf-8",
        )
        raise
