diff --git a/configuration.go b/configuration.go index c4046d9..d439c1f 100644 --- a/configuration.go +++ b/configuration.go @@ -4,6 +4,7 @@ import ( "log" "os" "strings" + "time" "github.com/fsnotify/fsnotify" "gopkg.in/yaml.v3" @@ -48,19 +49,28 @@ func WatchConfiguration(path string, callback func(*Configuration)) (func(), err done := make(chan struct{}) go func() { defer watcher.Close() + // A writer that truncates before it writes leaves an empty or + // partial file between its events, and empty YAML loads as a + // valid zero-value configuration. Loading on a short delay + // after the last event reads the final content instead of the + // window in between. + var reload <-chan time.Time 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) + reload = time.After(100 * time.Millisecond) } + case <-reload: + reload = nil + cfg, err := LoadConfiguration(path) + if err != nil { + log.Printf("reloading configuration: %v", err) + continue + } + callback(cfg) } } }() diff --git a/configuration_test.go b/configuration_test.go index 6ad5c62..e89258b 100644 --- a/configuration_test.go +++ b/configuration_test.go @@ -228,6 +228,43 @@ paths: } } +func TestWatchConfiguration_ignoresTruncationBeforeWrite(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() + + // A writer that truncates before it writes leaves an empty file + // between the two events. The pause lets the truncation's event + // arrive on its own, the way a slow writer would. + f, err := os.OpenFile(path, os.O_WRONLY, 0644) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + f.Truncate(0) + time.Sleep(20 * time.Millisecond) + f.Write([]byte("not: valid: yaml: [[[")) + f.Close() + + select { + case <-reloaded: + t.Fatal("callback should not be called for a half-written file") + case <-time.After(500 * time.Millisecond): + } +} + func TestWatchConfiguration_errorOnMissingFile(t *testing.T) { _, err := WatchConfiguration("/nonexistent/wicket.yaml", func(cfg *Configuration) {}) if err == nil {