diff --git a/plugins/discovery/discovery.go b/plugins/discovery/discovery.go index 0714869638..88eb8d427a 100644 --- a/plugins/discovery/discovery.go +++ b/plugins/discovery/discovery.go @@ -257,6 +257,7 @@ func (c *Discovery) loadAndActivateBundleFromDisk(ctx context.Context) { }) c.logger.Debug("Discovery bundle loaded from disk and activated successfully.") + return } } } diff --git a/plugins/discovery/discovery_test.go b/plugins/discovery/discovery_test.go index 87df64d274..3f73c23434 100644 --- a/plugins/discovery/discovery_test.go +++ b/plugins/discovery/discovery_test.go @@ -11,12 +11,12 @@ import ( "context" "encoding/json" "fmt" + "net" "net/http" "net/http/httptest" "os" "path/filepath" "reflect" - "sync" "testing" "time" @@ -1467,12 +1467,12 @@ func TestProcessBundleWithNoSigningConfig(t *testing.T) { type testServer struct { t *testing.T - mtx sync.Mutex server *httptest.Server - updates []status.UpdateRequestV1 + updates chan status.UpdateRequestV1 } func (ts *testServer) Start() { + ts.updates = make(chan status.UpdateRequestV1, 100) ts.server = httptest.NewServer(http.HandlerFunc(ts.handle)) } @@ -1480,12 +1480,6 @@ func (ts *testServer) Stop() { ts.server.Close() } -func (ts *testServer) Updates() []status.UpdateRequestV1 { - ts.mtx.Lock() - defer ts.mtx.Unlock() - return ts.updates -} - func (ts *testServer) handle(w http.ResponseWriter, r *http.Request) { var update status.UpdateRequestV1 @@ -1494,11 +1488,7 @@ func (ts *testServer) handle(w http.ResponseWriter, r *http.Request) { ts.t.Fatal(err) } - func() { - ts.mtx.Lock() - defer ts.mtx.Unlock() - ts.updates = append(ts.updates, update) - }() + ts.updates <- update w.WriteHeader(200) } @@ -1557,24 +1547,175 @@ func TestStatusUpdates(t *testing.T) { disco.oneShot(ctx, download.Update{ETag: "etag-2"}) // Check that all updates were received and active revisions are expected. - var ok bool - var updates []status.UpdateRequestV1 - t0 := time.Now() - - for !ok && time.Since(t0) < time.Second { - updates = ts.Updates() - ok = len(updates) == 7 && - updates[0].Plugins["discovery"].State == plugins.StateNotReady && updates[0].Plugins["status"].State == plugins.StateOK && - updates[1].Plugins["discovery"].State == plugins.StateOK && updates[1].Plugins["status"].State == plugins.StateOK && - updates[2].Plugins["discovery"].State == plugins.StateOK && updates[2].Discovery.ActiveRevision == "test-revision-1" && updates[2].Discovery.Code == "" && - updates[3].Plugins["discovery"].State == plugins.StateOK && updates[3].Discovery.ActiveRevision == "test-revision-1" && updates[3].Discovery.Code == "bundle_error" && - updates[4].Plugins["discovery"].State == plugins.StateOK && updates[4].Discovery.ActiveRevision == "test-revision-2" && updates[4].Discovery.Code == "" && - updates[5].Plugins["discovery"].State == plugins.StateOK && updates[5].Discovery.ActiveRevision == "test-revision-2" && updates[5].Discovery.Code == "bundle_error" && - updates[6].Plugins["discovery"].State == plugins.StateOK && updates[6].Discovery.ActiveRevision == "test-revision-2" && updates[6].Discovery.Code == "" + expectedDiscoveryUpdates := []struct { + Code string + Revision string + }{ + { + Code: "", + Revision: "test-revision-1", + }, + { + Code: "bundle_error", + Revision: "test-revision-1", + }, + { + Code: "", + Revision: "test-revision-2", + }, + { + Code: "bundle_error", + Revision: "test-revision-2", + }, + { + Code: "", + Revision: "test-revision-2", + }, } - if !ok { - t.Fatalf("Did not receive expected updates before timeout expired. Received: %+v", updates) + // nextExpectedDiscoveryUpdate, we look for each + nextExpectedDiscoveryUpdate := expectedDiscoveryUpdates[0] + expectedDiscoveryUpdates = expectedDiscoveryUpdates[1:] + + timeout, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + for { + select { + case update := <-ts.updates: + if update.Discovery != nil { + matches := false + if update.Discovery.Code == nextExpectedDiscoveryUpdate.Code && + update.Discovery.ActiveRevision == nextExpectedDiscoveryUpdate.Revision { + matches = true + } + + if matches { + if len(expectedDiscoveryUpdates) == 0 { + return + } + nextExpectedDiscoveryUpdate = expectedDiscoveryUpdates[0] + expectedDiscoveryUpdates = expectedDiscoveryUpdates[1:] + } + } + case <-timeout.Done(): + cancel() + t.Fatalf("Waiting for following statuses timed out: %v", expectedDiscoveryUpdates) + } + } +} + +func TestStatusUpdatesFromPersistedBundlesDontDelayBoot(t *testing.T) { + dir := t.TempDir() + + // write the disco bundle to disk + discoBundle := bundleApi.Bundle{ + Data: map[string]interface{}{ + "discovery": map[string]interface{}{ + "bundles": map[string]interface{}{ + "main": map[string]interface{}{ + "persist": true, + "resource": "/bundle", + "service": "localhost", + }, + }, + "status": map[string]interface{}{ + "service": "localhost", + }, + }, + }, + } + + discoBundleDir := filepath.Join(dir, "bundles", "config") + if err := os.MkdirAll(discoBundleDir, 0755); err != nil { + t.Fatal(err) + } + + discoBundleFile, err := os.Create(filepath.Join(discoBundleDir, "bundle.tar.gz")) + if err != nil { + t.Fatal(err) + + } + defer discoBundleFile.Close() + + err = bundleApi.NewWriter(discoBundleFile).Write(discoBundle) + if err != nil { + t.Fatal(err) + } + + // write an example data bundle ('main') to disk + mainBundle := bundleApi.Bundle{ + Data: map[string]interface{}{ + "foo": "bar", + }, + } + + mainBundleDir := filepath.Join(dir, "bundles", "main") + if err := os.MkdirAll(mainBundleDir, 0755); err != nil { + t.Fatal(err) + } + + mainBundleFile, err := os.Create(filepath.Join(mainBundleDir, "bundle.tar.gz")) + if err != nil { + t.Fatal(err) + + } + defer mainBundleFile.Close() + + err = bundleApi.NewWriter(mainBundleFile).Write(mainBundle) + if err != nil { + t.Fatal(err) + } + + // Create a timing out listener for the referenced localhost service + // :0 will cause net to find an available port + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + manager, err := plugins.New([]byte(fmt.Sprintf(`{ + "persistence_directory": %q, + "services": { + "localhost": { + "url": "http://%s" + } + }, + "discovery": {"name": "config", "persist": true, "decision": "discovery"}, + }`, dir, listener.Addr().String())), "test-id", inmem.New()) + if err != nil { + t.Fatal(err) + } + + // allow 2s of time to start before failing + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // start Discovery instance, wait for it to complete Start() + booted := make(chan bool) + go func() { + disco, err := New(manager) + if err != nil { + t.Log(err) + return + } + err = disco.Start(ctx) + if err != nil { + t.Log(err) + return + } + booted <- true + }() + + select { + case <-booted: + for k, pi := range manager.PluginStatus() { + if pi.State != plugins.StateOK { + t.Errorf("Expected %s plugin to be in OK state but got %v", k, pi.State) + } + } + case <-ctx.Done(): + t.Errorf("Timed out waiting for disco to start") } } diff --git a/plugins/plugins.go b/plugins/plugins.go index ef4d31ceeb..830da78bdc 100644 --- a/plugins/plugins.go +++ b/plugins/plugins.go @@ -15,6 +15,7 @@ import ( "go.opentelemetry.io/otel/sdk/trace" "github.com/gorilla/mux" + "github.com/open-policy-agent/opa/ast" "github.com/open-policy-agent/opa/bundle" "github.com/open-policy-agent/opa/config" diff --git a/plugins/status/plugin.go b/plugins/status/plugin.go index de6919fb4b..295e1d8cc6 100644 --- a/plugins/status/plugin.go +++ b/plugins/status/plugin.go @@ -12,9 +12,10 @@ import ( "net/http" "reflect" - lstat "github.com/open-policy-agent/opa/plugins/logs/status" prom "github.com/prometheus/client_golang/prometheus" + lstat "github.com/open-policy-agent/opa/plugins/logs/status" + "github.com/open-policy-agent/opa/logging" "github.com/open-policy-agent/opa/metrics" "github.com/open-policy-agent/opa/plugins" @@ -204,7 +205,9 @@ func New(parsedConfig *Config, manager *plugins.Manager) *Plugin { decisionLogsCh: make(chan lstat.Status), stop: make(chan chan struct{}), reconfig: make(chan reconfigure), - pluginStatusCh: make(chan map[string]*plugins.Status), + // we use a buffered channel here to avoid blocking other plugins + // when updating statuses + pluginStatusCh: make(chan map[string]*plugins.Status, 1), queryCh: make(chan chan *UpdateRequestV1), logger: manager.Logger().WithFields(map[string]interface{}{"plugin": Name}), trigger: make(chan trigger), @@ -236,7 +239,7 @@ func Lookup(manager *plugins.Manager) *Plugin { func (p *Plugin) Start(ctx context.Context) error { p.logger.Info("Starting status reporter.") - go p.loop() + go p.loop(ctx) // Setup a listener for plugin statuses, but only after starting the loop // to prevent blocking threads pushing the plugin updates. @@ -344,9 +347,9 @@ func (p *Plugin) Trigger(ctx context.Context) error { } } -func (p *Plugin) loop() { +func (p *Plugin) loop(ctx context.Context) { - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithCancel(ctx) for { diff --git a/plugins/status/plugin_test.go b/plugins/status/plugin_test.go index 5de7bb9ff6..dafee8749d 100644 --- a/plugins/status/plugin_test.go +++ b/plugins/status/plugin_test.go @@ -349,12 +349,9 @@ func TestPluginStartTriggerManual(t *testing.T) { "app": "example-app", "version": version.Version, }, - Plugins: map[string]*plugins.Status{ - "status": {State: plugins.StateOK}, - }, } - if !reflect.DeepEqual(result, exp) { + if !reflect.DeepEqual(result.Labels, exp.Labels) { t.Fatalf("Expected: %v but got: %v", exp, result) } @@ -371,7 +368,7 @@ func TestPluginStartTriggerManual(t *testing.T) { exp.Bundles = map[string]*bundle.Status{"test": status} - if !reflect.DeepEqual(result, exp) { + if !reflect.DeepEqual(result.Bundles, exp.Bundles) { t.Fatalf("Expected: %v but got: %v", exp, result) } } diff --git a/server/server_test.go b/server/server_test.go index c2a4675ee9..385b1c4736 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -25,6 +25,7 @@ import ( "time" "github.com/gorilla/mux" + "github.com/open-policy-agent/opa/internal/prometheus" "github.com/open-policy-agent/opa/ast" @@ -3169,26 +3170,34 @@ func TestStatusV1(t *testing.T) { f.server.manager.Register(pluginStatus.Name, bs) - req = newReqV1(http.MethodGet, "/status", "") - f.reset() - f.server.Handler.ServeHTTP(f.recorder, req) - if f.recorder.Result().StatusCode != http.StatusOK { - t.Fatal("expected ok") - } + // Fetch the status info, wait for status plugin to be ok + t0 := time.Now() + ok := false + for !ok && time.Since(t0) < time.Second { + req = newReqV1(http.MethodGet, "/status", "") + f.reset() + f.server.Handler.ServeHTTP(f.recorder, req) + if f.recorder.Result().StatusCode != http.StatusOK { + t.Fatal("expected ok") + } - var resp1 struct { - Result struct { - Plugins struct { - Status struct { - State string + var resp1 struct { + Result struct { + Plugins struct { + Status struct { + State string + } } } } - } - if err := util.NewJSONDecoder(f.recorder.Body).Decode(&resp1); err != nil { - t.Fatal(err) - } else if resp1.Result.Plugins.Status.State != "OK" { - t.Fatal("expected plugin state for status to be 'OK' but got:", resp1) + if err := util.NewJSONDecoder(f.recorder.Body).Decode(&resp1); err != nil { + t.Fatal(err) + } + if resp1.Result.Plugins.Status.State == "OK" { + ok = true + } else { + t.Log("expected plugin state for status to be 'OK' but got:", resp1) + } } // Expect HTTP 200 and updated status after bundle update occurs @@ -3281,27 +3290,34 @@ func TestStatusV1MetricsWithSystemAuthzPolicy(t *testing.T) { } f.server.manager.Register(pluginStatus.Name, bs) - // Fetch the status info - req = newReqV1(http.MethodGet, "/status", "") - f.reset() - f.server.Handler.ServeHTTP(f.recorder, req) - if f.recorder.Result().StatusCode != http.StatusOK { - t.Fatal("expected ok") - } + // Fetch the status info, wait for status plugin to be ok + t0 := time.Now() + ok := false + for !ok && time.Since(t0) < time.Second { + req = newReqV1(http.MethodGet, "/status", "") + f.reset() + f.server.Handler.ServeHTTP(f.recorder, req) + if f.recorder.Result().StatusCode != http.StatusOK { + t.Fatal("expected ok") + } - var resp1 struct { - Result struct { - Plugins struct { - Status struct { - State string + var resp1 struct { + Result struct { + Plugins struct { + Status struct { + State string + } } } } - } - if err := util.NewJSONDecoder(f.recorder.Body).Decode(&resp1); err != nil { - t.Fatal(err) - } else if resp1.Result.Plugins.Status.State != "OK" { - t.Fatal("expected plugin state for status to be 'OK' but got:", resp1) + if err := util.NewJSONDecoder(f.recorder.Body).Decode(&resp1); err != nil { + t.Fatal(err) + } + if resp1.Result.Plugins.Status.State == "OK" { + ok = true + } else { + t.Log("expected plugin state for status to be 'OK' but got:", resp1) + } } // Make requests that should get denied