package proxy_test import ( "bufio" "bytes" "encoding/json" "io" "maps" "net" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) const ( // shortTimeout is what a test sets a timeout to, to see it run out. shortTimeout = 300 * time.Millisecond // shortTimeoutSetting is shortTimeout as a setting's value. shortTimeoutSetting = "300ms" // longTimeoutSetting is a timeout that does not run out in a test. longTimeoutSetting = "10s" // waitLimit bounds how long a test waits for what should happen. waitLimit = 10 * time.Second // pollInterval is how often a test looks for a log line. pollInterval = 10 * time.Millisecond // localhost is where every test server listens, and so the address // smallwebwaf sees each test's requests come from. localhost = "127.0.0.1" ) // The settings the tests set. const ( clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT" clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT" upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT" upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT" requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES" responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" trustedProxies = "SWWAF_TRUSTED_PROXIES" rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE" ) // output collects what smallwebwaf writes on stdout. type output struct { mu sync.Mutex buf bytes.Buffer } // Write adds lines smallwebwaf writes. func (o *output) Write(p []byte) (int, error) { o.mu.Lock() defer o.mu.Unlock() return o.buf.Write(p) } // lines returns every line written so far, decoded. func (o *output) lines(t *testing.T) []map[string]any { t.Helper() o.mu.Lock() defer o.mu.Unlock() var lines []map[string]any for text := range strings.Lines(o.buf.String()) { var line map[string]any err := json.Unmarshal([]byte(text), &line) if err != nil { t.Fatalf("output line %q is not JSON: %v", text, err) } lines = append(lines, line) } return lines } // logLine is a request log line, as typed fields and as the JSON object // it was written as. type logLine struct { requestlog.Line fields map[string]any } // requestLines waits for count request log lines and returns them. func (o *output) requestLines(t *testing.T, count int) []logLine { t.Helper() deadline := time.Now().Add(waitLimit) for time.Now().Before(deadline) { var found []logLine for _, fields := range o.lines(t) { if fields["type"] == "request" { found = append(found, decodeLine(t, fields)) } } if len(found) >= count { return found } time.Sleep(pollInterval) } t.Fatalf("fewer than %d request log lines after %s", count, waitLimit) return nil } // requestLine waits for the request log line of a test's one request. func (o *output) requestLine(t *testing.T) logLine { t.Helper() return o.requestLines(t, 1)[0] } // decodeLine reads a request log line's fields into a logLine. func decodeLine(t *testing.T, fields map[string]any) logLine { t.Helper() encoded, err := json.Marshal(fields) if err != nil { t.Fatalf("encode %v: %v", fields, err) } line := logLine{fields: fields} err = json.Unmarshal(encoded, &line.Line) if err != nil { t.Fatalf("decode %s: %v", encoded, err) } return line } // startApp starts app as the app smallwebwaf passes requests to. func startApp(t *testing.T, app http.HandlerFunc) *httptest.Server { t.Helper() server := httptest.NewServer(app) t.Cleanup(server.Close) return server } // startProxy starts smallwebwaf in front of the app at appURL, with the // settings in env on top of the defaults, and returns where it listens and // what it writes. func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) { t.Helper() settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL} maps.Copy(settings, env) cfg, err := config.FromEnvironment(func(name string) (string, bool) { value, ok := settings[name] return value, ok }) if err != nil { t.Fatalf("settings %v: %v", settings, err) } out := &output{} server := proxy.New(proxy.Params{ Config: cfg, RequestLog: out, ProcessLog: requestlog.NewProcessLogger(out), }) listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") if err != nil { t.Fatalf("listen: %v", err) } go func() { _ = server.Serve(listener) }() t.Cleanup(func() { _ = server.Close() }) return listener.Addr().String(), out } // newClient returns an HTTP client that sends requests as they are made, // with no compression of its own. func newClient(t *testing.T) *http.Client { t.Helper() transport := &http.Transport{DisableCompression: true} t.Cleanup(transport.CloseIdleConnections) return &http.Client{Transport: transport} } // answer is a response as a test reads it: the status, the headers, as // much of the body as arrived, and the error that ended the reading, nil // when the whole body arrived. type answer struct { status int header http.Header body []byte err error } // readAnswer reads all of res, and closes its body. func readAnswer(res *http.Response) answer { body, err := io.ReadAll(res.Body) _ = res.Body.Close() return answer{status: res.StatusCode, header: res.Header, body: body, err: err} } // newRequest makes a request for path to smallwebwaf at addr. func newRequest(t *testing.T, method, addr, path string, body io.Reader) *http.Request { t.Helper() req, err := http.NewRequestWithContext(t.Context(), method, "http://"+addr+path, body) if err != nil { t.Fatalf("new request: %v", err) } return req } // do sends req and reads the answer. func do(t *testing.T, req *http.Request) answer { t.Helper() res, err := newClient(t).Do(req) if err != nil { t.Fatalf("%s %s: %v", req.Method, req.URL.Path, err) } return readAnswer(res) } // get sends a GET request for path to smallwebwaf at addr. func get(t *testing.T, addr, path string) answer { t.Helper() return do(t, newRequest(t, http.MethodGet, addr, path, http.NoBody)) } // dial opens a connection to smallwebwaf at addr, for requests the HTTP // client cannot make, such as one that stops sending halfway. func dial(t *testing.T, addr string) net.Conn { t.Helper() conn, err := (&net.Dialer{}).DialContext(t.Context(), "tcp", addr) if err != nil { t.Fatalf("dial %s: %v", addr, err) } t.Cleanup(func() { _ = conn.Close() }) return conn } // send writes text to conn. func send(t *testing.T, conn net.Conn, text string) { t.Helper() _, err := io.WriteString(conn, text) if err != nil { t.Fatalf("send: %v", err) } } // readResponse reads the answer to a request sent on conn. func readResponse(t *testing.T, conn net.Conn) answer { t.Helper() err := conn.SetReadDeadline(time.Now().Add(waitLimit)) if err != nil { t.Fatalf("set read deadline: %v", err) } res, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatalf("read response: %v", err) } return readAnswer(res) } // wantLine checks the request log line's status and action. func wantLine(t *testing.T, line logLine, status int, action string) { t.Helper() if line.Status != status || line.Action != action { t.Errorf("log line has status %d and action %q, want %d and %q", line.Status, line.Action, status, action) } } // wantStatus checks an answer's status. func wantStatus(t *testing.T, got answer, status int) { t.Helper() if got.status != status { t.Errorf("status %d, want %d", got.status, status) } } // wantTimedOut checks that what began at start ended once shortTimeout // had run out, and not much later. func wantTimedOut(t *testing.T, start time.Time) { t.Helper() took := time.Since(start) if took < shortTimeout || took > shortTimeout+waitLimit/2 { t.Errorf("took %s, want %s", took, shortTimeout) } }