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
+743 -314

No files matched your search

+51 -1
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
@@ -110,4 +160,4 @@ Value is a sequence of tags specified below:
> Custom signature packets does not carry any actual data, so there is no need to specify it on both sides. General recommendation is to use it on the client side only
> [!IMPORTANT]
> If the final size of any packet exceeds system MTU, it would be fractured into fragments, which looks suspicious
> If the final size of any packet exceeds system MTU, it would be fractured into fragments, which looks suspicious
+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
}