Skip to content
Open
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
26 changes: 16 additions & 10 deletions brute.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,21 +16,21 @@ import (
// testing because the correctness of this implementation can be easily guaranteed,
// and used as the ground truth when running a wider range of 'random' tests on
// other more sophisticated implementations.
type bruteRanger struct {
type bruteRanger[V any] struct {
ipV4Entries map[string]RangerEntry
ipV6Entries map[string]RangerEntry
}

// newBruteRanger returns a new Ranger.
func newBruteRanger() Ranger {
return &bruteRanger{
func newBruteRanger[V any]() Ranger[V] {
return &bruteRanger[V]{
ipV4Entries: make(map[string]RangerEntry),
ipV6Entries: make(map[string]RangerEntry),
}
}

// Insert inserts a RangerEntry into ranger.
func (b *bruteRanger) Insert(entry RangerEntry) error {
func (b *bruteRanger[V]) Insert(entry RangerEntry, value ...V) error {
network := entry.Network()
key := network.String()
if _, found := b.ipV4Entries[key]; !found {
Expand All @@ -44,7 +44,7 @@ func (b *bruteRanger) Insert(entry RangerEntry) error {
}

// Remove removes a RangerEntry identified by given network from ranger.
func (b *bruteRanger) Remove(network net.IPNet) (RangerEntry, error) {
func (b *bruteRanger[V]) Remove(network net.IPNet) (RangerEntry, error) {
networks, err := b.getEntriesByVersion(network.IP)
if err != nil {
return nil, err
Expand All @@ -59,7 +59,7 @@ func (b *bruteRanger) Remove(network net.IPNet) (RangerEntry, error) {

// Contains returns bool indicating whether given ip is contained by any
// network in ranger.
func (b *bruteRanger) Contains(ip net.IP) (bool, error) {
func (b *bruteRanger[V]) Contains(ip net.IP) (bool, error) {
entries, err := b.getEntriesByVersion(ip)
if err != nil {
return false, err
Expand All @@ -74,7 +74,7 @@ func (b *bruteRanger) Contains(ip net.IP) (bool, error) {
}

// ContainingNetworks returns all RangerEntry(s) that given ip contained in.
func (b *bruteRanger) ContainingNetworks(ip net.IP) ([]RangerEntry, error) {
func (b *bruteRanger[V]) ContainingNetworks(ip net.IP) ([]RangerEntry, error) {
entries, err := b.getEntriesByVersion(ip)
if err != nil {
return nil, err
Expand All @@ -89,10 +89,16 @@ func (b *bruteRanger) ContainingNetworks(ip net.IP) ([]RangerEntry, error) {
return results, nil
}

func (b *bruteRanger[V]) IterByIncomingNetworks(ip net.IP, fn func(network net.IPNet, value V) error) error {
// does not need to be implemented
// this Ranger is used for testing only
panic("implement me")
}

// CoveredNetworks returns the list of RangerEntry(s) the given ipnet
// covers. That is, the networks that are completely subsumed by the
// specified network.
func (b *bruteRanger) CoveredNetworks(network net.IPNet) ([]RangerEntry, error) {
func (b *bruteRanger[V]) CoveredNetworks(network net.IPNet) ([]RangerEntry, error) {
entries, err := b.getEntriesByVersion(network.IP)
if err != nil {
return nil, err
Expand All @@ -109,11 +115,11 @@ func (b *bruteRanger) CoveredNetworks(network net.IPNet) ([]RangerEntry, error)
}

// Len returns number of networks in ranger.
func (b *bruteRanger) Len() int {
func (b *bruteRanger[V]) Len() int {
return len(b.ipV4Entries) + len(b.ipV6Entries)
}

func (b *bruteRanger) getEntriesByVersion(ip net.IP) (map[string]RangerEntry, error) {
func (b *bruteRanger[V]) getEntriesByVersion(ip net.IP) (map[string]RangerEntry, error) {
if ip.To4() != nil {
return b.ipV4Entries, nil
}
Expand Down
19 changes: 12 additions & 7 deletions brute_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ import (
)

func TestInsert(t *testing.T) {
ranger := newBruteRanger().(*bruteRanger)
ranger := newBruteRanger[httpHeaders]().(*bruteRanger[httpHeaders])
_, networkIPv4, _ := net.ParseCIDR("0.0.1.0/24")
_, networkIPv6, _ := net.ParseCIDR("8000::/96")
entryIPv4 := NewBasicRangerEntry(*networkIPv4)
Expand All @@ -25,15 +25,15 @@ func TestInsert(t *testing.T) {
}

func TestInsertError(t *testing.T) {
bRanger := newBruteRanger().(*bruteRanger)
bRanger := newBruteRanger[httpHeaders]().(*bruteRanger[httpHeaders])
_, networkIPv4, _ := net.ParseCIDR("0.0.1.0/24")
networkIPv4.IP = append(networkIPv4.IP, byte(4))
err := bRanger.Insert(NewBasicRangerEntry(*networkIPv4))
assert.Equal(t, ErrInvalidNetworkInput, err)
}

func TestRemove(t *testing.T) {
ranger := newBruteRanger().(*bruteRanger)
ranger := newBruteRanger[httpHeaders]().(*bruteRanger[httpHeaders])
_, networkIPv4, _ := net.ParseCIDR("0.0.1.0/24")
_, networkIPv6, _ := net.ParseCIDR("8000::/96")
_, notInserted, _ := net.ParseCIDR("8000::/96")
Expand All @@ -60,7 +60,7 @@ func TestRemove(t *testing.T) {
}

func TestRemoveError(t *testing.T) {
r := newBruteRanger().(*bruteRanger)
r := newBruteRanger[httpHeaders]().(*bruteRanger[httpHeaders])
_, invalidNetwork, _ := net.ParseCIDR("0.0.1.0/24")
invalidNetwork.IP = append(invalidNetwork.IP, byte(4))

Expand All @@ -69,7 +69,7 @@ func TestRemoveError(t *testing.T) {
}

func TestContains(t *testing.T) {
r := newBruteRanger().(*bruteRanger)
r := newBruteRanger[httpHeaders]().(*bruteRanger[httpHeaders])
_, network, _ := net.ParseCIDR("0.0.1.0/24")
_, network1, _ := net.ParseCIDR("8000::/112")
r.Insert(NewBasicRangerEntry(*network))
Expand Down Expand Up @@ -101,8 +101,13 @@ func TestContains(t *testing.T) {
}
}

type httpHeaders struct {
key string
value string
}

func TestContainingNetworks(t *testing.T) {
r := newBruteRanger().(*bruteRanger)
r := newBruteRanger[httpHeaders]().(*bruteRanger[httpHeaders])
_, network1, _ := net.ParseCIDR("0.0.1.0/24")
_, network2, _ := net.ParseCIDR("0.0.1.0/25")
_, network3, _ := net.ParseCIDR("8000::/112")
Expand Down Expand Up @@ -149,7 +154,7 @@ func TestContainingNetworks(t *testing.T) {
func TestCoveredNetworks(t *testing.T) {
for _, tc := range coveredNetworkTests {
t.Run(tc.name, func(t *testing.T) {
ranger := newBruteRanger()
ranger := newBruteRanger[any]()
for _, insert := range tc.inserts {
_, network, _ := net.ParseCIDR(insert)
err := ranger.Insert(NewBasicRangerEntry(*network))
Expand Down
42 changes: 23 additions & 19 deletions cidranger.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,38 +4,37 @@ inclusion tests against it.

To create a new instance of the path-compressed trie:

ranger := NewPCTrieRanger()
ranger := NewPCTrieRanger()

To insert or remove an entry (any object that satisfies the RangerEntry
interface):

_, network, _ := net.ParseCIDR("192.168.0.0/24")
ranger.Insert(NewBasicRangerEntry(*network))
ranger.Remove(network)
_, network, _ := net.ParseCIDR("192.168.0.0/24")
ranger.Insert(NewBasicRangerEntry(*network))
ranger.Remove(network)

If you desire for any value to be attached to the entry, simply
create custom struct that satisfies the RangerEntry interface:

type RangerEntry interface {
Network() net.IPNet
}
type RangerEntry interface {
Network() net.IPNet
}

To test whether an IP is contained in the constructed networks ranger:

// returns bool, error
containsBool, err := ranger.Contains(net.ParseIP("192.168.0.1"))
// returns bool, error
containsBool, err := ranger.Contains(net.ParseIP("192.168.0.1"))

To get a list of CIDR blocks in constructed ranger that contains IP:

// returns []RangerEntry, error
entries, err := ranger.ContainingNetworks(net.ParseIP("192.168.0.1"))
// returns []RangerEntry, error
entries, err := ranger.ContainingNetworks(net.ParseIP("192.168.0.1"))

To get a list of all IPv4/IPv6 rangers respectively:

// returns []RangerEntry, error
entries, err := ranger.CoveredNetworks(*AllIPv4)
entries, err := ranger.CoveredNetworks(*AllIPv6)

// returns []RangerEntry, error
entries, err := ranger.CoveredNetworks(*AllIPv4)
entries, err := ranger.CoveredNetworks(*AllIPv6)
*/
package cidranger

Expand Down Expand Up @@ -83,17 +82,22 @@ func NewBasicRangerEntry(ipNet net.IPNet) RangerEntry {
}

// Ranger is an interface for cidr block containment lookups.
type Ranger interface {
Insert(entry RangerEntry) error
type Ranger[V any] interface {
Insert(entry RangerEntry, value ...V) error
Remove(network net.IPNet) (RangerEntry, error)
Contains(ip net.IP) (bool, error)
ContainingNetworks(ip net.IP) ([]RangerEntry, error)
CoveredNetworks(network net.IPNet) ([]RangerEntry, error)
Len() int
IterByIncomingNetworks(ip net.IP, fn func(network net.IPNet, value V) error) error
}

// NewPCTrieRanger returns a versionedRanger that supports both IPv4 and IPv6
// using the path compressed trie implemention.
func NewPCTrieRanger() Ranger {
return newVersionedRanger(newPrefixTree)
func NewPCTrieRanger[V any](defaultValue ...V) Ranger[V] {
var val V
if len(defaultValue) > 0 {
val = defaultValue[0]
}
return newVersionedRanger[V](newPrefixTree[V], val)
}
61 changes: 33 additions & 28 deletions cidranger_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,10 +49,11 @@ func testContainsAgainstBase(t *testing.T, iterations int, ipGen ipGenerator) {
if testing.Short() {
t.Skip("Skipping memory test in `-short` mode")
}
rangers := []Ranger{NewPCTrieRanger()}
baseRanger := newBruteRanger()
var a any
rangers := []Ranger[any]{NewPCTrieRanger[any](a)}
baseRanger := newBruteRanger[any]()
for _, ranger := range rangers {
configureRangerWithAWSRanges(t, ranger)
configureRangerWithAWSRanges[any](t, ranger)
}
configureRangerWithAWSRanges(t, baseRanger)

Expand All @@ -72,10 +73,11 @@ func testContainingNetworksAgainstBase(t *testing.T, iterations int, ipGen ipGen
if testing.Short() {
t.Skip("Skipping memory test in `-short` mode")
}
rangers := []Ranger{NewPCTrieRanger()}
baseRanger := newBruteRanger()
var a any
rangers := []Ranger[any]{NewPCTrieRanger[any](a)}
baseRanger := newBruteRanger[any]()
for _, ranger := range rangers {
configureRangerWithAWSRanges(t, ranger)
configureRangerWithAWSRanges[any](t, ranger)
}
configureRangerWithAWSRanges(t, baseRanger)

Expand All @@ -98,8 +100,9 @@ func testCoversNetworksAgainstBase(t *testing.T, iterations int, netGen networkG
if testing.Short() {
t.Skip("Skipping memory test in `-short` mode")
}
rangers := []Ranger{NewPCTrieRanger()}
baseRanger := newBruteRanger()
var a any
rangers := []Ranger[any]{NewPCTrieRanger[any](a)}
baseRanger := newBruteRanger[any]()
for _, ranger := range rangers {
configureRangerWithAWSRanges(t, ranger)
}
Expand Down Expand Up @@ -127,59 +130,60 @@ func testCoversNetworksAgainstBase(t *testing.T, iterations int, netGen networkG
*/

func BenchmarkPCTrieHitIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("52.95.110.1"), NewPCTrieRanger())

benchmarkContainsUsingAWSRanges(b, net.ParseIP("52.95.110.1"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerHitIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("52.95.110.1"), newBruteRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("52.95.110.1"), newBruteRanger[any]())
}

func BenchmarkPCTrieHitIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), NewPCTrieRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerHitIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), newBruteRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), newBruteRanger[any]())
}

func BenchmarkPCTrieMissIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("123.123.123.123"), NewPCTrieRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("123.123.123.123"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerMissIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("123.123.123.123"), newBruteRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("123.123.123.123"), newBruteRanger[any]())
}

func BenchmarkPCTrieHMissIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620::ffff"), NewPCTrieRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620::ffff"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerMissIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620::ffff"), newBruteRanger())
benchmarkContainsUsingAWSRanges(b, net.ParseIP("2620::ffff"), newBruteRanger[any]())
}

func BenchmarkPCTrieHitContainingNetworksIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("52.95.110.1"), NewPCTrieRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("52.95.110.1"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerHitContainingNetworksIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("52.95.110.1"), newBruteRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("52.95.110.1"), newBruteRanger[any]())
}

func BenchmarkPCTrieHitContainingNetworksIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), NewPCTrieRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerHitContainingNetworksIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), newBruteRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620:107:300f::36b7:ff81"), newBruteRanger[any]())
}

func BenchmarkPCTrieMissContainingNetworksIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("123.123.123.123"), NewPCTrieRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("123.123.123.123"), NewPCTrieRanger[httpHeaders]())
}
func BenchmarkBruteRangerMissContainingNetworksIPv4UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("123.123.123.123"), newBruteRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("123.123.123.123"), newBruteRanger[httpHeaders]())
}

func BenchmarkPCTrieHMissContainingNetworksIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620::ffff"), NewPCTrieRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620::ffff"), NewPCTrieRanger[any]())
}
func BenchmarkBruteRangerMissContainingNetworksIPv6UsingAWSRanges(b *testing.B) {
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620::ffff"), newBruteRanger())
benchmarkContainingNetworksUsingAWSRanges(b, net.ParseIP("2620::ffff"), newBruteRanger[any]())
}

func BenchmarkNewPathprefixTriev4(b *testing.B) {
Expand All @@ -190,14 +194,14 @@ func BenchmarkNewPathprefixTriev6(b *testing.B) {
benchmarkNewPathprefixTrie(b, "8000::/24")
}

func benchmarkContainsUsingAWSRanges(tb testing.TB, nn net.IP, ranger Ranger) {
func benchmarkContainsUsingAWSRanges[V any](tb testing.TB, nn net.IP, ranger Ranger[V]) {
configureRangerWithAWSRanges(tb, ranger)
for n := 0; n < tb.(*testing.B).N; n++ {
ranger.Contains(nn)
}
}

func benchmarkContainingNetworksUsingAWSRanges(tb testing.TB, nn net.IP, ranger Ranger) {
func benchmarkContainingNetworksUsingAWSRanges[V any](tb testing.TB, nn net.IP, ranger Ranger[V]) {
configureRangerWithAWSRanges(tb, ranger)
for n := 0; n < tb.(*testing.B).N; n++ {
ranger.ContainingNetworks(nn)
Expand All @@ -210,10 +214,11 @@ func benchmarkNewPathprefixTrie(b *testing.B, net1 string) {

n1 := rnet.NewNetwork(*ipNet1)
uOnes := uint(ones)
var a any

b.ResetTimer()
for n := 0; n < b.N; n++ {
newPathprefixTrie(n1, uOnes)
newPathprefixTrie[any](n1, uOnes, a)
}
}

Expand Down Expand Up @@ -286,7 +291,7 @@ func loadAWSRanges() *AWSRanges {
return &ranges
}

func configureRangerWithAWSRanges(tb testing.TB, ranger Ranger) {
func configureRangerWithAWSRanges[V any](tb testing.TB, ranger Ranger[V]) {
for _, prefix := range awsRanges.Prefixes {
_, network, err := net.ParseCIDR(prefix.IPPrefix)
assert.NoError(tb, err)
Expand Down
Loading