refactor(go.d/iprange): migrate from net to net/netip (#20636)
Ilya Mashchenko committed
Jul 5, 2025 at 16:23 UTC
a032a85e2c88ab2bab82ea46ba3cb88ea04cfc70
17 files changed
+1728
-510
src/go/go.mod
-1
@@ -8,7 +8,6 @@ require (
8
github.com/DATA-DOG/go-sqlmock v1.5.2
9
github.com/Masterminds/sprig/v3 v3.3.0
10
github.com/Wing924/ltsv v0.4.0
11
- github.com/apparentlymart/go-cidr v1.1.0
11
github.com/araddon/dateparse v0.0.0-20210429162001-6b43995a97de
12
github.com/axiomhq/hyperloglog v0.2.5
13
github.com/blang/semver/v4 v4.0.0
src/go/go.sum
-2
@@ -39,8 +39,6 @@ github.com/alecthomas/units v0.0.0-20240927000941-0f3dac36c52b h1:mimo19zliBX/vS
39
github.com/alecthomas/units v0.0.0-20240927000941-0f3dac36c52b/go.mod h1:fvzegU4vN3H1qMT+8wDmzjAcDONcgo2/SZ/TyfdUOFs=
40
github.com/alexbrainman/sspi v0.0.0-20231016080023-1a75b4708caa h1:LHTHcTQiSGT7VVbI0o4wBRNQIgn917usHWOd6VAffYI=
41
github.com/alexbrainman/sspi v0.0.0-20231016080023-1a75b4708caa/go.mod h1:cEWa1LVoE5KvSD9ONXsZrj0z6KqySlCCNKHlLzbqAt4=
42
-github.com/apparentlymart/go-cidr v1.1.0 h1:2mAhrMoF+nhXqxTzSZMUzDHkLjmIHC+Zzn4tdgBZjnU=
43
-github.com/apparentlymart/go-cidr v1.1.0/go.mod h1:EBcsNrHc3zQeuaeCeCtQruQm+n9/YjEn/vI25Lg7Gwc=
42
github.com/araddon/dateparse v0.0.0-20210429162001-6b43995a97de h1:FxWPpzIjnTlhPwqqXc4/vE0f7GvRjuAsbW+HOIe8KnA=
43
github.com/araddon/dateparse v0.0.0-20210429162001-6b43995a97de/go.mod h1:DCaWoUhZrYW9p1lxo/cm8EmUOOzAPSEZNGF2DK1dJgw=
44
github.com/aws/aws-sdk-go v1.55.6 h1:cSg4pvZ3m8dgYcgqB97MrcdjUmZ1BeMYKUxMMB89IPk=
src/go/plugin/go.d/collector/dnsmasq_dhcp/collect.go
+13
-13
@@ -9,7 +9,7 @@ import (
9
"io"
10
"math"
11
"math/big"
12
- "net"
12
+ "net/netip"
13
"os"
14
"strings"
15
"time"
@@ -70,7 +70,7 @@ func (c *Collector) collectV4V6Stats() {
70
71
c.mx["ipv4_dhcp_hosts"], c.mx["ipv6_dhcp_hosts"] = 0, 0
72
for _, ip := range c.dhcpHosts {
73
- if ip.To4() == nil {
73
+ if ip.Is6() {
74
c.mx["ipv6_dhcp_hosts"]++
75
} else {
76
c.mx["ipv4_dhcp_hosts"]++
@@ -78,24 +78,24 @@ func (c *Collector) collectV4V6Stats() {
78
}
79
}
80
81
-func (c *Collector) collectRangesStats(leases []net.IP) {
81
+func (c *Collector) collectRangesStats(leases []netip.Addr) {
82
for _, r := range c.dhcpRanges {
83
c.mx["dhcp_range_"+r.String()+"_allocated_leases"] = 0
84
c.mx["dhcp_range_"+r.String()+"_utilization"] = 0
85
}
86
87
- for _, ip := range leases {
87
+ for _, addr := range leases {
88
for _, r := range c.dhcpRanges {
89
- if r.Contains(ip) {
89
+ if r.Contains(addr) {
90
c.mx["dhcp_range_"+r.String()+"_allocated_leases"]++
91
break
92
}
93
}
94
}
95
96
- for _, ip := range c.dhcpHosts {
96
+ for _, addr := range c.dhcpHosts {
97
for _, r := range c.dhcpRanges {
98
- if r.Contains(ip) {
98
+ if r.Contains(addr) {
99
c.mx["dhcp_range_"+r.String()+"_allocated_leases"]++
100
break
101
}
@@ -134,13 +134,13 @@ func (c *Collector) updateCharts() bool {
134
return updated
135
}
136
137
-func findLeases(r io.Reader) []net.IP {
137
+func findLeases(r io.Reader) []netip.Addr {
138
/*
139
1560300536 08:00:27:61:3c:ee 2.2.2.3 debian8 *
140
duid 00:01:00:01:24:90:cf:5b:08:00:27:61:2e:2c
141
1560300414 660684014 1234::20b * 00:01:00:01:24:90:cf:a3:08:00:27:61:3c:ee
142
*/
143
- var ips []net.IP
143
+ var addrs []netip.Addr
144
s := bufio.NewScanner(r)
145
146
for s.Scan() {
@@ -149,14 +149,14 @@ func findLeases(r io.Reader) []net.IP {
149
continue
150
}
151
152
- ip := net.ParseIP(parts[2])
153
- if ip == nil {
152
+ addr, err := netip.ParseAddr(parts[2])
153
+ if err != nil {
154
continue
155
}
156
- ips = append(ips, ip)
156
+ addrs = append(addrs, addr)
157
}
158
159
- return ips
159
+ return addrs
160
}
161
162
func calcPercent(ips int64, hosts *big.Int) float64 {
src/go/plugin/go.d/collector/dnsmasq_dhcp/collector.go
+2
-2
@@ -9,7 +9,7 @@ import (
9
_ "embed"
10
"errors"
11
"fmt"
12
- "net"
12
+ "net/netip"
13
"time"
14
15
"github.com/netdata/netdata/go/plugins/plugin/go.d/agent/module"
@@ -59,7 +59,7 @@ type Collector struct {
59
parseConfigTime time.Time
60
parseConfigEvery time.Duration
61
dhcpRanges []iprange.Range
62
- dhcpHosts []net.IP
62
+ dhcpHosts []netip.Addr
63
cacheDHCPRanges map[string]bool
64
65
mx map[string]int64
src/go/plugin/go.d/collector/dnsmasq_dhcp/parse_configuration.go
+13
-12
@@ -7,7 +7,7 @@ package dnsmasq_dhcp
7
import (
8
"bufio"
9
"fmt"
10
- "net"
10
+ "net/netip"
11
"os"
12
"path/filepath"
13
"regexp"
@@ -17,7 +17,7 @@ import (
17
"github.com/netdata/netdata/go/plugins/plugin/go.d/pkg/iprange"
18
)
19
20
-func (c *Collector) parseDnsmasqDHCPConfiguration() ([]iprange.Range, []net.IP) {
20
+func (c *Collector) parseDnsmasqDHCPConfiguration() ([]iprange.Range, []netip.Addr) {
21
configs := findConfigurationFiles(c.ConfPath, c.ConfDir)
22
23
dhcpRanges := c.getDHCPRanges(configs)
@@ -58,8 +58,8 @@ func (c *Collector) getDHCPRanges(configs []*configFile) []iprange.Range {
58
return dhcpRanges
59
}
60
61
-func (c *Collector) getDHCPHosts(configs []*configFile) []net.IP {
62
- var dhcpHosts []net.IP
61
+func (c *Collector) getDHCPHosts(configs []*configFile) []netip.Addr {
62
+ var dhcpHosts []netip.Addr
63
seen := make(map[string]bool)
64
var parsed string
65
@@ -73,14 +73,14 @@ func (c *Collector) getDHCPHosts(configs []*configFile) []net.IP {
73
}
74
seen[parsed] = true
75
76
- v := net.ParseIP(parsed)
77
- if v == nil {
78
- c.Warningf("error on parsing dhcp-host '%s', skipping it", parsed)
76
+ addr, err := netip.ParseAddr(parsed)
77
+ if err != nil {
78
+ c.Warningf("error on parsing dhcp-host '%s': %v, skipping it", parsed, err)
79
continue
80
}
81
82
c.Debugf("adding dhcp-host '%s'", parsed)
83
- dhcpHosts = append(dhcpHosts, v)
83
+ dhcpHosts = append(dhcpHosts, addr)
84
}
85
}
86
return dhcpHosts
@@ -107,15 +107,16 @@ func parseDHCPRangeValue(s string) (r string) {
107
108
s = strings.ReplaceAll(s, " ", "")
109
110
- var start, end net.IP
110
+ var start, end netip.Addr
111
parts := strings.Split(s, ",")
112
113
for _, v := range parts {
114
- if start == nil {
115
- start = net.ParseIP(v)
114
+ if !start.IsValid() {
115
+ start, _ = netip.ParseAddr(v)
116
continue
117
}
118
- if end = net.ParseIP(v); end == nil || iprange.New(start, end) == nil {
118
+
119
+ if end, _ = netip.ParseAddr(v); !end.IsValid() || iprange.New(start, end) == nil {
120
return ""
121
}
122
return fmt.Sprintf("%s-%s", start, end)
src/go/plugin/go.d/collector/isc_dhcpd/collect.go
+1
-1
@@ -62,7 +62,7 @@ func collectPool(collected map[string]int64, pool ipPool, leases []leaseEntry) {
62
63
func calcPoolActiveLeases(pool ipPool, leases []leaseEntry) (num int64) {
64
for _, l := range leases {
65
- if pool.addresses.Contains(l.ip) {
65
+ if pool.addresses.Contains(l.addr) {
66
num++
67
}
68
}
src/go/plugin/go.d/collector/isc_dhcpd/init.go
+2
-2
@@ -15,7 +15,7 @@ import (
15
16
type ipPool struct {
17
name string
18
- addresses iprange.Pool
18
+ addresses *iprange.Pool
19
}
20
21
func (c *Collector) validateConfig() error {
@@ -48,7 +48,7 @@ func (c *Collector) initPools() ([]ipPool, error) {
48
continue
49
}
50
51
- pool := ipPool{name: cfg.Name, addresses: ipRange}
51
+ pool := ipPool{name: cfg.Name, addresses: iprange.NewPool(ipRange...)}
52
pools = append(pools, pool)
53
}
54
src/go/plugin/go.d/collector/isc_dhcpd/parse.go
+14
-10
@@ -7,7 +7,7 @@ package isc_dhcpd
7
import (
8
"bufio"
9
"bytes"
10
- "net"
10
+ "net/netip"
11
"os"
12
)
13
@@ -41,11 +41,11 @@ DHCPv6 prepare declaration:
41
*/
42
43
type leaseEntry struct {
44
- ip net.IP
44
+ addr netip.Addr
45
bindingState string
46
}
47
48
-func (l leaseEntry) hasIP() bool { return l.ip != nil }
48
+func (l leaseEntry) isAddrValid() bool { return l.addr.IsValid() }
49
func (l leaseEntry) hasBindingState() bool { return l.bindingState != "" }
50
51
func parseDHCPdLeasesFile(filepath string) ([]leaseEntry, error) {
@@ -62,21 +62,25 @@ func parseDHCPdLeasesFile(filepath string) ([]leaseEntry, error) {
62
for sc.Scan() {
63
bs := bytes.TrimSpace(sc.Bytes())
64
switch {
65
- case !l.hasIP() && bytes.HasPrefix(bs, []byte("lease")):
65
+ case !l.isAddrValid() && bytes.HasPrefix(bs, []byte("lease")):
66
// "lease 192.168.0.1 {" => "192.168.0.1"
67
s := string(bs)
68
- l.ip = net.ParseIP(s[6 : len(s)-2])
69
- case !l.hasIP() && bytes.HasPrefix(bs, []byte("iaaddr")):
68
+ if addr, err := netip.ParseAddr(s[6 : len(s)-2]); err == nil && addr.IsValid() {
69
+ l.addr = addr
70
+ }
71
+ case !l.isAddrValid() && bytes.HasPrefix(bs, []byte("iaaddr")):
72
// "iaaddr 1985:470:1f0b:c9a::001 {" => "1985:470:1f0b:c9a::001"
73
s := string(bs)
72
- l.ip = net.ParseIP(s[7 : len(s)-2])
73
- case l.hasIP() && !l.hasBindingState() && bytes.HasPrefix(bs, []byte("binding state")):
74
+ if addr, err := netip.ParseAddr(s[7 : len(s)-2]); err == nil && addr.IsValid() {
75
+ l.addr = addr
76
+ }
77
+ case l.isAddrValid() && !l.hasBindingState() && bytes.HasPrefix(bs, []byte("binding state")):
78
// "binding state active;" => "active"
79
s := string(bs)
80
l.bindingState = s[14 : len(s)-1]
81
case bytes.HasPrefix(bs, []byte("}")):
78
- if l.hasIP() && l.hasBindingState() {
79
- leasesSet[l.ip.String()] = l
82
+ if l.isAddrValid() && l.hasBindingState() {
83
+ leasesSet[l.addr.String()] = l
84
}
85
l = leaseEntry{}
86
}
src/go/plugin/go.d/collector/ntpd/collect.go
+9
-7
@@ -4,7 +4,7 @@ package ntpd
4
5
import (
6
"fmt"
7
- "net"
7
+ "net/netip"
8
"strconv"
9
"time"
10
)
@@ -125,18 +125,20 @@ func (c *Collector) findPeers() error {
125
continue
126
}
127
128
- addr, ok := info["srcadr"]
129
- if ip := net.ParseIP(addr); !ok || ip == nil || c.peerIPAddrFilter.Contains(ip) {
128
+ srcAddr, ok := info["srcadr"]
129
+ addr, err := netip.ParseAddr(srcAddr)
130
+
131
+ if !ok || err != nil || !addr.IsValid() || c.peerIPAddrFilter.Contains(addr) {
132
c.Debugf("skipping NTP peer id='%d', srcadr='%s'", id, addr)
133
continue
134
}
135
134
- seen[addr] = true
136
+ seen[addr.String()] = true
137
136
- if !c.peerAddr[addr] {
137
- c.peerAddr[addr] = true
138
+ if !c.peerAddr[addr.String()] {
139
+ c.peerAddr[addr.String()] = true
140
c.Debugf("new NTP peer id='%d', srcadr='%s': creating charts", id, addr)
139
- c.addPeerCharts(addr)
141
+ c.addPeerCharts(addr.String())
142
}
143
144
c.peerIDs = append(c.peerIDs, id)
src/go/plugin/go.d/collector/ntpd/collector.go
+2
-2
@@ -61,7 +61,7 @@ type Collector struct {
61
findPeersEvery time.Duration
62
peerAddr map[string]bool
63
peerIDs []uint16
64
- peerIPAddrFilter iprange.Pool
64
+ peerIPAddrFilter *iprange.Pool
65
}
66
67
func (c *Collector) Configuration() any {
@@ -79,7 +79,7 @@ func (c *Collector) Init(context.Context) error {
79
return fmt.Errorf("error on parsing ip range '%s': %v", txt, err)
80
}
81
82
- c.peerIPAddrFilter = r
82
+ c.peerIPAddrFilter = iprange.NewPool(r...)
83
84
return nil
85
}
src/go/plugin/go.d/pkg/iprange/iterator.go
+27
-24
@@ -4,38 +4,41 @@ package iprange
4
5
import (
6
"iter"
7
- "net"
7
+ "net/netip"
8
)
9
10
-func iterate(r Range) iter.Seq[net.IP] {
11
- ipCopy := make(net.IP, len(r.getStart()))
12
- nextBuf := make(net.IP, len(r.getStart()))
10
+// iterate returns an iterator that yields each IP address in the range.
11
+// It handles both IPv4 and IPv6 ranges efficiently.
12
+func iterate(r Range) iter.Seq[netip.Addr] {
13
+ return func(yield func(netip.Addr) bool) {
14
+ current := r.Start()
15
+ end := r.End()
16
14
- return func(yield func(net.IP) bool) {
15
- for ip := r.getStart(); ip != nil; ip = nextIP(ip, nextBuf) {
16
- copy(ipCopy, ip)
17
- if !yield(ipCopy) {
17
+ // Handle empty or invalid range
18
+ if !current.IsValid() || !end.IsValid() {
19
+ return
20
+ }
21
+
22
+ for {
23
+ // Yield current address
24
+ if !yield(current) {
25
return
26
}
20
- if ip.Equal(r.getEnd()) {
21
- break
27
+
28
+ // Check if we've reached the end
29
+ if current == end {
30
+ return
31
}
23
- }
24
- }
25
-}
32
27
-func nextIP(ip net.IP, buf net.IP) net.IP {
28
- ip = ip.To16()
29
- if ip == nil {
30
- return nil
31
- }
32
- copy(buf, ip)
33
+ // Move to next address
34
+ next := current.Next()
35
+
36
+ // Check for overflow or going past the end
37
+ if !next.IsValid() || next.Compare(end) > 0 {
38
+ return
39
+ }
40
34
- for i := len(buf) - 1; i >= 0; i-- {
35
- buf[i]++
36
- if buf[i] != 0 {
37
- break
41
+ current = next
42
}
43
}
40
- return buf
44
}
src/go/plugin/go.d/pkg/iprange/parse.go
+191
-76
@@ -3,136 +3,251 @@
3
package iprange
4
5
import (
6
- "bytes"
6
+ "errors"
7
"fmt"
8
- "net"
9
- "regexp"
8
+ "math/bits"
9
+ "net/netip"
10
"strings"
11
+)
12
+
13
+var (
14
+ // ErrInvalidSyntax is returned when the input string has invalid syntax.
15
+ ErrInvalidSyntax = errors.New("invalid IP range syntax")
16
12
- "github.com/apparentlymart/go-cidr/cidr"
17
+ // ErrInvalidRange is returned when the range is invalid (e.g., start > end).
18
+ ErrInvalidRange = errors.New("invalid IP range")
19
+
20
+ // ErrMixedAddressFamilies is returned when trying to create a range with mixed IP families.
21
+ ErrMixedAddressFamilies = errors.New("mixed address families in range")
22
)
23
15
-// ParseRanges parses s as a space separated list of IP Ranges, returning the result and an error if any.
16
-// IP Range can be in IPv4 address ("192.0.2.1"), IPv4 range ("192.0.2.0-192.0.2.10")
17
-// IPv4 CIDR ("192.0.2.0/24"), IPv4 subnet mask ("192.0.2.0/255.255.255.0"),
18
-// IPv6 address ("2001:db8::1"), IPv6 range ("2001:db8::-2001:db8::10"),
19
-// or IPv6 CIDR ("2001:db8::/64") form.
20
-// IPv4 CIDR, IPv4 subnet mask and IPv6 CIDR ranges don't include network and broadcast addresses.
24
+// ParseRanges parses s as a space-separated list of IP ranges.
25
+// Each range can be in one of the following formats:
26
+// - IPv4 address: "192.0.2.1"
27
+// - IPv4 range: "192.0.2.0-192.0.2.10"
28
+// - IPv4 CIDR: "192.0.2.0/24"
29
+// - IPv4 subnet mask: "192.0.2.0/255.255.255.0"
30
+// - IPv6 address: "2001:db8::1"
31
+// - IPv6 range: "2001:db8::-2001:db8::10"
32
+// - IPv6 CIDR: "2001:db8::/64"
33
+//
34
+// For CIDR notations, network and broadcast addresses are excluded
35
+// (except for /31, /32, /127, and /128 prefixes).
36
func ParseRanges(s string) ([]Range, error) {
22
- parts := strings.Fields(s)
23
- if len(parts) == 0 {
37
+ s = strings.TrimSpace(s)
38
+ if s == "" {
39
return nil, nil
40
}
41
27
- var ranges []Range
28
- for _, v := range parts {
29
- r, err := ParseRange(v)
42
+ parts := strings.Fields(s)
43
+ ranges := make([]Range, 0, len(parts))
44
+
45
+ for _, part := range parts {
46
+ r, err := ParseRange(part)
47
if err != nil {
31
- return nil, err
48
+ return nil, fmt.Errorf("parsing %q: %w", part, err)
49
}
33
-
50
if r != nil {
51
ranges = append(ranges, r)
52
}
53
}
54
+
55
return ranges, nil
56
}
57
41
-var (
42
- reRange = regexp.MustCompile("^[0-9a-f.:-]+$") // addr | addr-addr
43
- reCIDR = regexp.MustCompile("^[0-9a-f.:]+/[0-9]{1,3}$") // addr/prefix_length
44
- reSubnetMask = regexp.MustCompile("^[0-9.]+/[0-9.]{7,}$") // v4_addr/mask
45
-)
46
-
47
-// ParseRange parses s as an IP Range, returning the result and an error if any.
48
-// The string s can be in IPv4 address ("192.0.2.1"), IPv4 range ("192.0.2.0-192.0.2.10")
49
-// IPv4 CIDR ("192.0.2.0/24"), IPv4 subnet mask ("192.0.2.0/255.255.255.0"),
50
-// IPv6 address ("2001:db8::1"), IPv6 range ("2001:db8::-2001:db8::10"),
51
-// or IPv6 CIDR ("2001:db8::/64") form.
52
-// IPv4 CIDR, IPv4 subnet mask and IPv6 CIDR ranges don't include network and broadcast addresses.
58
+// ParseRange parses s as a single IP range.
59
+// See ParseRanges for supported formats.
60
func ParseRange(s string) (Range, error) {
54
- s = strings.ToLower(s)
61
+ s = strings.TrimSpace(strings.ToLower(s))
62
if s == "" {
63
return nil, nil
64
}
65
59
- var r Range
66
+ // Try different formats in order of likelihood
67
switch {
61
- case reRange.MatchString(s):
62
- r = parseRange(s)
63
- case reCIDR.MatchString(s):
64
- r = parseCIDR(s)
65
- case reSubnetMask.MatchString(s):
66
- r = parseSubnetMask(s)
68
+ case strings.Contains(s, "-"):
69
+ return parseIPRange(s)
70
+ case strings.Contains(s, "/"):
71
+ if strings.Count(s, ".") >= 3 && strings.Contains(s[strings.LastIndex(s, "/"):], ".") {
72
+ return parseSubnetMask(s)
73
+ }
74
+ return parseCIDR(s)
75
+ default:
76
+ return parseSingleIP(s)
77
}
78
+}
79
69
- if r == nil {
70
- return nil, fmt.Errorf("ip range (%s) invalid syntax", s)
80
+// parseSingleIP parses a single IP address as a range containing only that address.
81
+func parseSingleIP(s string) (Range, error) {
82
+ addr, err := netip.ParseAddr(s)
83
+ if err != nil {
84
+ return nil, fmt.Errorf("%w: %v", ErrInvalidSyntax, err)
85
}
72
- return r, nil
86
+ return New(addr, addr), nil
87
}
88
75
-func parseRange(s string) Range {
76
- var start, end net.IP
77
- if idx := strings.IndexByte(s, '-'); idx != -1 {
78
- start, end = net.ParseIP(s[:idx]), net.ParseIP(s[idx+1:])
79
- } else {
80
- start, end = net.ParseIP(s), net.ParseIP(s)
89
+// parseIPRange parses an IP range in "start-end" format.
90
+func parseIPRange(s string) (Range, error) {
91
+ parts := strings.SplitN(s, "-", 2)
92
+ if len(parts) != 2 {
93
+ return nil, fmt.Errorf("%w: invalid range format", ErrInvalidSyntax)
94
+ }
95
+
96
+ start, err := netip.ParseAddr(strings.TrimSpace(parts[0]))
97
+ if err != nil {
98
+ return nil, fmt.Errorf("%w: invalid start address: %v", ErrInvalidSyntax, err)
99
+ }
100
+
101
+ end, err := netip.ParseAddr(strings.TrimSpace(parts[1]))
102
+ if err != nil {
103
+ return nil, fmt.Errorf("%w: invalid end address: %v", ErrInvalidSyntax, err)
104
+ }
105
+
106
+ // Validate the range
107
+ if start.Is4() != end.Is4() {
108
+ return nil, ErrMixedAddressFamilies
109
+ }
110
+
111
+ if start.Compare(end) > 0 {
112
+ return nil, fmt.Errorf("%w: start address is greater than end address", ErrInvalidRange)
113
}
114
83
- return New(start, end)
115
+ return New(start, end), nil
116
}
117
86
-func parseCIDR(s string) Range {
87
- ip, network, err := net.ParseCIDR(s)
118
+// parseCIDR parses an IP range in CIDR notation.
119
+func parseCIDR(s string) (Range, error) {
120
+ prefix, err := netip.ParsePrefix(s)
121
if err != nil {
89
- return nil
122
+ return nil, fmt.Errorf("%w: %v", ErrInvalidSyntax, err)
123
}
124
92
- start, end := cidr.AddressRange(network)
93
- prefixLen, _ := network.Mask.Size()
125
+ // Normalize to network address
126
+ prefix = prefix.Masked()
127
95
- if isV4IP(ip) && prefixLen < 31 || isV6IP(ip) && prefixLen < 127 {
96
- start = cidr.Inc(start)
97
- end = cidr.Dec(end)
128
+ start := prefix.Addr()
129
+ end := lastAddrInPrefix(prefix)
130
+
131
+ // Exclude network and broadcast addresses for typical subnets
132
+ // Keep all addresses for /31, /32 (IPv4) and /127, /128 (IPv6)
133
+ prefixLen := prefix.Bits()
134
+ if shouldExcludeNetworkAndBroadcast(start, prefixLen) {
135
+ newStart := start.Next()
136
+ if newStart.IsValid() && newStart.Compare(end) <= 0 {
137
+ start = newStart
138
+ end = end.Prev()
139
+ }
140
}
141
100
- return parseRange(fmt.Sprintf("%s-%s", start, end))
142
+ return New(start, end), nil
143
}
144
103
-func parseSubnetMask(s string) Range {
104
- idx := strings.LastIndexByte(s, '/')
145
+// parseSubnetMask parses an IPv4 range with subnet mask notation.
146
+func parseSubnetMask(s string) (Range, error) {
147
+ idx := strings.LastIndex(s, "/")
148
if idx == -1 {
106
- return nil
149
+ return nil, fmt.Errorf("%w: invalid subnet mask format", ErrInvalidSyntax)
150
}
151
109
- address, mask := s[:idx], s[idx+1:]
152
+ addrStr := s[:idx]
153
+ maskStr := s[idx+1:]
154
111
- ip := net.ParseIP(mask).To4()
112
- if ip == nil {
113
- return nil
155
+ // Validate the address part
156
+ addr, err := netip.ParseAddr(addrStr)
157
+ if err != nil || !addr.Is4() {
158
+ return nil, fmt.Errorf("%w: invalid IPv4 address in subnet mask notation", ErrInvalidSyntax)
159
}
160
116
- prefixLen, bits := net.IPv4Mask(ip[0], ip[1], ip[2], ip[3]).Size()
117
- if prefixLen+bits == 0 {
118
- return nil
161
+ // Parse and validate the mask
162
+ mask, err := netip.ParseAddr(maskStr)
163
+ if err != nil || !mask.Is4() {
164
+ return nil, fmt.Errorf("%w: invalid subnet mask", ErrInvalidSyntax)
165
}
166
121
- return parseCIDR(fmt.Sprintf("%s/%d", address, prefixLen))
122
-}
167
+ // Convert mask to prefix length
168
+ prefixLen, ok := maskToPrefixLen(mask)
169
+ if !ok {
170
+ return nil, fmt.Errorf("%w: invalid subnet mask (not contiguous)", ErrInvalidSyntax)
171
+ }
172
124
-func isV4RangeValid(start, end net.IP) bool {
125
- return isV4IP(start) && isV4IP(end) && bytes.Compare(end, start) >= 0
173
+ // Create CIDR notation and parse it
174
+ return parseCIDR(fmt.Sprintf("%s/%d", addrStr, prefixLen))
175
}
176
128
-func isV6RangeValid(start, end net.IP) bool {
129
- return isV6IP(start) && isV6IP(end) && bytes.Compare(end, start) >= 0
177
+// shouldExcludeNetworkAndBroadcast determines if network and broadcast addresses
178
+// should be excluded from the range based on the prefix length.
179
+func shouldExcludeNetworkAndBroadcast(addr netip.Addr, prefixLen int) bool {
180
+ if addr.Is4() {
181
+ return prefixLen < 31
182
+ }
183
+ return prefixLen < 127
184
}
185
132
-func isV4IP(ip net.IP) bool {
133
- return ip.To4() != nil
186
+// lastAddrInPrefix returns the last address in the given prefix.
187
+func lastAddrInPrefix(p netip.Prefix) netip.Addr {
188
+ addr := p.Addr()
189
+ prefixBits := p.Bits()
190
+
191
+ if addr.Is4() {
192
+ a := addr.As4()
193
+ addrUint := uint32(a[0])<<24 | uint32(a[1])<<16 | uint32(a[2])<<8 | uint32(a[3])
194
+
195
+ // Set all host bits to 1
196
+ hostBits := 32 - prefixBits
197
+ if hostBits == 0 {
198
+ return addr
199
+ }
200
+ hostMask := (uint32(1) << hostBits) - 1
201
+ lastUint := addrUint | hostMask
202
+
203
+ return netip.AddrFrom4([4]byte{
204
+ byte(lastUint >> 24),
205
+ byte(lastUint >> 16),
206
+ byte(lastUint >> 8),
207
+ byte(lastUint),
208
+ })
209
+ }
210
+
211
+ // IPv6
212
+ a := addr.As16()
213
+ var last [16]byte
214
+ copy(last[:], a[:])
215
+
216
+ // Set all host bits to 1
217
+ for i := prefixBits / 8; i < 16; i++ {
218
+ if i == prefixBits/8 && prefixBits%8 != 0 {
219
+ // Partial byte: set remaining bits to 1
220
+ last[i] |= byte((1 << (8 - prefixBits%8)) - 1)
221
+ } else {
222
+ // Full byte: set all bits to 1
223
+ last[i] = 0xFF
224
+ }
225
+ }
226
+
227
+ return netip.AddrFrom16(last)
228
}
229
136
-func isV6IP(ip net.IP) bool {
137
- return !isV4IP(ip) && ip.To16() != nil
230
+// maskToPrefixLen converts an IPv4 subnet mask to a prefix length.
231
+// It returns the prefix length and whether the mask is valid (contiguous ones followed by zeros).
232
+func maskToPrefixLen(mask netip.Addr) (int, bool) {
233
+ if !mask.Is4() {
234
+ return 0, false
235
+ }
236
+
237
+ m := mask.As4()
238
+ maskUint := uint32(m[0])<<24 | uint32(m[1])<<16 | uint32(m[2])<<8 | uint32(m[3])
239
+
240
+ // Count the number of 1 bits
241
+ ones := bits.OnesCount32(maskUint)
242
+
243
+ // Valid netmask must have all 1s followed by all 0s
244
+ // This means if we have 'ones' 1-bits, they must all be leading bits
245
+ // So the mask should equal ^((1 << (32 - ones)) - 1)
246
+ expectedMask := ^uint32(0) << (32 - ones)
247
+
248
+ if maskUint != expectedMask {
249
+ return 0, false
250
+ }
251
+
252
+ return ones, true
253
}
src/go/plugin/go.d/pkg/iprange/parse_test.go
+298
-117
@@ -3,256 +3,437 @@
3
package iprange
4
5
import (
6
- "fmt"
7
- "net"
6
+ "errors"
7
+ "net/netip"
8
"testing"
9
10
"github.com/stretchr/testify/assert"
11
+ "github.com/stretchr/testify/require"
12
)
13
14
func TestParseRanges(t *testing.T) {
14
- tests := map[string]struct {
15
+ t.Parallel()
16
+
17
+ tests := []struct {
18
+ name string
19
input string
20
wantRanges []Range
21
wantErr bool
22
}{
19
- "single range": {
23
+ {
24
+ name: "empty string",
25
+ input: "",
26
+ },
27
+ {
28
+ name: "whitespace only",
29
+ input: " \t\n ",
30
+ },
31
+ {
32
+ name: "single range",
33
input: "192.0.2.0-192.0.2.10",
34
wantRanges: []Range{
22
- prepareRange("192.0.2.0", "192.0.2.10"),
35
+ mustParseRange(t, "192.0.2.0", "192.0.2.10"),
36
},
37
},
25
- "multiple ranges": {
38
+ {
39
+ name: "multiple ranges with different formats",
40
input: "2001:db8::0 192.0.2.0-192.0.2.10 2001:db8::0/126 192.0.2.0/255.255.255.0",
41
wantRanges: []Range{
28
- prepareRange("2001:db8::0", "2001:db8::0"),
29
- prepareRange("192.0.2.0", "192.0.2.10"),
30
- prepareRange("2001:db8::1", "2001:db8::2"),
31
- prepareRange("192.0.2.1", "192.0.2.254"),
42
+ mustParseRange(t, "2001:db8::0", "2001:db8::0"),
43
+ mustParseRange(t, "192.0.2.0", "192.0.2.10"),
44
+ mustParseRange(t, "2001:db8::1", "2001:db8::2"),
45
+ mustParseRange(t, "192.0.2.1", "192.0.2.254"),
46
},
47
},
34
- "single invalid syntax": {
48
+ {
49
+ name: "single invalid syntax",
50
input: "192.0.2.0-192.0.2.",
51
wantErr: true,
52
},
38
- "multiple invalid syntax": {
53
+ {
54
+ name: "multiple with one invalid",
55
input: "2001:db8::0 192.0.2.0-192.0.2.10 2001:db8::0/999 192.0.2.0/255.255.255.0",
56
wantErr: true,
57
},
58
+ {
59
+ name: "extra whitespace",
60
+ input: " 192.0.2.0 192.0.2.1-192.0.2.2 ",
61
+ wantRanges: []Range{
62
+ mustParseRange(t, "192.0.2.0", "192.0.2.0"),
63
+ mustParseRange(t, "192.0.2.1", "192.0.2.2"),
64
+ },
65
+ },
66
}
67
44
- for name, test := range tests {
45
- t.Run(name, func(t *testing.T) {
46
- rs, err := ParseRanges(test.input)
68
+ for _, tt := range tests {
69
+ t.Run(tt.name, func(t *testing.T) {
70
+ t.Parallel()
71
+
72
+ ranges, err := ParseRanges(tt.input)
73
48
- if test.wantErr {
74
+ if tt.wantErr {
75
assert.Error(t, err)
50
- assert.Nilf(t, rs, "want: nil, got: %s", rs)
76
+ assert.Nil(t, ranges)
77
} else {
78
assert.NoError(t, err)
53
- assert.Equalf(t, test.wantRanges, rs, "want: %s, got: %s", test.wantRanges, rs)
79
+ assert.Equal(t, tt.wantRanges, ranges)
80
}
81
})
82
}
83
}
84
85
func TestParseRange(t *testing.T) {
60
- tests := map[string]struct {
86
+ t.Parallel()
87
+
88
+ tests := []struct {
89
+ name string
90
input string
91
wantRange Range
63
- wantErr bool
92
+ wantErr error
93
}{
65
- "v4 IP": {
94
+ // Empty input
95
+ {
96
+ name: "empty string",
97
+ input: "",
98
+ },
99
+ {
100
+ name: "whitespace only",
101
+ input: " ",
102
+ },
103
+
104
+ // IPv4 single IP
105
+ {
106
+ name: "IPv4 single IP",
107
input: "192.0.2.0",
67
- wantRange: prepareRange("192.0.2.0", "192.0.2.0"),
108
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.0"),
109
},
69
- "v4 IP: invalid address": {
110
+ {
111
+ name: "IPv4 invalid address",
112
input: "192.0.2.",
71
- wantErr: true,
113
+ wantErr: ErrInvalidSyntax,
114
},
73
- "v4 Range": {
115
+
116
+ // IPv4 ranges
117
+ {
118
+ name: "IPv4 range",
119
input: "192.0.2.0-192.0.2.10",
75
- wantRange: prepareRange("192.0.2.0", "192.0.2.10"),
120
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.10"),
121
},
77
- "v4 Range: start == end": {
122
+ {
123
+ name: "IPv4 range start equals end",
124
input: "192.0.2.0-192.0.2.0",
79
- wantRange: prepareRange("192.0.2.0", "192.0.2.0"),
125
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.0"),
126
},
81
- "v4 Range: start > end": {
127
+ {
128
+ name: "IPv4 range start > end",
129
input: "192.0.2.10-192.0.2.0",
83
- wantErr: true,
130
+ wantErr: ErrInvalidRange,
131
},
85
- "v4 Range: invalid start": {
132
+ {
133
+ name: "IPv4 range invalid start",
134
input: "192.0.2.-192.0.2.10",
87
- wantErr: true,
135
+ wantErr: ErrInvalidSyntax,
136
},
89
- "v4 Range: invalid end": {
137
+ {
138
+ name: "IPv4 range invalid end",
139
input: "192.0.2.0-192.0.2.",
91
- wantErr: true,
140
+ wantErr: ErrInvalidSyntax,
141
},
93
- "v4 Range: v6 start": {
142
+ {
143
+ name: "IPv4 range with IPv6 start",
144
input: "2001:db8::0-192.0.2.10",
95
- wantErr: true,
145
+ wantErr: ErrMixedAddressFamilies,
146
},
97
- "v4 Range: v6 end": {
147
+ {
148
+ name: "IPv4 range with IPv6 end",
149
input: "192.0.2.0-2001:db8::0",
99
- wantErr: true,
150
+ wantErr: ErrMixedAddressFamilies,
151
+ },
152
+ {
153
+ name: "IPv4 range with spaces",
154
+ input: " 192.0.2.0 - 192.0.2.10 ",
155
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.10"),
156
},
101
- "v4 CIDR: /0": {
157
+
158
+ // IPv4 CIDR
159
+ {
160
+ name: "IPv4 CIDR /0",
161
input: "192.0.2.0/0",
103
- wantRange: prepareRange("0.0.0.1", "255.255.255.254"),
162
+ wantRange: mustParseRange(t, "0.0.0.1", "255.255.255.254"),
163
},
105
- "v4 CIDR: /24": {
164
+ {
165
+ name: "IPv4 CIDR /24",
166
input: "192.0.2.0/24",
107
- wantRange: prepareRange("192.0.2.1", "192.0.2.254"),
167
+ wantRange: mustParseRange(t, "192.0.2.1", "192.0.2.254"),
168
},
109
- "v4 CIDR: /30": {
169
+ {
170
+ name: "IPv4 CIDR /30",
171
input: "192.0.2.0/30",
111
- wantRange: prepareRange("192.0.2.1", "192.0.2.2"),
172
+ wantRange: mustParseRange(t, "192.0.2.1", "192.0.2.2"),
173
},
113
- "v4 CIDR: /31": {
174
+ {
175
+ name: "IPv4 CIDR /31",
176
input: "192.0.2.0/31",
115
- wantRange: prepareRange("192.0.2.0", "192.0.2.1"),
177
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.1"),
178
},
117
- "v4 CIDR: /32": {
179
+ {
180
+ name: "IPv4 CIDR /32",
181
input: "192.0.2.0/32",
119
- wantRange: prepareRange("192.0.2.0", "192.0.2.0"),
182
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.0"),
183
},
121
- "v4 CIDR: ip instead of host address": {
184
+ {
185
+ name: "IPv4 CIDR non-network address",
186
input: "192.0.2.10/24",
123
- wantRange: prepareRange("192.0.2.1", "192.0.2.254"),
187
+ wantRange: mustParseRange(t, "192.0.2.1", "192.0.2.254"),
188
},
125
- "v4 CIDR: missing prefix length": {
189
+ {
190
+ name: "IPv4 CIDR missing prefix",
191
input: "192.0.2.0/",
127
- wantErr: true,
192
+ wantErr: ErrInvalidSyntax,
193
},
129
- "v4 CIDR: invalid prefix length": {
194
+ {
195
+ name: "IPv4 CIDR invalid prefix",
196
input: "192.0.2.0/99",
131
- wantErr: true,
197
+ wantErr: ErrInvalidSyntax,
198
},
133
- "v4 Mask: /0": {
199
+
200
+ // IPv4 subnet mask
201
+ {
202
+ name: "IPv4 mask /0",
203
input: "192.0.2.0/0.0.0.0",
135
- wantRange: prepareRange("0.0.0.1", "255.255.255.254"),
204
+ wantRange: mustParseRange(t, "0.0.0.1", "255.255.255.254"),
205
},
137
- "v4 Mask: /24": {
206
+ {
207
+ name: "IPv4 mask /24",
208
input: "192.0.2.0/255.255.255.0",
139
- wantRange: prepareRange("192.0.2.1", "192.0.2.254"),
209
+ wantRange: mustParseRange(t, "192.0.2.1", "192.0.2.254"),
210
},
141
- "v4 Mask: /30": {
211
+ {
212
+ name: "IPv4 mask /30",
213
input: "192.0.2.0/255.255.255.252",
143
- wantRange: prepareRange("192.0.2.1", "192.0.2.2"),
214
+ wantRange: mustParseRange(t, "192.0.2.1", "192.0.2.2"),
215
},
145
- "v4 Mask: /31": {
216
+ {
217
+ name: "IPv4 mask /31",
218
input: "192.0.2.0/255.255.255.254",
147
- wantRange: prepareRange("192.0.2.0", "192.0.2.1"),
219
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.1"),
220
},
149
- "v4 Mask: /32": {
221
+ {
222
+ name: "IPv4 mask /32",
223
input: "192.0.2.0/255.255.255.255",
151
- wantRange: prepareRange("192.0.2.0", "192.0.2.0"),
152
- },
153
- "v4 Mask: missing prefix mask": {
154
- input: "192.0.2.0/",
155
- wantErr: true,
224
+ wantRange: mustParseRange(t, "192.0.2.0", "192.0.2.0"),
225
},
157
- "v4 Mask: invalid mask": {
226
+ {
227
+ name: "IPv4 mask invalid",
228
input: "192.0.2.0/mask",
159
- wantErr: true,
229
+ wantErr: ErrInvalidSyntax,
230
},
161
- "v4 Mask: not canonical form mask": {
231
+ {
232
+ name: "IPv4 mask non-contiguous",
233
input: "192.0.2.0/255.255.0.254",
163
- wantErr: true,
234
+ wantErr: ErrInvalidSyntax,
235
},
165
- "v4 Mask: v6 address": {
236
+ {
237
+ name: "IPv4 mask with IPv6 address",
238
input: "2001:db8::/255.255.255.0",
167
- wantErr: true,
239
+ wantErr: ErrInvalidSyntax,
240
},
241
170
- "v6 IP": {
242
+ // IPv6 single IP
243
+ {
244
+ name: "IPv6 single IP",
245
input: "2001:db8::0",
172
- wantRange: prepareRange("2001:db8::0", "2001:db8::0"),
246
+ wantRange: mustParseRange(t, "2001:db8::0", "2001:db8::0"),
247
},
174
- "v6 IP: invalid address": {
248
+ {
249
+ name: "IPv6 invalid address",
250
input: "2001:db8",
176
- wantErr: true,
251
+ wantErr: ErrInvalidSyntax,
252
},
178
- "v6 Range": {
253
+
254
+ // IPv6 ranges
255
+ {
256
+ name: "IPv6 range",
257
input: "2001:db8::-2001:db8::10",
180
- wantRange: prepareRange("2001:db8::", "2001:db8::10"),
258
+ wantRange: mustParseRange(t, "2001:db8::", "2001:db8::10"),
259
},
182
- "v6 Range: start == end": {
260
+ {
261
+ name: "IPv6 range start equals end",
262
input: "2001:db8::-2001:db8::",
184
- wantRange: prepareRange("2001:db8::", "2001:db8::"),
263
+ wantRange: mustParseRange(t, "2001:db8::", "2001:db8::"),
264
},
186
- "v6 Range: start > end": {
265
+ {
266
+ name: "IPv6 range start > end",
267
input: "2001:db8::10-2001:db8::",
188
- wantErr: true,
268
+ wantErr: ErrInvalidRange,
269
},
190
- "v6 Range: invalid start": {
270
+ {
271
+ name: "IPv6 range invalid start",
272
input: "2001:db8-2001:db8::10",
192
- wantErr: true,
273
+ wantErr: ErrInvalidSyntax,
274
},
194
- "v6 Range: invalid end": {
275
+ {
276
+ name: "IPv6 range invalid end",
277
input: "2001:db8::-2001:db8",
196
- wantErr: true,
278
+ wantErr: ErrInvalidSyntax,
279
},
198
- "v6 Range: v4 start": {
280
+ {
281
+ name: "IPv6 range with IPv4 start",
282
input: "192.0.2.0-2001:db8::10",
200
- wantErr: true,
283
+ wantErr: ErrMixedAddressFamilies,
284
},
202
- "v6 Range: v4 end": {
285
+ {
286
+ name: "IPv6 range with IPv4 end",
287
input: "2001:db8::-192.0.2.10",
204
- wantErr: true,
288
+ wantErr: ErrMixedAddressFamilies,
289
},
206
- "v6 CIDR: /0": {
290
+
291
+ // IPv6 CIDR
292
+ {
293
+ name: "IPv6 CIDR /0",
294
input: "2001:db8::/0",
208
- wantRange: prepareRange("::1", "ffff:ffff:ffff:ffff:ffff:ffff:ffff:fffe"),
295
+ wantRange: mustParseRange(t, "::1", "ffff:ffff:ffff:ffff:ffff:ffff:ffff:fffe"),
296
},
210
- "v6 CIDR: /64": {
297
+ {
298
+ name: "IPv6 CIDR /64",
299
input: "2001:db8::/64",
212
- wantRange: prepareRange("2001:db8::1", "2001:db8::ffff:ffff:ffff:fffe"),
300
+ wantRange: mustParseRange(t, "2001:db8::1", "2001:db8::ffff:ffff:ffff:fffe"),
301
},
214
- "v6 CIDR: /126": {
302
+ {
303
+ name: "IPv6 CIDR /126",
304
input: "2001:db8::/126",
216
- wantRange: prepareRange("2001:db8::1", "2001:db8::2"),
305
+ wantRange: mustParseRange(t, "2001:db8::1", "2001:db8::2"),
306
},
218
- "v6 CIDR: /127": {
307
+ {
308
+ name: "IPv6 CIDR /127",
309
input: "2001:db8::/127",
220
- wantRange: prepareRange("2001:db8::", "2001:db8::1"),
310
+ wantRange: mustParseRange(t, "2001:db8::", "2001:db8::1"),
311
},
222
- "v6 CIDR: /128": {
312
+ {
313
+ name: "IPv6 CIDR /128",
314
input: "2001:db8::/128",
224
- wantRange: prepareRange("2001:db8::", "2001:db8::"),
315
+ wantRange: mustParseRange(t, "2001:db8::", "2001:db8::"),
316
},
226
- "v6 CIDR: ip instead of host address": {
317
+ {
318
+ name: "IPv6 CIDR non-network address",
319
input: "2001:db8::10/64",
228
- wantRange: prepareRange("2001:db8::1", "2001:db8::ffff:ffff:ffff:fffe"),
320
+ wantRange: mustParseRange(t, "2001:db8::1", "2001:db8::ffff:ffff:ffff:fffe"),
321
},
230
- "v6 CIDR: missing prefix length": {
322
+ {
323
+ name: "IPv6 CIDR missing prefix",
324
input: "2001:db8::/",
232
- wantErr: true,
325
+ wantErr: ErrInvalidSyntax,
326
},
234
- "v6 CIDR: invalid prefix length": {
327
+ {
328
+ name: "IPv6 CIDR invalid prefix",
329
input: "2001:db8::/999",
236
- wantErr: true,
330
+ wantErr: ErrInvalidSyntax,
331
+ },
332
+
333
+ // Case sensitivity
334
+ {
335
+ name: "mixed case IPv6",
336
+ input: "2001:DB8::A-2001:DB8::F",
337
+ wantRange: mustParseRange(t, "2001:db8::a", "2001:db8::f"),
338
},
339
}
340
240
- for name, test := range tests {
241
- name = fmt.Sprintf("%s (%s)", name, test.input)
242
- t.Run(name, func(t *testing.T) {
243
- r, err := ParseRange(test.input)
341
+ for _, tt := range tests {
342
+ t.Run(tt.name, func(t *testing.T) {
343
+ t.Parallel()
344
245
- if test.wantErr {
345
+ r, err := ParseRange(tt.input)
346
+
347
+ if tt.wantErr != nil {
348
assert.Error(t, err)
247
- assert.Nilf(t, r, "want: nil, got: %s", r)
349
+ if !errors.Is(err, tt.wantErr) {
350
+ assert.ErrorContains(t, err, tt.wantErr.Error())
351
+ }
352
+ assert.Nil(t, r)
353
} else {
354
assert.NoError(t, err)
250
- assert.Equalf(t, test.wantRange, r, "want: %s, got: %s", test.wantRange, r)
355
+ assert.Equal(t, tt.wantRange, r)
356
}
357
})
358
}
359
}
360
256
-func prepareRange(start, end string) Range {
257
- return New(net.ParseIP(start), net.ParseIP(end))
361
+func TestMaskToPrefixLen(t *testing.T) {
362
+ t.Parallel()
363
+
364
+ tests := []struct {
365
+ name string
366
+ mask string
367
+ want int
368
+ wantOk bool
369
+ }{
370
+ {"all zeros", "0.0.0.0", 0, true},
371
+ {"all ones", "255.255.255.255", 32, true},
372
+ {"/8", "255.0.0.0", 8, true},
373
+ {"/16", "255.255.0.0", 16, true},
374
+ {"/24", "255.255.255.0", 24, true},
375
+ {"/25", "255.255.255.128", 25, true},
376
+ {"/30", "255.255.255.252", 30, true},
377
+ {"/31", "255.255.255.254", 31, true},
378
+ {"non-contiguous", "255.255.0.254", 0, false},
379
+ {"holes in mask", "255.0.255.0", 0, false},
380
+ }
381
+
382
+ for _, tt := range tests {
383
+ t.Run(tt.name, func(t *testing.T) {
384
+ t.Parallel()
385
+
386
+ mask := netip.MustParseAddr(tt.mask)
387
+ got, ok := maskToPrefixLen(mask)
388
+
389
+ assert.Equal(t, tt.wantOk, ok)
390
+ if ok {
391
+ assert.Equal(t, tt.want, got)
392
+ }
393
+ })
394
+ }
395
+}
396
+
397
+// Helper function to create a range for testing
398
+func mustParseRange(t *testing.T, start, end string) Range {
399
+ t.Helper()
400
+
401
+ startAddr, err := netip.ParseAddr(start)
402
+ require.NoError(t, err)
403
+
404
+ endAddr, err := netip.ParseAddr(end)
405
+ require.NoError(t, err)
406
+
407
+ r := New(startAddr, endAddr)
408
+ require.NotNil(t, r)
409
+
410
+ return r
411
+}
412
+
413
+// Benchmark tests
414
+func BenchmarkParseRange_IPv4(b *testing.B) {
415
+ inputs := []string{
416
+ "192.0.2.1",
417
+ "192.0.2.0-192.0.2.255",
418
+ "192.0.2.0/24",
419
+ "192.0.2.0/255.255.255.0",
420
+ }
421
+
422
+ b.ResetTimer()
423
+ for i := 0; i < b.N; i++ {
424
+ _, _ = ParseRange(inputs[i%len(inputs)])
425
+ }
426
+}
427
+
428
+func BenchmarkParseRange_IPv6(b *testing.B) {
429
+ inputs := []string{
430
+ "2001:db8::1",
431
+ "2001:db8::-2001:db8::ffff",
432
+ "2001:db8::/64",
433
+ }
434
+
435
+ b.ResetTimer()
436
+ for i := 0; i < b.N; i++ {
437
+ _, _ = ParseRange(inputs[i%len(inputs)])
438
+ }
439
}
src/go/plugin/go.d/pkg/iprange/pool.go
+239
-18
@@ -3,38 +3,259 @@
3
package iprange
4
5
import (
6
+ "iter"
7
"math/big"
7
- "net"
8
+ "net/netip"
9
+ "sort"
10
"strings"
11
)
12
11
-// Pool is a collection of IP Ranges.
12
-type Pool []Range
13
+// Pool represents a collection of IP ranges.
14
+// It provides methods to work with multiple ranges as a single unit.
15
+type Pool struct {
16
+ ranges []Range
17
+}
18
+
19
+// NewPool creates a new Pool from the given ranges.
20
+// It filters out nil ranges and makes a defensive copy of the slice.
21
+func NewPool(ranges ...Range) *Pool {
22
+ filtered := make([]Range, 0, len(ranges))
23
+ for _, r := range ranges {
24
+ if r != nil {
25
+ filtered = append(filtered, r)
26
+ }
27
+ }
28
+ return &Pool{ranges: filtered}
29
+}
30
+
31
+// ParsePool parses a string containing multiple IP ranges and returns a Pool.
32
+// The string should contain space-separated range specifications.
33
+func ParsePool(s string) (*Pool, error) {
34
+ ranges, err := ParseRanges(s)
35
+ if err != nil {
36
+ return nil, err
37
+ }
38
+ return NewPool(ranges...), nil
39
+}
40
+
41
+// Ranges returns a copy of the ranges in the pool.
42
+func (p *Pool) Ranges() []Range {
43
+ if p == nil || len(p.ranges) == 0 {
44
+ return nil
45
+ }
46
14
-// String returns the string form of the pool.
15
-func (p Pool) String() string {
16
- var b strings.Builder
17
- for _, r := range p {
18
- b.WriteString(r.String() + " ")
47
+ result := make([]Range, len(p.ranges))
48
+ copy(result, p.ranges)
49
+ return result
50
+}
51
+
52
+// Len returns the number of ranges in the pool.
53
+func (p *Pool) Len() int {
54
+ if p == nil {
55
+ return 0
56
}
20
- return strings.TrimSpace(b.String())
57
+ return len(p.ranges)
58
}
59
23
-// Size reports the number of IP addresses in the pool.
24
-func (p Pool) Size() *big.Int {
25
- size := big.NewInt(0)
26
- for _, r := range p {
27
- size.Add(size, r.Size())
60
+// IsEmpty reports whether the pool contains no ranges.
61
+func (p *Pool) IsEmpty() bool {
62
+ return p.Len() == 0
63
+}
64
+
65
+// String returns the string representation of the pool.
66
+// Ranges are separated by spaces.
67
+func (p *Pool) String() string {
68
+ if p == nil || len(p.ranges) == 0 {
69
+ return ""
70
+ }
71
+
72
+ parts := make([]string, len(p.ranges))
73
+ for i, r := range p.ranges {
74
+ parts[i] = r.String()
75
}
29
- return size
76
+ return strings.Join(parts, " ")
77
}
78
32
-// Contains reports whether the pool includes IP.
33
-func (p Pool) Contains(ip net.IP) bool {
34
- for _, r := range p {
79
+// Size returns the total number of IP addresses across all ranges in the pool.
80
+// Note: This does not account for overlapping ranges.
81
+func (p *Pool) Size() *big.Int {
82
+ if p == nil || len(p.ranges) == 0 {
83
+ return big.NewInt(0)
84
+ }
85
+
86
+ total := new(big.Int)
87
+ for _, r := range p.ranges {
88
+ total.Add(total, r.Size())
89
+ }
90
+ return total
91
+}
92
+
93
+// Contains reports whether any range in the pool includes the given IP address.
94
+func (p *Pool) Contains(ip netip.Addr) bool {
95
+ if p == nil || !ip.IsValid() {
96
+ return false
97
+ }
98
+
99
+ for _, r := range p.ranges {
100
if r.Contains(ip) {
101
return true
102
}
103
}
104
return false
105
}
106
+
107
+// ContainsRange reports whether the pool fully contains the given range.
108
+// This is true if every IP in the given range is contained in at least one range in the pool.
109
+// Note: For large ranges, this method may be expensive as it needs to verify coverage
110
+// of all addresses in the range.
111
+func (p *Pool) ContainsRange(r Range) bool {
112
+ if p == nil || r == nil {
113
+ return false
114
+ }
115
+
116
+ // Quick check: if start or end is not in pool, range can't be contained
117
+ if !p.Contains(r.Start()) || !p.Contains(r.End()) {
118
+ return false
119
+ }
120
+
121
+ // For efficiency, check if the range is fully contained within
122
+ // any single range in the pool
123
+ for _, poolRange := range p.ranges {
124
+ if rangeFullyContains(poolRange, r) {
125
+ return true
126
+ }
127
+ }
128
+
129
+ // If not in a single range, we need to check if the range is covered
130
+ // by multiple pool ranges without gaps.
131
+
132
+ // For small ranges (up to 256 IPs), check every address for completeness.
133
+ // This handles cases like a /24 network efficiently while being 100% accurate
134
+ size := r.Size()
135
+ if size.Cmp(big.NewInt(256)) <= 0 {
136
+ for addr := range r.Iterate() {
137
+ if !p.Contains(addr) {
138
+ return false
139
+ }
140
+ }
141
+ return true
142
+ }
143
+
144
+ // For larger ranges, we need a more efficient approach.
145
+ // Sort pool ranges and check for coverage without gaps.
146
+ return p.checkRangeCoverage(r)
147
+}
148
+
149
+// rangeFullyContains returns true if r1 fully contains r2
150
+func rangeFullyContains(r1, r2 Range) bool {
151
+ return r1.Start().Compare(r2.Start()) <= 0 && r1.End().Compare(r2.End()) >= 0
152
+}
153
+
154
+// checkRangeCoverage efficiently checks if a range is fully covered by the pool's ranges
155
+func (p *Pool) checkRangeCoverage(r Range) bool {
156
+ // Filter pool ranges that might overlap with r
157
+ var relevant []Range
158
+ for _, pr := range p.ranges {
159
+ // Check if pr overlaps with r
160
+ if pr.End().Compare(r.Start()) >= 0 && pr.Start().Compare(r.End()) <= 0 {
161
+ relevant = append(relevant, pr)
162
+ }
163
+ }
164
+
165
+ if len(relevant) == 0 {
166
+ return false
167
+ }
168
+
169
+ sort.Slice(relevant, func(i, j int) bool {
170
+ return relevant[i].Start().Compare(relevant[j].Start()) < 0
171
+ })
172
+
173
+ // Check if the sorted ranges cover r without gaps
174
+ currentCoverage := r.Start()
175
+
176
+ for _, pr := range relevant {
177
+ // If there's a gap between current coverage and this range
178
+ if pr.Start().Compare(currentCoverage) > 0 {
179
+ return false
180
+ }
181
+
182
+ // Extend coverage if this range extends it
183
+ if pr.End().Compare(currentCoverage) >= 0 {
184
+ currentCoverage = pr.End()
185
+
186
+ // If we've covered the entire range, we're done
187
+ if currentCoverage.Compare(r.End()) >= 0 {
188
+ return true
189
+ }
190
+
191
+ // Move to next address for continuous coverage check
192
+ next := currentCoverage.Next()
193
+ if next.IsValid() {
194
+ currentCoverage = next
195
+ }
196
+ }
197
+ }
198
+
199
+ // Check if we covered up to or past the end
200
+ return currentCoverage.Compare(r.End()) > 0
201
+}
202
+
203
+// Iterate returns an iterator over all IP addresses in all ranges.
204
+// Note: If ranges overlap, addresses in the overlap will be yielded multiple times.
205
+func (p *Pool) Iterate() iter.Seq[netip.Addr] {
206
+ return func(yield func(netip.Addr) bool) {
207
+ if p == nil {
208
+ return
209
+ }
210
+
211
+ for _, r := range p.ranges {
212
+ for addr := range r.Iterate() {
213
+ if !yield(addr) {
214
+ return
215
+ }
216
+ }
217
+ }
218
+ }
219
+}
220
+
221
+// IterateRanges returns an iterator over all ranges in the pool.
222
+func (p *Pool) IterateRanges() iter.Seq[Range] {
223
+ return func(yield func(Range) bool) {
224
+ if p == nil {
225
+ return
226
+ }
227
+
228
+ for _, r := range p.ranges {
229
+ if !yield(r) {
230
+ return
231
+ }
232
+ }
233
+ }
234
+}
235
+
236
+// Add appends one or more ranges to the pool.
237
+func (p *Pool) Add(ranges ...Range) {
238
+ for _, r := range ranges {
239
+ if r != nil {
240
+ p.ranges = append(p.ranges, r)
241
+ }
242
+ }
243
+}
244
+
245
+// AddString parses the string as IP ranges and adds them to the pool.
246
+func (p *Pool) AddString(s string) error {
247
+ ranges, err := ParseRanges(s)
248
+ if err != nil {
249
+ return err
250
+ }
251
+ p.Add(ranges...)
252
+ return nil
253
+}
254
+
255
+// Clone returns a deep copy of the pool.
256
+func (p *Pool) Clone() *Pool {
257
+ if p == nil {
258
+ return nil
259
+ }
260
+ return NewPool(p.ranges...)
261
+}
src/go/plugin/go.d/pkg/iprange/pool_test.go
+430
-46
@@ -5,100 +5,484 @@ package iprange
5
import (
6
"fmt"
7
"math/big"
8
- "net"
8
+ "net/netip"
9
"testing"
10
11
"github.com/stretchr/testify/assert"
12
"github.com/stretchr/testify/require"
13
)
14
15
+func TestNewPool(t *testing.T) {
16
+ t.Parallel()
17
+
18
+ r1 := mustParseRange(t, "192.0.2.0", "192.0.2.10")
19
+ r2 := mustParseRange(t, "192.0.2.20", "192.0.2.30")
20
+
21
+ tests := []struct {
22
+ name string
23
+ ranges []Range
24
+ wantCount int
25
+ }{
26
+ {
27
+ name: "empty pool",
28
+ ranges: nil,
29
+ wantCount: 0,
30
+ },
31
+ {
32
+ name: "single range",
33
+ ranges: []Range{r1},
34
+ wantCount: 1,
35
+ },
36
+ {
37
+ name: "multiple ranges",
38
+ ranges: []Range{r1, r2},
39
+ wantCount: 2,
40
+ },
41
+ {
42
+ name: "with nil ranges",
43
+ ranges: []Range{r1, nil, r2, nil},
44
+ wantCount: 2,
45
+ },
46
+ }
47
+
48
+ for _, tt := range tests {
49
+ t.Run(tt.name, func(t *testing.T) {
50
+ t.Parallel()
51
+
52
+ pool := NewPool(tt.ranges...)
53
+ assert.Equal(t, tt.wantCount, pool.Len())
54
+ })
55
+ }
56
+}
57
+
58
+func TestParsePool(t *testing.T) {
59
+ t.Parallel()
60
+
61
+ tests := []struct {
62
+ name string
63
+ input string
64
+ wantLen int
65
+ wantErr bool
66
+ }{
67
+ {
68
+ name: "single range",
69
+ input: "192.0.2.0-192.0.2.10",
70
+ wantLen: 1,
71
+ },
72
+ {
73
+ name: "multiple ranges",
74
+ input: "192.0.2.0-192.0.2.10 2001:db8::-2001:db8::10",
75
+ wantLen: 2,
76
+ },
77
+ {
78
+ name: "invalid range",
79
+ input: "192.0.2.0-192.0.2.10 invalid",
80
+ wantErr: true,
81
+ },
82
+ {
83
+ name: "empty string",
84
+ input: "",
85
+ wantLen: 0,
86
+ },
87
+ }
88
+
89
+ for _, tt := range tests {
90
+ t.Run(tt.name, func(t *testing.T) {
91
+ t.Parallel()
92
+
93
+ pool, err := ParsePool(tt.input)
94
+
95
+ if tt.wantErr {
96
+ assert.Error(t, err)
97
+ assert.Nil(t, pool)
98
+ } else {
99
+ assert.NoError(t, err)
100
+ require.NotNil(t, pool)
101
+ assert.Equal(t, tt.wantLen, pool.Len())
102
+ }
103
+ })
104
+ }
105
+}
106
+
107
func TestPool_String(t *testing.T) {
16
- tests := map[string]struct {
108
+ t.Parallel()
109
+
110
+ tests := []struct {
111
+ name string
112
input string
113
wantString string
114
}{
20
- "singe": {
115
+ {
116
+ name: "single range",
117
input: "192.0.2.0-192.0.2.10",
118
wantString: "192.0.2.0-192.0.2.10",
119
},
24
- "multiple": {
120
+ {
121
+ name: "multiple ranges",
122
input: "192.0.2.0-192.0.2.10 2001:db8::-2001:db8::10",
123
wantString: "192.0.2.0-192.0.2.10 2001:db8::-2001:db8::10",
124
},
125
+ {
126
+ name: "empty pool",
127
+ input: "",
128
+ wantString: "",
129
+ },
130
}
131
30
- for name, test := range tests {
31
- t.Run(name, func(t *testing.T) {
32
- rs, err := ParseRanges(test.input)
132
+ for _, tt := range tests {
133
+ t.Run(tt.name, func(t *testing.T) {
134
+ t.Parallel()
135
+
136
+ pool, err := ParsePool(tt.input)
137
require.NoError(t, err)
34
- p := Pool(rs)
138
36
- assert.Equal(t, test.wantString, p.String())
139
+ assert.Equal(t, tt.wantString, pool.String())
140
})
141
}
142
}
143
144
+func TestPool_NilSafety(t *testing.T) {
145
+ t.Parallel()
146
+
147
+ var pool *Pool
148
+
149
+ // All methods should handle nil pool gracefully
150
+ assert.Equal(t, 0, pool.Len())
151
+ assert.True(t, pool.IsEmpty())
152
+ assert.Equal(t, "", pool.String())
153
+ assert.Equal(t, big.NewInt(0), pool.Size())
154
+ assert.False(t, pool.Contains(netip.MustParseAddr("192.0.2.1")))
155
+ assert.Nil(t, pool.Ranges())
156
+ assert.Nil(t, pool.Clone())
157
+
158
+ // Iterators should not panic
159
+ for _ = range pool.Iterate() {
160
+ t.Fatal("nil pool should not yield any addresses")
161
+ }
162
+ for _ = range pool.IterateRanges() {
163
+ t.Fatal("nil pool should not yield any ranges")
164
+ }
165
+}
166
+
167
func TestPool_Size(t *testing.T) {
42
- tests := map[string]struct {
168
+ t.Parallel()
169
+
170
+ tests := []struct {
171
+ name string
172
input string
44
- wantSize *big.Int
173
+ wantSize int64
174
}{
46
- "singe": {
175
+ {
176
+ name: "single range",
177
input: "192.0.2.0-192.0.2.10",
48
- wantSize: big.NewInt(11),
178
+ wantSize: 11,
179
},
50
- "multiple": {
180
+ {
181
+ name: "multiple ranges",
182
input: "192.0.2.0-192.0.2.10 2001:db8::-2001:db8::10",
52
- wantSize: big.NewInt(11 + 17),
183
+ wantSize: 11 + 17,
184
+ },
185
+ {
186
+ name: "empty pool",
187
+ input: "",
188
+ wantSize: 0,
189
+ },
190
+ {
191
+ name: "overlapping ranges (counted separately)",
192
+ input: "192.0.2.0-192.0.2.10 192.0.2.5-192.0.2.15",
193
+ wantSize: 11 + 11, // Overlaps are counted twice
194
},
195
}
196
56
- for name, test := range tests {
57
- t.Run(name, func(t *testing.T) {
58
- rs, err := ParseRanges(test.input)
197
+ for _, tt := range tests {
198
+ t.Run(tt.name, func(t *testing.T) {
199
+ t.Parallel()
200
+
201
+ pool, err := ParsePool(tt.input)
202
require.NoError(t, err)
60
- p := Pool(rs)
203
62
- assert.Equal(t, test.wantSize, p.Size())
204
+ assert.Equal(t, big.NewInt(tt.wantSize), pool.Size())
205
})
206
}
207
}
208
209
func TestPool_Contains(t *testing.T) {
68
- tests := map[string]struct {
69
- input string
70
- ip string
71
- wantFail bool
210
+ t.Parallel()
211
+
212
+ tests := []struct {
213
+ name string
214
+ poolStr string
215
+ ip string
216
+ wantFound bool
217
}{
73
- "inside first": {
74
- input: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30 2001:db8::-2001:db8::10",
75
- ip: "192.0.2.5",
218
+ {
219
+ name: "IP in first range",
220
+ poolStr: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30 2001:db8::-2001:db8::10",
221
+ ip: "192.0.2.5",
222
+ wantFound: true,
223
+ },
224
+ {
225
+ name: "IP in last range",
226
+ poolStr: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30 2001:db8::-2001:db8::10",
227
+ ip: "2001:db8::5",
228
+ wantFound: true,
229
},
77
- "inside last": {
78
- input: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30 2001:db8::-2001:db8::10",
79
- ip: "2001:db8::5",
230
+ {
231
+ name: "IP not in any range",
232
+ poolStr: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30 2001:db8::-2001:db8::10",
233
+ ip: "192.0.2.100",
234
+ wantFound: false,
235
},
81
- "outside": {
82
- input: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30 2001:db8::-2001:db8::10",
83
- ip: "192.0.2.100",
84
- wantFail: true,
236
+ {
237
+ name: "IP in gap between ranges",
238
+ poolStr: "192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30",
239
+ ip: "192.0.2.15",
240
+ wantFound: false,
241
+ },
242
+ {
243
+ name: "empty pool",
244
+ poolStr: "",
245
+ ip: "192.0.2.1",
246
+ wantFound: false,
247
},
248
}
249
88
- for name, test := range tests {
89
- name = fmt.Sprintf("%s (range: %s, ip: %s)", name, test.input, test.ip)
90
- t.Run(name, func(t *testing.T) {
91
- rs, err := ParseRanges(test.input)
250
+ for _, tt := range tests {
251
+ t.Run(tt.name, func(t *testing.T) {
252
+ t.Parallel()
253
+
254
+ pool, err := ParsePool(tt.poolStr)
255
require.NoError(t, err)
93
- ip := net.ParseIP(test.ip)
94
- require.NotNil(t, ip)
95
- p := Pool(rs)
256
97
- if test.wantFail {
98
- assert.False(t, p.Contains(ip))
99
- } else {
100
- assert.True(t, p.Contains(ip))
101
- }
257
+ ip, err := netip.ParseAddr(tt.ip)
258
+ require.NoError(t, err)
259
+
260
+ assert.Equal(t, tt.wantFound, pool.Contains(ip))
261
+ })
262
+ }
263
+}
264
+
265
+func TestPool_ContainsRange(t *testing.T) {
266
+ t.Parallel()
267
+
268
+ tests := []struct {
269
+ name string
270
+ poolStr string
271
+ rangeStr string
272
+ wantContains bool
273
+ }{
274
+ {
275
+ name: "range fully within single pool range",
276
+ poolStr: "192.0.2.0-192.0.2.100",
277
+ rangeStr: "192.0.2.10-192.0.2.20",
278
+ wantContains: true,
279
+ },
280
+ {
281
+ name: "range equals pool range",
282
+ poolStr: "192.0.2.0-192.0.2.100",
283
+ rangeStr: "192.0.2.0-192.0.2.100",
284
+ wantContains: true,
285
+ },
286
+ {
287
+ name: "range extends beyond pool",
288
+ poolStr: "192.0.2.0-192.0.2.100",
289
+ rangeStr: "192.0.2.50-192.0.2.150",
290
+ wantContains: false,
291
+ },
292
+ {
293
+ name: "range not in pool",
294
+ poolStr: "192.0.2.0-192.0.2.100",
295
+ rangeStr: "192.0.3.0-192.0.3.100",
296
+ wantContains: false,
297
+ },
298
+ {
299
+ name: "empty pool",
300
+ poolStr: "",
301
+ rangeStr: "192.0.2.0-192.0.2.10",
302
+ wantContains: false,
303
+ },
304
+ {
305
+ name: "large range with gap in pool",
306
+ poolStr: "10.0.0.0-10.0.0.255 10.0.2.0-10.0.3.255", // Gap: 10.0.1.0-10.0.1.255
307
+ rangeStr: "10.0.0.128-10.0.2.128", // 513 addresses, spans the gap
308
+ wantContains: false,
309
+ },
310
+ {
311
+ name: "large range fully covered by multiple pool ranges",
312
+ poolStr: "10.0.0.0-10.0.1.255 10.0.2.0-10.0.3.255", // Combined: 1024 addresses
313
+ rangeStr: "10.0.0.0-10.0.1.255", // 512 addresses
314
+ wantContains: true,
315
+ },
316
+ {
317
+ name: "large range with adjacent pool ranges",
318
+ poolStr: "172.16.0.0-172.16.0.255 172.16.1.0-172.16.1.255 172.16.2.0-172.16.2.255",
319
+ rangeStr: "172.16.0.100-172.16.2.100", // >500 addresses, continuous coverage
320
+ wantContains: true,
321
+ },
322
+ {
323
+ name: "large IPv6 range with gap",
324
+ poolStr: "2001:db8::-2001:db8::fff 2001:db8::2000-2001:db8::2fff", // Gap from ::1000 to ::1fff
325
+ rangeStr: "2001:db8::500-2001:db8::2500", // Spans the gap
326
+ wantContains: false,
327
+ },
328
+ {
329
+ name: "large range at pool boundaries",
330
+ poolStr: "192.168.0.0-192.168.1.255 192.168.2.0-192.168.3.255",
331
+ rangeStr: "192.168.1.0-192.168.2.255", // 512 addresses, needs both ranges
332
+ wantContains: true,
333
+ },
334
+ {
335
+ name: "very large range /16 with small gap",
336
+ poolStr: "10.0.0.0/17 10.0.128.1-10.0.255.255", // Missing exactly 10.0.128.0
337
+ rangeStr: "10.0.0.0/16", // 65536 addresses
338
+ wantContains: false,
339
+ },
340
+ }
341
+
342
+ for _, tt := range tests {
343
+ t.Run(tt.name, func(t *testing.T) {
344
+ t.Parallel()
345
+
346
+ pool, err := ParsePool(tt.poolStr)
347
+ require.NoError(t, err)
348
+
349
+ r, err := ParseRange(tt.rangeStr)
350
+ require.NoError(t, err)
351
+
352
+ assert.Equal(t, tt.wantContains, pool.ContainsRange(r))
353
})
354
}
355
}
356
+
357
+func TestPool_Iterate(t *testing.T) {
358
+ t.Parallel()
359
+
360
+ pool, err := ParsePool("192.0.2.0-192.0.2.2 192.0.2.10-192.0.2.12")
361
+ require.NoError(t, err)
362
+
363
+ var addresses []string
364
+ for addr := range pool.Iterate() {
365
+ addresses = append(addresses, addr.String())
366
+ }
367
+
368
+ expected := []string{
369
+ "192.0.2.0", "192.0.2.1", "192.0.2.2",
370
+ "192.0.2.10", "192.0.2.11", "192.0.2.12",
371
+ }
372
+ assert.Equal(t, expected, addresses)
373
+}
374
+
375
+func TestPool_IterateRanges(t *testing.T) {
376
+ t.Parallel()
377
+
378
+ pool, err := ParsePool("192.0.2.0-192.0.2.10 2001:db8::-2001:db8::10")
379
+ require.NoError(t, err)
380
+
381
+ var ranges []string
382
+ for r := range pool.IterateRanges() {
383
+ ranges = append(ranges, r.String())
384
+ }
385
+
386
+ expected := []string{
387
+ "192.0.2.0-192.0.2.10",
388
+ "2001:db8::-2001:db8::10",
389
+ }
390
+ assert.Equal(t, expected, ranges)
391
+}
392
+
393
+func TestPool_Add(t *testing.T) {
394
+ t.Parallel()
395
+
396
+ pool := NewPool()
397
+ assert.Equal(t, 0, pool.Len())
398
+
399
+ r1 := mustParseRange(t, "192.0.2.0", "192.0.2.10")
400
+ pool.Add(r1)
401
+ assert.Equal(t, 1, pool.Len())
402
+
403
+ r2 := mustParseRange(t, "192.0.2.20", "192.0.2.30")
404
+ r3 := mustParseRange(t, "192.0.2.40", "192.0.2.50")
405
+ pool.Add(r2, nil, r3) // nil should be ignored
406
+ assert.Equal(t, 3, pool.Len())
407
+}
408
+
409
+func TestPool_AddString(t *testing.T) {
410
+ t.Parallel()
411
+
412
+ pool := NewPool()
413
+
414
+ err := pool.AddString("192.0.2.0-192.0.2.10 192.0.2.20/24")
415
+ assert.NoError(t, err)
416
+ assert.Equal(t, 2, pool.Len())
417
+
418
+ err = pool.AddString("invalid")
419
+ assert.Error(t, err)
420
+ assert.Equal(t, 2, pool.Len()) // Should not change
421
+}
422
+
423
+func TestPool_Clone(t *testing.T) {
424
+ t.Parallel()
425
+
426
+ original, err := ParsePool("192.0.2.0-192.0.2.10 192.0.2.20-192.0.2.30")
427
+ require.NoError(t, err)
428
+
429
+ clone := original.Clone()
430
+
431
+ // Should be equal
432
+ assert.Equal(t, original.String(), clone.String())
433
+ assert.Equal(t, original.Len(), clone.Len())
434
+
435
+ // But independent
436
+ r := mustParseRange(t, "192.0.2.40", "192.0.2.50")
437
+ clone.Add(r)
438
+
439
+ assert.NotEqual(t, original.Len(), clone.Len())
440
+}
441
+
442
+func TestPool_Ranges(t *testing.T) {
443
+ t.Parallel()
444
+
445
+ r1 := mustParseRange(t, "192.0.2.0", "192.0.2.10")
446
+ r2 := mustParseRange(t, "192.0.2.20", "192.0.2.30")
447
+
448
+ pool := NewPool(r1, r2)
449
+ ranges := pool.Ranges()
450
+
451
+ assert.Len(t, ranges, 2)
452
+ assert.Equal(t, []Range{r1, r2}, ranges)
453
+
454
+ // Modifying returned slice should not affect pool
455
+ ranges[0] = nil
456
+ assert.Equal(t, 2, pool.Len())
457
+}
458
+
459
+// Benchmark tests
460
+func BenchmarkPool_Contains(b *testing.B) {
461
+ // Create a pool with multiple ranges
462
+ pool := NewPool()
463
+ for i := 0; i < 10; i++ {
464
+ start := fmt.Sprintf("192.0.%d.0", i)
465
+ end := fmt.Sprintf("192.0.%d.255", i)
466
+ r, _ := ParseRange(fmt.Sprintf("%s-%s", start, end))
467
+ pool.Add(r)
468
+ }
469
+
470
+ ip := netip.MustParseAddr("192.0.5.100")
471
+
472
+ b.ResetTimer()
473
+ for i := 0; i < b.N; i++ {
474
+ _ = pool.Contains(ip)
475
+ }
476
+}
477
+
478
+func BenchmarkPool_Iterate(b *testing.B) {
479
+ pool, _ := ParsePool("192.0.2.0/24 192.0.3.0/24")
480
+
481
+ b.ResetTimer()
482
+ for i := 0; i < b.N; i++ {
483
+ count := 0
484
+ for _ = range pool.Iterate() {
485
+ count++
486
+ }
487
+ }
488
+}
src/go/plugin/go.d/pkg/iprange/range.go
+105
-53
@@ -3,11 +3,10 @@
3
package iprange
4
5
import (
6
- "bytes"
6
"fmt"
7
"iter"
8
"math/big"
10
- "net"
9
+ "net/netip"
10
)
11
12
// Family represents IP Range address-family.
@@ -20,100 +19,153 @@ const (
19
V6Family
20
)
21
23
-// Range represents an IP range.
22
+// String returns the string representation of the address family.
23
+func (f Family) String() string {
24
+ switch f {
25
+ case V4Family:
26
+ return "IPv4"
27
+ case V6Family:
28
+ return "IPv6"
29
+ default:
30
+ return fmt.Sprintf("Unknown(%d)", f)
31
+ }
32
+}
33
+
34
+// Range represents an IP address range.
35
+// It provides methods to check containment, iterate over addresses,
36
+// and calculate the size of the range.
37
type Range interface {
38
+ // Family returns the address family (IPv4 or IPv6) of the range.
39
Family() Family
26
- Contains(ip net.IP) bool
40
+
41
+ // Contains reports whether the range includes the given IP address.
42
+ Contains(ip netip.Addr) bool
43
+
44
+ // Size returns the number of IP addresses in the range.
45
Size() *big.Int
28
- Iterate() iter.Seq[net.IP]
46
+
47
+ // Iterate returns an iterator over all IP addresses in the range.
48
+ Iterate() iter.Seq[netip.Addr]
49
+
50
+ // String returns the string representation of the range in "start-end" format.
51
fmt.Stringer
52
31
- getStart() net.IP
32
- getEnd() net.IP
53
+ // Start returns the first IP address in the range.
54
+ Start() netip.Addr
55
+
56
+ // End returns the last IP address in the range.
57
+ End() netip.Addr
58
}
59
35
-// New returns new IP Range.
36
-// If it is not a valid range (start and end IPs have different address-families, or start > end),
37
-// New returns nil.
38
-func New(start, end net.IP) Range {
39
- if isV4RangeValid(start, end) {
40
- return v4Range{start: start, end: end}
60
+// New creates a new IP Range from start and end addresses.
61
+// It returns nil if the range is invalid (start and end have different
62
+// address families, or start > end).
63
+func New(start, end netip.Addr) Range {
64
+ if !start.IsValid() || !end.IsValid() {
65
+ return nil
66
}
42
- if isV6RangeValid(start, end) {
43
- return v6Range{start: start, end: end}
67
+
68
+ switch {
69
+ case start.Is4() && end.Is4():
70
+ if start.Compare(end) > 0 {
71
+ return nil
72
+ }
73
+ return &v4Range{start: start, end: end}
74
+ case start.Is6() && end.Is6() && !start.Is4In6() && !end.Is4In6():
75
+ if start.Compare(end) > 0 {
76
+ return nil
77
+ }
78
+ return &v6Range{start: start, end: end}
79
+ default:
80
+ return nil
81
}
45
- return nil
82
}
83
84
+// v4Range implements Range for IPv4 addresses.
85
type v4Range struct {
49
- start net.IP
50
- end net.IP
86
+ start netip.Addr
87
+ end netip.Addr
88
}
89
53
-func (r v4Range) getStart() net.IP { return r.start }
54
-func (r v4Range) getEnd() net.IP { return r.end }
90
+// compile-time check that v4Range implements Range
91
+var _ Range = (*v4Range)(nil)
92
+
93
+func (r *v4Range) Start() netip.Addr { return r.start }
94
+func (r *v4Range) End() netip.Addr { return r.end }
95
56
-func (r v4Range) Iterate() iter.Seq[net.IP] {
96
+func (r *v4Range) Iterate() iter.Seq[netip.Addr] {
97
return iterate(r)
98
}
99
60
-// String returns the string form of the range.
61
-func (r v4Range) String() string {
100
+func (r *v4Range) String() string {
101
return fmt.Sprintf("%s-%s", r.start, r.end)
102
}
103
65
-// Family returns the range address family.
66
-func (r v4Range) Family() Family {
104
+func (r *v4Range) Family() Family {
105
return V4Family
106
}
107
70
-// Contains reports whether the range includes IP.
71
-func (r v4Range) Contains(ip net.IP) bool {
72
- return bytes.Compare(ip, r.start) >= 0 && bytes.Compare(ip, r.end) <= 0
108
+func (r *v4Range) Contains(ip netip.Addr) bool {
109
+ if !ip.Is4() {
110
+ return false
111
+ }
112
+ return ip.Compare(r.start) >= 0 && ip.Compare(r.end) <= 0
113
}
114
75
-// Size reports the number of IP addresses in the range.
76
-func (r v4Range) Size() *big.Int {
77
- return big.NewInt(v4ToInt(r.end) - v4ToInt(r.start) + 1)
115
+func (r *v4Range) Size() *big.Int {
116
+ // For IPv4, we can safely use uint32 arithmetic
117
+ start := r.start.As4()
118
+ end := r.end.As4()
119
+
120
+ startUint := uint32(start[0])<<24 | uint32(start[1])<<16 | uint32(start[2])<<8 | uint32(start[3])
121
+ endUint := uint32(end[0])<<24 | uint32(end[1])<<16 | uint32(end[2])<<8 | uint32(end[3])
122
+
123
+ // Add 1 because the range is inclusive
124
+ return big.NewInt(int64(endUint - startUint + 1))
125
}
126
127
+// v6Range implements Range for IPv6 addresses.
128
type v6Range struct {
81
- start net.IP
82
- end net.IP
129
+ start netip.Addr
130
+ end netip.Addr
131
}
132
85
-func (r v6Range) getStart() net.IP { return r.start }
86
-func (r v6Range) getEnd() net.IP { return r.end }
133
+// compile-time check that v6Range implements Range
134
+var _ Range = (*v6Range)(nil)
135
88
-func (r v6Range) Iterate() iter.Seq[net.IP] {
136
+func (r *v6Range) Start() netip.Addr { return r.start }
137
+func (r *v6Range) End() netip.Addr { return r.end }
138
+
139
+func (r *v6Range) Iterate() iter.Seq[netip.Addr] {
140
return iterate(r)
141
}
142
92
-// String returns the string form of the range.
93
-func (r v6Range) String() string {
143
+func (r *v6Range) String() string {
144
return fmt.Sprintf("%s-%s", r.start, r.end)
145
}
146
97
-// Family returns the range address family.
98
-func (r v6Range) Family() Family {
147
+func (r *v6Range) Family() Family {
148
return V6Family
149
}
150
102
-// Contains reports whether the range includes IP.
103
-func (r v6Range) Contains(ip net.IP) bool {
104
- return bytes.Compare(ip, r.start) >= 0 && bytes.Compare(ip, r.end) <= 0
151
+func (r *v6Range) Contains(ip netip.Addr) bool {
152
+ if !ip.Is6() || ip.Is4In6() {
153
+ return false
154
+ }
155
+ return ip.Compare(r.start) >= 0 && ip.Compare(r.end) <= 0
156
}
157
107
-// Size reports the number of IP addresses in the range.
108
-func (r v6Range) Size() *big.Int {
109
- size := big.NewInt(0)
110
- size.Add(size, big.NewInt(0).SetBytes(r.end))
111
- size.Sub(size, big.NewInt(0).SetBytes(r.start))
158
+func (r *v6Range) Size() *big.Int {
159
+ // For IPv6, we must use big.Int to handle 128-bit arithmetic
160
+ startBytes := r.start.As16()
161
+ endBytes := r.end.As16()
162
+
163
+ startBig := new(big.Int).SetBytes(startBytes[:])
164
+ endBig := new(big.Int).SetBytes(endBytes[:])
165
+
166
+ // Calculate end - start + 1
167
+ size := new(big.Int).Sub(endBig, startBig)
168
size.Add(size, big.NewInt(1))
113
- return size
114
-}
169
116
-func v4ToInt(ip net.IP) int64 {
117
- ip = ip.To4()
118
- return int64(ip[0])<<24 | int64(ip[1])<<16 | int64(ip[2])<<8 | int64(ip[3])
170
+ return size
171
}
src/go/plugin/go.d/pkg/iprange/range_test.go
+382
-124
@@ -3,9 +3,8 @@
3
package iprange
4
5
import (
6
- "fmt"
6
"math/big"
8
- "net"
7
+ "net/netip"
8
"testing"
9
10
"github.com/stretchr/testify/assert"
@@ -13,40 +12,68 @@ import (
12
)
13
14
func TestV4Range_String(t *testing.T) {
16
- tests := map[string]struct {
15
+ t.Parallel()
16
+
17
+ tests := []struct {
18
+ name string
19
input string
20
wantString string
21
}{
20
- "IP": {input: "192.0.2.0", wantString: "192.0.2.0-192.0.2.0"},
21
- "Range": {input: "192.0.2.0-192.0.2.10", wantString: "192.0.2.0-192.0.2.10"},
22
- "CIDR": {input: "192.0.2.0/24", wantString: "192.0.2.1-192.0.2.254"},
23
- "Mask": {input: "192.0.2.0/255.255.255.0", wantString: "192.0.2.1-192.0.2.254"},
22
+ {
23
+ name: "single IP",
24
+ input: "192.0.2.0",
25
+ wantString: "192.0.2.0-192.0.2.0",
26
+ },
27
+ {
28
+ name: "IP range",
29
+ input: "192.0.2.0-192.0.2.10",
30
+ wantString: "192.0.2.0-192.0.2.10",
31
+ },
32
+ {
33
+ name: "CIDR /24",
34
+ input: "192.0.2.0/24",
35
+ wantString: "192.0.2.1-192.0.2.254",
36
+ },
37
+ {
38
+ name: "subnet mask",
39
+ input: "192.0.2.0/255.255.255.0",
40
+ wantString: "192.0.2.1-192.0.2.254",
41
+ },
42
}
43
26
- for name, test := range tests {
27
- t.Run(name, func(t *testing.T) {
28
- r, err := ParseRange(test.input)
44
+ for _, tt := range tests {
45
+ t.Run(tt.name, func(t *testing.T) {
46
+ t.Parallel()
47
+
48
+ r, err := ParseRange(tt.input)
49
require.NoError(t, err)
50
+ require.NotNil(t, r)
51
31
- assert.Equal(t, test.wantString, r.String())
52
+ assert.Equal(t, tt.wantString, r.String())
53
})
54
}
55
}
56
57
func TestV4Range_Family(t *testing.T) {
37
- tests := map[string]struct {
58
+ t.Parallel()
59
+
60
+ tests := []struct {
61
+ name string
62
input string
63
}{
40
- "IP": {input: "192.0.2.0"},
41
- "Range": {input: "192.0.2.0-192.0.2.10"},
42
- "CIDR": {input: "192.0.2.0/24"},
43
- "Mask": {input: "192.0.2.0/255.255.255.0"},
64
+ {"single IP", "192.0.2.0"},
65
+ {"IP range", "192.0.2.0-192.0.2.10"},
66
+ {"CIDR", "192.0.2.0/24"},
67
+ {"subnet mask", "192.0.2.0/255.255.255.0"},
68
}
69
46
- for name, test := range tests {
47
- t.Run(name, func(t *testing.T) {
48
- r, err := ParseRange(test.input)
70
+ for _, tt := range tests {
71
+ t.Run(tt.name, func(t *testing.T) {
72
+ t.Parallel()
73
+
74
+ r, err := ParseRange(tt.input)
75
require.NoError(t, err)
76
+ require.NotNil(t, r)
77
78
assert.Equal(t, V4Family, r.Family())
79
})
@@ -54,116 +81,184 @@ func TestV4Range_Family(t *testing.T) {
81
}
82
83
func TestV4Range_Size(t *testing.T) {
57
- tests := map[string]struct {
84
+ t.Parallel()
85
+
86
+ tests := []struct {
87
+ name string
88
input string
59
- wantSize *big.Int
89
+ wantSize int64
90
}{
61
- "IP": {input: "192.0.2.0", wantSize: big.NewInt(1)},
62
- "Range": {input: "192.0.2.0-192.0.2.10", wantSize: big.NewInt(11)},
63
- "CIDR": {input: "192.0.2.0/24", wantSize: big.NewInt(254)},
64
- "CIDR 31": {input: "192.0.2.0/31", wantSize: big.NewInt(2)},
65
- "CIDR 32": {input: "192.0.2.0/32", wantSize: big.NewInt(1)},
66
- "Mask": {input: "192.0.2.0/255.255.255.0", wantSize: big.NewInt(254)},
67
- "Mask 31": {input: "192.0.2.0/255.255.255.254", wantSize: big.NewInt(2)},
68
- "Mask 32": {input: "192.0.2.0/255.255.255.255", wantSize: big.NewInt(1)},
69
- }
70
-
71
- for name, test := range tests {
72
- t.Run(name, func(t *testing.T) {
73
- r, err := ParseRange(test.input)
91
+ {"single IP", "192.0.2.0", 1},
92
+ {"IP range", "192.0.2.0-192.0.2.10", 11},
93
+ {"CIDR /24", "192.0.2.0/24", 254},
94
+ {"CIDR /31", "192.0.2.0/31", 2},
95
+ {"CIDR /32", "192.0.2.0/32", 1},
96
+ {"subnet mask /24", "192.0.2.0/255.255.255.0", 254},
97
+ {"subnet mask /31", "192.0.2.0/255.255.255.254", 2},
98
+ {"subnet mask /32", "192.0.2.0/255.255.255.255", 1},
99
+ }
100
+
101
+ for _, tt := range tests {
102
+ t.Run(tt.name, func(t *testing.T) {
103
+ t.Parallel()
104
+
105
+ r, err := ParseRange(tt.input)
106
require.NoError(t, err)
107
+ require.NotNil(t, r)
108
76
- assert.Equal(t, test.wantSize, r.Size())
109
+ assert.Equal(t, big.NewInt(tt.wantSize), r.Size())
110
})
111
}
112
}
113
114
func TestV4Range_Contains(t *testing.T) {
82
- tests := map[string]struct {
83
- input string
84
- ip string
85
- wantFail bool
115
+ t.Parallel()
116
+
117
+ tests := []struct {
118
+ name string
119
+ rangeStr string
120
+ ip string
121
+ wantFound bool
122
}{
87
- "inside": {input: "192.0.2.0-192.0.2.10", ip: "192.0.2.5"},
88
- "outside": {input: "192.0.2.0-192.0.2.10", ip: "192.0.2.55", wantFail: true},
89
- "eq start": {input: "192.0.2.0-192.0.2.10", ip: "192.0.2.0"},
90
- "eq end": {input: "192.0.2.0-192.0.2.10", ip: "192.0.2.10"},
91
- "v6": {input: "192.0.2.0-192.0.2.10", ip: "2001:db8::", wantFail: true},
123
+ {
124
+ name: "IP inside range",
125
+ rangeStr: "192.0.2.0-192.0.2.10",
126
+ ip: "192.0.2.5",
127
+ wantFound: true,
128
+ },
129
+ {
130
+ name: "IP outside range",
131
+ rangeStr: "192.0.2.0-192.0.2.10",
132
+ ip: "192.0.2.55",
133
+ wantFound: false,
134
+ },
135
+ {
136
+ name: "IP equals start",
137
+ rangeStr: "192.0.2.0-192.0.2.10",
138
+ ip: "192.0.2.0",
139
+ wantFound: true,
140
+ },
141
+ {
142
+ name: "IP equals end",
143
+ rangeStr: "192.0.2.0-192.0.2.10",
144
+ ip: "192.0.2.10",
145
+ wantFound: true,
146
+ },
147
+ {
148
+ name: "IPv6 address in IPv4 range",
149
+ rangeStr: "192.0.2.0-192.0.2.10",
150
+ ip: "2001:db8::",
151
+ wantFound: false,
152
+ },
153
}
154
94
- for name, test := range tests {
95
- name = fmt.Sprintf("%s (range: %s, ip: %s)", name, test.input, test.ip)
96
- t.Run(name, func(t *testing.T) {
97
- r, err := ParseRange(test.input)
155
+ for _, tt := range tests {
156
+ t.Run(tt.name, func(t *testing.T) {
157
+ t.Parallel()
158
+
159
+ r, err := ParseRange(tt.rangeStr)
160
require.NoError(t, err)
99
- ip := net.ParseIP(test.ip)
100
- require.NotNil(t, ip)
161
+ require.NotNil(t, r)
162
102
- if test.wantFail {
103
- assert.False(t, r.Contains(ip))
104
- } else {
105
- assert.True(t, r.Contains(ip))
106
- }
163
+ ip, err := netip.ParseAddr(tt.ip)
164
+ require.NoError(t, err)
165
+
166
+ assert.Equal(t, tt.wantFound, r.Contains(ip))
167
})
168
}
169
}
170
171
func TestV4Range_Iterate(t *testing.T) {
112
- tests := map[string]struct {
172
+ t.Parallel()
173
+
174
+ tests := []struct {
175
+ name string
176
input string
177
}{
115
- "Single IP": {input: "192.0.2.0"},
116
- "IP range": {input: "192.0.2.0-192.0.2.10"},
117
- "IP CIDR": {input: "192.0.2.0/24"},
178
+ {"single IP", "192.0.2.0"},
179
+ {"small range", "192.0.2.0-192.0.2.10"},
180
+ {"CIDR /30", "192.0.2.0/30"},
181
}
182
120
- for name, test := range tests {
121
- t.Run(name, func(t *testing.T) {
122
- r, err := ParseRange(test.input)
123
- require.NoError(t, err)
183
+ for _, tt := range tests {
184
+ t.Run(tt.name, func(t *testing.T) {
185
+ t.Parallel()
186
125
- var n int64
126
- for range r.Iterate() {
127
- n++
187
+ r, err := ParseRange(tt.input)
188
+ require.NoError(t, err)
189
+ require.NotNil(t, r)
190
+
191
+ // Count addresses yielded by iterator
192
+ var count int64
193
+ for addr := range r.Iterate() {
194
+ // Verify the address is valid and in range
195
+ assert.True(t, addr.IsValid())
196
+ assert.True(t, r.Contains(addr))
197
+ count++
198
}
129
- assert.Equal(t, r.Size().Int64(), n)
199
+
200
+ assert.Equal(t, r.Size().Int64(), count)
201
})
202
}
203
}
204
205
func TestV6Range_String(t *testing.T) {
135
- tests := map[string]struct {
206
+ t.Parallel()
207
+
208
+ tests := []struct {
209
+ name string
210
input string
211
wantString string
212
}{
139
- "IP": {input: "2001:db8::", wantString: "2001:db8::-2001:db8::"},
140
- "Range": {input: "2001:db8::-2001:db8::10", wantString: "2001:db8::-2001:db8::10"},
141
- "CIDR": {input: "2001:db8::/126", wantString: "2001:db8::1-2001:db8::2"},
213
+ {
214
+ name: "single IP",
215
+ input: "2001:db8::",
216
+ wantString: "2001:db8::-2001:db8::",
217
+ },
218
+ {
219
+ name: "IP range",
220
+ input: "2001:db8::-2001:db8::10",
221
+ wantString: "2001:db8::-2001:db8::10",
222
+ },
223
+ {
224
+ name: "CIDR /126",
225
+ input: "2001:db8::/126",
226
+ wantString: "2001:db8::1-2001:db8::2",
227
+ },
228
}
229
144
- for name, test := range tests {
145
- t.Run(name, func(t *testing.T) {
146
- r, err := ParseRange(test.input)
230
+ for _, tt := range tests {
231
+ t.Run(tt.name, func(t *testing.T) {
232
+ t.Parallel()
233
+
234
+ r, err := ParseRange(tt.input)
235
require.NoError(t, err)
236
+ require.NotNil(t, r)
237
149
- assert.Equal(t, test.wantString, r.String())
238
+ assert.Equal(t, tt.wantString, r.String())
239
})
240
}
241
}
242
243
func TestV6Range_Family(t *testing.T) {
155
- tests := map[string]struct {
244
+ t.Parallel()
245
+
246
+ tests := []struct {
247
+ name string
248
input string
249
}{
158
- "IP": {input: "2001:db8::"},
159
- "Range": {input: "2001:db8::-2001:db8::10"},
160
- "CIDR": {input: "2001:db8::/126"},
250
+ {"single IP", "2001:db8::"},
251
+ {"IP range", "2001:db8::-2001:db8::10"},
252
+ {"CIDR", "2001:db8::/126"},
253
}
254
163
- for name, test := range tests {
164
- t.Run(name, func(t *testing.T) {
165
- r, err := ParseRange(test.input)
255
+ for _, tt := range tests {
256
+ t.Run(tt.name, func(t *testing.T) {
257
+ t.Parallel()
258
+
259
+ r, err := ParseRange(tt.input)
260
require.NoError(t, err)
261
+ require.NotNil(t, r)
262
263
assert.Equal(t, V6Family, r.Family())
264
})
@@ -171,76 +266,239 @@ func TestV6Range_Family(t *testing.T) {
266
}
267
268
func TestV6Range_Size(t *testing.T) {
174
- tests := map[string]struct {
269
+ t.Parallel()
270
+
271
+ tests := []struct {
272
+ name string
273
input string
176
- wantSize *big.Int
274
+ wantSize int64
275
}{
178
- "IP": {input: "2001:db8::", wantSize: big.NewInt(1)},
179
- "Range": {input: "2001:db8::-2001:db8::10", wantSize: big.NewInt(17)},
180
- "CIDR": {input: "2001:db8::/120", wantSize: big.NewInt(254)},
181
- "CIDR 127": {input: "2001:db8::/127", wantSize: big.NewInt(2)},
182
- "CIDR 128": {input: "2001:db8::/128", wantSize: big.NewInt(1)},
276
+ {"single IP", "2001:db8::", 1},
277
+ {"IP range", "2001:db8::-2001:db8::10", 17},
278
+ {"CIDR /120", "2001:db8::/120", 254},
279
+ {"CIDR /127", "2001:db8::/127", 2},
280
+ {"CIDR /128", "2001:db8::/128", 1},
281
}
282
185
- for name, test := range tests {
186
- t.Run(name, func(t *testing.T) {
187
- r, err := ParseRange(test.input)
283
+ for _, tt := range tests {
284
+ t.Run(tt.name, func(t *testing.T) {
285
+ t.Parallel()
286
+
287
+ r, err := ParseRange(tt.input)
288
require.NoError(t, err)
289
+ require.NotNil(t, r)
290
190
- assert.Equal(t, test.wantSize, r.Size())
291
+ assert.Equal(t, big.NewInt(tt.wantSize), r.Size())
292
})
293
}
294
}
295
296
func TestV6Range_Contains(t *testing.T) {
196
- tests := map[string]struct {
197
- input string
198
- ip string
199
- wantFail bool
297
+ t.Parallel()
298
+
299
+ tests := []struct {
300
+ name string
301
+ rangeStr string
302
+ ip string
303
+ wantFound bool
304
}{
201
- "inside": {input: "2001:db8::-2001:db8::10", ip: "2001:db8::5"},
202
- "outside": {input: "2001:db8::-2001:db8::10", ip: "2001:db8::ff", wantFail: true},
203
- "eq start": {input: "2001:db8::-2001:db8::10", ip: "2001:db8::"},
204
- "eq end": {input: "2001:db8::-2001:db8::10", ip: "2001:db8::10"},
205
- "v4": {input: "2001:db8::-2001:db8::10", ip: "192.0.2.0", wantFail: true},
305
+ {
306
+ name: "IP inside range",
307
+ rangeStr: "2001:db8::-2001:db8::10",
308
+ ip: "2001:db8::5",
309
+ wantFound: true,
310
+ },
311
+ {
312
+ name: "IP outside range",
313
+ rangeStr: "2001:db8::-2001:db8::10",
314
+ ip: "2001:db8::ff",
315
+ wantFound: false,
316
+ },
317
+ {
318
+ name: "IP equals start",
319
+ rangeStr: "2001:db8::-2001:db8::10",
320
+ ip: "2001:db8::",
321
+ wantFound: true,
322
+ },
323
+ {
324
+ name: "IP equals end",
325
+ rangeStr: "2001:db8::-2001:db8::10",
326
+ ip: "2001:db8::10",
327
+ wantFound: true,
328
+ },
329
+ {
330
+ name: "IPv4 address in IPv6 range",
331
+ rangeStr: "2001:db8::-2001:db8::10",
332
+ ip: "192.0.2.0",
333
+ wantFound: false,
334
+ },
335
}
336
208
- for name, test := range tests {
209
- name = fmt.Sprintf("%s (range: %s, ip: %s)", name, test.input, test.ip)
210
- t.Run(name, func(t *testing.T) {
211
- r, err := ParseRange(test.input)
337
+ for _, tt := range tests {
338
+ t.Run(tt.name, func(t *testing.T) {
339
+ t.Parallel()
340
+
341
+ r, err := ParseRange(tt.rangeStr)
342
require.NoError(t, err)
213
- ip := net.ParseIP(test.ip)
214
- require.NotNil(t, ip)
343
+ require.NotNil(t, r)
344
216
- if test.wantFail {
217
- assert.False(t, r.Contains(ip))
218
- } else {
219
- assert.True(t, r.Contains(ip))
220
- }
345
+ ip, err := netip.ParseAddr(tt.ip)
346
+ require.NoError(t, err)
347
+
348
+ assert.Equal(t, tt.wantFound, r.Contains(ip))
349
})
350
}
351
}
352
353
func TestV6Range_Iterate(t *testing.T) {
226
- tests := map[string]struct {
354
+ t.Parallel()
355
+
356
+ tests := []struct {
357
+ name string
358
input string
359
}{
229
- "Single IP": {input: "2001:db8::5"},
230
- "IP range": {input: "2001:db8::-2001:db8::10"},
231
- "IP CIDR": {input: "2001:db8::/124"},
360
+ {"single IP", "2001:db8::5"},
361
+ {"small range", "2001:db8::-2001:db8::10"},
362
+ {"CIDR /124", "2001:db8::/124"},
363
+ }
364
+
365
+ for _, tt := range tests {
366
+ t.Run(tt.name, func(t *testing.T) {
367
+ t.Parallel()
368
+
369
+ r, err := ParseRange(tt.input)
370
+ require.NoError(t, err)
371
+ require.NotNil(t, r)
372
+
373
+ // Count addresses yielded by iterator
374
+ var count int64
375
+ for addr := range r.Iterate() {
376
+ // Verify the address is valid and in range
377
+ assert.True(t, addr.IsValid())
378
+ assert.True(t, r.Contains(addr))
379
+ count++
380
+ }
381
+
382
+ assert.Equal(t, r.Size().Int64(), count)
383
+ })
384
+ }
385
+}
386
+
387
+func TestNew(t *testing.T) {
388
+ t.Parallel()
389
+
390
+ tests := []struct {
391
+ name string
392
+ start string
393
+ end string
394
+ wantNil bool
395
+ wantFamily Family
396
+ }{
397
+ {
398
+ name: "valid IPv4 range",
399
+ start: "192.0.2.0",
400
+ end: "192.0.2.10",
401
+ wantFamily: V4Family,
402
+ },
403
+ {
404
+ name: "valid IPv6 range",
405
+ start: "2001:db8::",
406
+ end: "2001:db8::10",
407
+ wantFamily: V6Family,
408
+ },
409
+ {
410
+ name: "IPv4 start > end",
411
+ start: "192.0.2.10",
412
+ end: "192.0.2.0",
413
+ wantNil: true,
414
+ },
415
+ {
416
+ name: "IPv6 start > end",
417
+ start: "2001:db8::10",
418
+ end: "2001:db8::",
419
+ wantNil: true,
420
+ },
421
+ {
422
+ name: "mixed families",
423
+ start: "192.0.2.0",
424
+ end: "2001:db8::",
425
+ wantNil: true,
426
+ },
427
}
428
234
- for name, test := range tests {
235
- t.Run(name, func(t *testing.T) {
236
- r, err := ParseRange(test.input)
429
+ for _, tt := range tests {
430
+ t.Run(tt.name, func(t *testing.T) {
431
+ t.Parallel()
432
+
433
+ start, err := netip.ParseAddr(tt.start)
434
require.NoError(t, err)
435
239
- var n int64
240
- for range r.Iterate() {
241
- n++
436
+ end, err := netip.ParseAddr(tt.end)
437
+ require.NoError(t, err)
438
+
439
+ r := New(start, end)
440
+
441
+ if tt.wantNil {
442
+ assert.Nil(t, r)
443
+ } else {
444
+ require.NotNil(t, r)
445
+ assert.Equal(t, tt.wantFamily, r.Family())
446
+ assert.Equal(t, start, r.Start())
447
+ assert.Equal(t, end, r.End())
448
}
243
- assert.Equal(t, r.Size().Int64(), n)
449
})
450
}
451
}
452
+
453
+func TestFamily_String(t *testing.T) {
454
+ t.Parallel()
455
+
456
+ tests := []struct {
457
+ family Family
458
+ want string
459
+ }{
460
+ {V4Family, "IPv4"},
461
+ {V6Family, "IPv6"},
462
+ {Family(99), "Unknown(99)"},
463
+ }
464
+
465
+ for _, tt := range tests {
466
+ t.Run(tt.want, func(t *testing.T) {
467
+ t.Parallel()
468
+ assert.Equal(t, tt.want, tt.family.String())
469
+ })
470
+ }
471
+}
472
+
473
+// Benchmark tests
474
+func BenchmarkV4Range_Contains(b *testing.B) {
475
+ r, _ := ParseRange("192.0.2.0/24")
476
+ ip := netip.MustParseAddr("192.0.2.100")
477
+
478
+ b.ResetTimer()
479
+ for i := 0; i < b.N; i++ {
480
+ _ = r.Contains(ip)
481
+ }
482
+}
483
+
484
+func BenchmarkV6Range_Contains(b *testing.B) {
485
+ r, _ := ParseRange("2001:db8::/64")
486
+ ip := netip.MustParseAddr("2001:db8::1234")
487
+
488
+ b.ResetTimer()
489
+ for i := 0; i < b.N; i++ {
490
+ _ = r.Contains(ip)
491
+ }
492
+}
493
+
494
+func BenchmarkV4Range_Iterate(b *testing.B) {
495
+ r, _ := ParseRange("192.0.2.0/24")
496
+
497
+ b.ResetTimer()
498
+ for i := 0; i < b.N; i++ {
499
+ count := 0
500
+ for range r.Iterate() {
501
+ count++
502
+ }
503
+ }
504
+}