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)
+		}
+	}
+}