Merge branch 'fix-multi-layer-magicdns-resolve'
# Conflicts: # core/utils.go # main.go
This commit is contained in:
+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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
+154
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+124
@@ -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))
|
||||
}
|
||||
+170
@@ -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)
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
}
|
||||
+29
-119
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user