package proxy_test import ( "bufio" "bytes" "errors" "io" "net/http" "slices" "strings" "sync/atomic" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) // A request target with an escaped slash and space in its path, and a // query with a parameter ReverseProxy cannot parse. const ( rawPath = "/some%2Fpath/with%20space" rawQuery = "b=2&a=1&bad=%zz;x" ) // chunkSize is the size of each part of a body a test sends in parts. const chunkSize = 1 << 10 var errNotStreamed = errors.New("the first part never reached the app") // appSaw is what the app received. type appSaw struct { method string target string header http.Header body []byte } func TestPassesRequestAndAnswerUnchanged(t *testing.T) { t.Parallel() requestBody := bytes.Repeat([]byte("request body "), 8000) answerBody := bytes.Repeat([]byte("answer body "), 8000) saw := make(chan appSaw, 1) app := startApp(t, func(w http.ResponseWriter, r *http.Request) { body, _ := io.ReadAll(r.Body) saw <- appSaw{r.Method, r.RequestURI, r.Header.Clone(), body} w.Header().Set("X-App", "yes") w.Header().Add("Set-Cookie", "a=1") w.Header().Add("Set-Cookie", "b=2") w.WriteHeader(http.StatusTeapot) _, _ = w.Write(answerBody) }) addr, out := startProxy(t, app.URL, nil) req := newRequest(t, http.MethodPatch, addr, rawPath+"?"+rawQuery, bytes.NewReader(requestBody)) req.Header.Add("X-Test", "one") req.Header.Add("X-Test", "two") req.Header.Set("User-Agent", "test-agent") got := do(t, req) wantAppSaw(t, <-saw, requestBody) wantAnswer(t, got, answerBody) line := out.requestLine(t) wantLine(t, line, http.StatusTeapot, requestlog.ActionForward) wantRequestFields(t, line, addr, len(requestBody), len(answerBody)) } // wantAppSaw checks that the app received the test's request unchanged. func wantAppSaw(t *testing.T, saw appSaw, body []byte) { t.Helper() if saw.method != http.MethodPatch || saw.target != rawPath+"?"+rawQuery { t.Errorf("app saw %s %s, want %s %s", saw.method, saw.target, http.MethodPatch, rawPath+"?"+rawQuery) } if !slices.Equal(saw.header.Values("X-Test"), []string{"one", "two"}) { t.Errorf("app saw X-Test %q", saw.header.Values("X-Test")) } if saw.header.Get("User-Agent") != "test-agent" { t.Errorf("app saw User-Agent %q", saw.header.Get("User-Agent")) } if !bytes.Equal(saw.body, body) { t.Errorf("app saw a body of %d bytes, want the %d sent", len(saw.body), len(body)) } } // wantAnswer checks that the client received the app's answer unchanged. func wantAnswer(t *testing.T, got answer, body []byte) { t.Helper() wantStatus(t, got, http.StatusTeapot) if got.header.Get("X-App") != "yes" { t.Errorf("client got X-App %q", got.header.Get("X-App")) } if !slices.Equal(got.header.Values("Set-Cookie"), []string{"a=1", "b=2"}) { t.Errorf("client got Set-Cookie %q", got.header.Values("Set-Cookie")) } if got.err != nil || !bytes.Equal(got.body, body) { t.Errorf("client got %d bytes (%v), want the %d the app sent", len(got.body), got.err, len(body)) } } // wantRequestFields checks the log line's fields about the request. func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) { t.Helper() want := requestlog.Line{ Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost, Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery, Protocol: "HTTP/1.1", Status: http.StatusTeapot, UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent), ResponseBytes: int64(received), UserAgent: "test-agent", Action: requestlog.ActionForward, DurationTotal: line.DurationTotal, DurationUpstreamTotal: line.DurationUpstreamTotal, } if line.Line != want { t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) } _, err := time.Parse(time.RFC3339, line.Time) if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 { t.Errorf("log line has time %q and durations %v and %v", line.Time, line.DurationTotal, line.DurationUpstreamTotal) } } func TestStreamsTheRequestBody(t *testing.T) { t.Parallel() chunk := bytes.Repeat([]byte("x"), chunkSize) firstArrived := make(chan struct{}) app := startApp(t, func(w http.ResponseWriter, r *http.Request) { first := make([]byte, len(chunk)) _, err := io.ReadFull(r.Body, first) if err != nil { return } close(firstArrived) rest, _ := io.ReadAll(r.Body) _, _ = w.Write(rest) }) addr, _ := startProxy(t, app.URL, nil) body, writer := io.Pipe() go func() { _, _ = writer.Write(chunk) select { case <-firstArrived: _, _ = writer.Write(chunk) _ = writer.Close() case <-time.After(waitLimit): _ = writer.CloseWithError(errNotStreamed) } }() got := do(t, newRequest(t, http.MethodPost, addr, "/upload", body)) if got.err != nil || !bytes.Equal(got.body, chunk) { t.Errorf("app read %d bytes after the first part (%v), want %d", len(got.body), got.err, len(chunk)) } } func TestStreamsTheAnswerBody(t *testing.T) { t.Parallel() chunk := bytes.Repeat([]byte("y"), chunkSize) firstArrived := make(chan struct{}) app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write(chunk) _ = http.NewResponseController(w).Flush() select { case <-firstArrived: _, _ = w.Write(chunk) case <-time.After(waitLimit): } }) addr, _ := startProxy(t, app.URL, nil) req := newRequest(t, http.MethodGet, addr, "/download", http.NoBody) res, err := newClient(t).Do(req) if err != nil { t.Fatalf("request: %v", err) } first := make([]byte, len(chunk)) _, err = io.ReadFull(res.Body, first) close(firstArrived) got := readAnswer(res) if err != nil || got.err != nil || !bytes.Equal(got.body, chunk) { t.Errorf("client read %d bytes after the first part (%v, %v), want %d", len(got.body), err, got.err, len(chunk)) } } func TestUpgradedConnectionOutlastsTheTimeouts(t *testing.T) { t.Parallel() app := startApp(t, echoAfterUpgrade) addr, out := startProxy(t, app.URL, map[string]string{ clientRequestTimeout: shortTimeoutSetting, clientResponseTimeout: shortTimeoutSetting, upstreamRequestTimeout: shortTimeoutSetting, upstreamResponseTimeout: shortTimeoutSetting, }) 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") reader := bufio.NewReader(conn) res, err := http.ReadResponse(reader, nil) if err != nil { t.Fatalf("read the answer to the upgrade: %v", err) } _ = res.Body.Close() if res.StatusCode != http.StatusSwitchingProtocols { t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols) } // Every timeout started before the upgrade was answered: wait past // them all, then use the connection. time.Sleep(3 * shortTimeout / 2) send(t, conn, "still here\n") echoed, err := reader.ReadString('\n') if err != nil || echoed != "still here\n" { t.Errorf("echo %q (%v), want %q", echoed, err, "still here\n") } _ = conn.Close() wantLine(t, out.requestLine(t), http.StatusSwitchingProtocols, requestlog.ActionForward) } // echoAfterUpgrade is an app that switches protocols on request, and then // sends back each line it receives. func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) { if r.Header.Get("Upgrade") != "websocket" { http.Error(w, "not an upgrade", http.StatusBadRequest) return } 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 { line, err := buffered.ReadString('\n') if err != nil { return } _, _ = buffered.WriteString(line) _ = buffered.Flush() } } func TestServerHasTheFixedLimits(t *testing.T) { t.Parallel() cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false }) if err != nil { t.Fatalf("default settings: %v", err) } server := proxy.New(proxy.Params{ Config: cfg, RequestLog: io.Discard, ProcessLog: requestlog.NewProcessLogger(io.Discard), }) if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 || server.IdleTimeout != 2*time.Minute || server.ReadHeaderTimeout != time.Minute { t.Errorf("server listens on %q with header limit %d, idle time %s and "+ "header timeout %s", server.Addr, server.MaxHeaderBytes, server.IdleTimeout, server.ReadHeaderTimeout) } } func TestRefusesHeadersOver32KiB(t *testing.T) { t.Parallel() var calls atomic.Int32 app := startApp(t, func(http.ResponseWriter, *http.Request) { calls.Add(1) }) addr, _ := startProxy(t, app.URL, nil) // size counts every byte of the request: the request line, the // headers and the blank line that ends them. const ( start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: " end = "\r\n\r\n" ) for _, tc := range []struct { size int want int }{ {size: 32 << 10, want: http.StatusOK}, {size: 32<<10 + 1, want: http.StatusRequestHeaderFieldsTooLarge}, } { conn := dial(t, addr) send(t, conn, start+strings.Repeat("a", tc.size-len(start)-len(end))+end) wantStatus(t, readResponse(t, conn), tc.want) } if calls.Load() != 1 { t.Errorf("the app was called %d times, want once", calls.Load()) } } func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) { t.Parallel() // No test can listen on port 1: listening on port 0 gets one from 32768 up. addr, out := startProxy(t, "http://"+localhost+":1", nil) wantStatus(t, get(t, addr, "/"), http.StatusBadGateway) wantLine(t, out.requestLine(t), http.StatusBadGateway, requestlog.ActionUpstreamError) logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool { return line["type"] == "process" && line["msg"] == "request to the app failed" }) if !logged { t.Errorf("no process line says the request to the app failed") } } func TestLogsAnAnswerThatBrokeOff(t *testing.T) { t.Parallel() app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(w, "the first part") _ = http.NewResponseController(w).Flush() panic(http.ErrAbortHandler) // drops the connection mid-answer }) addr, out := startProxy(t, app.URL, nil) got := get(t, addr, "/") if string(got.body) != "the first part" || !errors.Is(got.err, io.ErrUnexpectedEOF) { t.Errorf("client read %q (%v), want the first part cut off", got.body, got.err) } wantLine(t, out.requestLine(t), http.StatusOK, requestlog.ActionUpstreamError) } func TestLogsAClientThatWentAway(t *testing.T) { t.Parallel() arrived := make(chan struct{}) app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { close(arrived) <-r.Context().Done() }) addr, out := startProxy(t, app.URL, nil) conn := dial(t, addr) send(t, conn, "GET /slow HTTP/1.1\r\nHost: app\r\n\r\n") select { case <-arrived: case <-time.After(waitLimit): t.Fatal("the request never reached the app") } _ = conn.Close() line := out.requestLine(t) if !line.Aborted || line.Status != 0 || line.Action != requestlog.ActionForward { t.Errorf("log line has aborted %v, status %d and action %q, "+ "want true, 0 and %q", line.Aborted, line.Status, line.Action, requestlog.ActionForward) } }