feat(config): add config validator to avoid wrong settings
This commit is contained in:
+2
-1
@@ -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"
|
||||
dst_addr = "any-client-in.ts.net:24454"
|
||||
|
||||
+146
-12
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user