mirror of
https://github.com/maxlerebourg/crowdsec-bouncer-traefik-plugin.git
synced 2026-09-02 20:28:50 +02:00
🐛 Support range decision on stream mode
This commit is contained in:
@@ -0,0 +1,48 @@
|
||||
package ip
|
||||
|
||||
import (
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// CIDRKeys returns all possible CIDR prefixes of an IP, from the most specific (/32 for IPv4, /128 for IPv6) to the least specific (/0).
|
||||
func CIDRKeys(ipStr string) []string {
|
||||
ip := net.ParseIP(ipStr)
|
||||
if ip == nil {
|
||||
return nil
|
||||
}
|
||||
ip4 := ip.To4()
|
||||
if ip4 != nil {
|
||||
keys := make([]string, 0, 33)
|
||||
for bits := 32; bits >= 0; bits-- {
|
||||
mask := net.CIDRMask(bits, 32)
|
||||
n := make(net.IP, 4)
|
||||
for i := 0; i < 4; i++ {
|
||||
n[i] = ip4[i] & mask[i]
|
||||
}
|
||||
keys = append(keys, n.String()+"/"+strconv.Itoa(bits))
|
||||
}
|
||||
return keys
|
||||
}
|
||||
ip16 := ip.To16()
|
||||
keys := make([]string, 0, 129)
|
||||
for bits := 128; bits >= 0; bits-- {
|
||||
mask := net.CIDRMask(bits, 128)
|
||||
n := make(net.IP, 16)
|
||||
for i := 0; i < 16; i++ {
|
||||
n[i] = ip16[i] & mask[i]
|
||||
}
|
||||
keys = append(keys, n.String()+"/"+strconv.Itoa(bits))
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
// NormalizeCIDR parses a CIDR string and returns its normalized form, or an empty string if invalid.
|
||||
func NormalizeCIDR(cidrStr string) string {
|
||||
_, ipNet, err := net.ParseCIDR(strings.TrimSpace(cidrStr))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return ipNet.String()
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
package ip
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCIDRKeys(t *testing.T) {
|
||||
tests := []struct {
|
||||
ip string
|
||||
wantKeys int
|
||||
checks map[int]string
|
||||
}{
|
||||
{
|
||||
ip: "10.0.0.1",
|
||||
wantKeys: 33,
|
||||
checks: map[int]string{0: "10.0.0.1/32", 8: "10.0.0.0/24", 32: "0.0.0.0/0"},
|
||||
},
|
||||
{
|
||||
ip: "2001:db8::1",
|
||||
wantKeys: 129,
|
||||
checks: map[int]string{0: "2001:db8::1/128", 32: "2001:db8::/96", 128: "::/0"},
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.ip, func(t *testing.T) {
|
||||
keys := CIDRKeys(tt.ip)
|
||||
if keys == nil {
|
||||
t.Fatal("CIDRKeys returned nil")
|
||||
}
|
||||
if len(keys) != tt.wantKeys {
|
||||
t.Fatalf("expected %d keys, got %d", tt.wantKeys, len(keys))
|
||||
}
|
||||
for idx, want := range tt.checks {
|
||||
if keys[idx] != want {
|
||||
t.Errorf("keys[%d] should be %s, got %s", idx, want, keys[idx])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIDRKeys_MostToLeastSpecific(t *testing.T) {
|
||||
ips := []string{"10.0.0.1", "2001:db8::1"}
|
||||
for _, ip := range ips {
|
||||
t.Run(ip, func(t *testing.T) {
|
||||
keys := CIDRKeys(ip)
|
||||
for i := 1; i < len(keys); i++ {
|
||||
prevBits := strings.Split(keys[i-1], "/")[1]
|
||||
curBits := strings.Split(keys[i], "/")[1]
|
||||
prevN, _ := strconv.Atoi(prevBits)
|
||||
curN, _ := strconv.Atoi(curBits)
|
||||
if prevN <= curN {
|
||||
t.Errorf("keys should go from most specific to least specific at index %d: /%d <= /%d", i, prevN, curN)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIDRKeys_IPVariants(t *testing.T) {
|
||||
tests := []struct {
|
||||
ip string
|
||||
wantKeys int
|
||||
}{
|
||||
{"0.0.0.0", 33},
|
||||
{"255.255.255.255", 33},
|
||||
{"1.2.3.4", 33},
|
||||
{"10.0.0.1", 33},
|
||||
{"192.168.1.1", 33},
|
||||
{"invalid", 0},
|
||||
{"", 0},
|
||||
{" ", 0},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.ip, func(t *testing.T) {
|
||||
keys := CIDRKeys(tt.ip)
|
||||
if len(keys) != tt.wantKeys {
|
||||
t.Errorf("CIDRKeys(%q) returned %d keys, want %d", tt.ip, len(keys), tt.wantKeys)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCIDRKeys_VerifyNetworkAddress(t *testing.T) {
|
||||
keys := CIDRKeys("10.1.2.3")
|
||||
tests := []struct {
|
||||
bits int
|
||||
want string
|
||||
}{
|
||||
{24, "10.1.2.0/24"},
|
||||
{16, "10.1.0.0/16"},
|
||||
{8, "10.0.0.0/8"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := keys[32-tt.bits]; got != tt.want {
|
||||
t.Errorf("/%d network should be %s, got %s", tt.bits, tt.want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeCIDR(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"10.0.0.0/8", "10.0.0.0/8"},
|
||||
{"10.0.0.0/16", "10.0.0.0/16"},
|
||||
{"192.168.1.0/24", "192.168.1.0/24"},
|
||||
{"2001:db8::/32", "2001:db8::/32"},
|
||||
{"0.0.0.0/0", "0.0.0.0/0"},
|
||||
{"::/0", "::/0"},
|
||||
{"invalid", ""},
|
||||
{"", ""},
|
||||
{" 10.0.0.0/8 ", "10.0.0.0/8"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
got := NormalizeCIDR(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("NormalizeCIDR(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user