master
go 504 lines 9.44 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package iprange
4
5 import (
6 "math/big"
7 "net/netip"
8 "testing"
9
10 "github.com/stretchr/testify/assert"
11 "github.com/stretchr/testify/require"
12 )
13
14 func TestV4Range_String(t *testing.T) {
15 t.Parallel()
16
17 tests := []struct {
18 name string
19 input string
20 wantString string
21 }{
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
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
52 assert.Equal(t, tt.wantString, r.String())
53 })
54 }
55 }
56
57 func TestV4Range_Family(t *testing.T) {
58 t.Parallel()
59
60 tests := []struct {
61 name string
62 input string
63 }{
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
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 })
80 }
81 }
82
83 func TestV4Range_Size(t *testing.T) {
84 t.Parallel()
85
86 tests := []struct {
87 name string
88 input string
89 wantSize int64
90 }{
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
109 assert.Equal(t, big.NewInt(tt.wantSize), r.Size())
110 })
111 }
112 }
113
114 func TestV4Range_Contains(t *testing.T) {
115 t.Parallel()
116
117 tests := []struct {
118 name string
119 rangeStr string
120 ip string
121 wantFound bool
122 }{
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
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)
161 require.NotNil(t, r)
162
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) {
172 t.Parallel()
173
174 tests := []struct {
175 name string
176 input string
177 }{
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
183 for _, tt := range tests {
184 t.Run(tt.name, func(t *testing.T) {
185 t.Parallel()
186
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 }
199
200 assert.Equal(t, r.Size().Int64(), count)
201 })
202 }
203 }
204
205 func TestV6Range_String(t *testing.T) {
206 t.Parallel()
207
208 tests := []struct {
209 name string
210 input string
211 wantString string
212 }{
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
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
238 assert.Equal(t, tt.wantString, r.String())
239 })
240 }
241 }
242
243 func TestV6Range_Family(t *testing.T) {
244 t.Parallel()
245
246 tests := []struct {
247 name string
248 input string
249 }{
250 {"single IP", "2001:db8::"},
251 {"IP range", "2001:db8::-2001:db8::10"},
252 {"CIDR", "2001:db8::/126"},
253 }
254
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 })
265 }
266 }
267
268 func TestV6Range_Size(t *testing.T) {
269 t.Parallel()
270
271 tests := []struct {
272 name string
273 input string
274 wantSize int64
275 }{
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
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
291 assert.Equal(t, big.NewInt(tt.wantSize), r.Size())
292 })
293 }
294 }
295
296 func TestV6Range_Contains(t *testing.T) {
297 t.Parallel()
298
299 tests := []struct {
300 name string
301 rangeStr string
302 ip string
303 wantFound bool
304 }{
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
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)
343 require.NotNil(t, r)
344
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) {
354 t.Parallel()
355
356 tests := []struct {
357 name string
358 input string
359 }{
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
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
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 }
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 }