diff --git a/client/control.go b/client/control.go index 6e3002bc..5980ed34 100644 --- a/client/control.go +++ b/client/control.go @@ -47,6 +47,8 @@ type SessionContext struct { Connector MessageConnector // Virtual net controller VnetController *vnet.Controller + // UDPPacketCodec is immutable for the lifetime of this negotiated session. + UDPPacketCodec string } type Control struct { @@ -92,7 +94,14 @@ func NewControl(ctx context.Context, sessionCtx *SessionContext) (*Control, erro ctl.registerMsgHandlers() ctl.msgTransporter = transport.NewMessageTransporter(ctl.msgDispatcher) - ctl.pm = proxy.NewManager(ctl.ctx, sessionCtx.Common, sessionCtx.Auth.EncryptionKey(), ctl.msgTransporter, sessionCtx.VnetController) + ctl.pm = proxy.NewManager( + ctl.ctx, + sessionCtx.Common, + sessionCtx.Auth.EncryptionKey(), + ctl.msgTransporter, + sessionCtx.VnetController, + sessionCtx.UDPPacketCodec, + ) ctl.vm = visitor.NewManager(ctl.ctx, sessionCtx.RunID, sessionCtx.Common, ctl.connectServer, ctl.msgTransporter, sessionCtx.VnetController) return ctl, nil diff --git a/client/control_session.go b/client/control_session.go index d533ba2d..e3ed27fc 100644 --- a/client/control_session.go +++ b/client/control_session.go @@ -99,6 +99,7 @@ func (d *controlSessionDialer) Dial(previousRunID string) (*SessionContext, erro Auth: d.auth, Connector: newMessageConnector(connector, d.common.Transport.WireProtocol), VnetController: d.vnetController, + UDPPacketCodec: loginResult.udpPacketCodec, }, nil } @@ -127,8 +128,9 @@ func (d *controlSessionDialer) buildLoginMsg(previousRunID string) (*msg.Login, } type loginExchangeResult struct { - resp *msg.LoginResp - crypto *wire.CryptoContext + resp *msg.LoginResp + crypto *wire.CryptoContext + udpPacketCodec string } func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) (*loginExchangeResult, error) { @@ -172,6 +174,7 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) }() var cryptoContext *wire.CryptoContext + var udpPacketCodec string if wireConn != nil { serverHelloFrame, err := wireConn.ReadFrame() if err != nil { @@ -191,6 +194,7 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) if err != nil { return nil, err } + udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec } var loginRespMsg msg.LoginResp @@ -198,8 +202,9 @@ func (d *controlSessionDialer) exchangeLogin(conn net.Conn, loginMsg *msg.Login) return nil, err } return &loginExchangeResult{ - resp: &loginRespMsg, - crypto: cryptoContext, + resp: &loginRespMsg, + crypto: cryptoContext, + udpPacketCodec: udpPacketCodec, }, nil } diff --git a/client/control_session_test.go b/client/control_session_test.go index a0778fba..4bbcd453 100644 --- a/client/control_session_test.go +++ b/client/control_session_test.go @@ -117,6 +117,7 @@ func TestControlSessionDialerDialV1(t *testing.T) { defer sessionCtx.Connector.Close() require.Equal(t, "run-v1", sessionCtx.RunID) + require.Empty(t, sessionCtx.UDPPacketCodec) require.NotNil(t, sessionCtx.Conn) require.NotNil(t, sessionCtx.Connector) require.False(t, connector.closed.Load()) @@ -225,6 +226,7 @@ func TestControlSessionDialerDialV2(t *testing.T) { defer sessionCtx.Connector.Close() require.Equal(t, "run-v2", sessionCtx.RunID) + require.Equal(t, wire.UDPPacketCodecBinary, sessionCtx.UDPPacketCodec) require.NotNil(t, sessionCtx.Conn) require.NotNil(t, sessionCtx.Connector) require.False(t, connector.closed.Load()) diff --git a/client/control_udp_test.go b/client/control_udp_test.go new file mode 100644 index 00000000..f3c91232 --- /dev/null +++ b/client/control_udp_test.go @@ -0,0 +1,125 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//go:build !frps + +package client + +import ( + "context" + "encoding/binary" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" + + clientproxy "github.com/fatedier/frp/client/proxy" + "github.com/fatedier/frp/pkg/auth" + v1 "github.com/fatedier/frp/pkg/config/v1" + "github.com/fatedier/frp/pkg/msg" + "github.com/fatedier/frp/pkg/proto/wire" +) + +func TestControlPropagatesBinaryUDPPacketCodecToWorkConn(t *testing.T) { + echoConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + require.NoError(t, err) + t.Cleanup(func() { _ = echoConn.Close() }) + + echoDone := make(chan error, 1) + go func() { + buf := make([]byte, 64) + n, addr, err := echoConn.ReadFromUDP(buf) + if err == nil { + _, err = echoConn.WriteToUDP(buf[:n], addr) + } + echoDone <- err + }() + + authRuntime, err := auth.BuildClientAuth(&v1.AuthClientConfig{ + Method: v1.AuthMethodToken, + Token: "token", + }) + require.NoError(t, err) + + controlConn, controlPeer := net.Pipe() + t.Cleanup(func() { + _ = controlConn.Close() + _ = controlPeer.Close() + }) + common := &v1.ClientCommonConfig{ + Transport: v1.ClientTransportConfig{WireProtocol: wire.ProtocolV2}, + UDPPacketSize: 1500, + } + ctl, err := NewControl(context.Background(), &SessionContext{ + Common: common, + RunID: "binary-udp-test", + Conn: msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn)), + Auth: authRuntime, + UDPPacketCodec: wire.UDPPacketCodecBinary, + }) + require.NoError(t, err) + t.Cleanup(ctl.pm.Close) + + echoAddr := echoConn.LocalAddr().(*net.UDPAddr) + proxyCfg := &v1.UDPProxyConfig{ + ProxyBaseConfig: v1.ProxyBaseConfig{ + Name: "udp", + Type: string(v1.ProxyTypeUDP), + ProxyBackend: v1.ProxyBackend{ + LocalIP: "127.0.0.1", + LocalPort: echoAddr.Port, + }, + }, + } + ctl.pm.UpdateAll([]v1.ProxyConfigurer{proxyCfg}) + require.Eventually(t, func() bool { + status, ok := ctl.pm.GetProxyStatus("udp") + return ok && status.Phase == clientproxy.ProxyPhaseWaitStart + }, time.Second, 10*time.Millisecond) + require.NoError(t, ctl.pm.StartProxy("udp", "", "")) + + workClient, workServer := net.Pipe() + t.Cleanup(func() { + _ = workClient.Close() + _ = workServer.Close() + }) + deadline := time.Now().Add(3 * time.Second) + require.NoError(t, workClient.SetDeadline(deadline)) + require.NoError(t, workServer.SetDeadline(deadline)) + ctl.pm.HandleWorkConn("udp", workClient, &msg.StartWorkConn{ProxyName: "udp"}) + + serverRW, err := msg.NewUDPPacketReadWriter(workServer, wire.ProtocolV2, wire.UDPPacketCodecBinary) + require.NoError(t, err) + writeDone := make(chan error, 1) + in := &msg.UDPPacket{ + Content: []byte("binary udp"), + RemoteAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 12345}, + } + go func() { + writeDone <- serverRW.WriteMsg(in) + }() + + frame, err := wire.NewConn(workServer).ReadFrame() + require.NoError(t, err) + require.Equal(t, wire.FrameTypeMessage, frame.Type) + require.GreaterOrEqual(t, len(frame.Payload), 2) + require.Equal(t, msg.V2TypeUDPPacketBinary, binary.BigEndian.Uint16(frame.Payload[:2])) + out, err := msg.DecodeUDPPacketBinary(frame.Payload[2:]) + require.NoError(t, err) + require.Equal(t, in.Content, out.Content) + require.Equal(t, in.RemoteAddr.String(), out.RemoteAddr.String()) + require.NoError(t, <-writeDone) + require.NoError(t, <-echoDone) +} diff --git a/client/proxy/proxy.go b/client/proxy/proxy.go index 84ff49a1..5af56599 100644 --- a/client/proxy/proxy.go +++ b/client/proxy/proxy.go @@ -61,6 +61,7 @@ func NewProxy( encryptionKey []byte, msgTransporter transport.MessageTransporter, vnetController *vnet.Controller, + udpPacketCodec string, ) (pxy Proxy) { var limiter *rate.Limiter limitBytes := pxyConf.GetBaseConfig().Transport.BandwidthLimit.Bytes() @@ -77,6 +78,7 @@ func NewProxy( vnetController: vnetController, xl: xlog.FromContextSafe(ctx), ctx: ctx, + udpPacketCodec: udpPacketCodec, } factory := proxyFactoryRegistry[reflect.TypeOf(pxyConf)] @@ -98,9 +100,10 @@ type BaseProxy struct { proxyPlugin plugin.Plugin inWorkConnCallback func(*v1.ProxyBaseConfig, net.Conn, *msg.StartWorkConn) /* continue */ bool - mu sync.RWMutex - xl *xlog.Logger - ctx context.Context + mu sync.RWMutex + xl *xlog.Logger + ctx context.Context + udpPacketCodec string } func (pxy *BaseProxy) Run() error { diff --git a/client/proxy/proxy_manager.go b/client/proxy/proxy_manager.go index 42f9f589..6bf8d4a8 100644 --- a/client/proxy/proxy_manager.go +++ b/client/proxy/proxy_manager.go @@ -43,7 +43,8 @@ type Manager struct { encryptionKey []byte clientCfg *v1.ClientCommonConfig - ctx context.Context + ctx context.Context + udpPacketCodec string } func NewManager( @@ -52,6 +53,7 @@ func NewManager( encryptionKey []byte, msgTransporter transport.MessageTransporter, vnetController *vnet.Controller, + udpPacketCodec string, ) *Manager { return &Manager{ proxies: make(map[string]*Wrapper), @@ -61,6 +63,7 @@ func NewManager( encryptionKey: encryptionKey, clientCfg: clientCfg, ctx: ctx, + udpPacketCodec: udpPacketCodec, } } @@ -166,7 +169,7 @@ func (pm *Manager) UpdateAll(proxyCfgs []v1.ProxyConfigurer) { for _, cfg := range proxyCfgs { name := cfg.GetBaseConfig().Name if _, ok := pm.proxies[name]; !ok { - pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController) + pxy := NewWrapper(pm.ctx, cfg, pm.clientCfg, pm.encryptionKey, pm.HandleEvent, pm.msgTransporter, pm.vnetController, pm.udpPacketCodec) if pm.inWorkConnCallback != nil { pxy.SetInWorkConnCallback(pm.inWorkConnCallback) } diff --git a/client/proxy/proxy_wrapper.go b/client/proxy/proxy_wrapper.go index 718c02e6..37a71277 100644 --- a/client/proxy/proxy_wrapper.go +++ b/client/proxy/proxy_wrapper.go @@ -99,6 +99,7 @@ func NewWrapper( eventHandler event.Handler, msgTransporter transport.MessageTransporter, vnetController *vnet.Controller, + udpPacketCodec string, ) *Wrapper { baseInfo := cfg.GetBaseConfig() xl := xlog.FromContextSafe(ctx).Spawn().AppendPrefix(baseInfo.Name) @@ -127,7 +128,7 @@ func NewWrapper( xl.Tracef("enable health check monitor") } - pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController) + pw.pxy = NewProxy(pw.ctx, pw.Cfg, clientCfg, encryptionKey, pw.msgTransporter, pw.vnetController, udpPacketCodec) return pw } diff --git a/client/proxy/udp.go b/client/proxy/udp.go index 110c01ac..d01d2c8e 100644 --- a/client/proxy/udp.go +++ b/client/proxy/udp.go @@ -97,10 +97,17 @@ func (pxy *UDPProxy) InWorkConn(conn net.Conn, _ *msg.StartWorkConn) { return } - pxy.mu.Lock() - pxy.workConn = netpkg.WrapReadWriteCloserToConn(remote, conn) + workConn := netpkg.WrapReadWriteCloserToConn(remote, conn) // Plain UDP payload follows the configured wire protocol for message framing. - payloadRW := msg.NewReadWriter(pxy.workConn, pxy.clientCfg.Transport.WireProtocol) + payloadRW, err := msg.NewUDPPacketReadWriter(workConn, pxy.clientCfg.Transport.WireProtocol, pxy.udpPacketCodec) + if err != nil { + xl.Errorf("create UDP packet read writer: %v", err) + workConn.Close() + return + } + + pxy.mu.Lock() + pxy.workConn = workConn pxy.readCh = make(chan *msg.UDPPacket, 1024) pxy.sendCh = make(chan msg.Message, 1024) pxy.closed = false diff --git a/client/service_shutdown_test.go b/client/service_shutdown_test.go index 5c710bec..3b956726 100644 --- a/client/service_shutdown_test.go +++ b/client/service_shutdown_test.go @@ -34,7 +34,7 @@ func newGracefulCloseTestService() *Service { }, doneCh: make(chan struct{}), } - ctl.pm = proxy.NewManager(ctx, common, nil, nil, nil) + ctl.pm = proxy.NewManager(ctx, common, nil, nil, nil, "") ctl.vm = visitor.NewManager(ctx, "graceful-close-race", common, nil, nil, nil) return &Service{ctl: ctl, cancel: context.CancelCauseFunc(func(error) {})} } diff --git a/pkg/msg/udp_binary.go b/pkg/msg/udp_binary.go new file mode 100644 index 00000000..7493c70f --- /dev/null +++ b/pkg/msg/udp_binary.go @@ -0,0 +1,338 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package msg + +import ( + "encoding/binary" + "fmt" + "io" + "net" + "unicode/utf8" + + "github.com/fatedier/frp/pkg/proto/wire" +) + +const MaxUDPPayloadSize = 65507 + +const ( + udpPacketFlagLocalAddr byte = 1 << 0 + udpPacketFlagRemoteAddr byte = 1 << 1 + udpPacketValidFlags = udpPacketFlagLocalAddr | udpPacketFlagRemoteAddr +) + +type binaryUDPAddr struct { + family byte + ip []byte + port uint16 + zone string +} + +// EncodeUDPPacketBinary encodes the body of a V2 binary UDP packet message. +// RemoteAddr is required by the UDP forwarding path. +func EncodeUDPPacketBinary(packet *UDPPacket) ([]byte, error) { + if packet == nil { + return nil, fmt.Errorf("nil UDP packet") + } + if packet.RemoteAddr == nil { + return nil, fmt.Errorf("UDP packet missing remote address") + } + if len(packet.Content) > MaxUDPPayloadSize { + return nil, fmt.Errorf("UDP payload length %d exceeds limit %d", len(packet.Content), MaxUDPPayloadSize) + } + + var flags byte + var localAddr, remoteAddr binaryUDPAddr + bodyLen := 1 + 2 + len(packet.Content) + if packet.LocalAddr != nil { + flags |= udpPacketFlagLocalAddr + var err error + localAddr, err = validateBinaryUDPAddr(packet.LocalAddr) + if err != nil { + return nil, fmt.Errorf("local address: %w", err) + } + bodyLen += binaryUDPAddrLen(localAddr) + } + flags |= udpPacketFlagRemoteAddr + var err error + remoteAddr, err = validateBinaryUDPAddr(packet.RemoteAddr) + if err != nil { + return nil, fmt.Errorf("remote address: %w", err) + } + bodyLen += binaryUDPAddrLen(remoteAddr) + if 2+bodyLen > wire.DefaultMaxFramePayloadSize { + return nil, fmt.Errorf("v2 frame payload length %d exceeds limit %d", 2+bodyLen, wire.DefaultMaxFramePayloadSize) + } + + body := make([]byte, bodyLen) + body[0] = flags + offset := 1 + if flags&udpPacketFlagLocalAddr != 0 { + offset = putBinaryUDPAddr(body, offset, localAddr) + } + offset = putBinaryUDPAddr(body, offset, remoteAddr) + binary.BigEndian.PutUint16(body[offset:offset+2], uint16(len(packet.Content))) + offset += 2 + copy(body[offset:], packet.Content) + return body, nil +} + +// DecodeUDPPacketBinary decodes a V2 binary UDP packet body and returns data +// that does not alias the input frame buffer. +func DecodeUDPPacketBinary(body []byte) (*UDPPacket, error) { + if len(body) < 3 { + return nil, fmt.Errorf("UDP packet body too short: %d", len(body)) + } + if 2+len(body) > wire.DefaultMaxFramePayloadSize { + return nil, fmt.Errorf("v2 frame payload length %d exceeds limit %d", 2+len(body), wire.DefaultMaxFramePayloadSize) + } + + flags := body[0] + if flags&^udpPacketValidFlags != 0 { + return nil, fmt.Errorf("reserved UDP packet flags set: 0x%02x", flags) + } + if flags&udpPacketFlagRemoteAddr == 0 { + return nil, fmt.Errorf("UDP packet missing remote address") + } + + packet := &UDPPacket{} + offset := 1 + var err error + if flags&udpPacketFlagLocalAddr != 0 { + packet.LocalAddr, offset, err = readBinaryUDPAddr(body, offset) + if err != nil { + return nil, fmt.Errorf("local address: %w", err) + } + } + if flags&udpPacketFlagRemoteAddr != 0 { + packet.RemoteAddr, offset, err = readBinaryUDPAddr(body, offset) + if err != nil { + return nil, fmt.Errorf("remote address: %w", err) + } + } + if len(body)-offset < 2 { + return nil, fmt.Errorf("truncated UDP payload length") + } + payloadLen := int(binary.BigEndian.Uint16(body[offset : offset+2])) + offset += 2 + if payloadLen > MaxUDPPayloadSize { + return nil, fmt.Errorf("UDP payload length %d exceeds limit %d", payloadLen, MaxUDPPayloadSize) + } + remaining := len(body) - offset + if remaining < payloadLen { + return nil, fmt.Errorf("truncated UDP payload: have %d want %d", remaining, payloadLen) + } + if remaining > payloadLen { + return nil, fmt.Errorf("trailing UDP packet bytes: %d", remaining-payloadLen) + } + packet.Content = append([]byte(nil), body[offset:offset+payloadLen]...) + return packet, nil +} + +func validateBinaryUDPAddr(addr *net.UDPAddr) (binaryUDPAddr, error) { + if addr.Port < 0 || addr.Port > 65535 { + return binaryUDPAddr{}, fmt.Errorf("port out of range: %d", addr.Port) + } + if ip := addr.IP.To4(); ip != nil { + if addr.Zone != "" { + return binaryUDPAddr{}, fmt.Errorf("IPv4 zone is forbidden") + } + return binaryUDPAddr{family: 4, ip: ip, port: uint16(addr.Port)}, nil + } + ip := addr.IP.To16() + if ip == nil { + return binaryUDPAddr{}, fmt.Errorf("invalid IP") + } + if len(addr.Zone) > 255 { + return binaryUDPAddr{}, fmt.Errorf("zone exceeds 255 bytes") + } + if !utf8.ValidString(addr.Zone) { + return binaryUDPAddr{}, fmt.Errorf("zone is not valid UTF-8") + } + return binaryUDPAddr{family: 6, ip: ip, port: uint16(addr.Port), zone: addr.Zone}, nil +} + +func binaryUDPAddrLen(addr binaryUDPAddr) int { + return 1 + len(addr.ip) + 2 + 1 + len(addr.zone) +} + +func putBinaryUDPAddr(body []byte, offset int, addr binaryUDPAddr) int { + body[offset] = addr.family + offset++ + copy(body[offset:], addr.ip) + offset += len(addr.ip) + binary.BigEndian.PutUint16(body[offset:offset+2], addr.port) + offset += 2 + body[offset] = byte(len(addr.zone)) + offset++ + copy(body[offset:], addr.zone) + return offset + len(addr.zone) +} + +func readBinaryUDPAddr(body []byte, offset int) (*net.UDPAddr, int, error) { + if offset >= len(body) { + return nil, offset, fmt.Errorf("truncated address family") + } + family := body[offset] + offset++ + var ipLen int + switch family { + case 4: + ipLen = net.IPv4len + case 6: + ipLen = net.IPv6len + default: + return nil, offset, fmt.Errorf("unknown address family %d", family) + } + if len(body)-offset < ipLen+3 { + return nil, offset, fmt.Errorf("truncated address") + } + ip := append(net.IP(nil), body[offset:offset+ipLen]...) + offset += ipLen + port := binary.BigEndian.Uint16(body[offset : offset+2]) + offset += 2 + zoneLen := int(body[offset]) + offset++ + if len(body)-offset < zoneLen { + return nil, offset, fmt.Errorf("truncated zone") + } + zoneBytes := body[offset : offset+zoneLen] + if family == 4 && zoneLen != 0 { + return nil, offset, fmt.Errorf("IPv4 zone is forbidden") + } + if !utf8.Valid(zoneBytes) { + return nil, offset, fmt.Errorf("zone is not valid UTF-8") + } + offset += zoneLen + return &net.UDPAddr{IP: ip, Port: int(port), Zone: string(zoneBytes)}, offset, nil +} + +type V2BinaryUDPPacketReadWriter struct { + conn *wire.Conn +} + +func NewV2BinaryUDPPacketReadWriter(rw io.ReadWriter) *V2BinaryUDPPacketReadWriter { + return &V2BinaryUDPPacketReadWriter{conn: wire.NewConn(rw)} +} + +func (rw *V2BinaryUDPPacketReadWriter) ReadMsg() (Message, error) { + frame, err := rw.conn.ReadFrame() + if err != nil { + return nil, err + } + if isV2MessageType(frame, V2TypeUDPPacketBinary) { + return decodeV2BinaryUDPPacketFrame(frame) + } + if isV2MessageType(frame, V2TypeUDPPacket) { + return nil, fmt.Errorf("received JSON UDP packet after binary codec negotiation") + } + return DecodeV2MessageFrame(frame) +} + +func (rw *V2BinaryUDPPacketReadWriter) ReadMsgInto(out Message) error { + frame, err := rw.conn.ReadFrame() + if err != nil { + return err + } + if packetOut, ok := out.(*UDPPacket); ok { + if !isV2MessageType(frame, V2TypeUDPPacketBinary) { + return unexpectedV2UDPPacketType(frame) + } + packet, err := decodeV2BinaryUDPPacketFrame(frame) + if err != nil { + return err + } + *packetOut = *packet + return nil + } + return DecodeV2MessageFrameInto(frame, out) +} + +func (rw *V2BinaryUDPPacketReadWriter) WriteMsg(message Message) error { + var packet *UDPPacket + switch typed := message.(type) { + case *UDPPacket: + packet = typed + case UDPPacket: + packet = &typed + default: + frame, err := EncodeV2MessageFrame(message) + if err != nil { + return err + } + return rw.conn.WriteFrame(frame) + } + body, err := EncodeUDPPacketBinary(packet) + if err != nil { + return err + } + payload := make([]byte, 2+len(body)) + binary.BigEndian.PutUint16(payload[:2], V2TypeUDPPacketBinary) + copy(payload[2:], body) + return rw.conn.WriteFrame(&wire.Frame{Type: wire.FrameTypeMessage, Payload: payload}) +} + +func decodeV2BinaryUDPPacketFrame(frame *wire.Frame) (*UDPPacket, error) { + if frame.Type != wire.FrameTypeMessage { + return nil, fmt.Errorf("unexpected frame type %d, want %d", frame.Type, wire.FrameTypeMessage) + } + if len(frame.Payload) < 2 { + return nil, fmt.Errorf("message frame payload too short") + } + if binary.BigEndian.Uint16(frame.Payload[:2]) != V2TypeUDPPacketBinary { + return nil, unexpectedV2UDPPacketType(frame) + } + return DecodeUDPPacketBinary(frame.Payload[2:]) +} + +func isV2MessageType(frame *wire.Frame, typeID uint16) bool { + return frame.Type == wire.FrameTypeMessage && len(frame.Payload) >= 2 && binary.BigEndian.Uint16(frame.Payload[:2]) == typeID +} + +func unexpectedV2UDPPacketType(frame *wire.Frame) error { + if frame.Type != wire.FrameTypeMessage { + return fmt.Errorf("unexpected frame type %d, want %d", frame.Type, wire.FrameTypeMessage) + } + if len(frame.Payload) < 2 { + return fmt.Errorf("message frame payload too short") + } + typeID := binary.BigEndian.Uint16(frame.Payload[:2]) + if typeID == V2TypeUDPPacket { + return fmt.Errorf("received JSON UDP packet after binary codec negotiation") + } + return fmt.Errorf("unexpected message type %d, want %d", typeID, V2TypeUDPPacketBinary) +} + +// NewUDPPacketReadWriter selects the negotiated packet codec without changing +// the framing or codecs used by non-UDP messages on the work connection. +func NewUDPPacketReadWriter(rw io.ReadWriter, wireProtocol, udpPacketCodec string) (ReadWriter, error) { + switch wireProtocol { + case "", wire.ProtocolV1: + if udpPacketCodec != "" { + return nil, fmt.Errorf("UDP packet codec %q requires wire protocol v2", udpPacketCodec) + } + return NewV1ReadWriter(rw), nil + case wire.ProtocolV2: + switch udpPacketCodec { + case "": + return NewV2ReadWriter(rw), nil + case wire.UDPPacketCodecBinary: + return NewV2BinaryUDPPacketReadWriter(rw), nil + default: + return nil, fmt.Errorf("unsupported UDP packet codec %q", udpPacketCodec) + } + default: + return nil, fmt.Errorf("unsupported wire protocol %q", wireProtocol) + } +} diff --git a/pkg/msg/udp_binary_test.go b/pkg/msg/udp_binary_test.go new file mode 100644 index 00000000..7dac31fb --- /dev/null +++ b/pkg/msg/udp_binary_test.go @@ -0,0 +1,248 @@ +// Copyright 2026 The frp Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 + +package msg + +import ( + "bytes" + "encoding/binary" + "net" + "strconv" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/fatedier/frp/pkg/proto/wire" +) + +func TestUDPPacketBinaryRoundTrip(t *testing.T) { + payload := bytes.Repeat([]byte{0xa5}, 1472) + in := &UDPPacket{ + Content: payload, + LocalAddr: &net.UDPAddr{ + IP: net.ParseIP("2001:db8::1"), + Port: 1234, + Zone: "en0", + }, + RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}, + } + body, err := EncodeUDPPacketBinary(in) + require.NoError(t, err) + out, err := DecodeUDPPacketBinary(body) + require.NoError(t, err) + require.Equal(t, in.Content, out.Content) + require.Equal(t, in.LocalAddr.String(), out.LocalAddr.String()) + require.Equal(t, in.RemoteAddr.String(), out.RemoteAddr.String()) + body[len(body)-1] ^= 0xff + body[25] ^= 0xff + require.Equal(t, byte(0xa5), out.Content[len(out.Content)-1], "decoded payload must own frame bytes") + require.Equal(t, byte(203), out.RemoteAddr.IP.To4()[0], "decoded address must own frame bytes") +} + +func TestUDPPacketBinarySizesAndOptionalLocalAddress(t *testing.T) { + for _, size := range []int{0, 32, 128, 512, 1200, 1472, 4096, 49107, 65507} { + t.Run(strconv.Itoa(size), func(t *testing.T) { + in := &UDPPacket{ + Content: bytes.Repeat([]byte{byte(size)}, size), + RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}, + } + body, err := EncodeUDPPacketBinary(in) + require.NoError(t, err) + out, err := DecodeUDPPacketBinary(body) + require.NoError(t, err) + require.Equal(t, len(in.Content), len(out.Content)) + if size == 0 { + require.Empty(t, out.Content) + } else { + require.Equal(t, in.Content, out.Content) + } + }) + } +} + +func TestUDPPacketBinaryMalformed(t *testing.T) { + valid, err := EncodeUDPPacketBinary(&UDPPacket{ + Content: []byte("payload"), + RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}, + }) + require.NoError(t, err) + tests := [][]byte{ + {0x80, 0, 0}, + {0x02, 4, 1, 2}, + {0x02, 4, 1, 2, 3, 4, 0xd4}, + append(append([]byte(nil), valid...), 0), + } + for _, malformed := range tests { + _, err := DecodeUDPPacketBinary(malformed) + require.Error(t, err) + } + _, err = DecodeUDPPacketBinary([]byte{0, 0, 0}) + require.ErrorContains(t, err, "missing remote address") + payloadLengthOffset := len(valid) - len("payload") - 2 + invalidPayloadLength := append([]byte(nil), valid...) + binary.BigEndian.PutUint16(invalidPayloadLength[payloadLengthOffset:payloadLengthOffset+2], 0xffff) + _, err = DecodeUDPPacketBinary(invalidPayloadLength) + require.ErrorContains(t, err, "payload length") + truncatedPayload := append([]byte(nil), valid[:payloadLengthOffset+2]...) + binary.BigEndian.PutUint16(truncatedPayload[payloadLengthOffset:payloadLengthOffset+2], 1) + _, err = DecodeUDPPacketBinary(truncatedPayload) + require.ErrorContains(t, err, "truncated UDP payload") + _, err = DecodeUDPPacketBinary(make([]byte, wire.DefaultMaxFramePayloadSize)) + require.ErrorContains(t, err, "frame payload length") + + badIPv4Zone := []byte{2, 4, 203, 0, 113, 9, 0xd4, 0x31, 1, 'z', 0, 0} + _, err = DecodeUDPPacketBinary(badIPv4Zone) + require.ErrorContains(t, err, "IPv4 zone") + badFamily := []byte{2, 9, 0, 0} + _, err = DecodeUDPPacketBinary(badFamily) + require.ErrorContains(t, err, "unknown address family") + badUTF8 := make([]byte, 0, 24) + badUTF8 = append(badUTF8, 2, 6) + badUTF8 = append(badUTF8, make([]byte, 16)...) + badUTF8 = append(badUTF8, 0, 1, 1, 0xff, 0, 0) + _, err = DecodeUDPPacketBinary(badUTF8) + require.ErrorContains(t, err, "UTF-8") +} + +func TestUDPPacketBinaryEncodeRejectsInvalidPackets(t *testing.T) { + _, err := EncodeUDPPacketBinary(&UDPPacket{}) + require.ErrorContains(t, err, "missing remote address") + _, err = EncodeUDPPacketBinary(&UDPPacket{ + LocalAddr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 1234}, + }) + require.ErrorContains(t, err, "missing remote address") + _, err = EncodeUDPPacketBinary(&UDPPacket{ + Content: make([]byte, MaxUDPPayloadSize+1), + RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}, + }) + require.ErrorContains(t, err, "exceeds limit") + _, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 1, Zone: "bad"}}) + require.ErrorContains(t, err, "IPv4 zone") + _, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1, Zone: string(bytes.Repeat([]byte{'z'}, 256))}}) + require.ErrorContains(t, err, "zone exceeds") + _, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Port: 1, Zone: string([]byte{0xff})}}) + require.ErrorContains(t, err, "UTF-8") + _, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{Port: -1}}) + require.ErrorContains(t, err, "port out of range") + _, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{Port: 65536}}) + require.ErrorContains(t, err, "port out of range") + _, err = EncodeUDPPacketBinary(&UDPPacket{RemoteAddr: &net.UDPAddr{IP: net.IP{1, 2, 3}}}) + require.ErrorContains(t, err, "invalid IP") + _, err = EncodeUDPPacketBinary(&UDPPacket{ + Content: make([]byte, MaxUDPPayloadSize), + RemoteAddr: &net.UDPAddr{IP: net.ParseIP("2001:db8::1"), Zone: string(bytes.Repeat([]byte{'z'}, 255))}, + }) + require.ErrorContains(t, err, "frame payload length") +} + +func TestV2BinaryUDPPacketReadWriterPreservesOtherMessages(t *testing.T) { + var buf bytes.Buffer + rw := NewV2BinaryUDPPacketReadWriter(&buf) + in := &UDPPacket{Content: []byte("udp"), RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}} + require.NoError(t, rw.WriteMsg(in)) + require.NoError(t, rw.WriteMsg(&Ping{Timestamp: 7})) + frameConn := wire.NewConn(&buf) + frame, err := frameConn.ReadFrame() + require.NoError(t, err) + require.Equal(t, V2TypeUDPPacketBinary, binary.BigEndian.Uint16(frame.Payload[:2])) + frame, err = frameConn.ReadFrame() + require.NoError(t, err) + require.Equal(t, V2TypePing, binary.BigEndian.Uint16(frame.Payload[:2])) +} + +func TestV2BinaryUDPPacketReadWriterRoundTripAndCodecInvariant(t *testing.T) { + in := &UDPPacket{Content: []byte("udp"), RemoteAddr: &net.UDPAddr{IP: net.ParseIP("203.0.113.9"), Port: 54321}} + var binaryStream bytes.Buffer + binaryWriter, err := NewUDPPacketReadWriter(&binaryStream, wire.ProtocolV2, wire.UDPPacketCodecBinary) + require.NoError(t, err) + require.NoError(t, binaryWriter.WriteMsg(in)) + binaryReader, err := NewUDPPacketReadWriter(&binaryStream, wire.ProtocolV2, wire.UDPPacketCodecBinary) + require.NoError(t, err) + out, err := binaryReader.ReadMsg() + require.NoError(t, err) + require.Equal(t, in.Content, out.(*UDPPacket).Content) + + for _, read := range []func(ReadWriter) error{ + func(rw ReadWriter) error { + _, err := rw.ReadMsg() + return err + }, + func(rw ReadWriter) error { + return rw.ReadMsgInto(&UDPPacket{}) + }, + } { + var jsonStream bytes.Buffer + require.NoError(t, NewReadWriter(&jsonStream, wire.ProtocolV2).WriteMsg(in)) + negotiatedReader, err := NewUDPPacketReadWriter(&jsonStream, wire.ProtocolV2, wire.UDPPacketCodecBinary) + require.NoError(t, err) + require.ErrorContains(t, read(negotiatedReader), "JSON UDP packet after binary codec negotiation") + } + + var fallbackStream bytes.Buffer + fallbackWriter, err := NewUDPPacketReadWriter(&fallbackStream, wire.ProtocolV2, "") + require.NoError(t, err) + require.NoError(t, fallbackWriter.WriteMsg(in)) + frame, err := wire.NewConn(&fallbackStream).ReadFrame() + require.NoError(t, err) + require.Equal(t, V2TypeUDPPacket, binary.BigEndian.Uint16(frame.Payload[:2])) +} + +func TestNewUDPPacketReadWriterDefaultProtocolUsesV1(t *testing.T) { + var stream bytes.Buffer + rw, err := NewUDPPacketReadWriter(&stream, "", "") + require.NoError(t, err) + require.IsType(t, &V1ReadWriter{}, rw) + require.NoError(t, rw.WriteMsg(&UDPPacket{Content: []byte("legacy")})) + require.Equal(t, TypeUDPPacket, stream.Bytes()[0]) +} + +func TestNewUDPPacketReadWriterRejectsInvalidSelection(t *testing.T) { + for _, tc := range []struct { + name string + wireProtocol string + udpPacketCodec string + errorSubstring string + }{ + { + name: "binary codec over v1", + wireProtocol: wire.ProtocolV1, + udpPacketCodec: wire.UDPPacketCodecBinary, + errorSubstring: "requires wire protocol v2", + }, + { + name: "binary codec over default protocol", + udpPacketCodec: wire.UDPPacketCodecBinary, + errorSubstring: "requires wire protocol v2", + }, + { + name: "unknown v2 codec", + wireProtocol: wire.ProtocolV2, + udpPacketCodec: "unknown", + errorSubstring: "unsupported UDP packet codec", + }, + { + name: "unknown wire protocol", + wireProtocol: "unknown", + errorSubstring: "unsupported wire protocol", + }, + } { + t.Run(tc.name, func(t *testing.T) { + rw, err := NewUDPPacketReadWriter(&bytes.Buffer{}, tc.wireProtocol, tc.udpPacketCodec) + require.Nil(t, rw) + require.ErrorContains(t, err, tc.errorSubstring) + }) + } +} + +func FuzzDecodeUDPPacketBinary(f *testing.F) { + f.Add([]byte{0, 0, 0}) + f.Add([]byte{2, 4, 203, 0, 113, 9, 0xd4, 0x31, 0, 0, 1}) + f.Fuzz(func(t *testing.T, body []byte) { + _, _ = DecodeUDPPacketBinary(body) + }) +} diff --git a/pkg/msg/wire_v2.go b/pkg/msg/wire_v2.go index 8d2cd88d..f3da2330 100644 --- a/pkg/msg/wire_v2.go +++ b/pkg/msg/wire_v2.go @@ -43,6 +43,7 @@ const ( V2TypeNatHoleResp uint16 = 16 V2TypeNatHoleSid uint16 = 17 V2TypeNatHoleReport uint16 = 18 + V2TypeUDPPacketBinary uint16 = 19 ) var v2MsgTypeMap = map[uint16]any{ diff --git a/pkg/msg/wire_v2_test.go b/pkg/msg/wire_v2_test.go index f6f25e55..ea6e8d5d 100644 --- a/pkg/msg/wire_v2_test.go +++ b/pkg/msg/wire_v2_test.go @@ -84,6 +84,9 @@ func TestV2MessageTypeIDsAreStable(t *testing.T) { require.Equal(t, uint16(16), V2TypeNatHoleResp) require.Equal(t, uint16(17), V2TypeNatHoleSid) require.Equal(t, uint16(18), V2TypeNatHoleReport) + require.Equal(t, uint16(19), V2TypeUDPPacketBinary) + _, registered := v2MsgTypeMap[V2TypeUDPPacketBinary] + require.False(t, registered, "binary UDP has a dedicated codec and must not alter generic type registry") } func TestV2MessageFrameEncoding(t *testing.T) { diff --git a/pkg/proto/wire/crypto.go b/pkg/proto/wire/crypto.go index 69aedc69..38916a28 100644 --- a/pkg/proto/wire/crypto.go +++ b/pkg/proto/wire/crypto.go @@ -68,7 +68,8 @@ func NewServerHello(clientHello ClientHello) (ServerHello, error) { return ServerHello{ Selected: ServerSelection{ Message: MessageSelection{ - Codec: MessageCodecJSON, + Codec: MessageCodecJSON, + UDPPacketCodec: selectUDPPacketCodec(clientHello.Capabilities.Message.UDPPacketCodecs), }, Crypto: CryptoSelection{ Algorithm: algorithm, @@ -92,6 +93,15 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel if serverHello.Selected.Message.Codec != MessageCodecJSON { return fmt.Errorf("unsupported selected message codec: %s", serverHello.Selected.Message.Codec) } + udpPacketCodec := serverHello.Selected.Message.UDPPacketCodec + if udpPacketCodec != "" { + if udpPacketCodec != UDPPacketCodecBinary { + return fmt.Errorf("unsupported selected UDP packet codec: %s", udpPacketCodec) + } + if !Supports(clientHello.Capabilities.Message.UDPPacketCodecs, udpPacketCodec) { + return fmt.Errorf("selected UDP packet codec was not advertised by client: %s", udpPacketCodec) + } + } cryptoSelection := serverHello.Selected.Crypto if !IsSupportedAEADAlgorithm(cryptoSelection.Algorithm) { return fmt.Errorf("unknown selected crypto algorithm: %s", cryptoSelection.Algorithm) @@ -105,6 +115,13 @@ func ValidateServerHelloForClient(clientHello ClientHello, serverHello ServerHel return nil } +func selectUDPPacketCodec(codecs []string) string { + if Supports(codecs, UDPPacketCodecBinary) { + return UDPPacketCodecBinary + } + return "" +} + func NewCryptoContext(algorithm string, clientHelloPayload, serverHelloPayload []byte) *CryptoContext { return &CryptoContext{ Algorithm: algorithm, diff --git a/pkg/proto/wire/wire.go b/pkg/proto/wire/wire.go index 47bf5984..5a109d67 100644 --- a/pkg/proto/wire/wire.go +++ b/pkg/proto/wire/wire.go @@ -36,6 +36,7 @@ const ( FrameTypeMessage uint16 = 16 MessageCodecJSON = "json" + UDPPacketCodecBinary = "binary-v1" DefaultMaxFramePayloadSize = 64 * 1024 MagicV2 = "FRP\x00\x02\r\n" @@ -182,7 +183,8 @@ type ClientCapabilities struct { } type MessageCapabilities struct { - Codecs []string `json:"codecs,omitempty"` + Codecs []string `json:"codecs,omitempty"` + UDPPacketCodecs []string `json:"udpPacketCodecs,omitempty"` } type CryptoCapabilities struct { @@ -201,7 +203,8 @@ type ServerSelection struct { } type MessageSelection struct { - Codec string `json:"codec,omitempty"` + Codec string `json:"codec,omitempty"` + UDPPacketCodec string `json:"udpPacketCodec,omitempty"` } type CryptoSelection struct { @@ -214,7 +217,8 @@ func clientHelloWithCryptoRandom(bootstrap BootstrapInfo, clientRandom []byte) C Bootstrap: bootstrap, Capabilities: ClientCapabilities{ Message: MessageCapabilities{ - Codecs: []string{MessageCodecJSON}, + Codecs: []string{MessageCodecJSON}, + UDPPacketCodecs: []string{UDPPacketCodecBinary}, }, Crypto: CryptoCapabilities{ Algorithms: PreferredAEADAlgorithms(), diff --git a/pkg/proto/wire/wire_test.go b/pkg/proto/wire/wire_test.go index b564f712..b1c1b062 100644 --- a/pkg/proto/wire/wire_test.go +++ b/pkg/proto/wire/wire_test.go @@ -148,10 +148,40 @@ func TestNewServerHelloSelectsFirstSupportedAEADAlgorithm(t *testing.T) { serverHello, err := NewServerHello(hello) require.NoError(t, err) require.Equal(t, MessageCodecJSON, serverHello.Selected.Message.Codec) + require.Equal(t, UDPPacketCodecBinary, serverHello.Selected.Message.UDPPacketCodec) require.Equal(t, AEADAlgorithmXChaCha20Poly1305, serverHello.Selected.Crypto.Algorithm) require.Len(t, serverHello.Selected.Crypto.ServerRandom, CryptoRandomSize) } +func TestUDPPacketCodecNegotiationFallbackAndValidation(t *testing.T) { + hello := mustClientHello(t, BootstrapInfo{}) + serverHello, err := NewServerHello(hello) + require.NoError(t, err) + require.Equal(t, UDPPacketCodecBinary, serverHello.Selected.Message.UDPPacketCodec) + require.NoError(t, ValidateServerHelloForClient(hello, serverHello)) + + legacyHello := hello + legacyHello.Capabilities.Message.UDPPacketCodecs = nil + legacyServerHello, err := NewServerHello(legacyHello) + require.NoError(t, err) + require.Empty(t, legacyServerHello.Selected.Message.UDPPacketCodec) + require.NoError(t, ValidateServerHelloForClient(legacyHello, legacyServerHello)) + + unknownOffer := hello + unknownOffer.Capabilities.Message.UDPPacketCodecs = []string{"unknown"} + unknownServerHello, err := NewServerHello(unknownOffer) + require.NoError(t, err) + require.Empty(t, unknownServerHello.Selected.Message.UDPPacketCodec) + + rejected := serverHello + rejected.Selected.Message.UDPPacketCodec = "unknown" + require.ErrorContains(t, ValidateServerHelloForClient(hello, rejected), "unsupported selected UDP packet codec") + + unadvertised := serverHello + unadvertised.Selected.Message.UDPPacketCodec = UDPPacketCodecBinary + require.ErrorContains(t, ValidateServerHelloForClient(legacyHello, unadvertised), "was not advertised") +} + func TestNewClientCryptoContextValidatesServerHello(t *testing.T) { hello := mustClientHello(t, BootstrapInfo{}) serverHello, err := NewServerHello(hello) diff --git a/server/control.go b/server/control.go index 6a5d3ce1..ea680d63 100644 --- a/server/control.go +++ b/server/control.go @@ -368,7 +368,8 @@ type SessionContext struct { // server configuration ServerCfg *v1.ServerConfig // negotiated wire protocol for this client session - WireProtocol string + WireProtocol string + UDPPacketCodec string } type controlState uint8 @@ -821,6 +822,7 @@ func (ctl *Control) RegisterProxy(pxyMsg *msg.NewProxy) (remoteAddr string, err ServerCfg: ctl.sessionCtx.ServerCfg, EncryptionKey: ctl.sessionCtx.EncryptionKey, WireProtocol: ctl.sessionCtx.WireProtocol, + UDPPacketCodec: ctl.sessionCtx.UDPPacketCodec, }) if err != nil { return remoteAddr, err diff --git a/server/proxy/proxy.go b/server/proxy/proxy.go index e9ace046..046794b6 100644 --- a/server/proxy/proxy.go +++ b/server/proxy/proxy.go @@ -82,19 +82,20 @@ type Proxy interface { } type BaseProxy struct { - name string - rc *controller.ResourceController - listeners []net.Listener - usedPortsNum int - poolCount int - getWorkConnFn GetWorkConnFn - serverCfg *v1.ServerConfig - encryptionKey []byte - limiter *rate.Limiter - userInfo plugin.UserInfo - loginMsg *msg.Login - configurer v1.ProxyConfigurer - wireProtocol string + name string + rc *controller.ResourceController + listeners []net.Listener + usedPortsNum int + poolCount int + getWorkConnFn GetWorkConnFn + serverCfg *v1.ServerConfig + encryptionKey []byte + limiter *rate.Limiter + userInfo plugin.UserInfo + loginMsg *msg.Login + configurer v1.ProxyConfigurer + wireProtocol string + udpPacketCodec string mu sync.RWMutex xl *xlog.Logger @@ -469,6 +470,7 @@ type Options struct { ServerCfg *v1.ServerConfig EncryptionKey []byte WireProtocol string + UDPPacketCodec string } func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) { @@ -482,20 +484,21 @@ func NewProxy(ctx context.Context, options *Options) (pxy Proxy, err error) { } basePxy := BaseProxy{ - name: configurer.GetBaseConfig().Name, - rc: options.ResourceController, - listeners: make([]net.Listener, 0), - poolCount: options.PoolCount, - getWorkConnFn: options.GetWorkConnFn, - serverCfg: options.ServerCfg, - encryptionKey: options.EncryptionKey, - limiter: limiter, - xl: xl, - ctx: xlog.NewContext(ctx, xl), - userInfo: options.UserInfo, - loginMsg: options.LoginMsg, - configurer: configurer, - wireProtocol: options.WireProtocol, + name: configurer.GetBaseConfig().Name, + rc: options.ResourceController, + listeners: make([]net.Listener, 0), + poolCount: options.PoolCount, + getWorkConnFn: options.GetWorkConnFn, + serverCfg: options.ServerCfg, + encryptionKey: options.EncryptionKey, + limiter: limiter, + xl: xl, + ctx: xlog.NewContext(ctx, xl), + userInfo: options.UserInfo, + loginMsg: options.LoginMsg, + configurer: configurer, + wireProtocol: options.WireProtocol, + udpPacketCodec: options.UDPPacketCodec, } factory := proxyFactoryRegistry[reflect.TypeOf(configurer)] diff --git a/server/proxy/udp.go b/server/proxy/udp.go index 609d9c42..47bdfc8c 100644 --- a/server/proxy/udp.go +++ b/server/proxy/udp.go @@ -224,7 +224,13 @@ func (pxy *UDPProxy) Run() (remoteAddr string, err error) { pxy.workConn = netpkg.WrapReadWriteCloserToConn(rwc, workConn) // Plain UDP payload follows the negotiated wire protocol for message framing. - payloadConn := msg.NewConn(pxy.workConn, msg.NewReadWriter(pxy.workConn, pxy.wireProtocol)) + payloadRW, err := msg.NewUDPPacketReadWriter(pxy.workConn, pxy.wireProtocol, pxy.udpPacketCodec) + if err != nil { + xl.Errorf("create UDP packet read writer: %v", err) + pxy.workConn.Close() + continue + } + payloadConn := msg.NewConn(pxy.workConn, payloadRW) ctx, cancel := context.WithCancel(context.Background()) go workConnReaderFn(payloadConn) go workConnSenderFn(payloadConn, ctx) diff --git a/server/service.go b/server/service.go index 0a174052..e24acf66 100644 --- a/server/service.go +++ b/server/service.go @@ -470,7 +470,7 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna } } if err == nil { - ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol) + ctl, err = svr.RegisterControl(controlConn, m, internal, acceptedConn.wireProtocol, acceptedConn.udpPacketCodec) } } @@ -509,7 +509,12 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna return } case *msg.NewWorkConn: - if err := svr.RegisterWorkConn(acceptedConn.conn, m); err != nil { + if err := svr.RegisterWorkConn( + acceptedConn.conn, + m, + acceptedConn.wireProtocol, + acceptedConn.clientHelloPresent, + ); err != nil { _ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{ Error: util.GenerateResponseErrorString("invalid NewWorkConn", err, lo.FromPtr(svr.cfg.DetailedErrorsToClient)), }) @@ -547,10 +552,12 @@ func (svr *Service) completeControlLogin(ctl *Control, writeSuccess func() error } type acceptedConnection struct { - conn *msg.Conn - wireProtocol string - cryptoContext *wire.CryptoContext - firstMsg msg.Message + conn *msg.Conn + wireProtocol string + clientHelloPresent bool + udpPacketCodec string + cryptoContext *wire.CryptoContext + firstMsg msg.Message } func (svr *Service) acceptConnection(ctx context.Context, conn net.Conn) (*acceptedConnection, error) { @@ -618,6 +625,7 @@ func (ac *acceptedConnection) readFirstV2Msg(conn net.Conn, wireConn *wire.Conn) return nil, fmt.Errorf("read v2 frame: %w", err) } if frame.Type == wire.FrameTypeClientHello { + ac.clientHelloPresent = true if err := ac.handleClientHello(conn, wireConn, frame); err != nil { return nil, err } @@ -666,6 +674,7 @@ func (ac *acceptedConnection) handleClientHello(conn net.Conn, wireConn *wire.Co return fmt.Errorf("write ServerHello: %w", err) } ac.cryptoContext = cryptoContext + ac.udpPacketCodec = serverHello.Selected.Message.UDPPacketCodec return nil } @@ -759,7 +768,20 @@ func (svr *Service) RegisterControl( loginMsg *msg.Login, internal bool, wireProtocol string, + udpPacketCodec string, ) (*Control, error) { + switch wireProtocol { + case wire.ProtocolV1: + if udpPacketCodec != "" { + return nil, fmt.Errorf("UDP packet codec %q requires wire protocol v2", udpPacketCodec) + } + case wire.ProtocolV2: + if udpPacketCodec != "" && udpPacketCodec != wire.UDPPacketCodecBinary { + return nil, fmt.Errorf("unsupported UDP packet codec selection: %s", udpPacketCodec) + } + default: + return nil, fmt.Errorf("unsupported wire protocol: %s", wireProtocol) + } // If client's RunID is empty, it's a new client, we just create a new controller. // Otherwise, we check if there is one controller has the same run id. If so, we release previous controller and start new one. var err error @@ -787,15 +809,16 @@ func (svr *Service) RegisterControl( } ctl, err := NewControl(ctx, &SessionContext{ - RC: svr.rc, - PxyManager: svr.pxyManager, - PluginManager: svr.pluginManager, - AuthVerifier: authVerifier, - EncryptionKey: svr.auth.EncryptionKey(), - Conn: ctlConn, - LoginMsg: loginMsg, - ServerCfg: svr.cfg, - WireProtocol: wireProtocol, + RC: svr.rc, + PxyManager: svr.pxyManager, + PluginManager: svr.pluginManager, + AuthVerifier: authVerifier, + EncryptionKey: svr.auth.EncryptionKey(), + Conn: ctlConn, + LoginMsg: loginMsg, + ServerCfg: svr.cfg, + WireProtocol: wireProtocol, + UDPPacketCodec: udpPacketCodec, }) if err != nil { xl.Warnf("create new controller error: %v", err) @@ -820,13 +843,24 @@ func (svr *Service) RegisterControl( } // RegisterWorkConn register a new work connection to control and proxies need it. -func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn) error { +func (svr *Service) RegisterWorkConn( + workConn *msg.Conn, + newMsg *msg.NewWorkConn, + workWireProtocol string, + workClientHelloPresent bool, +) error { + if workClientHelloPresent { + return fmt.Errorf("ClientHello is not allowed on work connections") + } xl := netpkg.NewLogFromConn(workConn) ctl, exist := svr.ctlManager.GetByID(newMsg.RunID) if !exist { xl.Warnf("no client control found for run id [%s]", newMsg.RunID) return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID) } + if workWireProtocol != ctl.sessionCtx.WireProtocol { + return fmt.Errorf("work connection wire protocol mismatch: got %s want %s", workWireProtocol, ctl.sessionCtx.WireProtocol) + } // server plugin hook content := &plugin.NewWorkConnContent{ diff --git a/server/service_test.go b/server/service_test.go index 0ebec604..20769725 100644 --- a/server/service_test.go +++ b/server/service_test.go @@ -79,6 +79,70 @@ func TestWriteWithDeadlineTimesOutAndClearsDeadline(t *testing.T) { } } +func TestServiceAcceptConnectionTracksClientHelloPresence(t *testing.T) { + for _, tc := range []struct { + name string + clientHelloPresent bool + offeredCodecs []string + expectedCodec string + }{ + { + name: "absent Hello", + }, + { + name: "present Hello with JSON fallback", + clientHelloPresent: true, + }, + { + name: "present Hello with binary codec", + clientHelloPresent: true, + offeredCodecs: []string{wire.UDPPacketCodecBinary}, + expectedCodec: wire.UDPPacketCodecBinary, + }, + } { + t.Run(tc.name, func(t *testing.T) { + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + clientErrCh := make(chan error, 1) + go func() { + if err := wire.WriteMagic(clientConn); err != nil { + clientErrCh <- err + return + } + wireConn := wire.NewConn(clientConn) + if tc.clientHelloPresent { + hello, err := wire.NewClientHello(wire.BootstrapInfo{}) + if err != nil { + clientErrCh <- err + return + } + hello.Capabilities.Message.UDPPacketCodecs = tc.offeredCodecs + if err := wireConn.WriteJSONFrame(wire.FrameTypeClientHello, hello); err != nil { + clientErrCh <- err + return + } + var serverHello wire.ServerHello + if err := wireConn.ReadJSONFrame(wire.FrameTypeServerHello, &serverHello); err != nil { + clientErrCh <- err + return + } + } + clientErrCh <- msg.NewV2ReadWriterWithConn(wireConn).WriteMsg(&msg.NewWorkConn{RunID: "shared-run"}) + }() + + acceptedConn, err := (&Service{}).acceptConnection(t.Context(), serverConn) + require.NoError(t, err) + require.NoError(t, <-clientErrCh) + require.Equal(t, tc.clientHelloPresent, acceptedConn.clientHelloPresent) + require.Equal(t, tc.expectedCodec, acceptedConn.udpPacketCodec) + require.IsType(t, &msg.NewWorkConn{}, acceptedConn.firstMsg) + require.NoError(t, acceptedConn.conn.Close()) + }) + } +} + func TestSharedPortHTTPListenerProtocols(t *testing.T) { listener, err := net.Listen("tcp", "127.0.0.1:0") require.NoError(t, err) @@ -428,7 +492,7 @@ func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) { pendingConn := newCountingCloseConn() pendingMsgConn := msg.NewConn(pendingConn, msg.NewV1ReadWriter(pendingConn)) - err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"}) + err = registerWorkConnAsCaller(svr, pendingMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false) require.Error(t, err) require.Equal(t, int64(1), pendingConn.closeCount.Load()) require.Len(t, ctl.workConnCh, 0) @@ -442,7 +506,7 @@ func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) { runningConn := newCountingCloseConn() runningMsgConn := msg.NewConn(runningConn, msg.NewV1ReadWriter(runningConn)) - require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"})) + require.NoError(t, svr.RegisterWorkConn(runningMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false)) require.Len(t, ctl.workConnCh, 1) require.NoError(t, ctl.Close()) @@ -450,6 +514,123 @@ func TestServiceWorkConnRoutingRequiresCurrentRunningControl(t *testing.T) { require.Equal(t, int64(1), runningConn.closeCount.Load()) } +func TestServiceWorkConnRoutingRejectsWireProtocolMismatch(t *testing.T) { + svr := newControlTestService(t) + ctl, controlConn, err := registerLifecycleTestControl(svr) + require.NoError(t, err) + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + + workConn := newCountingCloseConn() + workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn)) + err = svr.RegisterWorkConn(workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV2, false) + require.ErrorContains(t, err, "wire protocol mismatch") + require.Len(t, ctl.workConnCh, 0) + _ = workMsgConn.Close() + require.NoError(t, ctl.Close()) +} + +func TestServiceWorkConnRoutingClientHelloPolicy(t *testing.T) { + for _, tc := range []struct { + name string + controlUDPPacketCodec string + workClientHelloPresent bool + errorSubstring string + }{ + { + name: "JSON control allows work connection without Hello", + }, + { + name: "binary control allows work connection without Hello", + controlUDPPacketCodec: wire.UDPPacketCodecBinary, + }, + { + name: "JSON control rejects work connection with Hello", + workClientHelloPresent: true, + errorSubstring: "ClientHello is not allowed", + }, + { + name: "binary control rejects work connection with Hello", + controlUDPPacketCodec: wire.UDPPacketCodecBinary, + workClientHelloPresent: true, + errorSubstring: "ClientHello is not allowed", + }, + } { + t.Run(tc.name, func(t *testing.T) { + svr := newControlTestService(t) + controlConn := newDeadlineReadConn() + controlMsgConn := msg.NewConn(controlConn, msg.NewV2ReadWriter(controlConn)) + ctl, err := svr.RegisterControl(controlMsgConn, &msg.Login{ + RunID: "shared-run", + ClientID: "client", + ClientSpec: msg.ClientSpec{ + AlwaysAuthPass: true, + }, + }, true, wire.ProtocolV2, tc.controlUDPPacketCodec) + require.NoError(t, err) + require.NoError(t, svr.completeControlLogin(ctl, func() error { return nil })) + waitForSignal(t, controlConn.readStarted, "control reader to start") + + workConn := newCountingCloseConn() + workMsgConn := msg.NewConn(workConn, msg.NewV2ReadWriter(workConn)) + err = svr.RegisterWorkConn( + workMsgConn, + &msg.NewWorkConn{RunID: "shared-run"}, + wire.ProtocolV2, + tc.workClientHelloPresent, + ) + if tc.errorSubstring != "" { + require.ErrorContains(t, err, tc.errorSubstring) + require.Len(t, ctl.workConnCh, 0) + require.NoError(t, workMsgConn.Close()) + } else { + require.NoError(t, err) + require.Len(t, ctl.workConnCh, 1) + } + + require.NoError(t, ctl.Close()) + waitForControlDone(t, ctl) + require.Equal(t, int64(1), workConn.closeCount.Load()) + }) + } +} + +func TestServiceRegisterControlRejectsInvalidCodecSelection(t *testing.T) { + for _, tc := range []struct { + name string + wireProtocol string + udpPacketCodec string + errorSubstring string + }{ + { + name: "binary codec over v1", + wireProtocol: wire.ProtocolV1, + udpPacketCodec: wire.UDPPacketCodecBinary, + errorSubstring: "requires wire protocol v2", + }, + { + name: "unknown v2 codec", + wireProtocol: wire.ProtocolV2, + udpPacketCodec: "unknown", + errorSubstring: "unsupported UDP packet codec", + }, + { + name: "unknown wire protocol", + wireProtocol: "unknown", + errorSubstring: "unsupported wire protocol", + }, + } { + t.Run(tc.name, func(t *testing.T) { + svr := newControlTestService(t) + conn := newDeadlineReadConn() + msgConn := msg.NewConn(conn, msg.NewV1ReadWriter(conn)) + ctl, err := svr.RegisterControl(msgConn, &msg.Login{}, true, tc.wireProtocol, tc.udpPacketCodec) + require.Nil(t, ctl) + require.ErrorContains(t, err, tc.errorSubstring) + }) + } +} + func TestServiceWorkConnRoutingRejectsLostGeneration(t *testing.T) { for _, action := range []string{"replace", "close"} { t.Run(action, func(t *testing.T) { @@ -465,7 +646,7 @@ func TestServiceWorkConnRoutingRejectsLostGeneration(t *testing.T) { workMsgConn := msg.NewConn(workConn, msg.NewV1ReadWriter(workConn)) routeDone := make(chan error, 1) go func() { - routeDone <- registerWorkConnAsCaller(svr, workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}) + routeDone <- registerWorkConnAsCaller(svr, workMsgConn, &msg.NewWorkConn{RunID: "shared-run"}, wire.ProtocolV1, false) }() waitForSignal(t, barrier.entered, "work connection plugin barrier") @@ -509,7 +690,7 @@ func TestServiceVisitorRoutingExcludesPendingUser(t *testing.T) { ClientSpec: msg.ClientSpec{ AlwaysAuthPass: true, }, - }, true, wire.ProtocolV1) + }, true, wire.ProtocolV1, "") require.NoError(t, err) timestamp := time.Now().Unix() @@ -567,7 +748,7 @@ func registerLifecycleTestControl(svr *Service) (*Control, *deadlineReadConn, er ClientSpec: msg.ClientSpec{ AlwaysAuthPass: true, }, - }, true, wire.ProtocolV1) + }, true, wire.ProtocolV1, "") return ctl, conn, err } @@ -584,8 +765,14 @@ func waitForDifferentCurrentControl(t *testing.T, manager *ControlManager, runID return nil } -func registerWorkConnAsCaller(svr *Service, workConn *msg.Conn, newMsg *msg.NewWorkConn) error { - err := svr.RegisterWorkConn(workConn, newMsg) +func registerWorkConnAsCaller( + svr *Service, + workConn *msg.Conn, + newMsg *msg.NewWorkConn, + wireProtocol string, + clientHelloPresent bool, +) error { + err := svr.RegisterWorkConn(workConn, newMsg, wireProtocol, clientHelloPresent) if err != nil { _ = workConn.Close() } diff --git a/test/e2e/compatibility/compatibility_test.go b/test/e2e/compatibility/compatibility_test.go index 246bcac1..ad505510 100644 --- a/test/e2e/compatibility/compatibility_test.go +++ b/test/e2e/compatibility/compatibility_test.go @@ -192,6 +192,41 @@ transport.wireProtocol = "v2" }) }) +var _ = ginkgo.Describe("[Compatibility: BinaryUDPPacket]", func() { + f := framework.NewDefaultFramework() + + ginkgo.BeforeEach(func() { + supportsV2, knownVersion := baselineSupportsControlWireProtocolV2(compatCtx.BaselineVersion) + if !knownVersion || !supportsV2 { + ginkgo.Skip(fmt.Sprintf("baseline version %q does not have known wire protocol v2 support", compatCtx.BaselineVersion)) + } + }) + + ginkgo.It("current frps falls back to JSON for baseline frpc", func() { + portName := port.GenName("CompatBinaryUDPBaselineFRPC") + clientConf := udpClientConfig("udp", portName, `transport.wireProtocol = "v2"`) + f.RunProcessesWithBinaries( + compatCtx.CurrentFRPSPath, + compatCtx.BaselineFRPCPath, + consts.DefaultServerConfig, + []string{clientConf}, + ) + framework.NewRequestExpect(f).Protocol("udp").PortName(portName).Ensure() + }) + + ginkgo.It("current frpc falls back to JSON for baseline frps", func() { + portName := port.GenName("CompatBinaryUDPBaselineFRPS") + clientConf := udpClientConfig("udp", portName, `transport.wireProtocol = "v2"`) + f.RunProcessesWithBinaries( + compatCtx.BaselineFRPSPath, + compatCtx.CurrentFRPCPath, + consts.DefaultServerConfig, + []string{clientConf}, + ) + framework.NewRequestExpect(f).Protocol("udp").PortName(portName).Ensure() + }) +}) + func tcpClientConfig(proxyName string, remotePortName string, extra string) string { return fmt.Sprintf(` serverAddr = "127.0.0.1" @@ -208,6 +243,22 @@ remotePort = {{ .%s }} `, consts.PortServerName, extra, proxyName, framework.TCPEchoServerPort, remotePortName) } +func udpClientConfig(proxyName string, remotePortName string, extra string) string { + return fmt.Sprintf(` +serverAddr = "127.0.0.1" +serverPort = {{ .%s }} +loginFailExit = true +log.level = "trace" +%s + +[[proxies]] +name = "%s" +type = "udp" +localPort = {{ .%s }} +remotePort = {{ .%s }} +`, consts.PortServerName, extra, proxyName, framework.UDPEchoServerPort, remotePortName) +} + func expectProcessExit(p *process.Process, timeout time.Duration) { select { case <-p.Done(): diff --git a/test/e2e/v1/basic/wire.go b/test/e2e/v1/basic/wire.go index 051e6a6c..cb61558f 100644 --- a/test/e2e/v1/basic/wire.go +++ b/test/e2e/v1/basic/wire.go @@ -5,6 +5,7 @@ import ( "fmt" "io" "net/http" + "time" "github.com/onsi/ginkgo/v2" @@ -204,6 +205,87 @@ var _ = ginkgo.Describe("[Feature: WireProtocol]", func() { }) }) +var _ = ginkgo.Describe("[Feature: BinaryUDPPacket]", func() { + f := framework.NewDefaultFramework() + + for _, tc := range []struct { + name string + protocol string + extraServer string + extraTransport string + }{ + {name: "tcp mux on", protocol: "tcp", extraTransport: "transport.tcpMux = true"}, + {name: "tcp mux off", protocol: "tcp", extraServer: "transport.tcpMux = false", extraTransport: "transport.tcpMux = false"}, + {name: "kcp", protocol: "kcp"}, + {name: "quic stream", protocol: "quic"}, + {name: "websocket", protocol: "websocket"}, + } { + ginkgo.It(tc.name, func() { + runClientServerTest(f, &generalTestConfigures{ + server: renderBindPortConfig(tc.protocol) + "\n" + tc.extraServer, + client: fmt.Sprintf(` + transport.wireProtocol = "v2" + transport.protocol = %q + %s + `, tc.protocol, tc.extraTransport), + }) + }) + } + + ginkgo.It("wss", func() { + wssPort := f.AllocPort() + runClientServerTest(f, &generalTestConfigures{ + clientPrefix: fmt.Sprintf(` + serverAddr = "127.0.0.1" + serverPort = %d + loginFailExit = false + transport.protocol = "wss" + transport.wireProtocol = "v2" + log.level = "trace" + `, wssPort), + client2: fmt.Sprintf(` + [[proxies]] + name = "wss2ws" + type = "tcp" + remotePort = %d + [proxies.plugin] + type = "https2http" + localAddr = "127.0.0.1:{{ .%s }}" + `, wssPort, consts.PortServerName), + testDelay: 10 * time.Second, + }) + }) + + for _, tc := range []struct { + name string + transport string + }{ + {name: "plain"}, + {name: "aes-cfb", transport: "transport.useEncryption = true"}, + {name: "snappy", transport: "transport.useCompression = true"}, + {name: "snappy and aes-cfb", transport: "transport.useEncryption = true\ntransport.useCompression = true"}, + {name: "limiter", transport: "transport.bandwidthLimit = \"1MB\""}, + } { + ginkgo.It(tc.name, func() { + serverConf := consts.DefaultServerConfig + udpPortName := port.GenName("BinaryUDPPacket") + clientConf := consts.DefaultClientConfig + fmt.Sprintf(` + transport.wireProtocol = "v2" + + [[proxies]] + name = "udp" + type = "udp" + localPort = {{ .%s }} + remotePort = {{ .%s }} + %s + `, framework.UDPEchoServerPort, udpPortName, tc.transport) + + f.RunProcesses(serverConf, []string{clientConf}) + framework.NewRequestExpect(f).Protocol("udp").PortName(udpPortName).Ensure() + }) + } +}) + type wireClientInfo struct { ClientID string `json:"clientID"` WireProtocol string `json:"wireProtocol"`