diff --git a/common.go b/common.go index 6b8ab80..fdc9f03 100644 --- a/common.go +++ b/common.go @@ -18,6 +18,7 @@ import ( "github.com/fosrl/newt/logger" "github.com/fosrl/newt/proxy" "github.com/fosrl/newt/websocket" + "github.com/fsnotify/fsnotify" "golang.org/x/net/icmp" "golang.org/x/net/ipv4" "golang.zx2c4.com/wireguard/tun/netstack" @@ -600,3 +601,61 @@ func sendBlueprint(client *websocket.Client, file string) error { return nil } + +func watchBlueprintFile(ctx context.Context, filePath string, send func() error) { + watcher, err := fsnotify.NewWatcher() + if err != nil { + logger.Error("blueprint watcher: failed to create: %v", err) + return + } + defer watcher.Close() + + if err := watcher.Add(filePath); err != nil { + logger.Error("blueprint watcher: failed to watch %s: %v", filePath, err) + return + } + + logger.Info("Watching blueprint file for changes: %s", filePath) + + var debounce *time.Timer + scheduleSend := func() { + if debounce != nil { + debounce.Stop() + } + debounce = time.AfterFunc(500*time.Millisecond, func() { + logger.Info("Blueprint file changed, resending...") + if err := send(); err != nil { + logger.Error("blueprint watcher: resend failed: %v", err) + } + }) + } + + for { + select { + case <-ctx.Done(): + if debounce != nil { + debounce.Stop() + } + return + case event, ok := <-watcher.Events: + if !ok { + return + } + switch { + case event.Has(fsnotify.Write) || event.Has(fsnotify.Create): + if event.Has(fsnotify.Create) { + _ = watcher.Add(filePath) + } + scheduleSend() + case event.Has(fsnotify.Remove) || event.Has(fsnotify.Rename): + _ = watcher.Add(filePath) + scheduleSend() + } + case err, ok := <-watcher.Errors: + if !ok { + return + } + logger.Error("blueprint watcher: %v", err) + } + } +} diff --git a/common_test.go b/common_test.go index 67c02cf..82765da 100644 --- a/common_test.go +++ b/common_test.go @@ -1,10 +1,151 @@ package main import ( + "context" "net" + "os" + "path/filepath" + "sync/atomic" "testing" + "time" ) +func TestWatchBlueprintFile_WriteTriggersSend(t *testing.T) { + f, err := os.CreateTemp(t.TempDir(), "blueprint-*.yaml") + if err != nil { + t.Fatal(err) + } + f.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var calls atomic.Int32 + go watchBlueprintFile(ctx, f.Name(), func() error { + calls.Add(1) + return nil + }) + + time.Sleep(50 * time.Millisecond) + + if err := os.WriteFile(f.Name(), []byte("content"), 0644); err != nil { + t.Fatal(err) + } + + time.Sleep(700 * time.Millisecond) + + if calls.Load() != 1 { + t.Errorf("expected 1 send call, got %d", calls.Load()) + } +} + +func TestWatchBlueprintFile_DebounceCoalescesEvents(t *testing.T) { + f, err := os.CreateTemp(t.TempDir(), "blueprint-*.yaml") + if err != nil { + t.Fatal(err) + } + f.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var calls atomic.Int32 + go watchBlueprintFile(ctx, f.Name(), func() error { + calls.Add(1) + return nil + }) + + time.Sleep(50 * time.Millisecond) + + for i := 0; i < 5; i++ { + if err := os.WriteFile(f.Name(), []byte("change"), 0644); err != nil { + t.Fatal(err) + } + time.Sleep(50 * time.Millisecond) + } + + time.Sleep(700 * time.Millisecond) + + if calls.Load() != 1 { + t.Errorf("expected 1 send call after debounce, got %d", calls.Load()) + } +} + +func TestWatchBlueprintFile_ContextCancellationStops(t *testing.T) { + f, err := os.CreateTemp(t.TempDir(), "blueprint-*.yaml") + if err != nil { + t.Fatal(err) + } + f.Close() + + ctx, cancel := context.WithCancel(context.Background()) + + done := make(chan struct{}) + go func() { + watchBlueprintFile(ctx, f.Name(), func() error { return nil }) + close(done) + }() + + time.Sleep(50 * time.Millisecond) + cancel() + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Error("watchBlueprintFile did not exit after context cancellation") + } +} + +func TestWatchBlueprintFile_AtomicWriteTriggersSend(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "blueprint.yaml") + if err := os.WriteFile(target, []byte("initial"), 0644); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var calls atomic.Int32 + go watchBlueprintFile(ctx, target, func() error { + calls.Add(1) + return nil + }) + + time.Sleep(50 * time.Millisecond) + + tmp := filepath.Join(dir, "blueprint.yaml.tmp") + if err := os.WriteFile(tmp, []byte("updated"), 0644); err != nil { + t.Fatal(err) + } + if err := os.Rename(tmp, target); err != nil { + t.Fatal(err) + } + + time.Sleep(700 * time.Millisecond) + + if calls.Load() < 1 { + t.Errorf("expected at least 1 send call after atomic write, got %d", calls.Load()) + } +} + +func TestWatchBlueprintFile_MissingFileReturnsGracefully(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + done := make(chan struct{}) + go func() { + watchBlueprintFile(ctx, "/nonexistent/path/blueprint.yaml", func() error { return nil }) + close(done) + }() + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Error("watchBlueprintFile did not return for missing file") + } +} + func TestParseTargetString(t *testing.T) { tests := []struct { name string @@ -13,7 +154,6 @@ func TestParseTargetString(t *testing.T) { wantTargetAddr string wantErr bool }{ - // IPv4 test cases { name: "valid IPv4 basic", input: "3001:192.168.1.1:80", @@ -35,8 +175,6 @@ func TestParseTargetString(t *testing.T) { wantTargetAddr: "10.0.0.1:443", wantErr: false, }, - - // IPv6 test cases { name: "valid IPv6 loopback", input: "3001:[::1]:8080", @@ -72,8 +210,6 @@ func TestParseTargetString(t *testing.T) { wantTargetAddr: "[::ffff:192.168.1.1]:6000", wantErr: false, }, - - // Hostname test cases { name: "valid hostname", input: "8080:example.com:80", @@ -95,8 +231,6 @@ func TestParseTargetString(t *testing.T) { wantTargetAddr: "localhost:3000", wantErr: false, }, - - // Error cases { name: "invalid - no colons", input: "invalid", @@ -169,7 +303,7 @@ func TestParseTargetString(t *testing.T) { } if tt.wantErr { - return // Don't check other values if we expected an error + return } if listenPort != tt.wantListenPort { @@ -183,7 +317,7 @@ func TestParseTargetString(t *testing.T) { } } -// TestParseTargetStringNetDialCompatibility verifies that the output is compatible with net.Dial +// TestParseTargetStringNetDialCompatibility verifies that the output is compatible with net.Dial. func TestParseTargetStringNetDialCompatibility(t *testing.T) { tests := []struct { name string @@ -201,8 +335,6 @@ func TestParseTargetStringNetDialCompatibility(t *testing.T) { t.Fatalf("parseTargetString(%q) unexpected error: %v", tt.input, err) } - // Verify the format is valid for net.Dial by checking it can be split back - // This doesn't actually dial, just validates the format _, _, err = net.SplitHostPort(targetAddr) if err != nil { t.Errorf("parseTargetString(%q) produced invalid net.Dial format %q: %v", tt.input, targetAddr, err) @@ -248,4 +380,4 @@ func TestShouldFireRecovery(t *testing.T) { } }) } -} +} \ No newline at end of file diff --git a/go.mod b/go.mod index f9e1603..14ba979 100644 --- a/go.mod +++ b/go.mod @@ -46,6 +46,7 @@ require ( github.com/docker/go-connections v0.7.0 // indirect github.com/docker/go-units v0.5.0 // indirect github.com/felixge/httpsnoop v1.0.4 // indirect + github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/google/btree v1.1.3 // indirect diff --git a/go.sum b/go.sum index c73ea67..0bdb3c1 100644 --- a/go.sum +++ b/go.sum @@ -24,6 +24,8 @@ github.com/docker/go-units v0.5.0 h1:69rxXcBk27SvSaaxTtLh/8llcHD8vYHT7WSdRZ/jvr4 github.com/docker/go-units v0.5.0/go.mod h1:fgPhTUdO+D/Jk86RDLlptpiXQzgHJF7gydDDbaIK4Dk= github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2Wg= github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U= +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/gaissmai/bart v0.26.1 h1:+w4rnLGNlA2GDVn382Tfe3jOsK5vOr5n4KmigJ9lbTo= github.com/gaissmai/bart v0.26.1/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c= github.com/go-crypt/crypt v0.14.15 h1:q1i5OMpL05r935IxWmXgpDAVF0nvi4SMoHhGXLBQUEQ= diff --git a/main.go b/main.go index ecdc3a6..546a12b 100644 --- a/main.go +++ b/main.go @@ -2245,6 +2245,12 @@ persistent_keepalive_interval=5`, util.FixKey(privateKey.String()), util.FixKey( } } + if blueprintFile != "" { + go watchBlueprintFile(ctx, blueprintFile, func() error { + return sendBlueprint(client, blueprintFile) + }) + } + // Wait for context cancellation (from signal or service stop) <-ctx.Done()