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 +} diff --git a/core/connector.go b/core/connector.go new file mode 100644 index 0000000..bd28ace --- /dev/null +++ b/core/connector.go @@ -0,0 +1,39 @@ +package core + +import ( + "context" + "log/slog" + + "tailscale.com/tsnet" +) + +func StartConnectors(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule) { + for tag, rrs := range rules { + for _, rule := range rrs { + args := []any{ + slog.String("tag", tag), + slog.String("protocol", rule.Protocol), + slog.Int("local_port", rule.LocalPort), + slog.String("dst_addr", rule.DstAddr), + } + if rule.LocalAddr != "" { + args = append(args, slog.String("local_addr", rule.LocalAddr)) + } + slog.Info("starting connector", args...) + go runConnector(ctx, srv, rule, tag) + } + } +} + +func runConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, tag string) { + logger := RuleLogger(rule, tag) + + switch rule.Protocol { + case "tcp", "minecraft": + runTCPConnector(ctx, srv, rule, logger) + case "udp": + runUDPConnector(ctx, srv, rule, logger) + default: + logger.Error("unsupported protocol, expected tcp or udp") + } +} diff --git a/core/dial.go b/core/dial.go new file mode 100644 index 0000000..d6284eb --- /dev/null +++ b/core/dial.go @@ -0,0 +1,27 @@ +package core + +import ( + "context" + "net" + "time" + + "tailscale.com/tsnet" +) + +const dialTimeout = 10 * time.Second + +func dialTCP(ctx context.Context, addr string) (net.Conn, error) { + dialer := net.Dialer{Timeout: dialTimeout} + return dialer.DialContext(ctx, "tcp", addr) +} + +func dialUDP(ctx context.Context, addr string) (net.Conn, error) { + dialer := net.Dialer{Timeout: dialTimeout} + return dialer.DialContext(ctx, "udp", addr) +} + +func dialTsnet(ctx context.Context, srv *tsnet.Server, network, addr string) (net.Conn, error) { + dialCtx, cancel := context.WithTimeout(ctx, dialTimeout) + defer cancel() + return srv.Dial(dialCtx, network, addr) +} diff --git a/core/dns.go b/core/dns.go new file mode 100644 index 0000000..97a310e --- /dev/null +++ b/core/dns.go @@ -0,0 +1,154 @@ +package core + +import ( + "context" + "errors" + "fmt" + "net/netip" + "strings" + "sync" + "time" + + "golang.org/x/net/dns/dnsmessage" + "tailscale.com/ipn/ipnstate" + "tailscale.com/net/dns" + "tailscale.com/tsnet" +) + +// magicdns suffix cache +var ( + magicDNSSuffixMu sync.RWMutex + magicDNSSuffix string +) + +func SetMagicDNSSuffix(raw string) { + magicDNSSuffixMu.Lock() + defer magicDNSSuffixMu.Unlock() + magicDNSSuffix = strings.Trim(raw, ".") +} + +func GetMagicDNSSuffix() (string, bool) { + magicDNSSuffixMu.RLock() + defer magicDNSSuffixMu.RUnlock() + if magicDNSSuffix == "" { + return "", false + } + return magicDNSSuffix, true +} + +func GetMagicDNSSuffixFromStatus(st *ipnstate.Status) (string, error) { + suffix := st.CurrentTailnet.MagicDNSSuffix + suffix = strings.Trim(suffix, ".") + if suffix == "" { + return "", errors.New("magic dns suffix not found in status") + } + return suffix, nil +} + +// addr(ip or domain) to tailscale ip +// check the address is in the tailscale network +func resolveAddr(ctx context.Context, srv *tsnet.Server, addr string) (*netip.Addr, error) { + lc, err := srv.LocalClient() + if err != nil { + return nil, err + } + stat, err := lc.Status(ctx) + if err != nil { + return nil, err + } + + if ip, err := netip.ParseAddr(addr); err == nil { + for _, peer := range stat.Peer { + for _, ipRange := range peer.AllowedIPs.All() { + if ipRange.Contains(ip) { + return &peer.TailscaleIPs[0], nil + } + } + } + } else { + suffix, ok := GetMagicDNSSuffix() + if ok { + if !strings.HasSuffix(addr, suffix) { + dnsMgr, ok := srv.Sys().DNSManager.GetOK() + if !ok { + return nil, errors.New("DNS manager not available") + } + ipaddr, err := resolveHostViaResolver(dnsMgr, addr) + if err != nil { + return nil, err + } + return resolveAddr(ctx, srv, ipaddr.String()) + } + } + // addr is tailscale domain, resolve it + for _, peer := range stat.Peer { + dnsName := strings.TrimSuffix(peer.DNSName, ".") + if dnsName == addr { + return &peer.TailscaleIPs[0], nil + } + } + } + + return nil, errors.New(fmt.Sprintf("addr '%s' not found in tsnet", addr)) +} + +// resolveHostViaResolver resolves a hostname to a netip.Addr using the +// Tailscale DNS resolver. It queries A and AAAA records in a single +// message and follows CNAME chains (up to 8 levels deep). +func resolveHostViaResolver(resolver *dns.Manager, host string) (netip.Addr, error) { + return resolveHostWithDepth(resolver, host, 0) +} + +func resolveHostWithDepth(r *dns.Manager, host string, depth int) (netip.Addr, error) { + const maxCNAMEChase = 8 + if depth > maxCNAMEChase { + return netip.Addr{}, fmt.Errorf("CNAME chain too deep for %s", host) + } + + name, err := dnsmessage.NewName(host + ".") + if err != nil { + return netip.Addr{}, fmt.Errorf("invalid hostname %s: %w", host, err) + } + + msg := dnsmessage.Message{ + Header: dnsmessage.Header{RecursionDesired: true}, + Questions: []dnsmessage.Question{ + {Name: name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}, + }, + } + queryBytes, err := msg.Pack() + if err != nil { + return netip.Addr{}, fmt.Errorf("failed to pack DNS query: %w", err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + respBytes, err := r.Query(ctx, queryBytes, "udp", netip.AddrPort{}) + if err != nil { + return netip.Addr{}, fmt.Errorf("DNS resolution failed for %s: %w", host, err) + } + + var resp dnsmessage.Message + if err := resp.Unpack(respBytes); err != nil { + return netip.Addr{}, fmt.Errorf("failed to unpack DNS response: %w", err) + } + + var cnameTarget string + for _, ans := range resp.Answers { + switch r := ans.Body.(type) { + case *dnsmessage.AResource: + if ip := netip.AddrFrom4(r.A); ip.IsValid() { + return ip, nil + } + case *dnsmessage.CNAMEResource: + cnameTarget = strings.TrimSuffix(r.CNAME.String(), ".") + } + } + + // Follow CNAME if no direct A found + if cnameTarget != "" { + return resolveHostWithDepth(r, cnameTarget, depth+1) + } + + return netip.Addr{}, fmt.Errorf("no A/AAAA record found for %s", host) +} diff --git a/core/forwarder.go b/core/forwarder.go index 50a1c95..c676841 100644 --- a/core/forwarder.go +++ b/core/forwarder.go @@ -2,21 +2,12 @@ package core import ( "context" - "errors" "fmt" - "io" "log/slog" - "net" - "net/netip" - "sync" - "time" - "tailscale.com/ipn/ipnstate" "tailscale.com/tsnet" ) -const udpForwardIdleTimeout = 2 * time.Minute - func StartForwarders(ctx context.Context, srv *tsnet.Server, rules map[string][]ForwardRule) { for tag, rrs := range rules { for _, rule := range rrs { @@ -31,24 +22,6 @@ func StartForwarders(ctx context.Context, srv *tsnet.Server, rules map[string][] } } -func StartConnectors(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule) { - for tag, rrs := range rules { - for _, rule := range rrs { - args := []any{ - slog.String("tag", tag), - slog.String("protocol", rule.Protocol), - slog.Int("local_port", rule.LocalPort), - slog.String("dst_addr", rule.DstAddr), - } - if rule.LocalAddr != "" { - args = append(args, slog.String("local_addr", rule.LocalAddr)) - } - slog.Info("starting connector", args...) - go runConnector(ctx, srv, rule, tag) - } - } -} - func RuleLogger(rule any, tag string) *slog.Logger { var args []any switch r := rule.(type) { @@ -90,567 +63,3 @@ func runForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, tag logger.Error("unsupported protocol, expected tcp or udp") } } - -func runTCPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, logger *slog.Logger) { - ip := getSelfTsnetAddr(srv) - ln, err := srv.Listen("tcp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort)) - if err != nil { - logger.Error("failed to listen", "error", err) - return - } - logger.Debug("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort))) - - go func() { - <-ctx.Done() - ln.Close() - }() - - for { - conn, err := ln.Accept() - if err != nil { - if ctx.Err() != nil { - return - } - logger.Error("accept error", "error", err) - continue - } - go handleTCPForward(ctx, srv, conn, rule, logger) - } -} - -func handleTCPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ForwardRule, logger *slog.Logger) { - remoteAddrStr := conn.RemoteAddr().String() - clog := logger.With(slog.String("remote", remoteAddrStr)) - - lc, err := srv.LocalClient() - if err == nil { - who, err := lc.WhoIs(ctx, remoteAddrStr) - if err == nil { - clog = clog.With(slog.String("user", who.UserProfile.LoginName)) - } - } - - connType := getConnType(ctx, srv, remoteAddrStr) - clog.Info("accepted connection", - slog.String("conn_type", connType), - slog.String("local_addr", rule.LocalAddr), - ) - - localConn, err := net.Dial("tcp", rule.LocalAddr) - if err != nil { - clog.Error("failed to dial local", "error", err) - conn.Close() - return - } - - stop := context.AfterFunc(ctx, func() { - conn.Close() - localConn.Close() - }) - defer stop() - - toLocal, toTs := pipeConns(conn, localConn) - clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs)) -} - -var statusCache struct { - mu sync.Mutex - status *ipnstate.Status - expires time.Time -} - -const statusCacheTTL = 5 * time.Second - -func getCachedStatus(ctx context.Context, srv *tsnet.Server) (*ipnstate.Status, error) { - statusCache.mu.Lock() - if statusCache.status != nil && time.Now().Before(statusCache.expires) { - st := statusCache.status - statusCache.mu.Unlock() - return st, nil - } - statusCache.mu.Unlock() - - lc, err := srv.LocalClient() - if err != nil { - return nil, err - } - st, err := lc.Status(ctx) - if err != nil { - return nil, err - } - - statusCache.mu.Lock() - statusCache.status = st - statusCache.expires = time.Now().Add(statusCacheTTL) - statusCache.mu.Unlock() - return st, nil -} - -func isTsnetTarget(host string) bool { - if ip, err := netip.ParseAddr(host); err == nil { - tsnetV4 := netip.MustParsePrefix("100.64.0.0/10") - tsnetV6 := netip.MustParsePrefix("fd7a:115c:a1e0::/48") - return tsnetV4.Contains(ip) || tsnetV6.Contains(ip) - } - return true -} - -func getConnType(ctx context.Context, srv *tsnet.Server, remoteAddrStr string) string { - st, err := getCachedStatus(ctx, srv) - if err != nil { - return "unknown" - } - - remoteHost, _, err := net.SplitHostPort(remoteAddrStr) - if err != nil { - return "unknown" - } - - for _, peer := range st.Peer { - for _, addr := range peer.TailscaleIPs { - if addr.String() == remoteHost { - if peer.CurAddr != "" { - return "direct" - } - if peer.Relay != "" { - return fmt.Sprintf("derp(%s)", peer.Relay) - } - return "direct" - } - } - } - return "unknown" -} - -func runUDPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, logger *slog.Logger) { - ip := getSelfTsnetAddr(srv) - ln, err := srv.Listen("udp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort)) - if err != nil { - logger.Error("failed to listen", "error", err) - return - } - logger.Debug("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort))) - - go func() { - <-ctx.Done() - ln.Close() - }() - - for { - conn, err := ln.Accept() - if err != nil { - if ctx.Err() != nil { - return - } - logger.Error("accept error", "error", err) - continue - } - go handleUDPForward(ctx, srv, conn, rule, logger) - } -} - -func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ForwardRule, logger *slog.Logger) { - remoteAddrStr := conn.RemoteAddr().String() - clog := logger.With(slog.String("remote", remoteAddrStr)) - - lc, err := srv.LocalClient() - if err == nil { - who, err := lc.WhoIs(ctx, remoteAddrStr) - if err == nil { - clog = clog.With(slog.String("user", who.UserProfile.LoginName)) - } - } - - connType := getConnType(ctx, srv, remoteAddrStr) - clog.Info("accepted connection", - slog.String("conn_type", connType), - slog.String("local_addr", rule.LocalAddr), - ) - - localConn, err := net.Dial("udp", rule.LocalAddr) - if err != nil { - clog.Error("failed to dial local", "error", err) - conn.Close() - return - } - - stop := context.AfterFunc(ctx, func() { - conn.Close() - localConn.Close() - }) - defer stop() - - remoteIP, _, _ := net.SplitHostPort(remoteAddrStr) - - var toTs, toLocal int64 - var wg sync.WaitGroup - wg.Add(2) - - go func() { - defer wg.Done() - buf := make([]byte, 65535) - for { - _ = conn.SetReadDeadline(time.Now().Add(udpForwardIdleTimeout)) - n, err := conn.Read(buf) - if err != nil { - if netErr, ok := err.(net.Error); ok && netErr.Timeout() { - clog.Debug("udp forward idle timeout on ts side") - } else if !errors.Is(err, net.ErrClosed) && !errors.Is(err, io.EOF) { - clog.Debug("udp forward read from ts", "error", err) - } - localConn.Close() - return - } - _ = conn.SetReadDeadline(time.Time{}) - toLocal += int64(n) - clog.Debug("inbound udp packet", - slog.String("from_ip", remoteIP), - slog.String("to_ip", rule.LocalAddr), - slog.Int("pkg_size", n), - ) - if _, err := localConn.Write(buf[:n]); err != nil { - conn.Close() - return - } - } - }() - - go func() { - defer wg.Done() - buf := make([]byte, 65535) - for { - _ = localConn.SetReadDeadline(time.Now().Add(udpForwardIdleTimeout)) - n, err := localConn.Read(buf) - if err != nil { - if netErr, ok := err.(net.Error); ok && netErr.Timeout() { - clog.Debug("udp forward idle timeout on local side") - } else if !errors.Is(err, net.ErrClosed) && !errors.Is(err, io.EOF) { - clog.Debug("udp forward read from local", "error", err) - } - conn.Close() - return - } - _ = localConn.SetReadDeadline(time.Time{}) - toTs += int64(n) - localIP, _, _ := net.SplitHostPort(localConn.RemoteAddr().String()) - clog.Debug("outbound udp packet", - slog.String("from_ip", localIP), - slog.String("to_ip", remoteIP), - slog.Int("pkg_size", n), - ) - if _, err := conn.Write(buf[:n]); err != nil { - localConn.Close() - return - } - } - }() - - wg.Wait() - clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs)) -} - -const udpRelayMaxSessions = 1024 - -type udpSession struct { - conn net.Conn - remote net.Addr - lastUse time.Time -} - -type udpRelay struct { - listenConn net.PacketConn - dialAddr string - logger *slog.Logger - direction string - srv *tsnet.Server - - mu sync.Mutex - sessions map[string]*udpSession -} - -func (r *udpRelay) run(ctx context.Context) { - go func() { - <-ctx.Done() - r.listenConn.Close() - }() - - go func() { - ticker := time.NewTicker(2 * time.Minute) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - r.cleanup() - } - } - }() - - buf := make([]byte, 65535) - for { - select { - case <-ctx.Done(): - return - default: - } - - n, from, err := r.listenConn.ReadFrom(buf) - if err != nil { - if ctx.Err() != nil { - return - } - r.logger.Error("udp read error", "error", err) - return - } - - key := from.String() - var toIP string - r.mu.Lock() - sess, exists := r.sessions[key] - if !exists { - if len(r.sessions) >= udpRelayMaxSessions { - r.mu.Unlock() - r.logger.Warn("udp relay session limit reached, dropping packet", - slog.Int("limit", udpRelayMaxSessions), - slog.String("remote", key), - ) - continue - } - host, _, err := net.SplitHostPort(r.dialAddr) - if err != nil { - r.mu.Unlock() - r.logger.Error("failed to parse dial addr", "error", err) - continue - } - inTsnet := isTsnetTarget(host) - - r.mu.Unlock() - var dialed net.Conn - if inTsnet { - dialed, err = r.srv.Dial(ctx, "udp", r.dialAddr) - } else { - dialed, err = net.Dial("udp", r.dialAddr) - } - if err != nil { - r.logger.Error("failed to dial", "error", err) - continue - } - sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now()} - r.mu.Lock() - if existing, dup := r.sessions[key]; dup { - dialed.Close() - sess = existing - sess.lastUse = time.Now() - } else { - r.sessions[key] = sess - } - toIP = sess.conn.RemoteAddr().String() - r.mu.Unlock() - - r.logger.Info("new udp session", slog.String("remote", key), slog.String("direction", r.direction)) - go r.readSession(key, sess) - } else { - sess.lastUse = time.Now() - toIP = sess.conn.RemoteAddr().String() - r.mu.Unlock() - } - - fromIP, _, _ := net.SplitHostPort(from.String()) - toIPHost, _, _ := net.SplitHostPort(toIP) - r.logger.Debug("outbound udp packet", - slog.String("from_ip", fromIP), - slog.String("to_ip", toIPHost), - slog.Int("pkg_size", n), - ) - if _, err := sess.conn.Write(buf[:n]); err != nil { - r.logger.Error("failed to write", "error", err) - r.removeSession(key) - } - } -} - -func (r *udpRelay) readSession(key string, sess *udpSession) { - buf := make([]byte, 65535) - for { - n, err := sess.conn.Read(buf) - if err != nil { - r.removeSession(key) - return - } - fromIP, _, _ := net.SplitHostPort(sess.conn.RemoteAddr().String()) - toIP, _, _ := net.SplitHostPort(sess.remote.String()) - r.logger.Info("udp packet", - slog.String("from_ip", fromIP), - slog.String("to_ip", toIP), - slog.Int("pkg_size", n), - ) - if _, err := r.listenConn.WriteTo(buf[:n], sess.remote); err != nil { - r.logger.Error("failed to write back", "error", err) - r.removeSession(key) - return - } - sess.lastUse = time.Now() - } -} - -func (r *udpRelay) removeSession(key string) { - r.mu.Lock() - defer r.mu.Unlock() - if s, ok := r.sessions[key]; ok { - remote := s.remote.String() - s.conn.Close() - delete(r.sessions, key) - r.logger.Debug("udp session closed", slog.String("remote", remote)) - } -} - -func (r *udpRelay) cleanup() { - r.mu.Lock() - defer r.mu.Unlock() - threshold := time.Now().Add(-5 * time.Minute) - for key, s := range r.sessions { - if s.lastUse.Before(threshold) { - remote := s.remote.String() - s.conn.Close() - delete(r.sessions, key) - r.logger.Debug("udp session cleaned up", slog.String("remote", remote)) - } - } -} - -func runConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, tag string) { - logger := RuleLogger(rule, tag) - - switch rule.Protocol { - case "tcp", "minecraft": - runTCPConnector(ctx, srv, rule, logger) - case "udp": - runUDPConnector(ctx, srv, rule, logger) - default: - logger.Error("unsupported protocol, expected tcp or udp") - } -} - -func runTCPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, logger *slog.Logger) { - bindIP := rule.LocalAddr - if bindIP == "" { - bindIP = "0.0.0.0" - } - if rule.LANEnabled() && bindIP != "0.0.0.0" { - logger.Warn("lan_enable forces local_addr to 0.0.0.0, overriding") - bindIP = "0.0.0.0" - } - addr := fmt.Sprintf("%s:%d", bindIP, rule.LocalPort) - ln, err := net.Listen("tcp", addr) - if err != nil { - logger.Error("failed to listen locally", "error", err) - return - } - logger.Info("listening", slog.String("on", addr)) - - go func() { - <-ctx.Done() - ln.Close() - }() - - for { - conn, err := ln.Accept() - if err != nil { - if ctx.Err() != nil { - return - } - logger.Error("accept error", "error", err) - continue - } - go handleTCPConnect(ctx, srv, conn, rule, logger) - } -} - -func handleTCPConnect(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ConnectRule, logger *slog.Logger) { - clog := logger.With(slog.String("local_client", conn.RemoteAddr().String())) - - tsConn, err := srv.Dial(ctx, "tcp", rule.DstAddr) - if err != nil { - clog.Error("failed to dial tailscale", "error", err) - conn.Close() - return - } - - stop := context.AfterFunc(ctx, func() { - conn.Close() - tsConn.Close() - }) - defer stop() - - clog.Info("accepted connection", slog.String("dst_addr", rule.DstAddr)) - toConn, toTs := pipeConns(conn, tsConn) - clog.Info("connection closed", slog.Int64("ts_rx_bytes", toTs), slog.Int64("ts_tx_bytes", toConn)) -} - -func runUDPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, logger *slog.Logger) { - bindIP := rule.LocalAddr - if bindIP == "" { - bindIP = "0.0.0.0" - } - addr := fmt.Sprintf("%s:%d", bindIP, rule.LocalPort) - addrUDP, err := net.ResolveUDPAddr("udp", addr) - if err != nil { - logger.Error("failed to resolve local addr", "error", err) - return - } - - pc, err := net.ListenUDP("udp", addrUDP) - if err != nil { - logger.Error("failed to listen locally", "error", err) - return - } - logger.Info("listening", slog.String("on", addr)) - - relay := &udpRelay{ - listenConn: pc, - dialAddr: rule.DstAddr, - logger: logger, - direction: "tailscale", - srv: srv, - sessions: make(map[string]*udpSession), - } - relay.run(ctx) -} - -func pipeConns(a, b net.Conn) (toA, toB int64) { - done := make(chan struct{}, 2) - var aToB, bToA int64 - - go func() { - defer func() { done <- struct{}{} }() - n, err := io.Copy(a, b) - aToB = n - if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) { - slog.Debug("pipe copy error", "direction", "b->a", "error", err) - } - if tc, ok := a.(*net.TCPConn); ok { - tc.CloseWrite() - } else { - a.Close() - } - }() - - go func() { - defer func() { done <- struct{}{} }() - n, err := io.Copy(b, a) - bToA = n - if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) { - slog.Debug("pipe copy error", "direction", "a->b", "error", err) - } - if tc, ok := b.(*net.TCPConn); ok { - tc.CloseWrite() - } else { - b.Close() - } - }() - - <-done - <-done - return aToB, bToA -} diff --git a/core/peer_status.go b/core/peer_status.go new file mode 100644 index 0000000..66ee7cc --- /dev/null +++ b/core/peer_status.go @@ -0,0 +1,82 @@ +package core + +import ( + "context" + "fmt" + "net" + "net/netip" + "sync" + "time" + + "tailscale.com/ipn/ipnstate" + "tailscale.com/tsnet" +) + +var statusCache struct { + mu sync.Mutex + status *ipnstate.Status + expires time.Time +} + +const statusCacheTTL = 5 * time.Second + +func getCachedStatus(ctx context.Context, srv *tsnet.Server) (*ipnstate.Status, error) { + statusCache.mu.Lock() + if statusCache.status != nil && time.Now().Before(statusCache.expires) { + st := statusCache.status + statusCache.mu.Unlock() + return st, nil + } + statusCache.mu.Unlock() + + lc, err := srv.LocalClient() + if err != nil { + return nil, err + } + st, err := lc.Status(ctx) + if err != nil { + return nil, err + } + + statusCache.mu.Lock() + statusCache.status = st + statusCache.expires = time.Now().Add(statusCacheTTL) + statusCache.mu.Unlock() + return st, nil +} + +func isTsnetTarget(host string) bool { + if ip, err := netip.ParseAddr(host); err == nil { + tsnetV4 := netip.MustParsePrefix("100.64.0.0/10") + tsnetV6 := netip.MustParsePrefix("fd7a:115c:a1e0::/48") + return tsnetV4.Contains(ip) || tsnetV6.Contains(ip) + } + return true +} + +func getConnType(ctx context.Context, srv *tsnet.Server, remoteAddrStr string) string { + st, err := getCachedStatus(ctx, srv) + if err != nil { + return "unknown" + } + + remoteHost, _, err := net.SplitHostPort(remoteAddrStr) + if err != nil { + return "unknown" + } + + for _, peer := range st.Peer { + for _, addr := range peer.TailscaleIPs { + if addr.String() == remoteHost { + if peer.CurAddr != "" { + return "direct" + } + if peer.Relay != "" { + return fmt.Sprintf("derp(%s)", peer.Relay) + } + return "direct" + } + } + } + return "unknown" +} diff --git a/core/pipe.go b/core/pipe.go new file mode 100644 index 0000000..9dfccfe --- /dev/null +++ b/core/pipe.go @@ -0,0 +1,45 @@ +package core + +import ( + "errors" + "io" + "log/slog" + "net" +) + +func pipeConns(a, b net.Conn) (toA, toB int64) { + done := make(chan struct{}, 2) + var aToB, bToA int64 + + go func() { + defer func() { done <- struct{}{} }() + n, err := io.Copy(a, b) + aToB = n + if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) { + slog.Debug("pipe copy error", "direction", "b->a", "error", err) + } + if tc, ok := a.(*net.TCPConn); ok { + tc.CloseWrite() + } else { + a.Close() + } + }() + + go func() { + defer func() { done <- struct{}{} }() + n, err := io.Copy(b, a) + bToA = n + if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) { + slog.Debug("pipe copy error", "direction", "a->b", "error", err) + } + if tc, ok := b.(*net.TCPConn); ok { + tc.CloseWrite() + } else { + b.Close() + } + }() + + <-done + <-done + return aToB, bToA +} diff --git a/core/tcp.go b/core/tcp.go new file mode 100644 index 0000000..02e82a5 --- /dev/null +++ b/core/tcp.go @@ -0,0 +1,124 @@ +package core + +import ( + "context" + "fmt" + "log/slog" + "net" + + "tailscale.com/tsnet" +) + +func runTCPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, logger *slog.Logger) { + ip := getSelfTsnetAddr(srv) + ln, err := srv.Listen("tcp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort)) + if err != nil { + logger.Error("failed to listen", "error", err) + return + } + logger.Debug("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort))) + + go func() { + <-ctx.Done() + ln.Close() + }() + + for { + conn, err := ln.Accept() + if err != nil { + if ctx.Err() != nil { + return + } + logger.Error("accept error", "error", err) + continue + } + go handleTCPForward(ctx, srv, conn, rule, logger) + } +} + +func handleTCPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ForwardRule, logger *slog.Logger) { + remoteAddrStr := conn.RemoteAddr().String() + clog := logger.With(slog.String("remote", remoteAddrStr)) + + lc, err := srv.LocalClient() + if err == nil { + who, err := lc.WhoIs(ctx, remoteAddrStr) + if err == nil { + clog = clog.With(slog.String("user", who.UserProfile.LoginName)) + } + } + + connType := getConnType(ctx, srv, remoteAddrStr) + clog.Info("accepted connection", + slog.String("conn_type", connType), + slog.String("local_addr", rule.LocalAddr), + ) + + localConn, err := dialTCP(ctx, rule.LocalAddr) + if err != nil { + clog.Error("failed to dial local", "error", err) + conn.Close() + return + } + + stop := context.AfterFunc(ctx, func() { + conn.Close() + localConn.Close() + }) + defer stop() + + toLocal, toTs := pipeConns(conn, localConn) + clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs)) +} + +func runTCPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, logger *slog.Logger) { + bindIP := rule.BindIP() + if rule.LANEnabled() && rule.LocalAddr != "" && rule.LocalAddr != "0.0.0.0" { + logger.Warn("lan_enable forces local_addr to 0.0.0.0, overriding") + } + addr := fmt.Sprintf("%s:%d", bindIP, rule.LocalPort) + ln, err := net.Listen("tcp", addr) + if err != nil { + logger.Error("failed to listen locally", "error", err) + return + } + logger.Debug("listening", slog.String("on", addr)) + + go func() { + <-ctx.Done() + ln.Close() + }() + + for { + conn, err := ln.Accept() + if err != nil { + if ctx.Err() != nil { + return + } + logger.Error("accept error", "error", err) + continue + } + go handleTCPConnect(ctx, srv, conn, rule, logger) + } +} + +func handleTCPConnect(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ConnectRule, logger *slog.Logger) { + clog := logger.With(slog.String("local_client", conn.RemoteAddr().String())) + + tsConn, err := dialTsnet(ctx, srv, "tcp", rule.DstAddr) + if err != nil { + clog.Error("failed to dial tailscale", "error", err) + conn.Close() + return + } + + stop := context.AfterFunc(ctx, func() { + conn.Close() + tsConn.Close() + }) + defer stop() + + clog.Info("accepted connection", slog.String("dst_addr", rule.DstAddr)) + toConn, toTs := pipeConns(conn, tsConn) + clog.Info("connection closed", slog.Int64("ts_rx_bytes", toTs), slog.Int64("ts_tx_bytes", toConn)) +} diff --git a/core/udp.go b/core/udp.go new file mode 100644 index 0000000..184ce1b --- /dev/null +++ b/core/udp.go @@ -0,0 +1,170 @@ +package core + +import ( + "context" + "errors" + "fmt" + "io" + "log/slog" + "net" + "sync" + "time" + + "tailscale.com/tsnet" +) + +const udpForwardIdleTimeout = 2 * time.Minute + +func runUDPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, logger *slog.Logger) { + ip := getSelfTsnetAddr(srv) + ln, err := srv.Listen("udp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort)) + if err != nil { + logger.Error("failed to listen", "error", err) + return + } + logger.Debug("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort))) + + go func() { + <-ctx.Done() + ln.Close() + }() + + for { + conn, err := ln.Accept() + if err != nil { + if ctx.Err() != nil { + return + } + logger.Error("accept error", "error", err) + continue + } + go handleUDPForward(ctx, srv, conn, rule, logger) + } +} + +func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ForwardRule, logger *slog.Logger) { + remoteAddrStr := conn.RemoteAddr().String() + clog := logger.With(slog.String("remote", remoteAddrStr)) + + lc, err := srv.LocalClient() + if err == nil { + who, err := lc.WhoIs(ctx, remoteAddrStr) + if err == nil { + clog = clog.With(slog.String("user", who.UserProfile.LoginName)) + } + } + + connType := getConnType(ctx, srv, remoteAddrStr) + clog.Info("accepted connection", + slog.String("conn_type", connType), + slog.String("local_addr", rule.LocalAddr), + ) + + localConn, err := dialUDP(ctx, rule.LocalAddr) + if err != nil { + clog.Error("failed to dial local", "error", err) + conn.Close() + return + } + + stop := context.AfterFunc(ctx, func() { + conn.Close() + localConn.Close() + }) + defer stop() + + remoteIP, _, _ := net.SplitHostPort(remoteAddrStr) + + var toTs, toLocal int64 + var wg sync.WaitGroup + wg.Add(2) + + go func() { + defer wg.Done() + buf := make([]byte, 65535) + for { + _ = conn.SetReadDeadline(time.Now().Add(udpForwardIdleTimeout)) + n, err := conn.Read(buf) + if err != nil { + if netErr, ok := err.(net.Error); ok && netErr.Timeout() { + clog.Debug("udp forward idle timeout on ts side") + } else if !errors.Is(err, net.ErrClosed) && !errors.Is(err, io.EOF) { + clog.Debug("udp forward read from ts", "error", err) + } + localConn.Close() + return + } + _ = conn.SetReadDeadline(time.Time{}) + toLocal += int64(n) + clog.Debug("inbound udp packet", + slog.String("from_ip", remoteIP), + slog.String("to_ip", rule.LocalAddr), + slog.Int("pkg_size", n), + ) + if _, err := localConn.Write(buf[:n]); err != nil { + conn.Close() + return + } + } + }() + + go func() { + defer wg.Done() + buf := make([]byte, 65535) + for { + _ = localConn.SetReadDeadline(time.Now().Add(udpForwardIdleTimeout)) + n, err := localConn.Read(buf) + if err != nil { + if netErr, ok := err.(net.Error); ok && netErr.Timeout() { + clog.Debug("udp forward idle timeout on local side") + } else if !errors.Is(err, net.ErrClosed) && !errors.Is(err, io.EOF) { + clog.Debug("udp forward read from local", "error", err) + } + conn.Close() + return + } + _ = localConn.SetReadDeadline(time.Time{}) + toTs += int64(n) + localIP, _, _ := net.SplitHostPort(localConn.RemoteAddr().String()) + clog.Debug("outbound udp packet", + slog.String("from_ip", localIP), + slog.String("to_ip", remoteIP), + slog.Int("pkg_size", n), + ) + if _, err := conn.Write(buf[:n]); err != nil { + localConn.Close() + return + } + } + }() + + wg.Wait() + clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs)) +} + +func runUDPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, logger *slog.Logger) { + bindIP := rule.BindIP() + addr := fmt.Sprintf("%s:%d", bindIP, rule.LocalPort) + addrUDP, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + logger.Error("failed to resolve local addr", "error", err) + return + } + + pc, err := net.ListenUDP("udp", addrUDP) + if err != nil { + logger.Error("failed to listen locally", "error", err) + return + } + logger.Debug("listening", slog.String("on", addr)) + + relay := &udpRelay{ + listenConn: pc, + dialAddr: rule.DstAddr, + logger: logger, + direction: "tailscale", + srv: srv, + sessions: make(map[string]*udpSession), + } + relay.run(ctx) +} diff --git a/core/udp_relay.go b/core/udp_relay.go new file mode 100644 index 0000000..7731745 --- /dev/null +++ b/core/udp_relay.go @@ -0,0 +1,194 @@ +package core + +import ( + "context" + "log/slog" + "net" + "sync" + "time" + + "tailscale.com/tsnet" +) + +const udpRelayMaxSessions = 1024 + +type udpSession struct { + conn net.Conn + remote net.Addr + mu sync.Mutex + lastUse time.Time +} + +func (s *udpSession) touch() { + s.mu.Lock() + s.lastUse = time.Now() + s.mu.Unlock() +} + +func (s *udpSession) idleSince(threshold time.Time) bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.lastUse.Before(threshold) +} + +type udpRelay struct { + listenConn net.PacketConn + dialAddr string + logger *slog.Logger + direction string + srv *tsnet.Server + + mu sync.Mutex + sessions map[string]*udpSession +} + +func (r *udpRelay) run(ctx context.Context) { + go func() { + <-ctx.Done() + r.listenConn.Close() + }() + + go func() { + ticker := time.NewTicker(2 * time.Minute) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + r.cleanup() + } + } + }() + + buf := make([]byte, 65535) + for { + select { + case <-ctx.Done(): + return + default: + } + + n, from, err := r.listenConn.ReadFrom(buf) + if err != nil { + if ctx.Err() != nil { + return + } + r.logger.Error("udp read error", "error", err) + return + } + + key := from.String() + var toIP string + r.mu.Lock() + sess, exists := r.sessions[key] + if !exists { + if len(r.sessions) >= udpRelayMaxSessions { + r.mu.Unlock() + r.logger.Warn("udp relay session limit reached, dropping packet", + slog.Int("limit", udpRelayMaxSessions), + slog.String("remote", key), + ) + continue + } + host, _, err := net.SplitHostPort(r.dialAddr) + if err != nil { + r.mu.Unlock() + r.logger.Error("failed to parse dial addr", "error", err) + continue + } + inTsnet := isTsnetTarget(host) + + r.mu.Unlock() + var dialed net.Conn + if inTsnet { + dialed, err = dialTsnet(ctx, r.srv, "udp", r.dialAddr) + } else { + dialed, err = dialUDP(ctx, r.dialAddr) + } + if err != nil { + r.logger.Error("failed to dial", "error", err) + continue + } + sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now()} + r.mu.Lock() + if existing, dup := r.sessions[key]; dup { + dialed.Close() + sess = existing + sess.touch() + } else { + r.sessions[key] = sess + } + toIP = sess.conn.RemoteAddr().String() + r.mu.Unlock() + + r.logger.Info("new udp session", slog.String("remote", key), slog.String("direction", r.direction)) + go r.readSession(key, sess) + } else { + sess.touch() + toIP = sess.conn.RemoteAddr().String() + r.mu.Unlock() + } + + fromIP, _, _ := net.SplitHostPort(from.String()) + toIPHost, _, _ := net.SplitHostPort(toIP) + r.logger.Debug("outbound udp packet", + slog.String("from_ip", fromIP), + slog.String("to_ip", toIPHost), + slog.Int("pkg_size", n), + ) + if _, err := sess.conn.Write(buf[:n]); err != nil { + r.logger.Error("failed to write", "error", err) + r.removeSession(key) + } + } +} + +func (r *udpRelay) readSession(key string, sess *udpSession) { + buf := make([]byte, 65535) + for { + n, err := sess.conn.Read(buf) + if err != nil { + r.removeSession(key) + return + } + fromIP, _, _ := net.SplitHostPort(sess.conn.RemoteAddr().String()) + toIP, _, _ := net.SplitHostPort(sess.remote.String()) + r.logger.Info("udp packet", + slog.String("from_ip", fromIP), + slog.String("to_ip", toIP), + slog.Int("pkg_size", n), + ) + if _, err := r.listenConn.WriteTo(buf[:n], sess.remote); err != nil { + r.logger.Error("failed to write back", "error", err) + r.removeSession(key) + return + } + sess.touch() + } +} + +func (r *udpRelay) removeSession(key string) { + r.mu.Lock() + defer r.mu.Unlock() + if s, ok := r.sessions[key]; ok { + remote := s.remote.String() + s.conn.Close() + delete(r.sessions, key) + r.logger.Debug("udp session closed", slog.String("remote", remote)) + } +} + +func (r *udpRelay) cleanup() { + r.mu.Lock() + defer r.mu.Unlock() + threshold := time.Now().Add(-5 * time.Minute) + for key, s := range r.sessions { + if s.idleSince(threshold) { + remote := s.remote.String() + s.conn.Close() + delete(r.sessions, key) + r.logger.Debug("udp session cleaned up", slog.String("remote", remote)) + } + } +} diff --git a/core/utils.go b/core/utils.go index bde9a52..8573124 100644 --- a/core/utils.go +++ b/core/utils.go @@ -10,10 +10,8 @@ import ( "strings" "time" - "golang.org/x/net/dns/dnsmessage" "tailscale.com/client/local" "tailscale.com/ipn/ipnstate" - "tailscale.com/net/dns/resolver" "tailscale.com/tailcfg" "tailscale.com/tsnet" ) @@ -47,48 +45,27 @@ func StartTimeWatchDog(ctx context.Context, logger *slog.Logger) <-chan struct{} return ch } -func resolveAddr(ctx context.Context, srv *tsnet.Server, addr string) (*netip.Addr, error) { - lc, err := srv.LocalClient() - if err != nil { - return nil, err - } - stat, err := lc.Status(ctx) - if err != nil { - return nil, err - } - - if ip, err := netip.ParseAddr(addr); err == nil { - for _, peer := range stat.Peer { - for _, ipRange := range peer.AllowedIPs.All() { - if ipRange.Contains(ip) { - return &peer.TailscaleIPs[0], nil - } - } - } - } else { - // addr is domain, resolve it - for _, peer := range stat.Peer { - dnsName := strings.TrimSuffix(peer.DNSName, ".") - if dnsName == addr { - return &peer.TailscaleIPs[0], nil - } - } - } - - return nil, errors.New(fmt.Sprintf("addr '%s' not found in tsnet", addr)) -} - func getPeerFromRules(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule, logger *slog.Logger) ([]netip.Addr, error) { peerSet := make(map[netip.Addr]struct{}) - for _, rrs := range rules { + for tag, rrs := range rules { for _, rule := range rrs { rule := rule - ap, err := netip.ParseAddrPort(rule.DstAddr) + tag := tag + + ap, _, err := net.SplitHostPort(rule.DstAddr) if err != nil { + logger.Debug("error parsing rule", "tag", tag, "dst", rule.DstAddr, "err", err) continue } - peerSet[ap.Addr()] = struct{}{} + addr, err := resolveAddr(ctx, srv, ap) + + if err != nil { + logger.Warn("failed to resolve address", "err", err) + continue + } + logger.Debug("address found", "dst_addr", rule.DstAddr, "tag", tag, "address", addr) + peerSet[*addr] = struct{}{} } } @@ -182,7 +159,7 @@ func getSelfTsnetAddr(srv *tsnet.Server) netip.Addr { return ip } -func PresolveDstAddrWithSuffix(dst string, srv *tsnet.Server) (string, bool, error) { +func NormalizeDstAddrWithSuffix(ctx context.Context, srv *tsnet.Server, dst string) (string, bool, error) { host, port, err := net.SplitHostPort(dst) if err != nil { return dst, false, err @@ -192,98 +169,31 @@ func PresolveDstAddrWithSuffix(dst string, srv *tsnet.Server) (string, bool, err return dst, false, nil } - dnsMgr, ok := srv.Sys().DNSManager.GetOK() + suffix, ok := GetMagicDNSSuffix() if !ok { - return dst, false, errors.New("DNS manager not available") + return dst, false, nil } - addr, err := resolveHostViaResolver(dnsMgr.Resolver(), host) - if err != nil { - // tsnet magicdns failed; fall back to system DNS for non-tailnet domains - addr, err = fallbackSystemDNS(host) + + normalized := net.JoinHostPort(host+"."+suffix, port) + + // check domain exists before use + if strings.Contains(host, ".") { + _, err = resolveAddr(ctx, srv, normalized) if err != nil { - return dst, false, err + return dst, false, nil } } - return net.JoinHostPort(addr.String(), port), true, nil + + return normalized, true, nil } -// fallbackSystemDNS resolves a hostname via the standard system resolver. -// Returns the first usable IPv4 address (preferred) or IPv6 address. -func fallbackSystemDNS(host string) (netip.Addr, error) { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - ips, err := net.DefaultResolver.LookupNetIP(ctx, "ip", host) - if err != nil { - return netip.Addr{}, fmt.Errorf("system DNS resolution failed for %s: %w", host, err) - } - - for _, ip := range ips { - if ip.Is4() { - return ip, nil - } - } - // no IPv4 found, pick the first IPv6 - for _, ip := range ips { - if ip.Is6() { - return ip, nil - } - } - - return netip.Addr{}, fmt.Errorf("no valid IPs returned for %s", host) -} - -// resolveHostViaResolver resolves a hostname to a netip.Addr using the -// Tailscale DNS resolver. It queries A and AAAA records in a single -// message and follows CNAME chains (up to 8 levels deep). -func resolveHostViaResolver(resolver *resolver.Resolver, host string) (netip.Addr, error) { - name, err := dnsmessage.NewName(host + ".") - if err != nil { - return netip.Addr{}, fmt.Errorf("invalid hostname %s: %w", host, err) - } - - msg := dnsmessage.Message{ - Header: dnsmessage.Header{RecursionDesired: true}, - Questions: []dnsmessage.Question{ - {Name: name, Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}, - }, - } - queryBytes, err := msg.Pack() - if err != nil { - return netip.Addr{}, fmt.Errorf("failed to pack DNS query: %w", err) - } - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - respBytes, err := resolver.Query(ctx, queryBytes, "udp", netip.AddrPort{}) - if err != nil { - return netip.Addr{}, fmt.Errorf("DNS resolution failed for %s: %w", host, err) - } - - var resp dnsmessage.Message - if err := resp.Unpack(respBytes); err != nil { - return netip.Addr{}, fmt.Errorf("failed to unpack DNS response: %w", err) - } - - for _, ans := range resp.Answers { - switch r := ans.Body.(type) { - case *dnsmessage.AResource: - if ip := netip.AddrFrom4(r.A); ip.IsValid() { - return ip, nil - } - } - } - - return netip.Addr{}, fmt.Errorf("no A/AAAA record found for %s", host) -} - -func PresolveConnectRulesDstAddr(rules map[string][]ConnectRule, logger *slog.Logger, srv *tsnet.Server) { +func NormalizeConnectRulesDstAddr(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule, logger *slog.Logger) { for tag, rrs := range rules { for i := range rrs { rule := &rrs[i] - normalized, changed, err := PresolveDstAddrWithSuffix(rule.DstAddr, srv) + normalized, changed, err := NormalizeDstAddrWithSuffix(ctx, srv, rule.DstAddr) if err != nil { - logger.Warn("failed to resolve dst_addr", + logger.Debug("failed to normalize dst_addr", slog.String("tag", tag), slog.String("dst", rule.DstAddr), slog.String("error", err.Error()), @@ -291,7 +201,7 @@ func PresolveConnectRulesDstAddr(rules map[string][]ConnectRule, logger *slog.Lo continue } if changed { - logger.Debug("dst_addr resolved", + logger.Debug("dst_addr normalized with MagicDNS suffix", slog.String("tag", tag), slog.String("original", rule.DstAddr), slog.String("normalized", normalized), diff --git a/main.go b/main.go index 5a50d1f..0045d14 100644 --- a/main.go +++ b/main.go @@ -43,7 +43,7 @@ func serviceLogic(configPath string, isTsnetDebug bool, configURL string, logger } logger.Info("tsnet server initialized") - core.PresolveConnectRulesDstAddr(cfg.Connect, logger, srv) + core.NormalizeConnectRulesDstAddr(ctx, srv, cfg.Connect, logger) core.StartForwarders(ctx, srv, cfg.Forward) core.StartConnectors(ctx, srv, cfg.Connect)