Files
smallwebwaf/internal/rules/rules.go
T
clawbot 2c01d7f98f
check / check (push) Waiting to run
Rule files, and bans for a clear sign of attack (closes #24)
Every *.rules file in SWWAF_RULES_DIR is read at start and on each
change. Each request is checked against the rules after the rate
limits: log notes a match, block refuses with 403, ban refuses and bans
the netblock for SWWAF_ATTACK_BAN_DURATION, made permanent by its next
request or attack. path, query and uri are matched as the request line
sent them; header:Host and header:Transfer-Encoding are refused. Bans
gain a cause. The image ships 00-default.rules.

Judgement call: a header sent twice is matched with its values joined
by ", ".
Judgement call: SWWAF_MAX_BAN_DURATION does not cap a ban for an attack.
Not in this unit: offences for rule matches, with the error burst.

Model: opus-5-5
2026-10-06 17:04:23 +00:00

435 lines
12 KiB
Go

// Package rules reads the rule files: the plain text files in
// SWWAF_RULES_DIR, one rule to a line, that each request is checked
// against, as the "Rule files" section of SPEC.md describes. They are read
// at start, and again whenever one is edited, added or removed.
package rules
import (
"context"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net/http"
"os"
"path/filepath"
"regexp"
"slices"
"strings"
"sync/atomic"
"github.com/fsnotify/fsnotify"
)
// The actions a rule takes when it matches.
const (
// ActionLog notes the match in the request log, and does nothing else.
ActionLog = "log"
// ActionBlock refuses the request with 403.
ActionBlock = "block"
// ActionBan refuses the request and bans the client's netblock: the
// request is a clear sign of attack.
ActionBan = "ban"
)
// extension ends the name of every rule file.
const extension = ".rules"
// headerTarget starts the target that is one request header,
// header:<Name>.
const headerTarget = "header:"
// escapeLength is the length of a percent escape, such as %2e.
const escapeLength = 3
var (
// ruleLine is a rule: four fields separated by spaces or tabs, of
// which the fourth, the regex, runs to the end of the line.
ruleLine = regexp.MustCompile(`^([^ \t]+)[ \t]+([^ \t]+)[ \t]+([^ \t]+)[ \t]+(.+)$`)
// idChars are the characters of a rule's id.
idChars = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
)
var (
errNotRule = errors.New(
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs")
errNotID = errors.New("is not an id of letters, digits, - and _")
errNotTarget = errors.New(
"is not path, query, uri, method, host, user_agent, referer or header:<Name>")
errHeaderTakenOut = errors.New(
"names a header that Go's HTTP server takes out of every request, " +
"so a rule never sees it")
errNotAction = errors.New("is not log, block or ban")
errNotRegex = errors.New("does not compile")
errUsedTwice = errors.New("is already the id of the rule at")
)
// Rule is one rule of a rule file.
type Rule struct {
// ID names the rule in the request log, the metrics and ban notes.
ID string
// Target is what the regex is matched against, such as path or
// header:Accept.
Target string
// Action is ActionLog, ActionBlock or ActionBan.
Action string
regex *regexp.Regexp
}
// Params are what Load needs.
type Params struct {
// Dir is the directory of the rule files (SWWAF_RULES_DIR).
Dir string
// Enabled is SWWAF_RULES_ENABLED: while it is false, no file is read
// and no rule loaded.
Enabled bool
// ProcessLog receives how many rules were read, and the error in a
// rule file edited while smallwebwaf runs.
ProcessLog *slog.Logger
}
// Files are the rule files of a running smallwebwaf, and the rules read
// from them. They are safe for concurrent use.
type Files struct {
params Params
// rules are the rules loaded, in the order of their files' names, and
// then of their lines.
rules atomic.Pointer[[]Rule]
}
// Load reads the rules of every *.rules file in Dir, in the order of the
// files' names, unless Enabled is false. A Dir that cannot be read is an
// error, and so is a line that is not a rule, a rule for the Host or the
// Transfer-Encoding header, which Go's HTTP server takes out of every
// request, a regex that does not compile and an id used twice, each named
// with its file and line.
func Load(params Params) (*Files, error) {
f := &Files{params: params}
f.rules.Store(&[]Rule{})
if !params.Enabled {
return f, nil
}
rules, err := read(params.Dir)
if err != nil {
return nil, err
}
f.rules.Store(&rules)
f.logRead(len(rules))
return f, nil
}
// Match checks r against the rules, in order, and returns those it
// matches, up to the first whose action refuses it, block or ban, which
// is then the last one returned.
func (f *Files) Match(r *http.Request) []Rule {
var matched []Rule
for _, rule := range *f.rules.Load() {
if !rule.matches(r) {
continue
}
matched = append(matched, rule)
if rule.Action != ActionLog {
break
}
}
return matched
}
// Len returns how many rules are loaded.
func (f *Files) Len() int {
return len(*f.rules.Load())
}
// Watch watches Dir until ctx is done, and reads the rule files again
// whenever one is edited, added or removed. If they then hold an error,
// the rules stay as they were, the error is logged with its file and
// line, and the files are read again at the next change. If Dir cannot be
// watched, that is logged, and the rules stay as they were loaded. While
// Enabled is false, Watch returns at once.
func (f *Files) Watch(ctx context.Context) {
if !f.params.Enabled {
return
}
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(f.params.Dir)
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the rule files for edits",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the rule files for edits",
"directory", f.params.Dir)
for {
select {
case <-ctx.Done():
return
case event := <-watcher.Events:
if filepath.Ext(event.Name) == extension {
f.readAgain()
}
case err = <-watcher.Errors:
f.params.ProcessLog.Warn("watching the rule files failed",
"error", err.Error())
}
}
}
// readAgain reads the rule files again, in place of the rules loaded, or
// logs the error that keeps the rules as they were.
func (f *Files) readAgain() {
rules, err := read(f.params.Dir)
if err != nil {
f.params.ProcessLog.Error(
"a rule file has an error, and the rules stay as they were",
"error", err.Error())
return
}
f.rules.Store(&rules)
f.logRead(len(rules))
}
// logRead logs that the rule files were read, and how many rules they
// hold, which can be none.
func (f *Files) logRead(count int) {
f.params.ProcessLog.Info("read the rule files",
"directory", f.params.Dir, "rules", count)
}
// read returns the rules of every rule file in dir, in the order of the
// files' names, and then of their lines.
func read(dir string) ([]Rule, error) {
entries, err := os.ReadDir(dir)
if err != nil {
return nil, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
}
var rules []Rule
// places are where each id is, as "<file>, line <n>".
places := map[string]string{}
for _, entry := range entries {
if entry.IsDir() || filepath.Ext(entry.Name()) != extension {
continue
}
rules, err = readFile(filepath.Join(dir, entry.Name()), rules, places)
if err != nil {
return nil, err
}
}
return rules, nil
}
// readFile appends the rules of the rule file at path to rules. places
// are where each id read so far is, and gain those of the file.
func readFile(path string, rules []Rule, places map[string]string) ([]Rule, error) {
data, err := os.ReadFile(path) //nolint:gosec // a rule file, in SWWAF_RULES_DIR
if err != nil {
return nil, err
}
number := 0
for line := range strings.Lines(string(data)) {
number++
place := fmt.Sprintf("%s, line %d", path, number)
text := strings.TrimSuffix(strings.TrimSuffix(line, "\n"), "\r")
rule, isRule, err := parse(text)
if err != nil {
return nil, fmt.Errorf("%s: %w", place, err)
}
if !isRule {
continue
}
first, used := places[rule.ID]
if used {
return nil, fmt.Errorf("%s: the id %q %w %s", place, rule.ID, errUsedTwice, first)
}
places[rule.ID] = place
rules = append(rules, rule)
}
return rules, nil
}
// parse reads a line of a rule file. It returns false for a blank line
// and for a comment, a line that starts with #.
func parse(line string) (Rule, bool, error) {
line = strings.TrimLeft(line, " \t")
if line == "" || strings.HasPrefix(line, "#") {
return Rule{}, false, nil
}
fields := ruleLine.FindStringSubmatch(line)
if fields == nil {
return Rule{}, false, errNotRule
}
rule := Rule{ID: fields[1], Target: fields[2], Action: fields[3]}
switch {
case !idChars.MatchString(rule.ID):
return Rule{}, false, fmt.Errorf("the id %q %w", rule.ID, errNotID)
case !isTarget(rule.Target):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotTarget)
case strings.EqualFold(rule.Target, headerTarget+"Host"):
return Rule{}, false, fmt.Errorf(
"the target %q %w; the request's host is the target host",
rule.Target, errHeaderTakenOut)
case strings.EqualFold(rule.Target, headerTarget+"Transfer-Encoding"):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errHeaderTakenOut)
case !slices.Contains([]string{ActionLog, ActionBlock, ActionBan}, rule.Action):
return Rule{}, false, fmt.Errorf("the action %q %w", rule.Action, errNotAction)
}
regex, err := regexp.Compile(fields[4])
if err != nil {
return Rule{}, false, fmt.Errorf("the regex %w: %w", errNotRegex, err)
}
rule.regex = regex
return rule, true, nil
}
// isTarget reports whether target is one a rule may have.
func isTarget(target string) bool {
switch target {
case "path", "query", "uri", "method", "host", "user_agent", "referer":
return true
}
name, isHeader := strings.CutPrefix(target, headerTarget)
return isHeader && name != ""
}
// matches reports whether the rule's regex matches its target in r. For
// uri it is matched against the path and query as received, and against
// them once percent-decoded, so that an encoded probe cannot slip past.
func (rule Rule) matches(r *http.Request) bool {
if rule.Target == "uri" {
uri := pathAndQuery(r)
return rule.regex.MatchString(uri) || rule.regex.MatchString(decodeOnce(uri))
}
return rule.regex.MatchString(value(rule.Target, r))
}
// value returns what a rule with target, other than uri, is matched
// against in r: the path and the query as the client sent them, before
// any decoding or re-encoding, split at the first ?, and a header's values
// joined by ", ", as HTTP joins those of a header sent more than once.
func value(target string, r *http.Request) string {
switch target {
case "path":
path, _, _ := strings.Cut(pathAndQuery(r), "?")
return path
case "query":
_, query, _ := strings.Cut(pathAndQuery(r), "?")
return query
case "method":
return r.Method
case "host":
return r.Host
case "user_agent":
return header(r, "User-Agent")
case "referer":
return header(r, "Referer")
default:
return header(r, strings.TrimPrefix(target, headerTarget))
}
}
// pathAndQuery returns the path and the query of r as the client sent
// them, as the app is sent them: the target of its request line,
// r.RequestURI, of which a target with a scheme gives what follows the
// scheme and its :, and the host when // follows. So http://host/path, as
// a client sends it to a proxy, gives /path, and so does http:/path,
// which Go reads as a target with a scheme and no host. r.URL is not
// used: when the path holds a character it escapes, such as \ or a
// non-ASCII byte, it decodes the whole path and escapes it again, so that
// \ becomes %5C and %2e a dot.
func pathAndQuery(r *http.Request) string {
if !r.URL.IsAbs() {
return r.RequestURI
}
_, afterScheme, _ := strings.Cut(r.RequestURI, ":")
hostAndRest, hasHost := strings.CutPrefix(afterScheme, "//")
if !hasHost {
return afterScheme
}
start := strings.IndexAny(hostAndRest, "/?")
if start < 0 {
return ""
}
return hostAndRest[start:]
}
// header returns the values of r's header name joined by ", ", or "" if
// r has no such header.
func header(r *http.Request, name string) string {
return strings.Join(r.Header.Values(name), ", ")
}
// decodeOnce returns s with each percent escape, such as %2e, replaced by
// the byte it stands for. A % that is not followed by two hex digits is
// left as it is, so that a malformed escape cannot keep the rest of s
// from being decoded.
func decodeOnce(s string) string {
var decoded strings.Builder
for i := 0; i < len(s); i++ {
if s[i] == '%' && i+escapeLength <= len(s) {
b, err := hex.DecodeString(s[i+1 : i+escapeLength])
if err == nil {
decoded.Write(b)
i += escapeLength - 1
continue
}
}
decoded.WriteByte(s[i])
}
return decoded.String()
}