fix(forwarder): potential error on udpSession

This commit is contained in:
nullcat
2026-06-01 15:25:14 +08:00
parent bad69e16cb
commit a9828b4461
+164 -48
View File
@@ -2,6 +2,7 @@ package core
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"io" "io"
"log/slog" "log/slog"
@@ -10,14 +11,15 @@ import (
"sync" "sync"
"time" "time"
"tailscale.com/ipn/ipnstate"
"tailscale.com/tsnet" "tailscale.com/tsnet"
) )
const udpForwardIdleTimeout = 2 * time.Minute
func StartForwarders(ctx context.Context, srv *tsnet.Server, rules map[string][]ForwardRule) { func StartForwarders(ctx context.Context, srv *tsnet.Server, rules map[string][]ForwardRule) {
for tag, rrs := range rules { for tag, rrs := range rules {
for _, rule := range rrs { for _, rule := range rrs {
rule := rule
tag := tag
slog.Info("starting forwarder", slog.Info("starting forwarder",
slog.String("tag", tag), slog.String("tag", tag),
slog.String("protocol", rule.Protocol), slog.String("protocol", rule.Protocol),
@@ -32,8 +34,6 @@ func StartForwarders(ctx context.Context, srv *tsnet.Server, rules map[string][]
func StartConnectors(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule) { func StartConnectors(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule) {
for tag, rrs := range rules { for tag, rrs := range rules {
for _, rule := range rrs { for _, rule := range rrs {
rule := rule
tag := tag
args := []any{ args := []any{
slog.String("tag", tag), slog.String("tag", tag),
slog.String("protocol", rule.Protocol), slog.String("protocol", rule.Protocol),
@@ -67,6 +67,10 @@ func RuleLogger(rule any, tag string) *slog.Logger {
if r.LocalAddr != "" { if r.LocalAddr != "" {
args = append(args, slog.String("local_addr", r.LocalAddr)) args = append(args, slog.String("local_addr", r.LocalAddr))
} }
default:
args = []any{
slog.String("type", fmt.Sprintf("%T", rule)),
}
} }
if tag != "" { if tag != "" {
args = append(args, slog.String("tag", tag)) args = append(args, slog.String("tag", tag))
@@ -132,25 +136,67 @@ func handleTCPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
slog.String("local_addr", rule.LocalAddr), slog.String("local_addr", rule.LocalAddr),
) )
defer conn.Close()
localConn, err := net.Dial("tcp", rule.LocalAddr) localConn, err := net.Dial("tcp", rule.LocalAddr)
if err != nil { if err != nil {
clog.Error("failed to dial local", "error", err) clog.Error("failed to dial local", "error", err)
conn.Close()
return return
} }
defer localConn.Close()
toTs, toLocal := pipeConns(conn, localConn) 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)) clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs))
} }
func getConnType(ctx context.Context, srv *tsnet.Server, remoteAddrStr string) string { 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() lc, err := srv.LocalClient()
if err != nil { if err != nil {
return "unknown" return nil, err
} }
status, err := lc.Status(ctx) 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 { if err != nil {
return "unknown" return "unknown"
} }
@@ -160,9 +206,12 @@ func getConnType(ctx context.Context, srv *tsnet.Server, remoteAddrStr string) s
return "unknown" return "unknown"
} }
for _, peer := range status.Peer { for _, peer := range st.Peer {
for _, addr := range peer.TailscaleIPs { for _, addr := range peer.TailscaleIPs {
if addr.String() == remoteHost { if addr.String() == remoteHost {
if peer.CurAddr != "" {
return "direct"
}
if peer.Relay != "" { if peer.Relay != "" {
return fmt.Sprintf("derp(%s)", peer.Relay) return fmt.Sprintf("derp(%s)", peer.Relay)
} }
@@ -180,7 +229,7 @@ func runUDPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, l
logger.Error("failed to listen", "error", err) logger.Error("failed to listen", "error", err)
return return
} }
logger.Info("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort))) logger.Debug("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort)))
go func() { go func() {
<-ctx.Done() <-ctx.Done()
@@ -218,14 +267,18 @@ func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
slog.String("local_addr", rule.LocalAddr), slog.String("local_addr", rule.LocalAddr),
) )
defer conn.Close()
localConn, err := net.Dial("udp", rule.LocalAddr) localConn, err := net.Dial("udp", rule.LocalAddr)
if err != nil { if err != nil {
clog.Error("failed to dial local", "error", err) clog.Error("failed to dial local", "error", err)
conn.Close()
return return
} }
defer localConn.Close()
stop := context.AfterFunc(ctx, func() {
conn.Close()
localConn.Close()
})
defer stop()
remoteIP, _, _ := net.SplitHostPort(remoteAddrStr) remoteIP, _, _ := net.SplitHostPort(remoteAddrStr)
@@ -237,10 +290,18 @@ func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
defer wg.Done() defer wg.Done()
buf := make([]byte, 65535) buf := make([]byte, 65535)
for { for {
_ = conn.SetReadDeadline(time.Now().Add(udpForwardIdleTimeout))
n, err := conn.Read(buf) n, err := conn.Read(buf)
if err != nil { 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 return
} }
_ = conn.SetReadDeadline(time.Time{})
toLocal += int64(n) toLocal += int64(n)
clog.Debug("inbound udp packet", clog.Debug("inbound udp packet",
slog.String("from_ip", remoteIP), slog.String("from_ip", remoteIP),
@@ -248,6 +309,7 @@ func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
slog.Int("pkg_size", n), slog.Int("pkg_size", n),
) )
if _, err := localConn.Write(buf[:n]); err != nil { if _, err := localConn.Write(buf[:n]); err != nil {
conn.Close()
return return
} }
} }
@@ -257,10 +319,18 @@ func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
defer wg.Done() defer wg.Done()
buf := make([]byte, 65535) buf := make([]byte, 65535)
for { for {
_ = localConn.SetReadDeadline(time.Now().Add(udpForwardIdleTimeout))
n, err := localConn.Read(buf) n, err := localConn.Read(buf)
if err != nil { 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 return
} }
_ = localConn.SetReadDeadline(time.Time{})
toTs += int64(n) toTs += int64(n)
localIP, _, _ := net.SplitHostPort(localConn.RemoteAddr().String()) localIP, _, _ := net.SplitHostPort(localConn.RemoteAddr().String())
clog.Debug("outbound udp packet", clog.Debug("outbound udp packet",
@@ -269,6 +339,7 @@ func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
slog.Int("pkg_size", n), slog.Int("pkg_size", n),
) )
if _, err := conn.Write(buf[:n]); err != nil { if _, err := conn.Write(buf[:n]); err != nil {
localConn.Close()
return return
} }
} }
@@ -278,11 +349,12 @@ func handleUDPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rul
clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs)) clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs))
} }
const udpRelayMaxSessions = 1024
type udpSession struct { type udpSession struct {
conn net.Conn conn net.Conn
remote net.Addr remote net.Addr
inTailscale bool lastUse time.Time
lastUse time.Time
} }
type udpRelay struct { type udpRelay struct {
@@ -291,7 +363,6 @@ type udpRelay struct {
logger *slog.Logger logger *slog.Logger
direction string direction string
srv *tsnet.Server srv *tsnet.Server
ctx context.Context
mu sync.Mutex mu sync.Mutex
sessions map[string]*udpSession sessions map[string]*udpSession
@@ -334,44 +405,62 @@ func (r *udpRelay) run(ctx context.Context) {
} }
key := from.String() key := from.String()
var toIP string
r.mu.Lock() r.mu.Lock()
sess, exists := r.sessions[key] sess, exists := r.sessions[key]
if !exists { if !exists {
inTsnet := false if len(r.sessions) >= udpRelayMaxSessions {
if host, _, err := net.SplitHostPort(r.dialAddr); err == nil { r.mu.Unlock()
if dialIP, err := netip.ParseAddr(host); err == nil { r.logger.Warn("udp relay session limit reached, dropping packet",
tsnetCIDR := netip.MustParsePrefix("100.64.0.0/10") slog.Int("limit", udpRelayMaxSessions),
inTsnet = tsnetCIDR.Contains(dialIP) 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 var dialed net.Conn
if inTsnet { if inTsnet {
dialed, err = r.srv.Dial(r.ctx, "udp", r.dialAddr) dialed, err = r.srv.Dial(ctx, "udp", r.dialAddr)
} else { } else {
dialed, err = net.Dial("udp", r.dialAddr) dialed, err = net.Dial("udp", r.dialAddr)
} }
if err != nil { if err != nil {
r.mu.Unlock()
r.logger.Error("failed to dial", "error", err) r.logger.Error("failed to dial", "error", err)
continue continue
} }
sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now(), inTailscale: inTsnet} sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now()}
r.sessions[key] = sess 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.mu.Unlock()
r.logger.Info("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) go r.readSession(key, sess)
} else { } else {
sess.lastUse = time.Now() sess.lastUse = time.Now()
toIP = sess.conn.RemoteAddr().String()
r.mu.Unlock() r.mu.Unlock()
} }
fromIP, _, _ := net.SplitHostPort(from.String()) fromIP, _, _ := net.SplitHostPort(from.String())
toIP, _, _ := net.SplitHostPort(sess.conn.RemoteAddr().String()) toIPHost, _, _ := net.SplitHostPort(toIP)
r.logger.Debug("outbound udp packet", r.logger.Debug("outbound udp packet",
slog.String("from_ip", fromIP), slog.String("from_ip", fromIP),
slog.String("to_ip", toIP), slog.String("to_ip", toIPHost),
slog.Int("pkg_size", n), slog.Int("pkg_size", n),
) )
if _, err := sess.conn.Write(buf[:n]); err != nil { if _, err := sess.conn.Write(buf[:n]); err != nil {
@@ -401,6 +490,7 @@ func (r *udpRelay) readSession(key string, sess *udpSession) {
r.removeSession(key) r.removeSession(key)
return return
} }
sess.lastUse = time.Now()
} }
} }
@@ -408,9 +498,10 @@ func (r *udpRelay) removeSession(key string) {
r.mu.Lock() r.mu.Lock()
defer r.mu.Unlock() defer r.mu.Unlock()
if s, ok := r.sessions[key]; ok { if s, ok := r.sessions[key]; ok {
remote := s.remote.String()
s.conn.Close() s.conn.Close()
delete(r.sessions, key) delete(r.sessions, key)
r.logger.Debug("udp session closed", slog.String("remote", key)) r.logger.Debug("udp session closed", slog.String("remote", remote))
} }
} }
@@ -420,9 +511,10 @@ func (r *udpRelay) cleanup() {
threshold := time.Now().Add(-5 * time.Minute) threshold := time.Now().Add(-5 * time.Minute)
for key, s := range r.sessions { for key, s := range r.sessions {
if s.lastUse.Before(threshold) { if s.lastUse.Before(threshold) {
remote := s.remote.String()
s.conn.Close() s.conn.Close()
delete(r.sessions, key) delete(r.sessions, key)
r.logger.Debug("udp session cleaned up", slog.String("remote", key)) r.logger.Debug("udp session cleaned up", slog.String("remote", remote))
} }
} }
} }
@@ -477,14 +569,19 @@ func runTCPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, l
func handleTCPConnect(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ConnectRule, logger *slog.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())) clog := logger.With(slog.String("local_client", conn.RemoteAddr().String()))
defer conn.Close()
tsConn, err := srv.Dial(ctx, "tcp", rule.DstAddr) tsConn, err := srv.Dial(ctx, "tcp", rule.DstAddr)
if err != nil { if err != nil {
clog.Error("failed to dial tailscale", "error", err) clog.Error("failed to dial tailscale", "error", err)
conn.Close()
return return
} }
defer tsConn.Close()
stop := context.AfterFunc(ctx, func() {
conn.Close()
tsConn.Close()
})
defer stop()
clog.Info("accepted connection", slog.String("dst_addr", rule.DstAddr)) clog.Info("accepted connection", slog.String("dst_addr", rule.DstAddr))
toConn, toTs := pipeConns(conn, tsConn) toConn, toTs := pipeConns(conn, tsConn)
@@ -516,25 +613,44 @@ func runUDPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, l
logger: logger, logger: logger,
direction: "tailscale", direction: "tailscale",
srv: srv, srv: srv,
ctx: ctx,
sessions: make(map[string]*udpSession), sessions: make(map[string]*udpSession),
} }
relay.run(ctx) relay.run(ctx)
} }
func pipeConns(a, b net.Conn) (toA, toB int64) { func pipeConns(a, b net.Conn) (toA, toB int64) {
var wg sync.WaitGroup done := make(chan struct{}, 2)
wg.Add(2) var aToB, bToA int64
go func() { go func() {
defer wg.Done() defer func() { done <- struct{}{} }()
n, _ := io.Copy(a, b) n, err := io.Copy(a, b)
toA = n 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() { go func() {
defer wg.Done() defer func() { done <- struct{}{} }()
n, _ := io.Copy(b, a) n, err := io.Copy(b, a)
toB = n 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()
}
}() }()
wg.Wait()
return <-done
<-done
return aToB, bToA
} }