feat: awg 3.1 features

* add RandomTrailers feature which appends random amount of bytes to the
end of each packet
* add DisableCookie feature which prohibits interface to send any cookie
replies
This commit is contained in:
Yaroslav Gurov authored and Yaroslav Gurov committed 2026-08-13 00:25:03 +02:00
1 parent 08d68cdae2
commit 1f50ad736e
6 files changed
+143 -42

No files matched your search

+1
View File
@@ -24,6 +24,7 @@ const (
CookieRefreshTime = time.Second * 120 CookieRefreshTime = time.Second * 120
HandshakeInitationRate = time.Second / 50 HandshakeInitationRate = time.Second / 50
PaddingMultiple = 16 PaddingMultiple = 16
DefaultUdpWindow = 500
) )
const ( const (
+3
View File
@@ -126,6 +126,9 @@ type Device struct {
keepaliveTimeoutSec AtomicUintRange keepaliveTimeoutSec AtomicUintRange
maxHandshakeAttemps AtomicUintRange maxHandshakeAttemps AtomicUintRange
} }
randomTrailers atomic.Bool
disableCookies atomic.Bool
} }
// deviceState represents the state of a Device. // deviceState represents the state of a Device.
+6
View File
@@ -57,6 +57,7 @@ type Peer struct {
cookieGenerator CookieGenerator cookieGenerator CookieGenerator
trieEntries list.List trieEntries list.List
persistentKeepaliveInterval AtomicUintRange persistentKeepaliveInterval AtomicUintRange
udpWindow atomic.Uint32
} }
func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) { func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
@@ -79,6 +80,8 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
// create peer // create peer
peer := new(Peer) peer := new(Peer)
peer.udpWindow.Store(DefaultUdpWindow)
peer.cookieGenerator.Init(pk) peer.cookieGenerator.Init(pk)
peer.device = device peer.device = device
peer.queue.outbound = newAutodrainingOutboundQueue(device) peer.queue.outbound = newAutodrainingOutboundQueue(device)
@@ -283,6 +286,9 @@ func (peer *Peer) SetEndpointFromPacket(endpoint conn.Endpoint) {
if peer.endpoint.disableRoaming { if peer.endpoint.disableRoaming {
return return
} }
if peer.endpoint.val != endpoint {
peer.udpWindow.Store(DefaultUdpWindow)
}
peer.endpoint.clearSrcOnTx = false peer.endpoint.clearSrcOnTx = false
peer.endpoint.val = endpoint peer.endpoint.val = endpoint
} }
+46 -36
View File
@@ -154,8 +154,12 @@ func (device *Device) RoutineReceiveIncoming(
} }
// get message padding and type based on information from S1-S4 and H1-H4 // get message padding and type based on information from S1-S4 and H1-H4
msgType, padding := device.DeterminePacketTypeAndPadding(packet, MessageUnknownType, typeHash) msgSize, msgType, padding := device.DeterminePacketTypeAndPadding(packet, typeHash)
packet = packet[padding:] packet = packet[padding:]
if msgType != MessageTransportType {
packet = packet[:msgSize]
}
if cip != nil { if cip != nil {
applyHash(packet[:4], packet[:4], typeHash) applyHash(packet[:4], packet[:4], typeHash)
@@ -510,6 +514,11 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
} }
rxBytesLen += uint64(len(elem.packet) + MinMessageSize) rxBytesLen += uint64(len(elem.packet) + MinMessageSize)
udpWindow := elem.padding + MessageTransportHeaderSize + uint32(len(elem.packet))
if peer.udpWindow.Load() < udpWindow {
peer.udpWindow.Store(udpWindow)
}
if len(elem.packet) == 0 || elem.packet[0] == 0 { if len(elem.packet) == 0 || elem.packet[0] == 0 {
device.log.Verbosef("%v - Receiving keepalive packet", peer) device.log.Verbosef("%v - Receiving keepalive packet", peer)
continue continue
@@ -558,7 +567,7 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
continue continue
} }
bufs = append(bufs, elem.buffer[int(elem.padding):int(elem.padding)+MessageTransportOffsetContent+len(elem.packet)]) bufs = append(bufs, elem.buffer[int(elem.padding):int(elem.padding)+MessageTransportHeaderSize+len(elem.packet)])
} }
peer.rxBytes.Add(rxBytesLen) peer.rxBytes.Add(rxBytesLen)
@@ -592,57 +601,58 @@ func applyHash(dst, src, hash []byte) {
} }
} }
func (device *Device) DeterminePacketTypeAndPadding(packet []byte, expectedType uint32, typeHash []byte) (uint32, uint32) { func (device *Device) DeterminePacketTypeAndPadding(packet []byte, typeHash []byte) (int, uint32, uint32) {
var headerBytes [4]byte var headerBytes [4]byte
var padding uint32
var header UintRange
var expectedSize int
size := len(packet) size := len(packet)
randomTrailers := device.randomTrailers.Load()
if expectedType == MessageUnknownType || expectedType == MessageInitiationType { padding = device.paddings.init.Load()
padding := device.paddings.init.Load() header = device.headers.init.Load()
header := device.headers.init.Load() expectedSize = int(padding) + MessageInitiationSize
if size == int(padding)+MessageInitiationSize { if size == expectedSize || randomTrailers && size > expectedSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash) applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) { if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageInitiationType, padding return MessageInitiationSize, MessageInitiationType, padding
}
} }
} }
if expectedType == MessageUnknownType || expectedType == MessageResponseType { padding = device.paddings.response.Load()
padding := device.paddings.response.Load() header = device.headers.response.Load()
header := device.headers.response.Load() expectedSize = int(padding) + MessageResponseSize
if size == int(padding)+MessageResponseSize { if size == expectedSize || randomTrailers && size > expectedSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash) applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) { if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageResponseType, padding return MessageResponseSize, MessageResponseType, padding
}
} }
} }
if expectedType == MessageUnknownType || expectedType == MessageCookieReplyType { padding = device.paddings.cookie.Load()
padding := device.paddings.cookie.Load() header = device.headers.cookie.Load()
header := device.headers.cookie.Load() expectedSize = int(padding) + MessageCookieReplySize
if size == int(padding)+MessageCookieReplySize { if size == expectedSize || randomTrailers && size > expectedSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash) applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) { if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageCookieReplyType, padding return MessageCookieReplySize, MessageCookieReplyType, padding
}
} }
} }
if expectedType == MessageUnknownType || expectedType == MessageTransportType { padding = device.paddings.transport.Load()
padding := device.paddings.transport.Load() header = device.headers.transport.Load()
header := device.headers.transport.Load() expectedSize = int(padding) + MessageTransportSize
if size >= int(padding)+MessageTransportHeaderSize { if size >= expectedSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash) applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) { if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageTransportType, padding return MessageTransportSize, MessageTransportType, padding
}
} }
} }
return MessageUnknownType, 0 return 0, MessageUnknownType, 0
} }
+57 -6
View File
@@ -146,8 +146,10 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
sendBuffer = append(sendBuffer, peer.device.JunkPackets()...) sendBuffer = append(sendBuffer, peer.device.JunkPackets()...)
padding := peer.device.paddings.init.Load() padding := int(peer.device.paddings.init.Load())
buf := make([]byte, padding+MessageInitiationSize) trailerLen := max(peer.randomTrailer(padding+MessageInitiationSize), 0)
buf := make([]byte, padding+MessageInitiationSize+trailerLen)
crypt := buf[:padding] crypt := buf[:padding]
rand.Read(crypt) rand.Read(crypt)
@@ -168,6 +170,9 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
cip.XORKeyStream(packet, packet) cip.XORKeyStream(packet, packet)
} }
trailer := buf[padding+MessageInitiationSize:]
rand.Read(trailer)
sendBuffer = append(sendBuffer, buf) sendBuffer = append(sendBuffer, buf)
err = peer.SendBuffers(sendBuffer) err = peer.SendBuffers(sendBuffer)
if err != nil { if err != nil {
@@ -191,8 +196,10 @@ func (peer *Peer) SendHandshakeResponse() error {
return err return err
} }
padding := peer.device.paddings.response.Load() padding := int(peer.device.paddings.response.Load())
buf := make([]byte, padding+MessageResponseSize) trailerLen := max(peer.randomTrailer(padding+MessageResponseSize), 0)
buf := make([]byte, padding+MessageResponseSize+trailerLen)
crypt := buf[:padding] crypt := buf[:padding]
rand.Read(crypt) rand.Read(crypt)
@@ -220,6 +227,9 @@ func (peer *Peer) SendHandshakeResponse() error {
cip.XORKeyStream(packet, packet) cip.XORKeyStream(packet, packet)
} }
trailer := buf[padding+MessageResponseSize:]
rand.Read(trailer)
// TODO: allocation could be avoided // TODO: allocation could be avoided
err = peer.SendBuffers([][]byte{buf}) err = peer.SendBuffers([][]byte{buf})
if err != nil { if err != nil {
@@ -229,6 +239,11 @@ func (peer *Peer) SendHandshakeResponse() error {
} }
func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) error { func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) error {
if device.disableCookies.Load() {
device.log.Verbosef("Sending cookie response blocked for %v due to disabled cookies", initiatingElem.endpoint.DstToString())
return nil
}
device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString()) device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString())
sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8]) sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8])
@@ -245,8 +260,10 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement)
return err return err
} }
padding := device.paddings.cookie.Load() padding := int(device.paddings.cookie.Load())
buf := make([]byte, padding+MessageCookieReplySize) trailerLen := max(device.randomTrailer(padding+MessageCookieReplySize), 0)
buf := make([]byte, padding+MessageCookieReplySize, trailerLen)
crypt := buf[:padding] crypt := buf[:padding]
rand.Read(crypt) rand.Read(crypt)
@@ -263,6 +280,9 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement)
cip.XORKeyStream(packet, packet) cip.XORKeyStream(packet, packet)
} }
trailer := buf[padding+MessageCookieReplySize:]
rand.Read(trailer)
// TODO: allocation could be avoided // TODO: allocation could be avoided
device.net.bind.Send([][]byte{buf}, initiatingElem.endpoint) device.net.bind.Send([][]byte{buf}, initiatingElem.endpoint)
return nil return nil
@@ -531,6 +551,29 @@ func (device *Device) randomPaddingAddition(packetSize, mtu int) int {
return add return add
} }
func (device *Device) randomTrailer(packetSize int) int {
if !device.randomTrailers.Load() {
return -1
}
if DefaultUdpWindow < packetSize {
return 0
}
return int(fastrandn(uint32(DefaultUdpWindow - packetSize)))
}
func (peer *Peer) randomTrailer(packetSize int) int {
if !peer.device.randomTrailers.Load() {
return -1
}
udpWindow := int(peer.udpWindow.Load())
if udpWindow < packetSize {
return 0
}
return int(fastrandn(uint32(udpWindow - packetSize)))
}
/* Encrypts the elements in the queue /* Encrypts the elements in the queue
* and marks them for sequential consumption (by releasing the mutex) * and marks them for sequential consumption (by releasing the mutex)
* *
@@ -544,6 +587,11 @@ func (device *Device) RoutineEncryption(id int) {
for elemsContainer := range device.queue.encryption.c { for elemsContainer := range device.queue.encryption.c {
for _, elem := range elemsContainer.elems { for _, elem := range elemsContainer.elems {
udpWindow := elem.padding + MinMessageSize + uint32(len(elem.packet))
if elem.peer.udpWindow.Load() < udpWindow {
elem.peer.udpWindow.Store(udpWindow)
}
// fill crypto padding // fill crypto padding
crypt := elem.buffer[:elem.padding] crypt := elem.buffer[:elem.padding]
rand.Read(crypt) rand.Read(crypt)
@@ -563,6 +611,9 @@ func (device *Device) RoutineEncryption(id int) {
mtu := int(device.tun.mtu.Load()) mtu := int(device.tun.mtu.Load())
paddingSize := device.randomPaddingAddition(packetSize, mtu) paddingSize := device.randomPaddingAddition(packetSize, mtu)
if paddingSize < 0 {
paddingSize = elem.peer.randomTrailer(packetSize + MinMessageSize + int(elem.padding))
}
if paddingSize < 0 { if paddingSize < 0 {
// pad content to multiple of 16 // pad content to multiple of 16
paddingSize = calculatePaddingSize(packetSize, mtu) paddingSize = calculatePaddingSize(packetSize, mtu)
+30
View File
@@ -70,6 +70,18 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
} }
buf.WriteByte('\n') buf.WriteByte('\n')
} }
boolf := func(prefix string, val bool) {
buf.Grow(3 + len(prefix))
buf.WriteString(prefix)
buf.WriteByte('=')
if val {
buf.WriteByte('1')
} else {
buf.WriteByte('0')
}
buf.WriteByte('\n')
}
func() { func() {
// lock required resources // lock required resources
@@ -173,6 +185,8 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
if rang := device.timings.maxHandshakeAttemps.Load(); !rang.IsZero() { if rang := device.timings.maxHandshakeAttemps.Load(); !rang.IsZero() {
sendf("max_handshake_attempts=%s", rang.ToString()) sendf("max_handshake_attempts=%s", rang.ToString())
} }
boolf("random_trailers", device.randomTrailers.Load())
boolf("disable_cookies", device.disableCookies.Load())
for _, peer := range device.peers.keyMap { for _, peer := range device.peers.keyMap {
// Serialize peer state. // Serialize peer state.
@@ -513,6 +527,22 @@ func (device *Device) handleDeviceLine(ipcDev *ipcSetDevice, key, value string)
device.log.Verbosef("UAPI: Updating max handshake attempts") device.log.Verbosef("UAPI: Updating max handshake attempts")
device.timings.maxHandshakeAttemps.Store(rang) device.timings.maxHandshakeAttemps.Store(rang)
case "random_trailers":
val, err := strconv.ParseBool(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse random trailers: %w", err)
}
device.log.Verbosef("UAPI: Updating random trailers")
device.randomTrailers.Store(val)
case "disable_cookies":
val, err := strconv.ParseBool(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse disable cookies: %w", err)
}
device.log.Verbosef("UAPI: Updating disable cookies")
device.disableCookies.Store(val)
default: default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key) return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
} }