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