#!/usr/bin/env python3
"""
ecmp-probe — 探测 / 选择 ECMP 路径

链路做 per-flow ECMP 时,同一目的地会有多条延迟不同的路,新连接按 5 元组哈希
随机落在其中一条。ICMP 通常只按 src/dst IP 哈希,所以 ping 永远只看得到固定
的那一条 —— 用 ping 评估这种链路会漏掉更快的路。

本工具用 TCP 握手(SYN→SYN/ACK)测延迟,并可指定源端口,从而:
  scan  扫一批源端口,自动识别有几条路、各自延迟、哪些端口落在快路
  pin   固定源端口反复测,验证该路径是否稳定

用法:
  ecmp-probe.py scan <host> <port> [-n 120]
  ecmp-probe.py pin  <host> <port> --sport 46001 [-c 10]

例:
  ./ecmp-probe.py scan 43.108.37.202 443 -n 200
  ./ecmp-probe.py pin  43.108.37.202 443 --sport 46001 -c 20

说明: 用 SO_LINGER=0 做 abortive close(发 RST),避免源端口卡在 TIME_WAIT
      而无法立即复用。只读探测,不发送任何应用层数据。
"""
import argparse
import socket
import statistics
import struct
import sys
import time

LINGER = struct.pack("ii", 1, 0)  # abortive close -> RST, 跳过 TIME_WAIT


def probe(host, port, sport=None, timeout=3.0):
    """返回 TCP 握手耗时(ms),失败返回 None。"""
    s = socket.socket()
    s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
    s.setsockopt(socket.SOL_SOCKET, socket.SO_LINGER, LINGER)
    s.settimeout(timeout)
    try:
        if sport is not None:
            s.bind(("0.0.0.0", sport))
        t0 = time.perf_counter()
        s.connect((host, port))
        return (time.perf_counter() - t0) * 1000
    except OSError:
        return None
    finally:
        try:
            s.close()
        except OSError:
            pass


def cluster(values, min_gap=5.0):
    """把延迟样本按空隙切成若干簇。min_gap 以内算同一条路。"""
    if not values:
        return []
    vs = sorted(values)
    groups, cur = [], [vs[0]]
    for v in vs[1:]:
        if v - cur[-1] > min_gap:
            groups.append(cur)
            cur = [v]
        else:
            cur.append(v)
    groups.append(cur)
    return groups


def cmd_scan(args):
    samples = []
    base = 40000
    print(f"探测 {args.host}:{args.port} — {args.n} 个源端口 ...", file=sys.stderr)
    sport = base
    while len(samples) < args.n and sport < base + args.n * 4:
        d = probe(args.host, args.port, sport)
        if d is not None:
            samples.append((sport, d))
        sport += 1
        time.sleep(0.03)

    if not samples:
        print("全部探测失败 — 目标不可达或端口未开放")
        return 1

    vals = [d for _, d in samples]
    groups = cluster(vals, args.gap)
    print(f"\n样本 {len(vals)}  min={min(vals):.1f}  中位={statistics.median(vals):.1f}  max={max(vals):.1f}")
    print(f"识别出 {len(groups)} 条路径 (簇间隔 >{args.gap}ms 视为不同路):\n")
    for i, g in enumerate(groups, 1):
        share = 100.0 * len(g) / len(vals)
        print(f"  路径{i}: {statistics.mean(g):6.1f}ms  "
              f"(范围 {min(g):.1f}~{max(g):.1f}, {len(g)}/{len(vals)} = {share:.0f}%)")

    if len(groups) == 1:
        print("\n→ 只有一条路径,这条链路没有可选的 ECMP 分支。")
        return 0

    fastest = groups[0]
    hi = max(fastest)
    fast_ports = [p for p, d in samples if d <= hi]
    print(f"\n→ 最快那条约 {statistics.mean(fastest):.1f}ms,比最慢的快 "
          f"{statistics.mean(groups[-1]) - statistics.mean(fastest):.1f}ms")
    print(f"→ 落在最快路径的源端口(前20个): {fast_ports[:20]}")
    print(f"\n用 pin 验证其中一个是否稳定:")
    print(f"  {sys.argv[0]} pin {args.host} {args.port} --sport {fast_ports[0]} -c 20")
    return 0


def cmd_pin(args):
    print(f"固定源端口 {args.sport} 探测 {args.host}:{args.port} × {args.c} 次 ...\n", file=sys.stderr)
    ds = []
    for _ in range(args.c):
        d = probe(args.host, args.port, args.sport)
        if d is not None:
            ds.append(d)
        time.sleep(0.3)
    if not ds:
        print("全部失败 — 换个源端口重试")
        return 1
    spread = max(ds) - min(ds)
    print("  " + " ".join(f"{d:.1f}" for d in ds))
    print(f"\n  成功 {len(ds)}/{args.c}  min={min(ds):.1f} 均值={statistics.mean(ds):.1f} max={max(ds):.1f}")
    print(f"  极差 {spread:.1f}ms → " +
          ("路径稳定,该源端口可用来钉住这条路 ✓" if spread < 5
           else "跳变,说明这个 5 元组没有稳定落在同一条路 ✗"))
    return 0


def main():
    ap = argparse.ArgumentParser(description="探测/选择 ECMP 路径(TCP 握手测延迟)")
    sub = ap.add_subparsers(dest="cmd", required=True)

    s = sub.add_parser("scan", help="扫源端口,识别有几条路")
    s.add_argument("host")
    s.add_argument("port", type=int)
    s.add_argument("-n", type=int, default=120, help="样本数(默认120)")
    s.add_argument("--gap", type=float, default=5.0, help="簇间隔阈值ms(默认5)")
    s.set_defaults(func=cmd_scan)

    p = sub.add_parser("pin", help="固定源端口,验证路径是否稳定")
    p.add_argument("host")
    p.add_argument("port", type=int)
    p.add_argument("--sport", type=int, required=True)
    p.add_argument("-c", type=int, default=10, help="重复次数(默认10)")
    p.set_defaults(func=cmd_pin)

    args = ap.parse_args()
    sys.exit(args.func(args))


if __name__ == "__main__":
    main()
