diff --git a/agent/agent.go b/agent/agent.go index d1d2be4..796d563 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -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{}), } } diff --git a/agent/forward.go b/agent/forward.go index 7ca43bd..d1fcd43 100644 --- a/agent/forward.go +++ b/agent/forward.go @@ -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"}) diff --git a/e2e/handshake_test.go b/e2e/handshake_test.go new file mode 100644 index 0000000..8f19788 --- /dev/null +++ b/e2e/handshake_test.go @@ -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 内未同步断开观看端") + } +} diff --git a/e2e/leak_test.go b/e2e/leak_test.go new file mode 100644 index 0000000..2e0652a --- /dev/null +++ b/e2e/leak_test.go @@ -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) + } +} diff --git a/internal/node/node.go b/internal/node/node.go index 8b75dce..5900d9c 100644 --- a/internal/node/node.go +++ b/internal/node/node.go @@ -24,26 +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) - TunnelPortRange PortRange // 隧道监听端口范围(未配置时 ListenPort=0 走系统随机分配) + 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), } } diff --git a/internal/node/tunnel.go b/internal/node/tunnel.go index 99a0fee..cf5004e 100644 --- a/internal/node/tunnel.go +++ b/internal/node/tunnel.go @@ -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)