Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion Makefile
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
.PHONY: all lint test vet tidy vuln fuzz gen
.PHONY: all lint test vet tidy vuln fuzz gen bench

all: test

Expand All @@ -22,6 +22,7 @@ test: tidy lint vet vuln
@echo "testing..."
@go test -v -count=1 -race ./...
@make fuzz
@make bench

fuzz:
@echo "fuzzing..."
Expand All @@ -33,6 +34,10 @@ fuzz:
@go test -fuzz=FuzzPartialCorruption ./crypto -fuzztime=3s
@go test -fuzz=FuzzDifferentKeyPairs ./crypto -fuzztime=3s

bench:
@echo "benchmarking..."
@go test -bench=. ./...

gen:
@echo "generating..."
@go generate ./... && go fmt ./...
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
module github.com/agentuity/go-common

go 1.25.0
go 1.25.1

require (
github.com/buger/goterm v1.0.4
Expand Down
182 changes: 182 additions & 0 deletions network/concurrent_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,182 @@
package network

import (
"fmt"
"net"
"sync"
"testing"

"github.com/stretchr/testify/assert"
)

func TestConcurrentSubnetGeneration(t *testing.T) {
const numGoroutines = 10
const subnetsPerGoroutine = 50

var wg sync.WaitGroup
var mu sync.Mutex
allSubnets := make(map[string]bool)
errors := make([]error, 0)

for i := range numGoroutines {
wg.Add(1)
go func(goroutineID int) {
defer wg.Done()

for range subnetsPerGoroutine {
subnet, gateway, err := GenerateNonOverlappingIPv4Subnet([]*net.IPNet{}, 24)

mu.Lock()
if err != nil {
errors = append(errors, err)
} else {
// Note: Duplicates are expected here because each call is independent
// The function only avoids overlaps with the provided existingNetworks slice
subnetStr := subnet.String()
allSubnets[subnetStr] = true

// Basic validation
if subnet == nil || gateway == nil {
t.Errorf("Got nil subnet or gateway")
}
}
mu.Unlock()
}
}(i)
}

wg.Wait()

// Check results
assert.Empty(t, errors, "Should have no errors in concurrent generation")

// With concurrent access and no shared state, duplicates are expected
// The important thing is that the function doesn't crash or corrupt state
totalCalls := numGoroutines * subnetsPerGoroutine
uniqueSubnets := len(allSubnets)

t.Logf("Generated %d unique subnets from %d total calls across %d goroutines",
uniqueSubnets, totalCalls, numGoroutines)

// We should have at least some subnets generated
assert.Greater(t, uniqueSubnets, 0, "Should generate at least some subnets")
assert.LessOrEqual(t, uniqueSubnets, totalCalls, "Unique subnets should not exceed total calls")
}

func TestConcurrentWithExistingSubnets(t *testing.T) {
// Pre-generate some existing subnets
var existingNetworks []*net.IPNet
for range 100 {
subnet, _, err := GenerateNonOverlappingIPv4Subnet(existingNetworks, 24)
assert.NoError(t, err)
existingNetworks = append(existingNetworks, subnet)
}

const numGoroutines = 5
const subnetsPerGoroutine = 20

var wg sync.WaitGroup
var mu sync.Mutex
allSubnets := make(map[string]bool)
errors := make([]error, 0)

for i := range numGoroutines {
wg.Add(1)
go func(goroutineID int) {
defer wg.Done()

for range subnetsPerGoroutine {
subnet, gateway, err := GenerateNonOverlappingIPv4Subnet(existingNetworks, 24)

mu.Lock()
if err != nil {
errors = append(errors, err)
} else {
subnetStr := subnet.String()

// Check against existing networks - this should never happen
for _, existing := range existingNetworks {
if existing.String() == subnetStr {
t.Errorf("Generated subnet overlaps with existing: %s", subnetStr)
}
}

// Note: Duplicates between concurrent calls are expected
// Each call is independent and doesn't know about other concurrent results
allSubnets[subnetStr] = true

// Basic validation
if subnet == nil || gateway == nil {
t.Errorf("Got nil subnet or gateway")
}
}
mu.Unlock()
}
}(i)
}

wg.Wait()

// Check results
assert.Empty(t, errors, "Should have no errors in concurrent generation with existing subnets")

// With concurrent access, some duplicates between goroutines are expected
totalCalls := numGoroutines * subnetsPerGoroutine
uniqueSubnets := len(allSubnets)

// We should have at least some subnets generated
assert.Greater(t, uniqueSubnets, 0, "Should generate at least some subnets")
assert.LessOrEqual(t, uniqueSubnets, totalCalls, "Unique subnets should not exceed total calls")

t.Logf("Generated %d unique subnets from %d total calls across %d goroutines with %d existing subnets",
uniqueSubnets, totalCalls, numGoroutines, len(existingNetworks))
}

func TestConcurrentThreadSafety(t *testing.T) {
// This test focuses on thread safety of the RNG, not uniqueness of results
const numGoroutines = 20
const iterations = 100

var wg sync.WaitGroup
var mu sync.Mutex
errors := make([]error, 0)
totalGenerated := 0

for i := range numGoroutines {
wg.Add(1)
go func(goroutineID int) {
defer wg.Done()

localErrors := make([]error, 0)
localCount := 0

for range iterations {
subnet, gateway, err := GenerateNonOverlappingIPv4Subnet([]*net.IPNet{}, 24)

if err != nil {
localErrors = append(localErrors, err)
} else if subnet == nil || gateway == nil {
localErrors = append(localErrors, fmt.Errorf("got nil subnet or gateway"))
} else {
localCount++
}
}

mu.Lock()
errors = append(errors, localErrors...)
totalGenerated += localCount
mu.Unlock()
}(i)
}

wg.Wait()

// Thread safety check - no errors should occur
assert.Empty(t, errors, "Thread-safe RNG should not cause any errors")

expectedTotal := numGoroutines * iterations
assert.Equal(t, expectedTotal, totalGenerated, "Should generate expected total number of subnets")

t.Logf("Successfully generated %d subnets across %d concurrent goroutines without errors",
totalGenerated, numGoroutines)
}
66 changes: 42 additions & 24 deletions network/ipv4.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,16 @@ import (
"fmt"
"math/rand"
"net"
"sync"
)

// rng is the random number generator used for subnet selection.
// Can be overridden in tests for deterministic behavior.
var rng = rand.New(rand.NewSource(rand.Int63()))

Comment thread
jhaynie marked this conversation as resolved.
// rngMu protects concurrent access to rng
var rngMu sync.Mutex

var privateRanges = []struct {
network *net.IPNet
}{
Expand All @@ -32,48 +40,59 @@ func generateRandomSubnet(parent *net.IPNet, prefixLen int) *net.IPNet {
parentIP := parent.IP.Mask(parent.Mask)
parentOnes, _ := parent.Mask.Size()

hostBits := 32 - prefixLen
if hostBits < 0 {
// Check if the requested prefix length is valid
if prefixLen <= parentOnes {
return nil
}

// Calculate how many subnets of the requested size can fit in the parent
subnetBits := prefixLen - parentOnes
maxSubnets := uint32(1) << subnetBits

// Generate random subnet index
rngMu.Lock()
subnetIndex := uint32(rng.Intn(int(maxSubnets)))
rngMu.Unlock()

// Convert parent IP to 32-bit big-endian integer
parentInt := uint32(parentIP[0])<<24 | uint32(parentIP[1])<<16 | uint32(parentIP[2])<<8 | uint32(parentIP[3])

// Mask to parent prefix to get the base network
parentMask := ^uint32(0) << (32 - parentOnes)
baseNetwork := parentInt & parentMask
// Calculate the size of each subnet in terms of IP addresses
subnetSize := uint32(1) << (32 - prefixLen)

// Generate random offset within the available host bits for the new prefix
maxOffset := uint32(1) << hostBits
offset := uint32(rand.Intn(int(maxOffset)))
// Calculate the starting IP of the chosen subnet
newIP := parentInt + (subnetIndex * subnetSize)

// Add offset to base network
newIP := baseNetwork + offset

// Convert back to 4 bytes
randBytes := []byte{
// Convert back to 4 bytes and build IPNet directly
ipBytes := net.IP{
byte(newIP >> 24),
byte(newIP >> 16),
byte(newIP >> 8),
byte(newIP),
}

newSubnet := fmt.Sprintf("%s/%d", net.IP(randBytes).String(), prefixLen)
_, result, _ := net.ParseCIDR(newSubnet)
return result
mask := net.CIDRMask(prefixLen, 32)
return &net.IPNet{IP: ipBytes.Mask(mask), Mask: mask}
}

// GenerateNonOverlappingIPv4Subnet generates a non-overlapping ipv4 subnet with the given prefix size within the given range.
func GenerateNonOverlappingIPv4Subnet(existingNetworks []*net.IPNet, prefixLen int) (*net.IPNet, *net.IP, error) {
// Validate prefix length is within acceptable range
if prefixLen < 8 || prefixLen > 30 {
return nil, nil, fmt.Errorf("invalid prefix length %d: must be between 8 and 30", prefixLen)
if prefixLen < 9 || prefixLen > 30 {
return nil, nil, fmt.Errorf("invalid prefix length %d: must be between 9 and 30", prefixLen)
}

for _, rng := range privateRanges {
// Create a shuffled copy of private ranges for better distribution
ranges := make([]struct{ network *net.IPNet }, len(privateRanges))
copy(ranges, privateRanges)
rngMu.Lock()
rng.Shuffle(len(ranges), func(i, j int) {
ranges[i], ranges[j] = ranges[j], ranges[i]
})
rngMu.Unlock()

for _, rangeSpec := range ranges {
for range 1000 { // Try 1000 times to find non-overlapping subnet
candidate := generateRandomSubnet(rng.network, prefixLen)
candidate := generateRandomSubnet(rangeSpec.network, prefixLen)
if candidate != nil {
hasOverlap := false
for _, existing := range existingNetworks {
Expand All @@ -83,10 +102,9 @@ func GenerateNonOverlappingIPv4Subnet(existingNetworks []*net.IPNet, prefixLen i
}
}
if !hasOverlap {
network := candidate.IP.To4()
ip := make(net.IP, 4)
copy(ip, network)
ip[3] = 0x1 // make the gateway the first ip and copy since we have a shared copy we need to change
copy(ip, candidate.IP.To4())
ip[3] = 0x1 // make the gateway the first ip
return candidate, &ip, nil
}
}
Expand Down
Loading
Loading