diff --git a/configuration.go b/configuration.go index 6b17d99..c401ea3 100644 --- a/configuration.go +++ b/configuration.go @@ -1,9 +1,11 @@ package main import ( + "log" "os" "strings" + "github.com/fsnotify/fsnotify" "gopkg.in/yaml.v3" ) @@ -30,6 +32,41 @@ func LoadConfiguration(path string) (*Configuration, error) { return &cfg, nil } +var newWatcher = fsnotify.NewWatcher + +func WatchConfiguration(path string, callback func(*Configuration)) (func(), error) { + watcher, err := newWatcher() + if err != nil { + return nil, err + } + if err := watcher.Add(path); err != nil { + watcher.Close() + return nil, err + } + + done := make(chan struct{}) + go func() { + defer watcher.Close() + for { + select { + case <-done: + return + case event := <-watcher.Events: + if event.Has(fsnotify.Write) || event.Has(fsnotify.Create) { + cfg, err := LoadConfiguration(path) + if err != nil { + log.Printf("reloading configuration: %v", err) + continue + } + callback(cfg) + } + } + } + }() + + return func() { close(done) }, nil +} + func (c *Configuration) LookupSubscribeSecret(path string) string { if c == nil { return "" diff --git a/configuration_test.go b/configuration_test.go index 3a7aa6f..2af5f67 100644 --- a/configuration_test.go +++ b/configuration_test.go @@ -1,9 +1,13 @@ package main import ( + "errors" "os" "path/filepath" "testing" + "time" + + "github.com/fsnotify/fsnotify" ) func TestLoadConfiguration_valid(t *testing.T) { @@ -131,6 +135,88 @@ func TestLookupPathConfiguration_exactMatchOnly(t *testing.T) { } } +func TestWatchConfiguration_reloadsOnChange(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "wicket.yaml") + os.WriteFile(path, []byte(` +paths: + test/path: + subscribe_secret: "original" +`), 0644) + + reloaded := make(chan *Configuration, 1) + stop, err := WatchConfiguration(path, func(cfg *Configuration) { + reloaded <- cfg + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer stop() + + os.WriteFile(path, []byte(` +paths: + test/path: + subscribe_secret: "updated" +`), 0644) + + select { + case cfg := <-reloaded: + secret := cfg.LookupSubscribeSecret("test/path") + if secret != "updated" { + t.Errorf("expected updated secret, got %s", secret) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for config reload") + } +} + +func TestWatchConfiguration_ignoresInvalidYAML(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "wicket.yaml") + os.WriteFile(path, []byte(` +paths: + test/path: + subscribe_secret: "original" +`), 0644) + + reloaded := make(chan *Configuration, 1) + stop, err := WatchConfiguration(path, func(cfg *Configuration) { + reloaded <- cfg + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer stop() + + os.WriteFile(path, []byte("not: valid: yaml: [[["), 0644) + + select { + case <-reloaded: + t.Fatal("callback should not be called for invalid YAML") + case <-time.After(500 * time.Millisecond): + } +} + +func TestWatchConfiguration_errorOnMissingFile(t *testing.T) { + _, err := WatchConfiguration("/nonexistent/wicket.yaml", func(cfg *Configuration) {}) + if err == nil { + t.Fatal("expected error watching nonexistent file") + } +} + +func TestWatchConfiguration_errorCreatingWatcher(t *testing.T) { + original := newWatcher + newWatcher = func() (*fsnotify.Watcher, error) { + return nil, errors.New("simulated watcher error") + } + defer func() { newWatcher = original }() + + _, err := WatchConfiguration("/any/path", func(cfg *Configuration) {}) + if err == nil { + t.Fatal("expected error when watcher creation fails") + } +} + func TestLookupSubscribeSecret_nilConfig(t *testing.T) { var cfg *Configuration secret := cfg.LookupSubscribeSecret("any/path") diff --git a/go.mod b/go.mod index 4a02013..e60be94 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,9 @@ module tangled.org/guid.foo/wicket go 1.24 require ( + github.com/fsnotify/fsnotify v1.9.0 github.com/google/uuid v1.6.0 gopkg.in/yaml.v3 v3.0.1 ) + +require golang.org/x/sys v0.13.0 // indirect diff --git a/go.sum b/go.sum index b4c5744..900f5b3 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,9 @@ +github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= +github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +golang.org/x/sys v0.13.0 h1:Af8nKPmuFypiUBjVoU9V20FiaFXOcuZI21p0ycVYYGE= +golang.org/x/sys v0.13.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/main.go b/main.go index f0c55d7..c2b7548 100644 --- a/main.go +++ b/main.go @@ -26,6 +26,15 @@ func main() { } cfgPtr.Store(cfg) log.Printf("loaded configuration from %s", *configPath) + + stop, err := WatchConfiguration(*configPath, func(cfg *Configuration) { + cfgPtr.Store(cfg) + log.Printf("reloaded configuration from %s", *configPath) + }) + if err != nil { + log.Fatalf("watching configuration: %v", err) + } + defer stop() } broker := NewBroker(*bufferSize) diff --git a/server.go b/server.go index 044cc4b..a7308a2 100644 --- a/server.go +++ b/server.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "io" + "log" "net/http" "strings" "sync/atomic" @@ -106,7 +107,11 @@ func (s *Server) handleSSE(w http.ResponseWriter, r *http.Request) { } } - flusher := w.(http.Flusher) + flusher, ok := w.(http.Flusher) + if !ok { + http.Error(w, "streaming unsupported", http.StatusInternalServerError) + return + } filters := ParseFilters(r.URL.Query()) lastEventID := r.Header.Get("Last-Event-ID") @@ -132,7 +137,11 @@ func (s *Server) handleSSE(w http.ResponseWriter, r *http.Request) { if !MatchAll(filters, event) { continue } - data, _ := json.Marshal(event) + data, err := json.Marshal(event) + if err != nil { + log.Printf("marshaling event %s: %v", event.ID, err) + continue + } fmt.Fprintf(w, "id: %s\ndata: %s\n\n", event.ID, data) flusher.Flush() } diff --git a/server_test.go b/server_test.go index 59e3a38..b285d11 100644 --- a/server_test.go +++ b/server_test.go @@ -236,6 +236,32 @@ func TestServer_postInvalidJSON(t *testing.T) { } } +type bareResponseWriter struct { + code int + headers http.Header + body strings.Builder +} + +func (w *bareResponseWriter) Header() http.Header { return w.headers } +func (w *bareResponseWriter) Write(b []byte) (int, error) { return w.body.Write(b) } +func (w *bareResponseWriter) WriteHeader(code int) { w.code = code } + +func TestServer_sseWithoutFlusher(t *testing.T) { + broker := NewBroker(100) + var cfgPtr atomic.Pointer[Configuration] + handler := NewServer(broker, &cfgPtr) + + w := &bareResponseWriter{headers: make(http.Header)} + req := httptest.NewRequest("GET", "/test/topic", nil) + req.Header.Set("Accept", "text/event-stream") + + handler.ServeHTTP(w, req) + + if w.code != http.StatusInternalServerError { + t.Errorf("expected 500, got %d", w.code) + } +} + func TestServer_postWithMissingSignatureOnSecuredPath(t *testing.T) { cfg := &Configuration{ Paths: map[string]PathConfiguration{ diff --git a/sse_test.go b/sse_test.go index fd96d9b..8af7b64 100644 --- a/sse_test.go +++ b/sse_test.go @@ -5,6 +5,7 @@ import ( "bytes" "context" "encoding/json" + "math" "net/http" "net/http/httptest" "strings" @@ -290,6 +291,37 @@ func TestServer_sseChannelClosed(t *testing.T) { } } +func TestServer_sseSkipsUnmarshalableEvent(t *testing.T) { + ts, broker := newTestServer(nil) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/test/topic", nil) + time.Sleep(50 * time.Millisecond) + + broker.Publish(&Event{ + ID: "bad-event", + Path: "test/topic", + Payload: math.NaN(), + }) + broker.Publish(&Event{ + ID: "good-event", + Path: "test/topic", + Payload: map[string]any{"ok": true}, + }) + + select { + case event := <-events: + if event.ID != "good-event" { + t.Errorf("expected good-event, got %s", event.ID) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for good event after unmarshalable one") + } +} + func TestServer_nonJSONPayloadBase64(t *testing.T) { ts, _ := newTestServer(nil) defer ts.Close()