diff --git a/script.go b/script.go index 557c06cd..f0307f19 100644 --- a/script.go +++ b/script.go @@ -818,6 +818,7 @@ func WatchCmd(db *DB) script.Cmd { if err != nil { return nil, err } + defer iter.Close() header := tbl.TableHeader() if header == nil { diff --git a/script_test.go b/script_test.go index 35be2024..84030d1c 100644 --- a/script_test.go +++ b/script_test.go @@ -9,6 +9,7 @@ import ( "slices" "strings" "testing" + "time" "github.com/cilium/hive" "github.com/cilium/hive/cell" @@ -48,6 +49,30 @@ func TestScript(t *testing.T) { ) } +func TestWatchCmdClosesChangeIterator(t *testing.T) { + db, table, metrics := newTestDB(t) + + ctx, cancel := context.WithCancel(t.Context()) + state, err := script.NewState(ctx, t.TempDir(), nil) + require.NoError(t, err) + engine := script.Engine{ + Cmds: map[string]script.Cmd{"db/watch": WatchCmd(db)}, + } + + done := make(chan error, 1) + go func() { + done <- engine.ExecuteLine(state, "db/watch test", &strings.Builder{}) + }() + + require.Eventually(t, func() bool { + return expvarInt(metrics.DeleteTrackerCountVar.Get(table.Name())) == 1 + }, time.Second, time.Millisecond) + + cancel() + require.NoError(t, <-done) + require.EqualValues(t, 0, expvarInt(metrics.DeleteTrackerCountVar.Get(table.Name()))) +} + func TestHeaderLine(t *testing.T) { type retrieval struct { header string