master
go 488 lines 10.6 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package iprange
4
5 import (
6 "fmt"
7 "math/big"
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) {
108 t.Parallel()
109
110 tests := []struct {
111 name string
112 input string
113 wantString string
114 }{
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 },
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
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)
138
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) {
168 t.Parallel()
169
170 tests := []struct {
171 name string
172 input string
173 wantSize int64
174 }{
175 {
176 name: "single range",
177 input: "192.0.2.0-192.0.2.10",
178 wantSize: 11,
179 },
180 {
181 name: "multiple ranges",
182 input: "192.0.2.0-192.0.2.10 2001:db8::-2001:db8::10",
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
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)
203
204 assert.Equal(t, big.NewInt(tt.wantSize), pool.Size())
205 })
206 }
207 }
208
209 func TestPool_Contains(t *testing.T) {
210 t.Parallel()
211
212 tests := []struct {
213 name string
214 poolStr string
215 ip string
216 wantFound bool
217 }{
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 },
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 },
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
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)
256
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 := range 10 {
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 }