// Package node 轻量穿透节点(zonat server 的裁剪版): // TCP 桥接登录(JWT)→ smux 会话 → 隧道注册表(纯内存)→ 纯 TCP 隧道监听。 // 相比 zonat server 裁掉:KCP 桥接、udp/http(s) 隧道、协议嗅探与兜底站、 // 管理 REST、数据库、限流与流量统计。 // 隧道生命周期由 agent 经控制流帧命令注册/注销,节点不再对外暴露任何管理端口。 // 一个 agent 的多条隧道各持一条桥接连接(可分布多节点),因此同一 agentId // 同时存在多个会话是常态;同名隧道顶替按"隧道 ID + agentId"粒度在注册时进行。 package node import ( "errors" "fmt" "log/slog" "net" "strconv" "sync" "time" "git.zeroonesoft.cn/golib/zonat/agent" ) // loginTimeout 登录握手读超时 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) listener net.Listener mu sync.Mutex stopped bool 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), } } // Start 启动桥接监听与 TTL 清扫 func (n *Node) Start(addr string) error { l, err := net.Listen("tcp", addr) if err != nil { return err } n.listener = l slog.Info("节点 桥接监听", "addr", l.Addr().String()) go n.acceptLoop() go n.sweepLoop() return nil } // Stop 停止节点,关闭所有隧道与会话 func (n *Node) Stop() { n.mu.Lock() if n.stopped { n.mu.Unlock() return } n.stopped = true sessions := make([]*Session, 0, len(n.sessions)) for s := range n.sessions { sessions = append(sessions, s) } tunnels := make([]*TunnelServer, 0, len(n.tunnels)) for _, t := range n.tunnels { tunnels = append(tunnels, t) } n.mu.Unlock() if n.listener != nil { n.listener.Close() } for _, t := range tunnels { t.close() } for _, s := range sessions { s.close() } } func (n *Node) isStopped() bool { n.mu.Lock() defer n.mu.Unlock() return n.stopped } // Addr 桥接监听地址 func (n *Node) Addr() string { if n.listener == nil { return "" } return n.listener.Addr().String() } func (n *Node) acceptLoop() { for { conn, err := n.listener.Accept() if err != nil { if !n.isStopped() { slog.Warn("节点 桥接接受连接错误", "err", err.Error()) } return } go n.handleConn(conn) } } // handleConn 登录握手(JWT 校验)+ 建立 smux 会话。 // 新版 agent 每条隧道一条连接,同一 agentId 的多条会话并存。 func (n *Node) handleConn(conn net.Conn) { conn.SetReadDeadline(time.Now().Add(loginTimeout)) f, err := agent.ReadFrame(conn) if err != nil { conn.Close() return } if f.Cmd() != agent.CmdLogin { conn.Close() return } lg := agent.Login{} if err := f.Unmarshal(&lg); err != nil || lg.AgentId == "" { _ = writeRet(conn, agent.CmdLogin, 0, agent.Ret{Code: 1, Msg: "登录失败"}) conn.Close() return } if err := verifyToken(lg.Token, n.JwtSecret, lg.AgentId); err != nil { slog.Warn("节点 agent登录拒绝", "agent", lg.AgentId, "err", err.Error()) _ = writeRet(conn, agent.CmdLogin, 0, agent.Ret{Code: 1, Msg: "登录失败"}) conn.Close() return } if err := writeRet(conn, agent.CmdLogin, 0, agent.Ret{Code: 0, Msg: "登录成功"}); err != nil { conn.Close() return } // smux 自管超时,清除登录用的读超时 conn.SetReadDeadline(time.Time{}) sess, err := smuxServer(conn) if err != nil { conn.Close() return } s := &Session{agentId: lg.AgentId, node: n, session: sess} n.mu.Lock() if n.stopped { n.mu.Unlock() conn.Close() return } n.sessions[s] = struct{}{} n.mu.Unlock() slog.Info("节点 agent登录", "agent", lg.AgentId, "remote", conn.RemoteAddr().String()) s.loop() n.detachSession(s) } // detachSession 会话断开后:摘除会话,关闭其名下全部隧道 func (n *Node) detachSession(s *Session) { n.mu.Lock() delete(n.sessions, s) var mine []*TunnelServer for id, t := range n.tunnels { if t.session == s { mine = append(mine, t) delete(n.tunnels, id) } } n.mu.Unlock() for _, t := range mine { t.close() } slog.Info("节点 agent会话结束", "agent", s.agentId, "关闭隧道", len(mine)) } // Kick 踢掉指定 agent 的全部会话(运维/测试用)。 // 新版 agent 每条隧道一条会话,Kick 会断其所有隧道,agent 侧各自重连重注册。 func (n *Node) Kick(agentId string) bool { n.mu.Lock() var targets []*Session for s := range n.sessions { if s.agentId == agentId { targets = append(targets, s) } } n.mu.Unlock() for _, s := range targets { s.close() } return len(targets) > 0 } // registerTunnel 为 agent 注册一条隧道并开始公网监听 func (n *Node) registerTunnel(s *Session, t agent.Tunnel) (agent.Tunnel, error) { if t.Id == "" { return t, errors.New("隧道ID不能为空") } if t.TargetIp == "" || t.TargetPort <= 0 || t.TargetPort > 65535 { return t, fmt.Errorf("目标地址错误:%s", t.TargetAddr()) } if t.ListenPort < 0 || t.ListenPort > 65535 { return t, fmt.Errorf("监听端口错误:%d", t.ListenPort) } // 凭据只用于登录,节点侧不留存也不回传(CmdTarget 不携带) t.NodeAddr, t.Token = "", "" n.mu.Lock() defer n.mu.Unlock() if n.stopped { return t, errors.New("节点已停止") } // 同 ID 隧道已存在:同 agent 重登顶替(更新配置)或死会话残留可替换, // 其它 agent 的活隧道拒绝 if old, ok := n.tunnels[t.Id]; ok { if old.session.agentId != s.agentId && !old.session.isDead() { return t, fmt.Errorf("隧道已存在:%s", t.Id) } old.close() old.session.close() // 顶替旧连接,避免旧 worker 与新 worker 争抢 delete(n.tunnels, t.Id) slog.Info("节点 同agent隧道顶替", "id", t.Id, "agent", s.agentId) } ts := &TunnelServer{tunnel: t, session: s} ts.touch() l, err := net.Listen("tcp", net.JoinHostPort(n.BindTunnel, strconv.Itoa(t.ListenPort))) if err != nil { return t, err } ts.listener = l _, portStr, _ := net.SplitHostPort(l.Addr().String()) ts.tunnel.ListenPort, _ = strconv.Atoi(portStr) n.tunnels[t.Id] = ts go ts.acceptLoop() slog.Info("节点 隧道注册", "id", t.Id, "listen", ts.tunnel.ListenPort, "target", t.TargetAddr(), "ttl", t.TTLSec) return ts.tunnel, nil } // unregisterTunnel 删除隧道(关闭公网监听) func (n *Node) unregisterTunnel(id string) bool { n.mu.Lock() ts, ok := n.tunnels[id] if ok { delete(n.tunnels, id) } n.mu.Unlock() if !ok { return false } ts.close() slog.Info("节点 隧道删除", "id", id) return true } func (n *Node) sweepLoop() { ticker := time.NewTicker(n.SweepInterval) defer ticker.Stop() for range ticker.C { if n.isStopped() { return } n.sweepOnce() } } // sweepOnce 清扫 TTL 过期且空闲的隧道,并向所属 agent 推送关闭通知 func (n *Node) sweepOnce() { now := time.Now() n.mu.Lock() var expired []*TunnelServer for id, ts := range n.tunnels { ttl := time.Duration(ts.tunnel.TTLSec) * time.Second if ttl > 0 && ts.connCount.Load() == 0 && now.Sub(ts.lastActiveTime()) > ttl { expired = append(expired, ts) delete(n.tunnels, id) } } n.mu.Unlock() for _, ts := range expired { ts.close() ts.session.pushClosed(ts.tunnel.Id) slog.Info("节点 隧道TTL过期删除", "id", ts.tunnel.Id) } } func writeRet(conn net.Conn, cmd byte, sid uint32, ret agent.Ret) error { f := agent.NewFrame(agent.FrameVersion, cmd, sid) if err := f.Marshal(ret); err != nil { return err } return agent.WriteFrame(conn, f) }