Files

392 lines
10 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 端到端穿透测试:节点 + agent + 内网目标全部进程内启动,
// 验证 公网口 → 桥接(smux/TCP) → 内网目标 的完整链路,不依赖 capricorn。
// 每条隧道自带节点地址与 JWT,独立连接登录;agent 不再有集中连接。
package e2e
import (
"bytes"
"crypto/rand"
"encoding/binary"
"fmt"
"io"
"net"
"strconv"
"sync"
"testing"
"time"
"git.zeroonesoft.cn/golib/zonat/agent"
"git.zeroonesoft.cn/golib/zonat/internal/node"
)
// startEchoServer 启动 TCP 回显服务(模拟 VNC/RDP/3000 等内网 TCP 目标)
func startEchoServer(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 {
conn, err := l.Accept()
if err != nil {
return
}
go func(c net.Conn) {
defer c.Close()
io.Copy(c, c)
}(conn)
}
}()
return l.Addr().String()
}
// startNode 启动一个节点(随机端口,JWT 密钥 test-jwt-secret)
func startNode(t *testing.T) *node.Node {
t.Helper()
n := node.New()
n.JwtSecret = "test-jwt-secret"
n.BindTunnel = "127.0.0.1"
n.SweepInterval = 200 * time.Millisecond
if err := n.Start("127.0.0.1:0"); err != nil {
t.Fatal(err)
}
t.Cleanup(n.Stop)
return n
}
// startAgent 启动一个 agent(不连接任何节点,隧道注册时才连)
func startAgent(t *testing.T, id string) *agent.Agent {
t.Helper()
a := agent.New(id)
a.Backoff = 100 * time.Millisecond
a.PingEvery = 200 * time.Millisecond
go a.Run()
t.Cleanup(a.Close)
return a
}
// startEnv 单节点 + agent-1
func startEnv(t *testing.T) (*node.Node, *agent.Agent) {
t.Helper()
n := startNode(t)
return n, startAgent(t, "agent-1")
}
// mintToken 签发该节点可认的登录 JWT
func mintToken(t *testing.T, n *node.Node, agentId string) string {
t.Helper()
tok, err := node.SignToken(n.JwtSecret, agentId, time.Hour)
if err != nil {
t.Fatal(err)
}
return tok
}
// buildTunnel 构造一条登记到指定节点的隧道定义
func buildTunnel(t *testing.T, n *node.Node, a *agent.Agent, id, target string, ttl int) agent.Tunnel {
t.Helper()
host, portStr, _ := net.SplitHostPort(target)
portNum, _ := strconv.Atoi(portStr)
return agent.Tunnel{
Id: id, NodeAddr: n.Addr(), Token: mintToken(t, n, a.AgentId),
TargetIp: host, TargetPort: portNum, TTLSec: ttl,
}
}
func mustRegister(t *testing.T, n *node.Node, a *agent.Agent, target string) (id string, port int) {
t.Helper()
id = fmt.Sprintf("t-%d", time.Now().UnixNano())
port, err := a.RegisterTunnel(buildTunnel(t, n, a, id, target, 0))
if err != nil {
t.Fatal(err)
}
return id, port
}
// roundTrip 建立一条穿透连接,发送 payload 并读回,返回实际回读内容
func roundTrip(t *testing.T, port int, payload []byte) []byte {
t.Helper()
conn, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(port)), 2*time.Second)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
if err := conn.SetDeadline(time.Now().Add(10 * time.Second)); err != nil {
t.Fatal(err)
}
if _, err := conn.Write(payload); err != nil {
t.Fatal(err)
}
got := make([]byte, len(payload))
if _, err := io.ReadFull(conn, got); err != nil {
t.Fatal(err)
}
return got
}
func TestTCPEchoRoundTrip(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
_, port := mustRegister(t, n, a, target)
payload := []byte("hello 内网穿透")
if got := roundTrip(t, port, payload); !bytes.Equal(got, payload) {
t.Fatalf("回显不匹配 got=%q want=%q", got, payload)
}
}
// TestMultiNodeTunnels 本次多节点重构的验收用例:
// 一个 agent 的两条隧道分别登录两个不同节点,各自独立可用。
func TestMultiNodeTunnels(t *testing.T) {
target1 := startEchoServer(t)
target2 := startEchoServer(t)
n1 := startNode(t)
n2 := startNode(t)
a := startAgent(t, "multi-1")
p1, err := a.RegisterTunnel(buildTunnel(t, n1, a, "t-node1", target1, 0))
if err != nil {
t.Fatal(err)
}
p2, err := a.RegisterTunnel(buildTunnel(t, n2, a, "t-node2", target2, 0))
if err != nil {
t.Fatal(err)
}
if !a.Connected("t-node1") || !a.Connected("t-node2") {
t.Fatal("两条隧道应同时保持连接")
}
if got := roundTrip(t, p1, []byte("via-node-1")); !bytes.Equal(got, []byte("via-node-1")) {
t.Fatal("节点1的隧道回显失败")
}
if got := roundTrip(t, p2, []byte("via-node-2")); !bytes.Equal(got, []byte("via-node-2")) {
t.Fatal("节点2的隧道回显失败")
}
// 踢掉其中一个节点上的会话,只影响该节点的隧道
n1.Kick("multi-1")
time.Sleep(300 * time.Millisecond)
if !a.Connected("t-node2") {
t.Fatal("节点2的隧道不应受节点1踢会话影响")
}
}
func TestTCPEchoConcurrent(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
_, port := mustRegister(t, n, a, target)
const workers, rounds = 20, 10
var wg sync.WaitGroup
errCh := make(chan error, workers)
for i := 0; i < workers; i++ {
wg.Add(1)
go func(seed byte) {
defer wg.Done()
payload := bytes.Repeat([]byte{seed}, 1024)
for r := 0; r < rounds; r++ {
conn, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(port)), 2*time.Second)
if err != nil {
errCh <- err
return
}
if _, err := conn.Write(payload); err != nil {
conn.Close()
errCh <- err
return
}
got := make([]byte, len(payload))
if _, err := io.ReadFull(conn, got); err != nil {
conn.Close()
errCh <- err
return
}
conn.Close()
if !bytes.Equal(got, payload) {
errCh <- fmt.Errorf("回显不匹配 seed=%d", seed)
return
}
}
}(byte(i + 1))
}
wg.Wait()
select {
case err := <-errCh:
t.Fatal(err)
default:
}
}
func TestTCPEchoLargePayload(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
_, port := mustRegister(t, n, a, target)
// 4MB:远大于 smux 单帧(32KB),验证分帧与重组
payload := make([]byte, 4<<20)
if _, err := rand.Read(payload); err != nil {
t.Fatal(err)
}
if got := roundTrip(t, port, payload); !bytes.Equal(got, payload) {
t.Fatal("大包回显不匹配")
}
}
func TestUnregisterStopsListener(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
id, port := mustRegister(t, n, a, target)
if got := roundTrip(t, port, []byte("ok")); !bytes.Equal(got, []byte("ok")) {
t.Fatal("注销前回显失败")
}
if err := a.UnregisterTunnel(id); err != nil {
t.Fatal(err)
}
time.Sleep(200 * time.Millisecond)
if conn, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(port)), 1*time.Second); err == nil {
conn.Close()
t.Fatal("注销后端口仍可连接")
}
}
func TestTTLAutoDelete(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
port, err := a.RegisterTunnel(buildTunnel(t, n, a, "ttl-tunnel", target, 1))
if err != nil {
t.Fatal(err)
}
if got := roundTrip(t, port, []byte("x")); len(got) != 1 {
t.Fatal("TTL 隧道不可用")
}
// ttl=1s + 清扫间隔 200ms,2s 后应已删除并收到关闭推送
time.Sleep(2 * time.Second)
if conn, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(port)), 1*time.Second); err == nil {
conn.Close()
t.Fatal("TTL 过期后端口仍可连接")
}
}
func TestDuplicateTunnelId(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
tun := buildTunnel(t, n, a, "dup", target, 0)
if _, err := a.RegisterTunnel(tun); err != nil {
t.Fatal(err)
}
// 同 agent 重复注册同 ID 视为更新配置:停旧起新,应成功
if _, err := a.RegisterTunnel(tun); err != nil {
t.Fatalf("同 agent 重复注册应顶替成功: %v", err)
}
// 其它 agent 的活隧道占用同 ID:应拒绝
b := startAgent(t, "agent-2")
if _, err := b.RegisterTunnel(buildTunnel(t, n, b, "dup", target, 0)); err == nil {
t.Fatal("不同 agent 注册同一隧道ID应失败")
}
}
// TestJWTAuthRejected JWT 认证拒绝:错误密钥签名 / 已过期 / agentId 不一致
func TestJWTAuthRejected(t *testing.T) {
target := startEchoServer(t)
n := startNode(t)
expired, err := node.SignToken(n.JwtSecret, "agent-x", -time.Minute)
if err != nil {
t.Fatal(err)
}
cases := []struct {
name string
token func(t *testing.T) string
}{
{"错误密钥签名", func(t *testing.T) string {
tok, err := node.SignToken("wrong-secret", "agent-x", time.Hour)
if err != nil {
t.Fatal(err)
}
return tok
}},
{"过期token", func(*testing.T) string { return expired }},
{"agentId不一致", func(t *testing.T) string { return mintToken(t, n, "someone-else") }},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
a := startAgent(t, "agent-x")
host, portStr, _ := net.SplitHostPort(target)
portNum, _ := strconv.Atoi(portStr)
_, err := a.RegisterTunnel(agent.Tunnel{
Id: "bad-jwt", NodeAddr: n.Addr(), Token: c.token(t),
TargetIp: host, TargetPort: portNum,
})
if err == nil {
t.Fatal("JWT 校验未通过仍登录成功")
}
})
}
}
func TestAgentReRegisterAfterKick(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
portCh := make(chan int, 8)
a.OnTunnelPort = func(id string, port int) {
if port > 0 {
portCh <- port
}
}
if _, err := a.RegisterTunnel(buildTunnel(t, n, a, "keep", target, 0)); err != nil {
t.Fatal(err)
}
firstPort := <-portCh
// 节点踢掉会话,该隧道的 worker 应自动重连并重注册
if !n.Kick("agent-1") {
t.Fatal("Kick 失败")
}
var newPort int
select {
case newPort = <-portCh:
case <-time.After(5 * time.Second):
t.Fatal("重连后未重注册隧道")
}
if newPort == firstPort {
t.Log("重注册端口与原端口相同(随机分配碰撞,可接受)")
}
if got := roundTrip(t, newPort, []byte("after-kick")); !bytes.Equal(got, []byte("after-kick")) {
t.Fatal("重连后穿透失败")
}
}
// TestBinarySafePayload 首字节为 0x16(TLS 握手特征)等敏感字节的数据
// 必须原样通过——完整版的嗅探劫持问题在本库中已不存在。
func TestBinarySafePayload(t *testing.T) {
target := startEchoServer(t)
n, a := startEnv(t)
_, port := mustRegister(t, n, a, target)
// 0x16 0x03 0x01 开头:zonat IsHttps 会把它劫持到兜底站
payload := make([]byte, 0, 4096)
head := []byte{0x16, 0x03, 0x01, 0x00}
payload = append(payload, head...)
var l uint16 = 4000
binary.BigEndian.PutUint16(payload[3:5], l)
rest := make([]byte, 4096-len(payload))
if _, err := rand.Read(rest); err != nil {
t.Fatal(err)
}
payload = append(payload, rest...)
if got := roundTrip(t, port, payload); !bytes.Equal(got, payload) {
t.Fatal("二进制载荷被篡改(疑似协议嗅探干扰)")
}
}