From c23ca5025c19844618cebfee2b5de3e3f333eba9 Mon Sep 17 00:00:00 2001 From: nullcat Date: Mon, 25 May 2026 00:07:01 +0800 Subject: [PATCH] fix(forward): udp relay not using tsnet server context --- core/forwarder.go | 147 +++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 133 insertions(+), 14 deletions(-) diff --git a/core/forwarder.go b/core/forwarder.go index e119ed1..fde3085 100644 --- a/core/forwarder.go +++ b/core/forwarder.go @@ -6,6 +6,7 @@ import ( "io" "log/slog" "net" + "net/netip" "sync" "time" @@ -182,27 +183,114 @@ func runUDPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, l if !ip.IsValid() { ip = ip6 } - pc, err := srv.ListenPacket("udp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort)) + 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.Info("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort))) - relay := &udpRelay{ - listenConn: pc, - dialAddr: rule.LocalAddr, - logger: logger, - direction: "local", - sessions: make(map[string]*udpSession), + 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) } - relay.run(ctx) +} + +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), + ) + + defer conn.Close() + + localConn, err := net.Dial("udp", rule.LocalAddr) + if err != nil { + clog.Error("failed to dial local", "error", err) + return + } + defer localConn.Close() + + 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 { + n, err := conn.Read(buf) + if err != nil { + return + } + 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 { + return + } + } + }() + + go func() { + defer wg.Done() + buf := make([]byte, 65535) + for { + n, err := localConn.Read(buf) + if err != nil { + return + } + 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 { + return + } + } + }() + + wg.Wait() + clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs)) } type udpSession struct { - conn net.Conn - remote net.Addr - lastUse time.Time + conn net.Conn + remote net.Addr + inTailscale bool + lastUse time.Time } type udpRelay struct { @@ -210,6 +298,8 @@ type udpRelay struct { dialAddr string logger *slog.Logger direction string + srv *tsnet.Server + ctx context.Context mu sync.Mutex sessions map[string]*udpSession @@ -255,23 +345,43 @@ func (r *udpRelay) run(ctx context.Context) { r.mu.Lock() sess, exists := r.sessions[key] if !exists { - dialed, err := net.Dial("udp", r.dialAddr) + inTsnet := false + if host, _, err := net.SplitHostPort(r.dialAddr); err == nil { + if dialIP, err := netip.ParseAddr(host); err == nil { + tsnetCIDR := netip.MustParsePrefix("100.64.0.0/10") + inTsnet = tsnetCIDR.Contains(dialIP) + } + } + + var dialed net.Conn + if inTsnet { + dialed, err = r.srv.Dial(r.ctx, "udp", r.dialAddr) + } else { + dialed, err = net.Dial("udp", r.dialAddr) + } if err != nil { r.mu.Unlock() r.logger.Error("failed to dial", "error", err) continue } - sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now()} + sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now(), inTailscale: inTsnet} r.sessions[key] = sess r.mu.Unlock() - r.logger.Debug("new udp session", slog.String("remote", key), slog.String("direction", r.direction)) + 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() r.mu.Unlock() } + fromIP, _, _ := net.SplitHostPort(from.String()) + toIP, _, _ := net.SplitHostPort(sess.conn.RemoteAddr().String()) + r.logger.Debug("outbound udp packet", + slog.String("from_ip", fromIP), + slog.String("to_ip", toIP), + 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) @@ -287,6 +397,13 @@ func (r *udpRelay) readSession(key string, sess *udpSession) { 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) @@ -406,6 +523,8 @@ func runUDPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, l dialAddr: rule.DstAddr, logger: logger, direction: "tailscale", + srv: srv, + ctx: ctx, sessions: make(map[string]*udpSession), } relay.run(ctx)