mirror of
https://github.com/amnezia-vpn/amneziawg-go.git
synced 2026-10-02 21:36:06 +03:00
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:
1 parent
08d68cdae2
commit
1f50ad736e
6 files changed
+143
-42
No files matched your search
@@ -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 (
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in new issue
Block a user