diff --git a/agent/forward.go b/agent/forward.go index f399b1e..d1fcd43 100644 --- a/agent/forward.go +++ b/agent/forward.go @@ -3,6 +3,7 @@ package agent import ( "io" "net" + "time" ) // handleData 处理节点主动打开的数据流: @@ -35,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 index 9d4625a..8f19788 100644 --- a/e2e/handshake_test.go +++ b/e2e/handshake_test.go @@ -8,6 +8,7 @@ package e2e import ( "errors" "fmt" + "io" "net" "strconv" "testing" @@ -134,3 +135,63 @@ func TestAgentDialTimeoutCutsViewer(t *testing.T) { 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 内未同步断开观看端") + } +}