package rules import ( "context" "log/slog" "net/http" "net/http/httptest" "os" "path/filepath" "slices" "testing" "testing/synctest" "time" "github.com/fsnotify/fsnotify" "sneak.berlin/go/smallwebwaf/internal/alerts" ) // The tests below run readAfterChanges in a synctest bubble, where time is // a clock of the test's own: time.Sleep moves it on at once, and // synctest.Wait returns once readAfterChanges waits again, so that every // reading due by then is done. The test sends the changes itself, as the // watch of a directory cannot run in a bubble. func TestFileWrittenInTwoPartsTakenInOnlyWhole(t *testing.T) { t.Parallel() synctest.Test(t, func(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "50-app.rules") writeFile(t, path, "first path block ^/first\n") files := load(t, dir) changes := run(t, files) file, err := os.Create(path) //nolint:gosec // a file the test wrote if err != nil { t.Fatalf("create: %v", err) } defer func() { _ = file.Close() }() // The first part ends in the middle of a ban rule's regex, which, // read then, would ban every request. write(t, file, "first path block ^/first\nprobe path ban ^/") changes <- fsnotify.Event{Name: path, Op: fsnotify.Write} time.Sleep(quietTime - time.Nanosecond) synctest.Wait() wantMatched(t, files, "/anything") // The second part starts the wait again. write(t, file, `\.env$`+"\n") changes <- fsnotify.Event{Name: path, Op: fsnotify.Write} time.Sleep(quietTime - time.Nanosecond) synctest.Wait() wantMatched(t, files, "/.env") time.Sleep(time.Nanosecond) synctest.Wait() wantMatched(t, files, "/.env", "probe") wantMatched(t, files, "/anything") }) } func TestEditSavedBeforeTheWatchStartsTakenIn(t *testing.T) { t.Parallel() synctest.Test(t, func(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "50-app.rules") writeFile(t, path, "first path block ^/first\n") files := load(t, dir) // Saved after Load read the files, and before the directory was // watched, so that no change is seen for it. writeFile(t, path, "first path block ^/edited\n") run(t, files) time.Sleep(quietTime) synctest.Wait() wantMatched(t, files, "/edited", "first") }) } // load loads the rules in dir. func load(t *testing.T, dir string) *Files { t.Helper() files, err := Load(Params{ Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler), Alerts: alerts.New(alerts.Params{}), }) if err != nil { t.Fatalf("load: %v", err) } return files } // run runs files' readAfterChanges until the test ends, and returns the // channel that sends it changes. func run(t *testing.T, files *Files) chan<- fsnotify.Event { t.Helper() changes := make(chan fsnotify.Event) ctx, stop := context.WithCancel(t.Context()) stopped := make(chan struct{}) go func() { files.readAfterChanges(ctx, changes, nil) close(stopped) }() t.Cleanup(func() { stop() <-stopped }) return changes } // writeFile writes content to the file at path. func writeFile(t *testing.T, path, content string) { t.Helper() err := os.WriteFile(path, []byte(content), 0o600) if err != nil { t.Fatalf("write %s: %v", path, err) } } // write writes text to the end of file. func write(t *testing.T, file *os.File, text string) { t.Helper() _, err := file.WriteString(text) if err != nil { t.Fatalf("write: %v", err) } } // wantMatched checks the ids of the rules that a GET request for path // matches, in order. func wantMatched(t *testing.T, files *Files, path string, want ...string) { t.Helper() r := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "http://app.example"+path, nil) matched := files.Match(r) got := make([]string, 0, len(matched)) for _, rule := range matched { got = append(got, rule.ID) } if !slices.Equal(got, want) { t.Errorf("%s matched %v, want %v", path, got, want) } }