| 1 | // SPDX-License-Identifier: GPL-3.0-or-later |
| 2 | |
| 3 | package iprange |
| 4 | |
| 5 | import ( |
| 6 | "errors" |
| 7 | "fmt" |
| 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 | |
| 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 | |
| 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) { |
| 37 | s = strings.TrimSpace(s) |
| 38 | if s == "" { |
| 39 | return nil, nil |
| 40 | } |
| 41 | |
| 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 { |
| 48 | return nil, fmt.Errorf("parsing %q: %w", part, err) |
| 49 | } |
| 50 | if r != nil { |
| 51 | ranges = append(ranges, r) |
| 52 | } |
| 53 | } |
| 54 | |
| 55 | return ranges, nil |
| 56 | } |
| 57 | |
| 58 | // ParseRange parses s as a single IP range. |
| 59 | // See ParseRanges for supported formats. |
| 60 | func ParseRange(s string) (Range, error) { |
| 61 | s = strings.TrimSpace(strings.ToLower(s)) |
| 62 | if s == "" { |
| 63 | return nil, nil |
| 64 | } |
| 65 | |
| 66 | // Try different formats in order of likelihood |
| 67 | switch { |
| 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 | |
| 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 | } |
| 86 | return New(addr, addr), nil |
| 87 | } |
| 88 | |
| 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 | |
| 115 | return New(start, end), nil |
| 116 | } |
| 117 | |
| 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 { |
| 122 | return nil, fmt.Errorf("%w: %v", ErrInvalidSyntax, err) |
| 123 | } |
| 124 | |
| 125 | // Normalize to network address |
| 126 | prefix = prefix.Masked() |
| 127 | |
| 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 | |
| 142 | return New(start, end), nil |
| 143 | } |
| 144 | |
| 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 { |
| 149 | return nil, fmt.Errorf("%w: invalid subnet mask format", ErrInvalidSyntax) |
| 150 | } |
| 151 | |
| 152 | addrStr := s[:idx] |
| 153 | maskStr := s[idx+1:] |
| 154 | |
| 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 | |
| 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 | |
| 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 | |
| 173 | // Create CIDR notation and parse it |
| 174 | return parseCIDR(fmt.Sprintf("%s/%d", addrStr, prefixLen)) |
| 175 | } |
| 176 | |
| 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 | |
| 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 | |
| 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 | } |