mirror of
https://github.com/amnezia-vpn/amneziawg-go.git
synced 2026-10-02 21:36:06 +03:00
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:
1 parent
f8b2d2ba5a
commit
9965d6d761
5 files changed
+436
-7
No files matched your search
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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.")
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in new issue
Block a user