internal/poller: package for polling and updating a value Move the logic for updating experiments into a separate package. A subsequent CL will replace the logic in internal/postgres/excluded.go with a use of this package. Motivation: currently, multiple concurrent requests to IsExcluded will result in multiple DB queries, which is wasteful. It makes more sense to poll IsExcluded periodically, as we do for experiments. Change-Id: I33202c1ce1d94a5b1c99fe6a332d89174517fe08 Reviewed-on: https://go-review.googlesource.com/c/pkgsite/+/261818 Trust: Jonathan Amsterdam <jba@google.com> Run-TryBot: Jonathan Amsterdam <jba@google.com> TryBot-Result: kokoro <noreply+kokoro@google.com> Reviewed-by: Julie Qiu <julie@golang.org>
diff --git a/internal/middleware/experiment.go b/internal/middleware/experiment.go index a977952..2789a84 100644 --- a/internal/middleware/experiment.go +++ b/internal/middleware/experiment.go
@@ -9,7 +9,6 @@ "fmt" "hash/fnv" "net/http" - "sync" "time" "cloud.google.com/go/errorreporting" @@ -17,6 +16,7 @@ "golang.org/x/pkgsite/internal/derrors" "golang.org/x/pkgsite/internal/experiment" "golang.org/x/pkgsite/internal/log" + "golang.org/x/pkgsite/internal/poller" ) const experimentQueryParamKey = "experiment" @@ -32,27 +32,36 @@ // An Experimenter contains information about active experiments from the // experiment source. type Experimenter struct { - getExperiments ExperimentGetter - reporter Reporter - pollEvery time.Duration - mu sync.Mutex - snapshot []*internal.Experiment + p *poller.Poller } // NewExperimenter returns an Experimenter for use in the middleware. The // experimenter regularly polls for updates to the snapshot in the background. func NewExperimenter(ctx context.Context, pollEvery time.Duration, getter ExperimentGetter, rep Reporter) (_ *Experimenter, err error) { defer derrors.Wrap(&err, "middleware.NewExperimenter") - e := &Experimenter{ - getExperiments: getter, - reporter: rep, - pollEvery: pollEvery, - } + + initial, err := getter(ctx) // If we can't load the initial state, then fail. - if err := e.loadNextSnapshot(ctx); err != nil { + if err != nil { return nil, err } - go e.pollUpdates(ctx) + e := &Experimenter{ + p: poller.New( + initial, + func(ctx context.Context) (interface{}, error) { + return getter(ctx) + }, + func(err error) { + // Log and report // the error. + log.Error(ctx, err) + if rep != nil { + rep.Report(errorreporting.Entry{ + Error: fmt.Errorf("loading experiments: %v", err), + }) + } + }), + } + e.p.Start(ctx, pollEvery) return e, nil } @@ -70,10 +79,11 @@ // Experiments returns the experiments currently in use. func (e *Experimenter) Experiments() []*internal.Experiment { // Make a copy so the caller can't modify our state. - e.mu.Lock() - defer e.mu.Unlock() - exps := make([]*internal.Experiment, len(e.snapshot)) - for i, x := range e.snapshot { + snapshot := e.p.Current().([]*internal.Experiment) + // We don't need a lock here because e.p.current will be updated + // without modification. + exps := make([]*internal.Experiment, len(snapshot)) + for i, x := range snapshot { // Assume internal.Experiment has no pointers to mutable data. nx := *x exps[i] = &nx @@ -84,11 +94,9 @@ // setExperimentsForRequest sets the experiments for a given request. // Experiments should be stable for a given IP address. func (e *Experimenter) setExperimentsForRequest(r *http.Request) *http.Request { - e.mu.Lock() - defer e.mu.Unlock() - + snapshot := e.p.Current().([]*internal.Experiment) var exps []string - for _, exp := range e.snapshot { + for _, exp := range snapshot { if shouldSetExperiment(r, exp) { exps = append(exps, exp.Name) } @@ -97,47 +105,6 @@ return r.WithContext(experiment.NewContext(r.Context(), exps...)) } -// pollUpdates polls the experiment source for updates to the snapshot, until -// e.closeChan is closed. -func (e *Experimenter) pollUpdates(ctx context.Context) { - ticker := time.NewTicker(e.pollEvery) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - ctx2, cancel := context.WithTimeout(ctx, e.pollEvery) - if err := e.loadNextSnapshot(ctx2); err != nil { - // We already have a snapshot to fall back on, so log and report - // the error, but don't fail. - log.Error(ctx, err) - if e.reporter != nil { - e.reporter.Report(errorreporting.Entry{ - Error: fmt.Errorf("loading experiments: %v", err), - }) - } - } - cancel() - } - } -} - -// loadNextSnapshot loads and sets the current state of experiments from the -// experiment source. -func (e *Experimenter) loadNextSnapshot(ctx context.Context) (err error) { - defer derrors.Wrap(&err, "loadNextSnapshot") - snapshot, err := e.getExperiments(ctx) - if err != nil { - return err - } - e.mu.Lock() - e.snapshot = snapshot - e.mu.Unlock() - return nil -} - // shouldSetExperiment reports whether a given request should be enrolled in // the experiment, based on the ip. e.Name, and e.Rollout. //
diff --git a/internal/poller/poller.go b/internal/poller/poller.go new file mode 100644 index 0000000..19ab64c --- /dev/null +++ b/internal/poller/poller.go
@@ -0,0 +1,76 @@ +// Copyright 2020 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// Package poller supports periodic polling to load a value. +package poller + +import ( + "context" + "sync" + "time" +) + +// A Getter returns a value. +type Getter func(context.Context) (interface{}, error) + +// A Poller maintains a current value, and refreshes it by periodically +// polling for a new value. +type Poller struct { + getter Getter + onError func(error) + mu sync.Mutex + current interface{} +} + +// New creates a new poller with an initial value. The getter is invoked +// to obtain updated values. Errors returned from the getter are passed +// to onError. +func New(initial interface{}, getter Getter, onError func(error)) *Poller { + return &Poller{ + getter: getter, + onError: onError, + current: initial, + } +} + +// Start begins polling in a separate goroutine, at the given period. To stop +// the goroutine, cancel the context passed to Start. +func (p *Poller) Start(ctx context.Context, period time.Duration) { + ticker := time.NewTicker(period) + + go func() { + for { + select { + case <-ctx.Done(): + ticker.Stop() + return + case <-ticker.C: + ctx2, cancel := context.WithTimeout(ctx, period) + p.Poll(ctx2) + cancel() + } + } + }() +} + +// Poll calls p's getter immediately and synchronously. +func (p *Poller) Poll(ctx context.Context) { + next, err := p.getter(ctx) + if err != nil { + p.onError(err) + } else { + p.mu.Lock() + p.current = next + p.mu.Unlock() + } +} + +// Current returns the current value. Initially, this is the value passed to New. +// After each successful poll, the value is updated. +// If a poll fails, the value remains unchanged. +func (p *Poller) Current() interface{} { + p.mu.Lock() + defer p.mu.Unlock() + return p.current +}
diff --git a/internal/poller/poller_test.go b/internal/poller/poller_test.go new file mode 100644 index 0000000..090b6d0 --- /dev/null +++ b/internal/poller/poller_test.go
@@ -0,0 +1,61 @@ +// Copyright 2020 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package poller + +import ( + "context" + "strconv" + "testing" + "time" +) + +type numError struct { + num int +} + +func (e numError) Error() string { return strconv.Itoa(e.num) } + +func Test(t *testing.T) { + var goods, bads []int + + cur := -1 + getter := func(context.Context) (interface{}, error) { + // Even: success; odd: failure. + cur++ + if cur%2 == 0 { + return cur, nil + } + return nil, numError{cur} + } + + onError := func(err error) { + bads = append(bads, err.(numError).num) + } + + p := New(cur, getter, onError) + if got, want := p.Current(), cur; got != want { + t.Fatalf("got %v, want %v", got, want) + } + ctx, cancel := context.WithCancel(context.Background()) + p.Start(ctx, 50*time.Millisecond) + time.Sleep(100 * time.Millisecond) // wait for first poll + for i := 0; i < 10; i++ { + goods = append(goods, p.Current().(int)) + time.Sleep(60 * time.Millisecond) + } + cancel() + // Expect goods to be all even and non-decreasing. + for i, g := range goods { + if g%2 != 0 || (i > 0 && goods[i-1] > g) { + t.Errorf("incorrect 'good' value %d", g) + } + } + // Expect bads to be consecutive odd numbers. + for i, b := range bads { + if b%2 == 0 || (i > 0 && bads[i-1]+2 != b) { + t.Errorf("incorrect 'bad' value %d", b) + } + } +}