diff --git a/brute.go b/brute.go index 37a68be..dde2010 100644 --- a/brute.go +++ b/brute.go @@ -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 { @@ -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 @@ -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 @@ -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 @@ -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 @@ -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 } diff --git a/brute_test.go b/brute_test.go index 71ee637..b4b7757 100644 --- a/brute_test.go +++ b/brute_test.go @@ -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) @@ -25,7 +25,7 @@ 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)) @@ -33,7 +33,7 @@ func TestInsertError(t *testing.T) { } 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") @@ -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)) @@ -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)) @@ -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") @@ -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)) diff --git a/cidranger.go b/cidranger.go index 2e8f118..8fa58c2 100644 --- a/cidranger.go +++ b/cidranger.go @@ -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 @@ -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) } diff --git a/cidranger_test.go b/cidranger_test.go index c1c741e..ca6cf35 100644 --- a/cidranger_test.go +++ b/cidranger_test.go @@ -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) @@ -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) @@ -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) } @@ -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) { @@ -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) @@ -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) } } @@ -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) diff --git a/go.mod b/go.mod index a35ea91..3f1987d 100644 --- a/go.mod +++ b/go.mod @@ -1,8 +1,11 @@ module github.com/yl2chen/cidranger -go 1.13 +go 1.18 + +require github.com/stretchr/testify v1.6.1 require ( - github.com/stretchr/testify v1.6.1 - gopkg.in/yaml.v2 v2.2.2 // indirect + github.com/davecgh/go-spew v1.1.0 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c // indirect ) diff --git a/go.sum b/go.sum index d063842..afe7890 100644 --- a/go.sum +++ b/go.sum @@ -3,13 +3,9 @@ github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSs 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/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= -github.com/stretchr/testify v1.4.0 h1:2E4SXV/wtOkTonXsotYi4li6zVWxYlZuYNCXe9XRJyk= -github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= github.com/stretchr/testify v1.6.1 h1:hDPOHmpOpP40lSULcqw7IrRb/u7w6RpDC9399XyoNd0= github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/yaml.v2 v2.2.2 h1:ZCJp+EgiOT7lHqUV2J862kp8Qj64Jo6az82+3Td9dZw= -gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/trie.go b/trie.go index 31976bd..f4bc093 100644 --- a/trie.go +++ b/trie.go @@ -33,54 +33,66 @@ import ( // prefix trie, use versionedRanger wrapper instead. // // TODO: Implement level-compressed component of the LPC trie. -type prefixTrie struct { - parent *prefixTrie - children []*prefixTrie +type prefixTrie[V any] struct { + parent *prefixTrie[V] + children []*prefixTrie[V] numBitsSkipped uint numBitsHandled uint network rnet.Network entry RangerEntry + value V size int // This is only maintained in the root trie. } // newPrefixTree creates a new prefixTrie. -func newPrefixTree(version rnet.IPVersion) Ranger { +func newPrefixTree[V any](version rnet.IPVersion, defaultValue ...V) Ranger[V] { _, rootNet, _ := net.ParseCIDR("0.0.0.0/0") if version == rnet.IPv6 { _, rootNet, _ = net.ParseCIDR("0::0/0") } - return &prefixTrie{ - children: make([]*prefixTrie, 2, 2), + + var value V + if len(defaultValue) > 0 { + value = defaultValue[0] + } + return &prefixTrie[V]{ + children: make([]*prefixTrie[V], 2, 2), numBitsSkipped: 0, numBitsHandled: 1, network: rnet.NewNetwork(*rootNet), + value: value, } } -func newPathprefixTrie(network rnet.Network, numBitsSkipped uint) *prefixTrie { - path := &prefixTrie{ - children: make([]*prefixTrie, 2, 2), +func newPathprefixTrie[V any](network rnet.Network, numBitsSkipped uint, value V) *prefixTrie[V] { + path := &prefixTrie[V]{ + children: make([]*prefixTrie[V], 2, 2), numBitsSkipped: numBitsSkipped, numBitsHandled: 1, network: network.Masked(int(numBitsSkipped)), + value: value, } return path } -func newEntryTrie(network rnet.Network, entry RangerEntry) *prefixTrie { +func newEntryTrie[V any](network rnet.Network, entry RangerEntry, value V) *prefixTrie[V] { ones, _ := network.IPNet.Mask.Size() - leaf := newPathprefixTrie(network, uint(ones)) + leaf := newPathprefixTrie(network, uint(ones), value) leaf.entry = entry return leaf } // Insert inserts a RangerEntry into prefix trie. -func (p *prefixTrie) Insert(entry RangerEntry) error { +func (p *prefixTrie[V]) Insert(entry RangerEntry, value ...V) error { network := entry.Network() - sizeIncreased, err := p.insert(rnet.NewNetwork(network), entry) + var val V + if len(value) > 0 { + val = value[0] + } + sizeIncreased, err := p.insert(rnet.NewNetwork(network), entry, val) if sizeIncreased { p.size++ } @@ -88,7 +100,7 @@ func (p *prefixTrie) Insert(entry RangerEntry) error { } // Remove removes RangerEntry identified by given network from trie. -func (p *prefixTrie) Remove(network net.IPNet) (RangerEntry, error) { +func (p *prefixTrie[V]) Remove(network net.IPNet) (RangerEntry, error) { entry, err := p.remove(rnet.NewNetwork(network)) if entry != nil { p.size-- @@ -98,7 +110,7 @@ func (p *prefixTrie) Remove(network net.IPNet) (RangerEntry, error) { // Contains returns boolean indicating whether given ip is contained in any // of the inserted networks. -func (p *prefixTrie) Contains(ip net.IP) (bool, error) { +func (p *prefixTrie[V]) Contains(ip net.IP) (bool, error) { nn := rnet.NewNetworkNumber(ip) if nn == nil { return false, ErrInvalidNetworkNumberInput @@ -108,7 +120,7 @@ func (p *prefixTrie) Contains(ip net.IP) (bool, error) { // ContainingNetworks returns the list of RangerEntry(s) the given ip is // contained in in ascending prefix order. -func (p *prefixTrie) ContainingNetworks(ip net.IP) ([]RangerEntry, error) { +func (p *prefixTrie[V]) ContainingNetworks(ip net.IP) ([]RangerEntry, error) { nn := rnet.NewNetworkNumber(ip) if nn == nil { return nil, ErrInvalidNetworkNumberInput @@ -116,22 +128,34 @@ func (p *prefixTrie) ContainingNetworks(ip net.IP) ([]RangerEntry, error) { return p.containingNetworks(nn) } +// IterByIncomingNetworks iterates over all networks that the transmitted IP is included in. +func (p *prefixTrie[V]) IterByIncomingNetworks(ip net.IP, f func(network net.IPNet, value V) error) error { + if err := f(p.network.IPNet, p.value); err != nil { + return err + } + nn := rnet.NewNetworkNumber(ip) + if nn == nil { + return ErrInvalidNetworkNumberInput + } + return p.iterByIncomingNetworks(nn, f) +} + // CoveredNetworks returns the list of RangerEntry(s) the given ipnet // covers. That is, the networks that are completely subsumed by the // specified network. -func (p *prefixTrie) CoveredNetworks(network net.IPNet) ([]RangerEntry, error) { +func (p *prefixTrie[V]) CoveredNetworks(network net.IPNet) ([]RangerEntry, error) { net := rnet.NewNetwork(network) return p.coveredNetworks(net) } // Len returns number of networks in ranger. -func (p *prefixTrie) Len() int { +func (p *prefixTrie[V]) Len() int { return p.size } // String returns string representation of trie, mainly for visualization and // debugging. -func (p *prefixTrie) String() string { +func (p *prefixTrie[V]) String() string { children := []string{} padding := strings.Repeat("| ", p.level()+1) for bits, child := range p.children { @@ -145,7 +169,7 @@ func (p *prefixTrie) String() string { p.targetBitPosition(), p.hasEntry(), strings.Join(children, "")) } -func (p *prefixTrie) contains(number rnet.NetworkNumber) (bool, error) { +func (p *prefixTrie[V]) contains(number rnet.NetworkNumber) (bool, error) { if !p.network.Contains(number) { return false, nil } @@ -166,7 +190,7 @@ func (p *prefixTrie) contains(number rnet.NetworkNumber) (bool, error) { return false, nil } -func (p *prefixTrie) containingNetworks(number rnet.NetworkNumber) ([]RangerEntry, error) { +func (p *prefixTrie[V]) containingNetworks(number rnet.NetworkNumber) ([]RangerEntry, error) { results := []RangerEntry{} if !p.network.Contains(number) { return results, nil @@ -198,7 +222,35 @@ func (p *prefixTrie) containingNetworks(number rnet.NetworkNumber) ([]RangerEntr return results, nil } -func (p *prefixTrie) coveredNetworks(network rnet.Network) ([]RangerEntry, error) { +func (p *prefixTrie[V]) iterByIncomingNetworks(number rnet.NetworkNumber, + f func(network net.IPNet, value V) error) error { + if !p.network.Contains(number) { + return nil + } + + if p.hasEntry() { + if err := f(p.network.IPNet, p.value); err != nil { + return err + } + } + if p.targetBitPosition() < 0 { + return nil + } + bit, err := p.targetBitFromIP(number) + if err != nil { + return err + } + child := p.children[bit] + if child != nil { + err = child.iterByIncomingNetworks(number, f) + if err != nil { + return err + } + } + return nil +} + +func (p *prefixTrie[V]) coveredNetworks(network rnet.Network) ([]RangerEntry, error) { var results []RangerEntry if network.Covers(p.network) { for entry := range p.walkDepth() { @@ -217,7 +269,7 @@ func (p *prefixTrie) coveredNetworks(network rnet.Network) ([]RangerEntry, error return results, nil } -func (p *prefixTrie) insert(network rnet.Network, entry RangerEntry) (bool, error) { +func (p *prefixTrie[V]) insert(network rnet.Network, entry RangerEntry, value V) (bool, error) { if p.network.Equal(network) { sizeIncreased := p.entry == nil p.entry = entry @@ -232,7 +284,7 @@ func (p *prefixTrie) insert(network rnet.Network, entry RangerEntry) (bool, erro // No existing child, insert new leaf trie. if existingChild == nil { - p.appendTrie(bit, newEntryTrie(network, entry)) + p.appendTrie(bit, newEntryTrie(network, entry, value)) return true, nil } @@ -241,7 +293,7 @@ func (p *prefixTrie) insert(network rnet.Network, entry RangerEntry) (bool, erro lcb, err := network.LeastCommonBitPosition(existingChild.network) divergingBitPos := int(lcb) - 1 if divergingBitPos > existingChild.targetBitPosition() { - pathPrefix := newPathprefixTrie(network, p.totalNumberOfBits()-lcb) + pathPrefix := newPathprefixTrie(network, p.totalNumberOfBits()-lcb, value) err := p.insertPrefix(bit, pathPrefix, existingChild) if err != nil { return false, err @@ -249,15 +301,15 @@ func (p *prefixTrie) insert(network rnet.Network, entry RangerEntry) (bool, erro // Update new child existingChild = pathPrefix } - return existingChild.insert(network, entry) + return existingChild.insert(network, entry, value) } -func (p *prefixTrie) appendTrie(bit uint32, prefix *prefixTrie) { +func (p *prefixTrie[V]) appendTrie(bit uint32, prefix *prefixTrie[V]) { p.children[bit] = prefix prefix.parent = p } -func (p *prefixTrie) insertPrefix(bit uint32, pathPrefix, child *prefixTrie) error { +func (p *prefixTrie[V]) insertPrefix(bit uint32, pathPrefix, child *prefixTrie[V]) error { // Set parent/child relationship between current trie and inserted pathPrefix p.children[bit] = pathPrefix pathPrefix.parent = p @@ -272,7 +324,7 @@ func (p *prefixTrie) insertPrefix(bit uint32, pathPrefix, child *prefixTrie) err return nil } -func (p *prefixTrie) remove(network rnet.Network) (RangerEntry, error) { +func (p *prefixTrie[V]) remove(network rnet.Network) (RangerEntry, error) { if p.hasEntry() && p.network.Equal(network) { entry := p.entry p.entry = nil @@ -297,7 +349,7 @@ func (p *prefixTrie) remove(network rnet.Network) (RangerEntry, error) { return nil, nil } -func (p *prefixTrie) qualifiesForPathCompression() bool { +func (p *prefixTrie[V]) qualifiesForPathCompression() bool { // Current prefix trie can be path compressed if it meets all following. // 1. records no CIDR entry // 2. has single or no child @@ -305,14 +357,14 @@ func (p *prefixTrie) qualifiesForPathCompression() bool { return !p.hasEntry() && p.childrenCount() <= 1 && p.parent != nil } -func (p *prefixTrie) compressPathIfPossible() error { +func (p *prefixTrie[V]) compressPathIfPossible() error { if !p.qualifiesForPathCompression() { // Does not qualify to be compressed return nil } // Find lone child. - var loneChild *prefixTrie + var loneChild *prefixTrie[V] for _, child := range p.children { if child != nil { loneChild = child @@ -335,7 +387,7 @@ func (p *prefixTrie) compressPathIfPossible() error { return parent.compressPathIfPossible() } -func (p *prefixTrie) childrenCount() int { +func (p *prefixTrie[V]) childrenCount() int { count := 0 for _, child := range p.children { if child != nil { @@ -345,25 +397,25 @@ func (p *prefixTrie) childrenCount() int { return count } -func (p *prefixTrie) totalNumberOfBits() uint { +func (p *prefixTrie[V]) totalNumberOfBits() uint { return rnet.BitsPerUint32 * uint(len(p.network.Number)) } -func (p *prefixTrie) targetBitPosition() int { +func (p *prefixTrie[V]) targetBitPosition() int { return int(p.totalNumberOfBits()-p.numBitsSkipped) - 1 } -func (p *prefixTrie) targetBitFromIP(n rnet.NetworkNumber) (uint32, error) { +func (p *prefixTrie[V]) targetBitFromIP(n rnet.NetworkNumber) (uint32, error) { // This is a safe uint boxing of int since we should never attempt to get // target bit at a negative position. return n.Bit(uint(p.targetBitPosition())) } -func (p *prefixTrie) hasEntry() bool { +func (p *prefixTrie[V]) hasEntry() bool { return p.entry != nil } -func (p *prefixTrie) level() int { +func (p *prefixTrie[V]) level() int { if p.parent == nil { return 0 } @@ -371,7 +423,7 @@ func (p *prefixTrie) level() int { } // walkDepth walks the trie in depth order, for unit testing. -func (p *prefixTrie) walkDepth() <-chan RangerEntry { +func (p *prefixTrie[V]) walkDepth() <-chan RangerEntry { entries := make(chan RangerEntry) go func() { if p.hasEntry() { diff --git a/trie_test.go b/trie_test.go index 04f2900..d97431e 100644 --- a/trie_test.go +++ b/trie_test.go @@ -2,8 +2,10 @@ package cidranger import ( "encoding/binary" + "fmt" "math/rand" "net" + "reflect" "runtime" "testing" "time" @@ -72,7 +74,7 @@ func TestPrefixTrieInsert(t *testing.T) { } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - trie := newPrefixTree(tc.version).(*prefixTrie) + trie := newPrefixTree[httpHeader](tc.version).(*prefixTrie[httpHeader]) for _, insert := range tc.inserts { _, network, _ := net.ParseCIDR(insert) err := trie.Insert(NewBasicRangerEntry(*network)) @@ -103,7 +105,7 @@ func TestPrefixTrieInsert(t *testing.T) { func TestPrefixTrieString(t *testing.T) { inserts := []string{"192.168.0.1/24", "192.168.1.1/24", "192.168.1.1/30"} - trie := newPrefixTree(rnet.IPv4).(*prefixTrie) + trie := newPrefixTree[httpHeader](rnet.IPv4).(*prefixTrie[httpHeader]) for _, insert := range inserts { _, network, _ := net.ParseCIDR(insert) trie.Insert(NewBasicRangerEntry(*network)) @@ -204,7 +206,7 @@ func TestPrefixTrieRemove(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - trie := newPrefixTree(tc.version).(*prefixTrie) + trie := newPrefixTree[httpHeader](tc.version).(*prefixTrie[httpHeader]) for _, insert := range tc.inserts { _, network, _ := net.ParseCIDR(insert) err := trie.Insert(NewBasicRangerEntry(*network)) @@ -253,13 +255,31 @@ func TestToReplicateIssue(t *testing.T) { inserts []string ip net.IP networks []string + headers [][]httpHeader name string }{ { rnet.IPv4, - []string{"192.168.0.1/32"}, + []string{ + "192.168.0.0/16", + "172.16.0.0/16", + "192.168.0.0/24", + "192.168.0.1/32", + }, net.ParseIP("192.168.0.1"), - []string{"192.168.0.1/32"}, + []string{ + "192.168.0.0/16", + "192.168.0.0/24", + "192.168.0.1/32", + }, + [][]httpHeader{ + { + {Name: "Host", Value: "example.com"}, + }, + { + {Name: "Host", Value: "yandex.ru"}, + }, + }, "basic containing network for /32 mask", }, { @@ -267,15 +287,25 @@ func TestToReplicateIssue(t *testing.T) { []string{"a::1/128"}, net.ParseIP("a::1"), []string{"a::1/128"}, + [][]httpHeader{ + { + {Name: "Host", Value: "example.com"}, + }, + }, "basic containing network for /128 mask", }, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - trie := newPrefixTree(tc.version) - for _, insert := range tc.inserts { + trie := newPrefixTree[[]httpHeader](tc.version) + for i, insert := range tc.inserts { _, network, _ := net.ParseCIDR(insert) - err := trie.Insert(NewBasicRangerEntry(*network)) + var err error + if len(tc.headers) > i { + err = trie.Insert(NewBasicRangerEntry(*network), tc.headers[i]) + } else { + err = trie.Insert(NewBasicRangerEntry(*network)) + } assert.NoError(t, err) } expectedEntries := []RangerEntry{} @@ -293,11 +323,153 @@ func TestToReplicateIssue(t *testing.T) { } } +func TestIterByIncomingNetworks(t *testing.T) { + type NetHeaders struct { + ipNet string + headers []httpHeader + } + + cases := []struct { + version rnet.IPVersion + inserts []NetHeaders + ip net.IP + want []httpHeader + name string + }{ + { + rnet.IPv4, + []NetHeaders{ + { + "172.16.0.0/16", + []httpHeader{ + {Name: "Host", Value: "172.16.0.0/16"}, + }, + }, + { + "192.168.0.0/16", + []httpHeader{ + {Name: "Host", Value: "192.168.0.0/16"}, + }, + }, + { + "192.0.0.0/8", + []httpHeader{ + {Name: "Host", Value: "192.0.0.0/8"}, + }, + }, + { + "172.16.99.0/24", + []httpHeader{ + {Name: "Host", Value: "172.16.99.0/24"}, + }, + }, + { + "192.168.99.0/24", + []httpHeader{ + {Name: "Host", Value: "192.168.99.0/24"}, + }, + }, + { + "192.168.99.1/32", + []httpHeader{ + {Name: "Host", Value: "192.168.99.1/32"}, + }, + }, + }, + net.ParseIP("192.168.99.1"), + []httpHeader{ + {Name: "Host", Value: "192.0.0.0/8"}, + {Name: "Host", Value: "192.168.0.0/16"}, + {Name: "Host", Value: "192.168.99.0/24"}, + {Name: "Host", Value: "192.168.99.1/32"}, + }, + "iterOver4IPv4Networks", + }, + { + rnet.IPv6, + []NetHeaders{ + { + "2001:db8:1234::/48", + []httpHeader{ + {Name: "Host", Value: "2001:db8:1234::/48"}, + }, + }, + { + "2001:db8:1234:5678::/64", + []httpHeader{ + {Name: "Host", Value: "2001:db8:1234:5678::/64"}, + }, + }, + { + "2001:db8:1234:5678:abcd::/80", + []httpHeader{ + {Name: "Host", Value: "2001:db8:1234:5678:abcd::/80"}, + }, + }, + { + "2001:db8:1234:5178:abcd::/80", + []httpHeader{ + {Name: "Host", Value: "2001:db8:1234:5178:abcd::/80"}, + }, + }, + { + "2001:db8:1274:5678::/64", + []httpHeader{ + {Name: "Host", Value: "2001:db8:1274:5678::/64"}, + }, + }, + }, + net.ParseIP("2001:db8:1234:5678:abcd::1"), + []httpHeader{ + {Name: "Host", Value: "2001:db8:1234::/48"}, + {Name: "Host", Value: "2001:db8:1234:5678::/64"}, + {Name: "Host", Value: "2001:db8:1234:5678:abcd::/80"}, + }, + "iterOver3IPv6Networks", + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + trie := newPrefixTree[httpHeader](tc.version) + for _, insert := range tc.inserts { + _, network, _ := net.ParseCIDR(insert.ipNet) + err := trie.Insert(NewBasicRangerEntry(*network), insert.headers...) + assert.NoError(t, err) + } + + f, got := collectHeaders[httpHeader]() + + err := trie.IterByIncomingNetworks(tc.ip, f) + assert.NoError(t, err) + if !reflect.DeepEqual(tc.want, *got) { + t.Errorf("want: \n%v, \ngot: \n%v", tc.want, *got) + } + }) + } +} + +func collectHeaders[V any]() (func(network net.IPNet, value V) error, *[]V) { + got := make([]V, 0) + return func(network net.IPNet, value V) error { + if reflect.ValueOf(value).IsZero() { + fmt.Println(network.String()) + return nil + } + got = append(got, value) + return nil + }, &got +} + type expectedIPRange struct { start net.IP end net.IP } +type httpHeader struct { + Name string + Value string +} + func TestPrefixTrieContains(t *testing.T) { cases := []struct { version rnet.IPVersion @@ -326,7 +498,7 @@ func TestPrefixTrieContains(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - trie := newPrefixTree(tc.version) + trie := newPrefixTree[httpHeader](tc.version) for _, insert := range tc.inserts { _, network, _ := net.ParseCIDR(insert) err := trie.Insert(NewBasicRangerEntry(*network)) @@ -379,7 +551,7 @@ func TestPrefixTrieContainingNetworks(t *testing.T) { } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - trie := newPrefixTree(tc.version) + trie := newPrefixTree[any](tc.version) for _, insert := range tc.inserts { _, network, _ := net.ParseCIDR(insert) err := trie.Insert(NewBasicRangerEntry(*network)) @@ -472,7 +644,7 @@ var coveredNetworkTests = []coveredNetworkTest{ func TestPrefixTrieCoveredNetworks(t *testing.T) { for _, tc := range coveredNetworkTests { t.Run(tc.name, func(t *testing.T) { - trie := newPrefixTree(tc.version) + trie := newPrefixTree[any](tc.version) for _, insert := range tc.inserts { _, network, _ := net.ParseCIDR(insert) err := trie.Insert(NewBasicRangerEntry(*network)) @@ -503,7 +675,7 @@ func TestTrieMemUsage(t *testing.T) { // by threshold, picking 1% as sane number for detecting memory leak. thresh := 1.01 - trie := newPrefixTree(rnet.IPv4) + trie := newPrefixTree[any](rnet.IPv4) var baseLineHeap, totalHeapAllocOverRuns uint64 for i := 0; i < runs; i++ { diff --git a/version.go b/version.go index 2c3fe2b..963e923 100644 --- a/version.go +++ b/version.go @@ -6,30 +6,34 @@ import ( rnet "github.com/yl2chen/cidranger/net" ) -type rangerFactory func(rnet.IPVersion) Ranger +type rangerFactory[V any] func(v rnet.IPVersion, value ...V) Ranger[V] -type versionedRanger struct { - ipV4Ranger Ranger - ipV6Ranger Ranger +type versionedRanger[V any] struct { + ipV4Ranger Ranger[V] + ipV6Ranger Ranger[V] } -func newVersionedRanger(factory rangerFactory) Ranger { - return &versionedRanger{ - ipV4Ranger: factory(rnet.IPv4), - ipV6Ranger: factory(rnet.IPv6), +func newVersionedRanger[V any](factory rangerFactory[V], defaultValue V) Ranger[V] { + return &versionedRanger[V]{ + ipV4Ranger: factory(rnet.IPv4, defaultValue), + ipV6Ranger: factory(rnet.IPv6, defaultValue), } } -func (v *versionedRanger) Insert(entry RangerEntry) error { +func (v *versionedRanger[V]) Insert(entry RangerEntry, value ...V) error { + var val V + if len(value) > 0 { + val = value[0] + } network := entry.Network() ranger, err := v.getRangerForIP(network.IP) if err != nil { return err } - return ranger.Insert(entry) + return ranger.Insert(entry, val) } -func (v *versionedRanger) Remove(network net.IPNet) (RangerEntry, error) { +func (v *versionedRanger[V]) Remove(network net.IPNet) (RangerEntry, error) { ranger, err := v.getRangerForIP(network.IP) if err != nil { return nil, err @@ -37,7 +41,7 @@ func (v *versionedRanger) Remove(network net.IPNet) (RangerEntry, error) { return ranger.Remove(network) } -func (v *versionedRanger) Contains(ip net.IP) (bool, error) { +func (v *versionedRanger[V]) Contains(ip net.IP) (bool, error) { ranger, err := v.getRangerForIP(ip) if err != nil { return false, err @@ -45,7 +49,7 @@ func (v *versionedRanger) Contains(ip net.IP) (bool, error) { return ranger.Contains(ip) } -func (v *versionedRanger) ContainingNetworks(ip net.IP) ([]RangerEntry, error) { +func (v *versionedRanger[V]) ContainingNetworks(ip net.IP) ([]RangerEntry, error) { ranger, err := v.getRangerForIP(ip) if err != nil { return nil, err @@ -53,7 +57,16 @@ func (v *versionedRanger) ContainingNetworks(ip net.IP) ([]RangerEntry, error) { return ranger.ContainingNetworks(ip) } -func (v *versionedRanger) CoveredNetworks(network net.IPNet) ([]RangerEntry, error) { +func (v *versionedRanger[V]) IterByIncomingNetworks(ip net.IP, fn func(network net.IPNet, value V) error) error { + ranger, err := v.getRangerForIP(ip) + if err != nil { + return err + } + + return ranger.IterByIncomingNetworks(ip, fn) +} + +func (v *versionedRanger[V]) CoveredNetworks(network net.IPNet) ([]RangerEntry, error) { ranger, err := v.getRangerForIP(network.IP) if err != nil { return nil, err @@ -62,11 +75,11 @@ func (v *versionedRanger) CoveredNetworks(network net.IPNet) ([]RangerEntry, err } // Len returns number of networks in ranger. -func (v *versionedRanger) Len() int { +func (v *versionedRanger[V]) Len() int { return v.ipV4Ranger.Len() + v.ipV6Ranger.Len() } -func (v *versionedRanger) getRangerForIP(ip net.IP) (Ranger, error) { +func (v *versionedRanger[V]) getRangerForIP(ip net.IP) (Ranger[V], error) { if ip.To4() != nil { return v.ipV4Ranger, nil }