feat: amneziawg 3.0

* feat: add header protection mechanism, random transport payload trailing size, and handshake timings randomization

* fix: use uint range instead of the int one

* fix: uapi typo

* fix: simultaneous padding access

* feat: prohibit Sx < 8 if headerProtection is set

* fix: do not count bytes on crypt

* fix: trailing size calculation

* fix: get rid of buffer copying on receival

* chore: rename random padding multiple

* fix: atomize junk packets

* fix: use padding from device by default

* chore: readme changes

* chore: add readme annotation about client-side params

* chore: add both side recommendation for content padding

* feat: change ContentPaddingMultiple to ContentPaddingAddition

* feat: bring chacha20 in there

* chore: use UintRange instead of magicHeader

* feat: make all timers parameters

* fix: reissue max_handshake_amount on successful handshake

* chore: add readme info about new timers

* fix: use ParseUint for UintRange

* feat: make persistent keepalive a range

* feat: use atomics everywhere where possible

* chore: readme changes

* fix: typo

* fix: minor timer fixes
This commit is contained in:
Yaroslav Gurov authored and GitHub committed 2026-07-24 14:18:33 +02:00
1 parent c1e9bb3758
commit 457d920a1a
10 files changed
+742 -313

No files matched your search

+50
View File
@@ -53,9 +53,59 @@ $ make
## Configuration
### Data types and definitions
`client-side` means the param is not required to be the same on both server and client, while
`server-side` means this is mandatory the values are the same on server and on client
`range,x`, where x is the underlying type of left and right
* Format: "a-b", or "a", or "(off)"
* Example: PersistentKeepalive = 22-30
> [!NOTE]
> If there is no value specified (for any param), AWG treats it as 0
### Header protection [AWG 3+]
Header protection is the mechanism of protecting low-entropy values of packets' headers. The idea is to apply fast encryption to the specific fields which WireGuard does use for authentication and its own encryption. The cipher uses `S1-S4` crypto padding as nonce for each incoming packet.
```
[Device]
+ HeaderProtectionKey: key,string - server-side # the key to be used in header protection
```
> [!TIP]
> Use `awg genkey` to generate header protection key
> [!IMPORTANT]
> Header protection requires `S1-S4` value to be 8 at least
### Content padding [AWG 3+]
```
[Device]
+ ContentPaddingAddition: uint32,range - client-side # the range to be used as a custom padding
```
> [!TIP]
> It's important to specify content padding on both sides. However, this is not strictly required and could be omitted.
### Timings [AWG 3+]
This param could be used to customize default Wireguard's timings
```
[Device]
+ RekeyAfterTime = "uint32,range - client-side - seconds" # time, after which client tries to handshake
+ RekeyTimeout = "uint32,range - client-side - seconds" # timeout, after which handshake is repeated
+ RejectAfterTime = "uint32,range - client-side - seconds" # time, after which client forces handshake, and declines all incoming data
+ KeepaliveTimeout = "uint32,range - client-side - seconds" # time from last data sending, after which keepalive is sent
+ MaxHandshakeAttempts = "uint32,range - client-side - amount" # maximum attempts of handshake repetition
[Peer]
M PersistentKeepalive = "uint32,range - client-side - seconds" # interval of persistent keepalive
```
### Junk packets
The amount of junk packets specified in `Jc` with a random size between `Jmin` and `Jmax` would be generated and sent prior every handshake
+40 -17
View File
@@ -91,26 +91,41 @@ type Device struct {
log *Logger
junk struct {
min int
max int
count int
min atomic.Uint32
max atomic.Uint32
count atomic.Uint32
}
headers struct {
init *magicHeader
cookie *magicHeader
response *magicHeader
transport *magicHeader
init AtomicUintRange
cookie AtomicUintRange
response AtomicUintRange
transport AtomicUintRange
}
paddings struct {
init int
response int
cookie int
transport int
init atomic.Uint32
response atomic.Uint32
cookie atomic.Uint32
transport atomic.Uint32
}
ipackets [5]*obfChain
headerProtection struct {
sync.RWMutex
key HeaderCipherKey
}
contentPaddingAddition AtomicUintRange
timings struct {
rekeyAfterTimeSec AtomicUintRange
rekeyTimeoutSec AtomicUintRange
rejectAfterTimeSec AtomicUintRange
keepaliveTimeoutSec AtomicUintRange
maxHandshakeAttemps AtomicUintRange
}
}
// deviceState represents the state of a Device.
@@ -205,7 +220,7 @@ func (device *Device) upLocked() error {
device.peers.RLock()
for _, peer := range device.peers.keyMap {
peer.Start()
if peer.persistentKeepaliveInterval.Load() > 0 {
if !peer.persistentKeepaliveInterval.Load().IsZero() {
peer.SendKeepalive()
}
}
@@ -305,6 +320,8 @@ func (device *Device) SetPrivateKey(sk NoisePrivateKey) error {
}
func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger) *Device {
var rang UintRange
device := new(Device)
device.state.state.Store(uint32(deviceStateDown))
device.closed = make(chan struct{})
@@ -321,10 +338,14 @@ func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger) *Device {
device.rate.limiter.Init()
device.indexTable.Init()
device.headers.init = &magicHeader{start: MessageInitiationType, end: MessageInitiationType}
device.headers.response = &magicHeader{start: MessageResponseType, end: MessageResponseType}
device.headers.cookie = &magicHeader{start: MessageCookieReplyType, end: MessageCookieReplyType}
device.headers.transport = &magicHeader{start: MessageTransportType, end: MessageTransportType}
rang.FromUint32(MessageInitiationType, MessageInitiationType)
device.headers.init.Store(rang)
rang.FromUint32(MessageResponseType, MessageResponseType)
device.headers.response.Store(rang)
rang.FromUint32(MessageCookieReplyType, MessageCookieReplyType)
device.headers.cookie.Store(rang)
rang.FromUint32(MessageTransportType, MessageTransportType)
device.headers.transport.Store(rang)
device.PopulatePools()
@@ -436,10 +457,12 @@ func (device *Device) SendKeepalivesToPeersWithCurrentKeypair() {
return
}
timeout := device.keychainExpireTime()
device.peers.RLock()
for _, peer := range device.peers.keyMap {
peer.keypairs.RLock()
sendKeepalive := peer.keypairs.current != nil && !peer.keypairs.current.created.Add(RejectAfterTime).Before(time.Now())
sendKeepalive := peer.keypairs.current != nil && !peer.keypairs.current.created.Add(timeout).Before(time.Now())
peer.keypairs.RUnlock()
if sendKeepalive {
peer.SendKeepalive()
-63
View File
@@ -1,63 +0,0 @@
package device
import (
"crypto/rand"
"errors"
"fmt"
"math/big"
"strconv"
"strings"
)
type magicHeader struct {
start uint32
end uint32
}
func newMagicHeader(spec string) (*magicHeader, error) {
parts := strings.Split(spec, "-")
if len(parts) < 1 || len(parts) > 2 {
return nil, errors.New("bad format")
}
start, err := strconv.ParseUint(parts[0], 10, 32)
if err != nil {
return nil, fmt.Errorf("failed to parse %s: %w", parts[0], err)
}
var end uint64
if len(parts) > 1 {
end, err = strconv.ParseUint(parts[1], 10, 32)
if err != nil {
return nil, fmt.Errorf("failed to parse %s: %w", parts[1], err)
}
} else {
end = start
}
if end < start {
return nil, errors.New("wrong range specified")
}
return &magicHeader{
start: uint32(start),
end: uint32(end),
}, nil
}
func (h *magicHeader) GenSpec() string {
if h.start == h.end {
return fmt.Sprintf("%d", h.start)
}
return fmt.Sprintf("%d-%d", h.start, h.end)
}
func (h *magicHeader) Validate(val uint32) bool {
return h.start <= val && val <= h.end
}
func (h *magicHeader) Generate() uint32 {
high := int64(h.end - h.start + 1)
r, _ := rand.Int(rand.Reader, big.NewInt(high))
return h.start + uint32(r.Int64())
}
+30 -2
View File
@@ -6,12 +6,14 @@
package device
import (
"crypto/rand"
"errors"
"fmt"
"sync"
"time"
"golang.org/x/crypto/blake2s"
"golang.org/x/crypto/chacha20"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/crypto/poly1305"
@@ -194,7 +196,7 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e
handshake.mixHash(handshake.remoteStatic[:])
msgType := device.headers.init.Generate()
msgType := device.headers.init.Load().PickOne()
msg := MessageInitiation{
Type: msgType,
@@ -370,7 +372,7 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
var msg MessageResponse
msg.Type = device.headers.response.Generate()
msg.Type = device.headers.response.Load().PickOne()
msg.Sender = handshake.localIndex
msg.Receiver = handshake.remoteIndex
@@ -626,3 +628,29 @@ func (peer *Peer) ReceivedWithKeypair(receivedKeypair *Keypair) bool {
keypairs.next.Store(nil)
return true
}
func (device *Device) JunkPackets() [][]byte {
var bufs [][]byte
min := device.junk.min.Load()
max := device.junk.max.Load()
for range device.junk.count.Load() {
buf := make([]byte, min+fastrandn(max-min))
rand.Read(buf)
bufs = append(bufs, buf)
}
return bufs
}
func (device *Device) HeaderProtectionCipher(salt []byte) (*chacha20.Cipher, error) {
device.headerProtection.RLock()
defer device.headerProtection.RUnlock()
if device.headerProtection.key.IsZero() {
return nil, nil
}
return chacha20.NewUnauthenticatedCipher(device.headerProtection.key[:], salt)
}
+108
View File
@@ -9,12 +9,18 @@ import (
"crypto/subtle"
"encoding/hex"
"errors"
"fmt"
"strconv"
"strings"
"sync/atomic"
)
const (
NoisePublicKeySize = 32
NoisePrivateKeySize = 32
NoisePresharedKeySize = 32
HeaderCipherKeySize = 32
HeaderCipherNonceSize = 12
)
type (
@@ -22,6 +28,7 @@ type (
NoisePrivateKey [NoisePrivateKeySize]byte
NoisePresharedKey [NoisePresharedKeySize]byte
NoiseNonce uint64 // padded to 12-bytes
HeaderCipherKey [HeaderCipherKeySize]byte
)
func loadExactHex(dst []byte, src string) error {
@@ -76,3 +83,104 @@ func (key NoisePublicKey) Equals(tar NoisePublicKey) bool {
func (key *NoisePresharedKey) FromHex(src string) error {
return loadExactHex(key[:], src)
}
func (key HeaderCipherKey) IsZero() bool {
var zero HeaderCipherKey
return key.Equals(zero)
}
func (key HeaderCipherKey) Equals(tar HeaderCipherKey) bool {
return subtle.ConstantTimeCompare(key[:], tar[:]) == 1
}
func (key *HeaderCipherKey) FromHex(src string) error {
return loadExactHex(key[:], src)
}
type UintRange uint64
func (r *UintRange) FromUint32(lo, hi uint32) {
*r = UintRange(uint64(hi)<<32 | uint64(lo))
}
func (r *UintRange) FromString(str string) error {
parts := strings.Split(str, "-")
if len(parts) < 1 || len(parts) > 2 {
return errors.New("wrong format")
}
lo, err := strconv.ParseUint(parts[0], 10, 32)
if err != nil {
return err
}
hi := lo
if len(parts) > 1 {
hi, err = strconv.ParseUint(parts[1], 10, 32)
if err != nil {
return err
}
}
if hi < lo {
return errors.New("wrong range specified")
}
r.FromUint32(uint32(lo), uint32(hi))
return nil
}
func (r UintRange) Contains(num uint32) bool {
lo, hi := uint32(r), uint32(r>>32)
return lo <= num && num <= hi
}
func (r UintRange) IsZero() bool {
return r == 0
}
func (r UintRange) PickOne() uint32 {
lo, hi := uint32(r), uint32(r>>32)
return lo + fastrandn(hi-lo+1)
}
func (r UintRange) ToString() string {
lo, hi := uint32(r), uint32(r>>32)
if lo == hi {
return fmt.Sprintf("%d", lo)
} else {
return fmt.Sprintf("%d-%d", lo, hi)
}
}
func (r UintRange) Overlap(right UintRange) bool {
l_lo, l_hi := uint32(r), uint32(r>>32)
r_lo, r_hi := uint32(right), uint32(right>>32)
return l_lo <= r_hi && r_lo <= l_hi
}
func (r UintRange) Lo() uint32 {
return uint32(r)
}
func (r UintRange) Hi() uint32 {
return uint32(r >> 32)
}
type AtomicUintRange struct {
v atomic.Uint64
}
func (a *AtomicUintRange) Load() UintRange {
return UintRange(a.v.Load())
}
func (a *AtomicUintRange) Store(r UintRange) {
a.v.Store(uint64(r))
}
func (a *AtomicUintRange) Swap(r UintRange) UintRange {
return UintRange(a.v.Swap(uint64(r)))
}
+4 -3
View File
@@ -39,6 +39,7 @@ type Peer struct {
zeroKeyMaterial *Timer
persistentKeepalive *Timer
handshakeAttempts atomic.Uint32
maxHandshakeAttempts atomic.Uint32
needAnotherKeepalive atomic.Bool
sentLastMinuteHandshake atomic.Bool
}
@@ -55,7 +56,7 @@ type Peer struct {
cookieGenerator CookieGenerator
trieEntries list.List
persistentKeepaliveInterval atomic.Uint32
persistentKeepaliveInterval AtomicUintRange
}
func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
@@ -192,7 +193,7 @@ func (peer *Peer) Start() {
peer.stopping.Add(2)
peer.handshake.mutex.Lock()
peer.handshake.lastSentHandshake = time.Now().Add(-(RekeyTimeout + time.Second))
peer.handshake.lastSentHandshake = time.Now().Add(-(peer.device.rekeyMinTimeout() + time.Second))
peer.handshake.mutex.Unlock()
peer.device.queue.encryption.wg.Add(1) // keep encryption queue open for our writes
@@ -242,7 +243,7 @@ func (peer *Peer) ExpireCurrentKeypairs() {
handshake.mutex.Lock()
peer.device.indexTable.Delete(handshake.localIndex)
handshake.Clear()
peer.handshake.lastSentHandshake = time.Now().Add(-(RekeyTimeout + time.Second))
peer.handshake.lastSentHandshake = time.Now().Add(-(peer.device.rekeyMinTimeout() + time.Second))
handshake.mutex.Unlock()
keypairs := &peer.keypairs
+64 -31
View File
@@ -32,6 +32,7 @@ type QueueInboundElement struct {
counter uint64
keypair *Keypair
endpoint conn.Endpoint
padding uint32
}
type QueueInboundElementsContainer struct {
@@ -59,7 +60,8 @@ func (peer *Peer) keepKeyFreshReceiving() {
return
}
keypair := peer.keypairs.Current()
if keypair != nil && keypair.isInitiator && time.Since(keypair.created) > (RejectAfterTime-KeepaliveTimeout-RekeyTimeout) {
if keypair != nil && keypair.isInitiator && time.Since(keypair.created) > peer.device.keyRefreshTimeoutReceiving() {
peer.timers.sentLastMinuteHandshake.Store(true)
peer.SendHandshakeInitiation(false)
}
@@ -95,6 +97,7 @@ func (device *Device) RoutineReceiveIncoming(
endpoints = make([]conn.Endpoint, maxBatchSize)
deathSpiral int
elemsByPeer = make(map[*Peer]*QueueInboundElementsContainer, maxBatchSize)
typeHashBuf [4]byte
)
for i := range maxBatchSize {
@@ -138,11 +141,24 @@ func (device *Device) RoutineReceiveIncoming(
// check size of packet
packet := bufsArrs[i][:size]
cip, err := device.HeaderProtectionCipher(packet[:HeaderCipherNonceSize])
if err != nil {
device.log.Errorf("Failed to initialize header cipher")
continue
}
typeHash := typeHashBuf[:]
clear(typeHash)
if cip != nil {
cip.XORKeyStream(typeHash, typeHash)
}
// get message padding and type based on information from S1-S4 and H1-H4
msgType, padding := device.DeterminePacketTypeAndPadding(packet, MessageUnknownType)
if padding > 0 {
copy(packet, packet[padding:])
packet = packet[:len(packet)-padding]
msgType, padding := device.DeterminePacketTypeAndPadding(packet, MessageUnknownType, typeHash)
packet = packet[padding:]
if cip != nil {
applyHash(packet[:4], packet[:4], typeHash)
}
switch msgType {
@@ -156,6 +172,9 @@ func (device *Device) RoutineReceiveIncoming(
if len(packet) < MessageTransportSize {
continue
}
if cip != nil {
cip.XORKeyStream(packet[4:MessageTransportHeaderSize], packet[4:MessageTransportHeaderSize])
}
// lookup key pair
@@ -170,7 +189,7 @@ func (device *Device) RoutineReceiveIncoming(
// check keypair expiry
if keypair.created.Add(RejectAfterTime).Before(time.Now()) {
if keypair.created.Add(device.keychainExpireTime()).Before(time.Now()) {
continue
}
@@ -182,6 +201,7 @@ func (device *Device) RoutineReceiveIncoming(
elem.keypair = keypair
elem.endpoint = endpoints[i]
elem.counter = 0
elem.padding = padding
elemsForPeer, ok := elemsByPeer[peer]
if !ok {
@@ -200,16 +220,25 @@ func (device *Device) RoutineReceiveIncoming(
if len(packet) != MessageInitiationSize {
continue
}
if cip != nil {
cip.XORKeyStream(packet[4:MessageInitiationSize], packet[4:MessageInitiationSize])
}
case MessageResponseType:
if len(packet) != MessageResponseSize {
continue
}
if cip != nil {
cip.XORKeyStream(packet[4:MessageResponseSize], packet[4:MessageResponseSize])
}
case MessageCookieReplyType:
if len(packet) != MessageCookieReplySize {
continue
}
if cip != nil {
cip.XORKeyStream(packet[4:MessageCookieReplySize], packet[4:MessageCookieReplySize])
}
default:
device.log.Verbosef("Received message with unknown type")
@@ -529,10 +558,7 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
continue
}
bufs = append(
bufs,
elem.buffer[:MessageTransportOffsetContent+len(elem.packet)],
)
bufs = append(bufs, elem.buffer[int(elem.padding):int(elem.padding)+MessageTransportOffsetContent+len(elem.packet)])
}
peer.rxBytes.Add(rxBytesLen)
@@ -560,52 +586,59 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
}
}
func (device *Device) DeterminePacketTypeAndPadding(packet []byte, expectedType uint32) (uint32, int) {
func applyHash(dst, src, hash []byte) {
for i := range len(dst) {
dst[i] = src[i] ^ hash[i]
}
}
func (device *Device) DeterminePacketTypeAndPadding(packet []byte, expectedType uint32, typeHash []byte) (uint32, uint32) {
var headerBytes [4]byte
size := len(packet)
if expectedType == MessageUnknownType || expectedType == MessageInitiationType {
padding := device.paddings.init
header := device.headers.init
padding := device.paddings.init.Load()
header := device.headers.init.Load()
if size == padding+MessageInitiationSize {
data := packet[padding:]
if header.Validate(binary.LittleEndian.Uint32(data)) {
if size == int(padding)+MessageInitiationSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageInitiationType, padding
}
}
}
if expectedType == MessageUnknownType || expectedType == MessageResponseType {
padding := device.paddings.response
header := device.headers.response
padding := device.paddings.response.Load()
header := device.headers.response.Load()
if size == padding+MessageResponseSize {
data := packet[padding:]
if header.Validate(binary.LittleEndian.Uint32(data)) {
if size == int(padding)+MessageResponseSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageResponseType, padding
}
}
}
if expectedType == MessageUnknownType || expectedType == MessageCookieReplyType {
padding := device.paddings.cookie
header := device.headers.cookie
padding := device.paddings.cookie.Load()
header := device.headers.cookie.Load()
if size == padding+MessageCookieReplySize {
data := packet[padding:]
if header.Validate(binary.LittleEndian.Uint32(data)) {
if size == int(padding)+MessageCookieReplySize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageCookieReplyType, padding
}
}
}
if expectedType == MessageUnknownType || expectedType == MessageTransportType {
padding := device.paddings.transport
header := device.headers.transport
padding := device.paddings.transport.Load()
header := device.headers.transport.Load()
if size >= padding+MessageTransportHeaderSize {
data := packet[padding:]
if header.Validate(binary.LittleEndian.Uint32(data)) {
if size >= int(padding)+MessageTransportHeaderSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageTransportType, padding
}
}
+108 -60
View File
@@ -10,9 +10,9 @@ import (
"crypto/rand"
"encoding/binary"
"errors"
"math/big"
"net"
"os"
"slices"
"sync"
"time"
@@ -53,6 +53,7 @@ type QueueOutboundElement struct {
nonce uint64 // nonce for encryption
keypair *Keypair // keypair for encryption
peer *Peer // related peer
padding uint32
}
type QueueOutboundElementsContainer struct {
@@ -64,6 +65,7 @@ func (device *Device) NewOutboundElement() *QueueOutboundElement {
elem := device.GetOutboundElement()
elem.buffer = device.GetMessageBuffer()
elem.nonce = 0
elem.padding = device.paddings.transport.Load()
// keypair and peer were cleared (if necessary) by clearPointers.
return elem
}
@@ -101,17 +103,20 @@ func (peer *Peer) SendKeepalive() {
func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
if !isRetry {
peer.timers.handshakeAttempts.Store(0)
peer.timers.maxHandshakeAttempts.Store(peer.device.maxHandshakeAttemps())
}
timeout := peer.device.rekeyMinTimeout()
peer.handshake.mutex.RLock()
if time.Since(peer.handshake.lastSentHandshake) < RekeyTimeout {
if time.Since(peer.handshake.lastSentHandshake) < timeout {
peer.handshake.mutex.RUnlock()
return nil
}
peer.handshake.mutex.RUnlock()
peer.handshake.mutex.Lock()
if time.Since(peer.handshake.lastSentHandshake) < RekeyTimeout {
if time.Since(peer.handshake.lastSentHandshake) < timeout {
peer.handshake.mutex.Unlock()
return nil
}
@@ -136,21 +141,15 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
}
}
jc := peer.device.junk.count
jmin := peer.device.junk.min
jmax := peer.device.junk.max
sendBuffer = append(sendBuffer, peer.device.JunkPackets()...)
for i := 0; i < jc; i++ {
nBig, _ := rand.Int(rand.Reader, big.NewInt(int64(jmax-jmin+1)))
n := int(nBig.Int64()) + jmin
padding := peer.device.paddings.init.Load()
buf := make([]byte, padding+MessageInitiationSize)
buf := make([]byte, n)
rand.Read(buf)
sendBuffer = append(sendBuffer, buf)
}
crypt := buf[:padding]
rand.Read(crypt)
var buf [MessageInitiationSize]byte
writer := bytes.NewBuffer(buf[:0])
writer := bytes.NewBuffer(buf[padding:padding])
binary.Write(writer, binary.LittleEndian, msg)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
@@ -158,15 +157,15 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
if padding := peer.device.paddings.init; padding > 0 {
buf := make([]byte, padding+len(packet))
rand.Read(buf[:padding])
copy(buf[padding:], packet)
packet = buf
cip, err := peer.device.HeaderProtectionCipher(crypt[:HeaderCipherNonceSize])
if err != nil {
return err
}
if cip != nil {
cip.XORKeyStream(packet, packet)
}
sendBuffer = append(sendBuffer, packet)
sendBuffer = append(sendBuffer, buf)
err = peer.SendBuffers(sendBuffer)
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err)
@@ -189,9 +188,13 @@ func (peer *Peer) SendHandshakeResponse() error {
return err
}
var buf [MessageResponseSize]byte
writer := bytes.NewBuffer(buf[:0])
padding := peer.device.paddings.response.Load()
buf := make([]byte, padding+MessageResponseSize)
crypt := buf[:padding]
rand.Read(crypt)
writer := bytes.NewBuffer(buf[padding:padding])
binary.Write(writer, binary.LittleEndian, response)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
@@ -206,15 +209,16 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
if padding := peer.device.paddings.response; padding > 0 {
buf := make([]byte, padding+len(packet))
rand.Read(buf[:padding])
copy(buf[padding:], packet)
packet = buf
cip, err := peer.device.HeaderProtectionCipher(crypt[:HeaderCipherNonceSize])
if err != nil {
return err
}
if cip != nil {
cip.XORKeyStream(packet, packet)
}
// TODO: allocation could be avoided
err = peer.SendBuffers([][]byte{packet})
err = peer.SendBuffers([][]byte{buf})
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err)
}
@@ -225,7 +229,7 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement)
device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString())
sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8])
msgType := device.headers.cookie.Generate()
msgType := device.headers.cookie.Load().PickOne()
reply, err := device.cookieChecker.CreateReply(
initiatingElem.packet,
@@ -238,20 +242,26 @@ func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement)
return err
}
var buf [MessageCookieReplySize]byte
writer := bytes.NewBuffer(buf[:0])
padding := device.paddings.cookie.Load()
buf := make([]byte, padding+MessageCookieReplySize)
crypt := buf[:padding]
rand.Read(crypt)
writer := bytes.NewBuffer(buf[padding:padding])
binary.Write(writer, binary.LittleEndian, reply)
packet := writer.Bytes()
if padding := device.paddings.cookie; padding > 0 {
buf := make([]byte, padding+len(packet))
rand.Read(buf[:padding])
copy(buf[padding:], packet)
packet = buf
cip, err := device.HeaderProtectionCipher(crypt[:HeaderCipherNonceSize])
if err != nil {
return err
}
if cip != nil {
cip.XORKeyStream(packet, packet)
}
// TODO: allocation could be avoided
device.net.bind.Send([][]byte{packet}, initiatingElem.endpoint)
device.net.bind.Send([][]byte{buf}, initiatingElem.endpoint)
return nil
}
@@ -261,7 +271,7 @@ func (peer *Peer) keepKeyFreshSending() {
return
}
nonce := keypair.sendNonce.Load()
if nonce > RekeyAfterMessages || (keypair.isInitiator && time.Since(keypair.created) > RekeyAfterTime) {
if nonce > RekeyAfterMessages || (keypair.isInitiator && time.Since(keypair.created) > peer.device.keyRefreshTimeoutSending()) {
peer.SendHandshakeInitiation(false)
}
}
@@ -283,7 +293,6 @@ func (device *Device) RoutineReadFromTUN() {
elemsByPeer = make(map[*Peer]*QueueOutboundElementsContainer, batchSize)
count = 0
sizes = make([]int, batchSize)
offset = MessageTransportHeaderSize
)
for i := range elems {
@@ -301,6 +310,9 @@ func (device *Device) RoutineReadFromTUN() {
}()
for {
padding := device.paddings.transport.Load()
offset := MessageTransportHeaderSize + int(padding)
// read packets
count, readErr = device.tun.device.Read(bufs, sizes, offset)
for i := 0; i < count; i++ {
@@ -310,6 +322,7 @@ func (device *Device) RoutineReadFromTUN() {
elem := elems[i]
elem.packet = bufs[i][offset : offset+sizes[i]]
elem.padding = padding
// lookup peer
var peer *Peer
@@ -404,7 +417,7 @@ top:
}
keypair := peer.keypairs.Current()
if keypair == nil || keypair.sendNonce.Load() >= RejectAfterMessages || time.Since(keypair.created) >= RejectAfterTime {
if keypair == nil || keypair.sendNonce.Load() >= RejectAfterMessages || time.Since(keypair.created) >= peer.device.keychainExpireTime() {
peer.SendHandshakeInitiation(false)
return
}
@@ -494,13 +507,33 @@ func calculatePaddingSize(packetSize, mtu int) int {
return paddedSize - lastUnit
}
func (device *Device) randomPaddingAddition(packetSize, mtu int) int {
addition := device.contentPaddingAddition.Load()
if addition.IsZero() {
return -1
}
add := int(addition.PickOne())
if mtu != 0 {
if packetSize > mtu {
packetSize %= mtu
}
space := mtu - packetSize
if add > space {
add = space
}
}
return add
}
/* Encrypts the elements in the queue
* and marks them for sequential consumption (by releasing the mutex)
*
* Obs. One instance per core
*/
func (device *Device) RoutineEncryption(id int) {
var paddingZeros [PaddingMultiple]byte
var nonce [chacha20poly1305.NonceSize]byte
defer device.log.Verbosef("Routine: encryption worker %d - stopped", id)
@@ -508,32 +541,55 @@ func (device *Device) RoutineEncryption(id int) {
for elemsContainer := range device.queue.encryption.c {
for _, elem := range elemsContainer.elems {
// fill crypto padding
crypt := elem.buffer[:elem.padding]
rand.Read(crypt)
// populate header fields
header := elem.buffer[:MessageTransportHeaderSize]
header := elem.buffer[elem.padding : elem.padding+MessageTransportHeaderSize]
fieldType := header[0:4]
fieldReceiver := header[4:8]
fieldNonce := header[8:16]
msgType := device.headers.transport.Generate()
binary.LittleEndian.PutUint32(fieldType, msgType)
binary.LittleEndian.PutUint32(fieldType, device.headers.transport.Load().PickOne())
binary.LittleEndian.PutUint32(fieldReceiver, elem.keypair.remoteIndex)
binary.LittleEndian.PutUint64(fieldNonce, elem.nonce)
// pad content to multiple of 16
paddingSize := calculatePaddingSize(len(elem.packet), int(device.tun.mtu.Load()))
elem.packet = append(elem.packet, paddingZeros[:paddingSize]...)
packetSize := len(elem.packet)
mtu := int(device.tun.mtu.Load())
paddingSize := device.randomPaddingAddition(packetSize, mtu)
if paddingSize < 0 {
// pad content to multiple of 16
paddingSize = calculatePaddingSize(packetSize, mtu)
}
// append trailing zeroes
oldLen := len(elem.packet)
elem.packet = slices.Grow(elem.packet, paddingSize)
elem.packet = elem.packet[:oldLen+paddingSize]
clear(elem.packet[oldLen:])
// encrypt content and release to consumer
binary.LittleEndian.PutUint64(nonce[4:], elem.nonce)
elem.packet = elem.keypair.send.Seal(
header,
elem.buffer[:elem.padding+MessageTransportHeaderSize],
nonce[:],
elem.packet,
nil,
)
cip, err := device.HeaderProtectionCipher(crypt[:HeaderCipherNonceSize])
if err != nil {
device.log.Errorf("Routing: header obfuscation failed - packet dropped")
elem.packet = nil
continue
}
if cip != nil {
cip.XORKeyStream(header, header)
}
}
elemsContainer.Unlock()
}
@@ -575,15 +631,7 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
if len(elem.packet) != MessageKeepaliveSize {
dataSent = true
}
if padding := device.paddings.transport; padding > 0 {
// elem.packet is stored at the start of elem.buffer
// with zero padding
for i := len(elem.packet) - 1; i >= 0; i-- {
elem.buffer[i+padding] = elem.buffer[i]
}
rand.Read(elem.buffer[:padding])
elem.packet = elem.buffer[:padding+len(elem.packet)]
}
bufs = append(bufs, elem.packet)
}
+124 -27
View File
@@ -22,24 +22,25 @@ type Timer struct {
*time.Timer
modifyingLock sync.RWMutex
runningLock sync.Mutex
isPending bool
duration time.Duration
}
func (peer *Peer) NewTimer(expirationFunction func(*Peer)) *Timer {
func (peer *Peer) NewTimer(expirationFunction func(*Peer, time.Duration)) *Timer {
timer := &Timer{}
timer.Timer = time.AfterFunc(time.Hour, func() {
timer.runningLock.Lock()
defer timer.runningLock.Unlock()
timer.modifyingLock.Lock()
if !timer.isPending {
if timer.duration == 0 {
timer.modifyingLock.Unlock()
return
}
timer.isPending = false
duration := timer.duration
timer.modifyingLock.Unlock()
timer.duration = 0
expirationFunction(peer)
expirationFunction(peer, duration)
})
timer.Stop()
return timer
@@ -47,14 +48,14 @@ func (peer *Peer) NewTimer(expirationFunction func(*Peer)) *Timer {
func (timer *Timer) Mod(d time.Duration) {
timer.modifyingLock.Lock()
timer.isPending = true
timer.duration = d
timer.Reset(d)
timer.modifyingLock.Unlock()
}
func (timer *Timer) Del() {
timer.modifyingLock.Lock()
timer.isPending = false
timer.duration = 0
timer.Stop()
timer.modifyingLock.Unlock()
}
@@ -69,16 +70,18 @@ func (timer *Timer) DelSync() {
func (timer *Timer) IsPending() bool {
timer.modifyingLock.RLock()
defer timer.modifyingLock.RUnlock()
return timer.isPending
return timer.duration > 0
}
func (peer *Peer) timersActive() bool {
return peer.isRunning.Load() && peer.device != nil && peer.device.isUp()
}
func expiredRetransmitHandshake(peer *Peer) {
if peer.timers.handshakeAttempts.Load() > MaxTimerHandshakes {
peer.device.log.Verbosef("%s - Handshake did not complete after %d attempts, giving up", peer, MaxTimerHandshakes+2)
func expiredRetransmitHandshake(peer *Peer, d time.Duration) {
maxAttempts := peer.timers.maxHandshakeAttempts.Load()
if peer.timers.handshakeAttempts.Load() > maxAttempts {
peer.device.log.Verbosef("%s - Handshake did not complete after %d attempts, giving up", peer, maxAttempts+2)
if peer.timersActive() {
peer.timers.sendKeepalive.Del()
@@ -93,11 +96,11 @@ func expiredRetransmitHandshake(peer *Peer) {
* of a partial exchange.
*/
if peer.timersActive() && !peer.timers.zeroKeyMaterial.IsPending() {
peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3)
peer.timers.zeroKeyMaterial.Mod(peer.device.keychainExpireTime() * 3)
}
} else {
peer.timers.handshakeAttempts.Add(1)
peer.device.log.Verbosef("%s - Handshake did not complete after %d seconds, retrying (try %d)", peer, int(RekeyTimeout.Seconds()), peer.timers.handshakeAttempts.Load()+1)
peer.device.log.Verbosef("%s - Handshake did not complete after %d seconds, retrying (try %d)", peer, int(d.Seconds()), peer.timers.handshakeAttempts.Load()+1)
/* We clear the endpoint address src address, in case this is the cause of trouble. */
peer.markEndpointSrcForClearing()
@@ -106,30 +109,30 @@ func expiredRetransmitHandshake(peer *Peer) {
}
}
func expiredSendKeepalive(peer *Peer) {
func expiredSendKeepalive(peer *Peer, d time.Duration) {
peer.SendKeepalive()
if peer.timers.needAnotherKeepalive.Load() {
peer.timers.needAnotherKeepalive.Store(false)
if peer.timersActive() {
peer.timers.sendKeepalive.Mod(KeepaliveTimeout)
peer.timers.sendKeepalive.Mod(peer.sendKeepaliveTimeout())
}
}
}
func expiredNewHandshake(peer *Peer) {
peer.device.log.Verbosef("%s - Retrying handshake because we stopped hearing back after %d seconds", peer, int((KeepaliveTimeout + RekeyTimeout).Seconds()))
func expiredNewHandshake(peer *Peer, d time.Duration) {
peer.device.log.Verbosef("%s - Retrying handshake because we stopped hearing back after %d seconds", peer, int(d.Seconds()))
/* We clear the endpoint address src address, in case this is the cause of trouble. */
peer.markEndpointSrcForClearing()
peer.SendHandshakeInitiation(false)
}
func expiredZeroKeyMaterial(peer *Peer) {
peer.device.log.Verbosef("%s - Removing all keys, since we haven't received a new one in %d seconds", peer, int((RejectAfterTime * 3).Seconds()))
func expiredZeroKeyMaterial(peer *Peer, d time.Duration) {
peer.device.log.Verbosef("%s - Removing all keys, since we haven't received a new one in %d seconds", peer, int(d.Seconds()))
peer.ZeroAndFlushAll()
}
func expiredPersistentKeepalive(peer *Peer) {
if peer.persistentKeepaliveInterval.Load() > 0 {
func expiredPersistentKeepalive(peer *Peer, d time.Duration) {
if !peer.persistentKeepaliveInterval.Load().IsZero() {
peer.SendKeepalive()
}
}
@@ -137,7 +140,7 @@ func expiredPersistentKeepalive(peer *Peer) {
/* Should be called after an authenticated data packet is sent. */
func (peer *Peer) timersDataSent() {
if peer.timersActive() && !peer.timers.newHandshake.IsPending() {
peer.timers.newHandshake.Mod(KeepaliveTimeout + RekeyTimeout + time.Millisecond*time.Duration(fastrandn(RekeyTimeoutJitterMaxMs)))
peer.timers.newHandshake.Mod(peer.newHandshakeTimeout() + time.Millisecond*time.Duration(fastrandn(RekeyTimeoutJitterMaxMs)))
}
}
@@ -145,7 +148,7 @@ func (peer *Peer) timersDataSent() {
func (peer *Peer) timersDataReceived() {
if peer.timersActive() {
if !peer.timers.sendKeepalive.IsPending() {
peer.timers.sendKeepalive.Mod(KeepaliveTimeout)
peer.timers.sendKeepalive.Mod(peer.sendKeepaliveTimeout())
} else {
peer.timers.needAnotherKeepalive.Store(true)
}
@@ -169,7 +172,7 @@ func (peer *Peer) timersAnyAuthenticatedPacketReceived() {
/* Should be called after a handshake initiation message is sent. */
func (peer *Peer) timersHandshakeInitiated() {
if peer.timersActive() {
peer.timers.retransmitHandshake.Mod(RekeyTimeout + time.Millisecond*time.Duration(fastrandn(RekeyTimeoutJitterMaxMs)))
peer.timers.retransmitHandshake.Mod(peer.retransmitHandshakeTimeout() + time.Millisecond*time.Duration(fastrandn(RekeyTimeoutJitterMaxMs)))
}
}
@@ -179,6 +182,7 @@ func (peer *Peer) timersHandshakeComplete() {
peer.timers.retransmitHandshake.Del()
}
peer.timers.handshakeAttempts.Store(0)
peer.timers.maxHandshakeAttempts.Store(peer.device.maxHandshakeAttemps())
peer.timers.sentLastMinuteHandshake.Store(false)
peer.lastHandshakeNano.Store(time.Now().UnixNano())
}
@@ -186,15 +190,15 @@ func (peer *Peer) timersHandshakeComplete() {
/* Should be called after an ephemeral key is created, which is before sending a handshake response or after receiving a handshake response. */
func (peer *Peer) timersSessionDerived() {
if peer.timersActive() {
peer.timers.zeroKeyMaterial.Mod(RejectAfterTime * 3)
peer.timers.zeroKeyMaterial.Mod(peer.device.keychainExpireTime() * 3)
}
}
/* Should be called before a packet with authentication -- keepalive, data, or handshake -- is sent, or after one is received. */
func (peer *Peer) timersAnyAuthenticatedPacketTraversal() {
keepalive := peer.persistentKeepaliveInterval.Load()
if keepalive > 0 && peer.timersActive() {
peer.timers.persistentKeepalive.Mod(time.Duration(keepalive) * time.Second)
if !keepalive.IsZero() && peer.timersActive() {
peer.timers.persistentKeepalive.Mod(time.Duration(keepalive.PickOne()) * time.Second)
}
}
@@ -208,6 +212,7 @@ func (peer *Peer) timersInit() {
func (peer *Peer) timersStart() {
peer.timers.handshakeAttempts.Store(0)
peer.timers.maxHandshakeAttempts.Store(peer.device.maxHandshakeAttemps())
peer.timers.sentLastMinuteHandshake.Store(false)
peer.timers.needAnotherKeepalive.Store(false)
}
@@ -219,3 +224,95 @@ func (peer *Peer) timersStop() {
peer.timers.zeroKeyMaterial.DelSync()
peer.timers.persistentKeepalive.DelSync()
}
func (peer *Peer) retransmitHandshakeTimeout() time.Duration {
timeout := RekeyTimeout
if t := peer.device.timings.rekeyTimeoutSec.Load(); !t.IsZero() {
timeout = time.Duration(t.PickOne()) * time.Second
}
return timeout
}
func (peer *Peer) sendKeepaliveTimeout() time.Duration {
timeout := KeepaliveTimeout
if t := peer.device.timings.keepaliveTimeoutSec.Load(); !t.IsZero() {
timeout = time.Duration(t.PickOne()) * time.Second
}
return timeout
}
func (peer *Peer) newHandshakeTimeout() time.Duration {
keepaliveTimeout := KeepaliveTimeout
rekeyTimeout := RekeyTimeout
if t := peer.device.timings.keepaliveTimeoutSec.Load(); !t.IsZero() {
keepaliveTimeout = time.Duration(t.Hi()) * time.Second
}
if t := peer.device.timings.rekeyTimeoutSec.Load(); !t.IsZero() {
rekeyTimeout = time.Duration(t.PickOne()) * time.Second
}
return keepaliveTimeout + rekeyTimeout
}
func (device *Device) keyRefreshTimeoutSending() time.Duration {
rekeyAfterTime := RekeyAfterTime
if t := device.timings.rekeyAfterTimeSec.Load(); !t.IsZero() {
rekeyAfterTime = time.Duration(t.PickOne()) * time.Second
}
return rekeyAfterTime
}
func (device *Device) keyRefreshTimeoutReceiving() time.Duration {
rejectAfterTime := RejectAfterTime
keepaliveTimeout := KeepaliveTimeout
rekeyTimeout := RekeyTimeout
if t := device.timings.rejectAfterTimeSec.Load(); !t.IsZero() {
rejectAfterTime = time.Duration(t.PickOne()) * time.Second
}
if t := device.timings.keepaliveTimeoutSec.Load(); !t.IsZero() {
keepaliveTimeout = time.Duration(t.Lo()) * time.Second
}
if t := device.timings.rekeyTimeoutSec.Load(); !t.IsZero() {
rekeyTimeout = time.Duration(t.Lo()) * time.Second
}
return max(0, rejectAfterTime-keepaliveTimeout-rekeyTimeout)
}
func (device *Device) keychainExpireTime() time.Duration {
rejectAfterTime := RejectAfterTime
if t := device.timings.rejectAfterTimeSec.Load(); !t.IsZero() {
rejectAfterTime = time.Duration(t.Hi()) * time.Second
}
return rejectAfterTime
}
func (device *Device) rekeyMinTimeout() time.Duration {
rekeyTimeout := RekeyTimeout
if t := device.timings.rekeyTimeoutSec.Load(); !t.IsZero() {
rekeyTimeout = time.Duration(t.Lo()) * time.Second
}
return rekeyTimeout
}
func (device *Device) maxHandshakeAttemps() uint32 {
res := uint32(MaxTimerHandshakes)
if t := device.timings.maxHandshakeAttemps.Load(); !t.IsZero() {
res = t.PickOne()
}
return res
}
+214 -110
View File
@@ -83,6 +83,9 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
device.peers.RLock()
defer device.peers.RUnlock()
device.headerProtection.RLock()
defer device.headerProtection.RUnlock()
// serialize device related values
if !device.staticIdentity.privateKey.IsZero() {
@@ -97,48 +100,48 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("fwmark=%d", device.net.fwmark)
}
if device.junk.count != 0 {
sendf("jc=%d", device.junk.count)
if count := device.junk.count.Load(); count != 0 {
sendf("jc=%d", count)
}
if device.junk.min != 0 {
sendf("jmin=%d", device.junk.min)
if min := device.junk.min.Load(); min != 0 {
sendf("jmin=%d", min)
}
if device.junk.max != 0 {
sendf("jmax=%d", device.junk.max)
if max := device.junk.max.Load(); max != 0 {
sendf("jmax=%d", max)
}
if device.paddings.init != 0 {
sendf("s1=%d", device.paddings.init)
if padding := device.paddings.init.Load(); padding != 0 {
sendf("s1=%d", padding)
}
if device.paddings.response != 0 {
sendf("s2=%d", device.paddings.response)
if padding := device.paddings.response.Load(); padding != 0 {
sendf("s2=%d", padding)
}
if device.paddings.cookie != 0 {
sendf("s3=%d", device.paddings.cookie)
if padding := device.paddings.cookie.Load(); padding != 0 {
sendf("s3=%d", padding)
}
if device.paddings.transport != 0 {
sendf("s4=%d", device.paddings.transport)
if padding := device.paddings.transport.Load(); padding != 0 {
sendf("s4=%d", padding)
}
if device.headers.init != nil {
sendf("h1=%s", device.headers.init.GenSpec())
if header := device.headers.init.Load(); !header.IsZero() {
sendf("h1=%s", header.ToString())
}
if device.headers.response != nil {
sendf("h2=%s", device.headers.response.GenSpec())
if header := device.headers.response.Load(); !header.IsZero() {
sendf("h2=%s", header.ToString())
}
if device.headers.cookie != nil {
sendf("h3=%s", device.headers.cookie.GenSpec())
if header := device.headers.cookie.Load(); !header.IsZero() {
sendf("h3=%s", header.ToString())
}
if device.headers.transport != nil {
sendf("h4=%s", device.headers.transport.GenSpec())
if header := device.headers.transport.Load(); !header.IsZero() {
sendf("h4=%s", header.ToString())
}
for i, ipacket := range device.ipackets {
@@ -147,6 +150,30 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
}
}
if !device.headerProtection.key.IsZero() {
keyf("header_protection_key", (*[32]byte)(&device.headerProtection.key))
}
if addition := device.contentPaddingAddition.Load(); !addition.IsZero() {
sendf("content_padding_addition=%s", addition.ToString())
}
if timing := device.timings.rekeyAfterTimeSec.Load(); !timing.IsZero() {
sendf("rekey_after_time=%s", timing.ToString())
}
if timing := device.timings.rekeyTimeoutSec.Load(); !timing.IsZero() {
sendf("rekey_timeout=%s", timing.ToString())
}
if timing := device.timings.rejectAfterTimeSec.Load(); !timing.IsZero() {
sendf("reject_after_time=%s", timing.ToString())
}
if timing := device.timings.keepaliveTimeoutSec.Load(); !timing.IsZero() {
sendf("keepalive_timeout=%s", timing.ToString())
}
if rang := device.timings.maxHandshakeAttemps.Load(); !rang.IsZero() {
sendf("max_handshake_attempts=%s", rang.ToString())
}
for _, peer := range device.peers.keyMap {
// Serialize peer state.
peer.handshake.mutex.RLock()
@@ -168,7 +195,10 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("last_handshake_time_nsec=%d", nano)
sendf("tx_bytes=%d", peer.txBytes.Load())
sendf("rx_bytes=%d", peer.rxBytes.Load())
sendf("persistent_keepalive_interval=%d", peer.persistentKeepaliveInterval.Load())
if keepalive := peer.persistentKeepaliveInterval.Load(); !keepalive.IsZero() {
sendf("persistent_keepalive_interval=%s", keepalive.ToString())
}
device.allowedips.EntriesForPeer(peer, func(prefix netip.Prefix) bool {
sendf("allowed_ip=%s", prefix.String())
@@ -198,6 +228,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
}()
ipcDev := new(ipcSetDevice)
ipcDev.fromDevice(device)
peer := new(ipcSetPeer)
deviceConfig := true
@@ -237,7 +268,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
var err error
if deviceConfig {
err = device.handleDeviceLine(key, value)
err = device.handleDeviceLine(ipcDev, key, value)
} else {
err = device.handlePeerLine(peer, key, value)
}
@@ -257,7 +288,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return nil
}
func (device *Device) handleDeviceLine(key, value string) error {
func (device *Device) handleDeviceLine(ipcDev *ipcSetDevice, key, value string) error {
switch key {
case "private_key":
var sk NoisePrivateKey
@@ -308,109 +339,87 @@ func (device *Device) handleDeviceLine(key, value string) error {
device.RemoveAllPeers()
case "jc":
jc, err := strconv.Atoi(value)
jc, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jc: %w", err)
}
if jc <= 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "jc must be a positive value")
}
device.log.Verbosef("UAPI: Updating junk count")
device.junk.count = jc
device.junk.count.Store(uint32(jc))
case "jmin":
jmin, err := strconv.Atoi(value)
jmin, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmin: %w", err)
}
if jmin <= 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "jmin must be a positive value")
}
device.log.Verbosef("UAPI: Updating junk min")
device.junk.min = jmin
device.junk.min.Store(uint32(jmin))
case "jmax":
jmax, err := strconv.Atoi(value)
jmax, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmax: %w", err)
}
if jmax <= 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "jmax must be a positive value")
}
device.log.Verbosef("UAPI: Updating junk max")
device.junk.max = jmax
device.junk.max.Store(uint32(jmax))
case "s1":
padding, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s1: %w", err)
}
if padding < 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "s1 must be non-negative")
}
device.log.Verbosef("UAPI: Updating s1 padding")
device.paddings.init = padding
ipcDev.paddings.init = uint32(padding)
case "s2":
padding, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s2: %w", err)
}
if padding < 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "s2 must be non-negative")
}
device.log.Verbosef("UAPI: Updating s2 padding")
device.paddings.response = padding
ipcDev.paddings.response = uint32(padding)
case "s3":
padding, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s3: %w", err)
}
if padding < 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "s3 must be non-negative")
}
device.log.Verbosef("UAPI: Updating s3 padding")
device.paddings.cookie = padding
ipcDev.paddings.cookie = uint32(padding)
case "s4":
padding, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s4: %w", err)
}
if padding < 0 {
return ipcErrorf(ipc.IpcErrorInvalid, "s4 must be non-negative")
}
device.log.Verbosef("UAPI: Updating s4 padding")
device.paddings.transport = padding
ipcDev.paddings.transport = uint32(padding)
case "h1":
header, err := newMagicHeader(value)
if err != nil {
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H1: %w", err)
}
device.headers.init = header
ipcDev.headers.init = rang
case "h2":
header, err := newMagicHeader(value)
if err != nil {
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H2: %w", err)
}
device.headers.response = header
ipcDev.headers.response = rang
case "h3":
header, err := newMagicHeader(value)
if err != nil {
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H3: %w", err)
}
device.headers.cookie = header
ipcDev.headers.cookie = rang
case "h4":
header, err := newMagicHeader(value)
if err != nil {
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H4: %w", err)
}
device.headers.transport = header
ipcDev.headers.transport = rang
case "i1":
chain, err := newObfChain(value)
@@ -447,6 +456,63 @@ func (device *Device) handleDeviceLine(key, value string) error {
}
device.ipackets[4] = chain
case "header_protection_key":
var key HeaderCipherKey
err := key.FromHex(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to set header_protection_key: %w", err)
}
ipcDev.headerProtectionKey = key
case "content_padding_addition":
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse content_padding_addition: %w", err)
}
device.log.Verbosef("UAPI: Updating content padding addition")
device.contentPaddingAddition.Store(rang)
case "rekey_after_time":
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse rekey after time: %w", err)
}
device.log.Verbosef("UAPI: Updating rekey after time")
device.timings.rekeyAfterTimeSec.Store(rang)
case "rekey_timeout":
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse rekey timeout: %w", err)
}
device.log.Verbosef("UAPI: Updating rekey timeout")
device.timings.rekeyTimeoutSec.Store(rang)
case "reject_after_time":
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse reject after time: %w", err)
}
device.log.Verbosef("UAPI: Updating reject after time")
device.timings.rejectAfterTimeSec.Store(rang)
case "keepalive_timeout":
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse keepalive timeout: %w", err)
}
device.log.Verbosef("UAPI: Updating keepalive timeout")
device.timings.keepaliveTimeoutSec.Store(rang)
case "max_handshake_attempts":
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse max handshake attempts: %w", err)
}
device.log.Verbosef("UAPI: Updating max handshake attempts")
device.timings.maxHandshakeAttemps.Store(rang)
default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
@@ -567,19 +633,15 @@ func (device *Device) handlePeerLine(
case "persistent_keepalive_interval":
device.log.Verbosef("%v - UAPI: Updating persistent keepalive interval", peer.Peer)
secs, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(
ipc.IpcErrorInvalid,
"failed to set persistent keepalive interval: %w",
err,
)
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to set persistent keepalive interval: %w", err)
}
old := peer.persistentKeepaliveInterval.Swap(uint32(secs))
old := peer.persistentKeepaliveInterval.Swap(rang)
// Send immediate keepalive if we're turning it on and before it wasn't on.
peer.pkaOn = old == 0 && secs != 0
peer.pkaOn = old.IsZero() && !rang.IsZero()
case "replace_allowed_ips":
device.log.Verbosef("%v - UAPI: Removing all allowedips", peer.Peer)
@@ -698,46 +760,88 @@ func (device *Device) IpcHandle(socket net.Conn) {
type ipcSetDevice struct {
headers struct {
init *magicHeader
response *magicHeader
cookie *magicHeader
transport *magicHeader
init UintRange
response UintRange
cookie UintRange
transport UintRange
}
paddings struct {
init uint32
response uint32
cookie uint32
transport uint32
}
headerProtectionKey HeaderCipherKey
}
func (d *ipcSetDevice) fromDevice(device *Device) {
device.headerProtection.RLock()
defer device.headerProtection.RUnlock()
d.headers.init = device.headers.init.Load()
d.headers.response = device.headers.response.Load()
d.headers.cookie = device.headers.cookie.Load()
d.headers.transport = device.headers.transport.Load()
d.paddings.init = device.paddings.init.Load()
d.paddings.response = device.paddings.response.Load()
d.paddings.cookie = device.paddings.cookie.Load()
d.paddings.transport = device.paddings.transport.Load()
d.headerProtectionKey = device.headerProtection.key
}
func (d *ipcSetDevice) mergeWithDevice(device *Device) error {
if d.headers.init == nil {
d.headers.init = device.headers.init
}
device.headerProtection.Lock()
defer device.headerProtection.Unlock()
if d.headers.response == nil {
d.headers.response = device.headers.response
}
if d.headers.cookie == nil {
d.headers.cookie = device.headers.cookie
}
if d.headers.transport == nil {
d.headers.transport = device.headers.transport
}
headers := []*magicHeader{d.headers.init, d.headers.response, d.headers.cookie, d.headers.transport}
headers := []UintRange{d.headers.init, d.headers.response, d.headers.cookie, d.headers.transport}
for i := 0; i < len(headers); i++ {
for j := i + 1; j < len(headers); j++ {
left := headers[i]
right := headers[j]
if left.start <= right.end && right.start <= left.end {
if left.Overlap(right) {
return errors.New("headers must not overlap")
}
}
}
device.headers.init = d.headers.init
device.headers.response = d.headers.response
device.headers.cookie = d.headers.cookie
device.headers.transport = d.headers.transport
device.log.Verbosef("UAPI: Updating h1 padding")
device.headers.init.Store(d.headers.init)
device.log.Verbosef("UAPI: Updating h2 padding")
device.headers.response.Store(d.headers.response)
device.log.Verbosef("UAPI: Updating h3 padding")
device.headers.cookie.Store(d.headers.cookie)
device.log.Verbosef("UAPI: Updating h4 padding")
device.headers.transport.Store(d.headers.transport)
if !d.headerProtectionKey.IsZero() {
paddings := []uint32{d.paddings.init, d.paddings.response, d.paddings.cookie, d.paddings.transport}
for i, padding := range paddings {
if padding < HeaderCipherNonceSize {
return fmt.Errorf("S%d must be more then 8 to use headerProtection", i)
}
}
}
device.log.Verbosef("UAPI: Updating s1 padding")
device.paddings.init.Store(d.paddings.init)
device.log.Verbosef("UAPI: Updating s2 padding")
device.paddings.response.Store(d.paddings.response)
device.log.Verbosef("UAPI: Updating s3 padding")
device.paddings.cookie.Store(d.paddings.cookie)
device.log.Verbosef("UAPI: Updating s4 padding")
device.paddings.transport.Store(d.paddings.transport)
device.log.Verbosef("UAPI: Updating header protection key")
device.headerProtection.key = d.headerProtectionKey
return nil
}