package proxy_test import ( "bufio" "io" "net" "net/http" "net/netip" "strconv" "strings" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) // The byte limit settings. const ( bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE" bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR" bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY" bytesCount = "SWWAF_BYTES_COUNT" ) // The values of SWWAF_BYTES_COUNT. const ( countResponse = "response" countRequest = "request" countBoth = "both" ) const ( // bodyBytes is the size of the body of each request these tests send // with one, and answerBytes that of each answer of the app. bodyBytes = 30 answerBytes = 70 // byteLimit is the byte limit these tests set, as a setting: a request // with a body and its answer, 100 bytes, go over it. byteLimit = "99" // minuteBytes is limit_hit for SWWAF_BYTES_LIMIT_PER_MINUTE. minuteBytes = "minute_bytes" ) func TestEachByteLimitBansOnceTheResponseHasEnded(t *testing.T) { t.Parallel() const scraper = "192.0.2.200" for _, tc := range []struct { setting, window string // apart is the time between the two requests, which the window // still covers. apart time.Duration }{ {bytesLimitPerMinute, minute, 0}, {bytesLimitPerHour, "hour", 2 * time.Minute}, {bytesLimitPerDay, "day", 2 * time.Hour}, } { t.Run(tc.setting, func(t *testing.T) { t.Parallel() s, clk := startWithAnswers(t, map[string]string{ tc.setting: byteLimit, metricsToken: token, }) // 70 bytes are within the limit of 99. line, _ := s.download() if line.LimitHit != "" || line.Offence != "" { t.Errorf("log line has limit_hit %q and offence %q, want neither", line.LimitHit, line.Offence) } // 140 bytes are over it. The response is passed on whole, and // then bans the client for an hour. clk.advance(tc.apart) expires := requestlog.FormatTime(clk.Now().Add(time.Hour)) line, got := s.download() if got.err != nil || len(got.body) != answerBytes || line.ResponseBytes != answerBytes { t.Errorf("got %d bytes (%v), and the log line has response_bytes %d, "+ "want %d", len(got.body), got.err, line.ResponseBytes, answerBytes) } if line.LimitHit != tc.window+"_bytes" || line.Offence != requestlog.OffenceLimit || line.BanExpires != expires { t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+ "want %s_bytes, limit and %s", line.LimitHit, line.Offence, line.BanExpires, tc.window, expires) } s.get(client, http.StatusForbidden, requestlog.ActionBanned) wantMetric(t, s.scrape(scraper), `smallwebwaf_rate_limit_hits_total{`+ `instance="`+alertInstance+`",kind="bytes",window="`+tc.window+`"}`, 1) }) } } func TestResponseOverAByteLimitByItselfIsPassedOnWhole(t *testing.T) { t.Parallel() s, _ := startWithAnswers(t, map[string]string{bytesLimitPerMinute: "50"}) // The answer's 70 bytes are over the limit of 50 on their own. line, got := s.download() if got.err != nil || len(got.body) != answerBytes || line.LimitHit != minuteBytes { t.Errorf("got %d bytes (%v), and the log line has limit_hit %q, want %d and %s", len(got.body), got.err, line.LimitHit, answerBytes, minuteBytes) } s.get(client, http.StatusForbidden, requestlog.ActionBanned) } func TestBytesOfAnAnswerThatBreaksOffAreCounted(t *testing.T) { t.Parallel() s, clk, _, _ := startAppWithAlerts(t, breakOff, map[string]string{ bytesLimitPerMinute: "50", }) expires := requestlog.FormatTime(clk.Now().Add(time.Hour)) // The 70 bytes passed on before the app broke off are over the limit of // 50, and ban the client for an hour. line, got := s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK, requestlog.ActionUpstreamError) if len(got.body) != answerBytes || line.LimitHit != minuteBytes || line.BanExpires != expires { t.Errorf("got %d bytes, and the log line has limit_hit %q and ban_expires %q, "+ "want %d, %s and %s", len(got.body), line.LimitHit, line.BanExpires, answerBytes, minuteBytes, expires) } s.get(client, http.StatusForbidden, requestlog.ActionBanned) } func TestWebSocketBytesAreCountedOnceItCloses(t *testing.T) { t.Parallel() for _, tc := range []struct { setting string counted float64 }{ {countResponse, answerBytes}, {countRequest, bodyBytes}, {countBoth, bodyBytes + answerBytes}, } { t.Run(tc.setting, func(t *testing.T) { t.Parallel() s, _, _, _ := startAppWithAlerts(t, answerAfterUpgrade, map[string]string{ bytesLimitPerMinute: "29", bytesCount: tc.setting, }) // The client sends 30 bytes and the app 70, each over the limit // of 29, which bans the client once the WebSocket has closed. line := s.webSocket() if line.LimitHit != minuteBytes || line.Counts.MinuteBytes != tc.counted { t.Errorf("log line has limit_hit %q and minute_bytes %v, want %s and %v", line.LimitHit, line.Counts.MinuteBytes, minuteBytes, tc.counted) } s.get(client, http.StatusForbidden, requestlog.ActionBanned) }) } } func TestWebSocketPassesTheAnswerAfterTheClientStopsSending(t *testing.T) { t.Parallel() app := startApp(t, echoOnceTheClientStops) addr, out := startProxy(t, app.URL, map[string]string{trustedProxies: trustLocalhost}) s := &sender{t: t, addr: addr, out: out} conn, reader := s.openWebSocket() send(t, conn, uploadBody) // The client closes its sending side and waits for the answer, which the // app sends only once it has seen the client stop. smallwebwaf passes the // close on to the app through CloseWrite on upgradedConn; without that, // it closes both connections, and the answer is lost. tcp, ok := conn.(*net.TCPConn) if !ok { t.Fatalf("connection is a %T, want a *net.TCPConn", conn) } err := tcp.CloseWrite() if err != nil { t.Fatalf("close the sending side: %v", err) } got, err := io.ReadAll(reader) if err != nil || string(got) != uploadBody { t.Errorf("got %q (%v), want %q", got, err, uploadBody) } s.closeWebSocket(conn) } func TestBytesCountSaysWhichBytesCount(t *testing.T) { t.Parallel() for _, tc := range []struct { setting string // each is the bytes each request counts, and breaking the request // that goes over the limit of 99. each float64 breaking int }{ {countResponse, answerBytes, 2}, {countRequest, bodyBytes, 4}, {countBoth, bodyBytes + answerBytes, 1}, } { t.Run(tc.setting, func(t *testing.T) { t.Parallel() s, _ := startWithAnswers(t, map[string]string{ bytesLimitPerMinute: byteLimit, bytesCount: tc.setting, }) for i := 1; i <= tc.breaking; i++ { line := s.upload() want := "" if i == tc.breaking { want = minuteBytes } counted := float64(i) * tc.each if line.LimitHit != want || line.Counts.MinuteBytes != counted { t.Errorf("request %d: log line has limit_hit %q and minute_bytes %v, "+ "want %q and %v", i, line.LimitHit, line.Counts.MinuteBytes, want, counted) } } s.get(client, http.StatusForbidden, requestlog.ActionBanned) }) } } func TestByteLimitsLeaveOutWhatTheRateLimitsLeaveOut(t *testing.T) { t.Parallel() const ( allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS ) s, _ := startWithAnswers(t, map[string]string{ bytesLimitPerMinute: byteLimit, allowNets: allowed, rateLimitExemptNets: exempt, rateLimitExemptPaths: "/assets/", }) // Each sends 200 bytes, none of which is counted. for _, sent := range []struct{ from, path string }{ {allowed, "/"}, {exempt, "/"}, {client, "/assets/app.js"}, } { for range 2 { line, _ := s.requestWithBody(http.MethodPost, sent.from, sent.path, uploadHeader, uploadBody, http.StatusOK, requestlog.ActionForward) if _, counted := line.fields["counts"]; counted || line.LimitHit != "" { t.Errorf("%s %s: log line has counts %v and limit_hit %q, want neither", sent.from, sent.path, line.fields["counts"], line.LimitHit) } } } // A path that is not exempt is counted, and breaks the limit. line := s.upload() if line.LimitHit != minuteBytes { t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes) } } func TestByteLimitsOffCountTheBytesAndBanNoOne(t *testing.T) { t.Parallel() const off = "off" s, _ := startWithAnswers(t, map[string]string{ bytesLimitPerMinute: off, bytesLimitPerHour: off, bytesLimitPerDay: off, }) for i := 1; i <= 3; i++ { line := s.upload() counted := float64(i * (bodyBytes + answerBytes)) if line.LimitHit != "" || line.Counts.MinuteBytes != counted || line.Counts.HourBytes != counted || line.Counts.DayBytes != counted { t.Errorf("request %d: log line has limit_hit %q and counts %+v, "+ "want none and %v bytes in each window", i, line.LimitHit, line.Counts, counted) } } } func TestBanForABrokenByteLimitHasItsNotesAndItsAlert(t *testing.T) { t.Parallel() s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{ bytesLimitPerMinute: byteLimit, }) start := clk.Now() s.requestWithBody(http.MethodPost, client, "/upload?part=1", uploadHeader, uploadBody, http.StatusOK, requestlog.ActionForward) netblock := netip.MustParsePrefix(client + "/32") want := bans.Ban{ Netblock: netblock, Start: start, Expires: start.Add(time.Hour), Cause: bans.CauseLimit, Reason: "bytes per minute over the limit of " + byteLimit, Notes: bans.Notes{ Kind: "bytes", Limit: 99, Window: minute, Count: bodyBytes + answerBytes, // The request as it was answered, by the app. Request: bans.Request{ Time: start, Method: http.MethodPost, Host: appHost, Path: "/upload?part=1", Status: http.StatusOK, UserAgent: userAgent, }, Requests: 1, }, } got := server.Ledger.Bans(netblock) if len(got) != 1 || got[0] != want { t.Fatalf("bans\n%+v\nwant\n%+v", got, want) } wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, want, requestlog.FormatTime(want.Expires))) if offences := historyOf(t, server, client).Offences.Limit; offences != 1 { t.Errorf("history counts %d offences for a limit, want 1", offences) } } func TestObserveModeLogsAndAlertsAByteLimitAndBansNoOne(t *testing.T) { t.Parallel() s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{ mode: observe, bytesLimitPerMinute: byteLimit, }) start := clk.Now() // No ban sets the client's counters back to zero, so each request // breaks the limit again. The answer is the app's either way, and the // alert for the ban is not sent twice within the cooldown. for range 2 { line := s.upload() wantWouldAction(t, line, "") if line.LimitHit != minuteBytes || line.Offence != requestlog.OffenceLimit || line.BanExpires != "" { t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+ "want %s, limit and none", line.LimitHit, line.Offence, line.BanExpires, minuteBytes) } } if held := server.Ledger.Snapshot(); len(held) != 0 { t.Errorf("the ledger holds %+v, want no ban", held) } waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] if len(waiting) != 1 { t.Fatalf("%d alerts wait, want 1: %+v", len(waiting), waiting) } notes, _ := waiting[0].Detail["notes"].(bans.Notes) alert := banAlert(alerts.EventBan, start, client, bans.Ban{ Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit, Reason: "bytes per minute over the limit of " + byteLimit, Notes: notes, }, requestlog.FormatTime(start.Add(time.Hour))) alert.Detail["mode"] = observe wantAlerts(t, queue, alert) } func TestObserveModeLeavesOutTheBytesOfARequestEnforceModeRefuses(t *testing.T) { t.Parallel() s, _ := startWithAnswers(t, map[string]string{ mode: observe, rateLimitPerMinute: "1", bytesLimitPerMinute: "150", }) s.upload() // The second request breaks the rate limit, which in enforce mode would // refuse it before the app sent anything, so its 100 bytes are not // counted, and the byte limit is not broken. Its line gives the bytes // counted before it. line := s.upload() wantWouldAction(t, line, requestlog.ActionRateLimited) if line.LimitHit != minute || line.Counts.MinuteBytes != bodyBytes+answerBytes { t.Errorf("log line has limit_hit %q and minute_bytes %v, want minute and %d", line.LimitHit, line.Counts.MinuteBytes, bodyBytes+answerBytes) } } // uploadHeader and uploadBody are the header and the body of a request // with a body of bodyBytes. // //nolint:gochecknoglobals // a constant cannot call strings.Repeat var ( uploadHeader = "Content-Length: " + strconv.Itoa(bodyBytes) uploadBody = strings.Repeat("u", bodyBytes) ) // readAndAnswer is the app of these tests: it reads each request's whole // body and answers with answerBytes bytes. func readAndAnswer(w http.ResponseWriter, r *http.Request) { _, _ = io.Copy(io.Discard, r.Body) _, _ = io.WriteString(w, strings.Repeat("a", answerBytes)) } // breakOff is an app that announces an answer of twice answerBytes, and // breaks off after answerBytes. func breakOff(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Length", strconv.Itoa(2*answerBytes)) _, _ = io.WriteString(w, strings.Repeat("a", answerBytes)) } // answerAfterUpgrade is an app that switches protocols, as for a // WebSocket, and then answers each line it receives with a line of // answerBytes. func answerAfterUpgrade(w http.ResponseWriter, _ *http.Request) { conn, buffered, err := http.NewResponseController(w).Hijack() if err != nil { return } defer func() { _ = conn.Close() }() _, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" + "Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n") _ = buffered.Flush() for { _, err := buffered.ReadString('\n') if err != nil { return } _, _ = buffered.WriteString(strings.Repeat("a", answerBytes-1) + "\n") _ = buffered.Flush() } } // echoOnceTheClientStops is an app that switches protocols, as for a // WebSocket, reads what the client sends until the client stops sending, // and then sends it all back. func echoOnceTheClientStops(w http.ResponseWriter, _ *http.Request) { conn, buffered, err := http.NewResponseController(w).Hijack() if err != nil { return } defer func() { _ = conn.Close() }() _, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" + "Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n") _ = buffered.Flush() received, _ := io.ReadAll(buffered) _, _ = buffered.Write(received) _ = buffered.Flush() } // webSocket opens a WebSocket from client to answerAfterUpgrade, sends a // line of bodyBytes on it, reads the answer, and closes it. It checks the // answer, and the log line as request does, and returns the log line. func (s *sender) webSocket() logLine { s.t.Helper() conn, reader := s.openWebSocket() send(s.t, conn, strings.Repeat("u", bodyBytes-1)+"\n") got, err := reader.ReadString('\n') if err != nil || len(got) != answerBytes { s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes) } return s.closeWebSocket(conn) } // openWebSocket sends a request from client to switch protocols, as for a // WebSocket, and checks that the app switches. It returns the connection, // on which reading fails once waitLimit has passed, and a reader of what // the app sends on it. func (s *sender) openWebSocket() (net.Conn, *bufio.Reader) { s.t.Helper() conn := dial(s.t, s.addr) send(s.t, conn, "GET /socket HTTP/1.1\r\nHost: "+appHost+"\r\n"+forwardedFor+ ": "+client+"\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\n") err := conn.SetReadDeadline(time.Now().Add(waitLimit)) if err != nil { s.t.Fatalf("set read deadline: %v", err) } reader := bufio.NewReader(conn) res, err := http.ReadResponse(reader, nil) if err != nil { s.t.Fatalf("read the answer to the upgrade: %v", err) } _ = res.Body.Close() if res.StatusCode != http.StatusSwitchingProtocols { s.t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols) } return conn, reader } // closeWebSocket closes conn, a WebSocket openWebSocket opened, checks its // log line as request does, and returns it. func (s *sender) closeWebSocket(conn net.Conn) logLine { s.t.Helper() _ = conn.Close() line := s.out.requestLines(s.t, s.sent+1)[s.sent] s.sent++ wantLine(s.t, line, http.StatusSwitchingProtocols, requestlog.ActionForward) return line } // startWithAnswers is startAppWithAlerts in front of readAndAnswer, for a // test that looks at neither the server nor the alerts. func startWithAnswers(t *testing.T, env map[string]string) (*sender, *clock) { t.Helper() s, clk, _, _ := startAppWithAlerts(t, readAndAnswer, env) return s, clk } // download sends a GET request for / from client, and checks that the // app's answer is passed on, as request does. It returns the log line and // the answer. func (s *sender) download() (logLine, answer) { s.t.Helper() return s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK, requestlog.ActionForward) } // upload is download for a POST request with a body of bodyBytes, and // returns the log line. func (s *sender) upload() logLine { s.t.Helper() line, _ := s.requestWithBody(http.MethodPost, client, "/", uploadHeader, uploadBody, http.StatusOK, requestlog.ActionForward) return line }