package progress import ( "context" "fmt" "io" "maps" "slices" "testing" "time" "github.com/pkg/errors" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/sync/errgroup" ) func TestProgress(t *testing.T) { t.Parallel() s, err := calc(context.TODO(), 4, "calc") require.NoError(t, err) assert.Equal(t, 10, s) eg, ctx := errgroup.WithContext(context.Background()) pr, ctx, cancelProgress := NewContext(ctx) var trace trace eg.Go(func() error { return saveProgress(ctx, pr, &trace) }) pw, _, ctx2 := NewFromContext(ctx, WithMetadata("tag", "foo")) s, err = calc(ctx2, 5, "calc") pw.Close() require.NoError(t, err) assert.Equal(t, 15, s) cancelProgress(errors.WithStack(context.Canceled)) err = eg.Wait() require.NoError(t, err) assert.Greater(t, len(trace.items), 5) assert.LessOrEqual(t, len(trace.items), 7) for _, p := range trace.items { v, ok := p.Meta("tag") assert.True(t, ok) assert.Equal(t, "foo", v.(string)) } } func TestProgressNested(t *testing.T) { t.Parallel() eg, ctx := errgroup.WithContext(context.Background()) pr, ctx, cancelProgress := NewContext(ctx) var trace trace eg.Go(func() error { return saveProgress(ctx, pr, &trace) }) s, err := reduceCalc(ctx, 3) require.NoError(t, err) assert.Equal(t, 6, s) cancelProgress(errors.WithStack(context.Canceled)) err = eg.Wait() require.NoError(t, err) last := map[string]*Progress{} for _, p := range trace.items { prev, ok := last[p.ID] if !ok || p.Timestamp.After(prev.Timestamp) { last[p.ID] = p } } require.ElementsMatch(t, []string{"reduce", "synccalc", "calc-0", "calc-1"}, slices.Collect(maps.Keys(last))) assert.Equal(t, Status{Action: "starting"}, last["reduce"].Sys) for _, id := range []string{"synccalc", "calc-0", "calc-1"} { assert.Equal(t, Status{Action: "done", Current: 3, Total: 3}, last[id].Sys) } } func calc(ctx context.Context, total int, name string) (int, error) { pw, _, ctx := NewFromContext(ctx) defer pw.Close() sum := 0 pw.Write(name, Status{Action: "starting", Total: total}) for i := 1; i <= total; i++ { select { case <-ctx.Done(): return 0, context.Cause(ctx) case <-time.After(10 * time.Millisecond): } if i == total { pw.Write(name, Status{Action: "done", Total: total, Current: total}) } else { pw.Write(name, Status{Action: "calculating", Total: total, Current: i}) } sum += i } return sum, nil } func reduceCalc(ctx context.Context, total int) (int, error) { eg, ctx := errgroup.WithContext(ctx) pw, _, ctx := NewFromContext(ctx) defer pw.Close() pw.Write("reduce", Status{Action: "starting"}) // sync step sum, err := calc(ctx, total, "synccalc") if err != nil { return 0, err } // parallel steps for i := range 2 { func(i int) { eg.Go(func() error { _, err := calc(ctx, total, fmt.Sprintf("calc-%d", i)) return err }) }(i) } if err := eg.Wait(); err != nil { return 0, err } return sum, nil } type trace struct { items []*Progress } func saveProgress(ctx context.Context, pr Reader, t *trace) error { for { p, err := pr.Read(ctx) if err != nil { if errors.Is(err, io.EOF) { return nil } return err } t.items = append(t.items, p...) } }