@cryptotaxi247 / netdata-1 / commits / a032a85e2

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 +}