fix(conn): route UDP conceal format errors to fallback

Add tuntest coverage for TCP/UDP TUN packets through UDP and TCP masquerade
pipelines. Cover TCP and UDP fallback on the device path, and preserve
format errors from UDP datagram decoding so invalid packets can be rerouted.
This commit is contained in:
Frog Rocky authored and Yaroslav Gurov committed 2026-05-11 16:59:04 +02:00
1 parent f8b2d2ba5a
commit 9965d6d761
5 files changed
+436 -7

No files matched your search

+13 -4
View File
@@ -2,6 +2,7 @@ package conceal
import (
"encoding/binary"
"errors"
"sync"
)
@@ -85,11 +86,19 @@ func (p *UDPDatagramPipeline) Encode(dst, src []byte) (int, error) {
}
func (p *UDPDatagramPipeline) DecodeInPlace(buf []byte, n int) (int, bool) {
n, err := p.DecodeInPlaceErr(buf, n)
return n, err == nil
}
func (p *UDPDatagramPipeline) DecodeInPlaceErr(buf []byte, n int) (int, error) {
if p.masqueradeActive {
var err error
n, err = p.masquerade.DecodeInPlace(buf, n)
if err != nil {
return 0, false
if !errors.Is(err, ErrFormat) {
err = NewFormatError(buf[:n], err)
}
return 0, err
}
}
@@ -97,15 +106,15 @@ func (p *UDPDatagramPipeline) DecodeInPlace(buf []byte, n int) (int, bool) {
var err error
n, err = p.framing.Decode(buf[:n])
if err != nil {
return 0, false
return 0, err
}
}
if !p.classifier.IsValid(buf[:n]) {
return 0, false
return 0, NewFormatError(buf[:n], errInvalidData)
}
return n, true
return n, nil
}
func (p *UDPDatagramPipeline) EmitPrelude(packet []byte, emit func([]byte) error) error {
+8
View File
@@ -27,6 +27,14 @@ func fallbackUDPAddress(addr net.Addr, port uint16) *net.UDPAddr {
return &net.UDPAddr{IP: ip, Port: int(port)}
}
func fallbackUDPAddressForEndpoint(ep Endpoint, port uint16) *net.UDPAddr {
ip := net.IPv4(127, 0, 0, 1)
if ep != nil && ep.DstIP().Is6() {
ip = net.IPv6loopback
}
return &net.UDPAddr{IP: ip, Port: int(port)}
}
func (b *BindStream) proxyStreamFallback(ep *streamEndpoint, first []byte) {
rawConn := ep.rawConn
if rawConn == nil {
+142 -2
View File
@@ -2,6 +2,8 @@ package conn
import (
"bytes"
"errors"
"net"
"sync"
"sync/atomic"
@@ -13,6 +15,7 @@ var (
_ Framable = (*ConcealBind)(nil)
_ Preludable = (*ConcealBind)(nil)
_ Masqueradable = (*ConcealBind)(nil)
_ Fallbackable = (*ConcealBind)(nil)
)
type ConcealBind struct {
@@ -25,6 +28,9 @@ type ConcealBind struct {
framedOpts conceal.FramedOpts
preludeOpts conceal.PreludeOpts
masqueradeOpts conceal.MasqueradeOpts
fallbackPort uint16
fallbackSessions map[string]*concealBindFallbackSession
pipeline atomic.Pointer[conceal.UDPDatagramPipeline]
}
@@ -92,8 +98,15 @@ func (b *ConcealBind) wrapReceiveFunc(fn ReceiveFunc) ReceiveFunc {
continue
}
size, ok := pipeline.DecodeInPlace(packets[i], sizes[i])
if !ok {
size, err := pipeline.DecodeInPlaceErr(packets[i], sizes[i])
if err != nil {
if errors.Is(err, conceal.ErrFormat) {
data := conceal.FormatErrorData(err)
if len(data) == 0 {
data = bytes.Clone(packets[i][:sizes[i]])
}
_ = b.forwardFallbackUDP(eps[i], data)
}
sizes[i] = 0
eps[i] = nil
continue
@@ -106,6 +119,9 @@ func (b *ConcealBind) wrapReceiveFunc(fn ReceiveFunc) ReceiveFunc {
}
func (b *ConcealBind) Close() error {
b.mu.Lock()
b.closeFallbackSessionsLocked()
b.mu.Unlock()
return b.inner.Close()
}
@@ -228,3 +244,127 @@ func (b *ConcealBind) SetMasqueradeOpts(opts conceal.MasqueradeOpts) {
b.masqueradeOpts = opts
b.rebuildPipelineLocked()
}
func (b *ConcealBind) SetFallbackPort(port uint16) {
b.mu.Lock()
if b.fallbackPort != port {
b.closeFallbackSessionsLocked()
}
b.fallbackPort = port
b.mu.Unlock()
if fallbackable, ok := b.inner.(Fallbackable); ok {
fallbackable.SetFallbackPort(port)
}
}
func (b *ConcealBind) forwardFallbackUDP(ep Endpoint, data []byte) error {
if ep == nil || len(data) == 0 {
return nil
}
port := b.currentFallbackPort()
if port == 0 {
return nil
}
session, err := b.fallbackSession(ep, port)
if err != nil {
return err
}
_, err = session.conn.Write(data)
return err
}
func (b *ConcealBind) currentFallbackPort() uint16 {
b.mu.Lock()
defer b.mu.Unlock()
return b.fallbackPort
}
func (b *ConcealBind) fallbackSession(ep Endpoint, port uint16) (*concealBindFallbackSession, error) {
key := ep.DstToString()
b.mu.Lock()
if session := b.fallbackSessions[key]; session != nil {
b.mu.Unlock()
return session, nil
}
b.mu.Unlock()
fallbackConn, err := net.DialUDP("udp", nil, fallbackUDPAddressForEndpoint(ep, port))
if err != nil {
return nil, err
}
session := &concealBindFallbackSession{
parent: b,
key: key,
ep: ep,
conn: fallbackConn,
}
b.mu.Lock()
if b.fallbackPort != port || b.fallbackPort == 0 {
b.mu.Unlock()
fallbackConn.Close()
return nil, net.ErrClosed
}
if b.fallbackSessions == nil {
b.fallbackSessions = make(map[string]*concealBindFallbackSession)
}
if existing := b.fallbackSessions[key]; existing != nil {
b.mu.Unlock()
fallbackConn.Close()
return existing, nil
}
b.fallbackSessions[key] = session
b.mu.Unlock()
go session.relay()
return session, nil
}
func (b *ConcealBind) removeFallbackSession(key string, session *concealBindFallbackSession) {
b.mu.Lock()
if b.fallbackSessions[key] == session {
delete(b.fallbackSessions, key)
}
b.mu.Unlock()
}
func (b *ConcealBind) closeFallbackSessionsLocked() {
for key, session := range b.fallbackSessions {
session.close()
delete(b.fallbackSessions, key)
}
}
type concealBindFallbackSession struct {
parent *ConcealBind
key string
ep Endpoint
conn *net.UDPConn
once sync.Once
}
func (s *concealBindFallbackSession) relay() {
defer s.parent.removeFallbackSession(s.key, s)
buf := make([]byte, 65535)
for {
n, err := s.conn.Read(buf)
if err != nil {
return
}
if err := s.parent.inner.Send([][]byte{buf[:n]}, s.ep); err != nil {
return
}
}
}
func (s *concealBindFallbackSession) close() {
s.once.Do(func() {
s.conn.Close()
})
}
+213 -1
View File
@@ -179,6 +179,16 @@ func (pair *testPair) Send(
tb testing.TB,
ping SendDirection,
done chan struct{},
) {
tb.Helper()
pair.SendPacket(tb, ping, tuntest.Ping, done)
}
func (pair *testPair) SendPacket(
tb testing.TB,
ping SendDirection,
packet func(dst, src netip.Addr) []byte,
done chan struct{},
) {
tb.Helper()
p0, p1 := pair[0], pair[1]
@@ -187,7 +197,7 @@ func (pair *testPair) Send(
p0, p1 = p1, p0
}
msg := tuntest.Ping(p0.ip, p1.ip)
msg := packet(p0.ip, p1.ip)
p1.tun.Outbound <- msg
timer := time.NewTimer(6 * time.Second)
defer timer.Stop()
@@ -381,6 +391,208 @@ func TestAWGDevicePingTCPConceal(t *testing.T) {
})
}
func TestAWGDeviceUDPMasqueradeTransportsTUNTCPAndUDP(t *testing.T) {
goroutineLeakCheck(t)
pair := genTestPair(t, true,
"format_in", "<b 0xfeed><dz be 2><d>",
"format_out", "<b 0xfeed><dz be 2><d>",
)
for _, tc := range []struct {
name string
packet func(dst, src netip.Addr) []byte
}{
{name: "udp", packet: tuntest.UDP},
{name: "tcp", packet: tuntest.TCP},
} {
t.Run(tc.name+"_ping", func(t *testing.T) {
pair.SendPacket(t, Ping, tc.packet, nil)
})
t.Run(tc.name+"_pong", func(t *testing.T) {
pair.SendPacket(t, Pong, tc.packet, nil)
})
}
}
func TestAWGDeviceTCPMasqueradeTransportsTUNTCPAndUDP(t *testing.T) {
goroutineLeakCheck(t)
pair := genTestPairTCP(t,
"network", "tcp",
"format_in", "<b 0xfeed><dz be 2><d>",
"format_out", "<b 0xfeed><dz be 2><d>",
)
for _, tc := range []struct {
name string
packet func(dst, src netip.Addr) []byte
}{
{name: "udp", packet: tuntest.UDP},
{name: "tcp", packet: tuntest.TCP},
} {
t.Run(tc.name+"_ping", func(t *testing.T) {
pair.SendPacket(t, Ping, tc.packet, nil)
})
t.Run(tc.name+"_pong", func(t *testing.T) {
pair.SendPacket(t, Pong, tc.packet, nil)
})
}
}
func TestAWGDeviceTCPFallbackWithTUNDevice(t *testing.T) {
fallbackListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen fallback: %v", err)
}
defer fallbackListener.Close()
probe := []byte("GET / HTTP/1.1\r\n")
response := []byte("HTTP/1.1 200 OK\r\n\r\n")
received := make(chan []byte, 1)
served := make(chan error, 1)
go func() {
conn, err := fallbackListener.Accept()
if err != nil {
served <- err
return
}
defer conn.Close()
buf := make([]byte, len(probe))
if _, err := io.ReadFull(conn, buf); err != nil {
served <- err
return
}
received <- bytes.Clone(buf)
_, err = conn.Write(response)
served <- err
}()
device := newFallbackTestDevice(t,
"tcp",
getFreeTCPPort(t),
fallbackListener.Addr().(*net.TCPAddr).Port,
)
client, err := net.Dial("tcp", net.JoinHostPort("127.0.0.1", strconv.Itoa(int(device.net.port))))
if err != nil {
t.Fatalf("dial device: %v", err)
}
defer client.Close()
if err := client.SetDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set client deadline: %v", err)
}
if _, err := client.Write(probe); err != nil {
t.Fatalf("write probe: %v", err)
}
got := make([]byte, len(response))
if _, err := io.ReadFull(client, got); err != nil {
t.Fatalf("read fallback response: %v", err)
}
if !bytes.Equal(got, response) {
t.Fatalf("fallback response = %q, want %q", got, response)
}
client.Close()
assertFallbackReceived(t, received, served, probe)
}
func TestAWGDeviceUDPFallbackWithTUNDevice(t *testing.T) {
fallbackConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatalf("listen fallback udp: %v", err)
}
defer fallbackConn.Close()
probe := []byte("GE")
response := []byte("udp-ok")
received := make(chan []byte, 1)
served := make(chan error, 1)
go func() {
buf := make([]byte, 64)
n, addr, err := fallbackConn.ReadFromUDP(buf)
if err != nil {
served <- err
return
}
received <- bytes.Clone(buf[:n])
_, err = fallbackConn.WriteToUDP(response, addr)
served <- err
}()
device := newFallbackTestDevice(t, "udp", 0, fallbackConn.LocalAddr().(*net.UDPAddr).Port)
client, err := net.DialUDP("udp4", nil, &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: int(device.net.port)})
if err != nil {
t.Fatalf("dial device udp: %v", err)
}
defer client.Close()
if err := client.SetDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set udp client deadline: %v", err)
}
if _, err := client.Write(probe); err != nil {
t.Fatalf("write udp probe: %v", err)
}
got := make([]byte, len(response))
if _, err := io.ReadFull(client, got); err != nil {
t.Fatalf("read udp fallback response: %v", err)
}
if !bytes.Equal(got, response) {
t.Fatalf("udp fallback response = %q, want %q", got, response)
}
assertFallbackReceived(t, received, served, probe)
}
func newFallbackTestDevice(t *testing.T, network string, listenPort, fallbackPort int) *Device {
t.Helper()
var key NoisePrivateKey
if _, err := rand.Read(key[:]); err != nil {
t.Fatalf("generate private key: %v", err)
}
tunDevice := tuntest.NewChannelTUN()
device := NewDevice(tunDevice.TUN(), conn.NewDefaultBind(), NewLogger(LogLevelError, "fallback-test: "))
cfg := uapiCfg(
"private_key", hex.EncodeToString(key[:]),
"listen_port", strconv.Itoa(listenPort),
"network", network,
"format_in", "<b 0xfeed><dz be 2><d>",
"fallback_port", strconv.Itoa(fallbackPort),
)
if err := device.IpcSet(cfg); err != nil {
device.Close()
t.Fatalf("configure fallback device: %v", err)
}
if err := device.Up(); err != nil {
device.Close()
t.Fatalf("bring up fallback device: %v", err)
}
t.Cleanup(device.Close)
return device
}
func assertFallbackReceived(t *testing.T, received <-chan []byte, served <-chan error, want []byte) {
t.Helper()
select {
case got := <-received:
if !bytes.Equal(got, want) {
t.Fatalf("fallback received = %q, want %q", got, want)
}
case <-time.After(2 * time.Second):
t.Fatal("fallback service did not receive probe")
}
if err := <-served; err != nil {
t.Fatalf("fallback service failed: %v", err)
}
}
// Needs to be stopped with Ctrl-C
func TestAWGHandshakeDevicePing(t *testing.T) {
t.Skip("This test is intended to be run manually, not as part of the test suite.")
+60
View File
@@ -25,6 +25,48 @@ func Ping(dst, src netip.Addr) []byte {
return genICMPv4(payload, dst, src)
}
func UDP(dst, src netip.Addr) []byte {
payload := []byte("tuntest-udp")
const (
udpProtocolNumber = 17
ipv4Size = 20
udpSize = 8
)
pkt := make([]byte, ipv4Size+udpSize+len(payload))
udp := pkt[ipv4Size : ipv4Size+udpSize]
binary.BigEndian.PutUint16(udp[0:], 1337)
binary.BigEndian.PutUint16(udp[2:], 7331)
binary.BigEndian.PutUint16(udp[4:], uint16(udpSize+len(payload)))
copy(pkt[ipv4Size+udpSize:], payload)
fillIPv4Header(pkt[:ipv4Size], dst, src, udpProtocolNumber, len(pkt))
return pkt
}
func TCP(dst, src netip.Addr) []byte {
payload := []byte("tuntest-tcp")
const (
tcpProtocolNumber = 6
ipv4Size = 20
tcpSize = 20
)
pkt := make([]byte, ipv4Size+tcpSize+len(payload))
tcp := pkt[ipv4Size : ipv4Size+tcpSize]
binary.BigEndian.PutUint16(tcp[0:], 1337)
binary.BigEndian.PutUint16(tcp[2:], 7331)
binary.BigEndian.PutUint32(tcp[4:], 1)
binary.BigEndian.PutUint32(tcp[8:], 1)
tcp[12] = tcpSize / 4 << 4
tcp[13] = 0x18 // PSH|ACK
binary.BigEndian.PutUint16(tcp[14:], 65535)
copy(pkt[ipv4Size+tcpSize:], payload)
fillIPv4Header(pkt[:ipv4Size], dst, src, tcpProtocolNumber, len(pkt))
return pkt
}
// Checksum is the "internet checksum" from https://tools.ietf.org/html/rfc1071.
func checksum(buf []byte, initial uint16) uint16 {
v := uint32(initial)
@@ -79,6 +121,24 @@ func genICMPv4(payload []byte, dst, src netip.Addr) []byte {
return pkt
}
func fillIPv4Header(ip []byte, dst, src netip.Addr, protocol, length int) {
const (
ipv4Size = 20
ipv4TotalLenOffset = 2
ipv4ChecksumOffset = 10
ttl = 65
)
ip[0] = (4 << 4) | (ipv4Size / 4)
binary.BigEndian.PutUint16(ip[ipv4TotalLenOffset:], uint16(length))
ip[8] = ttl
ip[9] = byte(protocol)
copy(ip[12:], src.AsSlice())
copy(ip[16:], dst.AsSlice())
chksum := ^checksum(ip[:], 0)
binary.BigEndian.PutUint16(ip[ipv4ChecksumOffset:], chksum)
}
type ChannelTUN struct {
Inbound chan []byte // incoming packets, closed on TUN close
Outbound chan []byte // outbound packets, blocks forever on TUN close