Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4f78b8c15e | ||
|
|
3db891bdcf | ||
|
|
560a641809 | ||
|
|
7a0ec8ded9 | ||
|
|
c174a2ed3a | ||
|
|
4dad7fbbbe |
+14
-9
@@ -28,6 +28,9 @@ import (
|
||||
const (
|
||||
defaultBackoff = 5 * time.Second
|
||||
defaultPing = 30 * time.Second
|
||||
// defaultDialTimeout 内网目标拨号上限:目标离线且防火墙丢包(不回 RST)时
|
||||
// 裸 Dial 在 Windows 上要 SYN 重传 20s+, 观看端在节点侧只能干等。
|
||||
defaultDialTimeout = 5 * time.Second
|
||||
)
|
||||
|
||||
// Logger agent 日志接口。logrus.Logger 的 Debugf/Infof/Warnf/Errorf 签名
|
||||
@@ -59,9 +62,10 @@ func (DefaultLogger) Errorf(format string, args ...any) {
|
||||
// Agent 被控端:隧道注册表。每条 Tunnel 携带自己的 NodeAddr+Token,
|
||||
// 由独立 worker(tunnelConn)连接登录各自节点;Agent 不再有集中连接与登录。
|
||||
type Agent struct {
|
||||
AgentId string
|
||||
Backoff time.Duration // 永久隧道重连退避(默认 5s)
|
||||
PingEvery time.Duration // 各隧道控制流心跳间隔(默认 30s)
|
||||
AgentId string
|
||||
Backoff time.Duration // 永久隧道重连退避(默认 5s)
|
||||
PingEvery time.Duration // 各隧道控制流心跳间隔(默认 30s)
|
||||
DialTimeout time.Duration // 内网目标拨号上限(默认 5s; 超时向节点回执失败, 观看端立即断开)
|
||||
|
||||
// Logger 日志注入点。需要纳入宿主统一日志体系时设置
|
||||
// (logrus 直接注入 logrus.StandardLogger() 即可);未设置用 DefaultLogger。
|
||||
@@ -82,12 +86,13 @@ type Agent struct {
|
||||
// New 创建被控端。节点地址与令牌随各隧道下发(Tunnel.NodeAddr/Token)。
|
||||
func New(agentId string) *Agent {
|
||||
return &Agent{
|
||||
AgentId: agentId,
|
||||
Backoff: defaultBackoff,
|
||||
PingEvery: defaultPing,
|
||||
Logger: DefaultLogger{},
|
||||
tunnels: make(map[string]*tunnelConn),
|
||||
done: make(chan struct{}),
|
||||
AgentId: agentId,
|
||||
Backoff: defaultBackoff,
|
||||
PingEvery: defaultPing,
|
||||
DialTimeout: defaultDialTimeout,
|
||||
Logger: DefaultLogger{},
|
||||
tunnels: make(map[string]*tunnelConn),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+12
-1
@@ -3,6 +3,7 @@ package agent
|
||||
import (
|
||||
"io"
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
// handleData 处理节点主动打开的数据流:
|
||||
@@ -22,7 +23,11 @@ func (a *Agent) handleData(stream net.Conn) {
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := net.Dial("tcp", t.TargetAddr())
|
||||
dialTimeout := a.DialTimeout
|
||||
if dialTimeout <= 0 { // 零值 Agent 兜底
|
||||
dialTimeout = defaultDialTimeout
|
||||
}
|
||||
conn, err := net.DialTimeout("tcp", t.TargetAddr(), dialTimeout)
|
||||
if err != nil {
|
||||
a.logger().Warnf("agent 拨号本地目标失败 target=%s err=%s", t.TargetAddr(), err.Error())
|
||||
rf := NewFrame(FrameVersion, CmdTarget, f.StreamID())
|
||||
@@ -31,6 +36,12 @@ func (a *Agent) handleData(stream net.Conn) {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
// agent↔目标段 TCP 保活: 目标假死(进程在但不响应/不发 FIN)时是
|
||||
// 全链路唯一没有存活检测的段, 不开 keepalive 观看端会无限挂死画面
|
||||
if tc, ok := conn.(*net.TCPConn); ok {
|
||||
tc.SetKeepAlive(true)
|
||||
tc.SetKeepAlivePeriod(5 * time.Second)
|
||||
}
|
||||
|
||||
rf := NewFrame(FrameVersion, CmdTarget, f.StreamID())
|
||||
_ = rf.Marshal(Ret{Code: 0, Msg: "ok"})
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
// 节点配置文件:`-config node.ini` 启用;显式命令行 flag 优先于配置项,配置项优先于内置默认。
|
||||
// 零依赖 ini:[节] 下 key = value,`#`/`;` 注释行忽略,等号两侧空白忽略。
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// nodeConfig 配置文件映射(零值表示未配置,由调用方回退 flag/默认值)
|
||||
type nodeConfig struct {
|
||||
Addr string // [server] Addr 桥接监听地址
|
||||
TunnelBind string // [server] TunnelBind 隧道监听绑定地址
|
||||
JwtSecret string // [server] JwtSecret agent 登录 JWT 校验密钥
|
||||
SweepSeconds int // [server] SweepSeconds TTL 清扫间隔(秒)
|
||||
PortRange string // [tunnel] PortRange 隧道监听端口范围(如 20000-21000)
|
||||
}
|
||||
|
||||
// loadConfigFile 解析 ini 文件;文件不存在返回错误(-config 显式指定了就应当存在)。
|
||||
func loadConfigFile(path string) (*nodeConfig, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
cfg := &nodeConfig{}
|
||||
section := ""
|
||||
sc := bufio.NewScanner(f)
|
||||
lineNo := 0
|
||||
for sc.Scan() {
|
||||
lineNo++
|
||||
line := strings.TrimSpace(sc.Text())
|
||||
if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, ";") {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(line, "[") && strings.HasSuffix(line, "]") {
|
||||
section = strings.ToLower(strings.TrimSpace(line[1 : len(line)-1]))
|
||||
continue
|
||||
}
|
||||
key, value, ok := strings.Cut(line, "=")
|
||||
if !ok {
|
||||
continue // 容错:无等号的行忽略
|
||||
}
|
||||
key = strings.ToLower(strings.TrimSpace(key))
|
||||
value = strings.TrimSpace(value)
|
||||
switch section + "." + key {
|
||||
case "server.addr":
|
||||
cfg.Addr = value
|
||||
case "server.tunnelbind":
|
||||
cfg.TunnelBind = value
|
||||
case "server.jwtsecret":
|
||||
cfg.JwtSecret = value
|
||||
case "server.sweepseconds":
|
||||
cfg.SweepSeconds, err = strconv.Atoi(value)
|
||||
if err != nil || cfg.SweepSeconds <= 0 {
|
||||
return nil, errors.New("server.sweepseconds 需为正整数, 行 " + strconv.Itoa(lineNo))
|
||||
}
|
||||
case "tunnel.portrange":
|
||||
cfg.PortRange = value
|
||||
}
|
||||
}
|
||||
return cfg, sc.Err()
|
||||
}
|
||||
|
||||
// sweepDuration 秒数转 Duration;0 表示未配置
|
||||
func (c *nodeConfig) sweepDuration() time.Duration {
|
||||
if c.SweepSeconds <= 0 {
|
||||
return 0
|
||||
}
|
||||
return time.Duration(c.SweepSeconds) * time.Second
|
||||
}
|
||||
+87
-14
@@ -2,6 +2,8 @@
|
||||
// 无管理 API、无数据库;隧道由 agent 经桥接控制流注册。
|
||||
// agent 登录用 JWT 校验(HS256);token 正式部署由 cloud 签发,
|
||||
// 手工/测试场景可用 -print-token 现场签发。
|
||||
//
|
||||
// 配置来源优先级:显式命令行 flag > -config 指定的配置文件 > 内置默认。
|
||||
package main
|
||||
|
||||
import (
|
||||
@@ -17,14 +19,18 @@ import (
|
||||
)
|
||||
|
||||
func main() {
|
||||
addr := flag.String("addr", ":5212", "桥接监听地址")
|
||||
tunnelBind := flag.String("tunnel-bind", "0.0.0.0", "隧道监听绑定地址")
|
||||
jwtSecret := flag.String("jwt-secret", "", "agent 登录 JWT 校验密钥(HS256)")
|
||||
sweep := flag.Duration("sweep", 30*time.Second, "TTL 清扫间隔")
|
||||
// 便捷签发:node -jwt-secret s3cret -print-token -agent a1 -ttl 24h
|
||||
printToken := flag.Bool("print-token", false, "用 -jwt-secret 签发一张 agent 登录 token 后退出")
|
||||
tokenAgent := flag.String("agent", "", "print-token 时写入 agentId claim")
|
||||
tokenTTL := flag.Duration("ttl", 24*time.Hour, "print-token 签发的有效期")
|
||||
var (
|
||||
config = flag.String("config", "node.ini", "配置文件路径; 默认自动加载本目录 node.ini, 文件不存在且未显式 -config 时按无配置启动")
|
||||
addr = flag.String("addr", ":5212", "桥接监听地址")
|
||||
tunnelBind = flag.String("tunnel-bind", "0.0.0.0", "隧道监听绑定地址")
|
||||
jwtSecret = flag.String("jwt-secret", "", "agent 登录 JWT 校验密钥(HS256)")
|
||||
sweep = flag.Duration("sweep", 30*time.Second, "TTL 清扫间隔")
|
||||
portRange = flag.String("port-range", "", "隧道监听端口范围(如 20000-21000); 未指定端口的隧道在范围内取空闲端口, 空=系统随机")
|
||||
// 便捷签发:node -jwt-secret s3cret -print-token -agent a1 -ttl 24h
|
||||
printToken = flag.Bool("print-token", false, "用 -jwt-secret 签发一张 agent 登录 token 后退出")
|
||||
tokenAgent = flag.String("agent", "", "print-token 时写入 agentId claim")
|
||||
tokenTTL = flag.Duration("ttl", 24*time.Hour, "print-token 签发的有效期")
|
||||
)
|
||||
flag.Parse()
|
||||
|
||||
if *printToken {
|
||||
@@ -41,16 +47,83 @@ func main() {
|
||||
return
|
||||
}
|
||||
|
||||
if *jwtSecret == "" {
|
||||
slog.Error("缺少 -jwt-secret(agent 登录 JWT 校验密钥)")
|
||||
// 终值 = 内置默认 ← 配置文件(非零项) ← 显式 flag
|
||||
explicit := map[string]bool{}
|
||||
flag.Visit(func(f *flag.Flag) { explicit[f.Name] = true })
|
||||
|
||||
// 自动加载: -config 未显式指定且默认路径不存在 → 按无配置启动(兼容纯 flag/老部署)
|
||||
if !explicit["config"] {
|
||||
if _, err := os.Stat(*config); err != nil {
|
||||
*config = ""
|
||||
}
|
||||
}
|
||||
|
||||
var (
|
||||
finalAddr = *addr
|
||||
finalTunnelBind = *tunnelBind
|
||||
finalSecret = *jwtSecret
|
||||
finalSweep = *sweep
|
||||
finalRange node.PortRange
|
||||
)
|
||||
if *config != "" {
|
||||
cfg, err := loadConfigFile(*config)
|
||||
if err != nil {
|
||||
slog.Error("读取配置文件失败", "path", *config, "err", err.Error())
|
||||
os.Exit(1)
|
||||
}
|
||||
if cfg.Addr != "" {
|
||||
finalAddr = cfg.Addr
|
||||
}
|
||||
if cfg.TunnelBind != "" {
|
||||
finalTunnelBind = cfg.TunnelBind
|
||||
}
|
||||
if cfg.JwtSecret != "" {
|
||||
finalSecret = cfg.JwtSecret
|
||||
}
|
||||
if d := cfg.sweepDuration(); d > 0 {
|
||||
finalSweep = d
|
||||
}
|
||||
if cfg.PortRange != "" {
|
||||
r, err := node.ParsePortRange(cfg.PortRange)
|
||||
if err != nil {
|
||||
slog.Error("配置文件端口范围无效", "err", err.Error())
|
||||
os.Exit(1)
|
||||
}
|
||||
finalRange = r
|
||||
}
|
||||
}
|
||||
if explicit["addr"] {
|
||||
finalAddr = *addr
|
||||
}
|
||||
if explicit["tunnel-bind"] {
|
||||
finalTunnelBind = *tunnelBind
|
||||
}
|
||||
if explicit["jwt-secret"] {
|
||||
finalSecret = *jwtSecret
|
||||
}
|
||||
if explicit["sweep"] {
|
||||
finalSweep = *sweep
|
||||
}
|
||||
if explicit["port-range"] {
|
||||
r, err := node.ParsePortRange(*portRange)
|
||||
if err != nil {
|
||||
slog.Error("端口范围无效", "err", err.Error())
|
||||
os.Exit(1)
|
||||
}
|
||||
finalRange = r
|
||||
}
|
||||
|
||||
if finalSecret == "" {
|
||||
slog.Error("缺少 jwt-secret(agent 登录 JWT 校验密钥; -jwt-secret、-config 指定文件或本目录 node.ini 的 server.jwt-secret)")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
n := node.New()
|
||||
n.JwtSecret = *jwtSecret
|
||||
n.BindTunnel = *tunnelBind
|
||||
n.SweepInterval = *sweep
|
||||
if err := n.Start(*addr); err != nil {
|
||||
n.JwtSecret = finalSecret
|
||||
n.BindTunnel = finalTunnelBind
|
||||
n.SweepInterval = finalSweep
|
||||
n.TunnelPortRange = finalRange
|
||||
if err := n.Start(finalAddr); err != nil {
|
||||
slog.Error("节点启动失败", "err", err.Error())
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
// 转发链路"慢失败"测试:两处历史缺陷的回归锚——
|
||||
// 1. 节点转发握手无上限:agent 会话半死(收得到帧但永不应答)时观看端无限挂起
|
||||
// → HandshakeTimeout 到点必须切断观看端, 让其快速失败可重试;
|
||||
// 2. agent 拨内网目标无上限:目标离线且防火墙丢包(不回 RST)时裸 Dial 吊 20s+
|
||||
// → DialTimeout 到点必须向节点回执失败, 观看端立即断开。
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.zeroonesoft.cn/golib/zonat/agent"
|
||||
"git.zeroonesoft.cn/golib/zonat/internal/node"
|
||||
|
||||
"github.com/xtaci/smux"
|
||||
)
|
||||
|
||||
// deafAgent 手工模拟"半死 agent":完成登录+注册隧道后, 对节点打开的数据流
|
||||
// (CmdTarget)永不应答。返回节点为该隧道分配的公网端口。
|
||||
func deafAgent(t *testing.T, n *node.Node, tunnelId, targetIp string, targetPort int) int {
|
||||
t.Helper()
|
||||
conn, err := net.DialTimeout("tcp", n.Addr(), 2*time.Second)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
|
||||
tok, err := node.SignToken(n.JwtSecret, "deaf-agent", time.Hour)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
lf := agent.NewFrame(agent.FrameVersion, agent.CmdLogin, 0)
|
||||
if err := lf.Marshal(agent.Login{AgentId: "deaf-agent", Token: tok}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := agent.WriteFrame(conn, lf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rf, err := agent.ReadFrame(conn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ret := agent.Ret{}
|
||||
if err := rf.Unmarshal(&ret); err != nil || ret.Code != 0 {
|
||||
t.Fatalf("登录失败: %+v", ret)
|
||||
}
|
||||
|
||||
sess, err := smux.Client(conn, agent.SmuxConfig())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
control, err := sess.OpenStream()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reg := agent.NewFrame(agent.FrameVersion, agent.CmdRegisterTunnel, 0)
|
||||
if err := reg.Marshal(agent.Tunnel{Id: tunnelId, TargetIp: targetIp, TargetPort: targetPort, TTLSec: 0}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := agent.WriteFrame(control, reg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rr, err := agent.ReadFrame(control)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rret := agent.Ret{}
|
||||
if err := rr.Unmarshal(&rret); err != nil || rret.Code != 0 {
|
||||
t.Fatalf("注册失败: %+v", rret)
|
||||
}
|
||||
return rret.Port
|
||||
}
|
||||
|
||||
// assertViewerCutBy 断言观看端在时限内被节点断开(读到 EOF/重置, 而非超时)。
|
||||
func assertViewerCutBy(t *testing.T, port int, within time.Duration) {
|
||||
t.Helper()
|
||||
viewer, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(port)), 2*time.Second)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer viewer.Close()
|
||||
_ = viewer.SetReadDeadline(time.Now().Add(within))
|
||||
buf := make([]byte, 16)
|
||||
_, err = viewer.Read(buf)
|
||||
if err == nil {
|
||||
t.Fatalf("观看端读到了数据, 不应发生")
|
||||
}
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
||||
t.Fatalf("观看端 %v 内未被节点断开(超时切断未生效)", within)
|
||||
}
|
||||
}
|
||||
|
||||
// TestHandshakeTimeoutCutsViewer 半死 agent: 节点必须在握手超时内切断观看端
|
||||
// (修复前观看端无限挂起, 只能等 smux keepalive 判死整个会话 ~40s)。
|
||||
func TestHandshakeTimeoutCutsViewer(t *testing.T) {
|
||||
target := startEchoServer(t)
|
||||
n := node.New()
|
||||
n.JwtSecret = "test-jwt-secret"
|
||||
n.BindTunnel = "127.0.0.1"
|
||||
n.HandshakeTimeout = 500 * time.Millisecond
|
||||
if err := n.Start("127.0.0.1:0"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(n.Stop)
|
||||
|
||||
host, portStr, _ := net.SplitHostPort(target)
|
||||
portNum, _ := strconv.Atoi(portStr)
|
||||
listen := deafAgent(t, n, fmt.Sprintf("deaf-%d", time.Now().UnixNano()), host, portNum)
|
||||
|
||||
assertViewerCutBy(t, listen, 3*time.Second)
|
||||
}
|
||||
|
||||
// TestAgentDialTimeoutCutsViewer 目标黑洞(不回 RST): agent 拨号到点回执失败,
|
||||
// 观看端立即断开(修复前裸 Dial 在 Windows 吊 20s+, 观看端干等)。
|
||||
func TestAgentDialTimeoutCutsViewer(t *testing.T) {
|
||||
// 本机网络若把 TEST-NET-3 判定为可达(隧道/虚拟网卡), 黑洞前提不成立
|
||||
if c, err := net.DialTimeout("tcp", "203.0.113.1:5900", 300*time.Millisecond); err == nil {
|
||||
c.Close()
|
||||
t.Skip("本机可将 TEST-NET-3 判定连通, 黑洞拨号场景不成立")
|
||||
}
|
||||
|
||||
n, a := startEnv(t)
|
||||
a.DialTimeout = 300 * time.Millisecond
|
||||
id := fmt.Sprintf("blackhole-%d", time.Now().UnixNano())
|
||||
port, err := a.RegisterTunnel(buildTunnel(t, n, a, id, "203.0.113.1:5900", 0))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
assertViewerCutBy(t, port, 3*time.Second)
|
||||
}
|
||||
|
||||
// startCloseAfterEchoServer 回显一段即关闭的目标:模拟 tvnserver 进程退出。
|
||||
func startCloseAfterEchoServer(t *testing.T) string {
|
||||
t.Helper()
|
||||
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { l.Close() })
|
||||
go func() {
|
||||
for {
|
||||
c, err := l.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go func(c net.Conn) {
|
||||
defer c.Close()
|
||||
buf := make([]byte, 512)
|
||||
n, _ := c.Read(buf)
|
||||
if n > 0 {
|
||||
c.Write(buf[:n]) // 回显一次随即断开
|
||||
}
|
||||
}(c)
|
||||
}
|
||||
}()
|
||||
return l.Addr().String()
|
||||
}
|
||||
|
||||
// TestTargetCloseCutsViewer 目标端→观看端断开同步:目标(如 tvnserver)关闭后,
|
||||
// 观看端必须在时限内被同步断开, 不允许残留悬挂连接。
|
||||
func TestTargetCloseCutsViewer(t *testing.T) {
|
||||
n, a := startEnv(t)
|
||||
target := startCloseAfterEchoServer(t)
|
||||
_, port := mustRegister(t, n, a, target)
|
||||
|
||||
viewer, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(port)), 2*time.Second)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer viewer.Close()
|
||||
_ = viewer.SetDeadline(time.Now().Add(3 * time.Second))
|
||||
payload := []byte("probe")
|
||||
if _, err := viewer.Write(payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := make([]byte, len(payload))
|
||||
if _, err := io.ReadFull(viewer, got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 回显已到, 目标随即关闭 → 观看端应读到 EOF/重置(而非超时)
|
||||
buf := make([]byte, 16)
|
||||
_, err = viewer.Read(buf)
|
||||
if err == nil {
|
||||
t.Fatal("目标已关闭, 观看端却读到数据")
|
||||
}
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
||||
t.Fatal("目标关闭 3s 内未同步断开观看端")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package e2e
|
||||
|
||||
// 并发残留压测:多观看端并发穿透后全部离场(混合断开方式), 目标端必须
|
||||
// 零连接残留, 且隧道立即恢复可用——断开同步的正确性验收。
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// countedEcho 计数回显目标: 跟踪存活连接数与峰值, 连接保持到对端关闭。
|
||||
type countedEcho struct {
|
||||
ln net.Listener
|
||||
live atomic.Int64
|
||||
peak atomic.Int64
|
||||
}
|
||||
|
||||
func startCountedEcho(t *testing.T) *countedEcho {
|
||||
t.Helper()
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { ln.Close() })
|
||||
ce := &countedEcho{ln: ln}
|
||||
go func() {
|
||||
for {
|
||||
c, err := ln.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go func(c net.Conn) {
|
||||
defer c.Close()
|
||||
cur := ce.live.Add(1)
|
||||
defer ce.live.Add(-1)
|
||||
for { // 峰值 CAS 刷高
|
||||
p := ce.peak.Load()
|
||||
if cur <= p || ce.peak.CompareAndSwap(p, cur) {
|
||||
break
|
||||
}
|
||||
}
|
||||
io.Copy(c, c)
|
||||
}(c)
|
||||
}
|
||||
}()
|
||||
return ce
|
||||
}
|
||||
|
||||
// TestConcurrentViewersNoTargetLeak 50 观看端并发穿透、三种离场方式
|
||||
// (干净 FIN / RST 崩溃 / 迟走), 全部离场后目标端零残留, 隧道立即可复用。
|
||||
func TestConcurrentViewersNoTargetLeak(t *testing.T) {
|
||||
n, a := startEnv(t)
|
||||
ce := startCountedEcho(t)
|
||||
_, port := mustRegister(t, n, a, ce.ln.Addr().String())
|
||||
|
||||
const viewers = 50
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex
|
||||
var late []net.Conn // 迟走者: 并发段结束后统一关闭
|
||||
|
||||
for i := 0; i < viewers; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", port), 3*time.Second)
|
||||
if err != nil {
|
||||
t.Errorf("观看端%d 拨号失败: %v", i, err)
|
||||
return
|
||||
}
|
||||
_ = conn.SetDeadline(time.Now().Add(5 * time.Second))
|
||||
payload := []byte(fmt.Sprintf("viewer-%02d", i))
|
||||
if _, err := conn.Write(payload); err != nil {
|
||||
t.Errorf("观看端%d 写失败: %v", i, err)
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
got := make([]byte, len(payload))
|
||||
if _, err := io.ReadFull(conn, got); err != nil {
|
||||
t.Errorf("观看端%d 读回失败: %v", i, err)
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
switch i % 3 {
|
||||
case 0: // 干净 FIN(页面正常关闭)
|
||||
conn.Close()
|
||||
case 1: // RST 硬断(观看端进程崩溃)
|
||||
if tc, ok := conn.(*net.TCPConn); ok {
|
||||
tc.SetLinger(0)
|
||||
}
|
||||
conn.Close()
|
||||
default: // 迟走: 压测段结束后统一关
|
||||
mu.Lock()
|
||||
late = append(late, conn)
|
||||
mu.Unlock()
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
t.Logf("并发 %d 观看端完成往返: 目标端存活=%d 峰值=%d", viewers, ce.live.Load(), ce.peak.Load())
|
||||
|
||||
for _, c := range late {
|
||||
c.Close()
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(10 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if ce.live.Load() == 0 {
|
||||
break
|
||||
}
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
if left := ce.live.Load(); left != 0 {
|
||||
t.Fatalf("全部观看端离场后目标端残留 %d 条连接(应为 0)", left)
|
||||
}
|
||||
|
||||
// 隧道压测后立即可复用: 新观看端一轮完整往返
|
||||
if got := roundTrip(t, port, []byte("after-storm")); string(got) != "after-storm" {
|
||||
t.Fatalf("压测后隧道不可用, 回读=%q", got)
|
||||
}
|
||||
}
|
||||
+28
-9
@@ -24,25 +24,28 @@ const loginTimeout = 10 * time.Second
|
||||
|
||||
// Node 穿透节点
|
||||
type Node struct {
|
||||
JwtSecret string // agent 登录 JWT 校验密钥(HS256 系)
|
||||
BindTunnel string // 隧道监听绑定地址(默认 0.0.0.0)
|
||||
SweepInterval time.Duration // TTL 清扫间隔(默认 30s)
|
||||
JwtSecret string // agent 登录 JWT 校验密钥(HS256 系)
|
||||
BindTunnel string // 隧道监听绑定地址(默认 0.0.0.0)
|
||||
SweepInterval time.Duration // TTL 清扫间隔(默认 30s)
|
||||
HandshakeTimeout time.Duration // 单条观看连接的 agent 拨号应答上限(默认 5s; 超时立即断开观看端, 数据阶段不受限)
|
||||
TunnelPortRange PortRange // 隧道监听端口范围(未配置时 ListenPort=0 走系统随机分配)
|
||||
|
||||
listener net.Listener
|
||||
|
||||
mu sync.Mutex
|
||||
stopped bool
|
||||
sessions map[*Session]struct{} // 该节点上的全部 agent 会话(一个 agent 多隧道=多会话)
|
||||
sessions map[*Session]struct{} // 该节点上的全部 agent 会话(一个 agent 多隧道=多会话)
|
||||
tunnels map[string]*TunnelServer // tunnelId -> 隧道
|
||||
}
|
||||
|
||||
// New 创建节点
|
||||
func New() *Node {
|
||||
return &Node{
|
||||
BindTunnel: "0.0.0.0",
|
||||
SweepInterval: 30 * time.Second,
|
||||
sessions: make(map[*Session]struct{}),
|
||||
tunnels: make(map[string]*TunnelServer),
|
||||
BindTunnel: "0.0.0.0",
|
||||
SweepInterval: 30 * time.Second,
|
||||
HandshakeTimeout: 5 * time.Second,
|
||||
sessions: make(map[*Session]struct{}),
|
||||
tunnels: make(map[string]*TunnelServer),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -54,6 +57,9 @@ func (n *Node) Start(addr string) error {
|
||||
}
|
||||
n.listener = l
|
||||
slog.Info("节点 桥接监听", "addr", l.Addr().String())
|
||||
if n.TunnelPortRange.Valid() {
|
||||
slog.Info("节点 隧道端口范围", "range", n.TunnelPortRange.String())
|
||||
}
|
||||
go n.acceptLoop()
|
||||
go n.sweepLoop()
|
||||
return nil
|
||||
@@ -236,7 +242,7 @@ func (n *Node) registerTunnel(s *Session, t agent.Tunnel) (agent.Tunnel, error)
|
||||
|
||||
ts := &TunnelServer{tunnel: t, session: s}
|
||||
ts.touch()
|
||||
l, err := net.Listen("tcp", net.JoinHostPort(n.BindTunnel, strconv.Itoa(t.ListenPort)))
|
||||
l, err := n.listenTunnel(t.ListenPort)
|
||||
if err != nil {
|
||||
return t, err
|
||||
}
|
||||
@@ -251,6 +257,19 @@ func (n *Node) registerTunnel(s *Session, t agent.Tunnel) (agent.Tunnel, error)
|
||||
return ts.tunnel, nil
|
||||
}
|
||||
|
||||
// listenTunnel 建立隧道公网监听:指定端口直接监听;端口 0 时若配置了范围则在范围内取空闲端口,
|
||||
// 未配置范围回退系统随机分配。
|
||||
func (n *Node) listenTunnel(listenPort int) (net.Listener, error) {
|
||||
if listenPort == 0 && n.TunnelPortRange.Valid() {
|
||||
l, err := n.TunnelPortRange.listen(n.BindTunnel)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("分配隧道端口失败: %w", err)
|
||||
}
|
||||
return l, nil
|
||||
}
|
||||
return net.Listen("tcp", net.JoinHostPort(n.BindTunnel, strconv.Itoa(listenPort)))
|
||||
}
|
||||
|
||||
// unregisterTunnel 删除隧道(关闭公网监听)
|
||||
func (n *Node) unregisterTunnel(id string) bool {
|
||||
n.mu.Lock()
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
// 端口范围:临时/永久隧道注册时未指定监听端口(ListenPort=0)的取值范围。
|
||||
// 云主机防火墙按范围放行,避免隧道端口落在系统随机高位段而无法预开。
|
||||
package node
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// PortRange 端口范围(含两端);Min==0 表示未配置(回退系统随机分配)。
|
||||
type PortRange struct {
|
||||
Min, Max int
|
||||
}
|
||||
|
||||
// Valid 是否已配置有效范围
|
||||
func (r PortRange) Valid() bool { return r.Min > 0 && r.Max >= r.Min }
|
||||
|
||||
// ParsePortRange 解析 "20000-21000"(单端口 "20000" 视为 min=max)。
|
||||
func ParsePortRange(s string) (PortRange, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return PortRange{}, nil
|
||||
}
|
||||
minStr, maxStr, ok := strings.Cut(s, "-")
|
||||
if !ok {
|
||||
maxStr = minStr
|
||||
}
|
||||
minP, err := strconv.Atoi(strings.TrimSpace(minStr))
|
||||
if err != nil {
|
||||
return PortRange{}, fmt.Errorf("端口范围格式错误:%s", s)
|
||||
}
|
||||
maxP, err := strconv.Atoi(strings.TrimSpace(maxStr))
|
||||
if err != nil {
|
||||
return PortRange{}, fmt.Errorf("端口范围格式错误:%s", s)
|
||||
}
|
||||
r := PortRange{Min: minP, Max: maxP}
|
||||
if !r.Valid() || r.Max > 65535 {
|
||||
return PortRange{}, fmt.Errorf("端口范围无效:%s", s)
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
// String 还原配置形式(便于日志)
|
||||
func (r PortRange) String() string {
|
||||
if !r.Valid() {
|
||||
return ""
|
||||
}
|
||||
if r.Min == r.Max {
|
||||
return strconv.Itoa(r.Min)
|
||||
}
|
||||
return fmt.Sprintf("%d-%d", r.Min, r.Max)
|
||||
}
|
||||
|
||||
// listen 从范围内随机起点顺延探测,返回第一个监听成功的 listener(占用即生效,避免探测与使用竞态)。
|
||||
// 全部被占返回错误。bind 为隧道监听绑定地址。
|
||||
func (r PortRange) listen(bind string) (net.Listener, error) {
|
||||
if !r.Valid() {
|
||||
return nil, fmt.Errorf("端口范围未配置")
|
||||
}
|
||||
start := r.Min + rand.Intn(r.Max-r.Min+1)
|
||||
for i := 0; i <= r.Max-r.Min; i++ {
|
||||
port := r.Min + (start-r.Min+i)%(r.Max-r.Min+1)
|
||||
l, err := net.Listen("tcp", net.JoinHostPort(bind, strconv.Itoa(port)))
|
||||
if err == nil {
|
||||
return l, nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("端口范围 %s 内无可用端口", r)
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package node
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParsePortRange(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
wantMin int
|
||||
wantMax int
|
||||
wantErr bool
|
||||
}{
|
||||
{"", 0, 0, false}, // 空 = 未配置
|
||||
{"20000-21000", 20000, 21000, false},
|
||||
{"20000 - 21000", 20000, 21000, false}, // 两侧空白
|
||||
{"5900", 5900, 5900, false}, // 单端口
|
||||
{"abc", 0, 0, true},
|
||||
{"21000-20000", 0, 0, true}, // 逆序
|
||||
{"0-100", 0, 0, true}, // 零起点无效
|
||||
{"65536", 0, 0, true}, // 越界
|
||||
}
|
||||
for _, c := range cases {
|
||||
r, err := ParsePortRange(c.in)
|
||||
if c.wantErr {
|
||||
if err == nil {
|
||||
t.Errorf("ParsePortRange(%q) 期望报错, 得到 %v", c.in, r)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
t.Errorf("ParsePortRange(%q) 意外报错: %v", c.in, err)
|
||||
continue
|
||||
}
|
||||
if r.Min != c.wantMin || r.Max != c.wantMax {
|
||||
t.Errorf("ParsePortRange(%q) = %v, 期望 %d-%d", c.in, r, c.wantMin, c.wantMax)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPortRangeListen(t *testing.T) {
|
||||
r, err := ParsePortRange("20000-20002")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 占满 3 个端口后第 4 次应失败
|
||||
var ls []net.Listener
|
||||
defer func() {
|
||||
for _, l := range ls {
|
||||
l.Close()
|
||||
}
|
||||
}()
|
||||
for i := 0; i < 3; i++ {
|
||||
l, err := r.listen("127.0.0.1")
|
||||
if err != nil {
|
||||
t.Fatalf("第 %d 次取端口失败: %v", i+1, err)
|
||||
}
|
||||
ls = append(ls, l)
|
||||
port := l.Addr().(*net.TCPAddr).Port
|
||||
if port < 20000 || port > 20002 {
|
||||
t.Fatalf("取得端口 %d 超出范围", port)
|
||||
}
|
||||
}
|
||||
if _, err := r.listen("127.0.0.1"); err == nil {
|
||||
t.Fatal("范围占满后应返回错误")
|
||||
}
|
||||
// 释放后应可再取
|
||||
ls[0].Close()
|
||||
l, err := r.listen("127.0.0.1")
|
||||
if err != nil {
|
||||
t.Fatalf("释放后再取失败: %v", err)
|
||||
}
|
||||
l.Close()
|
||||
}
|
||||
@@ -58,6 +58,11 @@ func (t *TunnelServer) handleConn(conn net.Conn) {
|
||||
}
|
||||
defer stream.Close()
|
||||
|
||||
// 转发握手(告知目标+等 agent 回执)限时: agent 会话半死/目标拨号不回时
|
||||
// 观看端不再无限挂起, 到点立即断开让观看端快速失败可重试
|
||||
if hs := t.session.node.HandshakeTimeout; hs > 0 {
|
||||
_ = stream.SetDeadline(time.Now().Add(hs))
|
||||
}
|
||||
tf := agent.NewFrame(agent.FrameVersion, agent.CmdTarget, 0)
|
||||
if err := tf.Marshal(t.tunnel); err != nil {
|
||||
return
|
||||
@@ -67,8 +72,10 @@ func (t *TunnelServer) handleConn(conn net.Conn) {
|
||||
}
|
||||
rf, err := agent.ReadFrame(stream)
|
||||
if err != nil {
|
||||
slog.Warn("节点 agent拨号应答失败", "id", t.tunnel.Id, "err", err.Error())
|
||||
return
|
||||
}
|
||||
_ = stream.SetDeadline(time.Time{}) // 数据阶段不受握手超时约束
|
||||
ret := agent.Ret{}
|
||||
if err := rf.Unmarshal(&ret); err != nil || ret.Code != 0 {
|
||||
slog.Warn("节点 agent拨号失败", "id", t.tunnel.Id, "msg", ret.Msg)
|
||||
|
||||
Reference in New Issue
Block a user