Compare commits

..
Author SHA1 Message Date
mhxandClaude Opus 5 f359d5d935 cache: keep redis readers by pointer
A pooled SimpleRedis holds a sync.Mutex, so appending one into rc.readers
by value copies the lock and trips go vet's copylocks check. Keep the
readers by pointer instead; the round-robin over replicas is unchanged.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01D95Nh68xKXhynPXHzozrXp
2026-08-25 19:40:35 +02:00
10 changed files with 28 additions and 754 deletions
+1 -19
View File
@@ -330,7 +330,7 @@ func New(_ context.Context, next http.Handler, config *configuration.Config, nam
// ServeHTTP principal function of plugin. // ServeHTTP principal function of plugin.
// //
//nolint:nestif,gocognit //nolint:nestif
func (bouncer *Bouncer) ServeHTTP(rw http.ResponseWriter, req *http.Request) { func (bouncer *Bouncer) ServeHTTP(rw http.ResponseWriter, req *http.Request) {
if !bouncer.enabled { if !bouncer.enabled {
bouncer.next.ServeHTTP(rw, req) bouncer.next.ServeHTTP(rw, req)
@@ -392,16 +392,6 @@ func (bouncer *Bouncer) ServeHTTP(rw http.ResponseWriter, req *http.Request) {
// Right here if we cannot join the stream we forbid the request to go on. // Right here if we cannot join the stream we forbid the request to go on.
if bouncer.crowdsecMode == configuration.StreamMode || bouncer.crowdsecMode == configuration.AloneMode { if bouncer.crowdsecMode == configuration.StreamMode || bouncer.crowdsecMode == configuration.AloneMode {
if isCrowdsecStreamHealthy { if isCrowdsecStreamHealthy {
cidrValue, cidrErr := bouncer.cacheClient.GetCIDR(remoteIP)
if cidrErr == nil {
bouncer.log.Debug(fmt.Sprintf("ServeHTTP ip:%s cidr:hit isBanned:%v", remoteIP, cidrValue))
if cidrValue == cache.NoBannedValue {
bouncer.handleNextServeHTTP(rw, req, remoteIP)
} else {
bouncer.handleRemediationServeHTTP(rw, req, remoteIP, cidrValue)
}
return
}
bouncer.handleNextServeHTTP(rw, req, remoteIP) bouncer.handleNextServeHTTP(rw, req, remoteIP)
} else { } else {
bouncer.log.Debug(fmt.Sprintf("ServeHTTP isCrowdsecStreamHealthy:false ip:%s updateFailure:%d", remoteIP, updateFailure)) bouncer.log.Debug(fmt.Sprintf("ServeHTTP isCrowdsecStreamHealthy:false ip:%s updateFailure:%d", remoteIP, updateFailure))
@@ -684,20 +674,12 @@ func handleStreamCache(bouncer *Bouncer) error {
default: default:
bouncer.log.Info("handleStreamCache:unknownType " + decision.Type) bouncer.log.Info("handleStreamCache:unknownType " + decision.Type)
} }
if strings.Contains(decision.Value, "/") {
bouncer.cacheClient.SetCIDR(decision.Value, value, int64(duration.Seconds()))
} else {
bouncer.cacheClient.Set(decision.Value, value, int64(duration.Seconds())) bouncer.cacheClient.Set(decision.Value, value, int64(duration.Seconds()))
} }
} }
}
for _, decision := range stream.Deleted { for _, decision := range stream.Deleted {
if strings.Contains(decision.Value, "/") {
bouncer.cacheClient.DeleteCIDR(decision.Value)
} else {
bouncer.cacheClient.Delete(decision.Value) bouncer.cacheClient.Delete(decision.Value)
} }
}
bouncer.log.Debug("handleStreamCache:updated") bouncer.log.Debug("handleStreamCache:updated")
isCrowdsecStreamStartup = false isCrowdsecStreamStartup = false
return nil return nil
+5 -90
View File
@@ -6,14 +6,10 @@ import (
"errors" "errors"
"fmt" "fmt"
"log/slog" "log/slog"
"strconv"
"strings"
"sync/atomic" "sync/atomic"
ttl_map "github.com/leprosus/golang-ttl-map" ttl_map "github.com/leprosus/golang-ttl-map"
simpleredis "github.com/maxlerebourg/simpleredis" simpleredis "github.com/maxlerebourg/simpleredis"
"github.com/maxlerebourg/crowdsec-bouncer-traefik-plugin/pkg/ip"
) )
const ( const (
@@ -29,17 +25,6 @@ const (
CacheMiss = "cache:miss" CacheMiss = "cache:miss"
// CacheUnreachable error string when cache is unreachable. // CacheUnreachable error string when cache is unreachable.
CacheUnreachable = "cache:unreachable" CacheUnreachable = "cache:unreachable"
// cidrPrefix namespaces a CIDR decision, one cache key per CIDR.
cidrPrefix = "cidr:"
// cidrPrefixLensKey holds the prefix lengths that have a decision, so a lookup probes
// only those. Its absence means no CIDR decision was ever stored.
cidrPrefixLensKey = "cidrprefixlens"
// cidrPrefixLensSeparator separates the prefix lengths in cidrPrefixLensKey.
cidrPrefixLensSeparator = ","
// cidrPrefixLensDuration has to outlive every decision it describes, so it is
// effectively infinite. It cannot be zero: the local cache ignores a zero duration
// and redis rejects a non positive EX.
cidrPrefixLensDuration = 10 * 365 * 24 * 60 * 60
) )
//nolint:gochecknoglobals //nolint:gochecknoglobals
@@ -67,7 +52,7 @@ func (localCache) delete(key string) {
type redisCache struct { type redisCache struct {
log *slog.Logger log *slog.Logger
writer simpleredis.SimpleRedis writer simpleredis.SimpleRedis
readers []simpleredis.SimpleRedis readers []*simpleredis.SimpleRedis
counter atomic.Uint64 counter atomic.Uint64
} }
@@ -77,7 +62,7 @@ func (rc *redisCache) nextReader() *simpleredis.SimpleRedis {
return &rc.writer return &rc.writer
} }
idx := rc.counter.Add(1) % uint64(n) idx := rc.counter.Add(1) % uint64(n)
return &rc.readers[idx] return rc.readers[idx]
} }
func (rc *redisCache) get(key string) (string, error) { func (rc *redisCache) get(key string) (string, error) {
@@ -130,7 +115,9 @@ func (c *Client) New(log *slog.Logger, isRedis bool, writeHost string, readHosts
rc := &redisCache{log: log} rc := &redisCache{log: log}
rc.writer.Init(writeHost, pass, database) rc.writer.Init(writeHost, pass, database)
for _, h := range readHosts { for _, h := range readHosts {
var r simpleredis.SimpleRedis // A pooled SimpleRedis holds a mutex, so it is kept by pointer:
// appending it by value would copy the lock along with it.
r := &simpleredis.SimpleRedis{}
r.Init(h, pass, database) r.Init(h, pass, database)
rc.readers = append(rc.readers, r) rc.readers = append(rc.readers, r)
} }
@@ -159,75 +146,3 @@ func (c *Client) Set(key string, value string, duration int64) {
c.log.Debug(fmt.Sprintf("cache:Set key:%v value:%v duration:%vs", key, value, duration)) c.log.Debug(fmt.Sprintf("cache:Set key:%v value:%v duration:%vs", key, value, duration))
c.cache.set(key, value, duration) c.cache.set(key, value, duration)
} }
// DeleteCIDR removes a CIDR decision from the cache.
func (c *Client) DeleteCIDR(cidr string) {
normalized := ip.NormalizeCIDR(cidr)
if normalized == "" {
c.log.Error(fmt.Sprintf("cache:DeleteCIDR:invalidCIDR cidr:%v decision is left in cache", cidr))
return
}
cidr = normalized
c.cache.delete(cidrPrefix + cidr)
c.log.Debug(fmt.Sprintf("cache:DeleteCIDR cidr:%v", cidr))
}
// GetCIDR checks if an IP matches a CIDR decision in the cache.
// Only probes the prefix lengths that have a decision: it is on the request path.
func (c *Client) GetCIDR(ipStr string) (string, error) {
prefixLens, err := c.cache.get(cidrPrefixLensKey)
if err != nil {
return "", err
}
for _, key := range ip.CIDRLookupKeys(ipStr, parsePrefixLens(prefixLens)) {
value, getErr := c.cache.get(cidrPrefix + key)
if getErr == nil {
return value, nil
}
}
return "", errors.New(CacheMiss)
}
// SetCIDR stores a CIDR decision in the cache.
func (c *Client) SetCIDR(cidr, value string, duration int64) {
normalized := ip.NormalizeCIDR(cidr)
prefixLen := ip.CIDRPrefixLen(cidr)
if normalized == "" || prefixLen < 0 {
c.log.Error(fmt.Sprintf("cache:SetCIDR:invalidCIDR cidr:%v value:%v decision is not enforced", cidr, value))
return
}
// Publish the length first, or a concurrent lookup misses the decision.
c.addCIDRPrefixLen(prefixLen)
c.cache.set(cidrPrefix+normalized, value, duration)
c.log.Debug(fmt.Sprintf("cache:SetCIDR cidr:%v value:%v duration:%vs", normalized, value, duration))
}
// addCIDRPrefixLen records a prefix length in the set probed on lookup. The set only grows:
// a stale length costs one extra read, dropping one too early leaves decisions unmatched.
func (c *Client) addCIDRPrefixLen(prefixLen int) {
prefixLens, err := c.cache.get(cidrPrefixLensKey)
if err == nil {
for _, known := range parsePrefixLens(prefixLens) {
if known == prefixLen {
return
}
}
prefixLens += cidrPrefixLensSeparator + strconv.Itoa(prefixLen)
} else {
prefixLens = strconv.Itoa(prefixLen)
}
c.cache.set(cidrPrefixLensKey, prefixLens, cidrPrefixLensDuration)
c.log.Debug(fmt.Sprintf("cache:addCIDRPrefixLen prefixLens:%v", prefixLens))
}
// parsePrefixLens decodes the set of prefix lengths stored in cidrPrefixLensKey.
func parsePrefixLens(value string) []int {
fields := strings.Split(value, cidrPrefixLensSeparator)
prefixLens := make([]int, 0, len(fields))
for _, field := range fields {
if prefixLen, err := strconv.Atoi(field); err == nil {
prefixLens = append(prefixLens, prefixLen)
}
}
return prefixLens
}
+5 -186
View File
@@ -3,7 +3,6 @@
package cache package cache
import ( import (
"errors"
"testing" "testing"
logger "github.com/maxlerebourg/crowdsec-bouncer-traefik-plugin/pkg/logger" logger "github.com/maxlerebourg/crowdsec-bouncer-traefik-plugin/pkg/logger"
@@ -131,7 +130,7 @@ func indexOfReader(rc *redisCache, r *simpleredis.SimpleRedis) int {
return -1 return -1
} }
for i := range rc.readers { for i := range rc.readers {
if r == &rc.readers[i] { if r == rc.readers[i] {
return i return i
} }
} }
@@ -152,7 +151,10 @@ func Test_nextReader(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
rc := &redisCache{log: logger.New("INFO", "")} rc := &redisCache{log: logger.New("INFO", "")}
rc.readers = make([]simpleredis.SimpleRedis, tt.readers) rc.readers = make([]*simpleredis.SimpleRedis, tt.readers)
for i := range rc.readers {
rc.readers[i] = &simpleredis.SimpleRedis{}
}
for call, want := range tt.want { for call, want := range tt.want {
if got := indexOfReader(rc, rc.nextReader()); got != want { if got := indexOfReader(rc, rc.nextReader()); got != want {
t.Errorf("call %d: nextReader() -> reader[%d], want reader[%d]", call, got, want) t.Errorf("call %d: nextReader() -> reader[%d], want reader[%d]", call, got, want)
@@ -161,186 +163,3 @@ func Test_nextReader(t *testing.T) {
}) })
} }
} }
// countingCache is an isolated cacheInterface recording how many reads a lookup costs,
// so a CIDR lookup can be checked for both its result and its price.
type countingCache struct {
values map[string]string
reads int
}
func newCountingCache() *countingCache {
return &countingCache{values: map[string]string{}}
}
func (c *countingCache) get(key string) (string, error) {
c.reads++
if value, found := c.values[key]; found && value != "" {
return value, nil
}
return "", errors.New(CacheMiss)
}
func (c *countingCache) set(key, value string, _ int64) {
c.values[key] = value
}
func (c *countingCache) delete(key string) {
delete(c.values, key)
}
func newCIDRClient(decisions map[string]string) (*Client, *countingCache) {
counting := newCountingCache()
client := &Client{cache: counting, log: logger.New("INFO", "")}
for cidr, value := range decisions {
client.SetCIDR(cidr, value, 60)
}
return client, counting
}
func Test_GetCIDR(t *testing.T) {
decisions := map[string]string{
"10.0.0.0/24": BannedValue,
"192.168.1.42/24": CaptchaValue, // not a network address, host bits are dropped
"2001:db8::/32": BannedValue,
}
tests := []struct {
name string
clientIP string
want string
wantErr bool
}{
{name: "IP inside a banned range", clientIP: "10.0.0.7", want: BannedValue},
{name: "network address itself", clientIP: "10.0.0.0", want: BannedValue},
{name: "broadcast address of the range", clientIP: "10.0.0.255", want: BannedValue},
{name: "IP just outside the range", clientIP: "10.0.1.0", wantErr: true},
{name: "IP inside a captcha range", clientIP: "192.168.1.7", want: CaptchaValue},
{name: "IP inside an IPv6 range", clientIP: "2001:db8::dead:beef", want: BannedValue},
{name: "IP outside the IPv6 range", clientIP: "2001:db9::1", wantErr: true},
{name: "IPv4 mapped client against an IPv4 range", clientIP: "::ffff:10.0.0.7", want: BannedValue},
{name: "unknown IP", clientIP: "8.8.8.8", wantErr: true},
{name: "invalid IP", clientIP: "not-an-ip", wantErr: true},
{name: "empty IP", clientIP: "", wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
client, _ := newCIDRClient(decisions)
got, err := client.GetCIDR(tt.clientIP)
if (err != nil) != tt.wantErr {
t.Fatalf("GetCIDR(%q) error = %v, wantErr %v", tt.clientIP, err, tt.wantErr)
}
if got != tt.want {
t.Errorf("GetCIDR(%q) = %q, want %q", tt.clientIP, got, tt.want)
}
})
}
}
// Test_GetCIDR_MostSpecific pins precedence: the narrowest range wins, so a captcha on
// a /24 is not overruled by a ban on its /8.
func Test_GetCIDR_MostSpecific(t *testing.T) {
client, _ := newCIDRClient(map[string]string{
"10.0.0.0/8": BannedValue,
"10.1.0.0/16": CaptchaValue,
"10.1.2.0/24": BannedValue,
"2001:db8::/32": BannedValue,
"2001:db8::/48": CaptchaValue,
})
tests := []struct {
clientIP string
want string
}{
{clientIP: "10.1.2.3", want: BannedValue},
{clientIP: "10.1.3.3", want: CaptchaValue},
{clientIP: "10.2.3.4", want: BannedValue},
{clientIP: "2001:db8::1", want: CaptchaValue},
{clientIP: "2001:db8:1::1", want: BannedValue},
}
for _, tt := range tests {
t.Run(tt.clientIP, func(t *testing.T) {
got, err := client.GetCIDR(tt.clientIP)
if err != nil {
t.Fatalf("GetCIDR(%q) unexpected error %v", tt.clientIP, err)
}
if got != tt.want {
t.Errorf("GetCIDR(%q) = %q, want %q", tt.clientIP, got, tt.want)
}
})
}
}
func Test_DeleteCIDR(t *testing.T) {
client, _ := newCIDRClient(map[string]string{
"10.0.0.0/8": BannedValue,
"10.1.2.0/24": CaptchaValue,
})
client.DeleteCIDR("10.1.2.0/24")
// The wider decision is untouched and takes over.
if got, err := client.GetCIDR("10.1.2.3"); err != nil || got != BannedValue {
t.Errorf("after deleting the /24, GetCIDR = %q %v, want %q", got, err, BannedValue)
}
client.DeleteCIDR("10.0.0.0/8")
if _, err := client.GetCIDR("10.1.2.3"); err == nil {
t.Error("GetCIDR should miss once every decision is deleted")
}
}
func Test_SetCIDR_InvalidIsNotStored(t *testing.T) {
for _, cidr := range []string{"", "garbage", "10.0.0.1", "10.0.0.0/33", "10.0.0.0/-1"} {
t.Run(cidr, func(t *testing.T) {
_, counting := newCIDRClient(map[string]string{cidr: BannedValue})
if len(counting.values) != 0 {
t.Errorf("SetCIDR(%q) stored %v, want nothing", cidr, counting.values)
}
})
}
}
// Test_GetCIDR_Reads guards the cost of the lookup: it must probe only the prefix
// lengths that have a decision, not every possible one.
func Test_GetCIDR_Reads(t *testing.T) {
tests := []struct {
name string
decisions map[string]string
clientIP string
wantReads int
}{
{name: "no decision at all, IPv4", decisions: nil, clientIP: "10.0.0.1", wantReads: 1},
{name: "no decision at all, IPv6", decisions: nil, clientIP: "2001:db8::1", wantReads: 1},
{
name: "one prefix length, hit",
decisions: map[string]string{"10.0.0.0/24": BannedValue},
clientIP: "10.0.0.1",
wantReads: 2,
},
{
name: "one prefix length, miss",
decisions: map[string]string{"10.0.0.0/24": BannedValue},
clientIP: "11.0.0.1",
wantReads: 2,
},
{
name: "three prefix lengths, miss probes each once",
decisions: map[string]string{"10.0.0.0/8": BannedValue, "10.1.0.0/16": BannedValue, "10.1.2.0/24": BannedValue},
clientIP: "11.0.0.1",
wantReads: 4,
},
{
name: "IPv6 client does not probe every length",
decisions: map[string]string{"2001:db8::/32": BannedValue},
clientIP: "2001:dead::1",
wantReads: 2,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
client, counting := newCIDRClient(tt.decisions)
counting.reads = 0
// Only the number of reads matters here, the result is covered above.
_, _ = client.GetCIDR(tt.clientIP)
if counting.reads != tt.wantReads {
t.Errorf("GetCIDR(%q) did %d cache reads, want %d", tt.clientIP, counting.reads, tt.wantReads)
}
})
}
}
-85
View File
@@ -1,85 +0,0 @@
package ip
import (
"net"
"strings"
)
const (
maxIPv4PrefixLen = 32
maxIPv6PrefixLen = 128
)
// 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 {
parsed, maxBits := parseForPrefix(ipStr)
if parsed == nil {
return nil
}
keys := make([]string, 0, maxBits+1)
for bits := maxBits; bits >= 0; bits-- {
keys = append(keys, cidrKey(parsed, bits, maxBits))
}
return keys
}
// CIDRLookupKeys returns the keys of the CIDRs containing an IP for the given prefix lengths
// only, most specific first. Duplicates and lengths of the other family are skipped.
func CIDRLookupKeys(ipStr string, prefixLens []int) []string {
parsed, maxBits := parseForPrefix(ipStr)
if parsed == nil {
return nil
}
var wanted [maxIPv6PrefixLen + 1]bool
for _, bits := range prefixLens {
if bits >= 0 && bits <= maxBits {
wanted[bits] = true
}
}
keys := make([]string, 0, len(prefixLens))
for bits := maxBits; bits >= 0; bits-- {
if wanted[bits] {
keys = append(keys, cidrKey(parsed, bits, maxBits))
}
}
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()
}
// CIDRPrefixLen returns the prefix length of a CIDR, or -1 if it is not a valid CIDR.
func CIDRPrefixLen(cidrStr string) int {
_, ipNet, err := net.ParseCIDR(strings.TrimSpace(cidrStr))
if err != nil {
return -1
}
prefixLen, _ := ipNet.Mask.Size()
return prefixLen
}
// parseForPrefix returns the IP in the native form of its family, and that family's bit length.
func parseForPrefix(ipStr string) (net.IP, int) {
parsed := net.ParseIP(ipStr)
if parsed == nil {
return nil, 0
}
if parsed4 := parsed.To4(); parsed4 != nil {
return parsed4, maxIPv4PrefixLen
}
return parsed.To16(), maxIPv6PrefixLen
}
// cidrKey builds the key of the CIDR of bits length containing the IP.
// It formats through net.IPNet like NormalizeCIDR, so writes and lookups agree.
func cidrKey(parsed net.IP, bits, maxBits int) string {
mask := net.CIDRMask(bits, maxBits)
ipNet := net.IPNet{IP: parsed.Mask(mask), Mask: mask}
return ipNet.String()
}
-301
View File
@@ -1,301 +0,0 @@
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 TestCIDRKeys_VerifyNetworkAddressIPv6(t *testing.T) {
keys := CIDRKeys("2001:db8:1:2:3:4:5:6")
tests := []struct {
bits int
want string
}{
{64, "2001:db8:1:2::/64"},
{48, "2001:db8:1::/48"},
{32, "2001:db8::/32"},
}
for _, tt := range tests {
if got := keys[128-tt.bits]; got != tt.want {
t.Errorf("/%d network should be %s, got %s", tt.bits, tt.want, got)
}
}
}
// TestCIDRKeys_MatchNormalizeCIDR covers the invariant range support rests on: the key
// written for a decision is a key looked up for the IPs it covers, and only those.
func TestCIDRKeys_MatchNormalizeCIDR(t *testing.T) {
tests := []struct {
cidr string
ip string
match bool
}{
{cidr: "10.0.0.0/8", ip: "10.1.2.3", match: true},
{cidr: "10.0.0.0/24", ip: "10.0.0.1", match: true},
{cidr: "10.0.0.0/24", ip: "10.0.1.1", match: false},
{cidr: "1.2.3.4/32", ip: "1.2.3.4", match: true},
{cidr: "1.2.3.4/32", ip: "1.2.3.5", match: false},
{cidr: "0.0.0.0/0", ip: "8.8.8.8", match: true},
// LAPI does not have to send a network address, the host bits are dropped.
{cidr: "10.0.0.5/24", ip: "10.0.0.9", match: true},
{cidr: " 192.168.1.0/24 ", ip: "192.168.1.42", match: true},
{cidr: "2001:db8::/32", ip: "2001:db8::1", match: true},
{cidr: "2001:db8::/32", ip: "2001:db9::1", match: false},
{cidr: "::/0", ip: "2001:db8::1", match: true},
// An IPv4 range and an IPv4 mapped client still have to meet.
{cidr: "10.0.0.0/8", ip: "::ffff:10.1.2.3", match: true},
{cidr: "::ffff:10.0.0.0/104", ip: "10.1.2.3", match: true},
// Families do not mix.
{cidr: "::/0", ip: "8.8.8.8", match: false},
{cidr: "2001:db8::/32", ip: "::ffff:10.0.0.1", match: false},
}
for _, tt := range tests {
t.Run(tt.cidr+"_"+tt.ip, func(t *testing.T) {
key := NormalizeCIDR(tt.cidr)
if key == "" {
t.Fatalf("NormalizeCIDR(%q) returned nothing", tt.cidr)
}
found := false
for _, candidate := range CIDRKeys(tt.ip) {
if candidate == key {
found = true
break
}
}
if found != tt.match {
t.Errorf("key %q of %q found in CIDRKeys(%q) = %v, want %v", key, tt.cidr, tt.ip, found, tt.match)
}
})
}
}
func TestCIDRLookupKeys(t *testing.T) {
tests := []struct {
name string
ip string
prefixLens []int
want []string
}{
{
name: "most specific first",
ip: "10.1.2.3",
prefixLens: []int{8, 32, 16},
want: []string{"10.1.2.3/32", "10.1.0.0/16", "10.0.0.0/8"},
},
{
name: "duplicates are dropped",
ip: "10.1.2.3",
prefixLens: []int{24, 24, 24},
want: []string{"10.1.2.0/24"},
},
{
name: "lengths of the other family are skipped",
ip: "10.1.2.3",
prefixLens: []int{48, 64, 24},
want: []string{"10.1.2.0/24"},
},
{
name: "out of range lengths are skipped",
ip: "10.1.2.3",
prefixLens: []int{-1, 33, 129, 8},
want: []string{"10.0.0.0/8"},
},
{
name: "ipv6 keeps its own lengths",
ip: "2001:db8::1",
prefixLens: []int{32, 64},
want: []string{"2001:db8::/64", "2001:db8::/32"},
},
{
name: "no length gives no key",
ip: "10.1.2.3",
prefixLens: []int{},
want: []string{},
},
{
name: "invalid ip gives no key",
ip: "invalid",
prefixLens: []int{24},
want: nil,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := CIDRLookupKeys(tt.ip, tt.prefixLens)
if len(got) != len(tt.want) {
t.Fatalf("CIDRLookupKeys(%q, %v) = %v, want %v", tt.ip, tt.prefixLens, got, tt.want)
}
for i := range got {
if got[i] != tt.want[i] {
t.Errorf("CIDRLookupKeys(%q, %v)[%d] = %q, want %q", tt.ip, tt.prefixLens, i, got[i], tt.want[i])
}
}
})
}
}
// TestCIDRLookupKeys_SubsetOfCIDRKeys: restricting the lengths only removes candidates,
// it never changes the key of a length that is kept.
func TestCIDRLookupKeys_SubsetOfCIDRKeys(t *testing.T) {
for _, ipStr := range []string{"10.1.2.3", "2001:db8::1", "::ffff:10.1.2.3"} {
t.Run(ipStr, func(t *testing.T) {
all := CIDRKeys(ipStr)
maxBits := len(all) - 1
for bits := 0; bits <= maxBits; bits++ {
got := CIDRLookupKeys(ipStr, []int{bits})
if len(got) != 1 {
t.Fatalf("CIDRLookupKeys(%q, [%d]) returned %d keys", ipStr, bits, len(got))
}
if want := all[maxBits-bits]; got[0] != want {
t.Errorf("CIDRLookupKeys(%q, [%d]) = %q, want %q", ipStr, bits, got[0], want)
}
}
})
}
}
func TestCIDRPrefixLen(t *testing.T) {
tests := []struct {
input string
want int
}{
{"10.0.0.0/8", 8},
{"10.0.0.0/32", 32},
{"0.0.0.0/0", 0},
{"10.0.0.5/24", 24},
{" 10.0.0.0/16 ", 16},
{"2001:db8::/32", 32},
{"2001:db8::/128", 128},
{"::/0", 0},
{"10.0.0.1", -1},
{"invalid", -1},
{"", -1},
}
for _, tt := range tests {
t.Run(tt.input, func(t *testing.T) {
if got := CIDRPrefixLen(tt.input); got != tt.want {
t.Errorf("CIDRPrefixLen(%q) = %d, want %d", tt.input, got, tt.want)
}
})
}
}
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)
}
})
}
}
+9 -17
View File
@@ -78,11 +78,11 @@ ensure_mock() {
# Poll a URL until it returns the expected status code, or fail. # Poll a URL until it returns the expected status code, or fail.
# Usage: wait_for_status URL CODE [TIMEOUT_SECONDS] [curl args...] # Usage: wait_for_status URL CODE [TIMEOUT_SECONDS] [curl args...]
wait_for_status() { wait_for_status() {
local url="$1" expected="$2" timeout="${3:-15}" local url="$1" expected="$2" timeout="${3:-30}"
shift 3 || true shift 3 || true
local elapsed=0 got="" local elapsed=0 got=""
while (( elapsed < timeout )); do while (( elapsed < timeout )); do
got=$(curl -s -m 1 -o /dev/null -w '%{http_code}' "$@" "$url" || true) got=$(curl -s -o /dev/null -w '%{http_code}' "$@" "$url" || true)
if [[ "$got" == "$expected" ]]; then if [[ "$got" == "$expected" ]]; then
return 0 return 0
fi fi
@@ -98,11 +98,11 @@ wait_for_status() {
# code alone can't tell the states apart (e.g. captcha page vs backend, both 200). # code alone can't tell the states apart (e.g. captcha page vs backend, both 200).
# Usage: wait_for_body_contains URL NEEDLE [TIMEOUT_SECONDS] [curl args...] # Usage: wait_for_body_contains URL NEEDLE [TIMEOUT_SECONDS] [curl args...]
wait_for_body_contains() { wait_for_body_contains() {
local url="$1" needle="$2" timeout="${3:-15}" local url="$1" needle="$2" timeout="${3:-30}"
shift 3 || true shift 3 || true
local elapsed=0 body="" local elapsed=0 body=""
while (( elapsed < timeout )); do while (( elapsed < timeout )); do
body=$(curl -s -m 1 "$@" "$url" || true) body=$(curl -s "$@" "$url" || true)
if grep -q "$needle" <<<"$body"; then if grep -q "$needle" <<<"$body"; then
return 0 return 0
fi fi
@@ -119,7 +119,7 @@ assert_status() {
local url="$1" expected="$2" local url="$1" expected="$2"
shift 2 || true shift 2 || true
local got local got
got=$(curl -s --connect-timeout 1 -m 5 -o /dev/null -w '%{http_code}' "$@" "$url") got=$(curl -s -o /dev/null -w '%{http_code}' "$@" "$url")
if [[ "$got" != "$expected" ]]; then if [[ "$got" != "$expected" ]]; then
echo "assert_status: $url expected $expected, got $got" >&2 echo "assert_status: $url expected $expected, got $got" >&2
return 1 return 1
@@ -132,7 +132,7 @@ assert_header() {
local url="$1" header="$2" expected="$3" local url="$1" header="$2" expected="$3"
shift 3 || true shift 3 || true
local got local got
got=$(curl -s --connect-timeout 1 -m 5 -D - -o /dev/null "$@" "$url" | tr -d '\r' \ got=$(curl -s -D - -o /dev/null "$@" "$url" | tr -d '\r' \
| awk -v h="${header,,}" -F': ' 'tolower($1) == h { print $2; exit }') | awk -v h="${header,,}" -F': ' 'tolower($1) == h { print $2; exit }')
if [[ "$got" != "$expected" ]]; then if [[ "$got" != "$expected" ]]; then
echo "assert_header: $url header $header expected \"$expected\", got \"$got\"" >&2 echo "assert_header: $url header $header expected \"$expected\", got \"$got\"" >&2
@@ -146,7 +146,7 @@ assert_body_contains() {
local url="$1" needle="$2" local url="$1" needle="$2"
shift 2 || true shift 2 || true
local body local body
body=$(curl -s --connect-timeout 1 -m 5 "$@" "$url") body=$(curl -s "$@" "$url")
if ! grep -q "$needle" <<<"$body"; then if ! grep -q "$needle" <<<"$body"; then
echo "assert_body_contains: $url expected to contain \"$needle\", got:" >&2 echo "assert_body_contains: $url expected to contain \"$needle\", got:" >&2
echo "$body" >&2 echo "$body" >&2
@@ -158,20 +158,12 @@ assert_body_contains() {
lapi_add_decision() { lapi_add_decision() {
local ip="$1" type="${2:-ban}" duration="${3:-4h}" local ip="$1" type="${2:-ban}" duration="${3:-4h}"
curl -sS --connect-timeout 1 -m 5 -X POST "http://127.0.0.1:${LAPI_PORT}/admin/decisions?ip=${ip}&type=${type}&duration=${duration}" >/dev/null curl -sS -X POST "http://127.0.0.1:${LAPI_PORT}/admin/decisions?ip=${ip}&type=${type}&duration=${duration}" >/dev/null
} }
lapi_delete_decision() { lapi_delete_decision() {
local ip="$1" local ip="$1"
curl -sS --connect-timeout 1 -m 5 -X DELETE "http://127.0.0.1:${LAPI_PORT}/admin/decisions?ip=${ip}" >/dev/null curl -sS -X DELETE "http://127.0.0.1:${LAPI_PORT}/admin/decisions?ip=${ip}" >/dev/null
}
lapi_set_stream_fail() {
curl -sS --connect-timeout 1 -m 5 -X POST "http://127.0.0.1:${LAPI_PORT}/admin/stream-fail" >/dev/null
}
lapi_clear_stream_fail() {
curl -sS --connect-timeout 1 -m 5 -X DELETE "http://127.0.0.1:${LAPI_PORT}/admin/stream-fail" >/dev/null
} }
# --- stack lifecycle --------------------------------------------------------- # --- stack lifecycle ---------------------------------------------------------
-19
View File
@@ -21,7 +21,6 @@ import (
"net/http" "net/http"
"strings" "strings"
"sync" "sync"
"sync/atomic"
) )
// Decision is the subset of a LAPI decision the plugin actually reads. // Decision is the subset of a LAPI decision the plugin actually reads.
@@ -35,9 +34,6 @@ var (
mu sync.Mutex mu sync.Mutex
active = map[string]Decision{} // ip -> decision currently in force active = map[string]Decision{} // ip -> decision currently in force
deleted = map[string]Decision{} // ip -> decision to report in the stream "deleted" list deleted = map[string]Decision{} // ip -> decision to report in the stream "deleted" list
// streamFail makes /v1/decisions/stream return 500 when set, to exercise the
// bouncer's fail-closed behaviour on consecutive stream poll failures.
streamFail atomic.Bool
) )
func writeJSON(w http.ResponseWriter, v any) { func writeJSON(w http.ResponseWriter, v any) {
@@ -177,10 +173,6 @@ func main() {
// "deleted". Re-sending the same on every poll is harmless — the plugin just // "deleted". Re-sending the same on every poll is harmless — the plugin just
// re-adds to / re-deletes from its cache. // re-adds to / re-deletes from its cache.
mux.HandleFunc("/v1/decisions/stream", func(w http.ResponseWriter, _ *http.Request) { mux.HandleFunc("/v1/decisions/stream", func(w http.ResponseWriter, _ *http.Request) {
if streamFail.Load() {
w.WriteHeader(http.StatusInternalServerError)
return
}
mu.Lock() mu.Lock()
defer mu.Unlock() defer mu.Unlock()
writeJSON(w, map[string][]Decision{"new": list(active), "deleted": list(deleted)}) writeJSON(w, map[string][]Decision{"new": list(active), "deleted": list(deleted)})
@@ -191,17 +183,6 @@ func main() {
w.WriteHeader(http.StatusCreated) w.WriteHeader(http.StatusCreated)
}) })
// Test control plane: make the stream endpoint fail (POST) or recover (DELETE).
mux.HandleFunc("/admin/stream-fail", func(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodPost:
streamFail.Store(true)
case http.MethodDelete:
streamFail.Store(false)
}
w.WriteHeader(http.StatusOK)
})
// Test control plane: add / remove decisions instead of cscli. // Test control plane: add / remove decisions instead of cscli.
mux.HandleFunc("/admin/decisions", func(_ http.ResponseWriter, r *http.Request) { mux.HandleFunc("/admin/decisions", func(_ http.ResponseWriter, r *http.Request) {
q := r.URL.Query() q := r.URL.Query()
+3 -11
View File
@@ -11,25 +11,17 @@ body() {
echo "[$SCENARIO] no decision -> request passes (LAPI queried per request)" echo "[$SCENARIO] no decision -> request passes (LAPI queried per request)"
assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 1.2.3.4" assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 1.2.3.4"
echo "[$SCENARIO] adding ban decision for 1.2.3.4 and 2001:db8::1" echo "[$SCENARIO] adding ban decision for 1.2.3.4"
lapi_add_decision 1.2.3.4 ban 5m lapi_add_decision 1.2.3.4 ban 5m
lapi_add_decision "2001:db8::1" ban 5m
echo "[$SCENARIO] IP banned must be blocked (HTTP 403)" echo "[$SCENARIO] none mode has no cache -> next request must be blocked immediately"
assert_status "http://127.0.0.1:${WEB_PORT}/foo" 403 -H "X-Forwarded-For: 1.2.3.4" assert_status "http://127.0.0.1:${WEB_PORT}/foo" 403 -H "X-Forwarded-For: 1.2.3.4"
echo "[$SCENARIO] IPv6 banned must be blocked (HTTP 403)"
assert_status "http://127.0.0.1:${WEB_PORT}/foo" 403 -H "X-Forwarded-For: 2001:db8::1"
echo "[$SCENARIO] deleting decision" echo "[$SCENARIO] deleting decision"
lapi_delete_decision 1.2.3.4 lapi_delete_decision 1.2.3.4
lapi_delete_decision "2001:db8::1"
echo "[$SCENARIO] previously banned IP must pass again" echo "[$SCENARIO] previously banned IP must pass again immediately"
assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 1.2.3.4" assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 1.2.3.4"
echo "[$SCENARIO] previously banned IPv6 must pass again"
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 200 15 -H "X-Forwarded-For: 2001:db8::1"
} }
run_scenario "$SCENARIO" "$HERE" body run_scenario "$SCENARIO" "$HERE" body
@@ -18,8 +18,7 @@ http:
bouncer: bouncer:
enabled: "true" enabled: "true"
crowdsecMode: stream crowdsecMode: stream
updateIntervalSeconds: "1" updateIntervalSeconds: "2"
updateMaxFailure: "2"
crowdsecLapiScheme: http crowdsecLapiScheme: http
crowdsecLapiHost: "@@LAPI_HOST@@" crowdsecLapiHost: "@@LAPI_HOST@@"
crowdsecLapiKey: "@@APIKEY@@" crowdsecLapiKey: "@@APIKEY@@"
+2 -22
View File
@@ -11,9 +11,8 @@ body() {
echo "[$SCENARIO] no decision yet -> request allowed" echo "[$SCENARIO] no decision yet -> request allowed"
assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 1.2.3.4" assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 1.2.3.4"
echo "[$SCENARIO] adding ban decision for 1.2.3.4 and 10.0.0.0/8" echo "[$SCENARIO] adding ban decision for 1.2.3.4"
lapi_add_decision 1.2.3.4 ban 5m lapi_add_decision 1.2.3.4 ban 5m
lapi_add_decision 10.0.0.0/24 ban 5m
echo "[$SCENARIO] banned IP must be blocked once the next stream poll lands (HTTP 403)" echo "[$SCENARIO] banned IP must be blocked once the next stream poll lands (HTTP 403)"
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 403 15 -H "X-Forwarded-For: 1.2.3.4" wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 403 15 -H "X-Forwarded-For: 1.2.3.4"
@@ -21,30 +20,11 @@ body() {
echo "[$SCENARIO] non-banned IP must still pass (HTTP 200)" echo "[$SCENARIO] non-banned IP must still pass (HTTP 200)"
assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 5.6.7.8" assert_status "http://127.0.0.1:${WEB_PORT}/foo" 200 -H "X-Forwarded-For: 5.6.7.8"
echo "[$SCENARIO] banned IP in CIDR must be blocked once polled (HTTP 403)" echo "[$SCENARIO] deleting ban decision"
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 403 15 -H "X-Forwarded-For: 10.0.0.1"
echo "[$SCENARIO] deleting ban decision for 1.2.3.4 and 10.0.0.0/8"
lapi_delete_decision 1.2.3.4 lapi_delete_decision 1.2.3.4
lapi_delete_decision 10.0.0.0/24
echo "[$SCENARIO] previously banned IP must pass again once the deletion is polled" echo "[$SCENARIO] previously banned IP must pass again once the deletion is polled"
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 200 15 -H "X-Forwarded-For: 1.2.3.4" wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 200 15 -H "X-Forwarded-For: 1.2.3.4"
echo "[$SCENARIO] previously CIDR-banned IP must pass again once deletion is polled"
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 200 15 -H "X-Forwarded-For: 10.0.0.1"
echo "[$SCENARIO] making the stream endpoint fail -> bouncer must pass for one more cycle (updateMaxFailure: 2)"
lapi_set_stream_fail
sleep 2 # update cache is every 1 seconds then waiting for minimum 1 cycle
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 200 15 -H "X-Forwarded-For: 8.8.8.8"
echo "[$SCENARIO] bouncer must block everything (isStreamHealthy: false)"
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 403 15 -H "X-Forwarded-For: 8.8.8.8"
echo "[$SCENARIO] restoring the stream endpoint -> bouncer must recover and pass again"
lapi_clear_stream_fail
wait_for_status "http://127.0.0.1:${WEB_PORT}/foo" 200 15 -H "X-Forwarded-For: 8.8.8.8"
} }
run_scenario "$SCENARIO" "$HERE" body run_scenario "$SCENARIO" "$HERE" body