Lower limits for listed AS numbers and countries (closes #21)
check / check (push) Waiting to run

SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT give the clients
of the AS numbers and countries they list that percentage of every rate
and byte limit, rounded down; SWWAF_ASN_BYTES_PERCENT and
SWWAF_COUNTRY_BYTES_PERCENT take its place for the byte limits of those
they list; SWWAF_UNKNOWN_LIMIT_PERCENT (100) covers clients without a
country. The lowest applies. While one lowers a limit, a request waits
for its client's lookup, and SWWAF_LOOKUP_SOURCE=off stops the start. Log
lines give limit_percent and bytes_percent with their settings; ban
notes, and so alerts, give the broken limit's.

Judgement call: a client without a country is unknown, whatever its AS number.
Judgement call: bytes_percent and its setting are log fields SPEC does not name.
Rule suppressed: funlen on FromEnvironment, one line per setting.

Model: opus-5-5
This commit is contained in:
2026-10-07 12:11:02 +00:00
parent f35e3ddfe8
commit abf3b01ba9
17 changed files with 1102 additions and 109 deletions
+6
View File
@@ -108,6 +108,12 @@ type Notes struct {
Limit int64 `json:"limit,omitempty"`
Window string `json:"window,omitempty"`
Count float64 `json:"count,omitempty"`
// LimitPercent and LimitPercentSetting are, for a ban for a limit a
// biased threshold lowered, the client's percentage of that kind of
// limit, of which Limit is the result, and the setting that gave it.
// Both are left out for a limit that was not lowered.
LimitPercent *int64 `json:"limit_percent,omitempty"`
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
// RuleID and Target are, for a ban for a clear sign of attack, the id
// of the rule file rule that matched, and its target.
RuleID string `json:"rule_id,omitempty"`
+133 -5
View File
@@ -120,6 +120,22 @@ type Config struct {
// capitals, as GeoJS gives them.
DeniedCountries []string
ExclusivelyAllowedCountries []string
// The biased thresholds. ASNLimitPercent and CountryLimitPercent give
// the clients of the AS numbers and the countries they list that
// percentage of every rate limit and byte limit
// (SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT).
// ASNBytesPercent and CountryBytesPercent give those they list a
// percentage of the byte limits in place of that one
// (SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT). Each holds
// percentages from 0 to 100, by AS number, written as AS64496, or by
// country, a two-letter code in capitals, as the lookup gives them.
// UnknownLimitPercent is the percentage of every limit a client without
// a country gets (SWWAF_UNKNOWN_LIMIT_PERCENT).
ASNLimitPercent map[string]int64
CountryLimitPercent map[string]int64
ASNBytesPercent map[string]int64
CountryBytesPercent map[string]int64
UnknownLimitPercent int64
// BanResponse is the status a refused client is answered with, 403
// or 429, or 0 to close the connection without an answer
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
@@ -300,6 +316,11 @@ var (
"source_failure or file_error")
errNotNumberOrOff = errors.New("is not a whole number above zero, such as 60, or off")
errNotUTF8 = errors.New("is not valid UTF-8")
errNotASN = errors.New("is not an AS number such as AS64496")
errNotPercentItem = errors.New(
"is not a code, : and a percentage, such as AS64496:50 or cn:25")
errNotPercent = errors.New("is not a percentage, a whole number from 0 to 100")
errListedTwice = errors.New("is listed twice")
)
// FromEnvironment reads the settings with lookupEnv, normally
@@ -307,6 +328,8 @@ var (
// named by the setting's name with _FILE added names the file, which is
// read now (see lookup). A setting that is not set takes its default. A
// setting that is set but invalid is an error that names it.
//
//nolint:funlen // one line for each setting, a list that grows with them
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv}
cfg := &Config{
@@ -342,6 +365,11 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
ASNLimitPercent: env.percents("SWWAF_ASN_LIMIT_PERCENT", parseASN),
CountryLimitPercent: env.percents("SWWAF_COUNTRY_LIMIT_PERCENT", parseCountry),
ASNBytesPercent: env.percents("SWWAF_ASN_BYTES_PERCENT", parseASN),
CountryBytesPercent: env.percents("SWWAF_COUNTRY_BYTES_PERCENT", parseCountry),
UnknownLimitPercent: env.percent("SWWAF_UNKNOWN_LIMIT_PERCENT", "100"),
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
@@ -592,6 +620,25 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries
}
// percents reads a setting that is a list of AS numbers or countries,
// which parseCode reads, each with a percentage. It is empty by default.
func (e *environment) percents(
name string, parseCode func(string) (string, error),
) map[string]int64 {
percents, err := parsePercents(e.value(name, ""), parseCode)
e.check(name, err)
return percents
}
// percent reads a setting that is a percentage, from 0 to 100.
func (e *environment) percent(name, defaultValue string) int64 {
percent, err := parsePercent(e.value(name, defaultValue))
e.check(name, err)
return percent
}
// lookupSource reads the setting that is where clients are looked up:
// geojs, file, or off.
func (e *environment) lookupSource(name, defaultValue string) string {
@@ -619,7 +666,9 @@ func (e *environment) checkLookupDBPath(cfg *Config) {
// checkCountriesAndLookups refuses a country on both country lists, and,
// while SWWAF_LOOKUP_SOURCE is off, each setting that needs clients looked
// up: the country lists and SWWAF_ADD_LOOKUP_HEADERS.
// up: the country lists, SWWAF_ADD_LOOKUP_HEADERS, and the biased
// thresholds, of which SWWAF_UNKNOWN_LIMIT_PERCENT needs them only below
// 100, where it lowers a limit.
func (e *environment) checkCountriesAndLookups(cfg *Config) {
for _, country := range cfg.ExclusivelyAllowedCountries {
if slices.Contains(cfg.DeniedCountries, country) {
@@ -639,6 +688,11 @@ func (e *environment) checkCountriesAndLookups(cfg *Config) {
{"SWWAF_DENIED_COUNTRIES", len(cfg.DeniedCountries) > 0},
{"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", len(cfg.ExclusivelyAllowedCountries) > 0},
{"SWWAF_ADD_LOOKUP_HEADERS", cfg.AddLookupHeaders},
{"SWWAF_ASN_LIMIT_PERCENT", len(cfg.ASNLimitPercent) > 0},
{"SWWAF_COUNTRY_LIMIT_PERCENT", len(cfg.CountryLimitPercent) > 0},
{"SWWAF_ASN_BYTES_PERCENT", len(cfg.ASNBytesPercent) > 0},
{"SWWAF_COUNTRY_BYTES_PERCENT", len(cfg.CountryBytesPercent) > 0},
{"SWWAF_UNKNOWN_LIMIT_PERCENT", cfg.UnknownLimitPercent < 100},
} {
if setting.set {
e.check(setting.name, fmt.Errorf("is set while SWWAF_LOOKUP_SOURCE is off; %w",
@@ -1148,13 +1202,12 @@ func parseCountries(value string) ([]string, error) {
return nil, err
}
known := strings.Fields(countryCodes)
countries := make([]string, 0, len(items))
for _, item := range items {
country := strings.ToUpper(item)
if !slices.Contains(known, country) {
return nil, fmt.Errorf("%q %w", item, errNotCountry)
country, err := parseCountry(item)
if err != nil {
return nil, err
}
countries = append(countries, country)
@@ -1163,6 +1216,81 @@ func parseCountries(value string) ([]string, error) {
return countries, nil
}
// parseCountry reads a country code in either case, and returns it in
// capitals.
func parseCountry(value string) (string, error) {
country := strings.ToUpper(value)
if !slices.Contains(strings.Fields(countryCodes), country) {
return "", fmt.Errorf("%q %w", value, errNotCountry)
}
return country, nil
}
// parseASN reads an AS number such as AS64496, in either case, and
// returns it as the lookup gives it: AS and the number, in capitals and
// without leading zeros.
func parseASN(value string) (string, error) {
digits, hasAS := strings.CutPrefix(strings.ToUpper(value), "AS")
number, err := strconv.ParseUint(digits, 10, 32)
if !hasAS || err != nil {
return "", fmt.Errorf("%q %w", value, errNotASN)
}
return "AS" + strconv.FormatUint(number, 10), nil
}
// parsePercents reads a comma-separated list of items, each an AS number
// or a country, which parseCode reads, then : and a percentage, such as
// AS64496:50 or cn:25, and returns each one's percentage. An empty value
// is an empty list. An AS number or country listed twice is an error.
func parsePercents(
value string, parseCode func(string) (string, error),
) (map[string]int64, error) {
items, err := parseList(value)
if err != nil {
return nil, err
}
percents := make(map[string]int64, len(items))
for _, item := range items {
codeText, percentText, found := strings.Cut(item, ":")
if !found {
return nil, fmt.Errorf("%q %w", item, errNotPercentItem)
}
code, err := parseCode(codeText)
if err != nil {
return nil, err
}
percent, err := parsePercent(percentText)
if err != nil {
return nil, err
}
if _, listed := percents[code]; listed {
return nil, fmt.Errorf("%q %w", codeText, errListedTwice)
}
percents[code] = percent
}
return percents, nil
}
// parsePercent reads a percentage, a whole number from 0 to 100.
func parsePercent(value string) (int64, error) {
percent, err := strconv.ParseInt(value, 10, 64)
if err != nil || percent < 0 || percent > 100 {
return 0, fmt.Errorf("%q %w", value, errNotPercent)
}
return percent, nil
}
// headerNameChars are the characters RFC 9110 allows in a header name:
// letters, digits and these marks.
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
+109 -4
View File
@@ -50,6 +50,11 @@ const (
addLookupHeaders = "SWWAF_ADD_LOOKUP_HEADERS"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
@@ -892,9 +897,14 @@ func TestSettingNeedingLookupsStopsTheStartWhileTheyAreOff(t *testing.T) {
t.Parallel()
for name, value := range map[string]string{
deniedCountries: "kp",
allowedCountries: "de",
addLookupHeaders: enabled,
deniedCountries: "kp",
allowedCountries: "de",
addLookupHeaders: enabled,
asnLimitPercent: "AS64496:50",
countryLimitPercent: "cn:25",
asnBytesPercent: "AS64496:50",
countryBytesPercent: "cn:25",
unknownLimitPercent: "99",
} {
t.Run(name, func(t *testing.T) {
t.Parallel()
@@ -909,12 +919,94 @@ func TestSettingNeedingLookupsStopsTheStartWhileTheyAreOff(t *testing.T) {
})
}
// Set empty, the country lists need nothing looked up.
// Set empty, the lists need nothing looked up, and nor does
// SWWAF_UNKNOWN_LIMIT_PERCENT at 100, which lowers no limit.
fromEnvironment(t, environment{
lookupSource: off, deniedCountries: "", allowedCountries: "",
asnLimitPercent: "", countryLimitPercent: "", asnBytesPercent: "",
countryBytesPercent: "", unknownLimitPercent: "100",
})
}
func TestBiasedThresholdsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
if len(cfg.ASNLimitPercent) != 0 || len(cfg.CountryLimitPercent) != 0 ||
len(cfg.ASNBytesPercent) != 0 || len(cfg.CountryBytesPercent) != 0 ||
cfg.UnknownLimitPercent != 100 {
t.Errorf("biased thresholds %v, %v, %v, %v and %d by default, "+
"want four empty lists and 100", cfg.ASNLimitPercent, cfg.CountryLimitPercent,
cfg.ASNBytesPercent, cfg.CountryBytesPercent, cfg.UnknownLimitPercent)
}
// AS numbers and countries in either case, an AS number with leading
// zeros, 0 and 100.
cfg = fromEnvironment(t, environment{
asnLimitPercent: "AS14061:50, as16276:0,AS045102:100",
countryLimitPercent: "cn:25,RU:50",
asnBytesPercent: "as16276:75",
countryBytesPercent: "ru:10",
unknownLimitPercent: "0",
})
for name, tc := range map[string]struct{ got, want map[string]int64 }{
asnLimitPercent: {
cfg.ASNLimitPercent,
map[string]int64{"AS14061": 50, "AS16276": 0, "AS45102": 100},
},
countryLimitPercent: {cfg.CountryLimitPercent, map[string]int64{"CN": 25, "RU": 50}},
asnBytesPercent: {cfg.ASNBytesPercent, map[string]int64{"AS16276": 75}},
countryBytesPercent: {cfg.CountryBytesPercent, map[string]int64{"RU": 10}},
} {
if !maps.Equal(tc.got, tc.want) {
t.Errorf("%s gave %v, want %v", name, tc.got, tc.want)
}
}
if cfg.UnknownLimitPercent != 0 {
t.Errorf("%s gave %d, want 0", unknownLimitPercent, cfg.UnknownLimitPercent)
}
}
func TestInvalidBiasedThresholdStopsTheStartSayingWhatIsWrong(t *testing.T) {
t.Parallel()
const (
notASN = " is not an AS number such as AS64496"
notItem = " is not a code, : and a percentage, such as AS64496:50 or cn:25"
notPercent = " is not a percentage, a whole number from 0 to 100"
)
for _, tc := range []struct{ name, value, want string }{
{asnLimitPercent, "14061:50", `"14061"` + notASN},
{asnLimitPercent, "AS4294967296:50", `"AS4294967296"` + notASN},
{asnLimitPercent, "AS14061", `"AS14061"` + notItem},
{asnLimitPercent, "AS14061:101", `"101"` + notPercent},
{asnLimitPercent, "AS14061:50,as14061:25", `"as14061" is listed twice`},
{
countryLimitPercent, "nk:25",
`"nk" is not a two-letter country code such as de or kp`,
},
{countryLimitPercent, "cn:25,CN:50", `"CN" is listed twice`},
{asnBytesPercent, "AS14061:-1", `"-1"` + notPercent},
{countryBytesPercent, "cn:50%", `"50%"` + notPercent},
{unknownLimitPercent, "101", `"101"` + notPercent},
{unknownLimitPercent, off, `"off"` + notPercent},
} {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
want := tc.name + ": " + tc.want
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
func TestSizesAndOff(t *testing.T) {
t.Parallel()
@@ -1054,6 +1146,14 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowedCountries, "uk"},
{allowedCountries, "zz"},
{allowedCountries, "de,germany"},
{asnLimitPercent, "AS14061:50,,AS16276:50"}, {asnLimitPercent, "ASX:50"},
{asnLimitPercent, "AS14061:"}, {asnLimitPercent, "AS14061 :50"},
{asnLimitPercent, "AS14061:1.5"}, {asnLimitPercent, "AS-1:50"},
{countryLimitPercent, "cn"}, {countryLimitPercent, "cn:"},
{countryLimitPercent, "cn:25:50"}, {countryLimitPercent, "china:25"},
{asnBytesPercent, "AS14061:101"}, {countryBytesPercent, "su:50"},
{unknownLimitPercent, ""}, {unknownLimitPercent, "-1"},
{unknownLimitPercent, "50%"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
@@ -1325,6 +1425,11 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
addLookupHeaders: "false",
deniedCountries: "",
allowedCountries: "",
asnLimitPercent: "",
countryLimitPercent: "",
asnBytesPercent: "",
countryBytesPercent: "",
unknownLimitPercent: "100",
banResponse: "403",
limitBanDuration: "1h",
limitBanRepeatWindow: "24h",
+20 -8
View File
@@ -42,9 +42,11 @@ func (rq *request) banned(now time.Time) bool {
// limitBroken counts the request for the rate limits at now, notes the
// client's counts for the log line, and reports whether the request takes
// the client over a rate limit, which breaks it.
// the client over a rate limit, as its limit percentage lowers it, which
// breaks it.
func (rq *request) limitBroken(now time.Time) bool {
counts, hit, over := rq.h.limiter.Count(clientGroup(rq.client), now)
counts, hit, over := rq.h.limiter.Count(clientGroup(rq.client), now,
rq.limitPercent.percent)
rq.line.Counts = counts
if over {
@@ -63,8 +65,9 @@ func (rq *request) limitBroken(now time.Time) bool {
// what it carried from the client with the request's. Only a request
// passed to the app has them counted, and only one the rate limits
// counted; in observe mode, not one that enforce mode would have refused.
// Bytes that take the client over a byte limit break it; the response was
// passed on whole.
// Bytes that take the client over a byte limit, as its limit percentage
// for the byte limits lowers it, break it; the response was passed on
// whole.
func (rq *request) countBytes() {
if !rq.counted || rq.line.WouldAction != "" {
return
@@ -89,7 +92,8 @@ func (rq *request) countBytes() {
now := rq.h.now()
counts, hit, over := rq.h.limiter.CountBytes(clientGroup(rq.client), now, bytes)
counts, hit, over := rq.h.limiter.CountBytes(clientGroup(rq.client), now, bytes,
rq.bytesPercent.percent)
rq.line.Counts.MinuteBytes = counts.MinuteBytes
rq.line.Counts.HourBytes = counts.HourBytes
rq.line.Counts.DayBytes = counts.DayBytes
@@ -103,9 +107,10 @@ func (rq *request) countBytes() {
// one hit names, and notes the offence for the log line. status is what
// the client was sent, or is sent: SWWAF_BAN_RESPONSE for a request over
// a rate limit, the app's answer for one whose bytes broke a byte limit.
// The ban sets the client's counters back to zero. In observe mode it
// makes no ban and sets nothing back, and raises the alert for the ban it
// would have made, if that alert would be sent.
// The ban's notes give the client's limit percentage for that kind of
// limit. The ban sets the client's counters back to zero. In observe mode
// it makes no ban and sets nothing back, and raises the alert for the ban
// it would have made, if that alert would be sent.
func (rq *request) banForLimit(now time.Time, hit ratelimit.Hit, status int) {
rq.line.LimitHit = hit.Window
if hit.Kind == ratelimit.KindBytes {
@@ -131,6 +136,13 @@ func (rq *request) banForLimit(now time.Time, hit ratelimit.Hit, status int) {
Requests: rq.netblockRequests(netblock),
}
percent := rq.limitPercent
if hit.Kind == ratelimit.KindBytes {
percent = rq.bytesPercent
}
notes.LimitPercent, notes.LimitPercentSetting = percent.logged()
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
if wouldBan {
+96
View File
@@ -0,0 +1,96 @@
package proxy
import (
"sneak.berlin/go/smallwebwaf/internal/config"
)
// whole is the percentage of each limit a client gets when no biased
// threshold lowers its limits.
const whole = 100
// percentage is a client's limit percentage for the rate limits or for
// the byte limits, as the biased thresholds give it, and the setting that
// gave it: "" with whole when none lowers that kind of limit.
type percentage struct {
percent int64
setting string
}
// biasedThresholdsSet reports whether a biased threshold can lower a
// client's limits: one of its lists is not empty, or
// SWWAF_UNKNOWN_LIMIT_PERCENT is below 100. The client's lookup is then
// needed before its request goes on.
func biasedThresholdsSet(cfg *config.Config) bool {
return len(cfg.ASNLimitPercent) > 0 || len(cfg.CountryLimitPercent) > 0 ||
len(cfg.ASNBytesPercent) > 0 || len(cfg.CountryBytesPercent) > 0 ||
cfg.UnknownLimitPercent < whole
}
// limitPercentages returns a client's limit percentages, for the rate
// limits and for the byte limits, by its AS number and country as looked
// up, each "" when unknown. Each is the lowest of those the settings give
// it, the first of them in the order below when several are lowest: the
// percentage SWWAF_ASN_LIMIT_PERCENT gives its AS number, the one
// SWWAF_COUNTRY_LIMIT_PERCENT gives its country, and, for a client
// without a country, SWWAF_UNKNOWN_LIMIT_PERCENT. For the byte limits,
// SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT take the place
// of the first two for an AS number or a country they list.
func limitPercentages(
cfg *config.Config, asn, country string,
) (percentage, percentage) {
unknown := percentage{percent: whole}
if country == "" {
unknown = percentage{cfg.UnknownLimitPercent, "SWWAF_UNKNOWN_LIMIT_PERCENT"}
}
asnRequests := given(cfg.ASNLimitPercent, asn, "SWWAF_ASN_LIMIT_PERCENT")
countryRequests := given(cfg.CountryLimitPercent, country,
"SWWAF_COUNTRY_LIMIT_PERCENT")
asnBytes, countryBytes := asnRequests, countryRequests
if _, listed := cfg.ASNBytesPercent[asn]; listed {
asnBytes = given(cfg.ASNBytesPercent, asn, "SWWAF_ASN_BYTES_PERCENT")
}
if _, listed := cfg.CountryBytesPercent[country]; listed {
countryBytes = given(cfg.CountryBytesPercent, country, "SWWAF_COUNTRY_BYTES_PERCENT")
}
return lowest(asnRequests, countryRequests, unknown),
lowest(asnBytes, countryBytes, unknown)
}
// given returns the percentage percents, the setting named setting, gives
// code, an AS number or a country, or whole when it does not list code.
func given(percents map[string]int64, code, setting string) percentage {
percent, listed := percents[code]
if !listed {
return percentage{percent: whole}
}
return percentage{percent, setting}
}
// lowest returns the lowest of percentages below whole, the first of them
// when several are lowest, or whole when none is below it.
func lowest(percentages ...percentage) percentage {
low := percentage{percent: whole}
for _, p := range percentages {
if p.percent < low.percent {
low = p
}
}
return low
}
// logged returns p as the log line and the notes of a ban give it: its
// percent and setting, or nil and "" for whole, which they leave out.
func (p percentage) logged() (*int64, string) {
if p.percent == whole {
return nil, ""
}
return &p.percent, p.setting
}
+494
View File
@@ -0,0 +1,494 @@
package proxy_test
import (
"fmt"
"io"
"maps"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"path/filepath"
"strings"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The biased thresholds.
const (
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
)
const (
// asnDEHalf and countryDEHalf give fromDE's AS number and its country
// half of every limit, and asnDEQuarter gives its AS number a quarter.
asnDEHalf = asnDE + ":50"
asnDEQuarter = asnDE + ":25"
countryDEHalf = "de:50"
// noCountry is in an AS of its own, AS64500, and in no country.
noCountry = "192.0.2.80"
// fourAMinute is the rate limit these tests set: half of it is 2
// requests a minute, a quarter of it 1.
fourAMinute = "4"
// twoUploads is the byte limit these tests set: 199 bytes, which an
// upload, a request with a body and its answer, 100 bytes, is within,
// and half of which, 99 bytes, it is over.
twoUploads = "199"
// none is how percentText gives a percentage left out.
none = "none"
)
func TestEachBiasedThresholdLowersTheRateLimits(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
setting, value, from string
}{
{asnLimitPercent, asnDEHalf, fromDE},
{countryLimitPercent, countryDEHalf, fromDE},
{unknownLimitPercent, "50", unplaced},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, tc.setting: tc.value,
})
// Half of 4 requests a minute: the third breaks the limit.
for _, sent := range []struct {
status int
action string
}{
{http.StatusOK, requestlog.ActionForward},
{http.StatusOK, requestlog.ActionForward},
{http.StatusForbidden, requestlog.ActionRateLimited},
} {
line := s.get(tc.from, sent.status, sent.action)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
"50 from "+tc.setting)
}
// fromKP, which no setting lists, has the whole limit.
for range 3 {
line := s.get(fromKP, http.StatusOK, requestlog.ActionForward)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
none)
}
})
}
}
func TestEachBiasedThresholdLowersTheByteLimits(t *testing.T) {
t.Parallel()
// The AS numbers and countries are given in either case.
for _, tc := range []struct {
setting, value, from string
}{
{asnLimitPercent, asnDEHalf, fromDE},
{countryLimitPercent, "DE:50", fromDE},
{unknownLimitPercent, "50", unplaced},
{asnBytesPercent, "as64496:50", fromDE},
{countryBytesPercent, countryDEHalf, fromDE},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
bytesLimitPerMinute: twoUploads, tc.setting: tc.value,
})
// The upload's 100 bytes are over half of 199, 99.
line := s.uploadFrom(tc.from)
if line.LimitHit != minuteBytes {
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
}
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
"50 from "+tc.setting)
// fromKP, which no setting lists, has the whole limit.
line = s.uploadFrom(fromKP)
if line.LimitHit != "" {
t.Errorf("log line for %s has limit_hit %q, want none", fromKP, line.LimitHit)
}
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting, none)
})
}
}
func TestBytesPercentSettingsTakeThePlaceOfTheOthersForByteLimits(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// limitPercent and bytesPercent are the log line's, as percentText
// gives them, and limitHit is its limit_hit.
limitPercent, bytesPercent, limitHit string
}{
{
"lowering the byte limits alone",
map[string]string{asnBytesPercent: asnDEHalf},
none, "50 from " + asnBytesPercent, minuteBytes,
},
{
"lowering the byte limits alone, by country",
map[string]string{countryBytesPercent: countryDEHalf},
none, "50 from " + countryBytesPercent, minuteBytes,
},
{
"raising the byte limits back",
map[string]string{asnLimitPercent: asnDEHalf, asnBytesPercent: asnDE + ":100"},
"50 from " + asnLimitPercent, none, "",
},
{
"raising the byte limits back, by country",
map[string]string{countryLimitPercent: countryDEHalf, countryBytesPercent: "de:100"},
"50 from " + countryLimitPercent, none, "",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{bytesLimitPerMinute: twoUploads}
maps.Copy(env, tc.env)
s, _, _ := startWithLookups(t, env)
// The upload's 100 bytes are over 99, half of 199, and within 199.
line := s.uploadFrom(fromDE)
if line.LimitHit != tc.limitHit {
t.Errorf("log line has limit_hit %q, want %q", line.LimitHit, tc.limitHit)
}
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.limitPercent)
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
tc.bytesPercent)
})
}
}
func TestZeroPercentIsAZeroAllowance(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{asnLimitPercent: asnDE + ":0"})
// The first request breaks the limit, and bans the client; the log line
// gives the 0.
line := s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
if line.fields["limit_percent"] != float64(0) ||
line.fields["limit_percent_setting"] != asnLimitPercent {
t.Errorf("log line has limit_percent %v from %v, want 0 from %s",
line.fields["limit_percent"], line.fields["limit_percent_setting"],
asnLimitPercent)
}
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
}
func TestLowestPercentageApplies(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
from string
// want is the log line's limit_percent, as percentText gives it.
want string
}{
{
"the country's",
map[string]string{asnLimitPercent: asnDEHalf, countryLimitPercent: "de:25"},
fromDE, "25 from " + countryLimitPercent,
},
{
"the AS number's",
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: countryDEHalf},
fromDE, "25 from " + asnLimitPercent,
},
{
"the AS number's, the first of two alike",
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: "de:25"},
fromDE, "25 from " + asnLimitPercent,
},
{
"that for a client without a country",
map[string]string{asnLimitPercent: "AS64500:50", unknownLimitPercent: "25"},
noCountry, "25 from " + unknownLimitPercent,
},
{
// SWWAF_UNKNOWN_LIMIT_PERCENT is left at its default, 100.
"the AS number's, for a client without a country",
map[string]string{asnLimitPercent: "AS64500:25"},
noCountry, "25 from " + asnLimitPercent,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{rateLimitPerMinute: fourAMinute}
maps.Copy(env, tc.env)
s, _, _ := startWithLookups(t, env)
// A quarter of 4 requests a minute: the second breaks the limit.
s.get(tc.from, http.StatusOK, requestlog.ActionForward)
line := s.get(tc.from, http.StatusForbidden, requestlog.ActionRateLimited)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.want)
})
}
}
func TestUnknownLimitPercentGivesEveryClientWithoutACountryItsPercentage(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, unknownLimitPercent: "50",
})
// One the lookup database does not hold, and one on a private address,
// which is never looked up: the third request of each breaks half of 4.
for _, from := range []string{unplaced, "10.0.0.8"} {
s.get(from, http.StatusOK, requestlog.ActionForward)
s.get(from, http.StatusOK, requestlog.ActionForward)
s.get(from, http.StatusForbidden, requestlog.ActionRateLimited)
}
// One in a country has the whole limit.
for range 3 {
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
}
}
func TestClientWithoutAnAnswerInTimeHasTheUnknownLimitPercent(t *testing.T) {
t.Parallel()
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
// says, with a GeoJS that never answers.
synctest.Test(t, func(t *testing.T) {
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{unknownLimitPercent: "0"})
// Once the second the request waits for its answer is up, the client
// counts as without a country, and its zero allowance refuses the
// request before it reaches the app.
serveFromDE(t, server, http.MethodGet, http.NoBody)
line := out.requestLine(t)
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
"0 from "+unknownLimitPercent)
})
}
func TestRequestWaitsForItsLookupWhileABiasedThresholdIsSet(t *testing.T) {
t.Parallel()
const timeout = 3 * time.Second
for _, tc := range []struct {
setting, value string
waits bool
}{
{asnLimitPercent, asnDEHalf, true},
{countryLimitPercent, countryDEHalf, true},
{asnBytesPercent, asnDEHalf, true},
{countryBytesPercent, countryDEHalf, true},
{unknownLimitPercent, "99", true},
// At 100, its default, it lowers no limit.
{unknownLimitPercent, "100", false},
} {
t.Run(tc.setting+"="+tc.value, func(t *testing.T) {
t.Parallel()
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
// says, with a GeoJS that never answers.
synctest.Test(t, func(t *testing.T) {
// The request's body is over SWWAF_REQUEST_MAX_BYTES, so that it
// is refused after the checks, and never reaches the app.
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{
lookupTimeout: timeout.String(), requestMaxBytes: "1",
tc.setting: tc.value,
})
began := time.Now()
serveFromDE(t, server, http.MethodPost, strings.NewReader("ab"))
want := time.Duration(0)
if tc.waits {
want = timeout
}
if waited := time.Since(began); waited != want {
t.Errorf("the request waited %s for its answer, want %s", waited, want)
}
wantLine(t, out.requestLine(t), http.StatusRequestEntityTooLarge,
requestlog.ActionTooLarge)
// The bubble's clock stops once this function returns, so the
// request to GeoJS, which a request that did not wait leaves
// under way, has to be abandoned before then.
time.Sleep(timeout)
})
})
}
}
func TestBanForALoweredLimitGivesThePercentageInItsNotesAndItsAlert(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// before is how many uploads come before the one that breaks a
// limit, which is answered with status and logged with action.
before int
status int
action string
// reason and want are the ban's reason, and its notes' limit
// percentage, as percentText gives it.
reason, want string
}{
{
// A quarter of 12 requests a minute is 3: the fourth breaks it.
"a rate limit",
map[string]string{rateLimitPerMinute: "12", asnLimitPercent: asnDEQuarter},
3, http.StatusForbidden, requestlog.ActionRateLimited,
"requests per minute over the limit of 3", "25 from " + asnLimitPercent,
},
{
// The byte limits' percentage, not the rate limits'.
"a byte limit",
map[string]string{
bytesLimitPerMinute: twoUploads, asnLimitPercent: asnDEQuarter,
asnBytesPercent: asnDEHalf,
},
0, http.StatusOK, requestlog.ActionForward,
"bytes per minute over the limit of 99", "50 from " + asnBytesPercent,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, tc.env)
for range tc.before {
s.uploadFrom(fromDE)
}
s.requestWithBody(http.MethodPost, fromDE, "/", uploadHeader, uploadBody,
tc.status, tc.action)
held := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))
if len(held) != 1 {
t.Fatalf("bans %+v, want one", held)
}
notes := held[0].Notes
if held[0].Reason != tc.reason {
t.Errorf("the ban's reason is %q, want %q", held[0].Reason, tc.reason)
}
wantPercent(t, "the notes' limit_percent", notes.LimitPercent,
notes.LimitPercentSetting, tc.want)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("%d alerts wait, want the ban's alone: %+v", len(waiting), waiting)
}
alerted, _ := waiting[0].Detail["notes"].(bans.Notes)
wantPercent(t, "the alert's notes' limit_percent", alerted.LimitPercent,
alerted.LimitPercentSetting, tc.want)
})
}
}
// startWithLookups is startAppWithAlerts in front of readAndAnswer, with
// the settings in env on top of clients looked up in a lookup database,
// which places fromDE and fromKP in the AS numbers and countries the
// stand-in for GeoJS gives them, noCountry in AS64500 and no country, and
// no other address. It returns the sender, the server and the queue of
// the alerts.
func startWithLookups(
t *testing.T, env map[string]string,
) (*sender, *proxy.Server, *alerts.Queue) {
t.Helper()
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.Write(t, path, map[string]lookuptest.Network{
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
noCountry + "/32": {ASN: "AS64500", ASName: "Nowhere Net"},
})
settings := map[string]string{lookupSource: fileSource, lookupDBPath: path}
maps.Copy(settings, env)
s, _, server, queue := startAppWithAlerts(t, readAndAnswer, settings)
return s, server, queue
}
// uploadFrom is upload from the client at from.
func (s *sender) uploadFrom(from string) logLine {
s.t.Helper()
line, _ := s.requestWithBody(http.MethodPost, from, "/", uploadHeader, uploadBody,
http.StatusOK, requestlog.ActionForward)
return line
}
// serveFromDE hands a request from fromDE with method and body straight to
// server's handler, without the network, and returns once it is answered.
func serveFromDE(t *testing.T, server *proxy.Server, method string, body io.Reader) {
t.Helper()
req := httptest.NewRequestWithContext(t.Context(), method, "/", body)
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
}
// wantPercent checks a limit percentage that a log line or a ban's notes
// give, what, and the setting that gave it, against want, as percentText
// gives them.
func wantPercent(t *testing.T, what string, percent *int64, setting, want string) {
t.Helper()
if got := percentText(percent, setting); got != want {
t.Errorf("%s is %s, want %s", what, got, want)
}
}
// percentText gives a limit percentage and the setting that gave it as
// text, such as "50 from SWWAF_ASN_LIMIT_PERCENT", or none when both are
// left out.
func percentText(percent *int64, setting string) string {
switch {
case percent == nil && setting == "":
return none
case percent == nil:
return "none from " + setting
default:
return fmt.Sprintf("%d from %s", *percent, setting)
}
}
+3 -2
View File
@@ -22,8 +22,9 @@ const (
// database or through GeoJS, and notes them for the log line, unless
// SWWAF_LOOKUP_SOURCE is off or the client is on a private, loopback or
// link-local address, which no lookup can place. The lookup database
// answers at once. With GeoJS, while a setting needs the answer, a new
// client's request waits for it. ctx is the request's own context.
// answers at once. With GeoJS, while a setting needs the answer, such as a
// country list or a biased threshold, a new client's request waits for it.
// ctx is the request's own context.
func (rq *request) lookUp(ctx context.Context) {
if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) {
return
+5 -1
View File
@@ -22,6 +22,10 @@ import (
// and country.
type asnAndCountry struct{ asn, asName, country string }
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
// lookup database.
const fileSource = "file"
func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) {
t.Parallel()
@@ -178,7 +182,7 @@ func TestClientsAreLookedUpInTheLookupDatabaseAndGeoJSIsNotAsked(t *testing.T) {
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
})
s, clk, server := startWithClock(t, geojsURL, map[string]string{
lookupSource: "file",
lookupSource: fileSource,
lookupDBPath: path,
allowedCountries: "DE",
rateLimitPerMinute: "1",
+3 -3
View File
@@ -124,11 +124,11 @@ func New(params Params) *Server {
h.geojs = lookup.New(lookup.Params{
URL: params.GeoJSURL,
Timeout: params.Config.LookupTimeout,
// The country lists and the headers act on the answer before the
// request goes on.
// The country lists, the headers and the biased thresholds act on
// the answer before the request goes on.
Wait: len(params.Config.DeniedCountries) > 0 ||
len(params.Config.ExclusivelyAllowedCountries) > 0 ||
params.Config.AddLookupHeaders,
params.Config.AddLookupHeaders || biasedThresholdsSet(params.Config),
Answered: h.addLookup,
Now: params.Now,
ProcessLog: params.ProcessLog,
+1 -1
View File
@@ -315,7 +315,7 @@ func newProxy(
var lookupFile *lookup.File
if cfg.LookupSource == "file" {
if cfg.LookupSource == fileSource {
lookupFile, err = lookup.OpenFile(lookup.FileParams{
Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue,
})
+16 -6
View File
@@ -56,9 +56,12 @@ type request struct {
lookedUp bool
lookupAnswer lookup.Answer
// counted is true for a request the rate limits counted, whose bytes
// the byte limits count once it has ended.
counted bool
start time.Time
// the byte limits count once it has ended. limitPercent and
// bytesPercent are then its client's limit percentages for the rate
// limits and for the byte limits.
counted bool
limitPercent, bytesPercent percentage
start time.Time
// checked is when the checks were done, and upstreamStart when the
// request was handed to the app.
checked time.Time
@@ -212,9 +215,10 @@ func (rq *request) check(ctx context.Context) *refusal {
// them refuses is not counted for the rate limits. Then come the rate
// limits, unless the client is in SWWAF_RATE_LIMIT_EXEMPT_NETS or the
// request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that
// every other request is counted, and last the rule files. A request
// exempt from the rate limits is exempt from the byte limits too. ctx is
// the request's own context.
// every other request is counted, each of them by the client's limit
// percentages, and last the rule files. A request exempt from the rate
// limits is exempt from the byte limits too. ctx is the request's own
// context.
func (rq *request) checkClient(ctx context.Context) string {
cfg := rq.h.config
if isInside(rq.client, cfg.AllowNets) {
@@ -239,6 +243,12 @@ func (rq *request) checkClient(ctx context.Context) string {
rq.counted = !isInside(rq.client, cfg.RateLimitExemptNets) &&
!pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if rq.counted {
rq.limitPercent, rq.bytesPercent = limitPercentages(cfg, rq.line.ASN, rq.line.Country)
rq.line.LimitPercent, rq.line.LimitPercentSetting = rq.limitPercent.logged()
rq.line.BytesPercent, rq.line.BytesPercentSetting = rq.bytesPercent.logged()
}
if rq.counted && rq.limitBroken(now) {
return requestlog.ActionRateLimited
}
+30 -16
View File
@@ -174,7 +174,7 @@ type Hit struct {
Kind string
// Window is "minute", "hour" or "day".
Window string
// Limit is the window's limit.
// Limit is the window's limit, as the client's percentage of it.
Limit int64
// Count is the client's requests, or bytes, counted in the window,
// this request's included.
@@ -196,21 +196,24 @@ type Counts struct {
// Count counts a request from client at now, in every window, whether or
// not it is refused, and returns the client's counts in each window. It
// reports whether the request takes the client over a rate limit, and the
// hit: the window whose limit it goes over, the shortest if it is over
// several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
return l.count(client, now, 1, 0)
// reports whether the request takes the client over a rate limit, of
// which the client gets the percentage percent, rounded down, and the hit:
// the window whose limit it goes over, the shortest if it is over
// several. A limit that is off stays off.
func (l *Limiter) Count(
client netip.Prefix, now time.Time, percent int64,
) (Counts, Hit, bool) {
return l.count(client, now, 1, 0, percent)
}
// CountBytes counts bytes, those of a request from client that has ended,
// at now, in every window, and returns the client's counts in each window.
// It reports whether the bytes take the client over a byte limit, and the
// hit, as Count does.
// It reports whether the bytes take the client over a byte limit, of which
// the client gets the percentage percent, and the hit, as Count does.
func (l *Limiter) CountBytes(
client netip.Prefix, now time.Time, bytes int64,
client netip.Prefix, now time.Time, bytes, percent int64,
) (Counts, Hit, bool) {
return l.count(client, now, 0, bytes)
return l.count(client, now, 0, bytes, percent)
}
// Reset sets client's counts of requests and of bytes in every window
@@ -373,9 +376,10 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
// count adds requests and bytes from client at now to its buckets in
// every window, and returns its counts. A limit is broken only by what is
// added to it, so that a request whose bytes are counted after another of
// the client's requests broke a rate limit does not break it too.
// the client's requests broke a rate limit does not break it too. The
// client gets the percentage percent of each limit.
func (l *Limiter) count(
client netip.Prefix, now time.Time, requests, bytes int64,
client netip.Prefix, now time.Time, requests, bytes, percent int64,
) (Counts, Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
@@ -391,16 +395,17 @@ func (l *Limiter) count(
for i, w := range l.windows {
requestCounts[i] = requestBuckets[i].add(now, w.length, requests)
byteCounts[i] = byteBuckets[i].add(now, w.length, bytes)
limit, byteLimit := percentOf(w.limit, percent), percentOf(w.byteLimit, percent)
switch {
case hit.Window != "":
case requests > 0 && w.limit > 0 && requestCounts[i] > float64(w.limit):
case requests > 0 && w.limit > 0 && requestCounts[i] > float64(limit):
hit = Hit{
Kind: KindRequests, Window: w.name, Limit: w.limit, Count: requestCounts[i],
Kind: KindRequests, Window: w.name, Limit: limit, Count: requestCounts[i],
}
case bytes > 0 && w.byteLimit > 0 && byteCounts[i] > float64(w.byteLimit):
case bytes > 0 && w.byteLimit > 0 && byteCounts[i] > float64(byteLimit):
hit = Hit{
Kind: KindBytes, Window: w.name, Limit: w.byteLimit, Count: byteCounts[i],
Kind: KindBytes, Window: w.name, Limit: byteLimit, Count: byteCounts[i],
}
}
}
@@ -446,6 +451,15 @@ type window struct {
byteLimit int64
}
// percentOf returns the percentage percent of limit, rounded down. It is
// written as limit's hundreds times percent, plus the rest's share, since
// limit*percent can overflow for a byte limit.
func percentOf(limit, percent int64) int64 {
const hundred = 100
return limit/hundred*percent + limit%hundred*percent/hundred
}
// add counts n requests, or n bytes, at now in a window of length, and
// returns the client's count in the window that ends at now: what is in
// the bucket under way, and what is in the bucket before it weighted by
+71 -11
View File
@@ -1,6 +1,7 @@
package ratelimit_test
import (
"math"
"net/netip"
"testing"
"time"
@@ -11,6 +12,10 @@ import (
// limit is the limit the tests set.
const limit = 3
// whole is the percentage of each limit a client gets when nothing lowers
// its limits.
const whole = 100
// The windows, as Count names them.
const (
minute = "minute"
@@ -62,14 +67,14 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight()
for range limit {
_, _, over := limiter.Count(client, start)
_, _, over := limiter.Count(client, start, whole)
if over {
t.Fatal("a request within the limit is over it")
}
}
// Over both limits; the minute's is named, with the four requests.
_, hit, over := limiter.Count(client, start)
_, hit, over := limiter.Count(client, start, whole)
want := ratelimit.Hit{
Kind: ratelimit.KindRequests, Window: minute, Limit: limit, Count: limit + 1,
@@ -80,6 +85,61 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
}
}
func TestClientGetsItsPercentageOfEachLimitRoundedDown(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 5, BytesPerDay: math.MaxInt64})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
// Half of 5 requests is 2.5, rounded down to 2: the third is over.
for range 2 {
_, _, over := limiter.Count(client, start, 50)
if over {
t.Fatal("a request within half the limit is over it")
}
}
_, hit, over := limiter.Count(client, start, 50)
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: minute, Limit: 2, Count: 3}
if !over || hit != want {
t.Errorf("the third request gives %+v and %t, want %+v and true", hit, over, want)
}
// Half of the largest byte limit is still far above a TiB: working it
// out does not overflow.
_, hit, over = limiter.CountBytes(client, start, 1<<40, 50)
if over {
t.Errorf("a TiB is over half the largest byte limit: %+v", hit)
}
}
func TestZeroPercentIsAZeroAllowanceAndALimitOffStaysOff(t *testing.T) {
t.Parallel()
// Only the hour has limits: the minute's and the day's are off.
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit, BytesPerHour: 1000})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
// At 0 percent, the first request and the first byte are over the
// hour's limits, which are 0; the minute's, which are off, stay off.
_, hit, _ := limiter.Count(client, start, 0)
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: hour, Limit: 0, Count: 1}
if hit != want {
t.Errorf("the first request gives %+v, want %+v", hit, want)
}
_, hit, _ = limiter.CountBytes(client, start, 1, 0)
want = ratelimit.Hit{Kind: ratelimit.KindBytes, Window: hour, Limit: 0, Count: 1}
if hit != want {
t.Errorf("the first byte gives %+v, want %+v", hit, want)
}
}
func TestEachByteLimitIsBrokenByTheBytesCounted(t *testing.T) {
t.Parallel()
@@ -100,12 +160,12 @@ func TestEachByteLimitIsBrokenByTheBytesCounted(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32")
// 600 bytes are within the limit, 600 more over it.
_, _, over := limiter.CountBytes(client, midnight(), 600)
_, _, over := limiter.CountBytes(client, midnight(), 600, whole)
if over {
t.Fatal("600 bytes are over the limit of 1000")
}
_, hit, over := limiter.CountBytes(client, midnight(), 600)
_, hit, over := limiter.CountBytes(client, midnight(), 600, whole)
want := ratelimit.Hit{
Kind: ratelimit.KindBytes, Window: tc.window, Limit: byteLimit, Count: 1200,
@@ -148,16 +208,16 @@ func TestCountGivesTheBytesInEachWindow(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
limiter.CountBytes(client, start, 300)
limiter.CountBytes(client, start, 300, whole)
// A quarter into the next hour, the minute has only these 100 bytes.
// The hour still covers three quarters of the bucket before, whose 300
// bytes count 225, and these: 325. The day covers all 400.
later := start.Add(time.Hour + time.Hour/4)
limiter.CountBytes(client, later, 100)
limiter.CountBytes(client, later, 100, whole)
// A request's counts give the bytes counted so far too.
counts, _, _ := limiter.Count(client, later)
counts, _, _ := limiter.Count(client, later, whole)
want := ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 100, HourBytes: 325, DayBytes: 400,
@@ -189,14 +249,14 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
start := midnight()
for range 3 {
limiter.Count(client, start)
limiter.Count(client, start, whole)
}
// A quarter into the next hour, the minute has only this request. The
// hour still covers three quarters of the bucket before, with its three
// requests, which count 2.25, and this one: 3.25. The day covers all
// four.
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4), whole)
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
if counts != want {
@@ -364,7 +424,7 @@ func wantCount(
) {
t.Helper()
_, hit, _ := limiter.Count(client, now)
_, hit, _ := limiter.Count(client, now, whole)
if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want)
@@ -379,7 +439,7 @@ func wantBytesCount(
) {
t.Helper()
_, hit, _ := limiter.CountBytes(client, now, bytes)
_, hit, _ := limiter.CountBytes(client, now, bytes, whole)
if hit.Kind != want {
t.Errorf("%d bytes from %s at %s break a limit on %q, want %q",
bytes, client, now.Format(time.RFC3339), hit.Kind, want)
+3 -3
View File
@@ -16,7 +16,7 @@ func TestSnapshotListsTheClientsByAddress(t *testing.T) {
limiter := ratelimit.New(ratelimit.Limits{})
for _, i := range []int{2, 3, 0, 1} {
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
limiter.Count(netip.MustParsePrefix(want[i]), midnight(), whole)
}
snapshot := limiter.Snapshot()
@@ -63,8 +63,8 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
start := midnight()
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start)
limiter.CountBytes(client, start, 5)
limiter.Count(client, start, whole)
limiter.CountBytes(client, start, 5, whole)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client {
+9
View File
@@ -117,6 +117,15 @@ type Line struct {
// ActionBanned, ActionCountryDenied, ActionRateLimited or
// ActionRuleBlocked.
WouldAction string `json:"would_action,omitempty"`
// LimitPercent and LimitPercentSetting are, for a request the rate
// limits counted whose client a biased threshold gives a percentage of
// the rate limits below 100, that percentage and the setting that gave
// it. BytesPercent and BytesPercentSetting are the same for the byte
// limits.
LimitPercent *int64 `json:"limit_percent,omitempty"`
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
BytesPercent *int64 `json:"bytes_percent,omitempty"`
BytesPercentSetting string `json:"bytes_percent_setting,omitempty"`
// Counts are, for a request the rate limits counted, the client's
// requests as they counted them with this one, and its bytes as the
// byte limits counted them, with this request's once it has ended if
+6 -3
View File
@@ -45,6 +45,9 @@ const (
// maxLogLines is how many lines of the process log wait for a test to
// read them.
maxLogLines = 64
// whole is the percentage of each limit a client gets when nothing
// lowers its limits.
whole = 100
)
// permanentBansJSON is bans.json holding permanentBan.
@@ -713,7 +716,7 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight())
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight(), whole)
err = files.WriteAll()
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
@@ -1320,10 +1323,10 @@ func fill(params state.Params) {
bans.Notes{RuleID: "env-file", Target: "path"})
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
params.Limiter.Count(netip.MustParsePrefix(c), now)
params.Limiter.Count(netip.MustParsePrefix(c), now, whole)
}
params.Limiter.CountBytes(client, now, 8)
params.Limiter.CountBytes(client, now, 8, whole)
params.Limiter.AddToHistory(client, now, ratelimit.Request{
Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
})