package smallwebwaf_test import ( "bytes" "context" "encoding/json" "io" "net" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/smallwebwaf" ) const ( // waitLimit bounds how long a test waits for what should happen. waitLimit = 10 * time.Second // pollInterval is how often a test looks for a line. pollInterval = 10 * time.Millisecond // testVersion is the version the tests give smallwebwaf. testVersion = "test" // localhost is where the tests listen. localhost = "127.0.0.1" listenAddr = "SWWAF_LISTEN_ADDR" ) // 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) } // line returns the first line whose field key is value, waiting for it. func (o *output) line(t *testing.T, key, value string) map[string]any { t.Helper() deadline := time.Now().Add(waitLimit) for time.Now().Before(deadline) { o.mu.Lock() text := o.buf.String() o.mu.Unlock() for line := range strings.Lines(text) { var fields map[string]any err := json.Unmarshal([]byte(line), &fields) if err != nil { t.Fatalf("output line %q is not JSON: %v", line, err) } if fields[key] == value { return fields } } time.Sleep(pollInterval) } t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.buf.String()) return nil } // run runs smallwebwaf with the settings in env until ctx is done, and // returns its exit status. func run(ctx context.Context, env map[string]string, out *output) int { return smallwebwaf.Run(ctx, smallwebwaf.Params{ Version: testVersion, LookupEnv: func(name string) (string, bool) { value, ok := env[name] return value, ok }, Stdout: out, }) } func TestInvalidSettingStopsTheStart(t *testing.T) { t.Parallel() out := &output{} status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out) if status != 1 { t.Errorf("exit status %d, want 1", status) } line := out.line(t, "msg", "invalid setting") message, _ := line["error"].(string) if line["type"] != "process" || line["level"] != "ERROR" || !strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") { t.Errorf("start refused with %v", line) } } func TestAddressInUseStopsTheStart(t *testing.T) { t.Parallel() taken, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") if err != nil { t.Fatalf("listen: %v", err) } defer func() { _ = taken.Close() }() out := &output{} status := run(t.Context(), map[string]string{listenAddr: taken.Addr().String()}, out) if status != 1 { t.Errorf("exit status %d, want 1", status) } out.line(t, "msg", "cannot listen on SWWAF_LISTEN_ADDR") } func TestServesUntilToldToStop(t *testing.T) { t.Parallel() app := httptest.NewServer(http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(w, "hello from the app") })) defer app.Close() ctx, stop := context.WithCancel(t.Context()) out := &output{} exited := make(chan int, 1) go func() { exited <- run(ctx, map[string]string{ listenAddr: localhost + ":0", "SWWAF_UPSTREAM_URL": app.URL, }, out) }() starting := out.line(t, "msg", "starting") wantStartingLine(t, starting, app.URL) addr, _ := starting["address"].(string) wantGreeting(t, "http://"+addr+"/") out.line(t, "type", "request") stop() select { case status := <-exited: if status != 0 { t.Errorf("exit status %d, want 0", status) } case <-time.After(waitLimit): t.Fatal("still running after being told to stop") } out.line(t, "msg", "stopped") } // wantStartingLine checks that the line at start gives the version and // every setting's value. func wantStartingLine(t *testing.T, line map[string]any, appURL string) { t.Helper() settings, _ := line["settings"].(map[string]any) want := map[string]any{ listenAddr: localhost + ":0", "SWWAF_UPSTREAM_URL": appURL, "SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", "SWWAF_CLIENT_REQUEST_TIMEOUT": "60s", "SWWAF_CLIENT_RESPONSE_TIMEOUT": "30m", "SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s", "SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m", "SWWAF_REQUEST_MAX_BYTES": "100M", "SWWAF_RESPONSE_MAX_BYTES": "5G", } for name, value := range want { if settings[name] != value { t.Errorf("starting line gives %s=%v, want %v", name, settings[name], value) } } if line["version"] != testVersion || line["type"] != "process" { t.Errorf("starting line %v", line) } } // wantGreeting checks that a request to url gets the app's answer. func wantGreeting(t *testing.T, url string) { t.Helper() req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, http.NoBody) if err != nil { t.Fatalf("new request: %v", err) } transport := &http.Transport{} defer transport.CloseIdleConnections() res, err := (&http.Client{Transport: transport}).Do(req) if err != nil { t.Fatalf("request: %v", err) } body, err := io.ReadAll(res.Body) _ = res.Body.Close() if err != nil || string(body) != "hello from the app" { t.Errorf("got %q (%v), want the app's answer", body, err) } }