package proxy_test import ( "io" "maps" "math" "net/http" "reflect" "slices" "strings" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) const ( // requestIDHeader carries the request's id. requestIDHeader = "X-Request-ID" // instance is the SWWAF_INSTANCE_NAME a test sets. instance = "fsn1app1/gitea" // ipv6Client is a client on IPv6, and ipv6Group the netblock the rate // limits count it as. ipv6Client = "2001:db8::7" ipv6Group = "2001:db8::/64" ) func TestLogLineHasEachFieldWhereItApplies(t *testing.T) { t.Parallel() received := make(chan string, 2) // the request ids the app received app := startApp(t, func(w http.ResponseWriter, r *http.Request) { received <- r.Header.Get(requestIDHeader) _, _ = io.Copy(io.Discard, r.Body) if r.URL.Path != "/full" { w.WriteHeader(http.StatusNoContent) return } w.Header().Set("Content-Type", "text/html") w.Header().Set("Cache-Control", "no-store") w.Header().Set("Location", "/elsewhere") w.WriteHeader(http.StatusFound) _, _ = io.WriteString(w, "moved") }) addr, out := startProxy(t, app.URL, map[string]string{ trustedProxies: trustLocalhost, rateLimitExemptNets: localhost, instanceName: instance, logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE", }) // This request comes from ipv6Client through a trusted proxy, with a // body and each header the log line looks at, and is answered with a // redirect. conn := dial(t, addr) send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+ forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+ forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+ "Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+ "Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+ "Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+ "Set-Cookie: secret-set-cookie\r\n\r\na=b") wantStatus(t, readResponse(t, conn), http.StatusFound) // A request's log line can come after its answer: each is waited for // before the next request, so that the lines are in order. full := out.requestLines(t, 1)[0] // This one comes from 127.0.0.1, which the rate limits do not count, // with a body of 4 bytes whose length it does not announce, so that its // request_bytes is not its content_length, and no header the log line // looks at, and is answered with 204 and no header. conn = dial(t, addr) send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+ "Transfer-Encoding: chunked\r\n\r\n4\r\nbody\r\n0\r\n\r\n") wantStatus(t, readResponse(t, conn), http.StatusNoContent) bare := out.requestLines(t, 2)[1] wantFullLine(t, full) wantBareLine(t, bare) for _, line := range []logLine{full, bare} { got := <-received if got != line.RequestID { t.Errorf("the app received request id %q, the log line has %q", got, line.RequestID) } } if strings.Contains(out.text(), "secret") { t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s", out.text()) } } // wantFullLine checks the log line of the request with every header the // line looks at. Its timings are checked by TestTimingsAreInOrder. func wantFullLine(t *testing.T, line logLine) { t.Helper() headers := map[string]string{"accept": "text/html", "x-custom": "one, two"} want := withTimings(line, requestlog.Line{ Type: requestType, Time: line.Time, Instance: instance, ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure, Host: appHost, Path: "/full", Protocol: protocol, Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5, RequestID: "from-traefik", PeerIP: localhost, ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group, ContentType: "application/x-www-form-urlencoded", ContentLength: 3, RequestHeaders: headers, HasAuthorization: true, HasCookie: true, ResponseContentType: "text/html", UpstreamStatus: http.StatusFound, CacheControl: "no-store", Location: "/elsewhere", Action: requestlog.ActionForward, Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1}, }) if !reflect.DeepEqual(line.Line, want) { t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) } } // wantBareLine checks the log line of the request with none of them, and // that the fields that do not apply to it are left out. func wantBareLine(t *testing.T, line logLine) { t.Helper() want := withTimings(line, requestlog.Line{ Type: requestType, Time: line.Time, Instance: instance, ClientIP: localhost, Method: http.MethodPost, Scheme: plain, Host: appHost, Path: "/bare", Protocol: protocol, Status: http.StatusNoContent, RequestBytes: 4, RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32", UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward, }) if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" { t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want) } for _, name := range []string{ "forwarded_for", "content_type", "content_length", "request_headers", "has_authorization", "has_cookie", "websocket", "response_content_type", "cache_control", "location", "counts", } { _, present := line.fields[name] if present { t.Errorf("log line has %s, which does not apply", name) } } } // withTimings returns want with the timings of line. func withTimings(line logLine, want requestlog.Line) requestlog.Line { want.DurationTotal = line.DurationTotal want.DurationChecks = line.DurationChecks want.DurationUpstreamConnect = line.DurationUpstreamConnect want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte want.DurationUpstreamTotal = line.DurationUpstreamTotal return want } func TestHasAuthorizationAndHasCookieEachComeFromTheirOwnHeader(t *testing.T) { t.Parallel() const hasAuthorization, hasCookie = "has_authorization", "has_cookie" for _, tc := range []struct{ header, field, other string }{ {"Authorization", hasAuthorization, hasCookie}, {"Cookie", hasCookie, hasAuthorization}, } { t.Run("only "+tc.header, func(t *testing.T) { t.Parallel() app := startApp(t, func(http.ResponseWriter, *http.Request) {}) addr, out := startProxy(t, app.URL, nil) req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req.Header.Set(tc.header, "secret") wantStatus(t, do(t, req), http.StatusOK) line := out.requestLine(t) _, otherPresent := line.fields[tc.other] if line.fields[tc.field] != true || otherPresent { t.Errorf("log line has %s %v and %s %v, want true and none", tc.field, line.fields[tc.field], tc.other, line.fields[tc.other]) } }) } } func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) { t.Parallel() const sentID = "from-traefik" sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}} trusted := map[string]string{trustedProxies: trustLocalhost} for _, tc := range []struct { name string env map[string]string header http.Header // wantID is the request id logged, "" for a new one. wantID, wantScheme string }{ {"a trusted proxy's are kept", trusted, sent, sentID, secure}, {"without them, the id is new and the scheme http", trusted, nil, "", plain}, {"another peer's are replaced", nil, sent, "", plain}, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() received := make(chan string, 2) app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { received <- r.Header.Get(requestIDHeader) }) addr, out := startProxy(t, app.URL, tc.env) // Two requests, so that two new ids can be told apart. ids := make([]string, 0, 2) for i := range 2 { req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) maps.Copy(req.Header, tc.header) wantStatus(t, do(t, req), http.StatusOK) line := out.requestLines(t, i+1)[i] ids = append(ids, line.RequestID) got := <-received if line.RequestID != got || line.Scheme != tc.wantScheme { t.Errorf("log line has request_id %q and scheme %q, and the "+ "app received id %q; want the same id and scheme %q", line.RequestID, line.Scheme, got, tc.wantScheme) } } switch { case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID): t.Errorf("request ids %q, want %q", ids, tc.wantID) case tc.wantID == "" && (slices.Contains(ids, sentID) || slices.Contains(ids, "") || ids[0] == ids[1]): t.Errorf("request ids %q, want two new ones", ids) } }) } } func TestTimingsAreInOrder(t *testing.T) { t.Parallel() const denied = "192.0.2.50" // in SWWAF_DENY_NETS app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { // The pauses set the times apart; a hold-up of the test only // lengthens them. time.Sleep(time.Millisecond) w.WriteHeader(http.StatusOK) _ = http.NewResponseController(w).Flush() time.Sleep(time.Millisecond) _, _ = io.WriteString(w, "done") }) addr, out := startProxy(t, app.URL, map[string]string{ trustedProxies: trustLocalhost, denyNets: denied, }) // Each log line is waited for before the next request, so that the // lines are in order. wantStatus(t, get(t, addr, "/"), http.StatusOK) forwarded := out.requestLines(t, 1)[0] req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req.Header.Set(forwardedFor, denied) wantStatus(t, do(t, req), http.StatusForbidden) refused := out.requestLines(t, 2)[1] wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK) health := out.requestLines(t, 3)[2] // A request passed to the app has every timing; one refused, none of // the app's; the health check, which runs no check, only the total. wantTimings(t, forwarded, "duration_total", "duration_checks", "duration_upstream_connect", "duration_upstream_first_byte", "duration_upstream_total") wantTimings(t, refused, "duration_total", "duration_checks") wantTimings(t, health, "duration_total") if t.Failed() { return } // In whole microseconds, as they are logged, so that the sum below is // exact. total := microseconds(forwarded.DurationTotal) checks := microseconds(*forwarded.DurationChecks) connect := microseconds(*forwarded.DurationUpstreamConnect) firstByte := microseconds(*forwarded.DurationUpstreamFirstByte) upstream := microseconds(*forwarded.DurationUpstreamTotal) // The checks end before the request is handed to the app, and the // connection comes before the answer, which the app ends after a // pause. if checks+upstream > total || connect >= firstByte || firstByte >= upstream { t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+ "first byte %d, upstream total %d", total, checks, connect, firstByte, upstream) } if *refused.DurationChecks > refused.DurationTotal { t.Errorf("refused request's checks took %v of %v milliseconds", *refused.DurationChecks, refused.DurationTotal) } } // wantTimings checks that the timings named are the only ones line has. func wantTimings(t *testing.T, line logLine, want ...string) { t.Helper() var got []string for name := range line.fields { if strings.HasPrefix(name, "duration_") { got = append(got, name) } } slices.Sort(got) slices.Sort(want) if !slices.Equal(got, want) { t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want) } } // microseconds is a timing in whole microseconds. func microseconds(milliseconds float64) int64 { return int64(math.Round(milliseconds * 1000)) } func TestLogsAnUpgradedConnection(t *testing.T) { t.Parallel() app := startApp(t, echoAfterUpgrade) addr, out := startProxy(t, app.URL, nil) conn := dial(t, addr) send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+ "Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n") wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols) _ = conn.Close() line := out.requestLine(t) if line.fields["websocket"] != true { t.Errorf("log line has websocket %v, want true", line.fields["websocket"]) } }