diff --git a/config.example.toml b/config.example.toml index d8a15c1..6b2f9b3 100644 --- a/config.example.toml +++ b/config.example.toml @@ -13,6 +13,7 @@ local_addr = "127.0.0.1:9090" [[connect.web]] # others -> you protocol = "tcp" local_port = 9000 +local_addr = "127.0.0.1" # default; set to 0.0.0.0 to expose on LAN dst_addr = "any-client-in.ts.net:8080" [[connect.minecraft]] @@ -25,4 +26,4 @@ lan_motd = "Minecraft via Tailscale" [[connect.udp_example]] protocol = "udp" local_port = 24454 -dst_addr = "any-client-in.ts.net:24454" \ No newline at end of file +dst_addr = "any-client-in.ts.net:24454" diff --git a/core/config.go b/core/config.go index b882159..65ba94e 100644 --- a/core/config.go +++ b/core/config.go @@ -2,8 +2,10 @@ package core import ( "bytes" + "errors" "fmt" "io" + "net" "net/http" "os" "strings" @@ -53,12 +55,150 @@ func (r ConnectRule) LANMotdOr(def string) string { return def } +func (r ConnectRule) BindIP() string { + if r.LANEnabled() { + return "0.0.0.0" + } + if r.LocalAddr != "" { + return r.LocalAddr + } + return "127.0.0.1" +} + type Config struct { Core Core `toml:"core"` Forward map[string][]ForwardRule `toml:"forward"` Connect map[string][]ConnectRule `toml:"connect"` } +func (cfg *Config) ApplyDefaults() { + if cfg.Forward == nil { + cfg.Forward = make(map[string][]ForwardRule) + } + if cfg.Connect == nil { + cfg.Connect = make(map[string][]ConnectRule) + } + if cfg.Core.Hostname == "" { + hostname, err := os.Hostname() + if err != nil { + hostname = "unknown" + } + cfg.Core.Hostname = hostname + } +} + +func (cfg *Config) Validate() error { + var errs []error + + if strings.TrimSpace(cfg.Core.AuthKey) == "" { + errs = append(errs, errors.New("core.auth_key is required")) + } + + usedForwardListeners := make(map[string]string) + usedConnectListeners := make(map[string]string) + + for tag, rules := range cfg.Forward { + for i, rule := range rules { + path := fmt.Sprintf("forward.%s[%d]", tag, i) + if rule.Protocol != "tcp" && rule.Protocol != "udp" { + errs = append(errs, fmt.Errorf("%s.protocol must be tcp or udp", path)) + } + if !validPort(rule.TailscalePort) { + errs = append(errs, fmt.Errorf("%s.tailscale_port must be between 1 and 65535", path)) + } else { + key := fmt.Sprintf("%s:%d", rule.Protocol, rule.TailscalePort) + if prev, ok := usedForwardListeners[key]; ok { + errs = append(errs, fmt.Errorf("%s.tailscale_port duplicates %s", path, prev)) + } else { + usedForwardListeners[key] = path + } + } + if strings.TrimSpace(rule.LocalAddr) == "" { + errs = append(errs, fmt.Errorf("%s.local_addr is required", path)) + } else if err := validateHostPort(rule.LocalAddr); err != nil { + errs = append(errs, fmt.Errorf("%s.local_addr invalid: %w", path, err)) + } + } + } + + for tag, rules := range cfg.Connect { + for i, rule := range rules { + path := fmt.Sprintf("connect.%s[%d]", tag, i) + if rule.Protocol != "tcp" && rule.Protocol != "udp" && rule.Protocol != "minecraft" { + errs = append(errs, fmt.Errorf("%s.protocol must be tcp, udp, or minecraft", path)) + } + if !validPort(rule.LocalPort) { + errs = append(errs, fmt.Errorf("%s.local_port must be between 1 and 65535", path)) + } + if strings.TrimSpace(rule.DstAddr) == "" { + errs = append(errs, fmt.Errorf("%s.dst_addr is required", path)) + } else if err := validateHostPort(rule.DstAddr); err != nil { + errs = append(errs, fmt.Errorf("%s.dst_addr invalid: %w", path, err)) + } + if rule.LocalAddr != "" && net.ParseIP(rule.LocalAddr) == nil { + errs = append(errs, fmt.Errorf("%s.local_addr must be an IP address", path)) + } + if validPort(rule.LocalPort) && (rule.Protocol == "tcp" || rule.Protocol == "udp" || rule.Protocol == "minecraft") { + network := rule.Protocol + if network == "minecraft" { + network = "tcp" + } + if prev, ok := conflictingListener(usedConnectListeners, network, rule.BindIP(), rule.LocalPort); ok { + errs = append(errs, fmt.Errorf("%s local listener duplicates %s", path, prev)) + } + usedConnectListeners[listenerKey(network, rule.BindIP(), rule.LocalPort)] = path + } + } + } + + return errors.Join(errs...) +} + +func validPort(port int) bool { + return port > 0 && port <= 65535 +} + +func validateHostPort(addr string) error { + host, port, err := net.SplitHostPort(addr) + if err != nil { + return err + } + if strings.TrimSpace(host) == "" { + return errors.New("host is required") + } + if strings.TrimSpace(port) == "" { + return errors.New("port is required") + } + return nil +} + +func listenerKey(network, ip string, port int) string { + return fmt.Sprintf("%s/%s", network, net.JoinHostPort(ip, fmt.Sprintf("%d", port))) +} + +func conflictingListener(used map[string]string, network, ip string, port int) (string, bool) { + candidates := []string{ + listenerKey(network, ip, port), + } + if ip == "0.0.0.0" { + for key, path := range used { + prefix := network + "/" + _, usedPort, err := net.SplitHostPort(strings.TrimPrefix(key, prefix)) + if strings.HasPrefix(key, prefix) && err == nil && usedPort == fmt.Sprintf("%d", port) { + return path, true + } + } + } else { + candidates = append(candidates, listenerKey(network, "0.0.0.0", port)) + } + for _, key := range candidates { + if prev, ok := used[key]; ok { + return prev, true + } + } + return "", false +} + // LoadConfig loads configuration from a file path or URL. // If path starts with "http://" or "https://", it fetches the config from the URL. // Otherwise, it reads from the local file system. @@ -92,12 +232,9 @@ func LoadConfig(path string) (*Config, error) { return nil, err } - if cfg.Core.Hostname == "" { - hostname, err := os.Hostname() - if err != nil { - hostname = "unknown" - } - cfg.Core.Hostname = hostname + cfg.ApplyDefaults() + if err := cfg.Validate(); err != nil { + return nil, err } return cfg, nil @@ -139,12 +276,9 @@ func loadConfigFromURL(url string) (*Config, error) { return nil, fmt.Errorf("failed to decode TOML config from %s: %w", url, err) } - if cfg.Core.Hostname == "" { - hostname, err := os.Hostname() - if err != nil { - hostname = "unknown" - } - cfg.Core.Hostname = hostname + cfg.ApplyDefaults() + if err := cfg.Validate(); err != nil { + return nil, err } return cfg, nil diff --git a/core/config_test.go b/core/config_test.go new file mode 100644 index 0000000..615b789 --- /dev/null +++ b/core/config_test.go @@ -0,0 +1,111 @@ +package core + +import ( + "strings" + "testing" +) + +func TestConnectRuleBindIP(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + rule ConnectRule + want string + }{ + { + name: "default local only", + rule: ConnectRule{Protocol: "tcp"}, + want: "127.0.0.1", + }, + { + name: "explicit local addr", + rule: ConnectRule{Protocol: "udp", LocalAddr: "192.168.1.10"}, + want: "192.168.1.10", + }, + { + name: "minecraft exposes LAN by default", + rule: ConnectRule{Protocol: "minecraft"}, + want: "0.0.0.0", + }, + { + name: "lan enable exposes LAN", + rule: ConnectRule{Protocol: "tcp", LanEnable: boolPtr(true)}, + want: "0.0.0.0", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if got := tt.rule.BindIP(); got != tt.want { + t.Fatalf("BindIP() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestConfigValidateAcceptsValidConfig(t *testing.T) { + t.Parallel() + + cfg := Config{ + Core: Core{AuthKey: "tskey-auth-example"}, + Forward: map[string][]ForwardRule{ + "web": { + {Protocol: "tcp", TailscalePort: 8080, LocalAddr: "127.0.0.1:9090"}, + {Protocol: "udp", TailscalePort: 8080, LocalAddr: "127.0.0.1:9090"}, + }, + }, + Connect: map[string][]ConnectRule{ + "api": { + {Protocol: "tcp", LocalPort: 9000, DstAddr: "host.ts.net:8080"}, + {Protocol: "udp", LocalPort: 9000, DstAddr: "host.ts.net:8080"}, + }, + }, + } + + if err := cfg.Validate(); err != nil { + t.Fatalf("Validate() returned error: %v", err) + } +} + +func TestConfigValidateRejectsInvalidConfig(t *testing.T) { + t.Parallel() + + cfg := Config{ + Core: Core{}, + Forward: map[string][]ForwardRule{ + "bad": { + {Protocol: "icmp", TailscalePort: 70000, LocalAddr: "127.0.0.1"}, + }, + }, + Connect: map[string][]ConnectRule{ + "bad": { + {Protocol: "tcp", LocalPort: 9000, DstAddr: "host.ts.net:8080"}, + {Protocol: "minecraft", LocalPort: 9000, DstAddr: "host.ts.net:25565"}, + }, + }, + } + + err := cfg.Validate() + if err == nil { + t.Fatal("Validate() returned nil, want error") + } + + for _, want := range []string{ + "core.auth_key is required", + "forward.bad[0].protocol must be tcp or udp", + "forward.bad[0].tailscale_port must be between 1 and 65535", + "forward.bad[0].local_addr invalid", + "connect.bad[1] local listener duplicates connect.bad[0]", + } { + if !strings.Contains(err.Error(), want) { + t.Fatalf("Validate() error %q does not contain %q", err.Error(), want) + } + } +} + +func boolPtr(v bool) *bool { + return &v +}