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>
This commit is contained in:
parent
1abd24b5b9
commit
f6542209f4
22 changed files with 1352 additions and 603 deletions
|
|
@ -3,142 +3,88 @@ 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 Cfg struct {
|
||||
IsSet bool
|
||||
JunkPacketCount int
|
||||
JunkPacketMinSize int
|
||||
JunkPacketMaxSize int
|
||||
InitHeaderJunkSize int
|
||||
ResponseHeaderJunkSize int
|
||||
CookieReplyHeaderJunkSize int
|
||||
TransportHeaderJunkSize int
|
||||
|
||||
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)
|
||||
MagicHeaders MagicHeaders
|
||||
}
|
||||
|
||||
type Protocol struct {
|
||||
IsASecOn abool.AtomicBool
|
||||
IsOn abool.AtomicBool
|
||||
// TODO: revision the need of the mutex
|
||||
ASecMux sync.RWMutex
|
||||
ASecCfg aSecCfgType
|
||||
JunkCreator junkCreator
|
||||
Mux sync.RWMutex
|
||||
Cfg Cfg
|
||||
JunkCreator JunkCreator
|
||||
|
||||
HandshakeHandler SpecialHandshakeHandler
|
||||
}
|
||||
|
||||
func (protocol *Protocol) CreateInitHeaderJunk() ([]byte, error) {
|
||||
return protocol.createHeaderJunk(protocol.ASecCfg.InitHeaderJunkSize)
|
||||
protocol.Mux.RLock()
|
||||
defer protocol.Mux.RUnlock()
|
||||
|
||||
return protocol.createHeaderJunk(protocol.Cfg.InitHeaderJunkSize, 0)
|
||||
}
|
||||
|
||||
func (protocol *Protocol) CreateResponseHeaderJunk() ([]byte, error) {
|
||||
return protocol.createHeaderJunk(protocol.ASecCfg.ResponseHeaderJunkSize)
|
||||
protocol.Mux.RLock()
|
||||
defer protocol.Mux.RUnlock()
|
||||
|
||||
return protocol.createHeaderJunk(protocol.Cfg.ResponseHeaderJunkSize, 0)
|
||||
}
|
||||
|
||||
func (protocol *Protocol) CreateCookieReplyHeaderJunk() ([]byte, error) {
|
||||
return protocol.createHeaderJunk(protocol.ASecCfg.CookieReplyHeaderJunkSize)
|
||||
protocol.Mux.RLock()
|
||||
defer protocol.Mux.RUnlock()
|
||||
|
||||
return protocol.createHeaderJunk(protocol.Cfg.CookieReplyHeaderJunkSize, 0)
|
||||
}
|
||||
|
||||
func (protocol *Protocol) CreateTransportHeaderJunk(packetSize int) ([]byte, error) {
|
||||
return protocol.createHeaderJunk(protocol.ASecCfg.TransportHeaderJunkSize, packetSize)
|
||||
protocol.Mux.RLock()
|
||||
defer protocol.Mux.RUnlock()
|
||||
|
||||
return protocol.createHeaderJunk(protocol.Cfg.TransportHeaderJunkSize, packetSize)
|
||||
}
|
||||
|
||||
func (protocol *Protocol) createHeaderJunk(junkSize int, optExtraSize ...int) ([]byte, error) {
|
||||
extraSize := 0
|
||||
if len(optExtraSize) == 1 {
|
||||
extraSize = optExtraSize[0]
|
||||
func (protocol *Protocol) createHeaderJunk(junkSize int, extraSize int) ([]byte, error) {
|
||||
if junkSize == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
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
|
||||
buf := make([]byte, 0, junkSize+extraSize)
|
||||
writer := bytes.NewBuffer(buf[:0])
|
||||
|
||||
err := protocol.JunkCreator.AppendJunk(writer, junkSize)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("append junk: %w", err)
|
||||
}
|
||||
|
||||
return writer.Bytes(), nil
|
||||
}
|
||||
|
||||
func (protocol *Protocol) GetMagicHeaderMinFor(msgType uint32) (uint32, error) {
|
||||
for _, magicHeader := range protocol.Cfg.MagicHeaders.Values {
|
||||
if magicHeader.Min <= msgType && msgType <= magicHeader.Max {
|
||||
return magicHeader.Min, nil
|
||||
}
|
||||
junk = writer.Bytes()
|
||||
}
|
||||
protocol.ASecMux.RUnlock()
|
||||
|
||||
return junk, nil
|
||||
return 0, fmt.Errorf("no header for value: %d", msgType)
|
||||
}
|
||||
|
||||
func (protocol *Protocol) GetMsgType(defaultMsgType uint32) (uint32, error) {
|
||||
return protocol.Cfg.MagicHeaders.Get(defaultMsgType)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,69 +2,49 @@ package awg
|
|||
|
||||
import (
|
||||
"bytes"
|
||||
crand "crypto/rand"
|
||||
"fmt"
|
||||
v2 "math/rand/v2"
|
||||
)
|
||||
|
||||
type junkCreator struct {
|
||||
aSecCfg aSecCfgType
|
||||
cha8Rand *v2.ChaCha8
|
||||
type JunkCreator struct {
|
||||
cfg Cfg
|
||||
randomGenerator PRNG[int]
|
||||
}
|
||||
|
||||
// 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
|
||||
func NewJunkCreator(cfg Cfg) JunkCreator {
|
||||
return JunkCreator{cfg: cfg, randomGenerator: NewPRNG[int]()}
|
||||
}
|
||||
|
||||
// Should be called with aSecMux RLocked
|
||||
func (jc *junkCreator) CreateJunkPackets(junks *[][]byte) error {
|
||||
if jc.aSecCfg.JunkPacketCount == 0 {
|
||||
return nil
|
||||
// Should be called with awg mux RLocked
|
||||
func (jc *JunkCreator) CreateJunkPackets(junks *[][]byte) {
|
||||
if jc.cfg.JunkPacketCount == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
for range jc.aSecCfg.JunkPacketCount {
|
||||
for range jc.cfg.JunkPacketCount {
|
||||
packetSize := jc.randomPacketSize()
|
||||
junk, err := jc.randomJunkWithSize(packetSize)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create junk packet: %v", err)
|
||||
}
|
||||
junk := jc.randomJunkWithSize(packetSize)
|
||||
*junks = append(*junks, junk)
|
||||
}
|
||||
return nil
|
||||
return
|
||||
}
|
||||
|
||||
// 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 awg mux RLocked
|
||||
func (jc *JunkCreator) randomPacketSize() int {
|
||||
return jc.randomGenerator.RandomSizeInRange(jc.cfg.JunkPacketMinSize, jc.cfg.JunkPacketMaxSize)
|
||||
}
|
||||
|
||||
// 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)
|
||||
// Should be called with awg mux RLocked
|
||||
func (jc *JunkCreator) AppendJunk(writer *bytes.Buffer, size int) error {
|
||||
headerJunk := jc.randomJunkWithSize(size)
|
||||
_, 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
|
||||
// Should be called with awg mux RLocked
|
||||
func (jc *JunkCreator) randomJunkWithSize(size int) []byte {
|
||||
return jc.randomGenerator.ReadSize(size)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,43 +6,34 @@ import (
|
|||
"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,
|
||||
func setUpJunkCreator() JunkCreator {
|
||||
mh, _ := NewMagicHeaders(
|
||||
[]MagicHeader{
|
||||
NewMagicHeaderSameValue(123456),
|
||||
NewMagicHeaderSameValue(67543),
|
||||
NewMagicHeaderSameValue(32345),
|
||||
NewMagicHeaderSameValue(123123),
|
||||
},
|
||||
)
|
||||
|
||||
jc := NewJunkCreator(Cfg{
|
||||
IsSet: true,
|
||||
JunkPacketCount: 5,
|
||||
JunkPacketMinSize: 500,
|
||||
JunkPacketMaxSize: 1000,
|
||||
InitHeaderJunkSize: 30,
|
||||
ResponseHeaderJunkSize: 40,
|
||||
MagicHeaders: mh,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("failed to create junk creator %v", err)
|
||||
return junkCreator{}, err
|
||||
}
|
||||
|
||||
return jc, nil
|
||||
return jc
|
||||
}
|
||||
|
||||
func Test_junkCreator_createJunkPackets(t *testing.T) {
|
||||
jc, err := setUpJunkCreator(t)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
jc := setUpJunkCreator()
|
||||
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
|
||||
}
|
||||
got := make([][]byte, 0, jc.cfg.JunkPacketCount)
|
||||
jc.CreateJunkPackets(&got)
|
||||
seen := make(map[string]bool)
|
||||
for _, junk := range got {
|
||||
key := string(junk)
|
||||
|
|
@ -61,34 +52,28 @@ func Test_junkCreator_createJunkPackets(t *testing.T) {
|
|||
|
||||
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)
|
||||
jc := setUpJunkCreator()
|
||||
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)
|
||||
t.Errorf("same junks")
|
||||
return
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func Test_junkCreator_randomPacketSize(t *testing.T) {
|
||||
jc, err := setUpJunkCreator(t)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
jc := setUpJunkCreator()
|
||||
for range [30]struct{}{} {
|
||||
t.Run("valid", func(t *testing.T) {
|
||||
if got := jc.randomPacketSize(); jc.aSecCfg.JunkPacketMinSize > got ||
|
||||
got > jc.aSecCfg.JunkPacketMaxSize {
|
||||
if got := jc.randomPacketSize(); jc.cfg.JunkPacketMinSize > got ||
|
||||
got > jc.cfg.JunkPacketMaxSize {
|
||||
t.Errorf(
|
||||
"junkCreator.randomPacketSize() = %v, not between range [%v,%v]",
|
||||
got,
|
||||
jc.aSecCfg.JunkPacketMinSize,
|
||||
jc.aSecCfg.JunkPacketMaxSize,
|
||||
jc.cfg.JunkPacketMinSize,
|
||||
jc.cfg.JunkPacketMaxSize,
|
||||
)
|
||||
}
|
||||
})
|
||||
|
|
@ -96,10 +81,7 @@ func Test_junkCreator_randomPacketSize(t *testing.T) {
|
|||
}
|
||||
|
||||
func Test_junkCreator_appendJunk(t *testing.T) {
|
||||
jc, err := setUpJunkCreator(t)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
jc := setUpJunkCreator()
|
||||
t.Run("valid", func(t *testing.T) {
|
||||
s := "apple"
|
||||
buffer := bytes.NewBuffer([]byte(s))
|
||||
|
|
|
|||
97
device/awg/magic_header.go
Normal file
97
device/awg/magic_header.go
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
package awg
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type MagicHeader struct {
|
||||
Min uint32
|
||||
Max uint32
|
||||
}
|
||||
|
||||
func NewMagicHeaderSameValue(value uint32) MagicHeader {
|
||||
return MagicHeader{Min: value, Max: value}
|
||||
}
|
||||
|
||||
func NewMagicHeader(min, max uint32) (MagicHeader, error) {
|
||||
if min > max {
|
||||
return MagicHeader{}, fmt.Errorf("min (%d) cannot be greater than max (%d)", min, max)
|
||||
}
|
||||
|
||||
return MagicHeader{Min: min, Max: max}, nil
|
||||
}
|
||||
|
||||
func ParseMagicHeader(key, value string) (MagicHeader, error) {
|
||||
hyphenIdx := strings.Index(value, "-")
|
||||
if hyphenIdx == -1 {
|
||||
// if there is no hyphen, we treat it as single magic header value
|
||||
magicHeader, err := strconv.ParseUint(value, 10, 32)
|
||||
if err != nil {
|
||||
return MagicHeader{}, fmt.Errorf("parse key: %s; value: %s; %w", key, value, err)
|
||||
}
|
||||
|
||||
return NewMagicHeader(uint32(magicHeader), uint32(magicHeader))
|
||||
}
|
||||
|
||||
minStr := value[:hyphenIdx]
|
||||
maxStr := value[hyphenIdx+1:]
|
||||
if len(minStr) == 0 || len(maxStr) == 0 {
|
||||
return MagicHeader{}, fmt.Errorf("invalid value for key: %s; value: %s; expected format: min-max", key, value)
|
||||
}
|
||||
|
||||
min, err := strconv.ParseUint(minStr, 10, 32)
|
||||
if err != nil {
|
||||
return MagicHeader{}, fmt.Errorf("parse min key: %s; value: %s; %w", key, minStr, err)
|
||||
}
|
||||
|
||||
max, err := strconv.ParseUint(maxStr, 10, 32)
|
||||
if err != nil {
|
||||
return MagicHeader{}, fmt.Errorf("parse max key: %s; value: %s; %w", key, maxStr, err)
|
||||
}
|
||||
|
||||
magicHeader, err := NewMagicHeader(uint32(min), uint32(max))
|
||||
if err != nil {
|
||||
return MagicHeader{}, fmt.Errorf("new magicHeader key: %s; value: %s-%s; %w", key, minStr, maxStr, err)
|
||||
}
|
||||
|
||||
return magicHeader, nil
|
||||
}
|
||||
|
||||
type MagicHeaders struct {
|
||||
Values []MagicHeader
|
||||
randomGenerator RandomNumberGenerator[uint32]
|
||||
}
|
||||
|
||||
func NewMagicHeaders(headerValues []MagicHeader) (MagicHeaders, error) {
|
||||
if len(headerValues) != 4 {
|
||||
return MagicHeaders{}, fmt.Errorf("all header types should be included: %v", headerValues)
|
||||
}
|
||||
|
||||
sortedMagicHeaders := slices.SortedFunc(slices.Values(headerValues), func(lhs MagicHeader, rhs MagicHeader) int {
|
||||
return cmp.Compare(lhs.Min, rhs.Min)
|
||||
})
|
||||
|
||||
for i := range 3 {
|
||||
if sortedMagicHeaders[i].Max >= sortedMagicHeaders[i+1].Min {
|
||||
return MagicHeaders{}, fmt.Errorf(
|
||||
"magic headers shouldn't overlap; %v > %v",
|
||||
sortedMagicHeaders[i].Max,
|
||||
sortedMagicHeaders[i+1].Min,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return MagicHeaders{Values: headerValues, randomGenerator: NewPRNG[uint32]()}, nil
|
||||
}
|
||||
|
||||
func (mh *MagicHeaders) Get(defaultMsgType uint32) (uint32, error) {
|
||||
if defaultMsgType == 0 || defaultMsgType > 4 {
|
||||
return 0, fmt.Errorf("invalid msg type: %d", defaultMsgType)
|
||||
}
|
||||
|
||||
return mh.randomGenerator.RandomSizeInRange(mh.Values[defaultMsgType-1].Min, mh.Values[defaultMsgType-1].Max), nil
|
||||
}
|
||||
488
device/awg/magic_header_test.go
Normal file
488
device/awg/magic_header_test.go
Normal file
|
|
@ -0,0 +1,488 @@
|
|||
package awg
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNewMagicHeaderSameValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
value uint32
|
||||
expected MagicHeader
|
||||
}{
|
||||
{
|
||||
name: "zero value",
|
||||
value: 0,
|
||||
expected: MagicHeader{Min: 0, Max: 0},
|
||||
},
|
||||
{
|
||||
name: "small value",
|
||||
value: 1,
|
||||
expected: MagicHeader{Min: 1, Max: 1},
|
||||
},
|
||||
{
|
||||
name: "large value",
|
||||
value: 4294967295, // max uint32
|
||||
expected: MagicHeader{Min: 4294967295, Max: 4294967295},
|
||||
},
|
||||
{
|
||||
name: "medium value",
|
||||
value: 1000,
|
||||
expected: MagicHeader{Min: 1000, Max: 1000},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result := NewMagicHeaderSameValue(tt.value)
|
||||
require.Equal(t, tt.expected, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewMagicHeader(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
min uint32
|
||||
max uint32
|
||||
expected MagicHeader
|
||||
errorMsg string
|
||||
}{
|
||||
{
|
||||
name: "valid range",
|
||||
min: 1,
|
||||
max: 10,
|
||||
expected: MagicHeader{Min: 1, Max: 10},
|
||||
},
|
||||
{
|
||||
name: "equal values",
|
||||
min: 5,
|
||||
max: 5,
|
||||
expected: MagicHeader{Min: 5, Max: 5},
|
||||
},
|
||||
{
|
||||
name: "zero range",
|
||||
min: 0,
|
||||
max: 0,
|
||||
expected: MagicHeader{Min: 0, Max: 0},
|
||||
},
|
||||
{
|
||||
name: "max uint32 range",
|
||||
min: 4294967294,
|
||||
max: 4294967295,
|
||||
expected: MagicHeader{Min: 4294967294, Max: 4294967295},
|
||||
},
|
||||
{
|
||||
name: "min greater than max",
|
||||
min: 10,
|
||||
max: 5,
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "min (10) cannot be greater than max (5)",
|
||||
},
|
||||
{
|
||||
name: "large min greater than max",
|
||||
min: 4294967295,
|
||||
max: 1,
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "min (4294967295) cannot be greater than max (1)",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result, err := NewMagicHeader(tt.min, tt.max)
|
||||
|
||||
if tt.errorMsg != "" {
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), tt.errorMsg)
|
||||
require.Equal(t, MagicHeader{}, result)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.expected, result)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseMagicHeader(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
key string
|
||||
value string
|
||||
expected MagicHeader
|
||||
errorMsg string
|
||||
}{
|
||||
{
|
||||
name: "single value",
|
||||
key: "header1",
|
||||
value: "100",
|
||||
expected: MagicHeader{Min: 100, Max: 100},
|
||||
},
|
||||
{
|
||||
name: "valid range",
|
||||
key: "header2",
|
||||
value: "10-20",
|
||||
expected: MagicHeader{Min: 10, Max: 20},
|
||||
},
|
||||
{
|
||||
name: "zero single value",
|
||||
key: "header3",
|
||||
value: "0",
|
||||
expected: MagicHeader{Min: 0, Max: 0},
|
||||
},
|
||||
{
|
||||
name: "zero range",
|
||||
key: "header4",
|
||||
value: "0-0",
|
||||
expected: MagicHeader{Min: 0, Max: 0},
|
||||
},
|
||||
{
|
||||
name: "max uint32 single",
|
||||
key: "header5",
|
||||
value: "4294967295",
|
||||
expected: MagicHeader{Min: 4294967295, Max: 4294967295},
|
||||
},
|
||||
{
|
||||
name: "max uint32 range",
|
||||
key: "header6",
|
||||
value: "4294967294-4294967295",
|
||||
expected: MagicHeader{Min: 4294967294, Max: 4294967295},
|
||||
},
|
||||
{
|
||||
name: "invalid single value - not number",
|
||||
key: "header7",
|
||||
value: "abc",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "parse key: header7; value: abc;",
|
||||
},
|
||||
{
|
||||
name: "invalid single value - negative",
|
||||
key: "header8",
|
||||
value: "-5",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "invalid value for key: header8; value: -5;",
|
||||
},
|
||||
{
|
||||
name: "invalid single value - too large",
|
||||
key: "header9",
|
||||
value: "4294967296",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "parse key: header9; value: 4294967296;",
|
||||
},
|
||||
{
|
||||
name: "invalid range - min not number",
|
||||
key: "header10",
|
||||
value: "abc-10",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "parse min key: header10; value: abc;",
|
||||
},
|
||||
{
|
||||
name: "invalid range - max not number",
|
||||
key: "header11",
|
||||
value: "10-abc",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "parse max key: header11; value: abc;",
|
||||
},
|
||||
{
|
||||
name: "invalid range - min greater than max",
|
||||
key: "header12",
|
||||
value: "20-10",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "new magicHeader key: header12; value: 20-10;",
|
||||
},
|
||||
{
|
||||
name: "invalid range - too many parts",
|
||||
key: "header13",
|
||||
value: "10-20-30",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "parse key: header13; value: 10-20-30;",
|
||||
},
|
||||
{
|
||||
name: "empty value",
|
||||
key: "header14",
|
||||
value: "",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "parse key: header14; value: ;",
|
||||
},
|
||||
{
|
||||
name: "hyphen only",
|
||||
key: "header15",
|
||||
value: "-",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "invalid value for key: header15; value: -;",
|
||||
},
|
||||
{
|
||||
name: "empty min",
|
||||
key: "header16",
|
||||
value: "-10",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "invalid value for key: header16; value: -10;",
|
||||
},
|
||||
{
|
||||
name: "empty max",
|
||||
key: "header17",
|
||||
value: "10-",
|
||||
expected: MagicHeader{},
|
||||
errorMsg: "invalid value for key: header17; value: 10-;",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result, err := ParseMagicHeader(tt.key, tt.value)
|
||||
|
||||
if tt.errorMsg != "" {
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), tt.errorMsg)
|
||||
require.Equal(t, MagicHeader{}, result)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.expected, result)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewMagicHeaders(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
magicHeaders []MagicHeader
|
||||
errorMsg string
|
||||
}{
|
||||
{
|
||||
name: "valid non-overlapping headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 11, Max: 20},
|
||||
{Min: 21, Max: 30},
|
||||
{Min: 31, Max: 40},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "valid adjacent headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 1},
|
||||
{Min: 2, Max: 2},
|
||||
{Min: 3, Max: 3},
|
||||
{Min: 4, Max: 4},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "valid zero-based headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 0, Max: 0},
|
||||
{Min: 1, Max: 1},
|
||||
{Min: 2, Max: 2},
|
||||
{Min: 3, Max: 3},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "valid large value headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 4294967290, Max: 4294967291},
|
||||
{Min: 4294967292, Max: 4294967293},
|
||||
{Min: 4294967294, Max: 4294967294},
|
||||
{Min: 4294967295, Max: 4294967295},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "too few headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 11, Max: 20},
|
||||
{Min: 21, Max: 30},
|
||||
},
|
||||
errorMsg: "all header types should be included:",
|
||||
},
|
||||
{
|
||||
name: "too many headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 11, Max: 20},
|
||||
{Min: 21, Max: 30},
|
||||
{Min: 31, Max: 40},
|
||||
{Min: 41, Max: 50},
|
||||
},
|
||||
errorMsg: "all header types should be included:",
|
||||
},
|
||||
{
|
||||
name: "empty headers",
|
||||
magicHeaders: []MagicHeader{},
|
||||
errorMsg: "all header types should be included:",
|
||||
},
|
||||
{
|
||||
name: "overlapping headers",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 15},
|
||||
{Min: 10, Max: 20},
|
||||
{Min: 25, Max: 30},
|
||||
{Min: 35, Max: 40},
|
||||
},
|
||||
errorMsg: "magic headers shouldn't overlap;",
|
||||
},
|
||||
{
|
||||
name: "overlapping headers at limit-first",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 10, Max: 20},
|
||||
{Min: 25, Max: 30},
|
||||
{Min: 35, Max: 40},
|
||||
},
|
||||
errorMsg: "magic headers shouldn't overlap;",
|
||||
},
|
||||
{
|
||||
name: "overlapping headers at limit-second",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 15, Max: 25},
|
||||
{Min: 25, Max: 30},
|
||||
{Min: 35, Max: 40},
|
||||
},
|
||||
errorMsg: "magic headers shouldn't overlap;",
|
||||
},
|
||||
{
|
||||
name: "overlapping headers at limit-third",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 15, Max: 25},
|
||||
{Min: 30, Max: 35},
|
||||
{Min: 35, Max: 40},
|
||||
},
|
||||
errorMsg: "magic headers shouldn't overlap;",
|
||||
},
|
||||
{
|
||||
name: "identical ranges",
|
||||
magicHeaders: []MagicHeader{
|
||||
{Min: 10, Max: 20},
|
||||
{Min: 10, Max: 20},
|
||||
{Min: 25, Max: 30},
|
||||
{Min: 35, Max: 40},
|
||||
},
|
||||
errorMsg: "magic headers shouldn't overlap;",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result, err := NewMagicHeaders(tt.magicHeaders)
|
||||
|
||||
if tt.errorMsg != "" {
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), tt.errorMsg)
|
||||
require.Equal(t, MagicHeaders{}, result)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.magicHeaders, result.Values)
|
||||
require.NotNil(t, result.randomGenerator)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Mock PRNG for testing
|
||||
type mockPRNG struct {
|
||||
returnValue uint32
|
||||
}
|
||||
|
||||
func (m *mockPRNG) RandomSizeInRange(min, max uint32) uint32 {
|
||||
return m.returnValue
|
||||
}
|
||||
|
||||
func (m *mockPRNG) Get() uint64 {
|
||||
return 0
|
||||
}
|
||||
func (m *mockPRNG) ReadSize(size int) []byte {
|
||||
return make([]byte, size)
|
||||
}
|
||||
|
||||
func TestMagicHeaders_Get(t *testing.T) {
|
||||
// Create test headers
|
||||
headers := []MagicHeader{
|
||||
{Min: 1, Max: 10},
|
||||
{Min: 11, Max: 20},
|
||||
{Min: 21, Max: 30},
|
||||
{Min: 31, Max: 40},
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
defaultMsgType uint32
|
||||
mockValue uint32
|
||||
expectedValue uint32
|
||||
errorMsg string
|
||||
}{
|
||||
{
|
||||
name: "valid type 1",
|
||||
defaultMsgType: 1,
|
||||
mockValue: 5,
|
||||
expectedValue: 5,
|
||||
},
|
||||
{
|
||||
name: "valid type 2",
|
||||
defaultMsgType: 2,
|
||||
mockValue: 15,
|
||||
expectedValue: 15,
|
||||
},
|
||||
{
|
||||
name: "valid type 3",
|
||||
defaultMsgType: 3,
|
||||
mockValue: 25,
|
||||
expectedValue: 25,
|
||||
},
|
||||
{
|
||||
name: "valid type 4",
|
||||
defaultMsgType: 4,
|
||||
mockValue: 35,
|
||||
expectedValue: 35,
|
||||
},
|
||||
{
|
||||
name: "invalid type 0",
|
||||
defaultMsgType: 0,
|
||||
mockValue: 0,
|
||||
expectedValue: 0,
|
||||
errorMsg: "invalid msg type: 0",
|
||||
},
|
||||
{
|
||||
name: "invalid type 5",
|
||||
defaultMsgType: 5,
|
||||
mockValue: 0,
|
||||
expectedValue: 0,
|
||||
errorMsg: "invalid msg type: 5",
|
||||
},
|
||||
{
|
||||
name: "invalid type max uint32",
|
||||
defaultMsgType: 4294967295,
|
||||
mockValue: 0,
|
||||
expectedValue: 0,
|
||||
errorMsg: "invalid msg type: 4294967295",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Create a new instance with mock PRNG for each test
|
||||
testMagicHeaders := MagicHeaders{
|
||||
Values: headers,
|
||||
randomGenerator: &mockPRNG{returnValue: tt.mockValue},
|
||||
}
|
||||
|
||||
result, err := testMagicHeaders.Get(tt.defaultMsgType)
|
||||
|
||||
if tt.errorMsg != "" {
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), tt.errorMsg)
|
||||
require.Equal(t, uint32(0), result)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.expectedValue, result)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
50
device/awg/prng.go
Normal file
50
device/awg/prng.go
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
package awg
|
||||
|
||||
import (
|
||||
crand "crypto/rand"
|
||||
v2 "math/rand/v2"
|
||||
|
||||
"golang.org/x/exp/constraints"
|
||||
)
|
||||
|
||||
type RandomNumberGenerator[T constraints.Integer] interface {
|
||||
RandomSizeInRange(min, max T) T
|
||||
Get() uint64
|
||||
ReadSize(size int) []byte
|
||||
}
|
||||
|
||||
type PRNG[T constraints.Integer] struct {
|
||||
cha8Rand *v2.ChaCha8
|
||||
}
|
||||
|
||||
func NewPRNG[T constraints.Integer]() PRNG[T] {
|
||||
buf := make([]byte, 32)
|
||||
_, _ = crand.Read(buf)
|
||||
|
||||
return PRNG[T]{
|
||||
cha8Rand: v2.NewChaCha8([32]byte(buf)),
|
||||
}
|
||||
}
|
||||
|
||||
func (p PRNG[T]) RandomSizeInRange(min, max T) T {
|
||||
if min > max {
|
||||
panic("min must be less than max")
|
||||
}
|
||||
|
||||
if min == max {
|
||||
return min
|
||||
}
|
||||
|
||||
return T(p.Get()%uint64(max-min)) + min
|
||||
}
|
||||
|
||||
func (p PRNG[T]) Get() uint64 {
|
||||
return p.cha8Rand.Uint64()
|
||||
}
|
||||
|
||||
func (p PRNG[T]) ReadSize(size int) []byte {
|
||||
// TODO: use a memory pool to allocate
|
||||
buf := make([]byte, size)
|
||||
_, _ = p.cha8Rand.Read(buf)
|
||||
return buf
|
||||
}
|
||||
|
|
@ -1,9 +1,6 @@
|
|||
package awg
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/tevino/abool"
|
||||
"go.uber.org/atomic"
|
||||
)
|
||||
|
|
@ -21,25 +18,13 @@ var WaitResponse = struct {
|
|||
}
|
||||
|
||||
type SpecialHandshakeHandler struct {
|
||||
isFirstDone bool
|
||||
SpecialJunk TagJunkPacketGenerators
|
||||
ControlledJunk TagJunkPacketGenerators
|
||||
|
||||
nextItime time.Time
|
||||
ITimeout time.Duration // seconds
|
||||
SpecialJunk TagJunkPacketGenerators
|
||||
|
||||
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...)
|
||||
return handler.SpecialJunk.Validate()
|
||||
}
|
||||
|
||||
func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
|
||||
|
|
@ -47,27 +32,5 @@ func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
|
|||
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()
|
||||
return handler.SpecialJunk.GeneratePackets()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -59,43 +59,110 @@ func hexToBytes(hexStr string) ([]byte, error) {
|
|||
return hex.DecodeString(hexStr)
|
||||
}
|
||||
|
||||
type RandomPacketGenerator struct {
|
||||
type randomGeneratorBase 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) {
|
||||
func newRandomGeneratorBase(param string) (*randomGeneratorBase, error) {
|
||||
size, err := strconv.Atoi(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("random packet parse int: %w", err)
|
||||
return nil, fmt.Errorf("parse int: %w", err)
|
||||
}
|
||||
|
||||
if size > 1000 {
|
||||
return nil, fmt.Errorf("random packet size must be less than 1000")
|
||||
return nil, fmt.Errorf("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 nil, fmt.Errorf("crand read: %w", err)
|
||||
}
|
||||
|
||||
return &RandomPacketGenerator{
|
||||
return &randomGeneratorBase{
|
||||
cha8Rand: v2.NewChaCha8([32]byte(buf)),
|
||||
size: size,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (rpg *randomGeneratorBase) generate() []byte {
|
||||
junk := make([]byte, rpg.size)
|
||||
rpg.cha8Rand.Read(junk)
|
||||
return junk
|
||||
}
|
||||
|
||||
func (rpg *randomGeneratorBase) Size() int {
|
||||
return rpg.size
|
||||
}
|
||||
|
||||
type RandomBytesGenerator struct {
|
||||
*randomGeneratorBase
|
||||
}
|
||||
|
||||
func newRandomBytesGenerator(param string) (Generator, error) {
|
||||
rpgBase, err := newRandomGeneratorBase(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("new random bytes generator: %w", err)
|
||||
}
|
||||
|
||||
return &RandomBytesGenerator{randomGeneratorBase: rpgBase}, nil
|
||||
}
|
||||
|
||||
func (rpg *RandomBytesGenerator) Generate() []byte {
|
||||
return rpg.generate()
|
||||
}
|
||||
|
||||
const alphanumericChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
|
||||
|
||||
type RandomASCIIGenerator struct {
|
||||
*randomGeneratorBase
|
||||
}
|
||||
|
||||
func newRandomASCIIGenerator(param string) (Generator, error) {
|
||||
rpgBase, err := newRandomGeneratorBase(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("new random ascii generator: %w", err)
|
||||
}
|
||||
|
||||
return &RandomASCIIGenerator{randomGeneratorBase: rpgBase}, nil
|
||||
}
|
||||
|
||||
func (rpg *RandomASCIIGenerator) Generate() []byte {
|
||||
junk := rpg.generate()
|
||||
|
||||
result := make([]byte, rpg.size)
|
||||
for i, b := range junk {
|
||||
result[i] = alphanumericChars[b%byte(len(alphanumericChars))]
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
type RandomDigitGenerator struct {
|
||||
*randomGeneratorBase
|
||||
}
|
||||
|
||||
func newRandomDigitGenerator(param string) (Generator, error) {
|
||||
rpgBase, err := newRandomGeneratorBase(param)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("new random digit generator: %w", err)
|
||||
}
|
||||
|
||||
return &RandomDigitGenerator{randomGeneratorBase: rpgBase}, nil
|
||||
}
|
||||
|
||||
func (rpg *RandomDigitGenerator) Generate() []byte {
|
||||
junk := rpg.generate()
|
||||
|
||||
result := make([]byte, rpg.size)
|
||||
for i, b := range junk {
|
||||
result[i] = '0' + (b % 10) // Convert to digit character
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
type TimestampGenerator struct {
|
||||
}
|
||||
|
||||
|
|
@ -117,34 +184,6 @@ func newTimestampGenerator(param string) (Generator, error) {
|
|||
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 {
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,9 @@ import (
|
|||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func Test_newBytesGenerator(t *testing.T) {
|
||||
func TestNewBytesGenerator(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
type args struct {
|
||||
param string
|
||||
}
|
||||
|
|
@ -63,6 +65,8 @@ func Test_newBytesGenerator(t *testing.T) {
|
|||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got, err := newBytesGenerator(tt.args.param)
|
||||
|
||||
if tt.wantErr != nil {
|
||||
|
|
@ -80,7 +84,9 @@ func Test_newBytesGenerator(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func Test_newRandomPacketGenerator(t *testing.T) {
|
||||
func TestNewRandomBytesGenerator(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
type args struct {
|
||||
param string
|
||||
}
|
||||
|
|
@ -117,9 +123,134 @@ func Test_newRandomPacketGenerator(t *testing.T) {
|
|||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := newRandomPacketGenerator(tt.args.param)
|
||||
t.Parallel()
|
||||
|
||||
got, err := newRandomBytesGenerator(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 TestNewRandomASCIIGenerator(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
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) {
|
||||
t.Parallel()
|
||||
|
||||
got, err := newRandomASCIIGenerator(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 TestNewRandomDigitGenerator(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
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) {
|
||||
t.Parallel()
|
||||
|
||||
got, err := newRandomDigitGenerator(tt.args.param)
|
||||
if tt.wantErr != nil {
|
||||
require.ErrorAs(t, err, &tt.wantErr)
|
||||
require.Nil(t, got)
|
||||
|
|
@ -137,6 +268,8 @@ func Test_newRandomPacketGenerator(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestPacketCounterGenerator(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
param string
|
||||
|
|
@ -155,7 +288,6 @@ func TestPacketCounterGenerator(t *testing.T) {
|
|||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
tc := tc // capture range variable
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
|
|
|||
|
|
@ -12,21 +12,21 @@ 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"
|
||||
BytesEnumTag EnumTag = "b"
|
||||
CounterEnumTag EnumTag = "c"
|
||||
TimestampEnumTag EnumTag = "t"
|
||||
RandomBytesEnumTag EnumTag = "r"
|
||||
RandomASCIIEnumTag EnumTag = "rc"
|
||||
RandomDigitEnumTag EnumTag = "rd"
|
||||
)
|
||||
|
||||
var generatorCreator = map[EnumTag]newGenerator{
|
||||
BytesEnumTag: newBytesGenerator,
|
||||
CounterEnumTag: newPacketCounterGenerator,
|
||||
TimestampEnumTag: newTimestampGenerator,
|
||||
RandomBytesEnumTag: newRandomPacketGenerator,
|
||||
WaitTimeoutEnumTag: newWaitTimeoutGenerator,
|
||||
// WaitResponseEnumTag: newWaitResponseGenerator,
|
||||
RandomBytesEnumTag: newRandomBytesGenerator,
|
||||
RandomASCIIEnumTag: newRandomASCIIGenerator,
|
||||
RandomDigitEnumTag: newRandomDigitGenerator,
|
||||
}
|
||||
|
||||
// helper map to determine enumTags are unique
|
||||
|
|
@ -55,7 +55,7 @@ func parseTag(input string) (Tag, error) {
|
|||
return tag, nil
|
||||
}
|
||||
|
||||
func Parse(name, input string) (TagJunkPacketGenerator, error) {
|
||||
func ParseTagJunkGenerator(name, input string) (TagJunkPacketGenerator, error) {
|
||||
inputSlice := strings.Split(input, "<")
|
||||
if len(inputSlice) <= 1 {
|
||||
return TagJunkPacketGenerator{}, fmt.Errorf("empty input: %s", input)
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ func TestParse(t *testing.T) {
|
|||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := Parse(tt.args.name, tt.args.input)
|
||||
_, err := ParseTagJunkGenerator(tt.args.name, tt.args.input)
|
||||
|
||||
// TODO: ErrorAs doesn't work as you think
|
||||
if tt.wantErr != nil {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue