diff --git a/cmd/node/config.go b/cmd/node/config.go new file mode 100644 index 0000000..cae387f --- /dev/null +++ b/cmd/node/config.go @@ -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 +} diff --git a/cmd/node/main.go b/cmd/node/main.go index 02e387a..01fbb88 100644 --- a/cmd/node/main.go +++ b/cmd/node/main.go @@ -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); 显式 flag 可覆盖配置项") + 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,76 @@ 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 }) + + 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 或配置文件 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) } diff --git a/internal/node/node.go b/internal/node/node.go index ec28753..8b75dce 100644 --- a/internal/node/node.go +++ b/internal/node/node.go @@ -24,9 +24,10 @@ 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) + TunnelPortRange PortRange // 隧道监听端口范围(未配置时 ListenPort=0 走系统随机分配) listener net.Listener @@ -54,6 +55,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 +240,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 +255,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() diff --git a/internal/node/portrange.go b/internal/node/portrange.go new file mode 100644 index 0000000..cee3d08 --- /dev/null +++ b/internal/node/portrange.go @@ -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) +} diff --git a/internal/node/portrange_test.go b/internal/node/portrange_test.go new file mode 100644 index 0000000..9336a7e --- /dev/null +++ b/internal/node/portrange_test.go @@ -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() +}