Compare commits

...
Author SHA1 Message Date
itfsdev 0527dfa476 fix: bump Docker builder image to golang:1.25.12 2026-07-28 18:33:16 +08:00
Yaroslav Gurov 9f5d948bc7 fix: use v3 versioning 2026-07-24 15:06:35 +02:00
Yaroslav Gurov 457d920a1a 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
2026-07-24 14:18:33 +02:00
Vinicius Fortuna c1e9bb3758 Update Outline SDK module path 2026-07-09 04:22:33 +02:00
Yaroslav Gurov 1cc94272ca fix: handle empty I1-I5 2026-06-18 00:38:58 +02:00
stereomonk 3610f21b75 doc: fix typo README.md 2026-06-17 22:22:36 +02:00
admin f4f4c99926 fix: apply S4 transport padding to keepalive packets
Keepalive packets were excluded from S4 padding because the padding
logic was nested inside the dataSent guard. The receiving side
(DeterminePacketTypeAndPadding) expects S4 padding on all transport
packets, so unpadded keepalives fail H4 header validation and are
silently dropped.

This prevents the responder from completing key confirmation —
lastHandshakeNano stays 0 until real data flows through the tunnel.
2026-05-13 11:11:21 +02:00
Yaroslav Gurov 12a012205e readme: actualize type for H1-H4 2026-03-31 17:56:17 +02:00
Yaroslav Gurov e7ef4339e7 readme: remove <c> tag from tag reference 2026-03-23 12:10:15 +01:00
Yaroslav Gurov 449d7cffd4 Feature/outline glue (#106)
* feat: added outline integration layer

* chore: make the function used in RegisterFallbackParser a standalone one

* fix: check if domain has a dot prior trimming it

* fix: use net.JoinHostPort instead of plain concat
2025-12-19 03:14:48 +01:00
vkamn e796d477d8 chore: update license (#105)
Signed-off-by: vkamn <vk@amnezia.org>
2025-12-11 18:56:42 +08:00
Yaroslav Gurov 730d6c39d0 chore: add docs for the params from awg2 2025-12-01 13:11:33 +01:00
Yaroslav Gurov 0361c54dca fix: refactor processing of junk packets (#103)
- fix the bug that transport packet interprets as init/resp/cookie with the same size
- cleanup error responses
- reduce buffer allocations
2025-12-01 20:07:48 +08:00
Mark PuhaandYaroslav Gurov f6542209f4 feat: awg 2.0 (#91)
* feat: ranged H1-H4
* feat: S3, S4 support
* chore: updated awg-tools version

---------

Co-authored-by: Yaroslav Gurov <ygurov@proton.me>
2025-09-01 14:04:52 +02:00
64 changed files with 1939 additions and 2247 deletions

No files matched your search

+2 -2
View File
@@ -1,4 +1,4 @@
FROM golang:1.24.4 as awg
FROM golang:1.25.12 as awg
COPY . /awg
WORKDIR /awg
RUN go mod download && \
@@ -6,7 +6,7 @@ RUN go mod download && \
go build -ldflags '-linkmode external -extldflags "-fno-PIC -static"' -v -o /usr/bin
FROM alpine:3.19
ARG AWGTOOLS_RELEASE="1.0.20241018"
ARG AWGTOOLS_RELEASE="1.0.20250901"
RUN apk --no-cache add iproute2 iptables bash && \
cd /usr/bin/ && \
+2
View File
@@ -1,3 +1,5 @@
Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
Permission is hereby granted, free of charge, to any person obtaining a copy of
this software and associated documentation files (the "Software"), to deal in
the Software without restriction, including without limitation the rights to
+111 -1
View File
@@ -50,4 +50,114 @@ $ git clone https://github.com/amnezia-vpn/amneziawg-go
$ cd amneziawg-go
$ 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
- `Jc: int`, recommended range is 4-12
- `Jmin: int` <= `Jmax:int`
> [!TIP]
> Junk packets do 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 Jmax >= system MTU (not the one specified in AWG), then the system can fracture this packet into fragments, which looks suspicious from the censor side
### Message paddings
- `S1: int` - padding of handshake initial message
- `S2: int` - padding of handshake response message
- `S3: int` - padding of handshake cookie message
- `S4: int` - padding of transport messages
### Message headers
Every message in wireguard has `uint32` type at the beginning of the packet. This field could be controlled by specifying the params below:
- `H1: string` - header range of handshake initial message
- `H2: string` - header range of handshake response message
- `H3: string` - header range of handshake cookie message
- `H4: string` - header range of transport message
Values could be specified as:
- range: `x-y`, x <= y; e.g. `123-456`
- single value `1234`
### Custom signature packets
These packets are being send prior to every handshake, in the same way as Junk packets do. The sending order is `I1`, `I2`, `I3`, `I4`, `I5`. If there is no value specified, the packet is skipped.
- `I1: string`
- `I2: string`
- `I3: string`
- `I4: string`
- `I5: string`
Value is a sequence of tags specified below:
- `<b 0x[seq]>` - static bytes tag. Dumps `[seq]` as-is to the packet. `[seq]` is hex-encoded sequence which represents bytes sequence (2 hex numbers per byte) and is always even-sized
- `<r [size]>` - random bytes tag. Dumps `[size]` amount of randomly-generated bytes to the packet
- `<rd [size]>` - random digits tag. Dumps `[size]` amount of randomly-generated bytes from `[0-9]` set to the packet
- `<rc [size]>` - random chars tag. Dumps `[size]` amount of randomly-generated bytes from `[a-zA-Z] set to the packet
- `<t>` - timestamp tag. Dumps 4-bytes long current system time in UNIX format
> [!TIP]
> 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
+1 -1
View File
@@ -17,7 +17,7 @@ import (
"golang.org/x/sys/windows"
"github.com/amnezia-vpn/amneziawg-go/conn/winrio"
"github.com/amnezia-vpn/amneziawg-go/v3/conn/winrio"
)
const (
+1 -1
View File
@@ -12,7 +12,7 @@ import (
"net/netip"
"os"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
)
type ChannelBind struct {
-144
View File
@@ -1,144 +0,0 @@
package awg
import (
"bytes"
"fmt"
"slices"
"strconv"
"strings"
"sync"
"github.com/tevino/abool"
)
type aSecCfgType struct {
IsSet bool
JunkPacketCount int
JunkPacketMinSize int
JunkPacketMaxSize int
InitHeaderJunkSize int
ResponseHeaderJunkSize int
CookieReplyHeaderJunkSize int
TransportHeaderJunkSize int
InitPacketMagicHeader uint32
ResponsePacketMagicHeader uint32
UnderloadPacketMagicHeader uint32
TransportPacketMagicHeader uint32
// InitPacketMagicHeader Limit
// ResponsePacketMagicHeader Limit
// UnderloadPacketMagicHeader Limit
// TransportPacketMagicHeader Limit
}
type Limit struct {
Min uint32
Max uint32
HeaderType uint32
}
func NewLimit(min, max, headerType uint32) (Limit, error) {
if min > max {
return Limit{}, fmt.Errorf("min (%d) cannot be greater than max (%d)", min, max)
}
return Limit{
Min: min,
Max: max,
HeaderType: headerType,
}, nil
}
func ParseMagicHeader(key, value string, defaultHeaderType uint32) (Limit, error) {
// tempAwg.ASecCfg.InitPacketMagicHeader, err = awg.NewLimit(uint32(initPacketMagicHeaderMin), uint32(initPacketMagicHeaderMax), DNewLimit(min, max, headerType)efaultMessageInitiationType)
// var min, max, headerType uint32
// _, err := fmt.Sscanf(value, "%d-%d:%d", &min, &max, &headerType)
// if err != nil {
// return Limit{}, fmt.Errorf("invalid magic header format: %s", value)
// }
limits := strings.Split(value, "-")
if len(limits) != 2 {
return Limit{}, fmt.Errorf("invalid format for key: %s; %s", key, value)
}
min, err := strconv.ParseUint(limits[0], 10, 32)
if err != nil {
return Limit{}, fmt.Errorf("parse min key: %s; value: ; %w", key, limits[0], err)
}
max, err := strconv.ParseUint(limits[1], 10, 32)
if err != nil {
return Limit{}, fmt.Errorf("parse max key: %s; value: ; %w", key, limits[0], err)
}
limit, err := NewLimit(uint32(min), uint32(max), defaultHeaderType)
if err != nil {
return Limit{}, fmt.Errorf("new lmit key: %s; value: ; %w", key, limits[0], err)
}
return limit, nil
}
type Limits []Limit
func NewLimits(limits []Limit) Limits {
slices.SortFunc(limits, func(a, b Limit) int {
if a.Min < b.Min {
return -1
} else if a.Min > b.Min {
return 1
}
return 0
})
return Limits(limits)
}
type Protocol struct {
IsASecOn abool.AtomicBool
// TODO: revision the need of the mutex
ASecMux sync.RWMutex
ASecCfg aSecCfgType
JunkCreator junkCreator
HandshakeHandler SpecialHandshakeHandler
}
func (protocol *Protocol) CreateInitHeaderJunk() ([]byte, error) {
return protocol.createHeaderJunk(protocol.ASecCfg.InitHeaderJunkSize)
}
func (protocol *Protocol) CreateResponseHeaderJunk() ([]byte, error) {
return protocol.createHeaderJunk(protocol.ASecCfg.ResponseHeaderJunkSize)
}
func (protocol *Protocol) CreateCookieReplyHeaderJunk() ([]byte, error) {
return protocol.createHeaderJunk(protocol.ASecCfg.CookieReplyHeaderJunkSize)
}
func (protocol *Protocol) CreateTransportHeaderJunk(packetSize int) ([]byte, error) {
return protocol.createHeaderJunk(protocol.ASecCfg.TransportHeaderJunkSize, packetSize)
}
func (protocol *Protocol) createHeaderJunk(junkSize int, optExtraSize ...int) ([]byte, error) {
extraSize := 0
if len(optExtraSize) == 1 {
extraSize = optExtraSize[0]
}
var junk []byte
protocol.ASecMux.RLock()
if junkSize != 0 {
buf := make([]byte, 0, junkSize+extraSize)
writer := bytes.NewBuffer(buf[:0])
err := protocol.JunkCreator.AppendJunk(writer, junkSize)
if err != nil {
protocol.ASecMux.RUnlock()
return nil, err
}
junk = writer.Bytes()
}
protocol.ASecMux.RUnlock()
return junk, nil
}
-37
View File
@@ -1,37 +0,0 @@
package internal
type mockGenerator struct {
size int
}
func NewMockGenerator(size int) mockGenerator {
return mockGenerator{size: size}
}
func (m mockGenerator) Generate() []byte {
return make([]byte, m.size)
}
func (m mockGenerator) Size() int {
return m.size
}
func (m mockGenerator) Name() string {
return "mock"
}
type mockByteGenerator struct {
data []byte
}
func NewMockByteGenerator(data []byte) mockByteGenerator {
return mockByteGenerator{data: data}
}
func (bg mockByteGenerator) Generate() []byte {
return bg.data
}
func (bg mockByteGenerator) Size() int {
return len(bg.data)
}
-70
View File
@@ -1,70 +0,0 @@
package awg
import (
"bytes"
crand "crypto/rand"
"fmt"
v2 "math/rand/v2"
)
type junkCreator struct {
aSecCfg aSecCfgType
cha8Rand *v2.ChaCha8
}
// TODO: refactor param to only pass the junk related params
func NewJunkCreator(aSecCfg aSecCfgType) (junkCreator, error) {
buf := make([]byte, 32)
_, err := crand.Read(buf)
if err != nil {
return junkCreator{}, err
}
return junkCreator{aSecCfg: aSecCfg, cha8Rand: v2.NewChaCha8([32]byte(buf))}, nil
}
// Should be called with aSecMux RLocked
func (jc *junkCreator) CreateJunkPackets(junks *[][]byte) error {
if jc.aSecCfg.JunkPacketCount == 0 {
return nil
}
for range jc.aSecCfg.JunkPacketCount {
packetSize := jc.randomPacketSize()
junk, err := jc.randomJunkWithSize(packetSize)
if err != nil {
return fmt.Errorf("create junk packet: %v", err)
}
*junks = append(*junks, junk)
}
return nil
}
// Should be called with aSecMux RLocked
func (jc *junkCreator) randomPacketSize() int {
return int(
jc.cha8Rand.Uint64()%uint64(
jc.aSecCfg.JunkPacketMaxSize-jc.aSecCfg.JunkPacketMinSize,
),
) + jc.aSecCfg.JunkPacketMinSize
}
// Should be called with aSecMux RLocked
func (jc *junkCreator) AppendJunk(writer *bytes.Buffer, size int) error {
headerJunk, err := jc.randomJunkWithSize(size)
if err != nil {
return fmt.Errorf("create header junk: %v", err)
}
_, err = writer.Write(headerJunk)
if err != nil {
return fmt.Errorf("write header junk: %v", err)
}
return nil
}
// Should be called with aSecMux RLocked
func (jc *junkCreator) randomJunkWithSize(size int) ([]byte, error) {
// TODO: use a memory pool to allocate
junk := make([]byte, size)
_, err := jc.cha8Rand.Read(junk)
return junk, err
}
-115
View File
@@ -1,115 +0,0 @@
package awg
import (
"bytes"
"fmt"
"testing"
)
func setUpJunkCreator(t *testing.T) (junkCreator, error) {
jc, err := NewJunkCreator(aSecCfgType{
IsSet: true,
JunkPacketCount: 5,
JunkPacketMinSize: 500,
JunkPacketMaxSize: 1000,
InitHeaderJunkSize: 30,
ResponseHeaderJunkSize: 40,
InitPacketMagicHeader: 123456,
ResponsePacketMagicHeader: 67543,
UnderloadPacketMagicHeader: 32345,
TransportPacketMagicHeader: 123123,
})
if err != nil {
t.Errorf("failed to create junk creator %v", err)
return junkCreator{}, err
}
return jc, nil
}
func Test_junkCreator_createJunkPackets(t *testing.T) {
jc, err := setUpJunkCreator(t)
if err != nil {
return
}
t.Run("valid", func(t *testing.T) {
got := make([][]byte, 0, jc.aSecCfg.JunkPacketCount)
err := jc.CreateJunkPackets(&got)
if err != nil {
t.Errorf(
"junkCreator.createJunkPackets() = %v; failed",
err,
)
return
}
seen := make(map[string]bool)
for _, junk := range got {
key := string(junk)
if seen[key] {
t.Errorf(
"junkCreator.createJunkPackets() = %v, duplicate key: %v",
got,
junk,
)
return
}
seen[key] = true
}
})
}
func Test_junkCreator_randomJunkWithSize(t *testing.T) {
t.Run("valid", func(t *testing.T) {
jc, err := setUpJunkCreator(t)
if err != nil {
return
}
r1, _ := jc.randomJunkWithSize(10)
r2, _ := jc.randomJunkWithSize(10)
fmt.Printf("%v\n%v\n", r1, r2)
if bytes.Equal(r1, r2) {
t.Errorf("same junks %v", err)
return
}
})
}
func Test_junkCreator_randomPacketSize(t *testing.T) {
jc, err := setUpJunkCreator(t)
if err != nil {
return
}
for range [30]struct{}{} {
t.Run("valid", func(t *testing.T) {
if got := jc.randomPacketSize(); jc.aSecCfg.JunkPacketMinSize > got ||
got > jc.aSecCfg.JunkPacketMaxSize {
t.Errorf(
"junkCreator.randomPacketSize() = %v, not between range [%v,%v]",
got,
jc.aSecCfg.JunkPacketMinSize,
jc.aSecCfg.JunkPacketMaxSize,
)
}
})
}
}
func Test_junkCreator_appendJunk(t *testing.T) {
jc, err := setUpJunkCreator(t)
if err != nil {
return
}
t.Run("valid", func(t *testing.T) {
s := "apple"
buffer := bytes.NewBuffer([]byte(s))
err := jc.AppendJunk(buffer, 30)
if err != nil &&
buffer.Len() != len(s)+30 {
t.Error("appendWithJunk() size don't match")
}
read := make([]byte, 50)
buffer.Read(read)
fmt.Println(string(read))
})
}
-73
View File
@@ -1,73 +0,0 @@
package awg
import (
"errors"
"time"
"github.com/tevino/abool"
"go.uber.org/atomic"
)
// TODO: atomic?/ and better way to use this
var PacketCounter *atomic.Uint64 = atomic.NewUint64(0)
// TODO
var WaitResponse = struct {
Channel chan struct{}
ShouldWait *abool.AtomicBool
}{
make(chan struct{}, 1),
abool.New(),
}
type SpecialHandshakeHandler struct {
isFirstDone bool
SpecialJunk TagJunkPacketGenerators
ControlledJunk TagJunkPacketGenerators
nextItime time.Time
ITimeout time.Duration // seconds
IsSet bool
}
func (handler *SpecialHandshakeHandler) Validate() error {
var errs []error
if err := handler.SpecialJunk.Validate(); err != nil {
errs = append(errs, err)
}
if err := handler.ControlledJunk.Validate(); err != nil {
errs = append(errs, err)
}
return errors.Join(errs...)
}
func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
if !handler.SpecialJunk.IsDefined() {
return nil
}
// TODO: create tests
if !handler.isFirstDone {
handler.isFirstDone = true
} else if !handler.isTimeToSendSpecial() {
return nil
}
rv := handler.SpecialJunk.GeneratePackets()
handler.nextItime = time.Now().Add(handler.ITimeout)
return rv
}
func (handler *SpecialHandshakeHandler) isTimeToSendSpecial() bool {
return time.Now().After(handler.nextItime)
}
func (handler *SpecialHandshakeHandler) GenerateControlledJunk() [][]byte {
if !handler.ControlledJunk.IsDefined() {
return nil
}
return handler.ControlledJunk.GeneratePackets()
}
-190
View File
@@ -1,190 +0,0 @@
package awg
import (
crand "crypto/rand"
"encoding/binary"
"encoding/hex"
"fmt"
"strconv"
"strings"
"time"
v2 "math/rand/v2"
// "go.uber.org/atomic"
)
type Generator interface {
Generate() []byte
Size() int
}
type newGenerator func(string) (Generator, error)
type BytesGenerator struct {
value []byte
size int
}
func (bg *BytesGenerator) Generate() []byte {
return bg.value
}
func (bg *BytesGenerator) Size() int {
return bg.size
}
func newBytesGenerator(param string) (Generator, error) {
hasPrefix := strings.HasPrefix(param, "0x") || strings.HasPrefix(param, "0X")
if !hasPrefix {
return nil, fmt.Errorf("not correct hex: %s", param)
}
hex, err := hexToBytes(param)
if err != nil {
return nil, fmt.Errorf("hexToBytes: %w", err)
}
return &BytesGenerator{value: hex, size: len(hex)}, nil
}
func hexToBytes(hexStr string) ([]byte, error) {
hexStr = strings.TrimPrefix(hexStr, "0x")
hexStr = strings.TrimPrefix(hexStr, "0X")
// Ensure even length (pad with leading zero if needed)
if len(hexStr)%2 != 0 {
hexStr = "0" + hexStr
}
return hex.DecodeString(hexStr)
}
type RandomPacketGenerator struct {
cha8Rand *v2.ChaCha8
size int
}
func (rpg *RandomPacketGenerator) Generate() []byte {
junk := make([]byte, rpg.size)
rpg.cha8Rand.Read(junk)
return junk
}
func (rpg *RandomPacketGenerator) Size() int {
return rpg.size
}
func newRandomPacketGenerator(param string) (Generator, error) {
size, err := strconv.Atoi(param)
if err != nil {
return nil, fmt.Errorf("random packet parse int: %w", err)
}
if size > 1000 {
return nil, fmt.Errorf("random packet size must be less than 1000")
}
buf := make([]byte, 32)
_, err = crand.Read(buf)
if err != nil {
return nil, fmt.Errorf("random packet crand read: %w", err)
}
return &RandomPacketGenerator{
cha8Rand: v2.NewChaCha8([32]byte(buf)),
size: size,
}, nil
}
type TimestampGenerator struct {
}
func (tg *TimestampGenerator) Generate() []byte {
buf := make([]byte, 8)
binary.BigEndian.PutUint64(buf, uint64(time.Now().Unix()))
return buf
}
func (tg *TimestampGenerator) Size() int {
return 8
}
func newTimestampGenerator(param string) (Generator, error) {
if len(param) != 0 {
return nil, fmt.Errorf("timestamp param needs to be empty: %s", param)
}
return &TimestampGenerator{}, nil
}
type WaitTimeoutGenerator struct {
waitTimeout time.Duration
}
func (wtg *WaitTimeoutGenerator) Generate() []byte {
time.Sleep(wtg.waitTimeout)
return []byte{}
}
func (wtg *WaitTimeoutGenerator) Size() int {
return 0
}
func newWaitTimeoutGenerator(param string) (Generator, error) {
timeout, err := strconv.Atoi(param)
if err != nil {
return nil, fmt.Errorf("timeout parse int: %w", err)
}
if timeout > 5000 {
return nil, fmt.Errorf("timeout must be less than 5000ms")
}
return &WaitTimeoutGenerator{
waitTimeout: time.Duration(timeout) * time.Millisecond,
}, nil
}
type PacketCounterGenerator struct {
}
func (c *PacketCounterGenerator) Generate() []byte {
buf := make([]byte, 8)
// TODO: better way to handle counter tag
binary.BigEndian.PutUint64(buf, PacketCounter.Load())
return buf
}
func (c *PacketCounterGenerator) Size() int {
return 8
}
func newPacketCounterGenerator(param string) (Generator, error) {
if len(param) != 0 {
return nil, fmt.Errorf("packet counter param needs to be empty: %s", param)
}
return &PacketCounterGenerator{}, nil
}
type WaitResponseGenerator struct {
}
func (c *WaitResponseGenerator) Generate() []byte {
WaitResponse.ShouldWait.Set()
<-WaitResponse.Channel
WaitResponse.ShouldWait.UnSet()
return []byte{}
}
func (c *WaitResponseGenerator) Size() int {
return 0
}
func newWaitResponseGenerator(param string) (Generator, error) {
if len(param) != 0 {
return nil, fmt.Errorf("wait response param needs to be empty: %s", param)
}
return &WaitResponseGenerator{}, nil
}
-189
View File
@@ -1,189 +0,0 @@
package awg
import (
"encoding/binary"
"fmt"
"testing"
"github.com/stretchr/testify/require"
)
func Test_newBytesGenerator(t *testing.T) {
type args struct {
param string
}
tests := []struct {
name string
args args
want []byte
wantErr error
}{
{
name: "empty",
args: args{
param: "",
},
wantErr: fmt.Errorf("not correct hex"),
},
{
name: "wrong start",
args: args{
param: "123456",
},
wantErr: fmt.Errorf("not correct hex"),
},
{
name: "not only hex value with X",
args: args{
param: "0X12345q",
},
wantErr: fmt.Errorf("not correct hex"),
},
{
name: "not only hex value with x",
args: args{
param: "0x12345q",
},
wantErr: fmt.Errorf("not correct hex"),
},
{
name: "valid hex",
args: args{
param: "0xf6ab3267fa",
},
want: []byte{0xf6, 0xab, 0x32, 0x67, 0xfa},
},
{
name: "valid hex with odd length",
args: args{
param: "0xfab3267fa",
},
want: []byte{0xf, 0xab, 0x32, 0x67, 0xfa},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := newBytesGenerator(tt.args.param)
if tt.wantErr != nil {
require.ErrorAs(t, err, &tt.wantErr)
require.Nil(t, got)
return
}
require.Nil(t, err)
require.NotNil(t, got)
gotValues := got.Generate()
require.Equal(t, tt.want, gotValues)
})
}
}
func Test_newRandomPacketGenerator(t *testing.T) {
type args struct {
param string
}
tests := []struct {
name string
args args
wantErr error
}{
{
name: "empty",
args: args{
param: "",
},
wantErr: fmt.Errorf("parse int"),
},
{
name: "not an int",
args: args{
param: "x",
},
wantErr: fmt.Errorf("parse int"),
},
{
name: "too large",
args: args{
param: "1001",
},
wantErr: fmt.Errorf("random packet size must be less than 1000"),
},
{
name: "valid",
args: args{
param: "12",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := newRandomPacketGenerator(tt.args.param)
if tt.wantErr != nil {
require.ErrorAs(t, err, &tt.wantErr)
require.Nil(t, got)
return
}
require.Nil(t, err)
require.NotNil(t, got)
first := got.Generate()
second := got.Generate()
require.NotEqual(t, first, second)
})
}
}
func TestPacketCounterGenerator(t *testing.T) {
tests := []struct {
name string
param string
wantErr bool
}{
{
name: "Valid empty param",
param: "",
wantErr: false,
},
{
name: "Invalid non-empty param",
param: "anything",
wantErr: true,
},
}
for _, tc := range tests {
tc := tc // capture range variable
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
gen, err := newPacketCounterGenerator(tc.param)
if tc.wantErr {
require.Error(t, err)
return
}
require.NoError(t, err)
require.Equal(t, 8, gen.Size())
// Reset counter to known value for test
initialCount := uint64(42)
PacketCounter.Store(initialCount)
output := gen.Generate()
require.Equal(t, 8, len(output))
// Verify counter value in output
counterValue := binary.BigEndian.Uint64(output)
require.Equal(t, initialCount, counterValue)
// Increment counter and verify change
PacketCounter.Add(1)
output = gen.Generate()
counterValue = binary.BigEndian.Uint64(output)
require.Equal(t, initialCount+1, counterValue)
})
}
}
-59
View File
@@ -1,59 +0,0 @@
package awg
import (
"fmt"
"strconv"
)
type TagJunkPacketGenerator struct {
name string
tagValue string
packetSize int
generators []Generator
}
func newTagJunkPacketGenerator(name, tagValue string, size int) TagJunkPacketGenerator {
return TagJunkPacketGenerator{
name: name,
tagValue: tagValue,
generators: make([]Generator, 0, size),
}
}
func (tg *TagJunkPacketGenerator) append(generator Generator) {
tg.generators = append(tg.generators, generator)
tg.packetSize += generator.Size()
}
func (tg *TagJunkPacketGenerator) generatePacket() []byte {
packet := make([]byte, 0, tg.packetSize)
for _, generator := range tg.generators {
packet = append(packet, generator.Generate()...)
}
return packet
}
func (tg *TagJunkPacketGenerator) Name() string {
return tg.name
}
func (tg *TagJunkPacketGenerator) nameIndex() (int, error) {
if len(tg.name) != 2 {
return 0, fmt.Errorf("name must be 2 character long: %s", tg.name)
}
index, err := strconv.Atoi(tg.name[1:2])
if err != nil {
return 0, fmt.Errorf("name 2 char should be an int %w", err)
}
return index, nil
}
func (tg *TagJunkPacketGenerator) IpcGetFields() IpcFields {
return IpcFields{
Key: tg.name,
Value: tg.tagValue,
}
}
@@ -1,210 +0,0 @@
package awg
import (
"testing"
"github.com/amnezia-vpn/amneziawg-go/device/awg/internal"
"github.com/stretchr/testify/require"
)
func TestNewTagJunkGenerator(t *testing.T) {
t.Parallel()
testCases := []struct {
name string
genName string
size int
expected TagJunkPacketGenerator
}{
{
name: "Create new generator with empty name",
genName: "",
size: 0,
expected: TagJunkPacketGenerator{
name: "",
packetSize: 0,
generators: make([]Generator, 0),
},
},
{
name: "Create new generator with valid name",
genName: "T1",
size: 0,
expected: TagJunkPacketGenerator{
name: "T1",
packetSize: 0,
generators: make([]Generator, 0),
},
},
{
name: "Create new generator with non-zero size",
genName: "T2",
size: 5,
expected: TagJunkPacketGenerator{
name: "T2",
packetSize: 0,
generators: make([]Generator, 5),
},
},
}
for _, tc := range testCases {
tc := tc // capture range variable
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
result := newTagJunkPacketGenerator(tc.genName, "", tc.size)
require.Equal(t, tc.expected.name, result.name)
require.Equal(t, tc.expected.packetSize, result.packetSize)
require.Equal(t, cap(result.generators), len(tc.expected.generators))
})
}
}
func TestTagJunkGeneratorAppend(t *testing.T) {
t.Parallel()
testCases := []struct {
name string
initialState TagJunkPacketGenerator
mockSize int
expectedLength int
expectedSize int
}{
{
name: "Append to empty generator",
initialState: newTagJunkPacketGenerator("T1", "", 0),
mockSize: 5,
expectedLength: 1,
expectedSize: 5,
},
{
name: "Append to non-empty generator",
initialState: TagJunkPacketGenerator{
name: "T2",
packetSize: 10,
generators: make([]Generator, 2),
},
mockSize: 7,
expectedLength: 3, // 2 existing + 1 new
expectedSize: 17, // 10 + 7
},
}
for _, tc := range testCases {
tc := tc // capture range variable
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
tg := tc.initialState
mockGen := internal.NewMockGenerator(tc.mockSize)
tg.append(mockGen)
require.Equal(t, tc.expectedLength, len(tg.generators))
require.Equal(t, tc.expectedSize, tg.packetSize)
})
}
}
func TestTagJunkGeneratorGenerate(t *testing.T) {
t.Parallel()
// Create mock generators for testing
mockGen1 := internal.NewMockByteGenerator([]byte{0x01, 0x02})
mockGen2 := internal.NewMockByteGenerator([]byte{0x03, 0x04, 0x05})
testCases := []struct {
name string
setupGenerator func() TagJunkPacketGenerator
expected []byte
}{
{
name: "Generate with empty generators",
setupGenerator: func() TagJunkPacketGenerator {
return newTagJunkPacketGenerator("T1", "", 0)
},
expected: []byte{},
},
{
name: "Generate with single generator",
setupGenerator: func() TagJunkPacketGenerator {
tg := newTagJunkPacketGenerator("T2", "", 0)
tg.append(mockGen1)
return tg
},
expected: []byte{0x01, 0x02},
},
{
name: "Generate with multiple generators",
setupGenerator: func() TagJunkPacketGenerator {
tg := newTagJunkPacketGenerator("T3", "", 0)
tg.append(mockGen1)
tg.append(mockGen2)
return tg
},
expected: []byte{0x01, 0x02, 0x03, 0x04, 0x05},
},
}
for _, tc := range testCases {
tc := tc // capture range variable
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
tg := tc.setupGenerator()
result := tg.generatePacket()
require.Equal(t, tc.expected, result)
})
}
}
func TestTagJunkGeneratorNameIndex(t *testing.T) {
t.Parallel()
testCases := []struct {
name string
generatorName string
expectedIndex int
expectError bool
}{
{
name: "Valid name with digit",
generatorName: "T5",
expectedIndex: 5,
expectError: false,
},
{
name: "Invalid name - too short",
generatorName: "T",
expectError: true,
},
{
name: "Invalid name - too long",
generatorName: "T55",
expectError: true,
},
{
name: "Invalid name - non-digit second character",
generatorName: "TX",
expectError: true,
},
}
for _, tc := range testCases {
tc := tc // capture range variable
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
tg := TagJunkPacketGenerator{name: tc.generatorName}
index, err := tg.nameIndex()
if tc.expectError {
require.Error(t, err)
} else {
require.NoError(t, err)
require.Equal(t, tc.expectedIndex, index)
}
})
}
}
-66
View File
@@ -1,66 +0,0 @@
package awg
import "fmt"
type TagJunkPacketGenerators struct {
tagGenerators []TagJunkPacketGenerator
length int
DefaultJunkCount int // Jc
}
func (generators *TagJunkPacketGenerators) AppendGenerator(
generator TagJunkPacketGenerator,
) {
generators.tagGenerators = append(generators.tagGenerators, generator)
generators.length++
}
func (generators *TagJunkPacketGenerators) IsDefined() bool {
return len(generators.tagGenerators) > 0
}
// validate that packets were defined consecutively
func (generators *TagJunkPacketGenerators) Validate() error {
seen := make([]bool, len(generators.tagGenerators))
for _, generator := range generators.tagGenerators {
index, err := generator.nameIndex()
if index > len(generators.tagGenerators) {
return fmt.Errorf("junk packet index should be consecutive")
}
if err != nil {
return fmt.Errorf("name index: %w", err)
} else {
seen[index-1] = true
}
}
for _, found := range seen {
if !found {
return fmt.Errorf("junk packet index should be consecutive")
}
}
return nil
}
func (generators *TagJunkPacketGenerators) GeneratePackets() [][]byte {
var rv = make([][]byte, 0, generators.length+generators.DefaultJunkCount)
for i, tagGenerator := range generators.tagGenerators {
rv = append(rv, make([]byte, tagGenerator.packetSize))
copy(rv[i], tagGenerator.generatePacket())
PacketCounter.Inc()
}
PacketCounter.Add(uint64(generators.DefaultJunkCount))
return rv
}
func (tg *TagJunkPacketGenerators) IpcGetFields() []IpcFields {
rv := make([]IpcFields, 0, len(tg.tagGenerators))
for _, generator := range tg.tagGenerators {
rv = append(rv, generator.IpcGetFields())
}
return rv
}
@@ -1,149 +0,0 @@
package awg
import (
"testing"
"github.com/amnezia-vpn/amneziawg-go/device/awg/internal"
"github.com/stretchr/testify/require"
)
func TestTagJunkGeneratorHandlerAppendGenerator(t *testing.T) {
tests := []struct {
name string
generator TagJunkPacketGenerator
}{
{
name: "append single generator",
generator: newTagJunkPacketGenerator("t1", "", 10),
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
generators := &TagJunkPacketGenerators{}
// Initial length should be 0
require.Equal(t, 0, generators.length)
require.Empty(t, generators.tagGenerators)
// After append, length should be 1 and generator should be added
generators.AppendGenerator(tt.generator)
require.Equal(t, 1, generators.length)
require.Len(t, generators.tagGenerators, 1)
require.Equal(t, tt.generator, generators.tagGenerators[0])
})
}
}
func TestTagJunkGeneratorHandlerValidate(t *testing.T) {
tests := []struct {
name string
generators []TagJunkPacketGenerator
wantErr bool
errMsg string
}{
{
name: "bad start",
generators: []TagJunkPacketGenerator{
newTagJunkPacketGenerator("t3", "", 10),
newTagJunkPacketGenerator("t4", "", 10),
},
wantErr: true,
errMsg: "junk packet index should be consecutive",
},
{
name: "non-consecutive indices",
generators: []TagJunkPacketGenerator{
newTagJunkPacketGenerator("t1", "", 10),
newTagJunkPacketGenerator("t3", "", 10), // Missing t2
},
wantErr: true,
errMsg: "junk packet index should be consecutive",
},
{
name: "consecutive indices",
generators: []TagJunkPacketGenerator{
newTagJunkPacketGenerator("t1", "", 10),
newTagJunkPacketGenerator("t2", "", 10),
newTagJunkPacketGenerator("t3", "", 10),
newTagJunkPacketGenerator("t4", "", 10),
newTagJunkPacketGenerator("t5", "", 10),
},
},
{
name: "nameIndex error",
generators: []TagJunkPacketGenerator{
newTagJunkPacketGenerator("error", "", 10),
},
wantErr: true,
errMsg: "name must be 2 character long",
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
generators := &TagJunkPacketGenerators{}
for _, gen := range tt.generators {
generators.AppendGenerator(gen)
}
err := generators.Validate()
if tt.wantErr {
require.Error(t, err)
require.Contains(t, err.Error(), tt.errMsg)
return
}
require.NoError(t, err)
})
}
}
func TestTagJunkGeneratorHandlerGenerate(t *testing.T) {
mockByte1 := []byte{0x01, 0x02}
mockByte2 := []byte{0x03, 0x04, 0x05}
mockGen1 := internal.NewMockByteGenerator(mockByte1)
mockGen2 := internal.NewMockByteGenerator(mockByte2)
tests := []struct {
name string
setupGenerator func() []TagJunkPacketGenerator
expected [][]byte
}{
{
name: "generate with no default junk",
setupGenerator: func() []TagJunkPacketGenerator {
tg1 := newTagJunkPacketGenerator("t1", "", 0)
tg1.append(mockGen1)
tg1.append(mockGen2)
tg2 := newTagJunkPacketGenerator("t2", "", 0)
tg2.append(mockGen2)
tg2.append(mockGen1)
return []TagJunkPacketGenerator{tg1, tg2}
},
expected: [][]byte{
append(mockByte1, mockByte2...),
append(mockByte2, mockByte1...),
},
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
generators := &TagJunkPacketGenerators{}
tagGenerators := tt.setupGenerator()
for _, gen := range tagGenerators {
generators.AppendGenerator(gen)
}
result := generators.GeneratePackets()
require.Equal(t, result, tt.expected)
})
}
}
-112
View File
@@ -1,112 +0,0 @@
package awg
import (
"fmt"
"maps"
"regexp"
"strings"
)
type IpcFields struct{ Key, Value string }
type EnumTag string
const (
BytesEnumTag EnumTag = "b"
CounterEnumTag EnumTag = "c"
TimestampEnumTag EnumTag = "t"
RandomBytesEnumTag EnumTag = "r"
WaitTimeoutEnumTag EnumTag = "wt"
WaitResponseEnumTag EnumTag = "wr"
)
var generatorCreator = map[EnumTag]newGenerator{
BytesEnumTag: newBytesGenerator,
CounterEnumTag: newPacketCounterGenerator,
TimestampEnumTag: newTimestampGenerator,
RandomBytesEnumTag: newRandomPacketGenerator,
WaitTimeoutEnumTag: newWaitTimeoutGenerator,
// WaitResponseEnumTag: newWaitResponseGenerator,
}
// helper map to determine enumTags are unique
var uniqueTags = map[EnumTag]bool{
CounterEnumTag: false,
TimestampEnumTag: false,
}
type Tag struct {
Name EnumTag
Param string
}
func parseTag(input string) (Tag, error) {
// Regular expression to match <tagname optional_param>
re := regexp.MustCompile(`([a-zA-Z]+)(?:\s+([^>]+))?>`)
match := re.FindStringSubmatch(input)
tag := Tag{
Name: EnumTag(match[1]),
}
if len(match) > 2 && match[2] != "" {
tag.Param = strings.TrimSpace(match[2])
}
return tag, nil
}
func Parse(name, input string) (TagJunkPacketGenerator, error) {
inputSlice := strings.Split(input, "<")
if len(inputSlice) <= 1 {
return TagJunkPacketGenerator{}, fmt.Errorf("empty input: %s", input)
}
uniqueTagCheck := make(map[EnumTag]bool, len(uniqueTags))
maps.Copy(uniqueTagCheck, uniqueTags)
// skip byproduct of split
inputSlice = inputSlice[1:]
rv := newTagJunkPacketGenerator(name, input, len(inputSlice))
for _, inputParam := range inputSlice {
if len(inputParam) <= 1 {
return TagJunkPacketGenerator{}, fmt.Errorf(
"empty tag in input: %s",
inputSlice,
)
} else if strings.Count(inputParam, ">") != 1 {
return TagJunkPacketGenerator{}, fmt.Errorf("ill formated input: %s", input)
}
tag, _ := parseTag(inputParam)
creator, ok := generatorCreator[tag.Name]
if !ok {
return TagJunkPacketGenerator{}, fmt.Errorf("invalid tag: %s", tag.Name)
}
if present, ok := uniqueTagCheck[tag.Name]; ok {
if present {
return TagJunkPacketGenerator{}, fmt.Errorf(
"tag %s needs to be unique",
tag.Name,
)
}
uniqueTagCheck[tag.Name] = true
}
generator, err := creator(tag.Param)
if err != nil {
return TagJunkPacketGenerator{}, fmt.Errorf("gen: %w", err)
}
// TODO: handle counter tag
// if tag.Name == CounterEnumTag {
// packetCounter, ok := generator.(*PacketCounterGenerator)
// if !ok {
// log.Fatalf("packet counter generator expected, got %T", generator)
// }
// PacketCounter = packetCounter.counter
// }
rv.append(generator)
}
return rv, nil
}
-77
View File
@@ -1,77 +0,0 @@
package awg
import (
"fmt"
"testing"
"github.com/stretchr/testify/require"
)
func TestParse(t *testing.T) {
type args struct {
name string
input string
}
tests := []struct {
name string
args args
wantErr error
}{
{
name: "invalid name",
args: args{name: "apple", input: ""},
wantErr: fmt.Errorf("ill formated input"),
},
{
name: "empty",
args: args{name: "i1", input: ""},
wantErr: fmt.Errorf("ill formated input"),
},
{
name: "extra >",
args: args{name: "i1", input: "<b 0xf6ab3267fa><c>>"},
wantErr: fmt.Errorf("ill formated input"),
},
{
name: "extra <",
args: args{name: "i1", input: "<<b 0xf6ab3267fa><c>"},
wantErr: fmt.Errorf("empty tag in input"),
},
{
name: "empty <>",
args: args{name: "i1", input: "<><b 0xf6ab3267fa><c>"},
wantErr: fmt.Errorf("empty tag in input"),
},
{
name: "invalid tag",
args: args{name: "i1", input: "<q 0xf6ab3267fa>"},
wantErr: fmt.Errorf("invalid tag"),
},
{
name: "counter uniqueness violation",
args: args{name: "i1", input: "<c><c>"},
wantErr: fmt.Errorf("parse tag needs to be unique"),
},
{
name: "timestamp uniqueness violation",
args: args{name: "i1", input: "<t><t>"},
wantErr: fmt.Errorf("parse tag needs to be unique"),
},
{
name: "valid",
args: args{input: "<b 0xf6ab3267fa><c><b 0xf6ab><t><r 10><wt 10>"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := Parse(tt.args.name, tt.args.input)
// TODO: ErrorAs doesn't work as you think
if tt.wantErr != nil {
require.ErrorAs(t, err, &tt.wantErr)
return
}
require.Nil(t, err)
})
}
}
+1 -1
View File
@@ -8,7 +8,7 @@ package device
import (
"errors"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
)
type DummyDatagram struct {
+2 -1
View File
@@ -118,6 +118,7 @@ func (st *CookieChecker) CreateReply(
msg []byte,
recv uint32,
src []byte,
msgType uint32,
) (*MessageCookieReply, error) {
st.RLock()
@@ -153,7 +154,7 @@ func (st *CookieChecker) CreateReply(
smac1 := smac2 - blake2s.Size128
reply := new(MessageCookieReply)
reply.Type = MessageCookieReplyType
reply.Type = msgType
reply.Receiver = recv
_, err := rand.Read(reply.Nonce[:])
+1 -1
View File
@@ -99,7 +99,7 @@ func TestCookieMAC1(t *testing.T) {
0x8c, 0xe1, 0xe8, 0xfa, 0x67, 0x20, 0x80, 0x6d,
}
generator.AddMacs(msg)
reply, err := checker.CreateReply(msg, 1377, src)
reply, err := checker.CreateReply(msg, 1377, src, MessageCookieReplyType)
if err != nil {
t.Fatal("Failed to create cookie reply:", err)
}
+55 -306
View File
@@ -6,55 +6,17 @@
package device
import (
"errors"
"runtime"
"sync"
"sync/atomic"
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device/awg"
"github.com/amnezia-vpn/amneziawg-go/ipc"
"github.com/amnezia-vpn/amneziawg-go/ratelimiter"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/ratelimiter"
"github.com/amnezia-vpn/amneziawg-go/v3/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
)
type Version uint8
const (
VersionDefault Version = iota
VersionAwg
VersionAwgSpecialHandshake
)
// TODO:
type AtomicVersion struct {
value atomic.Uint32
}
func NewAtomicVersion(v Version) *AtomicVersion {
av := &AtomicVersion{}
av.Store(v)
return av
}
func (av *AtomicVersion) Load() Version {
return Version(av.value.Load())
}
func (av *AtomicVersion) Store(v Version) {
av.value.Store(uint32(v))
}
func (av *AtomicVersion) CompareAndSwap(old, new Version) bool {
return av.value.CompareAndSwap(uint32(old), uint32(new))
}
func (av *AtomicVersion) Swap(new Version) Version {
return Version(av.value.Swap(uint32(new)))
}
type Device struct {
state struct {
// state holds the device's state. It is accessed atomically.
@@ -128,8 +90,42 @@ type Device struct {
closed chan struct{}
log *Logger
version Version
awg awg.Protocol
junk struct {
min atomic.Uint32
max atomic.Uint32
count atomic.Uint32
}
headers struct {
init AtomicUintRange
cookie AtomicUintRange
response AtomicUintRange
transport AtomicUintRange
}
paddings struct {
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.
@@ -224,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()
}
}
@@ -324,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{})
@@ -340,6 +338,15 @@ func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger) *Device {
device.rate.limiter.Init()
device.indexTable.Init()
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()
// create queues
@@ -437,8 +444,6 @@ func (device *Device) Close() {
device.rate.limiter.Close()
device.resetProtocol()
device.log.Verbosef("Device closed")
close(device.closed)
}
@@ -452,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()
@@ -578,261 +585,3 @@ func (device *Device) BindClose() error {
device.net.Unlock()
return err
}
func (device *Device) isAWG() bool {
return device.version >= VersionAwg
}
func (device *Device) resetProtocol() {
// restore default message type values
MessageInitiationType = DefaultMessageInitiationType
MessageResponseType = DefaultMessageResponseType
MessageCookieReplyType = DefaultMessageCookieReplyType
MessageTransportType = DefaultMessageTransportType
}
func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
if !tempAwg.ASecCfg.IsSet && !tempAwg.HandshakeHandler.IsSet {
return nil
}
var errs []error
isASecOn := false
device.awg.ASecMux.Lock()
if tempAwg.ASecCfg.JunkPacketCount < 0 {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"JunkPacketCount should be non negative",
),
)
}
device.awg.ASecCfg.JunkPacketCount = tempAwg.ASecCfg.JunkPacketCount
if tempAwg.ASecCfg.JunkPacketCount != 0 {
isASecOn = true
}
device.awg.ASecCfg.JunkPacketMinSize = tempAwg.ASecCfg.JunkPacketMinSize
if tempAwg.ASecCfg.JunkPacketMinSize != 0 {
isASecOn = true
}
if device.awg.ASecCfg.JunkPacketCount > 0 &&
tempAwg.ASecCfg.JunkPacketMaxSize == tempAwg.ASecCfg.JunkPacketMinSize {
tempAwg.ASecCfg.JunkPacketMaxSize++ // to make rand gen work
}
if tempAwg.ASecCfg.JunkPacketMaxSize >= MaxSegmentSize {
device.awg.ASecCfg.JunkPacketMinSize = 0
device.awg.ASecCfg.JunkPacketMaxSize = 1
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
tempAwg.ASecCfg.JunkPacketMaxSize,
MaxSegmentSize,
))
} else if tempAwg.ASecCfg.JunkPacketMaxSize < tempAwg.ASecCfg.JunkPacketMinSize {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"maxSize: %d; should be greater than minSize: %d",
tempAwg.ASecCfg.JunkPacketMaxSize,
tempAwg.ASecCfg.JunkPacketMinSize,
))
} else {
device.awg.ASecCfg.JunkPacketMaxSize = tempAwg.ASecCfg.JunkPacketMaxSize
}
if tempAwg.ASecCfg.JunkPacketMaxSize != 0 {
isASecOn = true
}
newInitSize := MessageInitiationSize + tempAwg.ASecCfg.InitHeaderJunkSize
if newInitSize >= MaxSegmentSize {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
`init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
tempAwg.ASecCfg.InitHeaderJunkSize,
MaxSegmentSize,
),
)
} else {
device.awg.ASecCfg.InitHeaderJunkSize = tempAwg.ASecCfg.InitHeaderJunkSize
}
if tempAwg.ASecCfg.InitHeaderJunkSize != 0 {
isASecOn = true
}
newResponseSize := MessageResponseSize + tempAwg.ASecCfg.ResponseHeaderJunkSize
if newResponseSize >= MaxSegmentSize {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
`response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
tempAwg.ASecCfg.ResponseHeaderJunkSize,
MaxSegmentSize,
),
)
} else {
device.awg.ASecCfg.ResponseHeaderJunkSize = tempAwg.ASecCfg.ResponseHeaderJunkSize
}
if tempAwg.ASecCfg.ResponseHeaderJunkSize != 0 {
isASecOn = true
}
newCookieSize := MessageCookieReplySize + tempAwg.ASecCfg.CookieReplyHeaderJunkSize
if newCookieSize >= MaxSegmentSize {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
`cookie reply size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
tempAwg.ASecCfg.CookieReplyHeaderJunkSize,
MaxSegmentSize,
),
)
} else {
device.awg.ASecCfg.CookieReplyHeaderJunkSize = tempAwg.ASecCfg.CookieReplyHeaderJunkSize
}
if tempAwg.ASecCfg.CookieReplyHeaderJunkSize != 0 {
isASecOn = true
}
newTransportSize := MessageTransportSize + tempAwg.ASecCfg.TransportHeaderJunkSize
if newTransportSize >= MaxSegmentSize {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
`transport size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
tempAwg.ASecCfg.TransportHeaderJunkSize,
MaxSegmentSize,
),
)
} else {
device.awg.ASecCfg.TransportHeaderJunkSize = tempAwg.ASecCfg.TransportHeaderJunkSize
}
if tempAwg.ASecCfg.TransportHeaderJunkSize != 0 {
isASecOn = true
}
if tempAwg.ASecCfg.InitPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating init_packet_magic_header")
device.awg.ASecCfg.InitPacketMagicHeader = tempAwg.ASecCfg.InitPacketMagicHeader
MessageInitiationType = device.awg.ASecCfg.InitPacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default init type")
MessageInitiationType = DefaultMessageInitiationType
}
if tempAwg.ASecCfg.ResponsePacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating response_packet_magic_header")
device.awg.ASecCfg.ResponsePacketMagicHeader = tempAwg.ASecCfg.ResponsePacketMagicHeader
MessageResponseType = device.awg.ASecCfg.ResponsePacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default response type")
MessageResponseType = DefaultMessageResponseType
}
if tempAwg.ASecCfg.UnderloadPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating underload_packet_magic_header")
device.awg.ASecCfg.UnderloadPacketMagicHeader = tempAwg.ASecCfg.UnderloadPacketMagicHeader
MessageCookieReplyType = device.awg.ASecCfg.UnderloadPacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default underload type")
MessageCookieReplyType = DefaultMessageCookieReplyType
}
if tempAwg.ASecCfg.TransportPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating transport_packet_magic_header")
device.awg.ASecCfg.TransportPacketMagicHeader = tempAwg.ASecCfg.TransportPacketMagicHeader
MessageTransportType = device.awg.ASecCfg.TransportPacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default transport type")
MessageTransportType = DefaultMessageTransportType
}
isSameHeaderMap := map[uint32]struct{}{
MessageInitiationType: {},
MessageResponseType: {},
MessageCookieReplyType: {},
MessageTransportType: {},
}
// size will be different if same values
if len(isSameHeaderMap) != 4 {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
`magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d`,
MessageInitiationType,
MessageResponseType,
MessageCookieReplyType,
MessageTransportType,
),
)
}
isSameSizeMap := map[int]struct{}{
newInitSize: {},
newResponseSize: {},
newCookieSize: {},
newTransportSize: {},
}
if len(isSameSizeMap) != 4 {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
`new sizes should differ; init: %d; response: %d; cookie: %d; trans: %d`,
newInitSize,
newResponseSize,
newCookieSize,
newTransportSize,
),
)
} else {
msgTypeToJunkSize = map[uint32]int{
MessageInitiationType: device.awg.ASecCfg.InitHeaderJunkSize,
MessageResponseType: device.awg.ASecCfg.ResponseHeaderJunkSize,
MessageCookieReplyType: device.awg.ASecCfg.CookieReplyHeaderJunkSize,
MessageTransportType: device.awg.ASecCfg.TransportHeaderJunkSize,
}
packetSizeToMsgType = map[int]uint32{
newInitSize: MessageInitiationType,
newResponseSize: MessageResponseType,
newCookieSize: MessageCookieReplyType,
newTransportSize: MessageTransportType,
}
}
device.awg.IsASecOn.SetTo(isASecOn)
var err error
device.awg.JunkCreator, err = awg.NewJunkCreator(device.awg.ASecCfg)
if err != nil {
errs = append(errs, err)
}
if tempAwg.HandshakeHandler.IsSet {
if err := tempAwg.HandshakeHandler.Validate(); err != nil {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid, "handshake handler validate: %w", err))
} else {
device.awg.HandshakeHandler = tempAwg.HandshakeHandler
device.awg.HandshakeHandler.ControlledJunk.DefaultJunkCount = tempAwg.ASecCfg.JunkPacketCount
device.awg.HandshakeHandler.SpecialJunk.DefaultJunkCount = tempAwg.ASecCfg.JunkPacketCount
device.version = VersionAwgSpecialHandshake
}
} else {
device.version = VersionAwg
}
device.awg.ASecMux.Unlock()
return errors.Join(errs...)
}
+16 -18
View File
@@ -23,10 +23,10 @@ import (
"go.uber.org/atomic"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/conn/bindtest"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/conn/bindtest"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/tun/tuntest"
)
// uapiCfg returns a string that contains cfg formatted use with IpcSet.
@@ -232,14 +232,14 @@ func TestAWGDevicePing(t *testing.T) {
"jc", "5",
"jmin", "500",
"jmax", "1000",
"s1", "30",
"s2", "40",
"s3", "50",
"s4", "5",
"h1", "123456",
"h2", "67543",
"h3", "123123",
"h4", "32345",
"s1", "15",
"s2", "18",
"s3", "20",
"s4", "25",
"h1", "123456-123500",
"h2", "67543-67550",
"h3", "123123-123200",
"h4", "32345-32350",
)
t.Run("ping 1.0.0.1", func(t *testing.T) {
pair.Send(t, Ping, nil)
@@ -264,12 +264,10 @@ func TestAWGHandshakeDevicePing(t *testing.T) {
goroutineLeakCheck(t)
pair := genTestPair(t, true,
"i1", "<b 0xf6ab3267fa><c><b 0xf6ab><t><r 10><wt 10>",
"i2", "<b 0xf6ab3267fa><r 100>",
"j1", "<b 0xffffffff><c><b 0xf6ab><t><r 10>",
"j2", "<c><b 0xf6ab><t><wt 1000>",
"j3", "<t><b 0xf6ab><c><r 10>",
"itime", "60",
"i1", "<b 0xf6ab3267fa><c><b 0xf6ab><t><r 10>",
"i2", "<b 0xf6ab3267fa><c><b 0xf6ab><t><rc 10>",
"i3", "<b 0xf6ab3267fa><c><b 0xf6ab><t><rd 10>",
"i4", "<b 0xf6ab3267fa><r 100>",
// "jc", "1",
// "jmin", "500",
// "jmax", "1000",
+1 -1
View File
@@ -11,7 +11,7 @@ import (
"sync/atomic"
"time"
"github.com/amnezia-vpn/amneziawg-go/replay"
"github.com/amnezia-vpn/amneziawg-go/v3/replay"
)
/* Due to limitations in Go and /x/crypto there is currently
+38 -29
View File
@@ -6,16 +6,18 @@
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"
"github.com/amnezia-vpn/amneziawg-go/tai64n"
"github.com/amnezia-vpn/amneziawg-go/v3/tai64n"
)
type handshakeState int
@@ -53,17 +55,11 @@ const (
)
const (
DefaultMessageInitiationType uint32 = 1
DefaultMessageResponseType uint32 = 2
DefaultMessageCookieReplyType uint32 = 3
DefaultMessageTransportType uint32 = 4
)
var (
MessageInitiationType uint32 = DefaultMessageInitiationType
MessageResponseType uint32 = DefaultMessageResponseType
MessageCookieReplyType uint32 = DefaultMessageCookieReplyType
MessageTransportType uint32 = DefaultMessageTransportType
MessageUnknownType uint32 = 0
MessageInitiationType uint32 = 1
MessageResponseType uint32 = 2
MessageCookieReplyType uint32 = 3
MessageTransportType uint32 = 4
)
const (
@@ -82,11 +78,6 @@ const (
MessageTransportOffsetContent = 16
)
var (
packetSizeToMsgType map[int]uint32
msgTypeToJunkSize map[uint32]int
)
/* Type is an 8-bit field, followed by 3 nul bytes,
* by marshalling the messages in little-endian byteorder
* we can treat these as a 32-bit unsigned int (for now)
@@ -205,12 +196,12 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e
handshake.mixHash(handshake.remoteStatic[:])
device.awg.ASecMux.RLock()
msgType := device.headers.init.Load().PickOne()
msg := MessageInitiation{
Type: MessageInitiationType,
Type: msgType,
Ephemeral: handshake.localEphemeral.publicKey(),
}
device.awg.ASecMux.RUnlock()
handshake.mixKey(msg.Ephemeral[:])
handshake.mixHash(msg.Ephemeral[:])
@@ -264,12 +255,9 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer {
chainKey [blake2s.Size]byte
)
device.awg.ASecMux.RLock()
if msg.Type != MessageInitiationType {
device.awg.ASecMux.RUnlock()
return nil
}
device.awg.ASecMux.RUnlock()
device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock()
@@ -384,9 +372,7 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
var msg MessageResponse
device.awg.ASecMux.RLock()
msg.Type = MessageResponseType
device.awg.ASecMux.RUnlock()
msg.Type = device.headers.response.Load().PickOne()
msg.Sender = handshake.localIndex
msg.Receiver = handshake.remoteIndex
@@ -436,12 +422,9 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
device.awg.ASecMux.RLock()
if msg.Type != MessageResponseType {
device.awg.ASecMux.RUnlock()
return nil
}
device.awg.ASecMux.RUnlock()
// lookup handshake by receiver
@@ -645,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)))
}
+2 -2
View File
@@ -10,8 +10,8 @@ import (
"encoding/binary"
"testing"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/tun/tuntest"
)
func TestCurveWrappers(t *testing.T) {
+143
View File
@@ -0,0 +1,143 @@
package device
import (
"errors"
"fmt"
"strings"
)
type obfBuilder func(val string) (obf, error)
var obfBuilders = map[string]obfBuilder{
"b": newBytesObf,
"t": newTimestampObf,
"r": newRandObf,
"rc": newRandCharObf,
"rd": newRandDigitsObf,
"d": newDataObf,
"ds": newDataStringObf,
"dz": newDataSizeObf,
}
type obf interface {
Obfuscate(dst, src []byte)
Deobfuscate(dst, src []byte) bool
ObfuscatedLen(srcLen int) int
DeobfuscatedLen(srcLen int) int
}
type obfChain struct {
Spec string
obfs []obf
}
func newObfChain(spec string) (*obfChain, error) {
var (
obfs []obf
errs []error
)
remaining := spec[:]
for {
start := strings.IndexByte(remaining, '<')
if start == -1 {
break
}
end := strings.IndexByte(remaining[start:], '>')
if end == -1 {
return nil, errors.New("missing enclosing >")
}
end += start
tag := remaining[start+1 : end]
parts := strings.Fields(tag)
if len(parts) == 0 {
errs = append(errs, errors.New("empty tag"))
remaining = remaining[end+1:]
continue
}
key := parts[0]
builder, ok := obfBuilders[key]
if !ok {
errs = append(errs, fmt.Errorf("unknown tag <%s>", key))
remaining = remaining[end+1:]
continue
}
val := ""
if len(parts) > 1 {
val = parts[1]
}
o, err := builder(val)
if err != nil {
errs = append(errs, fmt.Errorf("failed to build <%s>: %w", key, err))
remaining = remaining[end+1:]
continue
}
obfs = append(obfs, o)
remaining = remaining[end+1:]
}
if len(errs) > 0 {
return nil, errors.Join(errs...)
}
if len(obfs) == 0 {
return nil, nil
}
return &obfChain{
Spec: spec,
obfs: obfs,
}, nil
}
func (c *obfChain) Obfuscate(dst, src []byte) {
written := 0
for _, o := range c.obfs {
obfLen := o.ObfuscatedLen(len(src))
o.Obfuscate(dst[written:written+obfLen], src)
written += obfLen
}
}
func (c *obfChain) Deobfuscate(dst, src []byte) bool {
dynamicLen := len(src) - c.ObfuscatedLen(0)
written, read := 0, 0
for _, o := range c.obfs {
deobfLen := o.DeobfuscatedLen(dynamicLen)
obfLen := o.ObfuscatedLen(deobfLen)
if !o.Deobfuscate(dst[written:written+deobfLen], src[read:read+obfLen]) {
return false
}
written += deobfLen
read += obfLen
}
return true
}
func (c *obfChain) ObfuscatedLen(n int) int {
total := 0
for _, o := range c.obfs {
total += o.ObfuscatedLen(n)
}
return total
}
func (c *obfChain) DeobfuscatedLen(n int) int {
dynamicLen := n - c.ObfuscatedLen(0)
total := 0
for _, o := range c.obfs {
total += o.DeobfuscatedLen(dynamicLen)
}
return total
}
+47
View File
@@ -0,0 +1,47 @@
package device
import (
"bytes"
"encoding/hex"
"errors"
"strings"
)
func newBytesObf(val string) (obf, error) {
val = strings.TrimPrefix(val, "0x")
if len(val) == 0 {
return nil, errors.New("empty argument")
}
if len(val)%2 != 0 {
return nil, errors.New("odd amount of symbols")
}
bytes, err := hex.DecodeString(val)
if err != nil {
return nil, err
}
return &bytesObf{data: bytes}, nil
}
type bytesObf struct {
data []byte
}
func (o *bytesObf) Obfuscate(dst, src []byte) {
copy(dst, o.data)
}
func (o *bytesObf) Deobfuscate(dst, src []byte) bool {
return bytes.Equal(o.data, src[:o.ObfuscatedLen(0)])
}
func (o *bytesObf) ObfuscatedLen(srcLen int) int {
return len(o.data)
}
func (o *bytesObf) DeobfuscatedLen(srcLen int) int {
return 0
}
+25
View File
@@ -0,0 +1,25 @@
package device
func newDataObf(val string) (obf, error) {
return &dataObf{}, nil
}
type dataObf struct {
}
func (obf *dataObf) Obfuscate(dst, src []byte) {
copy(dst, src)
}
func (obf *dataObf) Deobfuscate(dst, src []byte) bool {
copy(dst, src)
return true
}
func (o *dataObf) ObfuscatedLen(n int) int {
return n
}
func (o *dataObf) DeobfuscatedLen(n int) int {
return n
}
+38
View File
@@ -0,0 +1,38 @@
package device
import "strconv"
func newDataSizeObf(val string) (obf, error) {
length, err := strconv.Atoi(val)
if err != nil {
return nil, err
}
return &dataSizeObf{
length: length,
}, nil
}
type dataSizeObf struct {
length int
}
func (o *dataSizeObf) Obfuscate(dst, src []byte) {
srcLen := len(src)
for i := o.length - 1; i >= 0; i-- {
dst[i] = byte(srcLen & 0xFF)
srcLen >>= 8
}
}
func (o *dataSizeObf) Deobfuscate(dst, src []byte) bool {
return true
}
func (o *dataSizeObf) ObfuscatedLen(srcLen int) int {
return o.length
}
func (o *dataSizeObf) DeobfuscatedLen(srcLen int) int {
return 0
}
+29
View File
@@ -0,0 +1,29 @@
package device
import (
"encoding/base64"
)
func newDataStringObf(val string) (obf, error) {
return &dataStringObf{}, nil
}
type dataStringObf struct {
}
func (o *dataStringObf) Obfuscate(dst, src []byte) {
base64.RawStdEncoding.Encode(dst, src)
}
func (o *dataStringObf) Deobfuscate(dst, src []byte) bool {
base64.RawStdEncoding.Decode(dst, src)
return true
}
func (o *dataStringObf) ObfuscatedLen(n int) int {
return base64.RawStdEncoding.EncodedLen(n)
}
func (o *dataStringObf) DeobfuscatedLen(n int) int {
return base64.RawStdEncoding.DecodedLen(n)
}
+39
View File
@@ -0,0 +1,39 @@
package device
import (
"crypto/rand"
"strconv"
)
func newRandObf(val string) (obf, error) {
length, err := strconv.Atoi(val)
if err != nil {
return nil, err
}
return &randObf{
length: length,
}, nil
}
type randObf struct {
length int
}
func (o *randObf) Obfuscate(dst, src []byte) {
rand.Read(dst[:o.length])
}
func (o *randObf) Deobfuscate(dst, src []byte) bool {
// there is no way to validate randomness :)
// assume that it is always true
return true
}
func (o *randObf) ObfuscatedLen(n int) int {
return o.length
}
func (o *randObf) DeobfuscatedLen(n int) int {
return 0
}
+48
View File
@@ -0,0 +1,48 @@
package device
import (
"crypto/rand"
"strconv"
"unicode"
)
const chars52 = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
func newRandCharObf(val string) (obf, error) {
length, err := strconv.Atoi(val)
if err != nil {
return nil, err
}
return &randCharObf{
length: length,
}, nil
}
type randCharObf struct {
length int
}
func (o *randCharObf) Obfuscate(dst, src []byte) {
rand.Read(dst[:o.length])
for i := range dst[:o.length] {
dst[i] = chars52[dst[i]%52]
}
}
func (o *randCharObf) Deobfuscate(dst, src []byte) bool {
for _, b := range src[:o.length] {
if !unicode.IsLetter(rune(b)) {
return false
}
}
return true
}
func (o *randCharObf) ObfuscatedLen(n int) int {
return o.length
}
func (o *randCharObf) DeobfuscatedLen(n int) int {
return 0
}
+48
View File
@@ -0,0 +1,48 @@
package device
import (
"crypto/rand"
"strconv"
"unicode"
)
const digits10 = "0123456789"
func newRandDigitsObf(val string) (obf, error) {
length, err := strconv.Atoi(val)
if err != nil {
return nil, err
}
return &randDigitObf{
length: length,
}, nil
}
type randDigitObf struct {
length int
}
func (o *randDigitObf) Obfuscate(dst, src []byte) {
rand.Read(dst[:o.length])
for i := range dst[:o.length] {
dst[i] = digits10[dst[i]%10]
}
}
func (o *randDigitObf) Deobfuscate(dst, src []byte) bool {
for _, b := range src[:o.length] {
if !unicode.IsDigit(rune(b)) {
return false
}
}
return true
}
func (o *randDigitObf) ObfuscatedLen(n int) int {
return o.length
}
func (o *randDigitObf) DeobfuscatedLen(n int) int {
return 0
}
+31
View File
@@ -0,0 +1,31 @@
package device
import (
"encoding/binary"
"time"
)
func newTimestampObf(_ string) (obf, error) {
return &timestampObf{}, nil
}
type timestampObf struct{}
func (o *timestampObf) Obfuscate(dst, src []byte) {
t := uint32(time.Now().Unix())
binary.BigEndian.PutUint32(dst, t)
}
func (o *timestampObf) Deobfuscate(dst, src []byte) bool {
// replay attack check?
// requires time to be always synchronized
return true
}
func (o *timestampObf) ObfuscatedLen(n int) int {
return 4
}
func (o *timestampObf) DeobfuscatedLen(n int) int {
return 0
}
+5 -15
View File
@@ -12,8 +12,7 @@ import (
"sync/atomic"
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device/awg"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
)
type Peer struct {
@@ -40,6 +39,7 @@ type Peer struct {
zeroKeyMaterial *Timer
persistentKeepalive *Timer
handshakeAttempts atomic.Uint32
maxHandshakeAttempts atomic.Uint32
needAnotherKeepalive atomic.Bool
sentLastMinuteHandshake atomic.Bool
}
@@ -56,7 +56,7 @@ type Peer struct {
cookieGenerator CookieGenerator
trieEntries list.List
persistentKeepaliveInterval atomic.Uint32
persistentKeepaliveInterval AtomicUintRange
}
func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
@@ -114,16 +114,6 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
return peer, nil
}
func (peer *Peer) SendAndCountBuffers(buffers [][]byte) error {
err := peer.SendBuffers(buffers)
if err == nil {
awg.PacketCounter.Add(uint64(len(buffers)))
return nil
}
return err
}
func (peer *Peer) SendBuffers(buffers [][]byte) error {
peer.device.net.RLock()
defer peer.device.net.RUnlock()
@@ -203,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
@@ -253,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
+1 -1
View File
@@ -5,7 +5,7 @@
package device
import "github.com/amnezia-vpn/amneziawg-go/conn"
import "github.com/amnezia-vpn/amneziawg-go/v3/conn"
/* Reduce memory consumption for Android */
+1 -1
View File
@@ -7,7 +7,7 @@
package device
import "github.com/amnezia-vpn/amneziawg-go/conn"
import "github.com/amnezia-vpn/amneziawg-go/v3/conn"
const (
QueueStagedSize = conn.IdealBatchSize
+106 -49
View File
@@ -13,7 +13,7 @@ import (
"sync"
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
@@ -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,15 +97,16 @@ func (device *Device) RoutineReceiveIncoming(
endpoints = make([]conn.Endpoint, maxBatchSize)
deathSpiral int
elemsByPeer = make(map[*Peer]*QueueInboundElementsContainer, maxBatchSize)
typeHashBuf [4]byte
)
for i := range bufsArrs {
for i := range maxBatchSize {
bufsArrs[i] = device.GetMessageBuffer()
bufs[i] = bufsArrs[i][:]
}
defer func() {
for i := 0; i < maxBatchSize; i++ {
for i := range maxBatchSize {
if bufsArrs[i] != nil {
device.PutMessageBuffer(bufsArrs[i])
}
@@ -129,7 +132,6 @@ func (device *Device) RoutineReceiveIncoming(
}
deathSpiral = 0
device.awg.ASecMux.RLock()
// handle each packet in the batch
for i, size := range sizes[:count] {
if size < MinMessageSize {
@@ -138,42 +140,25 @@ func (device *Device) RoutineReceiveIncoming(
// check size of packet
packet := bufsArrs[i][:size]
var msgType uint32
if device.isAWG() {
// TODO:
// if awg.WaitResponse.ShouldWait.IsSet() {
// awg.WaitResponse.Channel <- struct{}{}
// }
if assumedMsgType, ok := packetSizeToMsgType[size]; ok {
junkSize := msgTypeToJunkSize[assumedMsgType]
// transport size can align with other header types;
// making sure we have the right msgType
msgType = binary.LittleEndian.Uint32(packet[junkSize : junkSize+4])
if msgType == assumedMsgType {
packet = packet[junkSize:]
} else {
device.log.Verbosef("transport packet lined up with another msg type")
msgType = binary.LittleEndian.Uint32(packet[:4])
}
} else {
transportJunkSize := device.awg.ASecCfg.TransportHeaderJunkSize
msgType = binary.LittleEndian.Uint32(packet[transportJunkSize : transportJunkSize+4])
if msgType != MessageTransportType {
// probably a junk packet
device.log.Verbosef("aSec: Received message with unknown type: %d", msgType)
continue
}
cip, err := device.HeaderProtectionCipher(packet[:HeaderCipherNonceSize])
if err != nil {
device.log.Errorf("Failed to initialize header cipher")
continue
}
// remove junk from bufsArrs by shifting the packet
// this buffer is also used for decryption, so it needs to be corrected
copy(bufsArrs[i][:size], packet[transportJunkSize:])
size -= transportJunkSize
// need to reinitialize packet as well
packet = packet[:size]
}
} else {
msgType = binary.LittleEndian.Uint32(packet[:4])
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, typeHash)
packet = packet[padding:]
if cip != nil {
applyHash(packet[:4], packet[:4], typeHash)
}
switch msgType {
@@ -187,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
@@ -201,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
}
@@ -213,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 {
@@ -231,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")
@@ -259,7 +257,6 @@ func (device *Device) RoutineReceiveIncoming(
default:
}
}
device.awg.ASecMux.RUnlock()
for peer, elemsContainer := range elemsByPeer {
if peer.isRunning.Load() {
peer.queue.inbound.c <- elemsContainer
@@ -317,9 +314,6 @@ func (device *Device) RoutineHandshake(id int) {
device.log.Verbosef("Routine: handshake worker %d - started", id)
for elem := range device.queue.handshake.c {
device.awg.ASecMux.RLock()
// handle cookie fields and ratelimiting
switch elem.msgType {
@@ -405,6 +399,9 @@ func (device *Device) RoutineHandshake(id int) {
goto skip
}
// have to reassign msgType for ranged msgType to work
msg.Type = elem.msgType
// consume initiation
peer := device.ConsumeMessageInitiation(&msg)
if peer == nil {
@@ -437,6 +434,9 @@ func (device *Device) RoutineHandshake(id int) {
goto skip
}
// have to reassign msgType for ranged msgType to work
msg.Type = elem.msgType
// consume response
peer := device.ConsumeMessageResponse(&msg)
@@ -470,7 +470,6 @@ func (device *Device) RoutineHandshake(id int) {
peer.SendKeepalive()
}
skip:
device.awg.ASecMux.RUnlock()
device.PutMessageBuffer(elem.buffer)
}
}
@@ -559,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)
@@ -589,3 +585,64 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
device.PutInboundElementsContainer(elemsContainer)
}
}
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.Load()
header := device.headers.init.Load()
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.Load()
header := device.headers.response.Load()
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.Load()
header := device.headers.cookie.Load()
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.Load()
header := device.headers.transport.Load()
if size >= int(padding)+MessageTransportHeaderSize {
applyHash(headerBytes[:], packet[padding:padding+4], typeHash)
if header.Contains(binary.LittleEndian.Uint32(headerBytes[:])) {
return MessageTransportType, padding
}
}
}
return MessageUnknownType, 0
}
+128 -90
View File
@@ -7,15 +7,17 @@ package device
import (
"bytes"
"crypto/rand"
"encoding/binary"
"errors"
"net"
"os"
"slices"
"sync"
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
@@ -51,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 {
@@ -62,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
}
@@ -99,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
}
@@ -123,66 +130,43 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
peer.device.log.Errorf("%v - Failed to create initiation message: %v", peer, err)
return err
}
var sendBuffer [][]byte
// so only packet processed for cookie generation
var junkedHeader []byte
if peer.device.version >= VersionAwg {
var junks [][]byte
if peer.device.version == VersionAwgSpecialHandshake {
peer.device.awg.ASecMux.RLock()
// set junks depending on packet type
junks = peer.device.awg.HandshakeHandler.GenerateSpecialJunk()
if junks == nil {
junks = peer.device.awg.HandshakeHandler.GenerateControlledJunk()
if junks != nil {
peer.device.log.Verbosef("%v - Controlled junks sent", peer)
}
} else {
peer.device.log.Verbosef("%v - Special junks sent", peer)
}
peer.device.awg.ASecMux.RUnlock()
} else {
junks = make([][]byte, 0, peer.device.awg.ASecCfg.JunkPacketCount)
}
peer.device.awg.ASecMux.RLock()
err := peer.device.awg.JunkCreator.CreateJunkPackets(&junks)
peer.device.awg.ASecMux.RUnlock()
if err != nil {
peer.device.log.Errorf("%v - %v", peer, err)
return err
}
if len(junks) > 0 {
err = peer.SendBuffers(junks)
if err != nil {
peer.device.log.Errorf("%v - Failed to send junk packets: %v", peer, err)
return err
}
}
junkedHeader, err = peer.device.awg.CreateInitHeaderJunk()
if err != nil {
peer.device.log.Errorf("%v - %v", peer, err)
return err
for _, ipacket := range peer.device.ipackets {
if ipacket != nil {
buf := make([]byte, ipacket.ObfuscatedLen(0))
ipacket.Obfuscate(buf, nil)
sendBuffer = append(sendBuffer, buf)
}
}
var buf [MessageInitiationSize]byte
writer := bytes.NewBuffer(buf[:0])
sendBuffer = append(sendBuffer, peer.device.JunkPackets()...)
padding := peer.device.paddings.init.Load()
buf := make([]byte, padding+MessageInitiationSize)
crypt := buf[:padding]
rand.Read(crypt)
writer := bytes.NewBuffer(buf[padding:padding])
binary.Write(writer, binary.LittleEndian, msg)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
junkedHeader = append(junkedHeader, packet...)
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
sendBuffer = append(sendBuffer, junkedHeader)
cip, err := peer.device.HeaderProtectionCipher(crypt[:HeaderCipherNonceSize])
if err != nil {
return err
}
if cip != nil {
cip.XORKeyStream(packet, packet)
}
err = peer.SendAndCountBuffers(sendBuffer)
sendBuffer = append(sendBuffer, buf)
err = peer.SendBuffers(sendBuffer)
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err)
}
@@ -204,19 +188,16 @@ func (peer *Peer) SendHandshakeResponse() error {
return err
}
junkedHeader, err := peer.device.awg.CreateResponseHeaderJunk()
if err != nil {
peer.device.log.Errorf("%v - %v", peer, err)
return err
}
padding := peer.device.paddings.response.Load()
buf := make([]byte, padding+MessageResponseSize)
var buf [MessageResponseSize]byte
writer := bytes.NewBuffer(buf[:0])
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)
junkedHeader = append(junkedHeader, packet...)
err = peer.BeginSymmetricSession()
if err != nil {
@@ -228,43 +209,59 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
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.SendAndCountBuffers([][]byte{junkedHeader})
err = peer.SendBuffers([][]byte{buf})
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err)
}
return err
}
func (device *Device) SendHandshakeCookie(
initiatingElem *QueueHandshakeElement,
) error {
func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) error {
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.Load().PickOne()
reply, err := device.cookieChecker.CreateReply(
initiatingElem.packet,
sender,
initiatingElem.endpoint.DstToBytes(),
msgType,
)
if err != nil {
device.log.Errorf("Failed to create cookie reply: %v", err)
return err
}
junkedHeader, err := device.awg.CreateCookieReplyHeaderJunk()
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()
cip, err := device.HeaderProtectionCipher(crypt[:HeaderCipherNonceSize])
if err != nil {
device.log.Errorf("%v - %v", device, err)
return err
}
if cip != nil {
cip.XORKeyStream(packet, packet)
}
var buf [MessageCookieReplySize]byte
writer := bytes.NewBuffer(buf[:0])
binary.Write(writer, binary.LittleEndian, reply)
junkedHeader = append(junkedHeader, writer.Bytes()...)
// TODO: allocation could be avoided
device.net.bind.Send([][]byte{junkedHeader}, initiatingElem.endpoint)
device.net.bind.Send([][]byte{buf}, initiatingElem.endpoint)
return nil
}
@@ -274,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)
}
}
@@ -296,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 {
@@ -314,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++ {
@@ -323,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
@@ -417,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
}
@@ -507,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)
@@ -521,30 +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]
binary.LittleEndian.PutUint32(fieldType, MessageTransportType)
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()
}
@@ -585,22 +630,15 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
for _, elem := range elemsContainer.elems {
if len(elem.packet) != MessageKeepaliveSize {
dataSent = true
junkedHeader, err := device.awg.CreateTransportHeaderJunk(len(elem.packet))
if err != nil {
device.log.Errorf("%v - %v", device, err)
continue
}
elem.packet = append(junkedHeader, elem.packet...)
}
bufs = append(bufs, elem.packet)
}
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
err := peer.SendAndCountBuffers(bufs)
err := peer.SendBuffers(bufs)
if dataSent {
peer.timersDataSent()
}
+2 -2
View File
@@ -3,8 +3,8 @@
package device
import (
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/rwcancel"
)
func (device *Device) startRouteListener(_ conn.Bind) (*rwcancel.RWCancel, error) {
+2 -2
View File
@@ -20,8 +20,8 @@ import (
"golang.org/x/sys/unix"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/rwcancel"
)
func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) {
+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
}
+1 -1
View File
@@ -8,7 +8,7 @@ package device
import (
"fmt"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
)
const DefaultMTU = 1420
+304 -148
View File
@@ -18,8 +18,7 @@ import (
"sync"
"time"
"github.com/amnezia-vpn/amneziawg-go/device/awg"
"github.com/amnezia-vpn/amneziawg-go/ipc"
"github.com/amnezia-vpn/amneziawg-go/v3/ipc"
)
type IPCError struct {
@@ -84,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() {
@@ -98,54 +100,80 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("fwmark=%d", device.net.fwmark)
}
if device.isAWG() {
if device.awg.ASecCfg.JunkPacketCount != 0 {
sendf("jc=%d", device.awg.ASecCfg.JunkPacketCount)
}
if device.awg.ASecCfg.JunkPacketMinSize != 0 {
sendf("jmin=%d", device.awg.ASecCfg.JunkPacketMinSize)
}
if device.awg.ASecCfg.JunkPacketMaxSize != 0 {
sendf("jmax=%d", device.awg.ASecCfg.JunkPacketMaxSize)
}
if device.awg.ASecCfg.InitHeaderJunkSize != 0 {
sendf("s1=%d", device.awg.ASecCfg.InitHeaderJunkSize)
}
if device.awg.ASecCfg.ResponseHeaderJunkSize != 0 {
sendf("s2=%d", device.awg.ASecCfg.ResponseHeaderJunkSize)
}
if device.awg.ASecCfg.CookieReplyHeaderJunkSize != 0 {
sendf("s3=%d", device.awg.ASecCfg.CookieReplyHeaderJunkSize)
}
if device.awg.ASecCfg.TransportHeaderJunkSize != 0 {
sendf("s4=%d", device.awg.ASecCfg.TransportHeaderJunkSize)
}
if device.awg.ASecCfg.InitPacketMagicHeader != 0 {
sendf("h1=%d", device.awg.ASecCfg.InitPacketMagicHeader)
}
if device.awg.ASecCfg.ResponsePacketMagicHeader != 0 {
sendf("h2=%d", device.awg.ASecCfg.ResponsePacketMagicHeader)
}
if device.awg.ASecCfg.UnderloadPacketMagicHeader != 0 {
sendf("h3=%d", device.awg.ASecCfg.UnderloadPacketMagicHeader)
}
if device.awg.ASecCfg.TransportPacketMagicHeader != 0 {
sendf("h4=%d", device.awg.ASecCfg.TransportPacketMagicHeader)
}
if count := device.junk.count.Load(); count != 0 {
sendf("jc=%d", count)
}
specialJunkIpcFields := device.awg.HandshakeHandler.SpecialJunk.IpcGetFields()
for _, field := range specialJunkIpcFields {
sendf("%s=%s", field.Key, field.Value)
}
controlledJunkIpcFields := device.awg.HandshakeHandler.ControlledJunk.IpcGetFields()
for _, field := range controlledJunkIpcFields {
sendf("%s=%s", field.Key, field.Value)
}
if device.awg.HandshakeHandler.ITimeout != 0 {
sendf("itime=%d", device.awg.HandshakeHandler.ITimeout/time.Second)
if min := device.junk.min.Load(); min != 0 {
sendf("jmin=%d", min)
}
if max := device.junk.max.Load(); max != 0 {
sendf("jmax=%d", max)
}
if padding := device.paddings.init.Load(); padding != 0 {
sendf("s1=%d", padding)
}
if padding := device.paddings.response.Load(); padding != 0 {
sendf("s2=%d", padding)
}
if padding := device.paddings.cookie.Load(); padding != 0 {
sendf("s3=%d", padding)
}
if padding := device.paddings.transport.Load(); padding != 0 {
sendf("s4=%d", padding)
}
if header := device.headers.init.Load(); !header.IsZero() {
sendf("h1=%s", header.ToString())
}
if header := device.headers.response.Load(); !header.IsZero() {
sendf("h2=%s", header.ToString())
}
if header := device.headers.cookie.Load(); !header.IsZero() {
sendf("h3=%s", header.ToString())
}
if header := device.headers.transport.Load(); !header.IsZero() {
sendf("h4=%s", header.ToString())
}
for i, ipacket := range device.ipackets {
if ipacket != nil {
sendf("i%d=%s", i+1, ipacket.Spec)
}
}
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()
@@ -167,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())
@@ -196,18 +227,19 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
}
}()
ipcDev := new(ipcSetDevice)
ipcDev.fromDevice(device)
peer := new(ipcSetPeer)
deviceConfig := true
tempAwg := awg.Protocol{}
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := scanner.Text()
if line == "" {
// Blank line means terminate operation.
err := device.handlePostConfig(&tempAwg)
err := ipcDev.mergeWithDevice(device)
if err != nil {
return err
return ipcErrorf(ipc.IpcErrorInvalid, "failed to merge with device: %w", err)
}
peer.handlePostConfig()
return nil
@@ -236,7 +268,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
var err error
if deviceConfig {
err = device.handleDeviceLine(key, value, &tempAwg)
err = device.handleDeviceLine(ipcDev, key, value)
} else {
err = device.handlePeerLine(peer, key, value)
}
@@ -244,9 +276,9 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return err
}
}
err = device.handlePostConfig(&tempAwg)
err = ipcDev.mergeWithDevice(device)
if err != nil {
return err
return ipcErrorf(ipc.IpcErrorInvalid, "failed to merge with device: %w", err)
}
peer.handlePostConfig()
@@ -256,7 +288,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return nil
}
func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol) error {
func (device *Device) handleDeviceLine(ipcDev *ipcSetDevice, key, value string) error {
switch key {
case "private_key":
var sk NoisePrivateKey
@@ -307,140 +339,180 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
device.RemoveAllPeers()
case "jc":
junkPacketCount, err := strconv.Atoi(value)
jc, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_count %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jc: %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_count")
tempAwg.ASecCfg.JunkPacketCount = junkPacketCount
tempAwg.ASecCfg.IsSet = true
device.log.Verbosef("UAPI: Updating junk count")
device.junk.count.Store(uint32(jc))
case "jmin":
junkPacketMinSize, err := strconv.Atoi(value)
jmin, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_min_size %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmin: %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_min_size")
tempAwg.ASecCfg.JunkPacketMinSize = junkPacketMinSize
tempAwg.ASecCfg.IsSet = true
device.log.Verbosef("UAPI: Updating junk min")
device.junk.min.Store(uint32(jmin))
case "jmax":
junkPacketMaxSize, err := strconv.Atoi(value)
jmax, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_max_size %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmax: %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_max_size")
tempAwg.ASecCfg.JunkPacketMaxSize = junkPacketMaxSize
tempAwg.ASecCfg.IsSet = true
device.log.Verbosef("UAPI: Updating junk max")
device.junk.max.Store(uint32(jmax))
case "s1":
initPacketJunkSize, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_junk_size %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s1: %w", err)
}
device.log.Verbosef("UAPI: Updating init_packet_junk_size")
tempAwg.ASecCfg.InitHeaderJunkSize = initPacketJunkSize
tempAwg.ASecCfg.IsSet = true
ipcDev.paddings.init = uint32(padding)
case "s2":
responsePacketJunkSize, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_junk_size %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s2: %w", err)
}
device.log.Verbosef("UAPI: Updating response_packet_junk_size")
tempAwg.ASecCfg.ResponseHeaderJunkSize = responsePacketJunkSize
tempAwg.ASecCfg.IsSet = true
ipcDev.paddings.response = uint32(padding)
case "s3":
cookieReplyPacketJunkSize, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse cookie_reply_packet_junk_size %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s3: %w", err)
}
device.log.Verbosef("UAPI: Updating cookie_reply_packet_junk_size")
tempAwg.ASecCfg.CookieReplyHeaderJunkSize = cookieReplyPacketJunkSize
tempAwg.ASecCfg.IsSet = true
ipcDev.paddings.cookie = uint32(padding)
case "s4":
transportPacketJunkSize, err := strconv.Atoi(value)
padding, err := strconv.ParseUint(value, 10, 16)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_junk_size %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s4: %w", err)
}
device.log.Verbosef("UAPI: Updating transport_packet_junk_size")
tempAwg.ASecCfg.TransportHeaderJunkSize = transportPacketJunkSize
tempAwg.ASecCfg.IsSet = true
ipcDev.paddings.transport = uint32(padding)
case "h1":
initPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_magic_header %w", err)
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H1: %w", err)
}
tempAwg.ASecCfg.InitPacketMagicHeader = uint32(initPacketMagicHeader)
tempAwg.ASecCfg.IsSet = true
ipcDev.headers.init = rang
case "h2":
responsePacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_magic_header %w", err)
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H2: %w", err)
}
tempAwg.ASecCfg.ResponsePacketMagicHeader = uint32(responsePacketMagicHeader)
tempAwg.ASecCfg.IsSet = true
ipcDev.headers.response = rang
case "h3":
underloadPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse underload_packet_magic_header %w", err)
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H3: %w", err)
}
tempAwg.ASecCfg.UnderloadPacketMagicHeader = uint32(underloadPacketMagicHeader)
tempAwg.ASecCfg.IsSet = true
ipcDev.headers.cookie = rang
case "h4":
transportPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
var rang UintRange
if err := rang.FromString(value); err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H4: %w", err)
}
ipcDev.headers.transport = rang
case "i1":
chain, err := newObfChain(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_magic_header %w", err)
}
tempAwg.ASecCfg.TransportPacketMagicHeader = uint32(transportPacketMagicHeader)
tempAwg.ASecCfg.IsSet = true
case "i1", "i2", "i3", "i4", "i5":
if len(value) == 0 {
device.log.Verbosef("UAPI: received empty %s", key)
return nil
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I1: %w", err)
}
device.ipackets[0] = chain
generators, err := awg.Parse(key, value)
case "i2":
chain, err := newObfChain(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
}
device.log.Verbosef("UAPI: Updating %s", key)
tempAwg.HandshakeHandler.SpecialJunk.AppendGenerator(generators)
tempAwg.HandshakeHandler.IsSet = true
case "j1", "j2", "j3":
if len(value) == 0 {
device.log.Verbosef("UAPI: received empty %s", key)
return nil
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I2: %w", err)
}
device.ipackets[1] = chain
generators, err := awg.Parse(key, value)
case "i3":
chain, err := newObfChain(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I3: %w", err)
}
device.log.Verbosef("UAPI: Updating %s", key)
device.ipackets[2] = chain
tempAwg.HandshakeHandler.ControlledJunk.AppendGenerator(generators)
tempAwg.HandshakeHandler.IsSet = true
case "itime":
if len(value) == 0 {
device.log.Verbosef("UAPI: received empty itime")
return nil
}
itime, err := strconv.ParseInt(value, 10, 64)
case "i4":
chain, err := newObfChain(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "parse itime %w", err)
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I4: %w", err)
}
device.log.Verbosef("UAPI: Updating itime")
device.ipackets[3] = chain
case "i5":
chain, err := newObfChain(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I5: %w", err)
}
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)
tempAwg.HandshakeHandler.ITimeout = time.Duration(itime) * time.Second
tempAwg.HandshakeHandler.IsSet = true
default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
@@ -561,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)
@@ -689,3 +757,91 @@ func (device *Device) IpcHandle(socket net.Conn) {
buffered.Flush()
}
}
type ipcSetDevice struct {
headers struct {
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 {
device.headerProtection.Lock()
defer device.headerProtection.Unlock()
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.Overlap(right) {
return errors.New("headers must not overlap")
}
}
}
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
}
+23 -11
View File
@@ -1,23 +1,35 @@
module github.com/amnezia-vpn/amneziawg-go
module github.com/amnezia-vpn/amneziawg-go/v3
go 1.24.4
go 1.25.0
require (
github.com/stretchr/testify v1.10.0
github.com/tevino/abool v1.2.0
github.com/tevino/abool/v2 v2.1.0
github.com/goccy/go-yaml v1.17.1
go.uber.org/atomic v1.11.0
golang.org/x/crypto v0.39.0
golang.org/x/net v0.41.0
golang.org/x/sys v0.33.0
golang.getoutline.org/sdk v0.0.23
golang.getoutline.org/sdk/x v0.2.0
golang.org/x/crypto v0.42.0
golang.org/x/net v0.44.0
golang.org/x/sys v0.36.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489
)
require (
github.com/davecgh/go-spew v1.1.1 // indirect
github.com/go-task/slim-sprig v0.0.0-20230315185526-52ccab3ef572 // indirect
github.com/google/btree v1.1.3 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/google/pprof v0.0.0-20211214055906-6f57359322fd // indirect
github.com/gorilla/websocket v1.5.3 // indirect
github.com/onsi/ginkgo/v2 v2.12.0 // indirect
github.com/quic-go/qpack v0.5.1 // indirect
github.com/quic-go/quic-go v0.48.1 // indirect
github.com/shadowsocks/go-shadowsocks2 v0.1.5 // indirect
github.com/stretchr/testify v1.10.0 // indirect
go.uber.org/mock v0.4.0 // indirect
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 // indirect
golang.org/x/mobile v0.0.0-20240520174638-fa72addaaa1b // indirect
golang.org/x/mod v0.28.0 // indirect
golang.org/x/sync v0.17.0 // indirect
golang.org/x/text v0.29.0 // indirect
golang.org/x/time v0.9.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
golang.org/x/tools v0.37.0 // indirect
)
+74 -19
View File
@@ -1,40 +1,95 @@
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/go-logr/logr v1.2.4 h1:g01GSCwiDw2xSZfjJ2/T9M+S6pFdcNtFYsp+Y43HYDQ=
github.com/go-logr/logr v1.2.4/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
github.com/go-task/slim-sprig v0.0.0-20230315185526-52ccab3ef572 h1:tfuBGBXKqDEevZMzYi5KSi8KkcZtzBcTgAUUtapy0OI=
github.com/go-task/slim-sprig v0.0.0-20230315185526-52ccab3ef572/go.mod h1:9Pwr4B2jHnOSGXyyzV8ROjYa2ojvAY6HCGYYfMoC3Ls=
github.com/goccy/go-yaml v1.17.1 h1:LI34wktB2xEE3ONG/2Ar54+/HJVBriAGJ55PHls4YuY=
github.com/goccy/go-yaml v1.17.1/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA=
github.com/golang/protobuf v1.5.3 h1:KhyjKVUg7Usr/dYsdSqoFveMYd5ko72D+zANwlG1mmg=
github.com/golang/protobuf v1.5.3/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY=
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38=
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8=
github.com/google/gopacket v1.1.19/go.mod h1:iJ8V8n6KS+z2U1A8pUwu8bW5SyEMkXJB8Yo/Vo+TKTo=
github.com/google/pprof v0.0.0-20211214055906-6f57359322fd h1:1FjCyPC+syAzJ5/2S8fqdZK1R22vvA0J7JZKcuOIQ7Y=
github.com/google/pprof v0.0.0-20211214055906-6f57359322fd/go.mod h1:KgnwoLYCZ8IQu3XUZ8Nc/bM9CCZFOyjUNOSygVozoDg=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/ianlancetaylor/demangle v0.0.0-20210905161508-09a460cdf81d/go.mod h1:aYm2/VgdVmcIU8iMfdMvDMsRAQjcfZSKFby6HOFvi/w=
github.com/onsi/ginkgo/v2 v2.12.0 h1:UIVDowFPwpg6yMUpPjGkYvf06K3RAiJXUhCxEwQVHRI=
github.com/onsi/ginkgo/v2 v2.12.0/go.mod h1:ZNEzXISYlqpb8S36iN71ifqLi3vVD1rVJGvWRCJOUpQ=
github.com/onsi/gomega v1.27.10 h1:naR28SdDFlqrG6kScpT8VWpu1xWY5nJRCF3XaYyBjhI=
github.com/onsi/gomega v1.27.10/go.mod h1:RsS8tutOdbdgzbPtzzATp12yT7kM5I5aElG3evPbQ0M=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/quic-go/qpack v0.5.1 h1:giqksBPnT/HDtZ6VhtFKgoLOWmlyo9Ei6u9PqzIMbhI=
github.com/quic-go/qpack v0.5.1/go.mod h1:+PC4XFrEskIVkcLzpEkbLqq1uCoxPhQuvK5rH1ZgaEg=
github.com/quic-go/quic-go v0.48.1 h1:y/8xmfWI9qmGTc+lBr4jKRUWLGSlSigv847ULJ4hYXA=
github.com/quic-go/quic-go v0.48.1/go.mod h1:yBgs3rWBOADpga7F+jJsb6Ybg1LSYiQvwWlLX+/6HMs=
github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3 h1:f/FNXud6gA3MNr8meMVVGxhp+QBTqY91tM8HjEuMjGg=
github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3/go.mod h1:HgjTstvQsPGkxUsCd2KWxErBblirPizecHcpD3ffK+s=
github.com/shadowsocks/go-shadowsocks2 v0.1.5 h1:PDSQv9y2S85Fl7VBeOMF9StzeXZyK1HakRm86CUbr28=
github.com/shadowsocks/go-shadowsocks2 v0.1.5/go.mod h1:AGGpIoek4HRno4xzyFiAtLHkOpcoznZEkAccaI/rplM=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/tevino/abool v1.2.0 h1:heAkClL8H6w+mK5md9dzsuohKeXHUpY7Vw0ZCKW+huA=
github.com/tevino/abool v1.2.0/go.mod h1:qc66Pna1RiIsPa7O4Egxxs9OqkuxDX55zznh9K07Tzg=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
github.com/things-go/go-socks5 v0.0.5 h1:qvKaGcBkfDrUL33SchHN93srAmYGzb4CxSM2DPYufe8=
github.com/things-go/go-socks5 v0.0.5/go.mod h1:mtzInf8v5xmsBpHZVbIw2YQYhc4K0jRwzfsH64Uh0IQ=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
golang.org/x/crypto v0.39.0 h1:SHs+kF4LP+f+p14esP5jAoDpHU8Gu/v9lFRK6IT5imM=
golang.org/x/crypto v0.39.0/go.mod h1:L+Xg3Wf6HoL4Bn4238Z6ft6KfEpN0tJGo53AAPC632U=
golang.org/x/mod v0.13.0 h1:I/DsJXRlw/8l/0c24sM9yb0T4z9liZTduXvdAWYiysY=
golang.org/x/mod v0.21.0 h1:vvrHzRwRfVKSiLrG+d4FMl/Qi4ukBCE6kZlTUkDYRT0=
golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
golang.org/x/net v0.41.0 h1:vBTly1HeNPEn3wtREYfy4GZ/NECgw2Cnl+nK6Nz3uvw=
golang.org/x/net v0.41.0/go.mod h1:B/K4NNqkfmg07DQYrbwvSluqCJOOXwUjeb/5lOisjbA=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
go.uber.org/mock v0.4.0 h1:VcM4ZOtdbR4f6VXfiOpwpVJDL6lCReaZ6mw31wqh7KU=
go.uber.org/mock v0.4.0/go.mod h1:a6FSlNadKUHUa9IP5Vyt1zh4fC7uAwxMutEAscFbkZc=
golang.getoutline.org/sdk v0.0.23 h1:UKoKCrRH3Ed6Jpg8ODYwOo6O1B1lvfIHPreMrTkSTcc=
golang.getoutline.org/sdk v0.0.23/go.mod h1:nKZlO//e/sRFk+rp8gm8EJ5RasDSyY+fGDNSd3I2iaA=
golang.getoutline.org/sdk/x v0.2.0 h1:4kuT2SgkPXktwPwT6CXuF+Jwe+COAGqiBDZbgJkS5fM=
golang.getoutline.org/sdk/x v0.2.0/go.mod h1:vyWoHW0PUsnRHdLXiDGv1Z8X48Asp2ALgwIQfJfAaNs=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20210220033148-5ea612d1eb83/go.mod h1:jdWPYTVW3xRLrWPugEBEK3UY2ZEsg3UU495nc5E+M+I=
golang.org/x/crypto v0.42.0 h1:chiH31gIWm57EkTXpwnqf8qeuMUi0yekh6mT2AvFlqI=
golang.org/x/crypto v0.42.0/go.mod h1:4+rDnOTJhQCx2q7/j6rAN5XDw8kPjeaXEUR2eL94ix8=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
golang.org/x/mobile v0.0.0-20240520174638-fa72addaaa1b h1:WX7nnnLfCEXg+FmdYZPai2XuP3VqCP1HZVMST0n9DF0=
golang.org/x/mobile v0.0.0-20240520174638-fa72addaaa1b/go.mod h1:EiXZlVfUTaAyySFVJb9rsODuiO+WXu8HrUuySb7nYFw=
golang.org/x/mod v0.28.0 h1:gQBtGhjxykdjY9YhZpSlZIsbnaE2+PgjfLWUQTnoZ1U=
golang.org/x/mod v0.28.0/go.mod h1:yfB/L0NOf/kmEbXjzCPOx1iK1fRutOydrCMsqRhEBxI=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.44.0 h1:evd8IRDyfNBMBTTY5XRF1vaZlD+EmWx6x8PkhR04H/I=
golang.org/x/net v0.44.0/go.mod h1:ECOoLqd5U3Lhyeyo/QDCEVQ4sNgYsqvCZ722XogGieY=
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20191026070338-33540a1f6037/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20211007075335-d3039528d8ac/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.36.0 h1:KVRy2GtZBrk1cBYA7MKu5bEZFxQk4NIDV6RLVcC8o0k=
golang.org/x/sys v0.36.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/term v0.0.0-20201117132131-f5c789dd3221/go.mod h1:Nr5EML6q2oocZ2LXRh80K7BxOlk5/8JxuGnuhpl+muw=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
golang.org/x/time v0.9.0 h1:EsRrnYcQiGH+5FfbgvV4AP7qEZstoyrHB0DzarOQ4ZY=
golang.org/x/time v0.9.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/tools v0.37.0 h1:DVSRzp7FwePZW356yEAChSdNcQo6Nsp+fex1SUW09lE=
golang.org/x/tools v0.37.0/go.mod h1:MBN5QPQtLMHVdvsbtarmTNukZDdgwdwlO5qGacAzF0w=
golang.org/x/tools/go/expect v0.1.1-deprecated h1:jpBZDwmgPhXsKZC6WhL20P4b/wmnpsEAGHaNy0n/rJM=
golang.org/x/tools/go/expect v0.1.1-deprecated/go.mod h1:eihoPOH+FgIqa3FpoTwguz/bVUSGBlGQU67vpBeOrBY=
golang.org/x/tools/go/packages/packagestest v0.1.1-deprecated h1:1h2MnaIAIXISqTFKdENegdpAgUXz6NrPEsbIeWaBRvM=
golang.org/x/tools/go/packages/packagestest v0.1.1-deprecated/go.mod h1:RVAQXBGNv1ib0J382/DPCRS/BPnsGebyM1Gj5VSDpG8=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
google.golang.org/protobuf v1.33.0 h1:uNO2rsAINq/JlFpSdYEKIZ0uKD/R9cpdv0T+yoGwGmI=
google.golang.org/protobuf v1.33.0/go.mod h1:c6P6GXX6sHbq/GpV6MGZEdwhWPcYBgnhAHhKbcUYpos=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489 h1:ze1vwAdliUAr68RQ5NtufWaXaOg8WUO2OACzEV+TNdE=
gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489/go.mod h1:10sU+Uh5KKNv1+2x2A0Gvzt8FjD3ASIhorV3YsauXhk=
gvisor.dev/gvisor v0.0.0-20250428193742-2d800c3129d5 h1:sfK5nHuG7lRFZ2FdTT3RimOqWBg8IrVm+/Vko1FVOsk=
gvisor.dev/gvisor v0.0.0-20250428193742-2d800c3129d5/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
gvisor.dev/gvisor v0.0.0-20250606233247-e3c4c4cad86f h1:zmc4cHEcCudRt2O8VsCW7nYLfAsbVY2i910/DAop1TM=
gvisor.dev/gvisor v0.0.0-20250606233247-e3c4c4cad86f/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
+1 -1
View File
@@ -20,7 +20,7 @@ import (
"testing"
"time"
"github.com/amnezia-vpn/amneziawg-go/ipc/namedpipe"
"github.com/amnezia-vpn/amneziawg-go/v3/ipc/namedpipe"
"golang.org/x/sys/windows"
)
+1 -1
View File
@@ -9,7 +9,7 @@ import (
"net"
"os"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/v3/rwcancel"
"golang.org/x/sys/unix"
)
+1 -1
View File
@@ -8,7 +8,7 @@ package ipc
import (
"net"
"github.com/amnezia-vpn/amneziawg-go/ipc/namedpipe"
"github.com/amnezia-vpn/amneziawg-go/v3/ipc/namedpipe"
"golang.org/x/sys/windows"
)
+4 -4
View File
@@ -14,10 +14,10 @@ import (
"runtime"
"strconv"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device"
"github.com/amnezia-vpn/amneziawg-go/ipc"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/device"
"github.com/amnezia-vpn/amneziawg-go/v3/ipc"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
"golang.org/x/sys/unix"
)
+4 -4
View File
@@ -12,11 +12,11 @@ import (
"golang.org/x/sys/windows"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device"
"github.com/amnezia-vpn/amneziawg-go/ipc"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/device"
"github.com/amnezia-vpn/amneziawg-go/v3/ipc"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
)
const (
+77
View File
@@ -0,0 +1,77 @@
package outline
import (
"context"
"fmt"
"net"
"net/netip"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/device"
"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
"golang.getoutline.org/sdk/transport"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
)
type DialerOptions struct {
Ipc string
Prefixes []netip.Prefix
Mtu int
Dns []netip.Addr
}
func NewStreamDialer(opts DialerOptions) (*StreamDialer, error) {
var localAddresses []netip.Addr
for _, prefix := range opts.Prefixes {
localAddresses = append(localAddresses, prefix.Addr())
}
tun, tnet, err := netstack.CreateNetTUN(localAddresses, opts.Dns, opts.Mtu)
if err != nil {
return nil, fmt.Errorf("failed to create network tun: %v", err)
}
awgLogger := device.Logger{
Verbosef: func(format string, args ...any) {
},
Errorf: func(format string, args ...any) {
},
}
dev := device.NewDevice(tun, conn.NewDefaultBind(), &awgLogger)
if err := dev.IpcSet(opts.Ipc); err != nil {
return nil, fmt.Errorf("failed to configure device: %v", err)
}
if err := dev.Up(); err != nil {
return nil, fmt.Errorf("failed to start awg device: %v", err)
}
return &StreamDialer{
tnet: tnet,
}, nil
}
var _ transport.StreamDialer = (*StreamDialer)(nil)
type StreamDialer struct {
tnet *netstack.Net
}
func (d *StreamDialer) DialStream(ctx context.Context, raddr string) (transport.StreamConn, error) {
host, port, err := net.SplitHostPort(raddr)
if err != nil {
return nil, fmt.Errorf("failed to parse raddr: %v", err)
}
if l := len(host); l > 0 && host[l-1] == '.' {
host = host[:l-1]
raddr = net.JoinHostPort(host, port)
}
conn, err := d.tnet.DialContext(ctx, "tcp", raddr)
if err != nil {
return nil, err
}
return conn.(*gonet.TCPConn), nil
}
+224
View File
@@ -0,0 +1,224 @@
package outline
import (
"context"
"encoding/base64"
"encoding/hex"
"fmt"
"net/netip"
"strconv"
"strings"
"github.com/goccy/go-yaml"
"golang.getoutline.org/sdk/transport"
"golang.getoutline.org/sdk/x/mobileproxy"
"golang.getoutline.org/sdk/x/smart"
)
type DeviceConfig struct {
PrivateKey string `yaml:"private_key"`
Address []string `yaml:"address"`
Dns []string `yaml:"dns"`
Mtu int `yaml:"mtu,omitempty"`
Jc int `yaml:"jc,omitempty"`
Jmin int `yaml:"jmin,omitempty"`
Jmax int `yaml:"jmax,omitempty"`
S1 int `yaml:"s1,omitempty"`
S2 int `yaml:"s2,omitempty"`
S3 int `yaml:"s3,omitempty"`
S4 int `yaml:"s4,omitempty"`
H1 string `yaml:"h1,omitempty"`
H2 string `yaml:"h2,omitempty"`
H3 string `yaml:"h3,omitempty"`
H4 string `yaml:"h4,omitempty"`
I1 string `yaml:"i1,omitempty"`
I2 string `yaml:"i2,omitempty"`
I3 string `yaml:"i3,omitempty"`
I4 string `yaml:"i4,omitempty"`
I5 string `yaml:"i5,omitempty"`
Peers []PeerConfig `yaml:"peers,omitempty"`
}
type PeerConfig struct {
PublicKey string `yaml:"public_key"`
PresharedKey string `yaml:"preshared_key,omitempty"`
Endpoint string `yaml:"endpoint"`
AllowedIPs []string `yaml:"allowed_ips"`
PersistentKeepaliveInterval uint16 `yaml:"persistent_keepalive_interval,omitempty"`
}
func mapYamlToConfig(y smart.YAMLNode) (*DeviceConfig, error) {
bytes, err := yaml.Marshal(y)
if err != nil {
return nil, fmt.Errorf("failed to marshal yaml: %v", err)
}
var cfg DeviceConfig
if err = yaml.Unmarshal(bytes, &cfg); err != nil {
return nil, fmt.Errorf("failed to unmarshal yaml: %v", err)
}
return &cfg, nil
}
func genIpcString(cfg *DeviceConfig) (string, error) {
privateKeyBytes, err := base64.StdEncoding.DecodeString(cfg.PrivateKey)
if err != nil {
return "", fmt.Errorf("failed to decode private key: %v", err)
}
var b strings.Builder
b.WriteString("private_key=")
b.WriteString(hex.EncodeToString(privateKeyBytes))
if cfg.Jc != 0 {
b.WriteString("\njc=")
b.WriteString(strconv.Itoa(cfg.Jc))
}
if cfg.Jmin != 0 {
b.WriteString("\njmin=")
b.WriteString(strconv.Itoa(cfg.Jmin))
}
if cfg.Jmax != 0 {
b.WriteString("\njmax=")
b.WriteString(strconv.Itoa(cfg.Jmax))
}
if cfg.S1 != 0 {
b.WriteString("\ns1=")
b.WriteString(strconv.Itoa(cfg.S1))
}
if cfg.S2 != 0 {
b.WriteString("\ns2=")
b.WriteString(strconv.Itoa(cfg.S2))
}
if cfg.S3 != 0 {
b.WriteString("\ns3=")
b.WriteString(strconv.Itoa(cfg.S3))
}
if cfg.S4 != 0 {
b.WriteString("\ns4=")
b.WriteString(strconv.Itoa(cfg.S4))
}
if cfg.H1 != "" {
b.WriteString("\nh1=")
b.WriteString(cfg.H1)
}
if cfg.H2 != "" {
b.WriteString("\nh2=")
b.WriteString(cfg.H2)
}
if cfg.H3 != "" {
b.WriteString("\nh3=")
b.WriteString(cfg.H3)
}
if cfg.H4 != "" {
b.WriteString("\nh4=")
b.WriteString(cfg.H4)
}
if cfg.I1 != "" {
b.WriteString("\ni1=")
b.WriteString(cfg.I1)
}
if cfg.I2 != "" {
b.WriteString("\ni2=")
b.WriteString(cfg.I2)
}
if cfg.I3 != "" {
b.WriteString("\ni3=")
b.WriteString(cfg.I3)
}
if cfg.I4 != "" {
b.WriteString("\ni4=")
b.WriteString(cfg.I4)
}
if cfg.I5 != "" {
b.WriteString("\ni5=")
b.WriteString(cfg.I5)
}
for _, peer := range cfg.Peers {
publicKeyBytes, err := base64.StdEncoding.DecodeString(peer.PublicKey)
if err != nil {
return "", fmt.Errorf("failed to decode public key: %v", err)
}
b.WriteString("\npublic_key=")
b.WriteString(hex.EncodeToString(publicKeyBytes))
b.WriteString("\nendpoint=")
b.WriteString(peer.Endpoint)
for _, allowedIp := range peer.AllowedIPs {
b.WriteString("\nallowed_ip=")
b.WriteString(allowedIp)
}
if peer.PresharedKey != "" {
presharedKeyBytes, err := base64.StdEncoding.DecodeString(peer.PresharedKey)
if err != nil {
return "", fmt.Errorf("failed to decode preshared key: %v", err)
}
b.WriteString("\npreshared_key=")
b.WriteString(hex.EncodeToString(presharedKeyBytes))
}
if peer.PersistentKeepaliveInterval != 0 {
b.WriteString("\npersistent_keepalive_interval=")
b.WriteString(strconv.Itoa(int(peer.PersistentKeepaliveInterval)))
}
}
return b.String(), nil
}
func FallbackParser(ctx context.Context, y smart.YAMLNode) (transport.StreamDialer, string, error) {
cfg, err := mapYamlToConfig(y)
if err != nil {
return nil, "", fmt.Errorf("failed to map yaml to config: %v", err)
}
ipc, err := genIpcString(cfg)
if err != nil {
return nil, "", fmt.Errorf("faield to generate ipc config: %v", err)
}
var prefixes []netip.Prefix
for _, address := range cfg.Address {
prefix, err := netip.ParsePrefix(address)
if err != nil {
return nil, "", fmt.Errorf("failed to parse address: %v", err)
}
prefixes = append(prefixes, prefix)
}
var dns []netip.Addr
for _, saddr := range cfg.Dns {
addr, err := netip.ParseAddr(saddr)
if err != nil {
return nil, "", fmt.Errorf("failed to parse dns: %v", err)
}
dns = append(dns, addr)
}
if cfg.Mtu == 0 {
cfg.Mtu = 1408
}
dialer, err := NewStreamDialer(DialerOptions{
Ipc: ipc,
Prefixes: prefixes,
Mtu: cfg.Mtu,
Dns: dns,
})
if err != nil {
return nil, "", fmt.Errorf("failed to create dialer: %v", err)
}
return dialer, ipc, nil
}
func RegisterFallbackParser(opt *mobileproxy.SmartDialerOptions, name string) {
opt.RegisterFallbackParser(name, FallbackParser)
}
+52
View File
@@ -0,0 +1,52 @@
package outline_test
import (
"testing"
awg "github.com/amnezia-vpn/amneziawg-go/v3/outline"
"golang.getoutline.org/sdk/x/mobileproxy"
)
const cfg = `
dns:
- {system: {}}
tls:
- ""
fallback:
- awg:
address: [10.0.0.0/32]
dns: [8.8.8.8, 8.8.4.4]
private_key: +CdqlYvjqZ3OUr4mLWvGJo1h67CWpQwMIxA5OpyiJUM=
jc: 4
jmin: 50
jmax: 100
s1: 87
s2: 65
s3: 43
s4: 21
h1: 1000000000-1000000001
h2: 2000000000-2000000002
h3: 3000000000-3000000003
h4: 4000000000-4000000004
peers:
- public_key: EGxNYihRLKQ9nvdOE5j5aZ7rtw3ttzJS1xxaJpgYYHI=
preshared_key: 2OiSh6rP3t/g39jgJNGK70B+nize821yIFNtUqi8/XU=
endpoint: 123.123.123.123:51820
allowed_ips: [0.0.0.0/0, ::/0]
persistent_keepalive_interval: 25
`
var testDomains = mobileproxy.NewListFromLines("example.com")
func Test_outlineIntegration(t *testing.T) {
opts := mobileproxy.NewSmartDialerOptions(testDomains, cfg)
opts.SetLogWriter(mobileproxy.NewStderrLogWriter())
awg.RegisterFallbackParser(opts, "awg")
dialer, err := opts.NewStreamDialer()
if err != nil {
t.Fatal(err)
}
if _, err = mobileproxy.RunProxy("", dialer); err != nil {
t.Fatal(err)
}
}
+3 -3
View File
@@ -13,9 +13,9 @@ import (
"net/http"
"net/netip"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device"
"github.com/amnezia-vpn/amneziawg-go/tun/netstack"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/device"
"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
)
func main() {
+3 -3
View File
@@ -14,9 +14,9 @@ import (
"net/http"
"net/netip"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device"
"github.com/amnezia-vpn/amneziawg-go/tun/netstack"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/device"
"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
)
func main() {
+3 -3
View File
@@ -17,9 +17,9 @@ import (
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/device"
"github.com/amnezia-vpn/amneziawg-go/tun/netstack"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/device"
"github.com/amnezia-vpn/amneziawg-go/v3/tun/netstack"
)
func main() {
+1 -1
View File
@@ -22,7 +22,7 @@ import (
"syscall"
"time"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
"golang.org/x/net/dns/dnsmessage"
"gvisor.dev/gvisor/pkg/buffer"
+1 -1
View File
@@ -12,7 +12,7 @@ import (
"io"
"unsafe"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"golang.org/x/sys/unix"
)
+1 -1
View File
@@ -9,7 +9,7 @@ import (
"net/netip"
"testing"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
+2 -2
View File
@@ -17,8 +17,8 @@ import (
"time"
"unsafe"
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/v3/conn"
"github.com/amnezia-vpn/amneziawg-go/v3/rwcancel"
"golang.org/x/sys/unix"
)
+1 -1
View File
@@ -11,7 +11,7 @@ import (
"net/netip"
"os"
"github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/amnezia-vpn/amneziawg-go/v3/tun"
)
func Ping(dst, src netip.Addr) []byte {