From 43207dfc78bdab95099e31c865a69af4be7344ed Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Sat, 19 Sep 2026 19:03:55 -0600 Subject: [PATCH 1/7] test(suite): add the scenario tests The operator half of the multi-controller suite: the fakes, the cluster fixture, shard identity helpers, and the scenario tests themselves, written as a consumer of github.com/multigres/testkit/ctrltest. A scenario test exercises a joint between controllers, which is what the existing tier cannot reach. The deletion protocol runs shard to tablegroup to multigrescluster; the fan-out is asserted as a sequence rather than an end state; the race test pins two reconcilers writing one object. Others cover shard lifecycle, selector impostors, transitions and round-trip equality, and thrash. Seven live operator defects are pinned with KnownDefect, so each one reproduces its own evidence on every run rather than only in a document. A pin passes while its defect is present and fails the day it is fixed, so the fix has to replace the pin with a positive assertion in the same change. The pool scale-up pin constructs its race rather than sampling it. The defect needs the new pooler to register after the shard has converged, which happens naturally about a third of the time. poolerSim can hold new registrations for one namespace, so the test holds, scales up, waits for two pool pods and one entry in status.podRoles to stay stable, then releases. The pooler then appears in etcd with no Kubernetes event to announce it. That precondition is deliberately a stability window rather than RequireQuiescent: once the defect is fixed, the shard requeues while the pooler is held and the namespace never goes quiet. Tests open with newCase(t), which allocates the namespace, registers it at the reconcile gate and attaches the failure dump, and everything hangs off that receiver: the assertions, the harness, the client verbs, and this package's own vocabulary. Sub(t) is the subtest form, keeping the namespace while binding the subtest T. Bare(t) is for the pure-logic tests that never touch the cluster. Tests scope their work to the case namespace throughout. envtest never really deletes a namespace, so an unscoped List would see everything every earlier test in the run created. Signed-off-by: Brent Graveland --- Makefile | 22 +- go.mod | 11 +- go.sum | 10 +- test/suite/case.go | 68 ++ test/suite/fakes.go | 265 ++++++ test/suite/fixture.go | 113 +++ test/suite/golden_test.go | 67 ++ test/suite/identity.go | 98 ++ test/suite/identity_test.go | 239 +++++ test/suite/main_test.go | 66 ++ test/suite/scenario_deletion_test.go | 295 ++++++ test/suite/scenario_fanout_test.go | 159 ++++ test/suite/scenario_race_test.go | 335 +++++++ test/suite/scenario_selector_impostor_test.go | 643 +++++++++++++ test/suite/scenario_shard_lifecycle_test.go | 896 ++++++++++++++++++ test/suite/scenario_shard_quiescence_test.go | 26 + test/suite/scenario_thrash_test.go | 416 ++++++++ test/suite/scenario_transitions_test.go | 466 +++++++++ test/suite/shard_requeue_test.go | 183 ++++ test/suite/suite.go | 223 +++++ test/suite/suite_test.go | 123 +++ .../multigateway-deployment.golden.yaml | 89 ++ test/suite/types.go | 48 + 23 files changed, 4849 insertions(+), 12 deletions(-) create mode 100644 test/suite/case.go create mode 100644 test/suite/fakes.go create mode 100644 test/suite/fixture.go create mode 100644 test/suite/golden_test.go create mode 100644 test/suite/identity.go create mode 100644 test/suite/identity_test.go create mode 100644 test/suite/main_test.go create mode 100644 test/suite/scenario_deletion_test.go create mode 100644 test/suite/scenario_fanout_test.go create mode 100644 test/suite/scenario_race_test.go create mode 100644 test/suite/scenario_selector_impostor_test.go create mode 100644 test/suite/scenario_shard_lifecycle_test.go create mode 100644 test/suite/scenario_shard_quiescence_test.go create mode 100644 test/suite/scenario_thrash_test.go create mode 100644 test/suite/scenario_transitions_test.go create mode 100644 test/suite/shard_requeue_test.go create mode 100644 test/suite/suite.go create mode 100644 test/suite/suite_test.go create mode 100644 test/suite/testdata/multigateway-deployment.golden.yaml create mode 100644 test/suite/types.go diff --git a/Makefile b/Makefile index 235d2b31..9aef3644 100644 --- a/Makefile +++ b/Makefile @@ -278,22 +278,38 @@ build-installer: manifests generate kustomize ## Generate consolidated install Y ##@ Test +# test/suite is the multi-controller envtest suite. It carries no build tag, so +# every `go test ./...` call site has to exclude it by path or it lands in the +# required check before it is ready. That is one filter per call site, which is +# the deliberate trade against a tag that someone forgets on a new file. +# +# -v is load-bearing rather than cosmetic. Each KnownDefect pin logs the defect +# it is standing on while that defect is still present, and without -v go test +# discards the output of a passing test, so a green CI run shows none of them. +# The suite is meant to be readable as the operator's live defect list, and -v +# is what makes that list visible without waiting for a pin to expire. +.PHONY: test-suite +test-suite: manifests generate fmt vet setup-envtest ## Run the multi-controller test suite + KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ + go test -v -p 1 -timeout 20m ./test/suite/... + .PHONY: test test: manifests generate fmt vet ## Run tests (no integration testing) KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ - go test -p 1 $$(go list ./... | grep -v /e2e) -coverprofile=cover.out + go test -p 1 $$(go list ./... | grep -v /e2e | grep -v /test/suite) -coverprofile=cover.out .PHONY: test-integration test-integration: manifests generate fmt vet setup-envtest ## Run integration tests KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ - go test -p 1 -tags=integration,verbose $$(go list ./... | grep -v /e2e) -coverprofile=cover.out + go test -p 1 -tags=integration,verbose $$(go list ./... | grep -v /e2e | grep -v /test/suite) -coverprofile=cover.out .PHONY: test-coverage test-coverage: manifests generate fmt vet setup-envtest ## Generate coverage report with HTML @mkdir -p coverage @echo "==> Generating coverage..." KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ - go test -p 1 -tags=integration,verbose ./... -coverprofile=coverage/combined.out -covermode=atomic + go test -p 1 -tags=integration,verbose $$(go list ./... | grep -v /e2e | grep -v /test/suite) \ + -coverprofile=coverage/combined.out -covermode=atomic @echo "==> Generating HTML report..." @go tool cover -html=coverage/combined.out -o=coverage/combined.html @echo "Generated: coverage/combined.html" diff --git a/go.mod b/go.mod index afd81bb4..6b0d4ca9 100644 --- a/go.mod +++ b/go.mod @@ -1,11 +1,12 @@ module github.com/multigres/multigres-operator -go 1.26.6 +go 1.27 require ( github.com/go-logr/logr v1.4.4 github.com/google/go-cmp v0.7.0 github.com/multigres/multigres v0.0.0-20260925193740-522b90425a83 + github.com/multigres/testkit v0.1.0 github.com/prometheus/client_golang v1.24.1 github.com/prometheus/client_model v0.6.3 github.com/stretchr/testify v1.12.1 @@ -15,14 +16,16 @@ require ( go.opentelemetry.io/otel v1.46.0 go.opentelemetry.io/otel/sdk v1.46.0 go.opentelemetry.io/otel/trace v1.46.0 + go.uber.org/goleak v1.3.0 google.golang.org/grpc v1.83.2 google.golang.org/protobuf v1.36.12 k8s.io/api v0.37.0 k8s.io/apimachinery v0.37.0 k8s.io/client-go v0.37.0 - k8s.io/utils v0.0.0-20260626114624-be93311217bd - sigs.k8s.io/controller-runtime v0.25.0 + k8s.io/utils v0.0.0-20260707023825-cf1189d6abe3 + sigs.k8s.io/controller-runtime v0.25.1 sigs.k8s.io/e2e-framework v0.7.0 + sigs.k8s.io/structured-merge-diff/v6 v6.4.2 ) require ( @@ -128,7 +131,6 @@ require ( go.opentelemetry.io/otel/sdk/log v0.22.0 // indirect go.opentelemetry.io/otel/sdk/metric v1.46.0 // indirect go.opentelemetry.io/proto/otlp v1.11.0 // indirect - go.uber.org/goleak v1.3.0 // indirect go.uber.org/multierr v1.11.0 // indirect go.uber.org/zap v1.27.1 // indirect go.yaml.in/yaml/v2 v2.4.4 // indirect @@ -157,7 +159,6 @@ require ( sigs.k8s.io/apiserver-network-proxy/konnectivity-client v0.36.0 // indirect sigs.k8s.io/json v0.0.0-20250730193827-2d320260d730 // indirect sigs.k8s.io/randfill v1.0.0 // indirect - sigs.k8s.io/structured-merge-diff/v6 v6.4.2 // indirect sigs.k8s.io/yaml v1.6.0 // indirect ) diff --git a/go.sum b/go.sum index f263bbdf..bdc302ba 100644 --- a/go.sum +++ b/go.sum @@ -199,6 +199,8 @@ github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee h1:W5t00kpgFd github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= github.com/multigres/multigres v0.0.0-20260925193740-522b90425a83 h1:IdyFGtwc9pEZEzqs6h5JfavAnErWfuKwj3d42qkZUx8= github.com/multigres/multigres v0.0.0-20260925193740-522b90425a83/go.mod h1:Ov2hrkOguWSkCS2QIhAdguFeG5GlZ3v4WGqIdqkQ7Tg= +github.com/multigres/testkit v0.1.0 h1:i6DiCFZ9mVhEopBwakMmel5NTde1VErEPGtnoXn9CdU= +github.com/multigres/testkit v0.1.0/go.mod h1:3ONhsV/PNOUke7PID5HPlnxTLyQCcfOQ/JfLCRkSLOY= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/onsi/ginkgo/v2 v2.27.4 h1:fcEcQW/A++6aZAZQNUmNjvA9PSOzefMJBerHJ4t8v8Y= @@ -457,12 +459,12 @@ k8s.io/kube-openapi v0.0.0-20260721132016-d427ff9ee9ad h1:oXImqH8mQNk7PmvzKhmN3d k8s.io/kube-openapi v0.0.0-20260721132016-d427ff9ee9ad/go.mod h1:0/mqHCVhlumdJ3BhCfnjSZQE037nAhNodh1/hK0T8/I= k8s.io/streaming v0.37.0 h1:iPBUZLZiKt5bV+lxJurASMOV07VuBhNpiwJt2//AWrM= k8s.io/streaming v0.37.0/go.mod h1:APlJR26ZWRcVy5bIEj0QRrKUXROtBHPcxl2NT7EAzPU= -k8s.io/utils v0.0.0-20260626114624-be93311217bd h1:Ea7fgQ5we8Y9T0OX5o0dAHzQOBRI07D/dEYRaB9ZZEs= -k8s.io/utils v0.0.0-20260626114624-be93311217bd/go.mod h1:xDxuJ0whA3d0I4mf/C4ppKHxXynQ+fxnkmQH0vTHnuk= +k8s.io/utils v0.0.0-20260707023825-cf1189d6abe3 h1:jVkFFVfXdXP74B/zbO3hM3hpSFD0xvhQ5U686DPurkE= +k8s.io/utils v0.0.0-20260707023825-cf1189d6abe3/go.mod h1:M2s5JB1lIYP3jzZdorPLHXIPJzt9vv2muW5a6L9DtNM= sigs.k8s.io/apiserver-network-proxy/konnectivity-client v0.36.0 h1:/YpDJ4vReG7ZmzSpBGxduXgywWkJU9zHubgJG03MT+Y= sigs.k8s.io/apiserver-network-proxy/konnectivity-client v0.36.0/go.mod h1:tJo1aepTXyR+8Xs3sUsGBDk4Ub2AM5dPAPKJx0mpm5c= -sigs.k8s.io/controller-runtime v0.25.0 h1:44KgRUPew331KSJpNu8zJow3iTR5W0p/SfrHdw3lV40= -sigs.k8s.io/controller-runtime v0.25.0/go.mod h1:4QqLdT6z/L6Olj8JJCtvztid4/fnIiYsfaTFScegctc= +sigs.k8s.io/controller-runtime v0.25.1 h1:BKgU9OeE8xv8EbbM8cY0NVzTQs35rokkdq1jh12fMb4= +sigs.k8s.io/controller-runtime v0.25.1/go.mod h1:4QqLdT6z/L6Olj8JJCtvztid4/fnIiYsfaTFScegctc= sigs.k8s.io/e2e-framework v0.7.0 h1:AHkySTC6MvnnMbVSxaO4z1m2MhQKNFP+2Ihs5pRNLlM= sigs.k8s.io/e2e-framework v0.7.0/go.mod h1:1ZgXkUSjmnf18/JgHZNEATWjv48O5lJm9aI1QIsRdbw= sigs.k8s.io/json v0.0.0-20250730193827-2d320260d730 h1:IpInykpT6ceI+QxKBbEflcR5EXP7sU1kvOlxwZh5txg= diff --git a/test/suite/case.go b/test/suite/case.go new file mode 100644 index 00000000..f43e5091 --- /dev/null +++ b/test/suite/case.go @@ -0,0 +1,68 @@ +package suite + +import ( + "testing" + + "github.com/multigres/testkit/ctrltest" +) + +// C is this operator's test context: the generic harness handle from +// ctrltest, plus the multigres vocabulary. +// +// A local type because Go cannot add methods to another package's, which is +// the point rather than a workaround. pkg/ctrltest is meant to be copied into +// other operators, so it carries assertions and harness pointers and knows +// nothing about shards or poolers; each consumer wraps it and hangs its own +// domain on the same receiver. Everything ctrltest offers is promoted, so +// c.NoError and c.WaitForClusterHealthy read alike at the call site. +type C struct { + *ctrltest.C +} + +// newCase opens a test context on its own namespace. +// +// Every test in this package should start with one. It allocates the +// namespace, activates the reconcile gate for it, and registers the failure +// dump, so a failing test prints the interleaved op log and reconcile records +// rather than only the assertion message. +func newCase(t *testing.T) *C { + t.Helper() + return &C{C: Suite.Case(t)} +} + +// newBareCase opens a test context with no namespace, for the tests in this +// package that are pure logic and never touch the cluster. +// +// identity_test.go is all of them: MembersOf and ShardPVCOf take objects and +// return answers. newCase would allocate a real namespace against envtest and +// register it at the reconcile gate for each one, which buys nothing. +func newBareCase(t *testing.T) *C { + t.Helper() + return &C{C: ctrltest.Bare(t)} +} + +// Sub binds this case to a subtest's T while keeping its namespace. Use it +// for a t.Run that asserts about objects the parent test created; use newCase +// for a subtest that wants a namespace of its own. +// +// It shadows the embedded ctrltest.C.Sub so that one name always hands back +// this package's C, with the multigres vocabulary still on it. Without the +// shadow a subtest would silently drop to the generic type and lose every +// method below. +func (c *C) Sub(t *testing.T) *C { + t.Helper() + return &C{C: c.C.Sub(t)} +} + +// Check returns a C whose assertions report and continue rather than abort, +// shadowed for the same reason as Sub: without it c.Check() hands back a +// *ctrltest.C and a collecting assertion silently loses every method below. +func (c *C) Check() *C { + return &C{C: c.C.Check()} +} + +// Assert at compile time that both shadows hand back this package's type. The +// regression they guard against is silent: dropping to *ctrltest.C still +// compiles at every existing call site, and only stops compiling once someone +// chains a multigres method off one of them. +var _ = func(c *C) (*C, *C) { return c.Check(), c.Sub(nil) } diff --git a/test/suite/fakes.go b/test/suite/fakes.go new file mode 100644 index 00000000..6b6691ff --- /dev/null +++ b/test/suite/fakes.go @@ -0,0 +1,265 @@ +package suite + +import ( + "context" + "fmt" + "sort" + "strings" + "sync" + "time" + + "github.com/multigres/multigres/go/common/rpcclient" + "github.com/multigres/multigres/go/common/topoclient" + "github.com/multigres/multigres/go/common/topoclient/memorytopo" + cm "github.com/multigres/multigres/go/pb/clustermetadata" + md "github.com/multigres/multigres/go/pb/multipoolermanagerdata" + corev1 "k8s.io/api/core/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" + "github.com/multigres/multigres-operator/pkg/util/metadata" +) + +// defaultSimCell is the cell every fixture uses. The topo store needs its cells +// declared up front, so a test using a different cell name needs this widened. +const defaultSimCell = "zone-a" + +// topoRegistry hands out one in-memory topology store per namespace. +// +// Per namespace rather than one shared store, because namespace-per-test is the +// suite's isolation boundary and a single store would let one test's cluster +// see another's multipoolers. Both CreateTopoStore seams resolve through here: +// the shard's carries the object, the cluster's carries only a DNS address, so +// that one recovers the namespace by parsing it. +type topoRegistry struct { + ctx context.Context + + mu sync.Mutex + stores map[string]topoclient.Store + facts map[string]*memorytopo.Factory +} + +func newTopoRegistry(ctx context.Context) *topoRegistry { + return &topoRegistry{ + ctx: ctx, + stores: map[string]topoclient.Store{}, + facts: map[string]*memorytopo.Factory{}, + } +} + +// Store returns the namespace's store, creating it on first use. +func (r *topoRegistry) Store(ns string) topoclient.Store { + r.mu.Lock() + defer r.mu.Unlock() + if s, ok := r.stores[ns]; ok { + return s + } + store, factory := memorytopo.NewServerAndFactory(r.ctx, defaultSimCell) + r.stores[ns] = store + r.facts[ns] = factory + return store +} + +func (r *topoRegistry) client(ns string) (topoclient.Store, error) { + r.Store(ns) + r.mu.Lock() + factory := r.facts[ns] + r.mu.Unlock() + return topoclient.NewWithFactory( + factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), + ), nil +} + +// ForShard is ShardReconciler.CreateTopoStore. +func (r *topoRegistry) ForShard(shard *Shard) (topoclient.Store, error) { + return r.client(shard.Namespace) +} + +// ForClusterRef is MultigresClusterReconciler.CreateTopoStore. The ref carries +// no namespace, only the Service address the cluster controller built as +// "-global-topo..svc:2379", so the namespace comes back out +// of the address. Left unstubbed, this seam dials a real etcd and the cluster +// never reaches TopologyReady. +func (r *topoRegistry) ForClusterRef( + ref multigresv1alpha1.GlobalTopoServerRef, +) (topoclient.Store, error) { + ns, err := namespaceFromTopoAddress(ref.Address) + if err != nil { + return nil, err + } + return r.client(ns) +} + +func namespaceFromTopoAddress(address string) (string, error) { + parts := strings.Split(address, ".") + if len(parts) < 2 || parts[1] == "" { + return "", fmt.Errorf("cannot derive namespace from topo address %q", address) + } + return parts[1], nil +} + +// poolerSim is the data plane: it registers a multipooler per pool pod in that +// namespace's topology store and answers a healthy Status RPC for each. +// +// Without it the shard controller stalls at PostureConsistent=Unknown +// (AwaitingPoolerRegistration) and never reaches a terminal state. +// +// Deliberately generous: every pod is healthy, always, and the lowest-numbered +// pod is primary. It models no ordering, no failure, and no latency, so a test +// that needs any of those needs a better fake than this one. +type poolerSim struct { + c client.Client + rpc *rpcclient.FakeClient + topo *topoRegistry + interval time.Duration + + mu sync.Mutex + registered map[string]bool + held map[string]bool +} + +// HoldRegistrations stops this fake registering any *new* pooler in ns until +// the returned function is called. Poolers already registered keep answering. +// +// It exists to make a race deterministic instead of sampled. The defect it was +// built for needs a shard to converge having seen fewer poolers than pods, +// which happens on its own only when registration loses a race against the +// last reconcile: about half the time, measured. A test that waits for that by +// chance detects a regression about half the time too, which is what the +// twelve-attempt statistical pin it replaced was paying for. +// +// Holding lets a test construct the precondition on purpose: hold, scale up, +// wait until the shard has demonstrably converged short, then release. One +// attempt, and the regression either survives the release or it does not. +func (p *poolerSim) HoldRegistrations(ns string) func() { + p.mu.Lock() + if p.held == nil { + p.held = map[string]bool{} + } + p.held[ns] = true + p.mu.Unlock() + + var once sync.Once + return func() { + once.Do(func() { + p.mu.Lock() + delete(p.held, ns) + p.mu.Unlock() + }) + } +} + +func (p *poolerSim) run(ctx context.Context) { + p.registered = map[string]bool{} + t := time.NewTicker(p.interval) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + p.tick(ctx) + } + } +} + +func (p *poolerSim) tick(ctx context.Context) { + pods := &corev1.PodList{} + if err := p.c.List(ctx, pods, + client.MatchingLabels{metadata.LabelAppComponent: shardcontroller.PoolComponentName}, + ); err != nil { + return + } + byNamespace := map[string][]*corev1.Pod{} + for i := range pods.Items { + pod := &pods.Items[i] + if !pod.DeletionTimestamp.IsZero() { + continue + } + byNamespace[pod.Namespace] = append(byNamespace[pod.Namespace], pod) + } + for ns, group := range byNamespace { + p.tickNamespace(ctx, ns, group) + } +} + +func (p *poolerSim) tickNamespace(ctx context.Context, ns string, pods []*corev1.Pod) { + sort.Slice(pods, func(i, j int) bool { return pods[i].Name < pods[j].Name }) + + ids := make([]*cm.ID, 0, len(pods)) + for _, pod := range pods { + ids = append(ids, &cm.ID{ + Cell: pod.Labels[metadata.LabelMultigresCell], + Name: pod.Name, + }) + } + if len(ids) == 0 { + return + } + leader := ids[0] + rule := &cm.ShardRule{ + RuleNumber: &cm.RuleNumber{CoordinatorTerm: 2}, + LeaderId: leader, + CohortMembers: ids, + DurabilityPolicy: topoclient.AtLeastN(1), + } + store := p.topo.Store(ns) + + for i, pod := range pods { + id := ids[i] + role := cm.RoutingRole_ROUTING_ROLE_REPLICA + resp := &md.StatusResponse{ + Status: &md.Status{ + IsInitialized: true, + PostgresReady: true, + PostgresStatus: md.PostgresStatus_POSTGRES_STATUS_STANDBY, + }, + AvailabilityStatus: &cm.AvailabilityStatus{ + CohortEligibilityStatus: &cm.CohortEligibilityStatus{ + Signal: cm.CohortEligibilitySignal_COHORT_ELIGIBILITY_SIGNAL_ELIGIBLE, + }, + }, + ConsensusStatus: &cm.ConsensusStatus{ + Id: id, + CurrentPosition: &cm.PoolerPosition{Position: &cm.RulePosition{Decision: rule}}, + }, + } + if id.Name == leader.Name { + role = cm.RoutingRole_ROUTING_ROLE_PRIMARY + resp.Status.PostgresStatus = md.PostgresStatus_POSTGRES_STATUS_PRIMARY + resp.Status.PrimaryStatus = &md.PrimaryStatus{ + Ready: true, + ConnectedFollowers: ids[1:], + } + } + p.rpc.SetStatusResponse(topoclient.ComponentIDString(id), resp) + + key := ns + "/" + pod.Name + p.mu.Lock() + already := p.registered[key] + heldBack := p.held[ns] + p.mu.Unlock() + // A held namespace still gets its Status RPC answered above, so pods + // already registered stay healthy and the shard keeps converging. Only + // the new registration waits, which is the whole point. + if already || heldBack { + continue + } + pooler := &cm.Multipooler{ + Id: id, + Hostname: pod.Name, + ShardKey: &cm.ShardKey{ + Database: pod.Labels[metadata.LabelMultigresDatabase], + TableGroup: pod.Labels[metadata.LabelMultigresTableGroup], + Shard: pod.Labels[metadata.LabelMultigresShard], + }, + RoutingState: &cm.RoutingState{Role: role}, + } + if err := store.RegisterMultipooler(ctx, pooler, true); err == nil { + p.mu.Lock() + p.registered[key] = true + p.mu.Unlock() + } + } +} diff --git a/test/suite/fixture.go b/test/suite/fixture.go new file mode 100644 index 00000000..440b017a --- /dev/null +++ b/test/suite/fixture.go @@ -0,0 +1,113 @@ +package suite + +import ( + "fmt" + "time" + + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" +) + +// The name of a Secret object, not a credential. gosec matches on the string +// value rather than on how it is used, so the suppression has to be explicit. +// +//nolint:gosec // G101: this is a Secret's name; the password itself is set below +const adminSecretName = "multigres-admin-password" + +// MinimalCluster creates the equivalent of config/samples/minimal.yaml in ns, +// along with the password Secret it references, and returns it. +// +// Built in Go rather than read from the sample, because the e2e loader +// (framework.MustLoadCluster) is behind //go:build e2e and this package +// deliberately has no build tag. +// +// The Secret is intentionally unlabelled, matching what a user would create. +// Under the production cache config that makes it invisible to the cached +// client, so this fixture also exercises why the reconcilers hold an APIReader. +func (c *C) MinimalCluster(name string) *MultigresCluster { + c.Helper() + return c.newCluster(name) +} + +// newCluster creates the admin Secret every cluster in this package +// references, then the cluster itself, applying any spec adjustments in +// between. +// +// The four fixtures here were identical for twenty lines and diverged only +// at Spec.Databases, which is the argument for this existing: the cost of +// the copy grows with the test count rather than being a debt that stays +// fixed, and the next person writing a scenario test copies whichever +// fixture they happened to read. +// +// The adjustment is a callback rather than a returned unsaved object so +// that creating the cluster cannot be forgotten. A fixture that built an +// object and never persisted it would leave the test waiting on a +// convergence that had no reason to start. +// +// The Secret is intentionally unlabelled, matching what a user would create. +// Under the production cache config that makes it invisible to the cached +// client, so this fixture also exercises why the reconcilers hold an +// APIReader. +func (c *C) newCluster(name string, with ...func(*MultigresClusterSpec)) *MultigresCluster { + c.Helper() + + secret := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{Name: adminSecretName, Namespace: c.NS}, + StringData: map[string]string{"password": "postgres"}, + } + c.NoError(c.Create(secret), "create password secret") + + cluster := &MultigresCluster{ + ObjectMeta: metav1.ObjectMeta{Name: name, Namespace: c.NS}, + Spec: MultigresClusterSpec{ + PostgresPasswordSecretRef: PostgresPasswordSecretRef{ + Name: adminSecretName, + Key: "password", + }, + PVCDeletionPolicy: &PVCDeletionPolicy{ + WhenDeleted: multigresv1alpha1.DeletePVCRetentionPolicy, + WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, + }, + Cells: []CellConfig{ + {Name: defaultSimCell, ZoneID: "us-central1-a"}, + }, + }, + } + for _, adjust := range with { + adjust(&cluster.Spec) + } + c.NoError(c.Create(cluster), "create MultigresCluster") + return cluster +} + +// WaitForClusterHealthy blocks until the cluster reports PhaseHealthy. +// +// One method rather than the seven copies pass 1 left behind: two named +// helpers (waitForClusterHealthy in the thrash file and waitForHealthy in the +// transitions file, byte-identical to each other) and five inlined +// Eventually blocks. They were hard to see as duplicates while each was +// wrapped in its own t-and-namespace threading. +// +// It is the convergence check nearly every scenario test starts from, which +// the old comment on one of the copies said out loud without anyone acting +// on it. +// +// 30 seconds because that is what all seven used. It is a convergence wait +// for a whole cluster under five controllers, not a single object read, so it +// is deliberately far longer than any assertion budget. +func (c *C) WaitForClusterHealthy(cluster *MultigresCluster) { + c.Helper() + c.Eventually(30*time.Second, "cluster to report Healthy", func() error { + got := &MultigresCluster{} + if err := c.Get(client.ObjectKeyFromObject(cluster), got); err != nil { + return err + } + if got.Status.Phase != multigresv1alpha1.PhaseHealthy { + return fmt.Errorf("phase is %q", got.Status.Phase) + } + return nil + }) +} diff --git a/test/suite/golden_test.go b/test/suite/golden_test.go new file mode 100644 index 00000000..f68746ed --- /dev/null +++ b/test/suite/golden_test.go @@ -0,0 +1,67 @@ +package suite + +import ( + "testing" + + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + cellcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/cell" + "github.com/multigres/testkit/golden" +) + +// TestGoldenMultigatewayDeployment pilots golden.AssertYAML against +// BuildMultigatewayDeployment, called directly rather than through the +// running cell controller: the builder is exported and reachable from this +// package, so the pilot exercises a pure function with no generated fields to +// strip, rather than an object fetched from envtest. +// +// It builds its own scheme rather than reaching for Suite.Scheme: the +// builder call needs nothing envtest boots, and coupling to Suite would make +// this test pay for the whole suite's startup for no benefit. +// +// The fixture pins both Spec.Observability (via OTEL_EXPORTER_OTLP_ENDPOINT) +// and Spec.Images.Multigateway to fixed values: a golden over a builder's +// output must not depend on values that routine maintenance changes. +func TestGoldenMultigatewayDeployment(t *testing.T) { + // BuildMultigatewayDeployment resolves OTEL settings from the process + // environment when Spec.Observability is nil (as it is below), so the + // golden file is only stable once that read is pinned. "disabled" is the + // builder's own sentinel for suppressing every OTEL var. + t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "disabled") + c := newBareCase(t) + + scheme := runtime.NewScheme() + c.NoError(multigresv1alpha1.AddToScheme(scheme), "add to scheme") + + cell := &multigresv1alpha1.Cell{ + ObjectMeta: metav1.ObjectMeta{ + Name: "golden-cell", + Namespace: "default", + UID: "golden-cell-uid", + Labels: map[string]string{"multigres.com/cluster": "golden-cluster"}, + }, + Spec: multigresv1alpha1.CellSpec{ + Name: "zone1", + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "global-topo:2379", + RootPath: "/multigres/global", + Implementation: "etcd", + }, + LogLevels: multigresv1alpha1.ComponentLogLevels{ + Multigateway: "info", + }, + Images: multigresv1alpha1.CellImages{ + Multigateway: multigresv1alpha1.ImageRef( + "ghcr.io/multigres/multigres:golden-fixture", + ), + }, + }, + } + + got, err := cellcontroller.BuildMultigatewayDeployment(cell, scheme) + c.NoError(err, "BuildMultigatewayDeployment") + + golden.AssertYAML(t, got, "testdata/multigateway-deployment.golden.yaml") +} diff --git a/test/suite/identity.go b/test/suite/identity.go new file mode 100644 index 00000000..5112c0c6 --- /dev/null +++ b/test/suite/identity.go @@ -0,0 +1,98 @@ +package suite + +import ( + "context" + "fmt" + "sort" + "strings" + + "sigs.k8s.io/controller-runtime/pkg/client" + + shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" + "github.com/multigres/testkit/ctrltest" +) + +// Members is a snapshot of a Shard's pod roles, taken at the moment of a call. +type Members struct { + Primary string // pod name with role PRIMARY + Replicas []string // pod names with role REPLICA, sorted + Quarantined []string // pod names with role QUARANTINED, sorted +} + +// MembersOf reads the Shard's status.podRoles and classifies every pod. +// Error-returning because callers include KnownDefect bodies. +// +// The role set here is {PRIMARY, REPLICA, QUARANTINED}, not the +// {PRIMARY, REPLICA, DRAINED} the CRD doc comment on PodRoles still claims +// (api/v1alpha1/shard_types.go:329, stale since e3677f0). PodRoles has one +// writer, reconcile_data_plane.go, fed entirely by GetPoolerStatus in +// pkg/data-handler/topo/pooler.go, where roleName is one of exactly those +// three literals. DRAINED cannot be produced by the real operator today, so +// it falls to the default arm below like any other unrecognized value. +func MembersOf(ctx context.Context, c client.Client, key client.ObjectKey) (Members, error) { + shard := &Shard{} + if err := c.Get(ctx, key, shard); err != nil { + return Members{}, fmt.Errorf("get shard %s: %w", key, err) + } + + roles := shard.Status.PodRoles + if len(roles) == 0 { + return Members{}, fmt.Errorf("shard %s: status.podRoles is empty", key) + } + + var primaries, replicas, quarantined []string + for pod, role := range roles { + switch role { + case "PRIMARY": + primaries = append(primaries, pod) + case "REPLICA": + replicas = append(replicas, pod) + case "QUARANTINED": + quarantined = append(quarantined, pod) + default: + return Members{}, fmt.Errorf( + "shard %s: pod %s has unrecognized role %q", key, pod, role, + ) + } + } + + switch len(primaries) { + case 0: + return Members{}, fmt.Errorf("shard %s: no pod has role PRIMARY", key) + case 1: + default: + sort.Strings(primaries) + return Members{}, fmt.Errorf( + "shard %s: more than one pod has role PRIMARY: %s", + key, strings.Join(primaries, ", "), + ) + } + + sort.Strings(replicas) + sort.Strings(quarantined) + + return Members{ + Primary: primaries[0], + Replicas: replicas, + Quarantined: quarantined, + }, nil +} + +// ShardPVCOf returns the PVC bound by the named pool pod's data volume. +// +// Pool pods only. The toposerver controller declares its own +// DataVolumeName ("data", in its statefulset builder), so this returns a +// not-found error for a toposerver pod rather than that pod's data PVC. A +// loud error rather than a wrong answer, but the restriction is not +// visible in the signature. +// +// A pool pod can carry a second PVC-backed volume (the filesystem backup +// volume, when shard.Spec.Backup.Type is Filesystem; see +// buildSharedBackupVolume in the shard controller), which is why the +// underlying lookup selects by volume name rather than requiring the pod +// to have exactly one PVC volume. +func ShardPVCOf( + ctx context.Context, c client.Client, ns, pod string, +) (string, error) { + return ctrltest.PVCOf(ctx, c, ns, pod, shardcontroller.DataVolumeName) +} diff --git a/test/suite/identity_test.go b/test/suite/identity_test.go new file mode 100644 index 00000000..3be6874b --- /dev/null +++ b/test/suite/identity_test.go @@ -0,0 +1,239 @@ +package suite + +import ( + "testing" + + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" +) + +// identityFakeClient builds a fake client that knows about both the +// operator's CRDs and core types, since a Shard's status and a Pod's +// volumes both need to round-trip through it. +func identityFakeClient(objs ...client.Object) client.Client { + scheme := runtime.NewScheme() + if err := multigresv1alpha1.AddToScheme(scheme); err != nil { + panic(err) + } + if err := corev1.AddToScheme(scheme); err != nil { + panic(err) + } + return fake.NewClientBuilder(). + WithScheme(scheme). + WithObjects(objs...). + WithStatusSubresource(&Shard{}). + Build() +} + +func shardWithRoles(ns, name string, roles map[string]string) *Shard { + return &Shard{ + ObjectMeta: metav1.ObjectMeta{Namespace: ns, Name: name}, + Status: multigresv1alpha1.ShardStatus{ + PodRoles: roles, + }, + } +} + +func TestIdentityMembersOfClassifiesAndSortsReplicas(t *testing.T) { + c := newBareCase(t) + shard := shardWithRoles("ns1", "shard1", map[string]string{ + "pod-primary": "PRIMARY", + "pod-replica-b": "REPLICA", + "pod-replica-a": "REPLICA", + }) + fc := identityFakeClient(shard) + + members, err := MembersOf(c.Context(), fc, client.ObjectKeyFromObject(shard)) + c.NoError(err, "MembersOf") + + c.Check().Eq("pod-primary", members.Primary, "Primary") + want := []string{"pod-replica-a", "pod-replica-b"} + gotSorted := len(members.Replicas) == len(want) && + members.Replicas[0] == want[0] && members.Replicas[1] == want[1] + c.Check().True(gotSorted, "Replicas = %v, want %v (sorted)", members.Replicas, want) +} + +func TestIdentityMembersOfSeparatesQuarantinedFromReplicas(t *testing.T) { + c := newBareCase(t) + shard := shardWithRoles("ns1", "shard1", map[string]string{ + "pod-primary": "PRIMARY", + "pod-replica": "REPLICA", + "pod-quarantined": "QUARANTINED", + }) + fc := identityFakeClient(shard) + + members, err := MembersOf(c.Context(), fc, client.ObjectKeyFromObject(shard)) + c.NoError(err, "MembersOf") + + c.Check().EqDiff([]string{"pod-replica"}, members.Replicas, "Replicas") + c.Check().EqDiff([]string{"pod-quarantined"}, members.Quarantined, "Quarantined") + c.NotContains(members.Replicas, "pod-quarantined", "quarantined pod leaked into Replicas") +} + +// TestIdentityMembersOfUnrecognizedRoleErrors pins the guard that catches any +// role value outside {PRIMARY, REPLICA, QUARANTINED}, the set the operator's +// single writer (pkg/data-handler/topo/pooler.go) can actually produce. DRAINED +// is deliberately used as the unknown value here: the CRD doc comment on +// PodRoles still lists it, but commit e3677f0 removed it from the writer, so +// it is exactly the stale value a careless "known roles" list would still +// accept. +func TestIdentityMembersOfUnrecognizedRoleErrors(t *testing.T) { + c := newBareCase(t) + shard := shardWithRoles("ns1", "shard1", map[string]string{ + "pod-primary": "PRIMARY", + "pod-drained": "DRAINED", + }) + fc := identityFakeClient(shard) + + _, err := MembersOf(c.Context(), fc, client.ObjectKeyFromObject(shard)) + c.Error(err, "want an error for an unrecognized role") + + c.Check(). + ErrorContains(err, "unrecognized role", "want the error to call DRAINED an unrecognized role") + c.Check().ErrorContains(err, "DRAINED", "want the error to name DRAINED specifically") +} + +// TestIdentityMembersOfErrorCases pins the three distinct error cases the +// brief calls out. An absent primary and a duplicated primary are both bugs +// a test should be able to pin, but they are different bugs, so a test that +// only checked "err != nil" could not tell them apart. +func TestIdentityMembersOfErrorCases(t *testing.T) { + c := newBareCase(t) + cases := []struct { + name string + roles map[string]string + wantErrs []string + }{ + { + name: "empty PodRoles", + roles: map[string]string{}, + wantErrs: []string{"empty"}, + }, + { + name: "no primary", + roles: map[string]string{ + "pod-a": "REPLICA", + "pod-b": "QUARANTINED", + }, + wantErrs: []string{"no pod has role PRIMARY"}, + }, + { + name: "two primaries", + roles: map[string]string{ + "pod-a": "PRIMARY", + "pod-b": "PRIMARY", + }, + wantErrs: []string{"more than one pod has role PRIMARY", "pod-a", "pod-b"}, + }, + } + + // Each case's error must be distinguishable from the other two, so this + // collects every message actually produced and cross-checks that no + // case's message satisfies another case's expectation. + messages := make(map[string]string, len(cases)) + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + c := newBareCase(t) + shard := shardWithRoles("ns1", "shard1", tc.roles) + fc := identityFakeClient(shard) + + _, err := MembersOf(c.Context(), fc, client.ObjectKeyFromObject(shard)) + c.Error(err, "want an error for %s", tc.name) + for _, want := range tc.wantErrs { + c.Check().ErrorContains(err, want) + } + messages[tc.name] = err.Error() + }) + } + + c.Check().NotEq(messages["empty PodRoles"], messages["no primary"], + "empty PodRoles and no-primary produced the same error message") + c.Check().NotEq(messages["no primary"], messages["two primaries"], + "no-primary and two-primaries produced the same error message") + c.Check().NotEq(messages["empty PodRoles"], messages["two primaries"], + "empty PodRoles and two-primaries produced the same error message") +} + +// identityPVCVolume is a thin copy of relations_test.go's pvcVolume. Kept +// separate rather than shared, since after Task 2.1 the original lives in +// another package. +func identityPVCVolume(name, claim string) corev1.Volume { + return corev1.Volume{ + Name: name, + VolumeSource: corev1.VolumeSource{ + PersistentVolumeClaim: &corev1.PersistentVolumeClaimVolumeSource{ + ClaimName: claim, + }, + }, + } +} + +// identityPodWithVolumes is a thin copy of relations_test.go's +// podWithVolumes. Kept separate rather than shared, since after Task 2.1 the +// original lives in another package. +func identityPodWithVolumes(ns, name string, volumes ...corev1.Volume) *corev1.Pod { + return &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{Namespace: ns, Name: name}, + Spec: corev1.PodSpec{ + Containers: []corev1.Container{{Name: "c", Image: "busybox"}}, + Volumes: volumes, + }, + } +} + +// TestIdentityShardPVCOfUsesTheShardDataVolumeName pins which volume name +// the wrapper supplies. Nothing else does: PVCOf's own tests pass a +// literal, so a wrapper handing over the wrong constant is invisible. +func TestIdentityShardPVCOfUsesTheShardDataVolumeName(t *testing.T) { + c := newBareCase(t) + p := identityPodWithVolumes( + "ns1", "pod-a", + identityPVCVolume(shardcontroller.DataVolumeName, "pod-a-data"), + identityPVCVolume("backup-data", "pod-a-backup"), + ) + fc := identityFakeClient(p) + + claim, err := ShardPVCOf(c.Context(), fc, "ns1", "pod-a") + c.NoError(err, "ShardPVCOf") + c.Check().Eq("pod-a-data", claim) +} + +// TestIdentityMembersOfSortsEnoughReplicasToCatchMapOrder pins the sorts in +// MembersOf. PodRoles is a Go map and Go randomises map iteration, so two +// replicas would agree with insertion order half the time and the assertion +// would be a coin flip. Five in reverse order leaves a 1-in-120 chance of +// passing against an unsorted implementation. +func TestIdentityMembersOfSortsEnoughReplicasToCatchMapOrder(t *testing.T) { + c := newBareCase(t) + shard := &Shard{ + ObjectMeta: metav1.ObjectMeta{Name: "shard-0", Namespace: "ns"}, + Status: multigresv1alpha1.ShardStatus{PodRoles: map[string]string{ + "pool-primary": "PRIMARY", + "pool-e": "REPLICA", + "pool-d": "REPLICA", + "pool-c": "REPLICA", + "pool-b": "REPLICA", + "pool-a": "REPLICA", + "quar-e": "QUARANTINED", + "quar-d": "QUARANTINED", + "quar-c": "QUARANTINED", + "quar-b": "QUARANTINED", + "quar-a": "QUARANTINED", + }}, + } + fc := identityFakeClient(shard) + + got, err := MembersOf(c.Context(), fc, client.ObjectKeyFromObject(shard)) + c.NoError(err, "MembersOf") + wantReplicas := []string{"pool-a", "pool-b", "pool-c", "pool-d", "pool-e"} + wantQuarantined := []string{"quar-a", "quar-b", "quar-c", "quar-d", "quar-e"} + c.Check().EqDiff(wantReplicas, got.Replicas, "Replicas") + c.Check().EqDiff(wantQuarantined, got.Quarantined, "Quarantined") +} diff --git a/test/suite/main_test.go b/test/suite/main_test.go new file mode 100644 index 00000000..214d750e --- /dev/null +++ b/test/suite/main_test.go @@ -0,0 +1,66 @@ +package suite + +import ( + "fmt" + "os" + "runtime/pprof" + "testing" + + "go.uber.org/goleak" +) + +// TestMain boots the suite once, runs the package, tears it down, and only then +// checks for leaked goroutines. +// +// Deliberately not goleak.VerifyTestMain: that calls m.Run() and checks +// immediately afterwards, with no hook in between. Since envtest and the +// manager live for the whole package rather than for one test, the check would +// run while both were still up and report the entire manager as leaked. +// +// The ignore list is empty on purpose. A clean start/stop leaks nothing +// measurable (verified across 16 cycles before this suite existed), so every +// future entry should be justified against a real stack trace rather than +// pre-loaded against suspects. A pre-loaded ignore is a permanent hole. +func TestMain(m *testing.M) { + s, teardown, err := Boot() + if err != nil { + fmt.Fprintf(os.Stderr, "suite boot failed: %v\n", err) + os.Exit(1) + } + Suite = s + + code := m.Run() + + if err := teardown(); err != nil { + fmt.Fprintf(os.Stderr, "suite teardown failed: %v\n", err) + if code == 0 { + code = 1 + } + } + + // Only leak-check a passing run: a failed test may have left its own + // goroutines behind, and reporting those on top of a real failure buries + // the real failure. + if code == 0 { + if err := goleak.Find(); err != nil { + fmt.Fprintf(os.Stderr, "goroutine leak after suite teardown: %v\n", err) + dumpGoroutineLeakProfile() + code = 1 + } + } + os.Exit(code) +} + +// dumpGoroutineLeakProfile writes the runtime's own leak profile when the +// toolchain has one. It returns nil before Go 1.27, so this is a no-op today +// and becomes a diagnostic on the next toolchain bump, with no build tag. +// +// Count() does not run the leak detector but WriteTo() does, so the profile has +// to be driven through WriteTo and the count read afterwards, never before. +func dumpGoroutineLeakProfile() { + p := pprof.Lookup("goroutineleak") + if p == nil { + return + } + _ = p.WriteTo(os.Stderr, 1) +} diff --git a/test/suite/scenario_deletion_test.go b/test/suite/scenario_deletion_test.go new file mode 100644 index 00000000..2d2a872d --- /dev/null +++ b/test/suite/scenario_deletion_test.go @@ -0,0 +1,295 @@ +package suite + +import ( + "strings" + "testing" + "time" + + apierrors "k8s.io/apimachinery/pkg/api/errors" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + multigresclustercontroller "github.com/multigres/multigres-operator/pkg/cluster-handler/controller/multigrescluster" + tablegroupcontroller "github.com/multigres/multigres-operator/pkg/cluster-handler/controller/tablegroup" + "github.com/multigres/testkit/ctrltest" +) + +// TestReadyForDeletionProtocol asserts the three-hop ReadyForDeletion protocol +// found in pkg/resource-handler/controller/shard/reconcile_deletion.go, +// pkg/cluster-handler/controller/tablegroup/tablegroup_controller.go and +// pkg/cluster-handler/controller/multigrescluster/reconcile_databases.go: a +// Shard sets multigresv1alpha1.ConditionReadyForDeletion on itself once every +// pool pod has drained, its parent TableGroup sets the same condition on +// itself once every child Shard has, and MultigresCluster deletes a TableGroup +// only once that TableGroup reports it. +// +// This is orphan pruning, not whole-cluster teardown. Deleting a +// MultigresCluster goes through MultigresClusterReconciler.handleDeletion, +// which lists and Deletes its Cells and TableGroups directly and never +// consults ConditionReadyForDeletion at all; that path was added in +// da7d639177b0 as a narrower fix scoped explicitly to "orphan pruning" (its +// own commit message), for the case where a TableGroup or Cell falls out of a +// still-live cluster's spec and must drain before it is safe to remove. A +// whole-cluster delete has nothing left to protect by draining, so it tears +// down directly instead. The three-hop protocol is therefore reachable only +// through that narrower path, which this test drives directly: it builds an +// orphan TableGroup (one MultigresCluster.Spec.Databases entry will never +// name) with multigresclustercontroller.BuildTableGroup, the same builder +// production code uses, so multigrescluster's reconcileDatabases treats it +// exactly as it would treat a TableGroup a user just removed from spec. +func TestReadyForDeletionProtocol(t *testing.T) { + c := newCase(t) + ns := c.NS + cluster := c.MinimalCluster("scenario-del") + + c.WaitForClusterHealthy(cluster) + + // Copy the real TableGroup's resolved GlobalTopoServer ref and component + // Images rather than re-deriving them: both are resolved in-memory once + // per MultigresCluster reconcile (globalTopoRef by the unexported + // globalTopoRef method, Images by resolveImages) and never written back to + // MultigresCluster.Spec, so the cluster object this test already holds + // still has them blank. Re-deriving either by hand risks building an + // orphan that fails validation or that the topology store does not + // recognize, for reasons unrelated to what this test asserts. + realTGs := &TableGroupList{} + c.NoError(c.List(realTGs, + client.MatchingLabels{"multigres.com/cluster": cluster.Name}), + "list tablegroups") + c.Len(realTGs.Items, 1, + "want exactly one TableGroup before introducing an orphan") + globalTopoRef := realTGs.Items[0].Spec.GlobalTopoServer + cluster.Spec.Images = multigresv1alpha1.ClusterImages{ + Multiorch: realTGs.Items[0].Spec.Images.Multiorch, + Multipooler: realTGs.Items[0].Spec.Images.Multipooler, + Postgres: realTGs.Items[0].Spec.Images.Postgres, + ImagePullPolicy: realTGs.Items[0].Spec.Images.ImagePullPolicy, + ImagePullSecrets: realTGs.Items[0].Spec.Images.ImagePullSecrets, + } + + // The orphan's one shard has no pools and zero Multiorch replicas, so it + // creates no Pods. That keeps the pod drain state machine, which is its + // own protocol, out of this test's way: with zero pods, + // ShardReconciler.handlePendingDeletion takes the "no pods" branch and + // sets ConditionReadyForDeletion on its very first pass. What this test + // asserts is the condition handoff between the three controllers, not + // how long draining a pod takes. + // + // Multiorch.Cells is set explicitly because getMultiorchCells + // (shard_controller.go) falls back to the union of pool cells when it is + // empty, and errors out when that is empty too; with no pools, leaving + // Cells unset turns every normal (non-deletion) reconcile of this Shard + // into a reconcile error, which is retried on controller-runtime's own + // backoff rather than this suite's compressed one and made the protocol's + // timing depend on that backoff instead of on the protocol. + // + // Opened before the orphan TableGroup exists, not after: five reconcilers + // run concurrently (scenario_fanout_test.go documents the same discipline for + // cursors, for the same reason), and multigrescluster can notice and + // annotate an orphan TableGroup within the reconcile pass that follows its + // creation. A stream opened even one line later could start listening + // after that annotation, and the condition flips it drives, have already + // landed, which is exactly the 4ms-window flake this replaces. + st := c.Watch(&TableGroupList{}, &ShardList{}) + + orphanTG, err := multigresclustercontroller.BuildTableGroup( + cluster, + DatabaseConfig{Name: "postgres"}, + &TableGroupConfig{Name: "orphan"}, + []multigresv1alpha1.ShardResolvedSpec{{ + Name: "0-inf", + Multiorch: multigresv1alpha1.MultiorchSpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(0))}, + Cells: []CellName{defaultSimCell}, + }, + Pools: map[PoolName]PoolSpec{}, + }}, + globalTopoRef, + Suite.Scheme, + ) + c.NoError(err, "build orphan tablegroup") + c.NoError(c.Create(orphanTG), "create orphan tablegroup") + orphanKey := client.ObjectKeyFromObject(orphanTG) + + // Pre-create the child Shard, from the now-persisted orphanTG (so its + // owner reference carries a real UID), immediately after the TableGroup + // rather than leaving TableGroupReconciler to create it on its first + // normal pass. Without this, a real race exists: if multigrescluster + // annotates the brand-new TableGroup with AnnotationPendingDeletion before + // TableGroupReconciler has ever run stepApplyDesiredShards for it, + // handlePendingDeletion lists zero child Shards and reports + // ReadyForDeletion vacuously, having consulted nothing. Creating the Shard + // here, keyed identically to what stepApplyDesiredShards would build, + // guarantees the child the protocol is supposed to drain exists before + // either controller's watch can fire. AlreadyExists is fine rather than + // fatal: it means TableGroupReconciler's own first pass won the race and + // applied this same Shard first, which is the other safe ordering. + shardCR, err := tablegroupcontroller.BuildShard( + orphanTG, + &orphanTG.Spec.Shards[0], + Suite.Scheme, + ) + c.NoError(err, "build orphan shard") + if err := c.Create(shardCR); err != nil && + !apierrors.IsAlreadyExists(err) { + c.NoError(err, "create orphan shard") + } + orphanShardKey := client.ObjectKeyFromObject(shardCR) + clusterKey := client.ObjectKeyFromObject(cluster) + + // Corroborating evidence, independent of how fast the three controllers + // converge: the interceptor records every reconcile pass, including ones + // that wrote nothing, so a pass that asked to be woken again in 5s while + // something was still pending is permanent history even if the object it + // was about is deleted moments later. Both waits are scoped by object key, + // so neither can be satisfied by the pre-existing healthy + // TableGroup/Shard's own unrelated reconciles. + // + // What this actually proves is narrower than it looks: both + // TableGroupReconciler.handlePendingDeletion and reconcileDatabases also + // take a 5s-requeue path the first time they see an orphan (setting the + // PendingDeletion annotation itself sets `allReady`/`pendingDeletion` and + // requeues), so mutating away only the later + // `meta.IsStatusConditionTrue(...ConditionReadyForDeletion)` guard still + // leaves that earlier requeue in place and these two waits keep passing - + // verified by making exactly that mutation in each function and watching + // these two lines stay green while the poll below caught it instead. What + // these two waits do rule out is a version that deletes an orphan in the + // very same pass that first notices it, with no intervening wait at all. + Suite.Reconciles.WaitForRequeue(t, "tablegroup", orphanKey, 5*time.Second, 10*time.Second) + Suite.Reconciles.WaitForRequeue( + t, + "multigrescluster", + clusterKey, + 5*time.Second, + 10*time.Second, + ) + + // Primary evidence for the specific guard on each hop, read off the event + // stream rather than sampled by polling: the write-up measured the + // TableGroup's parent deleting it 4.3ms after ReadyForDeletion is set, + // against a 20ms poll, and showed no poll interval fixes that because one + // sample already costs about as long as the state persists. The stream is + // push rather than sample, so it sees the transition however briefly it + // held. sawX latches record having observed each condition true at least + // once, so that reaching a later state (the TableGroup gone) without ever + // having latched an earlier one (its own condition, or its child Shard's) + // is still caught even if every step landed inside one 50ms reorder + // window, or before this loop's first read. + // + // Mutation verified: deleting the + // `if !meta.IsStatusConditionTrue(s.Status.Conditions, + // multigresv1alpha1.ConditionReadyForDeletion) { allReady = false }` guard + // in TableGroupReconciler.handlePendingDeletion, or the equivalent guard + // over item.Status.Conditions in reconcileDatabases, each independently + // makes this fail (tried one at a time): the TableGroup reports + // ReadyForDeletion, or is deleted, before its Shard's own condition is + // ever observed true. + sawShardReady := false + sawTableGroupReady := false + tgGone := false + deadline := time.Now().Add(30 * time.Second) + for !tgGone { + ev, err := st.Next(time.Until(deadline)) + c.NoError(err, + "waiting for the protocol to reach Shard ready, then TableGroup ready, "+ + "then TableGroup deleted, in that order") + switch { + case ev.Key == orphanShardKey && ev.Kind == "Shard": + if conditionSetTrue(ev, string(multigresv1alpha1.ConditionReadyForDeletion)) { + sawShardReady = true + } + case ev.Key == orphanKey && ev.Kind == "TableGroup": + if ev.Type == "deleted" { + tgGone = true + continue + } + if conditionSetTrue(ev, string(multigresv1alpha1.ConditionReadyForDeletion)) { + c.True(sawShardReady, + "TableGroup %s reported ReadyForDeletion before Shard %s ever did", + orphanKey.Name, orphanShardKey.Name) + sawTableGroupReady = true + } + } + } + // A relist can win the race against the last blocked read: Next selects + // over the event channel and the failure channel, and Go picks at random + // when both are ready. If it hands back the deletion, the loop exits and + // the entry guard never runs again, so the terminal error would go + // unobserved and the assertions below would blame the operator for history + // the harness lost. + c.NoError(st.Terminal(), + "the event stream failed, so every assertion over it is void") + + c.True(sawShardReady, + "TableGroup %s was deleted before Shard %s ever reported ReadyForDeletion", + orphanKey.Name, orphanShardKey.Name) + c.True(sawTableGroupReady, + "TableGroup %s was deleted before it ever reported ReadyForDeletion", + orphanKey.Name) + + // Attribution check: confirm the delete that made the TableGroup + // disappear was actually issued by multigrescluster, using a static scan + // of the completed op log rather than a live cursor wait. A live + // CursorFor(ns, "multigrescluster") wait was not usable for any step + // above: a single MultigresCluster reconcile pass unconditionally + // re-applies the healthy default Cell and TableGroup before ever reaching + // the orphan-pruning loop, so the next op in that scope is legitimately + // something else almost every time, and WaitForNext does not skip ahead + // to find a match (that is WaitForMatching, which this suite marks as an + // escape hatch not to be reached for). The op log has already stopped + // growing with respect to this object by the time we reach this check, so + // a static scan carries none of Cursor's live-ordering caveats. + deletedByCluster := false + for _, op := range Suite.Ops.OpsInNamespace(ns) { + if op.Controller == "multigrescluster" && op.Verb == "delete" && + ctrltest.KindSuffix(op.Kind) == "TableGroup" && op.Key.Name == orphanKey.Name { + deletedByCluster = true + break + } + } + c.True(deletedByCluster, + "TableGroup %s disappeared without a recorded delete from multigrescluster", + orphanKey.Name) +} + +// conditionSetTrue reports whether ev records a status.conditions entry whose +// type field arrived at conditionType, with that same entry's status field +// arrived at "True", in the same event. +// +// Changed and Transitions carry a diff, not the object's current state, so +// the type and the status of one array slot have to be read out of the same +// event to know which condition moved; a status flip alone does not say +// which condition it belongs to. Requiring the type to change in the same +// event rather than looking it up separately is sound here because every +// condition this test watches for (Shard and TableGroup's own +// ConditionReadyForDeletion) is only ever set once, straight to True: neither +// controller ever writes it False first, so its type always appears fresh +// alongside the status that makes it true. +// +// If that premise ever breaks, and a controller sets the condition False +// before True, this helper stops recognising the transition and the test fails +// with an ordering complaint against the operator rather than against itself. +// So a sudden "reported ReadyForDeletion before X ever did" failure is worth +// checking here second: the Cell controller already writes a condition False +// first, so the premise holds by habit rather than by rule. +func conditionSetTrue(ev ctrltest.Event, conditionType string) bool { + wantType := `"` + conditionType + `"` + for path, transition := range ev.Transitions { + // The status.conditions prefix is checked as well as the leaf, so a + // future array of objects carrying both a type and a status field + // cannot start feeding this helper silently. No such array exists on + // either status today; the guard is what keeps the doc comment above + // true rather than merely true for now. + if !strings.HasPrefix(path, "status.conditions[") || + !strings.HasSuffix(path, "].type") || transition.To != wantType { + continue + } + statusPath := strings.TrimSuffix(path, "type") + "status" + if status, ok := ev.Transitions[statusPath]; ok && status.To == `"True"` { + return true + } + } + return false +} diff --git a/test/suite/scenario_fanout_test.go b/test/suite/scenario_fanout_test.go new file mode 100644 index 00000000..d0578763 --- /dev/null +++ b/test/suite/scenario_fanout_test.go @@ -0,0 +1,159 @@ +package suite + +import ( + "testing" + "time" + + "k8s.io/utils/ptr" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + "github.com/multigres/multigres-operator/pkg/util/name" + "github.com/multigres/testkit/ctrltest" +) + +// fanoutCluster builds a MultigresCluster with one cell and one database, +// table group and shard: enough for the cluster controller to fan out to +// every child kind it owns (TopoServer, Cell, TableGroup) and for the +// TableGroup controller it creates to fan out to a Shard of its own. +func (c *C) fanoutCluster(clusterName string) *MultigresCluster { + c.Helper() + return c.newCluster(clusterName, func(s *MultigresClusterSpec) { + s.Databases = []DatabaseConfig{ + { + Name: "postgres", + Default: true, + TableGroups: []TableGroupConfig{ + { + Name: "default", + Default: true, + Shards: []ShardConfig{{ + Name: "0-inf", + Spec: &ShardInlineSpec{ + Multiorch: multigresv1alpha1.MultiorchSpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(1)), + }, + }, + Pools: map[PoolName]PoolSpec{ + "primary": { + ReplicasPerCell: ptr.To(int32(1)), + Type: "readWrite", + Cells: []CellName{ + defaultSimCell, + }, + }, + }, + }, + }}, + }, + }, + }, + } + }) +} + +// findOp returns the index of the first op in ops matching kind and name. +// batch has already been validated as an exact multiset by WaitForAll, so a +// miss here would mean this helper's own matching is wrong, not that the op +// is absent. +func (c *C) findOp(ops []ctrltest.Op, kind, objName string) int { + c.Helper() + for i, op := range ops { + if ctrltest.KindSuffix(op.Kind) == kind && op.Key.Name == objName { + return i + } + } + c.Fatalf("no %s %q op among %v", kind, objName, ops) + return -1 +} + +// TestClusterFanOutSequence asserts that creating a MultigresCluster fans out +// to TopoServer, Cell and TableGroup attributed to the multigrescluster +// controller, in that relative order, and that the second-order fan-out to +// Shard is attributed to the tablegroup controller rather than to +// multigrescluster. +// +// Ordering choice: multigrescluster_controller.go calls +// reconcileGlobalComponents, then reconcileCells, then (after +// reconcileTopology, which makes no Kubernetes writes here) reconcileDatabases, +// as sequential statements inside one Reconcile call. That is a genuine, +// code-level guarantee, so TopoServer < Cell < TableGroup is asserted below. +// Nothing is asserted about the relative order of the Multiadmin/MultiadminWeb +// writes reconcileGlobalComponents also makes: they are unconditional (there +// is no way to disable them from the spec) and sit between the TopoServer and +// Cell writes, but their order relative to each other, or to TopoServer and +// Cell, is not what this test is about and is left unspecified. +// +// Matcher choice: WaitForNext cannot express this, because those +// Multiadmin/MultiadminWeb writes are real, deterministic, and land between +// TopoServer and Cell, so a strict next-op chain from TopoServer would hit a +// Deployment patch instead of Cell. WaitForAll is the right tool instead: one +// call consumes the whole first pass as a multiset, which (a) still proves +// each of TopoServer/Cell/TableGroup was written by multigrescluster and +// nothing else was (in particular, that multigrescluster never itself writes +// a Shard), and (b) hands back the ops in recorded order, which is what the +// ordering check below reads. WaitForMatching was avoided entirely: it would +// let the assertion silently skip past a misordered write instead of failing +// on it, which is exactly what an ordering test must not do. +func TestClusterFanOutSequence(t *testing.T) { + c := newCase(t) + const clusterName = "fanout" + + // Cursors are opened before the cluster exists, not after: five + // reconcilers run concurrently (one goroutine each, per suite.go), so a + // cursor opened even one line late can start its scan after another + // controller's reaction to the same write has already landed, and then + // miss the very op it was meant to catch. + clusterCur := c.Cursor("multigrescluster") + tgCur := c.Cursor("tablegroup") + + c.fanoutCluster(clusterName) + + topoServerName := clusterName + "-global-topo" + cellName := name.JoinWithConstraints(name.DefaultConstraints, clusterName, defaultSimCell) + tableGroupName := name.JoinWithConstraints( + name.DefaultConstraints, clusterName, "postgres", "default", + ) + shardName := name.JoinWithConstraints( + name.DefaultConstraints, clusterName, "postgres", "default", "0-inf", + ) + + batch := clusterCur.WaitForAll(t, []ctrltest.Expect{ + // ensureClusterFinalizer, then resolveImages recording the default + // image set: both patches of the cluster object itself, before any + // child is touched. + ctrltest.ExpectPatch("MultigresCluster", clusterName), + ctrltest.ExpectPatch("MultigresCluster", clusterName), + ctrltest.ExpectPatch("TopoServer", topoServerName), + ctrltest.ExpectPatch("Deployment", clusterName+"-multiadmin"), + ctrltest.ExpectPatch("Service", clusterName+"-multiadmin"), + ctrltest.ExpectPatch("Deployment", clusterName+"-multiadmin-web"), + ctrltest.ExpectPatch("Service", clusterName+"-multiadmin-web"), + ctrltest.ExpectPatch("Service", clusterName+"-multigateway"), + ctrltest.ExpectPatch("Service", clusterName+"-multigateway-replica"), + ctrltest.ExpectPatch("Cell", cellName), + ctrltest.ExpectPatch("TableGroup", tableGroupName), + ctrltest.ExpectStatusPatch("MultigresCluster", clusterName), + }, 30*time.Second) + + topoIdx := c.findOp(batch, "TopoServer", topoServerName) + cellIdx := c.findOp(batch, "Cell", cellName) + tgIdx := c.findOp(batch, "TableGroup", tableGroupName) + c.True(topoIdx < cellIdx, + "expected TopoServer before Cell in the multigrescluster controller's "+ + "writes, got positions %d, %d in %v", topoIdx, cellIdx, batch) + c.True(cellIdx < tgIdx, + "expected Cell before TableGroup in the multigrescluster controller's "+ + "writes, got positions %d, %d in %v", cellIdx, tgIdx, batch) + + // Second-order fan-out: the TableGroup controller, not the cluster + // controller, creates the Shard. Applying the desired Shard is the first + // write TableGroupReconciler makes (stepListChildShards only reads), so + // WaitForNext is sound here without any of the batching above. The + // attribution claim itself comes from the cursor's scope, not from a + // separate check: tgCur can only ever return an op whose Controller is + // "tablegroup" (see Recorder.firstMatchFrom), and the WaitForAll batch + // above already proved multigrescluster wrote no Shard of its own, since + // one would have shown up there as an unexpected op. + tgCur.WaitForNext(t, "patch", "Shard", shardName, 30*time.Second) +} diff --git a/test/suite/scenario_race_test.go b/test/suite/scenario_race_test.go new file mode 100644 index 00000000..8ac8a766 --- /dev/null +++ b/test/suite/scenario_race_test.go @@ -0,0 +1,335 @@ +package suite + +import ( + "fmt" + "sort" + "strings" + "testing" + "time" + + "k8s.io/apimachinery/pkg/api/equality" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/structured-merge-diff/v6/fieldpath" + + "github.com/multigres/testkit/ctrltest" +) + +// shardProbeAnnotation is a key neither controller's applied payload ever +// mentions: tablegroup's BuildShard only sets an annotation map at all when the +// TableGroup carries a project-ref annotation, which MinimalCluster's fixture +// never does. Writing it from the test is therefore a mutation neither +// manager's SSA apply will contend with or revert, and it is the only way this +// test can put a fresh, known-real change on the Shard once the cluster has +// converged, since a converged tablegroup no longer gets re-triggered by its own +// stale watch. +const shardProbeAnnotation = "scenario-race-test.multigres.com/probe" + +// TestTwoControllersWriteOneShard names the two-writer relationship between the +// shard and tablegroup controllers on one object, which every other test in +// this package runs past without seeing: each of them boots the suite with +// every reconciler live, but none asserts anything about a Shard being written +// by both. +// +// The relationship was measured, not assumed: before a fix, the shard +// controller wrote its own status about ten times a second forever, and while +// that ran, the tablegroup controller issued 3,641 patches against the Shard in +// three minutes, every one a no-op. Fixing the hot loop did not remove the +// two-writer relationship, only the churn it used to cause, so it is still +// worth a name. +func TestTwoControllersWriteOneShard(t *testing.T) { + c := newCase(t) + cluster := c.MinimalCluster("race") + + c.WaitForClusterHealthy(cluster) + key := c.shardKey() + + t.Run("shard and tablegroup both write the Shard", func(t *testing.T) { + c := c.Sub(t) + writers := map[string]bool{} + for _, op := range Suite.Ops.OpsInNamespace(c.NS) { + if ctrltest.KindSuffix(op.Kind) == "Shard" && op.Key.Name == key.Name { + writers[op.Controller] = true + } + } + for _, want := range []string{"shard", "tablegroup"} { + c.Check().True(writers[want], + "the %s controller has no recorded write to Shard %s; writers seen: %v", + want, key.Name, writers) + } + }) + + // Runs before the no-op check below, which mutates the Shard: a manager + // entry left over from that mutation would be a false-positive third writer + // and would blur what this is actually about, which is the two production + // managers. + t.Run("field ownership on the Shard is disjoint", func(t *testing.T) { + c := c.Sub(t) + // Several Shard status writes carry no field owner, so the API server + // attributes them to the manager that happens to be the process name, + // and that manager ends up co-owning fields the shard controller's own + // applier claims. Retire this pin with an explicit owner on every + // status write, and replace it with c.Empty on the conflicts. + conflicts := c.fieldOwnershipConflicts(key) + c.KnownDefect("MGO-SHARD-STATUS-WRITES-NO-FIELD-OWNER", func() error { + if len(conflicts) == 0 { + return nil + } + return fmt.Errorf( + "Shard %s has fields claimed by more than one field manager, "+ + "the same shape of defect as the status hot loop:\n %s", + key.Name, strings.Join(conflicts, "\n ")) + }) + }) + + t.Run("tablegroup's patches to the Shard never change it", func(t *testing.T) { + c.Sub(t).requireTableGroupPatchesAreNoOps(key) + }) +} + +// fieldOwnershipConflicts returns every field path on the Shard that more +// than one field manager claims, sorted. An empty result is the invariant, and +// it is the check that generalises: it would have caught the status hot-loop +// directly; the loop was two field managers each asserting a value for the +// same field, ObservedGeneration or a phase, and the API server obliging both. +// A dump of writes only shows that as churn; managedFields shows it as the +// overlap it is. +// +// Field ownership is expressed by the API server as a Set per manager, encoded +// in metadata.managedFields[].fieldsV1: this decodes each manager's Set with +// the same library the server itself uses and collects every field path that +// is a member of more than one. +func (c *C) fieldOwnershipConflicts(key client.ObjectKey) []string { + c.Helper() + + shard := &Shard{} + c.NoError(c.Get(key, shard), "get shard %s", key) + + claimants := map[string][]string{} + for _, mf := range shard.ManagedFields { + if mf.FieldsV1 == nil { + continue + } + set := &fieldpath.Set{} + c.NoError(set.FromJSON(mf.FieldsV1.GetRawReader()), + "decode managedFields for manager %s", mf.Manager) + set.Iterate(func(p fieldpath.Path) { + path := p.String() + claimants[path] = append(claimants[path], mf.Manager) + }) + } + + // A Shard with no decodable managedFields claim at all would make the + // disjointness check below pass having compared nothing, which is the one + // way this assertion can go green without the invariant holding. Nothing in + // this function's own reasoning rules that out, so it is asserted here + // rather than inherited from the writer check above. + c.True(len(claimants) > 0, + "Shard %s yielded no decodable managedFields claims, so field ownership "+ + "was not checked at all; %d managedFields entries were present", + key.Name, len(shard.ManagedFields)) + + var conflicts []string + for path, managers := range claimants { + distinct := map[string]bool{} + for _, m := range managers { + distinct[m] = true + } + if len(distinct) <= 1 { + continue + } + names := make([]string, 0, len(distinct)) + for m := range distinct { + names = append(names, m) + } + sort.Strings(names) + conflicts = append(conflicts, fmt.Sprintf("%s claimed by %v", path, names)) + } + + sort.Strings(conflicts) + return conflicts +} + +// requireTableGroupPatchesAreNoOps is assertion 2 from the brief, corrected to +// the harness as it exists now rather than as it was written: convergence is +// about 0.3s and requeues are compressed, so a fixed wall-clock window no +// longer separates from scheduling noise. It also has nothing to observe by the +// time a test could get around to sleeping: once the cluster is Healthy, +// tablegroup's own SSA patches to the Shard stop generating new watch events +// (a true no-op patch does not bump resourceVersion, so Owns(&Shard{}) never +// re-fires), so tablegroup simply stops reconciling and there is no ten-second +// window in which anything would happen anyway. +// +// So instead of waiting, this drives the relationship directly: it makes a +// small, known write of its own to the Shard (an annotation neither manager's +// apply payload mentions, see shardProbeAnnotation) to produce a fresh watch +// event, then waits for a COMPLETE tablegroup reconcile pass that both began +// after that write and wrote this Shard (see awaitTableGroupPassAfter), read +// from Suite.Reconciles rather than the raw op cursor. That matters: a +// reconcile record is only appended once the whole pass has returned, so +// waiting on it cannot observe half a pass the way a plain op-log cursor can, +// where a later step of the very pass that just satisfied the wait +// (tablegroup's own status-patch, a step after the Shard patch) can still land +// and be mistaken for the next pass's write. +// +// The no-op check itself compares content, not resourceVersion. The shard +// controller reconciles this same object roughly every clamp interval even at +// steady state (see ctrltest.RequeueClamp), so by the time tablegroup's pass +// has been observed, some other write to the Shard has almost always also +// landed in the same rough window; attributing a resourceVersion move to +// "whichever controller wrote most recently" is exactly as unsound as the +// wall-clock method this replaces; it was tried and produced a false pass +// under mutation (see the report). tablegroup's SSA apply payload is a +// complete statement of what it owns: spec, labels and ownerReferences, +// never status (see BuildShard), so comparing that payload's own fields +// before and after the pass answers the question directly and is immune to +// any concurrent, status-only write from the shard controller, no matter how +// often it fires. +// +// Repeated several times rather than once, since a single pass proves nothing +// about whether "never" holds. +func (c *C) requireTableGroupPatchesAreNoOps(key client.ObjectKey) { + c.Helper() + + const passesToObserve = 5 + for i := range passesToObserve { + before := &Shard{} + c.NoError(c.Get(key, before), "get shard %s", key) + + probedFrom := c.probeShard(key, i) + pass := awaitTableGroupPassAfter(c, c.NS, key, probedFrom) + + after := &Shard{} + c.NoError(c.Get(key, after), "get shard %s", key) + + c.True(equality.Semantic.DeepEqual(before.Spec, after.Spec), + "tablegroup's pass %s changed Shard %s's spec, the field surface its "+ + "SSA apply owns:\n before: %+v\n after: %+v", + pass, key.Name, before.Spec, after.Spec) + c.True( + equality.Semantic.DeepEqual(before.OwnerReferences, after.OwnerReferences), + "tablegroup's pass %s changed Shard %s's ownerReferences:\n before: %+v\n after: %+v", + pass, + key.Name, + before.OwnerReferences, + after.OwnerReferences, + ) + c.True(equality.Semantic.DeepEqual(before.Labels, after.Labels), + "tablegroup's pass %s changed Shard %s's labels:\n before: %+v\n after: %+v", + pass, key.Name, before.Labels, after.Labels) + } +} + +// awaitTableGroupPassAfter blocks until the tablegroup controller has completed +// a reconcile pass that started after notBefore and wrote the Shard at key, and +// returns that pass. +// +// The bar is Reconcile.Start measured against an instant captured BEFORE the +// probe write was issued, and both halves of that were arrived at by measuring +// a wrong version of it. +// +// Start, rather than a position in the log, because a record is appended when +// its pass RETURNS. A pass already in flight when the probe lands is therefore +// filed after the probe while having begun before it, so an index cursor +// admits it however freshly it was seeded. Five controllers converge this +// cluster before the first probe and leave dozens of finished passes behind, +// and a scan from index zero matched those exclusively: every iteration was +// satisfied by a pass that had started roughly half a second before the probe +// it was supposed to be reacting to. +// +// Before the write rather than after it, because a write becomes visible to +// watchers when the API server commits it, which is strictly before the +// client's own call returns. The gap is small but it is on the wrong side: the +// reacting tablegroup pass starts within a few tenths of a millisecond of the +// probe Patch returning, and measurably often starts just before it. A mark +// taken after the write returned then rejects the very pass it is waiting for, +// and since the probe is the only thing that writes this Shard once the cluster +// has converged, no later pass arrives to replace it. The wait cannot then +// succeed at any timeout, which is what made it fail about one run in three. A +// mark taken before the write has no such edge, because no pass can react to a +// write that has not been issued yet. +// +// The cost of moving the mark earlier is one API round trip of slack, in which a +// pass that did not see the probe would be accepted. That is bounded by a round +// trip instead of by the whole log, and at steady state nothing but this test +// writes the Shard, so the only candidate is a second pass caused by the +// previous iteration's probe. That pass has already been observed to completion +// before this iteration's mark is taken. +// +// Rescanning the namespace's whole log on each poll, rather than carrying a +// cursor across polls or across iterations, is deliberate for a related reason: +// the log is ordered by completion, not by start, so an index cursor and a +// start-time bar disagree about which entries are still candidates, and the +// cursor is the one that can step over the pass being waited for. The scan is +// bounded by one namespace's history and runs at most once per poll, so the +// cost of being right here is nothing worth optimising. +// +// No bookkeeping is needed to stop one iteration matching an earlier +// iteration's pass: that pass had already returned, and so had already started, +// before this iteration's mark was taken. +func awaitTableGroupPassAfter( + c *C, + ns string, + key client.ObjectKey, + notBefore time.Time, +) ctrltest.Reconcile { + c.Helper() + + var pass ctrltest.Reconcile + what := fmt.Sprintf("a tablegroup pass on Shard %s beginning after the probe", key.Name) + c.Eventually(10*time.Second, what, func() error { + for _, r := range Suite.Reconciles.InNamespace(ns) { + if r.Controller != "tablegroup" || !r.Start.After(notBefore) { + continue + } + for _, op := range Suite.Reconciles.Ops(r) { + if ctrltest.KindSuffix(op.Kind) == "Shard" && op.Key.Name == key.Name { + pass = r + return nil + } + } + } + return fmt.Errorf( + "no tablegroup pass touching Shard %s has begun since the probe write", + key.Name, + ) + }) + return pass +} + +// probeShard sets a test-owned annotation on the Shard to a fresh value and +// returns the instant just before that write was issued, which is the latest +// mark a pass reacting to it is guaranteed to start after. +// +// Returning the instant the write completed is the obvious choice and is the +// wrong one: the API server commits the write and dispatches the watch event +// before the client's Patch call returns, so the reacting pass is often already +// running by then. See awaitTableGroupPassAfter. +// +// Neither field manager's apply payload mentions the annotation (tablegroup's +// BuildShard only sets an annotation map at all when the TableGroup carries a +// project-ref annotation, which MinimalCluster's fixture never does), so it is +// a change neither manager's SSA apply will contend with or revert. It is a +// merge patch rather than a full Update, so it carries no resourceVersion +// precondition and cannot spuriously conflict with either controller's own +// concurrent write to the same object. Writing it is the only way this test +// can put a fresh, real change on the Shard once the cluster has converged, +// since a converged tablegroup no longer gets re-triggered by its own stale +// watch (see requireTableGroupPatchesAreNoOps). +func (c *C) probeShard(key client.ObjectKey, seq int) time.Time { + c.Helper() + shard := &Shard{} + c.NoError(c.Get(key, shard), "get shard %s", key) + base := shard.DeepCopy() + if shard.Annotations == nil { + shard.Annotations = map[string]string{} + } + shard.Annotations[shardProbeAnnotation] = fmt.Sprintf("%d", seq) + + issued := time.Now() + c.NoError( + c.Patch(shard, client.MergeFrom(base)), + "annotate shard %s", + key, + ) + return issued +} diff --git a/test/suite/scenario_selector_impostor_test.go b/test/suite/scenario_selector_impostor_test.go new file mode 100644 index 00000000..2ae0c28e --- /dev/null +++ b/test/suite/scenario_selector_impostor_test.go @@ -0,0 +1,643 @@ +// Selector inventory for the multigrescluster and shard controllers: every +// List call in pkg/cluster-handler/controller/multigrescluster/ and +// pkg/resource-handler/controller/shard/ that carries a label selector, the +// selector itself, and what this sweep found or judged about it. +// +// 35 List call sites in the two packages, 12 in multigrescluster and 23 in +// shard, of which 32 carry a label selector; the three that do not are +// recorded below as out of scope rather than omitted, so that the count can +// be rederived from this block. Line numbers point at the List( call, not at +// the MatchingLabels argument. Cross-checked for the other spellings a +// selector can take (MatchingLabelsSelector, HasLabels, a raw +// ListOptions{LabelSelector}): neither package uses any of them. +// +// multigrescluster (pkg/cluster-handler/controller/multigrescluster/): +// +// - reconcile_cells.go:23 CellList {cluster} +// Feeds the active/orphan diff for cells removed from spec. The eventual +// delete is gated behind the AnnotationPendingDeletion + ConditionReady- +// ForDeletion handshake (reconcile_cells.go:92-125), not a raw sweep. +// SAFE (condition-gated). Not tested here. +// +// - reconcile_databases.go:23 TableGroupList {cluster} +// Same shape as reconcile_cells.go:23, for TableGroups. SAFE (condition- +// gated); the protocol itself is already covered by +// TestReadyForDeletionProtocol in scenario_deletion_test.go. +// +// - reconcile_topology.go:245 CellList {cluster} +// Read-only: collects names of cells pending deletion, for topology +// pruning. Never mutates or deletes anything itself. SAFE. +// +// - reconcile_global.go:113 TopoServerList {cluster} +// CONFIRMED LIVE DEFECT (Defect 2 in +// tasks/multigres-operator-bugs-found-by-the-suite.md). No component +// filter, no owner-reference check: when global topology is external, +// every TopoServer carrying the cluster label is deleted, including a +// cell's own local TopoServer (component "local-topo"), which this +// selector cannot tell apart from the managed global one (component +// "global-topo"). TESTED below: +// TestSelectorImpostorGlobalTopoPruneDeletesCellOwnedLocalTopoServer. +// +// - multigrescluster_controller.go:401 CellList {cluster}, in +// handleDeletion (whole-cluster teardown). Raw delete, no owner-ref +// check. Reachable only while the cluster itself is being deleted, and +// an impostor sharing the label would first be picked up by +// reconcile_cells.go's own condition-gated path above (the cell +// controller reconciles any Cell object regardless of who owns it), +// which confounds an isolated impostor test for this exact line. +// JUDGED UNREACHABLE IN ISOLATION within this pass; not tested. See the +// task report for the reasoning in full. +// +// - multigrescluster_controller.go:419 TableGroupList {cluster}, in +// handleDeletion. Same shape and same confound as line 401, via +// reconcile_databases.go:23's condition-gated path. Not tested. +// +// - multigrescluster_controller.go:439 PersistentVolumeClaimList +// {cluster, component=toposerver}, in handleDeletion. Raw delete, no +// owner-ref check, but the function's own comment states the intent: +// these PVCs "may outlive their TopoServer when a cluster switches from +// managed to external topology," i.e. an unowned PVC with this label +// pair is the expected steady state this code exists to clean up, not +// an anomaly. SAFE BY DESIGN for this task's "should do nothing to an +// object it does not own" heuristic, because the whole point here is +// that ownership cannot be established for what it is meant to sweep. +// Not pinned as a defect; the design's blast radius (anything bearing +// these two labels, from any source, is eligible) is flagged in the +// report rather than pinned as a KnownDefect. +// +// This is SAFE while shard_controller.go:589 below is a DEFECT on what +// looks like the same evidence, an unowned PVC being acted on, and the +// two verdicts are worth reading together rather than one at a time. +// The difference is documented intent, which is the only thing that can +// separate them: :439 says an unowned PVC carrying these labels is +// precisely what it exists to sweep, whereas :589's own doc comment +// (shard_controller.go:571-574) and its call site +// (shard_controller.go:425-426) both say it exists to fix up ownerRefs +// on the shard's own PVCs across a mid-lifecycle policy change. Acting +// on a stranger is the job in one and an accident in the other. +// +// - multigrescluster_controller.go:603 MultigresClusterList, InNamespace +// only. Not a label selector (a map function for CoreTemplate/ +// CellTemplate/ShardTemplate change fanout). Out of scope. +// +// - certificate.go:394 (via pkg/util/certs.List) no label selector, +// InNamespace only. The eventual delete (certs.Prune) checks +// OwnedBy(cert, ownerUID) before deleting anything: the one place in +// this controller that already does what reconcile_global.go:113 does +// not. SAFE, and the contrast is the original bug write-up's own point. +// +// - status.go:125, :150, :220 CellList / TableGroupList / TopoServerList +// {cluster}. Purely read, to aggregate MultigresCluster.Status; nothing +// is ever mutated or deleted at these call sites. SAFE for this task's +// "does the controller act on it" question. A mislabelled object here +// would only ever skew the cluster's own reported status, a different +// risk this task's Quiet()-shaped assertion cannot express and which is +// not assessed here. +// +// shard (pkg/resource-handler/controller/shard/): +// +// - shard_controller.go:589 (reconcilePVCOwnerRefs) PersistentVolumeClaimList +// {cluster, database, tablegroup, shard} (no pool, no component). NEW +// DEFECT found by this sweep: the selector is the shard's four identity +// keys and nothing else, so it cannot tell the shard's own PVC from any +// other object carrying the same identity, and the only test applied to +// a match before adoption is whether it already carries a ref with this +// shard's UID (shard_controller.go:625-631). A PVC belonging to nobody +// fails that test, so it is adopted: SetControllerReference + Patch, +// whenever the shard's effective PVCDeletionPolicy resolves to Delete. +// +// There is a real ownership check here, and naming it correctly matters +// because it bounds the defect. ctrl.SetControllerReference returns +// AlreadyOwnedError when the object already carries a different +// controller ownerRef (controller-runtime v0.25.0, +// pkg/controller/controllerutil/controllerutil.go:97-99), so this code +// does not adopt a PVC that already belongs to someone else. What it +// adopts is a PVC with no controller owner at all. TESTED below: +// TestSelectorImpostorShardOwnerRefReconcileAdoptsUnrelatedPVC. +// +// - reconcile_deletion.go:57 DeploymentList {cluster, database, +// tablegroup, shard}, in the Shard's own handleDeletion. Raw delete, no +// owner-ref check. Same family as shard_controller.go:589. Not +// independently tested given this pass's budget. +// +// - reconcile_deletion.go:90 PodList same 4-key selector, same +// function. Raw delete, no owner-ref check. Lower interference risk +// than the multigrescluster Cell/TableGroup case above, since a Shard's +// own teardown is not cascaded through another controller's graceful +// orphan protocol. Not tested here. +// +// - reconcile_deletion.go:166 (cleanupShardPVCs) PersistentVolumeClaimList +// same 4-key selector. Every match is marked orphan or deleted with no +// owner-ref check, gated only by shardPVCShouldBeCleaned's policy read. +// Same family as shard_controller.go:589 at a different lifecycle +// point. Not independently tested. +// +// - reconcile_deletion.go:252 (handlePendingDeletion) PodList same +// 4-key selector. Runs the drain state machine (initiateDrain / +// clearDrainAnnotations / Delete) against any match once the Shard +// itself carries the PendingDeletion annotation. Same family; combining +// it with a graceful shard-level orphan flow adds the same entanglement +// seen in the multigrescluster Cell/TableGroup case. Not tested. +// +// - reconcile_data_plane.go:290, :367 PodList same 4-key selector. +// Read-only relative to the pods themselves (feeds +// shard.Status.PodRoles and posture.Evaluate). SAFE. +// +// - reconcile_data_plane.go:542 (reconcileDrainState) PodList same +// 4-key selector. MUTATES a matching pod (clears its drain annotations) +// when isDrainStale holds, which requires the pod's pool label to +// resolve to a real shard.Spec.Pools entry, a name that parses to an +// in-range replica ordinal, and a spec judged unchanged from desired. +// A real candidate, but reproducing that combination on a synthetic +// impostor is disproportionate for this pass; deferred. +// +// - reconcile_data_plane.go:673 (reconcilePoolerPrune) PodList same +// 4-key selector. The action it drives (topo.MarkDeadPoolers) writes to +// the fake topology store, not to the Kubernetes object, so this +// suite's k8s-event Stream cannot observe the outcome either way. +// UNTESTABLE WITH THIS HARNESS. +// +// - reconcile_quarantine.go:83 (reconcileQuarantineRemediation) PodList +// same 4-key selector. Deletes a pod and hard-deletes its PVC, but only +// for names the fake topology store reports as LIFECYCLE_QUARANTINED. +// Driving that requires reaching into the suite's internal topology +// registry; deferred. +// +// - disruption.go:30 PodList {cluster, database, tablegroup, shard, +// component=Pool} (adds the component key the multigrescluster prune +// lacks). Read-only (feeds canStartDisruption's decision). SAFE. +// +// - maintenance_surge.go:300, :338 PodList component+cell-scoped via +// shardPDBLabels/metadata.GetSelectorLabels. Read-only. SAFE. +// +// - postgres_config.go:266 ShardList, InNamespace only. Not a label +// selector (map function for ConfigMap-change fanout). Out of scope. +// +// - reconcile_shared_infra.go:373 PodList the PDB's own +// component+pool+cell selector. Read-only (sizes MinAvailable). SAFE. +// +// - reconcile_shared_infra.go:411 PodDisruptionBudgetList same PDB +// selector. Deletes an unmatched PDB, but only those that also pass +// metav1.IsControlledBy(pdb, shard): an explicit owner check right in +// the loop. SAFE, and the "done right" counterpart to +// shard_controller.go:589. +// +// - reconcile_pool_pods.go:50 PodList pool+cell-scoped +// (buildPoolLabelsWithCell). DEFECT of the same class as +// shard_controller.go:589, and the strongest untested candidate left in +// this inventory. The list populates existingPods +// (reconcile_pool_pods.go:69-73), which is then walked by name with no +// ownership check at all: +// +// Phase 0, syncDrainedLabels (:82, body :943-:974), iterates every map +// member and patches multigres.com/pod-role whenever +// resolvePodRole(shard, pod.Name) disagrees with the label the pod +// carries, so a pod this shard does not own has that label written or +// stripped. Phase 2, handleScaleDown (:118, body :499), classifies any +// member whose name does not parse as - (resolvePodIndex, +// :1240-:1250) or whose index is at or beyond effectiveReplicas as an +// extra pod (:537-:540), then drains (initiateDrain, :650) and deletes +// it (:567). A plausible impostor Pod carrying the pool and cell labels +// and any non-numeric name suffix is therefore drained and deleted. +// Secondary consequence at the same site: isPoolHealthy(existingPods, +// ...) (:585) counts an impostor as a pool member, so a non-ready one +// blocks legitimate scale-down of the real pool. +// +// NOT PINNED in this pass, and recorded here rather than left implied: +// a pin is a test, and this one needs a Pod-kind impostor against +// DataPlaneSim.tickPods, whose write set differs from tickPVCs's and +// has not been derived (see the TestSelectorImpostorShardOwnerRef... +// caveat below). Naming it SAFE, as an earlier revision of this block +// did, was the error worth correcting: an unexamined site is a gap, and +// a gap signed SAFE is worse than one left open. +// +// - reconcile_pool_pods.go:60 PersistentVolumeClaimList pool+cell- +// scoped. SAFE, but name-keyed rather than read-only, which is the +// accurate justification: pvcutil.ClearOrphan (:204) and +// expandPVCIfNeeded (:216) both write to list members, and what makes +// them safe is that every access is existingPVCs[pvcName] where +// pvcName comes from BuildPoolDataPVCName, a deterministic desired +// name. An impostor under any other name is never indexed, so it is +// genuinely untouched. Unlike :50 above, which walks the map itself. +// +// - reconcile_pool_pods.go:1116 PersistentVolumeClaimList pool+cell- +// scoped, counts non-orphan PVCs. Read-only. SAFE. +// +// - reload.go:69 PodList {cluster, database, tablegroup, shard, +// component=Pool}. Read-only (feeds the reload decision). SAFE. +// +// - status.go:221 PodList pool+cell-scoped (buildPoolLabelsWithCell). +// Read-only (status aggregation). SAFE. +// +// - status.go:365 PodList buildMultiorchLabelsWithCell selector. +// Read-only (crash-loop detection for status). SAFE. +// +// - reconcile_readiness.go:33 PodList {cluster, database, tablegroup, +// shard, component=Pool}. Read-only (readiness aggregation). SAFE. + +package suite + +import ( + "errors" + "fmt" + "testing" + "time" + + corev1 "k8s.io/api/core/v1" + "k8s.io/apimachinery/pkg/api/resource" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + topopkg "github.com/multigres/multigres-operator/pkg/data-handler/topo" + "github.com/multigres/multigres-operator/pkg/util/metadata" + "github.com/multigres/multigres-operator/pkg/util/name" + "github.com/multigres/testkit/ctrltest" +) + +// errLocalTopoDeleted is what the TopoServer script's invariant returns when +// it sees the deletion that test pins. +// +// A sentinel rather than a match on the harness's message text, because +// KnownDefect reads any non-nil error as "the pinned defect is still live". +// This pin's script can produce several other errors that are not that +// deletion, and every one of them would otherwise keep the pin green: a +// legitimate toposerver status write that the step's declaration failed to +// account for, a step that timed out because that declaration has drifted from +// what the controller now writes. Those are facts about this test, not about +// the operator, so the check body has to be able to tell them apart, and it +// cannot do that by reading a string the harness is free to reformat. +var errLocalTopoDeleted = errors.New( + "the cell's own local TopoServer was deleted", +) + +// selectorImpostorNudgeAnnotation is a key no controller's applied payload +// ever mentions, following the same reasoning as shardProbeAnnotation in +// scenario_race_test.go: tablegroup's BuildShard sets an annotation map on a Shard +// only when its TableGroup carries a project-ref annotation +// (pkg/cluster-handler/controller/tablegroup/builders.go:37-48), which +// MinimalCluster's fixture never does, so writing this key is a mutation +// neither manager's SSA apply contends with or reverts. +// +// It has to enqueue the Shard to be useful, and it does: the shard +// controller's For(&Shard{}) (shard_controller.go:693) +// carries no predicate, so a metadata-only patch is a reconcile trigger. That +// is the whole reason the key exists, since a converged Shard has nothing left +// to re-trigger it and reconcilePVCOwnerRefs only looks at an impostor on a +// pass that actually runs. +const selectorImpostorNudgeAnnotation = "scenario-selector-impostor-test.multigres.com/nudge" + +// externalGlobalTopoCluster creates a MultigresCluster whose global topology is +// external and whose one cell manages its own local TopoServer, the exact +// combination tasks/multigres-operator-external-topo-deletes-local.md +// reproduced on a live cluster: it is what makes reconcileGlobalTopoServer's +// desired-is-nil branch run on every reconcile, while still giving the cell +// controller a local TopoServer of its own to keep reapplying. +// +// It returns an error rather than calling t.Fatalf as MinimalCluster does, +// because its caller runs it inside a Script step's do: a Fatalf there would +// Goexit out of the middle of a script, whereas Script.TryStep routes a failing +// do to the test's own Fatalf and, crucially, never lets it reach KnownDefect +// as though it were evidence about the operator. +// +// The password Secret is created here rather than shared with MinimalCluster +// because the two fixtures differ in every other field; what is worth keeping +// in step is the deliberate choice to leave it unlabelled, which is what a user +// would create and what makes the reconcilers' APIReader necessary. +func (c *C) externalGlobalTopoCluster(clusterName string) error { + c.Helper() + + secret := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{Name: adminSecretName, Namespace: c.NS}, + StringData: map[string]string{"password": "postgres"}, + } + if err := c.Create(secret); err != nil { + return fmt.Errorf("create password secret: %w", err) + } + + cluster := &MultigresCluster{ + ObjectMeta: metav1.ObjectMeta{Name: clusterName, Namespace: c.NS}, + Spec: MultigresClusterSpec{ + PostgresPasswordSecretRef: PostgresPasswordSecretRef{ + Name: adminSecretName, + Key: "password", + }, + PVCDeletionPolicy: &PVCDeletionPolicy{ + WhenDeleted: multigresv1alpha1.DeletePVCRetentionPolicy, + WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, + }, + GlobalTopoServer: &multigresv1alpha1.GlobalTopoServerSpec{ + External: &multigresv1alpha1.ExternalTopoServerSpec{ + Endpoints: []multigresv1alpha1.EndpointUrl{ + "https://external-topo.invalid:2379", + }, + }, + }, + Cells: []CellConfig{ + { + Name: defaultSimCell, + ZoneID: "us-central1-a", + Spec: &multigresv1alpha1.CellInlineSpec{ + LocalTopoServer: &multigresv1alpha1.LocalTopoServerSpec{ + Etcd: &multigresv1alpha1.EtcdSpec{ + Replicas: ptr.To(int32(1)), + }, + }, + }, + }, + }, + }, + } + if err := c.Create(cluster); err != nil { + return fmt.Errorf("create MultigresCluster: %w", err) + } + return nil +} + +// TestSelectorImpostorGlobalTopoPruneDeletesCellOwnedLocalTopoServer pins +// Defect 2 from tasks/multigres-operator-bugs-found-by-the-suite.md: +// reconcile_global.go:113 lists TopoServers by cluster label alone (no +// component filter, no owner-reference check) and deletes every match +// whenever global topology is external. A cell's own local TopoServer +// carries that same cluster label, so it is not this controller's to +// manage, but the selector cannot tell the difference. +// +// The impostor here is not synthetic: it is the real local TopoServer the cell +// controller legitimately creates and keeps reapplying, which is exactly what +// makes the write-up call this "a permanent create/delete loop" rather than a +// one-off. That permanence is also what shapes the script below, because it +// rules out the move every other test in this suite makes first. An object +// caught in a create/delete loop never settles, so there is no converged +// namespace to open a watch onto: neither RequireQuiescent nor a poll for a +// stable TopoServer can be used here, and waiting a fixed margin for the +// toposerver controller to stop writing is a guess at how long another actor's +// work takes, which is the thing this suite exists to refuse. +// +// So the watch opens first, on an empty namespace, and the fixture is created +// inside the script's own step. Every legitimate write the toposerver +// controller then makes to the object is named as a permitted change, which +// leaves the deletion as the one event nothing accounts for, and leaves nothing +// to wait out. +func TestSelectorImpostorGlobalTopoPruneDeletesCellOwnedLocalTopoServer(t *testing.T) { + c := newCase(t) + const clusterName = "ext-global-topo" + cellResourceName := name.JoinWithConstraints( + name.DefaultConstraints, clusterName, string(defaultSimCell), + ) + localTopoName := topopkg.ManagedLocalTopoServerName(cellResourceName) + + script := c.NewScript(&TopoServerList{}) + + // The deletion is this test's entire claim, so it is asserted directly + // rather than inferred from being whatever event no step happened to + // permit. Script.TryStep and Script.TryFinish both run the invariants + // against an event before comparing it to the permitted set + // (script.go:224, :243, :294), so the delete is reported as this violation + // wherever it lands: while the step is still waiting on a status write, + // inside its settle window, or inside Finish's horizon. Combined with the + // sentinel above, that is what makes the pin's evidence the deletion on + // every run instead of whichever event happened to arrive first. + script.Invariant( + "the cell's own local TopoServer is never deleted", + func(ev ctrltest.Event) error { + if ev.Type == "deleted" && ev.Kind == "TopoServer" && + ev.Key.Name == localTopoName { + return errLocalTopoDeleted + } + return nil + }, + ) + + c.KnownDefect( + "pkg/cluster-handler/controller/multigrescluster/reconcile_global.go:113 "+ + "(external-global-topo prune selector has no component filter or "+ + "owner-reference check, so it also deletes a cell's own local TopoServer)", + func() error { + stepErr := script.TryStep( + "the cell controller creates its own local TopoServer and the "+ + "toposerver controller settles its status on it", + func() error { + return c.externalGlobalTopoCluster(clusterName) + }, + // The toposerver controller's whole settling sequence on a + // TopoServer it has just been handed, measured over four runs + // against a cluster whose global topology is managed so this + // prune never fires, which is the one way to observe what the + // object does when it is left alone: the first condition, then + // the client and peer endpoints once the etcd StatefulSet + // exists, then Ready once DataPlaneSim has ticked that + // StatefulSet ready. Three writes, in that order, then quiet + // indefinitely. + // + // The paths are the narrowest that pick out one write each. + // status.conditions[0] belongs only to the first and + // status.clientService only to the second; the third's paths + // are a subset of the second's, so it is matched by + // elimination, which is what assignEvents does a search rather + // than a greedy first match for. + // + // No ordering is declared between them even though one was + // observed, because ordering is opt-in for changes that follow + // from the code and nothing here needs it: the deletion is + // caught by the invariant above, not by an order violation. + ctrltest.Added("TopoServer", localTopoName), + ctrltest.Changed("TopoServer", localTopoName, "status.conditions[0].type"), + ctrltest.Changed("TopoServer", localTopoName, + "status.clientService", "status.peerService"), + ctrltest.Changed("TopoServer", localTopoName, "status.phase"), + ) + // TryFinish runs whatever the step returned, so that the script + // ends with Finish exactly once and the end-of-script backstop is + // satisfied on every path through this body. While the defect is + // live the step returns long before the object has finished + // settling, so the end of the script still has to be closed. + // + // The horizon is not load-bearing in either direction, which is the + // point of choosing it freely. While the defect is live nothing + // depends on it: the delete lands about 15ms after the create, + // inside the step. Once the defect is fixed this is the only window + // left in which a later prune pass could still be caught, and a + // longer horizon can only refuse more events, never permit one. + finishErr := script.TryFinish(10 * time.Second) + + switch { + case errors.Is(stepErr, errLocalTopoDeleted): + return stepErr + case errors.Is(finishErr, errLocalTopoDeleted): + return finishErr + case stepErr != nil: + // Anything else is this test's own declaration or pacing + // rather than evidence about the operator, and a pin that + // confirmed on it would survive the fix it is supposed to + // expire on. Fatalf is the right side of the line + // Script.fatalf already draws for the same reason. + c.Fatalf("the script's declaration of the toposerver "+ + "controller's settling sequence did not hold, which is a "+ + "fact about this test rather than about the prune it pins: %v", + stepErr) + case finishErr != nil: + c.Fatalf("the script's end was not quiet, and not because of "+ + "the deletion this test pins, which is a fact about this "+ + "test rather than about the prune: %v", finishErr) + } + return nil + }, + ) +} + +// TestSelectorImpostorShardOwnerRefReconcileAdoptsUnrelatedPVC pins a new +// defect found by this sweep: reconcilePVCOwnerRefs +// (shard_controller.go:589) lists PersistentVolumeClaims by the shard's four +// identity labels alone (cluster, database, tablegroup, shard; no pool, no +// component), so it cannot tell the shard's own PVC from any other object +// carrying the same identity, and the only test it applies to a match before +// adopting it is whether that match already carries a ref with this shard's +// UID. A PVC belonging to nobody fails that test and is adopted, whenever the +// shard's effective PVCDeletionPolicy resolves to Delete. +// +// The bound on the defect is worth stating precisely, because it is what a fix +// has to be aimed at. There is an ownership check in this path: +// ctrl.SetControllerReference returns AlreadyOwnedError when the object already +// carries a different controller ownerRef, so this code does not take a PVC +// that belongs to someone else. What it takes is a PVC with no controller owner +// at all. The missing check is not "does this belong to somebody else" but +// "does this belong to me", and the selector is what cannot answer it. +// +// The impostor is a bare PersistentVolumeClaim carrying just those four +// labels and no pool label, the same shape shard_controller.go's own +// "shared backup PVC" branch expects, and no owner reference at all: the +// shape a PVC left behind by some other process, or a since-recreated +// resource under the same identity, would plausibly have. +// +// RequireQuiescent runs before the script's watch opens so the baseline it +// replays is the shard's own already-settled PVCs (its pool data PVCs and +// its backup PVC), named explicitly rather than guessed: this suite's own +// discipline is that an already-populated namespace's replay is the first +// step's problem to permit, not something to dodge by racing the watch +// ahead of convergence. Unlike the TopoServer test above, that is available +// here, because the shard does converge. +// +// The impostor itself is created after the watch opens, and its own +// creation-then-bind is declared with Before: DataPlaneSim +// (pkg/ctrltest/datasim.go) binds every PersistentVolumeClaim in the cluster +// regardless of who it belongs to, as a stand-in for the volume provisioner +// envtest does not run, and that status patch (status.phase/accessModes/ +// capacity) has to be permitted explicitly or it is indistinguishable from +// the actual ownerRef adoption this test is pinning. Declaring the pair +// with Before, rather than waiting for Bound out-of-band first, is what +// keeps the watch open across the one window where the real defect could +// otherwise race in unobserved, immediately after creation and before this +// suite's own fake gets to it. +func TestSelectorImpostorShardOwnerRefReconcileAdoptsUnrelatedPVC(t *testing.T) { + c := newCase(t) + ns := c.NS + c.MinimalCluster("pvc-adopt") + + shard := c.awaitShard() + + c.RequireQuiescent(time.Second, 30*time.Second) + + existingPVCs := &corev1.PersistentVolumeClaimList{} + c.NoError(c.List(existingPVCs), "list existing PVCs") + + impostor := &corev1.PersistentVolumeClaim{ + ObjectMeta: metav1.ObjectMeta{ + Name: "impostor-shared-pvc", + Namespace: ns, + Labels: map[string]string{ + metadata.LabelMultigresCluster: shard.Labels[metadata.LabelMultigresCluster], + metadata.LabelMultigresDatabase: string(shard.Spec.DatabaseName), + metadata.LabelMultigresTableGroup: string(shard.Spec.TableGroupName), + metadata.LabelMultigresShard: string(shard.Spec.ShardName), + }, + }, + Spec: corev1.PersistentVolumeClaimSpec{ + AccessModes: []corev1.PersistentVolumeAccessMode{corev1.ReadWriteOnce}, + Resources: corev1.VolumeResourceRequirements{ + Requests: corev1.ResourceList{ + corev1.ResourceStorage: resource.MustParse("1Gi"), + }, + }, + }, + } + + allow := make([]ctrltest.Allow, 0, len(existingPVCs.Items)+1) + for _, pvc := range existingPVCs.Items { + allow = append(allow, ctrltest.Added("PersistentVolumeClaim", pvc.Name)) + } + allow = append(allow, ctrltest.Before( + ctrltest.Added("PersistentVolumeClaim", impostor.Name), + // Narrowed to status.phase rather than left open: DataPlaneSim's bind + // always touches it alongside accessModes/capacity, so naming it is + // enough to identify that write and that write only. Left open, this + // leaf would also happily absorb the ownerRef adoption this test + // exists to catch, since Changed with no paths matches any + // modification at all. + ctrltest.Changed("PersistentVolumeClaim", impostor.Name, "status.phase"), + )) + + script := c.NewScript(&corev1.PersistentVolumeClaimList{}) + + c.KnownDefect( + "pkg/resource-handler/controller/shard/shard_controller.go:589 "+ + "(reconcilePVCOwnerRefs selects on the shard's four identity labels "+ + "alone, so it adopts any PVC carrying them that has no controller "+ + "ownerRef, with nothing establishing the PVC is the shard's own)", + func() error { + stepErr := script.TryStep( + "baseline PVCs replay; the impostor is created, then bound like "+ + "any other PVC by this suite's data-plane fake", + func() error { + if err := c.Create(impostor); err != nil { + return err + } + // The impostor's own creation cannot trigger the shard's + // reconcile loop (it carries no owner reference, so + // Owns(&PersistentVolumeClaim{}) has nothing to map it back + // to), and RequireQuiescent above means nothing else is left + // to either: measured empirically, a fully quiesced shard + // does not reconcile again on its own. A metadata-only nudge + // on the Shard itself is what actually gets + // reconcilePVCOwnerRefs to run again and look at the + // impostor; this write is on ShardList, not the + // PersistentVolumeClaimList this script watches, so it needs + // no entry of its own in allow. + shardCopy := shard.DeepCopy() + patch := client.MergeFrom(shardCopy.DeepCopy()) + if shardCopy.Annotations == nil { + shardCopy.Annotations = map[string]string{} + } + shardCopy.Annotations[selectorImpostorNudgeAnnotation] = time.Now(). + UTC(). + Format(time.RFC3339Nano) + return c.Patch(shardCopy, patch) + }, + allow..., + ) + // Run unconditionally, and 10s, for the reasons given at the same + // call in the TopoServer test above. Unlike that test this one + // needs no sentinel to tell two permitted-set outcomes apart: the + // adoption is a modification of the impostor, and the only other + // modification anything makes to it is DataPlaneSim's bind, which + // the step permits by name and by path. So there is no second + // route to a non-nil error through a permitted-set mismatch. + // + // That is narrower than "no second route at all", and the + // difference matters for how much this pin can be trusted. A + // TryStep timeout is also a non-nil error, and KnownDefect reads + // any non-nil error as the defect still being live, so if the data + // plane fake never binds the impostor or a baseline PVC name + // drifts, this pin survives the operator fix that should have + // retired it. That is the general limitation of pinned steps + // stated on TryStep, and this call site is not exempt from it. The + // TopoServer test's sentinel-plus-Fatalf discrimination is the + // honest pattern if this ever needs to be tightened. + finishErr := script.TryFinish(10 * time.Second) + if stepErr != nil { + return stepErr + } + return finishErr + }, + ) +} diff --git a/test/suite/scenario_shard_lifecycle_test.go b/test/suite/scenario_shard_lifecycle_test.go new file mode 100644 index 00000000..d494a34d --- /dev/null +++ b/test/suite/scenario_shard_lifecycle_test.go @@ -0,0 +1,896 @@ +package suite + +import ( + "context" + "fmt" + "slices" + "strings" + "testing" + "time" + + corev1 "k8s.io/api/core/v1" + storagev1 "k8s.io/api/storage/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + "k8s.io/apimachinery/pkg/api/meta" + "k8s.io/apimachinery/pkg/api/resource" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + "github.com/multigres/multigres-operator/pkg/resolver" + shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" + "github.com/multigres/multigres-operator/pkg/util/metadata" + "github.com/multigres/multigres-operator/pkg/util/name" + "github.com/multigres/testkit/ctrltest" +) + +const lifecycleShardName multigresv1alpha1.ShardName = "0-inf" + +// blockedObservationWindow is how long steps 5 and 6 watch for a reaction +// that never comes. +// +// The shard controller requeues every disruptionRecoveryRequeue (5s, +// pkg/resource-handler/controller/shard/disruption.go:17) for as long as a +// disruption is refused, so a window several times that long is the +// difference between "the operator tried repeatedly and refused every time" +// and "the operator had not got round to trying yet". A Quiet() step on its +// own asserts silence only over the 250ms settle window, which for this +// question would be almost nothing. +const blockedObservationWindow = 20 * time.Second + +// lifecycleShardRef is a Shard carrying only the fields BuildPoolPodName, +// BuildPoolDataPVCName and BuildSharedBackupPVCName read: the cluster label +// and the three spec names. Those four values are known before the real +// Shard exists, since lifecycleCluster (below) chooses them, which lets the +// test predict a pod or PVC's name ahead of the create that produces it, +// using the operator's own name builders rather than a hand-rolled format +// string. It is never sent to the API server. +func lifecycleShardRef(clusterName string) *Shard { + return &Shard{ + ObjectMeta: metav1.ObjectMeta{ + Labels: map[string]string{metadata.LabelMultigresCluster: clusterName}, + }, + Spec: multigresv1alpha1.ShardSpec{ + DatabaseName: resolver.DefaultSystemDatabaseName, + TableGroupName: resolver.DefaultSystemTableGroupName, + ShardName: lifecycleShardName, + }, + } +} + +// lifecycleStorageClassName is the StorageClass lifecycleCluster's pool +// references, distinct per namespace so concurrently-running instances of +// this test never collide on the same cluster-scoped object. +func lifecycleStorageClassName(ns string) string { + return ns + "-lifecycle-expandable" +} + +// lifecycleCluster creates the same MultigresCluster MinimalCluster does +// (one cell, one database, one table group, one shard), plus a StorageClass +// with AllowVolumeExpansion set and the pool's Storage.Class pointed at it. +// +// Two facts force this rather than a plain call to MinimalCluster. First, +// PopulateClusterDefaults's own injection of Databases is commented +// "in-memory" for a reason: with no mutating webhook running in this suite, +// nothing ever writes it back to the MultigresCluster object itself, every +// reconcile recomputes it from scratch, and cluster.Spec.Databases stays +// permanently empty on the server unless a caller sets it explicitly. That +// alone would still allow starting from MinimalCluster's bare cluster and +// seeding Databases later, in updateLifecyclePool's first call. But second, +// storageClassName is immutable once a PersistentVolumeClaim exists, and step +// 3 needs the pool's existing data PVCs to already reference a StorageClass +// that allows expansion, or the API server's resize admission check refuses +// the request outright ("only dynamically provisioned pvc can be resized"). +// So the StorageClass has to be in place, and referenced, from this create. +// +// The shape mirrors what resolver.PopulateClusterDefaults would have +// injected for a MinimalCluster (one pool, one cell, the two-replica floor +// pkg/resolver/shard.go computes for a single-cell pool): every attribute +// this script mutates is still reached by changing a cluster that already +// exists, this just makes explicit at creation the one attribute (storage +// class) that cannot be introduced by a later mutation. +func (c *C) lifecycleCluster(clusterName string) *MultigresCluster { + c.Helper() + + scName := lifecycleStorageClassName(c.NS) + allowExpansion := true + sc := &storagev1.StorageClass{ + ObjectMeta: metav1.ObjectMeta{Name: scName}, + Provisioner: "multigres-test/no-op", + AllowVolumeExpansion: &allowExpansion, + } + c.NoError(c.Create(sc), "create StorageClass %s", scName) + // Cluster-scoped, so the per-test namespace delete does not reach it and + // every run would otherwise leave one behind in the shared envtest API + // server for the life of the package. context.Background rather than + // c.Context, which is already cancelled by the time cleanups run. + c.Cleanup(func() { + _ = c.Client().Delete(context.Background(), sc) + }) + + return c.newCluster(clusterName, func(s *MultigresClusterSpec) { + s.Databases = []DatabaseConfig{{ + Name: resolver.DefaultSystemDatabaseName, + Default: true, + TableGroups: []TableGroupConfig{{ + Name: resolver.DefaultSystemTableGroupName, + Default: true, + Shards: []ShardConfig{{ + Name: lifecycleShardName, + Spec: &ShardInlineSpec{ + Pools: map[PoolName]PoolSpec{ + resolver.DefaultPoolName: { + Type: "readWrite", + Cells: []CellName{defaultSimCell}, + ReplicasPerCell: ptr.To(int32(2)), + Storage: multigresv1alpha1.StorageSpec{Class: scName}, + }, + }, + }, + }}, + }}, + }} + }) +} + +// updateLifecyclePool re-reads the cluster and applies mutate to the +// "default" pool's spec, retrying on a conflict from a concurrent status +// write. The cluster's own status subresource is patched by the +// multigrescluster controller on a completely separate write path, but any +// spec Update still carries the resourceVersion it read, so a status patch +// landing between our Get and our Update aborts it. +func (c *C) updateLifecyclePool( + clusterName string, + mutate func(*PoolSpec), +) { + c.Helper() + key := client.ObjectKey{Namespace: c.NS, Name: clusterName} + for { + cluster := &MultigresCluster{} + c.NoError(c.Get(key, cluster), "get cluster %s", clusterName) + pools := cluster.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools + pool := pools[resolver.DefaultPoolName] + mutate(&pool) + pools[resolver.DefaultPoolName] = pool + + err := c.Update(cluster) + if err == nil { + return + } + if !apierrors.IsConflict(err) { + c.Fatalf("update cluster %s: %v", clusterName, err) + } + } +} + +// waitFor is eventually without the t.Fatalf, returning the last error instead. +// +// A KnownDefect body has to be able to contain a wait: KnownDefect reads a +// returned error as the pinned defect still being live, and eventually ends the +// goroutine through t.Fatalf rather than returning, so a pin whose subject is +// "this never happens" cannot be written with eventually at all. +func waitFor(timeout time.Duration, cond func() error) error { + deadline := time.Now().Add(timeout) + for { + err := cond() + if err == nil { + return nil + } + if time.Now().After(deadline) { + return err + } + time.Sleep(250 * time.Millisecond) + } +} + +// drainStateSeen reports nil once any pod in ns carries the drain state +// machine's annotation, and an error while none does. +// +// This is the first write a drain makes: initiateDrain patches the annotation +// to Requested (pkg/data-handler/drain/drain_helpers.go) before anything else +// moves, and ExecuteDrainStateMachine advances it one state per reconcile from +// there. So the annotation's mere presence is the earliest observable evidence +// that a drain started, which is exactly what step 2's pin is about, and unlike +// a count of pod events it does not depend on how many reconciles the operator +// took to get anywhere. +func (c *C) drainStateSeen() error { + c.Helper() + pods := &corev1.PodList{} + if err := c.List(pods); err != nil { + return err + } + for i := range pods.Items { + if _, ok := pods.Items[i].Annotations[metadata.AnnotationDrainState]; ok { + return nil + } + } + return fmt.Errorf("no pod in %s carries %s", c.NS, metadata.AnnotationDrainState) +} + +// poolPodFingerprints maps every pod in ns to its UID and resourceVersion. +// +// Comparing two of these across a window is how this file asserts that the +// operator left the pods alone, and it replaces what a Quiet() step used to say +// about pods before pods came out of the script's watch entirely (see the +// script's own comment in TestShardLifecycle). It is not the weaker claim: +// resourceVersion is monotonic per object and moves on every write the API +// server accepts, so any patch at all, by the operator or by the data-plane +// fake, shows up as a changed fingerprint. A pod created or deleted in the +// window changes the key set, and a delete-and-recreate at the same name +// changes the UID. What it does not do is care when any of that happened, which +// is the whole reason to read pods rather than watch them: the claim is about +// the operator's writes and not about the fake's pacing. +func (c *C) poolPodFingerprints() map[string]string { + c.Helper() + pods := &corev1.PodList{} + c.NoError(c.List(pods), "list pods in %s", c.NS) + out := make(map[string]string, len(pods.Items)) + for i := range pods.Items { + p := &pods.Items[i] + out[p.Name] = string(p.UID) + "@" + p.ResourceVersion + } + return out +} + +// podFingerprintDiff describes how two poolPodFingerprints snapshots differ, or +// returns "" when they are identical. Sorted, so a failure message is the same +// text on every run rather than whatever order the map iterated in. +func podFingerprintDiff(before, after map[string]string) string { + var notes []string + for name, was := range before { + now, ok := after[name] + switch { + case !ok: + notes = append(notes, fmt.Sprintf("%s was deleted", name)) + case now != was: + notes = append(notes, fmt.Sprintf("%s was written (%s to %s)", name, was, now)) + } + } + for name := range after { + if _, ok := before[name]; !ok { + notes = append(notes, fmt.Sprintf("%s was created", name)) + } + } + slices.Sort(notes) + return strings.Join(notes, "; ") +} + +// firstCreateIndex returns the position in ops of controller's first accepted +// create of kind/name, and whether it made one. +// +// This is how step 4's ordering claim survives pods leaving the script's watch. +// The op log's order is arrival at the recorder's mutex, which recorder.go is +// explicit is program order within one controller and not a causal order across +// controllers, so two indexes are only comparable when both name the same +// controller. Both creates here are the shard controller's, in one pass of +// createMissingResources, which is what makes the comparison sound and is why +// the controller is a parameter rather than left implicit. +func firstCreateIndex(ops []ctrltest.Op, controller, kind, name string) (int, bool) { + for i, op := range ops { + if op.Controller == controller && op.Verb == "create" && + ctrltest.KindSuffix(op.Kind) == kind && op.Key.Name == name { + return i, true + } + } + return 0, false +} + +// podRoleViolation reports whether err is one of the errors MembersOf returns +// about what status.podRoles actually says, as opposed to a transient "not +// reconciled yet" state (shard not found, status.podRoles still empty) that +// the standing invariant below must not treat as a violation. +// +// The unrecognized-role error has to count, not just the two primary-count +// ones. MembersOf returns it from its classification loop, before it counts +// primaries at all, so a snapshot holding one pod in a role this suite does +// not know (a future DRAINED, say) and two pods reporting PRIMARY comes back +// as the unrecognized-role error alone. Reading that as "not a violation" +// would wave the two primaries through, which is the one thing the invariant +// exists to catch. +// +// Matching on message text is the only option MembersOf offers, since it +// exports no sentinel errors. That coupling is invisible from identity.go, so +// rewording a message there disables this check silently; closing it needs +// either sentinels in identity.go or a unit test pinning these strings, and +// both are outside this file. +func podRoleViolation(err error) bool { + if err == nil { + return false + } + msg := err.Error() + return strings.Contains(msg, "no pod has role PRIMARY") || + strings.Contains(msg, "more than one pod has role PRIMARY") || + strings.Contains(msg, "has unrecognized role ") +} + +// rollingUpdateDrift reports whether the Shard says wantPods of its pool pods +// have drifted from their desired spec, through the RollingUpdate condition +// handleRollingUpdates writes (reconcile_pool_pods.go:735-751). +// +// This is the only object state the operator changes in response to a spec +// change it then refuses to act on: the drain it would start next is blocked +// before it writes anything to a pod, and the DisruptionBlocked Event that +// refusal records goes through client-go's per-object spam filter (burst 25, +// one refill per 300s), which a chatty Shard inside a short envtest run has +// already spent. Steps 5 and 6 need this because a step asserting that +// nothing happened asserts nothing at all unless the stimulus provably +// arrived first. +// +// The expected message is built from the same format string the operator +// uses, so this is coupled to that wording. The coupling is deliberate and +// fails in the safe direction: a reworded message makes the step fail loudly +// rather than quietly stop asserting. +func rollingUpdateDrift(key client.ObjectKey, wantPods int) error { + shard := &Shard{} + if err := Suite.Client.Get(context.Background(), key, shard); err != nil { + return err + } + cond := meta.FindStatusCondition(shard.Status.Conditions, "RollingUpdate") + if cond == nil { + return fmt.Errorf("shard %s has no RollingUpdate condition yet", key) + } + want := fmt.Sprintf("%d pods need update in pool %s", wantPods, resolver.DefaultPoolName) + if cond.Status != metav1.ConditionTrue || cond.Reason != "PodsDrifted" || + cond.Message != want { + return fmt.Errorf( + "shard %s reports RollingUpdate=%s reason=%s %q, want True PodsDrifted %q", + key, cond.Status, cond.Reason, cond.Message, want, + ) + } + return nil +} + +// lifecycleResources builds a concrete, distinguishable resource request/limit +// pair so successive calls with different label values are guaranteed to +// differ from both the resolver's own defaults and from each other. +func lifecycleResources(cpuReq, memReq, cpuLim, memLim string) corev1.ResourceRequirements { + return corev1.ResourceRequirements{ + Requests: corev1.ResourceList{ + corev1.ResourceCPU: resource.MustParse(cpuReq), + corev1.ResourceMemory: resource.MustParse(memReq), + }, + Limits: corev1.ResourceList{ + corev1.ResourceCPU: resource.MustParse(cpuLim), + corev1.ResourceMemory: resource.MustParse(memLim), + }, + } +} + +// TestShardLifecycle starts from the smallest cluster this suite can +// converge (one pool, two pods, under the two-replica floor a single-cell +// pool always resolves to) and mutates it forward: resource change, storage +// change, scale up, a second resource change with three members present, +// and scale back down. Every attribute is reached by changing a cluster +// that already exists, never by constructing the end state, so this +// exercises the operator's transition paths rather than its defaulting +// paths. +func TestShardLifecycle(t *testing.T) { + c := newCase(t) + ns := c.NS + const clusterName = "lifecycle" + + shardKey := client.ObjectKey{ + Namespace: ns, + Name: name.JoinWithConstraints( + name.DefaultConstraints, + clusterName, + string(resolver.DefaultSystemDatabaseName), + string(resolver.DefaultSystemTableGroupName), + string(lifecycleShardName), + ), + } + shardRef := lifecycleShardRef(clusterName) + const cellName = defaultSimCell + + poolName := string(resolver.DefaultPoolName) + pod0 := shardcontroller.BuildPoolPodName(shardRef, poolName, cellName, 0) + pvc0 := shardcontroller.BuildPoolDataPVCName(shardRef, poolName, cellName, 0) + pod1 := shardcontroller.BuildPoolPodName(shardRef, poolName, cellName, 1) + pvc1 := shardcontroller.BuildPoolDataPVCName(shardRef, poolName, cellName, 1) + backupPVC := shardcontroller.BuildSharedBackupPVCName(shardRef) + + // Step 1: converge, out of the script's sight, and then open the script + // over what converging produced. + // + // lifecycleCluster mirrors what MinimalCluster plus the resolver's own + // defaults would produce (one cell, one pool at the two-replica floor + // pkg/resolver/shard.go computes for a single-cell pool under the default + // AT_LEAST_2 durability policy), so "the smallest possible cluster" + // already has a primary and a replica the moment it is healthy: pods 0 + // and 1, their data PVCs, and the shard's one shared backup PVC. + // + // The watch therefore opens after its own fixture, against this suite's + // standing rule that a stream is established before the objects it + // watches are created. That rule exists so a reaction to the create + // cannot land before anyone is listening, and the events it protects here + // are ones no assertion reads: the real claims of this script are steps 2 + // through 6, each about one mutation of a cluster that already exists, + // and none of them looks at how the cluster got there. + // + // What the closed step it replaces did read was the convergence sequence + // itself, and that count is not the operator's to keep. A step's + // allow-list is an exact multiset with no "N events of this kind" form, + // so a declared count is sound only where the code fixes it. Pod + // readiness here is paced by the data-plane sim (pkg/ctrltest/datasim.go) + // racing reconcilePoolerReadiness's own readiness-gate patch on the same + // Pod, and measured over runs of this test a pod settles in three status + // modifications most of the time and four sometimes, which failed step 1 + // on an unpermitted event about one run in five. A fourth speculative + // Changed leaf would only move that boundary, since nothing bounds the + // sequence at three or four either. Do not restore the closed step: it + // pinned the harness, not the operator. + // + // RequireQuiescent is what stands in its place, and over the convergence + // it is strictly the stronger claim: no projected state change on any of + // watchedKinds() and no attempted write for five seconds, rather than a + // count of Pod and PVC events. It is also what makes the baseline below + // a fixed set rather than a race, because it establishes that nothing is + // still moving when the watch opens. + c.lifecycleCluster(clusterName) + c.Eventually(60*time.Second, "shard to report Healthy", func() error { + shard := &Shard{} + if err := c.Get(shardKey, shard); err != nil { + return err + } + if shard.Status.Phase != multigresv1alpha1.PhaseHealthy { + return fmt.Errorf("shard phase is %q", shard.Status.Phase) + } + return nil + }) + c.RequireQuiescent(5*time.Second, 60*time.Second) + + // PersistentVolumeClaim and not Pod, which is the general form of what + // step 1 above ran into rather than a second workaround for it. + // + // A closed step counts events, so it is sound only over writes whose number + // follows from the operator's code. Every PVC event in this namespace is + // one of those: the operator creates each PVC once, patches it once per + // resize, and the data-plane fake's bind is a single terminal transition + // from Pending to Bound rather than a progression, so it too is one write + // by construction. Pod status is the opposite. The fake walks a pod's + // readiness conditions in however many passes it takes while + // reconcilePoolerReadiness writes its readiness gate on the same object, + // and the number of modifications that takes is a property of that race: + // three most of the time, four often enough to fail one run in five, + // measured here twice, at step 1 and again at step 4. + // + // Note what is not available as a middle road: keeping Pod in the watch + // while declining to enumerate status modifications. There is no "N events + // of this kind" form, so a watched kind's events must each be permitted by + // name or they fail the step in flight. Watching pods therefore forces the + // enumeration, which is why pods come out of the watch altogether and every + // pod-level claim below is a direct read instead. Those reads are not the + // weaker choice; see poolPodFingerprints for why a state comparison is at + // least as strong here as an event count, and firstCreateIndex for how the + // one genuine ordering claim is made from the operator's own write log. + s := c.NewScript(&corev1.PersistentVolumeClaimList{}) + + // Standing invariant: exactly one PRIMARY whenever status.podRoles is + // non-empty, no pod in a role this suite does not know, and no pod + // quarantined. MembersOf distinguishes those errors from every other read + // failure (shard not found, podRoles still empty), and podRoleViolation is + // what keeps one of those from being reported as a violation. The script + // opens after convergence, so podRoles is populated by the time this is + // registered, but steps 2 through 6 mutate the pool and the shard rewrites + // podRoles as those land, so a read taken mid-rewrite is still ordinary. + // + // Quarantined is checked explicitly because it is its own bucket in + // Members, disjoint from Replicas: a script that only ever asserted + // primary-count and replica-count could watch a pod sit quarantined for + // its entire length without ever naming it. This script never triggers + // quarantine, so any appearance here is itself something to catch. + // + // Written as a closure and registered, rather than only registered, + // because an invariant is sampled once per event and this script now + // watches one kind: the steps that matter most to it, 2 and 5 and 6, are + // steps where the operator is refused and no event arrives at all, so + // registration alone would leave those windows unsampled. requirePodRoles + // is called directly at the end of each of them. + checkPodRoles := func() error { + members, err := MembersOf(context.Background(), c.Client(), shardKey) + if err != nil { + if podRoleViolation(err) { + return err + } + return nil + } + if len(members.Quarantined) > 0 { + return fmt.Errorf( + "pod(s) unexpectedly quarantined: %s", strings.Join(members.Quarantined, ", "), + ) + } + return nil + } + requirePodRoles := func(where string) { + c.Helper() + c.Check().NoError(checkPodRoles(), "pod roles after %s", where) + } + s.Invariant( + "shard has exactly one primary and no quarantined pod", + func(_ ctrltest.Event) error { return checkPodRoles() }, + ) + + // The watch carries no starting resourceVersion, so the API server replays + // the namespace's PVCs as "added" and this step permits exactly that + // baseline. Naming the three from the operator's own name builders, rather + // than from a List taken a moment earlier, is what keeps it an assertion: + // it says the converged pool's storage is those three claims and nothing + // else, so a fourth PVC or a misnamed one fails here instead of being + // permitted by whatever happened to exist. + s.Step("the converged pool's PVCs replay as the baseline", nil, + ctrltest.Added("PersistentVolumeClaim", pvc0), + ctrltest.Added("PersistentVolumeClaim", pvc1), + ctrltest.Added("PersistentVolumeClaim", backupPVC), + ) + + s.Step("cluster settles before any change", nil, ctrltest.Quiet()) + + // The pod half of that same claim, made as a read because pods are not + // watched: the converged pool is pods 0 and 1 and nothing else. This is + // what the baseline's two Added("Pod", ...) leaves used to say. + converged := c.poolPodFingerprints() + for _, want := range []string{pod0, pod1} { + c.HasKey(converged, want, "converged pool has no pod %s", want) + } + c.Eq(2, len(converged), "want exactly %s and %s", pod0, pod1) + + // The precondition that makes step 2 meaningful: a two-member pool, one of + // them primary, which is the cohort size the defect below is about. + // + // Read under Eventually rather than once. The convergence check above is + // paced by the data plane sim at 250ms, but status.podRoles is written + // through the pooler sim at 500ms, so both pods can be ready while the + // second pod's role has not propagated yet. A single read lands in that + // window often enough to matter: observed failing one full run in six, + // reporting one primary and zero replicas. + var initialMembers Members + c.Eventually( + 30*time.Second, + "the pool to report one primary and one replica", + func() error { + members, err := MembersOf(c.Context(), c.Client(), shardKey) + if err != nil { + return err + } + if len(members.Replicas) != 1 || members.Primary == "" { + return fmt.Errorf("got %+v", members) + } + initialMembers = members + return nil + }, + ) + + // Step 2: change CPU and memory on the pool. This is written as the + // positive assertion the brief describes, and it does not pass: the + // mutation lands in the Shard spec and podNeedsUpdate correctly flags both + // pool pods as drifted (the Shard reports RollingUpdate=True + // reason=PodsDrifted, "2 pods need update in pool default"), but + // handleRollingUpdates never drains either of them. canStartDisruption + // (disruption.go) calls posture.CheckDisruption, which calls + // consensus.CheckSufficientRecruitment against the two-pooler rule + // test/suite/fakes.go's poolerSim registers; that function's own majority + // rule is len(cohort)/2+1, which for a two-member cohort is 2, so + // excluding either pod to disrupt it always leaves 1 short. This is not + // gated by DurabilityPolicy at all: fakes.go's rule sets AtLeastN(1), not + // the cluster's AT_LEAST_2 default, so the floor here comes from the + // majority check that runs before any policy-specific one, and it blocks + // a two-member pool categorically, not just under this suite's default. + // + // This step is pinned, where steps 5 and 6 below are not, because this one + // is a claim about the operator. A single-cell pool takes ReplicasPerCell 2 + // from the resolver's own default (pkg/resolver/shard.go:102-113), and the + // operator's purpose-built escape hatch for a member that cannot be + // disrupted, reconcileCellMaintenanceSurge, is gated on + // MULTI_CELL_AT_LEAST_2 with exactly two cells + // (maintenance_surge.go:349-351), so the shape the operator defaults to + // gets no surge and no other route. The pin may never expire, which is + // worth saying plainly: it expires if the majority rule changes for a + // 2-cohort, if the single-cell default moves off 2, or if the surge gate + // widens to the condition that actually triggers it, and not otherwise. + // + // The pin is that no drain ever starts, and it is written as exactly that: + // the drain state annotation never appears on any pod in the namespace. + // What it replaced was a permitted set enumerating the whole nine-event + // cycle a drain would produce on each of the two pods, which pinned the + // same defect less precisely and could not survive the readiness leaves + // three of those nine were (see step 1). Naming the annotation is the + // better pin on its own terms: initiateDrain patches it before the drain + // machine does anything else, so a drain that started and then stalled + // halfway confirms the pin today and would be indistinguishable from + // "never started" under a count of events that never arrived. + // + // The positive half comes first and is not part of the pin. A step that + // asserts nothing happened asserts nothing at all unless the stimulus + // provably arrived, so the drift condition is waited on with eventually, + // which fails the test outright rather than confirming the pin, exactly as + // KnownDefect's contract requires of a setup step. + // + // The step's own Quiet() carries the closed-world half over PVCs: a rolling + // update rewrites pods and leaves their claims alone, so the PVC silence + // here is the assertion that the operator did not take some other action + // instead of the one it refused. + s.Step("change CPU and memory on the pool", func() error { + c.updateLifecyclePool(clusterName, func(p *PoolSpec) { + p.Postgres.Resources = lifecycleResources("100m", "128Mi", "200m", "256Mi") + }) + c.Eventually( + 30*time.Second, + "the shard to report both pool pods drifted", + func() error { + return rollingUpdateDrift(shardKey, len(initialMembers.Replicas)+1) + }, + ) + podsBefore := c.poolPodFingerprints() + c.KnownDefect("shard-lifecycle-two-member-pool-never-rolls", + func() error { + return waitFor(blockedObservationWindow, func() error { + return c.drainStateSeen() + }) + }, + ) + // Stronger than the pin and independent of it: not only did no drain + // annotation appear, no pod was written to at all while we watched. + c.Check().Eq("", podFingerprintDiff(podsBefore, c.poolPodFingerprints()), + "the refused rolling update still touched pods") + requirePodRoles("the refused rolling update") + return nil + }, ctrltest.Quiet()) + + // Step 3: change storage. expandPVCIfNeeded patches the PVC's storage + // request directly; it is not drift the pod's spec hash notices (the pod + // references the PVC by name, not by size), so no pod event follows. + // This mutation is independent of step 2's never-applied one and is + // unaffected by it. + s.Step("change storage", func() error { + c.updateLifecyclePool(clusterName, func(p *PoolSpec) { + p.Storage.Size = "2Gi" + }) + return nil + }, ctrltest.Changed("PersistentVolumeClaim", pvc0), ctrltest.Changed("PersistentVolumeClaim", pvc1)) + + // Step 4: add a replica. The closed step covers the new PVC, created once + // by the operator and bound once by the data-plane fake. The new pod is + // waited on as a read, and the one genuine code-level ordering here, that + // the PVC exists before the pod that binds it, is asserted below from the + // operator's own write log rather than from the arrival order of two + // watches. + // + // This step is where the second instance of step 1's problem was measured: + // it declared three Changed("Pod", pod2) leaves for the readiness settle + // and failed on a fourth, one run in eight, with the same + // status.conditions paths. Those leaves are gone rather than widened. + pod2 := shardcontroller.BuildPoolPodName(shardRef, poolName, cellName, 2) + pvc2 := shardcontroller.BuildPoolDataPVCName(shardRef, poolName, cellName, 2) + + s.Step("add a replica", func() error { + c.updateLifecyclePool(clusterName, func(p *PoolSpec) { + p.ReplicasPerCell = ptr.To(int32(3)) + }) + c.Eventually(30*time.Second, "new pod ready", func() error { + pod := &corev1.Pod{} + key := client.ObjectKey{Namespace: ns, Name: pod2} + if err := c.Get(key, pod); err != nil { + return err + } + for _, cond := range pod.Status.Conditions { + if cond.Type == corev1.PodReady && cond.Status == corev1.ConditionTrue { + return nil + } + } + return fmt.Errorf("pod %s not Ready yet", pod2) + }) + return nil + }, + ctrltest.Added("PersistentVolumeClaim", pvc2), + ctrltest.Changed("PersistentVolumeClaim", pvc2), + ) + + // The ordering claim, from the write log: createMissingResources creates + // the data PVC and then the pod that mounts it, in that order, within one + // pass (reconcile_pool_pods.go:178-198 and the Create that follows it). + // Both are the shard controller's writes, which is what makes their + // relative position in the log program order rather than arrival noise; + // see firstCreateIndex. + ops := Suite.Ops.OpsInNamespace(ns) + pvcAt, pvcCreated := firstCreateIndex(ops, "shard", "PersistentVolumeClaim", pvc2) + podAt, podCreated := firstCreateIndex(ops, "shard", "Pod", pod2) + c.Check().True(pvcCreated, "the shard controller never created PVC %s", pvc2) + c.Check().True(podCreated, "the shard controller never created pod %s", pod2) + // Only meaningful once both exist. firstCreateIndex reports a miss as + // index 0, so an absent pod would otherwise read as one created before + // its PVC, and the run would carry a confident ordering complaint about + // an object that was never created. The switch this replaced got that + // right by being mutually exclusive; three independent checks have to say + // it explicitly. + if pvcCreated && podCreated { + c.Check().True(pvcAt <= podAt, + "pod %s was created before the PVC %s it binds (ops %d and %d)", + pod2, pvc2, podAt, pvcAt) + } + + requirePodRoles("the scale-up") + + // The new pod must bind the new PVC, not either existing one: resolved + // through ShardPVCOf against the live pod rather than assumed from the names + // above. + c.Eventually(30*time.Second, "new pod bound to a PVC", func() error { + _, err := ShardPVCOf(c.Context(), c.Client(), ns, pod2) + return err + }) + boundTo, err := ShardPVCOf(c.Context(), c.Client(), ns, pod2) + c.NoError(err, "ShardPVCOf(%s)", pod2) + // One comparison, not two: pvc0, pvc1 and pvc2 are distinct names, so + // "it is pvc2" already says "it is not either existing PVC", and a second + // check against those two could never fire. + c.Eq(pvc2, boundTo, "pod %s is bound to the wrong PVC, want the new PVC rather than %s or %s", + pod2, pvc0, pvc1) + + // Steps 5 and 6 do not assert the guarantee this script was written to + // assert, and say so rather than implying otherwise. + // + // That guarantee, documented at + // pkg/resource-handler/controller/shard/reconcile_pool_pods.go:711, is + // that handleRollingUpdates drains drifted pods one at a time, replicas + // before the primary, with the primary's switchover requested before it is + // touched, and that handleScaleDown removes an extra replica by draining + // it. Nothing in this harness can observe any of it, because no drain ever + // starts at any cohort size. canStartDisruption calls + // posture.CheckDisruption, whose revocation check + // (consensus.CheckSufficientRecruitment) requires that the pod being + // excluded cannot satisfy the durability policy on its own, and + // test/suite/fakes.go's poolerSim registers DurabilityPolicy: + // topoclient.AtLeastN(1) regardless of the shard's configured policy, so + // any single excluded pod satisfies it alone and the check refuses every + // exclusion. handleScaleDown calls the identical gate + // (reconcile_pool_pods.go:629), which is why step 6 is in the same + // position as step 5. + // + // Neither step is a KnownDefect, and that is the point. A pin records a + // live operator defect and expires on the day the operator is fixed; this + // is a harness fidelity gap, so a pin here could never expire, which makes + // it a suppression wearing a pin's clothes. Nor is threading the shard's + // real policy through fakes.go known to be enough to reach these steps: at + // AtLeastN(2) the failure moves earlier instead, the shard loses its + // PRIMARY from status.podRoles altogether, the standing invariant above + // fires during step 4, and the operator hot-loops "No primary in podRoles, + // requeueing to re-read topology". Whether that is further harness + // infidelity (the fake models no election and no promotion, and registers + // each pooler's routing role once, at first sight) or an operator defect at + // AT_LEAST_2 with three members is unresolved, and settling it is the + // prerequisite for writing these two steps for real. + // + // What is left is the pair of facts this harness can establish, asserted + // positively: the operator sees the change, and then does nothing about it + // for as long as we watch. Both halves carry weight. Without the first the + // silence would also be satisfied by a mutation that never reached the + // shard controller at all, which is the vacuous assertion this suite + // exists to refuse; without the second there is no tripwire. The day + // either half changes, for a harness reason or an operator one, these + // steps fail and force the question open again. + // + // One note for whoever picks that up, because it is not obvious from the + // guarantee's wording: the switchover half has no Pod or PVC footprint to + // assert even in principle. handleRollingUpdates "requests a switchover" + // by calling the same initiateDrain annotation patch it uses for a replica + // (drain_helpers.go:58) and recording a RollingUpdateStarted Event on the + // Shard. So the only Kubernetes-visible difference between draining the + // primary and draining a replica is an Event, on an object this script + // does not watch, delivered over the spam-filtered path rollingUpdateDrift + // describes. Asserting it needs a different observation, not a better + // permitted set. + + // Step 5: change CPU and memory again, now with three members. Resolved + // through MembersOf rather than assumed from index, both as the + // precondition that makes the step meaningful (silence about a rolling + // update is only interesting over a pool that really does hold three + // members, one of them primary) and because the drifted-pod count + // asserted below is derived from it. + members, err := MembersOf(c.Context(), c.Client(), shardKey) + c.NoError(err, "MembersOf before step 5") + c.Eq(2, len(members.Replicas), "want two replicas before step 5, got %+v", members) + c.NotEq("", members.Primary, "want a primary before step 5, got %+v", members) + poolPods := len(members.Replicas) + 1 + + s.Step("change CPU and memory again, now with three members", func() error { + podsBefore := c.poolPodFingerprints() + c.updateLifecyclePool(clusterName, func(p *PoolSpec) { + p.Postgres.Resources = lifecycleResources("150m", "192Mi", "300m", "384Mi") + }) + // The count is what makes this non-vacuous. Step 2's change was never + // applied either, so two pods have been drifted since then and a bare + // "RollingUpdate is True" would have been satisfied before this step + // ran; only the third pod, created at step 4 from the then-current + // spec, drifts because of this mutation. + c.Eventually( + 30*time.Second, + "the shard to report every pool pod drifted", + func() error { + return rollingUpdateDrift(shardKey, poolPods) + }, + ) + // Slept inside the step's own action rather than after it, so the + // closed world covers the whole window: any event the operator + // produces while we wait is buffered by the stream and fails the + // Quiet() below as an unpermitted change. + time.Sleep(blockedObservationWindow) + // The pod half of the same silence, which the Quiet() cannot carry now + // that pods are read rather than watched. Snapshotted before the + // mutation rather than after the drift wait, so the window compared + // here is the whole step and not just its tail, which is the span the + // Quiet() covers. + c.Check().Eq("", podFingerprintDiff(podsBefore, c.poolPodFingerprints()), + "the refused rolling update still touched pods") + requirePodRoles("the refused three-member rolling update") + return nil + }, ctrltest.Quiet()) + + // Step 6: scale the pool back down, which should drain and remove one + // replica. Same shape as step 5: prove the request reached the Shard the + // controller reconciles, then watch it be refused. + // + // When this becomes assertable, do not resolve the pod being removed as + // the highest-named replica. selectShardScaleDownPod (disruption.go:144) + // selects on the pod index being at or above the desired replica count, + // and the two only agree here because this fake always elects the + // lowest-indexed pod primary: with pod 2 primary, the highest-named + // replica is pod 1 and the operator would remove pod 2. Derive it from the + // index, confirm through MembersOf that it is not the primary, and take + // its PVC from ShardPVCOf. The permitted set then wants the PVC's change before + // the pod's deletion, not after: cleanupDrainedPod patches the data PVC's + // orphan mark (reconcile_pool_pods.go:560) before the Delete at :567, and + // it patches rather than deletes because orphanByRemainingCount holds + // while the pool has three data PVCs and the threshold is 3 + // (shard_controller.go:49). + membersBeforeScaleDown, err := MembersOf(c.Context(), c.Client(), shardKey) + c.NoError(err, "MembersOf before step 6") + c.Eq(2, len(membersBeforeScaleDown.Replicas), + "want two replicas before step 6, got %+v", membersBeforeScaleDown) + c.NotEq("", membersBeforeScaleDown.Primary, + "want a primary before step 6, got %+v", membersBeforeScaleDown) + + s.Step("scale the pool back down by one replica", func() error { + podsBefore := c.poolPodFingerprints() + c.updateLifecyclePool(clusterName, func(p *PoolSpec) { + p.ReplicasPerCell = ptr.To(int32(2)) + }) + // Scale-down writes no condition of its own, so what is proved here is + // narrower than step 5's: the desired count reached the Shard spec, + // which is the input handleScaleDown reads, while three pods are still + // live for it to act on. + c.Eventually( + 30*time.Second, + "the scale-down to reach the Shard spec", + func() error { + shard := &Shard{} + if err := c.Get(shardKey, shard); err != nil { + return err + } + pool, ok := shard.Spec.Pools[resolver.DefaultPoolName] + if !ok { + return fmt.Errorf("shard %s has no %q pool", shardKey, resolver.DefaultPoolName) + } + if got := ptr.Deref(pool.ReplicasPerCell, -1); got != 2 { + return fmt.Errorf("pool %q wants %d replicas per cell, not 2", + resolver.DefaultPoolName, got) + } + return nil + }, + ) + time.Sleep(blockedObservationWindow) + // The claim step 6 exists to make, in the only form this harness can + // state it while the gate refuses every exclusion: no pod was removed, + // and no pod was written to either. When the gate opens, the primary's + // pod and PVC must still be here and the replica's must not, which is + // what the note above is about; the fingerprint comparison is the half + // of that which is assertable today, and it is resolved from the live + // pods rather than from the names, so it does not assume which pod the + // operator would have picked. + c.Check().Eq("", podFingerprintDiff(podsBefore, c.poolPodFingerprints()), + "the refused scale-down still touched pods") + requirePodRoles("the refused scale-down") + return nil + }, ctrltest.Quiet()) + + s.Finish(time.Second) +} diff --git a/test/suite/scenario_shard_quiescence_test.go b/test/suite/scenario_shard_quiescence_test.go new file mode 100644 index 00000000..dedc79ef --- /dev/null +++ b/test/suite/scenario_shard_quiescence_test.go @@ -0,0 +1,26 @@ +package suite + +import ( + "testing" + "time" +) + +// TestShardStatusQuiesces pins the fix for the shard status hot loop. It was +// written red, against the defect described below, and went green when the two +// server-side-apply defects behind it were fixed. +// +// A healthy Shard should stop writing once its status reflects reality, and it +// did not: two server-side-apply defects fought each other forever, so +// status.orchReady and status.poolsReady flipped false/true and the +// StorageClassValid condition's message alternated between two strings, each +// several times a second, with no terminal state. Keeping the measurement here +// is the point of the test: a status that converges is the property, and these +// are the fields that used to prove it did not. +func TestShardStatusQuiesces(t *testing.T) { + c := newCase(t) + cluster := c.MinimalCluster("quiesce") + + c.WaitForClusterHealthy(cluster) + + c.RequireQuiescent(10*time.Second, 30*time.Second) +} diff --git a/test/suite/scenario_thrash_test.go b/test/suite/scenario_thrash_test.go new file mode 100644 index 00000000..71d14d63 --- /dev/null +++ b/test/suite/scenario_thrash_test.go @@ -0,0 +1,416 @@ +package suite + +import ( + "fmt" + "testing" + "time" + + corev1 "k8s.io/api/core/v1" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + + shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" + "github.com/multigres/multigres-operator/pkg/util/metadata" +) + +// TestThrash adds and removes things faster than the operator can converge, +// then asserts it lands in the right end state and stops. Both subtests +// assert only two things: what the namespace looks like once it settles, and +// that it does settle at all. +// +// Neither subtest uses a Script. A Script's steps are a closed-world +// assertion of one interleaving of events, and thrash does not produce one +// interleaving: which reconciler wins a given race, how many redundant SSA +// applies land, and how many requeues fire along the way are all legitimately +// unpredictable once changes are issued faster than the operator can react to +// them. A Script step here would be asserting an ordering nobody can +// guarantee, which is exactly the kind of flake this suite exists to stop +// writing. So the claim is only about the end state and about quiescence, +// made with ordinary reads and RequireQuiescent, not about the path taken to +// get there. Do not "improve" this into a Script. +func TestThrash(t *testing.T) { + t.Run("pool replica thrash", testPoolReplicaThrash) + t.Run("child CR thrash", testChildCRThrash) +} + +// thrashPoolName matches resolver.DefaultPoolName, the name the operator +// gives the pool it injects when a cluster specifies none. Importing +// pkg/resolver for one string constant did not seem worth the dependency, so +// this is a literal deliberately kept next to its cross-check. +const thrashPoolName = PoolName("default") + +// testPoolReplicaThrash scales one pool 1 -> 2 -> 1 -> 2 with no wait between +// changes, then requires the namespace to go quiet and every pool PVC to be +// bound to a live pod. +// +// It does not assert on this thrashed shard's status.podRoles: whether the +// last scale-up's role lands there is a race against the defect that +// pinPoolScaleUpRoleStale constructs deterministically on a namespace of its +// own, so asserting it here would only sample that race. +func testPoolReplicaThrash(t *testing.T) { + c := newCase(t) + cluster := c.poolThrashCluster("pool-thrash", 1) + c.WaitForClusterHealthy(cluster) + + // Fired back to back, not waited on between calls: each Patch is an + // unconditional merge patch computed against the object's state as of the + // previous call in this loop, so it lands regardless of what the operator + // has or hasn't done with the prior one yet. That is the thrash. + for _, n := range []int32{2, 1, 2} { + c.scalePoolTo(cluster, n) + } + + // MGO-POOL-SCALEUP-ROLE-STALE: once a shard reconciled to Healthy with N + // poolers, raising replicasPerCell to add an (N+1)th sometimes never gets + // that pooler's role into shard.Status.PodRoles, permanently. + // + // Whether it bites depends on whether the pooler registers before or + // after the reconcile that declares the shard converged, so sampled + // naturally it reproduces about a third of the time. The pin constructs + // that ordering instead, which makes it deterministic. + pinPoolScaleUpRoleStale(t) + + c.RequireQuiescent(10*time.Second, 90*time.Second) + + // The end state, asserted on this shard rather than inferred from the + // pin's namespace. Pod readiness is a Kubernetes-level fact that does + // not travel through shard.Status.PodRoles, so this is immune to the defect + // pinned above: without it a lost or reverted final scale-up settles at one + // pod, one bound PVC and a quiet namespace, and every other assertion here + // is satisfied by that. + live := c.liveReadyPoolPodNames() + c.Check().Len(live, 2, "ready pods") + + c.requireNoOrphanedPoolPVCs(live) +} + +// pinPoolScaleUpRoleStale pins that scaling a pool from one to two does not +// land the new pooler's role in status.podRoles when the pooler registers +// after the shard has already converged. +// +// On its own namespace and its own cluster, with no thrash, because the +// defect never needed one: a bare single scale-up reproduces it, and the +// thrash above only found it first. +// +// Nothing wakes the shard once it is Healthy: a registration is a write to +// the topology store, with no Kubernetes event behind it, and the shard does +// not requeue itself while a managed pod is still awaiting its pooler. The fix +// is that requeue; with it the role lands well inside a second here, because +// the suite compresses requeues, so the window below is generous. +func pinPoolScaleUpRoleStale(t *testing.T) { + t.Helper() + + c := newCase(t) + cluster := c.poolThrashCluster("scaleup", 1) + c.WaitForClusterHealthy(cluster) + key := c.shardKey() + + // Hold registration before scaling, so the second pooler cannot register + // until this test says so. Without the hold the defect's precondition, a + // shard that converged having seen fewer poolers than pods, arrives only + // when registration loses a race against the last reconcile. Measured + // 2026-09-19 by disabling the fix and running this six times: three runs + // caught the regression and three did not. Constructing the state instead + // of waiting for it takes that from roughly half to always. + release := poolers.HoldRegistrations(c.NS) + defer release() + + c.scalePoolTo(cluster, 2) + + // The precondition itself, waited on rather than assumed: two pool pods + // exist and the shard has settled on a PodRoles that knows about one. A + // release before this point would prove nothing, because the reconcile + // that notices the new pooler might be one the scale-up was going to + // trigger anyway. + c.Eventually( + 60*time.Second, + "the shard to converge having seen fewer poolers than pods", + func() error { + pods := &corev1.PodList{} + if err := c.List(pods, client.MatchingLabels{ + metadata.LabelMultigresPool: string(thrashPoolName), + }); err != nil { + return err + } + if len(pods.Items) != 2 { + return fmt.Errorf("want 2 pool pods, got %d", len(pods.Items)) + } + shard := &Shard{} + if err := c.Get(key, shard); err != nil { + return err + } + if len(shard.Status.PodRoles) != 1 { + return fmt.Errorf("want 1 pod role while held, got %d", len(shard.Status.PodRoles)) + } + return nil + }, + ) + // Deliberately not RequireQuiescent. A held namespace never goes quiet + // once the fix is in, because the requeue this pins is firing on its + // backoff the whole time, so quiescence holds before the fix and cannot + // after it. A stability window is true either way: it confirms the shard + // has settled on one role rather than being mid-pass, which is all the + // release needs. + stable := time.Now().Add(3 * time.Second) + for time.Now().Before(stable) { + shard := &Shard{} + c.NoError(c.Get(key, shard), "read the shard while registration is held") + c.Eq(1, len(shard.Status.PodRoles), + "a held pooler reached status.podRoles, so the hold is not holding") + time.Sleep(250 * time.Millisecond) + } + + // Now the pooler appears, with no Kubernetes event to announce it: a + // registration is a write to etcd. Only a requeue the operator asked for + // itself can notice, which is the thing this pins. + release() + + // Retire this pin by replacing it with c.Eventually on the same condition. + c.KnownDefect("MGO-POOL-SCALEUP-ROLE-STALE", func() error { + var members Members + deadline := time.Now().Add(20 * time.Second) + for { + var err error + members, err = MembersOf(c.Context(), c.Client(), key) + // A read failure is the check's own setup failing, not the + // defect, so it fails the test rather than keeping the pin green. + c.NoError(err, "read shard members") + if len(members.Replicas) == 1 && len(members.Quarantined) == 0 { + return nil + } + if time.Now().After(deadline) { + break + } + time.Sleep(250 * time.Millisecond) + } + return fmt.Errorf( + "the scaled-up pooler registered after the shard converged and its role "+ + "never reached status.podRoles within 20s: got %+v", members) + }) +} + +func (c *C) poolThrashCluster( + name string, + replicasPerCell int32, +) *MultigresCluster { + c.Helper() + return c.newCluster(name, func(s *MultigresClusterSpec) { + s.Databases = []DatabaseConfig{{ + Name: "postgres", + Default: true, + TableGroups: []TableGroupConfig{{ + Name: "default", + Default: true, + Shards: []ShardConfig{{ + Name: "0-inf", + Spec: &ShardInlineSpec{ + Pools: map[PoolName]PoolSpec{ + thrashPoolName: poolSpecWithReplicas(replicasPerCell), + }, + }, + }}, + }}, + }} + }) +} + +func poolSpecWithReplicas(n int32) PoolSpec { + return PoolSpec{ + Type: "readWrite", + Cells: []CellName{defaultSimCell}, + ReplicasPerCell: ptr.To(n), + } +} + +// scalePoolTo patches cluster's pool to n replicas per cell via a merge patch +// against cluster's own in-memory state, not a fresh read of the server. That +// makes each call in a back-to-back thrash loop independent of whatever the +// operator has done with the previous one: the patch always states the full +// desired pool spec, so it lands regardless of the server's current state. +func (c *C) scalePoolTo(cluster *MultigresCluster, n int32) { + c.Helper() + base := cluster.DeepCopy() + pools := cluster.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools + pools[thrashPoolName] = poolSpecWithReplicas(n) + c.NoError( + c.Patch(cluster, client.MergeFrom(base)), + "scale pool %s to %d replicas per cell", + thrashPoolName, + n, + ) +} + +// liveReadyPoolPodNames lists Ready pods belonging to the thrashed pool. +// +// This deliberately does not go through MembersOf/shard.Status.PodRoles: that +// path is exactly what MGO-POOL-SCALEUP-ROLE-STALE (see +// testPoolReplicaThrash) breaks, and a pod being live is a Kubernetes-level +// fact independent of whether the operator's own role bookkeeping has caught +// up to it. +func (c *C) liveReadyPoolPodNames() []string { + c.Helper() + pods := &corev1.PodList{} + c.NoError( + c.List(pods, client.MatchingLabels{metadata.LabelMultigresPool: string(thrashPoolName)}), + "list pool pods", + ) + var names []string + for i := range pods.Items { + if podReady(&pods.Items[i]) { + names = append(names, pods.Items[i].Name) + } + } + return names +} + +func podReady(p *corev1.Pod) bool { + for _, c := range p.Status.Conditions { + if c.Type == corev1.PodReady { + return c.Status == corev1.ConditionTrue + } + } + return false +} + +// requireNoOrphanedPoolPVCs asserts every PVC belonging to the thrashed pool +// is bound to one of livePods. It resolves the binding with ShardPVCOf, reading it +// off each live pod's own volumes, rather than reconstructing a PVC name from +// the pool/cell/ordinal and asserting the two strings match: a pod bound to +// the wrong PVC would pass a name-arithmetic check and fail this one. +func (c *C) requireNoOrphanedPoolPVCs(livePods []string) { + c.Helper() + + bound := map[string]bool{} + for _, pod := range livePods { + pvcName, err := ShardPVCOf(c.Context(), c.Client(), c.NS, pod) + if err != nil { + c.Fatalf("resolve PVC bound to live pod %s: %v", pod, err) + } + bound[pvcName] = true + } + + pvcs := &corev1.PersistentVolumeClaimList{} + c.NoError( + c.List(pvcs, client.MatchingLabels{metadata.LabelMultigresPool: string(thrashPoolName)}), + "list pool PVCs", + ) + for _, pvc := range pvcs.Items { + if !bound[pvc.Name] { + c.Errorf( + "PVC %s belongs to pool %s but is not bound to any live pod; live pods: %v", + pvc.Name, thrashPoolName, livePods, + ) + } + } +} + +// testChildCRThrash deletes a Shard out from under its TableGroup and lets +// the parent recreate it, three times in a row, then requires the shard to +// reconverge and the namespace to go quiet. +// +// This is the likeliest spot in the wave to find a live defect: the +// ReadyForDeletion protocol between the shard and tablegroup controllers is +// already the subject of two filed defects (a vacuous ReadyForDeletion, and +// whole-cluster teardown skipping the drain), and it found a third here, pinned +// below. +func testChildCRThrash(t *testing.T) { + c := newCase(t) + ns := c.NS + cluster := c.MinimalCluster("cr-thrash") + c.WaitForClusterHealthy(cluster) + + key := c.shardKey() + shard := &Shard{} + + for i := 0; i < 3; i++ { + c.NoError(c.Get(key, shard), "cycle %d: get shard %s", i, key.Name) + oldUID := shard.UID + c.NoError(c.Delete(shard), "cycle %d: delete shard %s", i, key.Name) + + // The only wait in this loop: for the parent to have recreated a + // replacement (a new UID at the same name), which is a mechanical + // precondition for the next delete to hit a live object rather than a + // no-op against one already gone. It is not a wait for the replacement + // to converge, and the loop does not wait for that before deleting + // again: that is the thrash. + what := fmt.Sprintf("the tablegroup to recreate Shard %s after delete #%d", key.Name, i+1) + c.Eventually(30*time.Second, what, func() error { + got := &Shard{} + if err := c.Get(key, got); err != nil { + return err + } + if got.UID == oldUID { + return fmt.Errorf("shard %s not yet recreated", key.Name) + } + return nil + }) + } + + c.NoError(c.Get(key, shard), "get final shard incarnation") + + c.WaitForClusterHealthy(cluster) + + c.Eventually(60*time.Second, "the shard to report one primary and one replica", + func() error { + members, err := MembersOf(c.Context(), c.Client(), key) + if err != nil { + return err + } + if len(members.Replicas) != 1 || len(members.Quarantined) != 0 { + return fmt.Errorf( + "want 1 primary + 1 replica + 0 quarantined, got %+v", members, + ) + } + return nil + }, + ) + + c.RequireQuiescent(10*time.Second, 90*time.Second) + + // A second live defect, found by this test: the shared backup PVC never + // has its orphan label cleared when a torn-down Shard's replacement + // reclaims it. + // + // reconcileSharedBackupPVC (reconcile_shared_infra.go) reapplies the PVC by + // server-side apply from BuildSharedBackupPVC's payload, which never + // mentions multigres.com/orphan-since, so SSA leaves that label exactly as + // cleanupShardPVCs (reconcile_deletion.go) left it during the prior + // teardown: marked orphan. Contrast the per-pool data PVC path, which + // explicitly calls pvcutil.ClearOrphan on reuse + // (reconcile_pool_pods.go:204). The shared backup PVC has no equivalent + // call anywhere in the shard controller. + // + // Net effect: after any teardown-and-recreate of a Shard whose backup PVC + // survives (WhenDeleted=Delete, which MinimalCluster sets, still only + // orphans rather than deletes it in-line, because resolvePodIndex cannot + // parse an ordinal out of a backup PVC's name-hash suffix and the !hasIndex + // arm short-circuits before pvcOrphanReplicasThreshold is consulted at + // all), the backup PVC is left labeled orphan + // forever, even though it is immediately reclaimed and stays in active use + // by the reconverged, healthy shard. The multigres-gc CronJob acts on + // exactly that label, so in a real cluster this is a live backup volume + // scheduled for deletion out from under a running shard. + backupPVCKey := client.ObjectKey{ + Namespace: ns, + Name: shardcontroller.BuildSharedBackupPVCName(shard), + } + pvc := &corev1.PersistentVolumeClaim{} + c.NoError( + c.Get(backupPVCKey, pvc), + "get shared backup PVC %s", + backupPVCKey.Name, + ) + c.KnownDefect("MGO-BACKUP-PVC-ORPHAN-STALE", func() error { + since, stale := pvc.Labels[metadata.LabelOrphan] + if !stale { + return nil + } + return fmt.Errorf( + "shared backup PVC %s still carries %s=%s from an earlier teardown, though the "+ + "shard that owns it (uid %s) has reconverged healthy: reconcileSharedBackupPVC's "+ + "server-side apply never clears the label on reuse, unlike the per-pool data PVC "+ + "path (pvcutil.ClearOrphan in reconcile_pool_pods.go)", + backupPVCKey.Name, metadata.LabelOrphan, since, shard.UID, + ) + }) +} diff --git a/test/suite/scenario_transitions_test.go b/test/suite/scenario_transitions_test.go new file mode 100644 index 00000000..74a5986d --- /dev/null +++ b/test/suite/scenario_transitions_test.go @@ -0,0 +1,466 @@ +package suite + +import ( + "fmt" + "maps" + "strings" + "testing" + "time" + + corev1 "k8s.io/api/core/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + "github.com/multigres/testkit/ctrltest" +) + +// TestTransitions is the round-trip suite: for a handful of optional +// MultigresCluster fields, set the field, then unset it, and require that the +// cluster has nothing left to do. No subtest enumerates what cleanup it +// expects to see; the closing assertion fails on any activity at all, whatever +// it is, which is what finds a leak without anyone having to name it first. +// +// That closing assertion takes one of two forms below, and the choice is per +// field rather than stylistic. Where the field's consequence lands on a plain +// corev1 kind whose event count is deterministic, it is a Script step +// permitting nothing but Quiet(). Where it does not, it is RequireQuiescent, +// which makes the same closed-world claim over the twelve kinds of +// watchedKinds() plus the write recorder rather than over the one or two kinds +// a Script could usefully watch, and which opens and closes its own watch +// instead of needing one opened before the fixture exists. A Script opened +// late over two kinds for a second is strictly weaker than the primitive it +// would be standing in for, so a field that cannot use the first form uses the +// second rather than a token version of the first. +// +// Each subtest also compares state directly across the round trip, which is +// not redundant with either form. An object created during the "set" and then +// left alone emits no event on the way back, so residue that manifests as the +// absence of an expected deletion is invisible to an event stream by +// construction, whichever kinds it watches. +// +// How far that comparison reaches differs per subtest, and none of them reach +// every kind. Backup compares PVC names and storage requests, DurabilityPolicy +// compares the TableGroup and Shard mirrors, and PVCDeletionPolicy compares its +// own field plus the namespace's PVC requests. Residue of a kind a subtest does +// not read still passes it: an orphaned ConfigMap or Service would pass all +// three. +// +// The field list came from api/v1alpha1/multigrescluster_types.go rather than +// from this task's brief, as instructed. Three of the four named fields exist +// as optional fields on MultigresClusterSpec and are covered below: +// PVCDeletionPolicy, Backup and DurabilityPolicy. The fourth, "a pool's +// Replicas", does not exist under that name: PoolSpec (shard_types.go) has no +// Replicas field, only ReplicasPerCell *int32. Per this task's brief, a field +// that does not exist under the given name is reported rather than silently +// substituted, so there is no fourth subtest here. +func TestTransitions(t *testing.T) { + t.Run("PVCDeletionPolicy", testPVCDeletionPolicyRoundTrip) + t.Run("Backup", testBackupRoundTrip) + t.Run("DurabilityPolicy", testDurabilityPolicyRoundTrip) +} + +// updateCluster applies mutate to a fresh read of the cluster and writes it +// back, failing the test rather than returning an error: a write that does not +// land is this file's own setup breaking, never an observation about the +// operator. +func (c *C) updateCluster( + cluster *MultigresCluster, + mutate func(*MultigresCluster), +) { + c.Helper() + got := &MultigresCluster{} + c.NoError(c.Get(client.ObjectKeyFromObject(cluster), got), "get cluster") + mutate(got) + c.NoError(c.Update(got), "update cluster") +} + +// pvcRequests reads every PVC in ns with the storage request it carries, for +// the snapshot-and-compare half of a round trip. Comparing the whole map +// catches a PVC that appeared, a PVC that went away and was never recreated, +// and a request that moved and stayed moved, none of which the event stream +// can report once the object stops changing. +func (c *C) pvcRequests() map[string]string { + c.Helper() + pvcs := &corev1.PersistentVolumeClaimList{} + c.NoError(c.List(pvcs), "list PVCs") + out := make(map[string]string, len(pvcs.Items)) + for _, pvc := range pvcs.Items { + out[pvc.Name] = pvc.Spec.Resources.Requests.Storage().String() + } + return out +} + +// shardBackupSizes reads the backup storage size each Shard has resolved. That +// is the value BuildSharedBackupPVC applies the shared backup PVC from +// (pool_pvc.go), so it is where a Backup write has to arrive for the PVC to +// see it. +func (c *C) shardBackupSizes() map[string]string { + c.Helper() + shards := &ShardList{} + c.NoError(c.List(shards), "list Shards") + if len(shards.Items) == 0 { + c.Fatalf("no Shards in %s to read a resolved backup size from", c.NS) + } + out := make(map[string]string, len(shards.Items)) + for _, shard := range shards.Items { + size := "" + if shard.Spec.Backup != nil && shard.Spec.Backup.Filesystem != nil { + size = shard.Spec.Backup.Filesystem.Storage.Size + } + out[shard.Name] = size + } + return out +} + +// durabilityMirrors reads the DurabilityPolicy every TableGroup and Shard in +// ns currently carries. The field's only other consumer is the topology store, +// which this suite fakes in memory (fakes.go) and which writes a policy of its +// own regardless, so these mirrored spec fields are the whole of what this +// field does that anything here can observe. +func (c *C) durabilityMirrors() map[string]string { + c.Helper() + out := map[string]string{} + tgs := &TableGroupList{} + c.NoError(c.List(tgs), "list TableGroups") + for _, tg := range tgs.Items { + out["TableGroup/"+tg.Name] = tg.Spec.DurabilityPolicy + } + shards := &ShardList{} + c.NoError(c.List(shards), "list Shards") + for _, shard := range shards.Items { + out["Shard/"+shard.Name] = shard.Spec.DurabilityPolicy + } + if len(out) == 0 { + c.Fatalf("no TableGroups or Shards in %s to read a DurabilityPolicy from", c.NS) + } + return out +} + +// testPVCDeletionPolicyRoundTrip watches only PersistentVolumeClaim. +// +// PVC is a plain corev1 kind with no status.conditions of its own, so writing +// to it never sets off the generation-bump status-condition churn that +// TableGroup and Shard produce on every spec write (see +// testDurabilityPolicyRoundTrip for where that churn made a Step-based +// assertion unusable). PVCDeletionPolicy's real consequence, +// reconcilePVCOwnerRefs (pkg/resource-handler/controller/shard/shard_controller.go), +// lands on PVCs directly, so this narrower watch still sees it, and nothing +// wider is needed to catch a leak here. +func testPVCDeletionPolicyRoundTrip(t *testing.T) { + c := newCase(t) + sc := c.NewScript(&corev1.PersistentVolumeClaimList{}) + sc.StepTimeout = 10 * time.Second + + cluster := c.MinimalCluster("pvcdp") + c.WaitForClusterHealthy(cluster) + + // Snapshotted so the round trip is checked against the namespace's PVCs and + // not only against the field's own value. A PVC created during the set and + // then left alone emits nothing on the way back, so the closing Quiet step + // cannot see it. + pvcsAtStart := c.pvcRequests() + c.RequireQuiescent(5*time.Second, 30*time.Second) + + // The script opened before MinimalCluster, per NewScript's own contract, + // so its watch replays every PVC the fixture created as an Added event + // ("an already-populated namespace is replayed as a run of added + // events"). Permitting that baseline explicitly, by listing what actually + // exists now that convergence is independently confirmed, is the + // documented way to handle it, and it is a statement about the fixture, + // not a guess about this field's behaviour. + pvcs := &corev1.PersistentVolumeClaimList{} + c.NoError(c.List(pvcs), "list PVCs") + c.NotEmpty(pvcs.Items, "MinimalCluster created no PVCs to test PVCDeletionPolicy against") + baseline := make([]ctrltest.Allow, 0, len(pvcs.Items)*2) + pvcNames := make([]string, 0, len(pvcs.Items)) + for _, pvc := range pvcs.Items { + baseline = append(baseline, + ctrltest.Added("PersistentVolumeClaim", pvc.Name), + ctrltest.Changed("PersistentVolumeClaim", pvc.Name)) + pvcNames = append(pvcNames, pvc.Name) + } + sc.Step("PVCs created and bound during initial convergence", nil, baseline...) + + sc.Step("cluster settled", nil, ctrltest.Quiet()) + + original := cluster.Spec.PVCDeletionPolicy.DeepCopy() + + consequences := make([]ctrltest.Allow, 0, len(pvcNames)) + for _, name := range pvcNames { + consequences = append(consequences, ctrltest.Changed("PersistentVolumeClaim", name)) + } + + sc.Step("set PVCDeletionPolicy to Retain/Retain", func() error { + got := &MultigresCluster{} + if err := c.Get(client.ObjectKeyFromObject(cluster), got); err != nil { + return err + } + got.Spec.PVCDeletionPolicy = &PVCDeletionPolicy{ + WhenDeleted: multigresv1alpha1.RetainPVCRetentionPolicy, + WhenScaled: multigresv1alpha1.RetainPVCRetentionPolicy, + } + return c.Update(got) + }, consequences...) + + sc.Step("settled after setting the field", nil, ctrltest.Quiet()) + + sc.Step("unset PVCDeletionPolicy", func() error { + got := &MultigresCluster{} + if err := c.Get(client.ObjectKeyFromObject(cluster), got); err != nil { + return err + } + got.Spec.PVCDeletionPolicy = nil + return c.Update(got) + }, consequences...) + + sc.Step("nothing left to do after the round trip", nil, ctrltest.Quiet()) + + sc.Finish(time.Second) + + // PVCDeletionPolicy carries a CRD-level +kubebuilder:default, so a nil + // pointer never survives the API server: MinimalCluster's explicit + // {Delete, Delete} and an omitted field both resolve to the same stored + // value. The round trip is checked against what the field actually reads + // back as, not against a literal nil. + final := &MultigresCluster{} + c.NoError(c.Get(client.ObjectKeyFromObject(cluster), final), "get final cluster") + c.Check().EqDiff(original, final.Spec.PVCDeletionPolicy, + "PVCDeletionPolicy round trip not lossless") + c.Check().EqDiff(pvcsAtStart, c.pvcRequests(), + "PVCs differ across the round trip") +} + +// testBackupRoundTrip takes Backup out to an explicit backup storage size and +// back. It is the subtest that found a defect, and the pin below is the +// finding rather than an aside. +// +// Backup's only consequence any watch in this suite can see is the shared +// backup PVC's storage request (pool_pvc.go). The rest of what the field +// feeds is the topology store, faked in memory here (fakes.go), or needs a +// Secret the caller precreates (reconcile_shared_infra.go). So the round trip +// is driven through that size, and the size is where the operator breaks. +// +// Which size, measured against this harness rather than assumed, because once +// the PVC is bound and the data-plane fake has copied its request into +// status.capacity (datasim.go) every candidate is refused for a different +// reason: +// +// - smaller, 5Gi against the resolver's 10Gi default, is refused by PVC +// validation: "spec.resources.requests.storage: Forbidden: field can not +// be less than status.capacity". That rule is unconditional in Kubernetes, +// so any user of the operator can reach it. +// - larger, 20Gi, is refused by the PersistentVolumeClaimResize admission +// plugin: "only dynamically provisioned pvc can be resized and the +// storageclass that provisions the pvc must support resize", because this +// suite creates no StorageClass. A real cluster whose class sets +// allowVolumeExpansion would accept it, so that refusal is an artifact of +// the fixture and is deliberately not what gets pinned here. +// - equal but written in another unit, 10240Mi, is canonicalised back to +// 10Gi by the API server and bumps no resourceVersion, so it is invisible +// rather than illegal. +// +// The shrink is therefore the transition worth driving: legal on the +// MultigresCluster, propagated all the way to the PVC apply, and refused +// there. +func testBackupRoundTrip(t *testing.T) { + c := newCase(t) + ns := c.NS + cluster := c.MinimalCluster("backup") + c.WaitForClusterHealthy(cluster) + c.RequireQuiescent(5*time.Second, 30*time.Second) + + pvcsBefore := c.pvcRequests() + sizesBefore := c.shardBackupSizes() + original := cluster.Spec.Backup.DeepCopy() + + c.updateCluster(cluster, func(c *MultigresCluster) { + c.Spec.Backup = &multigresv1alpha1.BackupConfig{ + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Storage: multigresv1alpha1.StorageSpec{Size: "5Gi"}, + }, + } + }) + + // The defect: the refused PVC apply returns an error from Reconcile + // (shard_controller.go), so everything after that block is skipped for as + // long as the field stays lowered, the postgres ConfigMap render, the pool + // pods, PDB sizing and reconcilePVCOwnerRefs, and controller-runtime + // retries forever. The status update is NOT skipped: updateStatus runs + // early in Reconcile, deliberately, so a wedged shard keeps reporting a + // current observedGeneration and phase. That is what makes this defect + // silent, and it is why a triage that looks for a stale status will not + // find one. A user who lowers this field wedges + // the shard's whole reconcile loop, with no clamping, no rejection at + // admission, and no condition on the Shard saying why. + // + // Pinned as non-quiescence rather than as a missing PVC event, because a + // missing event is what the API server's refusal guarantees on every run + // forever: no operator change could ever retire that pin, and a pin that + // cannot expire is a suppression. This one expires the day the operator + // clamps the value, because then the namespace goes quiet and quiescent + // returns nil. + // + // Two neighbouring fixes do not expire it through quiescence, and the + // difference is worth knowing before trusting this pin as a tracker. An + // admission refusal fails the test earlier, at the update call, so the + // suite still goes red but by another route. A fix that gives up and + // records a terminal condition only after retrying past this horizon's + // activity budget leaves the pin green, so that one has to be noticed by a + // human reading this comment rather than by the pin flipping. + // + // What this pin rests on, measured 2026-09-18 rather than assumed, because + // it used to rest on something weaker than it looked. The wedged pass makes + // several accepted no-op writes (both pg_hba and exporter-queries + // ConfigMaps, the multiorch Deployment and Service, a Shard status patch) + // before it reaches the refused PVC apply, so before the recorder counted + // rejected writes, this pin observed non-quiescence only through those + // earlier writes. Reordering updateStatus, a refactor with no behavioural + // intent, would have flipped it to "appears fixed" with the defect fully + // present. The recorder now counts the refused apply itself, which is the + // one signal the defect cannot occur without, so that particular + // reordering can no longer fool it. + // + // It is not yet true that the rejection alone carries the pin. Counting + // only rejected writes, the same wedge measured non-quiescent on one run + // and quiet on the next, 12 refused applies being right at the edge of a + // 5s window inside an 11s horizon as the retry backoff spreads them out. + // So the margin still comes from the signals combined. Widening the + // horizon is not the fix (see below); if this pin ever needs to stand on + // the rejection by itself, count refused reconcile passes directly rather + // than inferring them from a quiet window. + // + // The horizon is short on purpose. controller-runtime retries a failing + // Reconcile with exponential backoff, so the gap between failing passes + // grows without bound and a long enough horizon would let the backoff + // itself supply the quiet window while the shard is still wedged. + // Calibrated both ways on this harness, three runs each: measured from + // immediately after a legal spec write this goes quiet in about 5.3s, + // measured from immediately after this write it never goes quiet inside + // 11s, over 12 failing reconcile passes. + c.KnownDefect("MGO-BACKUP-PVC-SHRINK-WEDGES-SHARD-RECONCILE", func() error { + err := c.TryQuiescent(5*time.Second, 11*time.Second) + // A lost watch voids the measurement instead of observing the + // defect, and a non-nil error here is read as the defect still + // being present, which would hold this pin green on an unrelated + // failure. + c.True(err == nil || !strings.Contains(err.Error(), "is void"), + "quiescence over %s is void, so it is no evidence either way: %v", ns, err) + return err + }) + + // Where the value actually got to, as an executable claim rather than a + // comment, because the mechanism is easy to misread: it does reach the + // Shard, so the PVC is applied from the size the user asked for and the + // refusal happens at the API server. A fix aimed at the resolver's backup + // defaulting would land on code that is behaving correctly. + c.Eventually(15*time.Second, "the lowered size to reach every Shard", func() error { + for name, size := range c.shardBackupSizes() { + if size != "5Gi" { + return fmt.Errorf("Shard %s resolved backup size is %q", name, size) + } + } + return nil + }) + + c.updateCluster(cluster, func(c *MultigresCluster) { + c.Spec.Backup = nil + }) + + // The gate that makes the closing assertion mean something. The namespace + // is not quiet when the unset lands, so silence afterwards would be + // ambiguous between the operator having processed it and the retry backoff + // having merely grown past the window. The resolved size returning to + // where it started is positive evidence that the unset propagated back + // down to where the PVC is applied from. It is a wait, not the assertion. + c.Eventually(30*time.Second, "the unset to reach every Shard", func() error { + if got := c.shardBackupSizes(); !maps.Equal(got, sizesBefore) { + return fmt.Errorf("resolved backup sizes are %v, want %v", got, sizesBefore) + } + return nil + }) + + // The round trip's assertion: once the cluster is back to the + // configuration it started in, a converged operator has nothing left to + // do, and this is that claim over every kind the operator writes plus + // every write it makes. + c.RequireQuiescent(5*time.Second, 30*time.Second) + + c.Check().EqDiff(pvcsBefore, c.pvcRequests(), "PVCs did not round trip") + final := &MultigresCluster{} + c.NoError(c.Get(client.ObjectKeyFromObject(cluster), final), "get final cluster") + c.Check().EqDiff(original, final.Spec.Backup, "Backup round trip not lossless") +} + +// testDurabilityPolicyRoundTrip has no Script, and the reason is worth +// recording because the shape it would need does not exist in the runner. +// +// DurabilityPolicy is mirrored into TableGroup.Spec and Shard.Spec +// (builders_tablegroup.go, tablegroup/builders.go), and any spec write to +// either bumps its generation, which sends every controller watching it back +// to re-stamp its own status conditions' observedGeneration. That catch-up +// took a different number of passes on every run measured while writing this +// test: watching only TableGroup and Shard, the same single field write +// settled after 19, then 13, then 9 total events across three otherwise +// identical runs. A Step's allow-list is an exact multiset, with no "N events +// of this kind" wildcard available, so declaring one against a count that +// moves between runs would flake on this suite's own harness rather than on +// the operator, which is a worse failure than not writing the assertion at +// all. +// +// What is left for a Script to assert over those kinds is that nothing further +// happened, and RequireQuiescent asserts that strictly better: five seconds +// over twelve kinds plus the write recorder, rather than about a second over +// two, and it needs no watch opened before the fixture, which over these kinds +// is not possible to combine with an exact Step anyway. So the closing +// RequireQuiescent is this subtest's closed-world assertion, standing where +// the other form's Quiet() step stands. +func testDurabilityPolicyRoundTrip(t *testing.T) { + c := newCase(t) + cluster := c.MinimalCluster("durability") + c.WaitForClusterHealthy(cluster) + c.RequireQuiescent(5*time.Second, 30*time.Second) + + original := cluster.Spec.DurabilityPolicy + mirrorsBefore := c.durabilityMirrors() + + c.updateCluster(cluster, func(c *MultigresCluster) { + c.Spec.DurabilityPolicy = "MULTI_CELL_AT_LEAST_2" + }) + + // Not an expectation about cleanup, which this test writes none of. It is + // the guard that keeps the round trip from being vacuous: if setting the + // field moved nothing anywhere, unsetting it could not leak anything and + // the closing assertion would be proving nothing about this field. + c.Eventually(30*time.Second, "the set policy to reach every mirror", func() error { + for name, policy := range c.durabilityMirrors() { + if policy != "MULTI_CELL_AT_LEAST_2" { + return fmt.Errorf("%s carries %q", name, policy) + } + } + return nil + }) + c.RequireQuiescent(5*time.Second, 30*time.Second) + + c.updateCluster(cluster, func(c *MultigresCluster) { + c.Spec.DurabilityPolicy = original + }) + + // The mirrors returning is the state half of the round trip, and no event + // assertion can make it: a mirror left holding the set value emits nothing + // once it stops changing, so silence and correctness would be the same + // observation. Also the gate that the operator processed the unset before + // the assertion below asks for silence. + c.Eventually(30*time.Second, "the unset policy to reach every mirror", func() error { + if got := c.durabilityMirrors(); !maps.Equal(got, mirrorsBefore) { + return fmt.Errorf("mirrors are %v, want %v", got, mirrorsBefore) + } + return nil + }) + + c.RequireQuiescent(5*time.Second, 30*time.Second) + + final := &MultigresCluster{} + c.NoError(c.Get(client.ObjectKeyFromObject(cluster), final), "get final cluster") + c.Check().Eq(original, final.Spec.DurabilityPolicy, "DurabilityPolicy round trip not lossless") +} diff --git a/test/suite/shard_requeue_test.go b/test/suite/shard_requeue_test.go new file mode 100644 index 00000000..6cd6311d --- /dev/null +++ b/test/suite/shard_requeue_test.go @@ -0,0 +1,183 @@ +package suite + +import ( + "fmt" + "testing" + "time" + + "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/ctrltest" +) + +// certBootstrapBudget is how long a test will wait for the shard controller to +// finish generating pgBackRest's CA and server certificates. +// +// It is the one wait in this package deliberately larger than 30s, and it is +// sized against the step it waits on rather than against convergence. +// reconcilePgBackRestCerts generates two RSA keys, which is the most expensive +// thing this suite does and by far the most variable under -race: measured at +// 14s in one run of the suite and 75s in another, on the same machine. A +// number sized for the 14s case turns the 75s case into a red suite that says +// nothing about the operator. +const certBootstrapBudget = 2 * time.Minute + +// waitForShardPastPKI blocks until the shard controller has written a pool Pod +// in ns, which is the evidence this suite has that the controller is past +// certificate generation. +// +// It exists to keep a slow precondition out of an assertion's budget. The +// shard reconciles fourteen steps in a fixed order: the pgBackRest certificate +// step is fifth, reconcilePool is twelfth, and reconcileDataPlane, the only +// step that returns the pooler-registration requeue, is last. A test that +// starts its requeue budget before the keys exist is timing RSA keygen, and it +// goes red when the crypto was slow rather than when the operator was wrong. +// Split in two, each wait is sized against what it actually waits for. +// +// A pool Pod is the signal rather than the certificate Secrets themselves for +// two reasons. It sits between the two steps that matter, so it proves the +// keygen is behind us without depending on it being the immediately preceding +// step. And a Shard creating Pods for its pools is a more stable fact than the +// name the operator builds its Secrets from, which a test has no business +// knowing. It is read from the recorder rather than from the apiserver so that +// what is observed is the shard controller having written, not an object that +// something else could have created. +func (c *C) waitForShardPastPKI() { + c.Helper() + c.Eventually( + certBootstrapBudget, + "the shard controller to get past certificate generation", + func() error { + for _, op := range Suite.Ops.OpsInNamespace(c.NS) { + if op.Controller == "shard" && ctrltest.KindSuffix(op.Kind) == "Pod" { + return nil + } + } + return fmt.Errorf("the shard controller has written no pool Pod yet") + }, + ) +} + +// holdPoolerRegistration keeps every pooler in this case's namespace out of +// the topology store for the rest of the test. +// +// The shard controller only asks for its one-minute requeue when it finds no +// poolers registered at all. Left to the data-plane fake, which registers on +// a 500ms tick as soon as a pool pod exists, whether the shard ever sees that +// empty topology is a race: when the first pooler registers before the +// shard's first data-plane pass, the shard goes straight to "some poolers", +// skips the requeue these tests wait for, and on today's operator returns no +// requeue at all (MGO-POOL-SCALEUP-ROLE-STALE). Measured at 1 run in 20 under +// full-suite load. Holding registration makes the empty topology the only +// state the shard can see. +func (c *C) holdPoolerRegistration() { + c.Helper() + c.Cleanup(poolers.HoldRegistrations(c.NS)) +} + +// awaitShard waits for the cluster controller to create this namespace's one +// Shard and returns it. +// +// Exactly one, not at least one. Both polls this replaces went on to read +// Items[0], and the weaker form picked element zero out of a set whose size +// it never checked. Every fixture in this package declares a single +// ShardConfig and suite_test.go asserts that, so the stronger claim is true +// today and costs nothing to state. +// +// It is what happens when that stops being true that decides it. Sharding is +// this project's whole point, so a multi-shard fixture is a matter of time, +// and at that moment "at least one" stays green while silently asserting +// about whichever Shard the API server happened to return first. "Exactly +// one" fails saying it got two, which points at the fixture that changed. A +// test that wants a particular shard out of several should name it rather +// than index into a list. +func (c *C) awaitShard() Shard { + c.Helper() + var shard Shard + c.Eventually(30*time.Second, "the cluster's one Shard to exist", func() error { + shards := &ShardList{} + if err := c.List(shards); err != nil { + return err + } + if len(shards.Items) != 1 { + return fmt.Errorf("want exactly one Shard, got %d", len(shards.Items)) + } + shard = shards.Items[0] + return nil + }) + return shard +} + +// shardKey is awaitShard's key, for the callers that only need to address it. +func (c *C) shardKey() client.ObjectKey { + c.Helper() + shard := c.awaitShard() + return client.ObjectKeyFromObject(&shard) +} + +// TestShardAsksForAMinuteAwaitingPoolerRegistration is the assertion that +// replaces waiting a minute for the same information, and the reason requeue +// compression does not hide the defect it compresses. +// +// While no multipooler has registered in the topology store the shard +// controller asks to be woken in a minute. Nothing in Kubernetes watches that +// store, so in production nothing can wake it sooner, and before compression +// every convergence test in this package paid that minute whenever the last pod +// event happened to land before the data plane fake registered its poolers. +// Which test paid was a coin flip and the suite's wall time swung by a minute +// between runs for no visible reason. +// +// Compressed, the poll comes back in 50ms and the minute survives here as a +// fact about the operator. Fixing it is not this suite's job, and this +// assertion is what should fail when somebody does fix it. +// +// This test asserts about the operator and nothing else. That the suite clamps +// what it saw here is a fact about the harness and belongs to +// TestSuiteCompressesRequeues, so that a clamp change reports itself as a +// clamp change rather than as the shard controller's polling having moved. +func TestShardAsksForAMinuteAwaitingPoolerRegistration(t *testing.T) { + c := newCase(t) + c.holdPoolerRegistration() + c.MinimalCluster("requeue") + + key := c.shardKey() + c.waitForShardPastPKI() + + got := Suite.Reconciles.WaitForRequeue(t, "shard", key, time.Minute, 30*time.Second) + + // Exactly a minute, not merely at least a minute. One minute is the shard + // controller's only requeue of that length (poolerRegistrationRetryDelay + // in reconcile_data_plane.go), so the duration identifies the code path + // that the reconcile boundary itself cannot: ctrl.Result carries no + // reason, and the AwaitingPoolerRegistration reason lives on the shard's + // PostureConsistent condition, which by the time a test can read it has + // usually already moved on. + c.Check().Eq(time.Minute, got.RequestedAfter, "the shard's requeue duration") +} + +// TestSuiteCompressesRequeues is the canary on the harness half of the +// bargain, and it is live rather than a unit test for a reason the unit tests +// cannot cover: TestInterceptorCompressesRequeue proves that an interceptor +// clamps, not that this suite wired one around the real controllers. If that +// wiring is ever broken, compression dies silently, every timeout in the +// package regains its old "maybe it is just waiting" ambiguity, and nothing +// fails except runtimes nobody reads. +// +// It deliberately does not name a duration the operator chose. Any requeue +// longer than the clamp will do, so that fixing the shard controller's minute +// changes what this test observes but not whether it passes. +func TestSuiteCompressesRequeues(t *testing.T) { + c := newCase(t) + c.holdPoolerRegistration() + c.MinimalCluster("clamp") + + key := c.shardKey() + c.waitForShardPastPKI() + + got := Suite.Reconciles.WaitForRequeue(t, "shard", key, 2*ctrltest.RequeueClamp, 30*time.Second) + + c.Check().Eq(ctrltest.RequeueClamp, got.Result.RequeueAfter, + "controller-runtime's requeue-after duration") + c.Check(). + True(got.Compressed(), "Compressed() = false on a pass that asked for %s", got.RequestedAfter) +} diff --git a/test/suite/suite.go b/test/suite/suite.go new file mode 100644 index 00000000..428d78a3 --- /dev/null +++ b/test/suite/suite.go @@ -0,0 +1,223 @@ +// Package suite is the multi-controller envtest harness for this operator: one +// manager running every reconciler the operator runs, so behaviour that is a +// protocol between controllers becomes testable. +// +// The generic half lives in pkg/ctrltest. What is left here is everything that +// is about this operator specifically: its scheme, its CRDs, its cache config, +// its five reconcilers, and the data plane doubles those reconcilers need. +// +// It has no build tag. Exclusion from the ordinary test targets is by path +// filter, and the entrypoint is `make test-suite`. +package suite + +import ( + "context" + "fmt" + "path/filepath" + "time" + + "github.com/multigres/multigres/go/common/rpcclient" + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + networkingv1 "k8s.io/api/networking/v1" + policyv1 "k8s.io/api/policy/v1" + storagev1 "k8s.io/api/storage/v1" + "k8s.io/apimachinery/pkg/runtime" + clientgoscheme "k8s.io/client-go/kubernetes/scheme" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/controller" + "sigs.k8s.io/controller-runtime/pkg/manager" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + "github.com/multigres/multigres-operator/pkg/cacheopts" + multigresclustercontroller "github.com/multigres/multigres-operator/pkg/cluster-handler/controller/multigrescluster" + tablegroupcontroller "github.com/multigres/multigres-operator/pkg/cluster-handler/controller/tablegroup" + "github.com/multigres/multigres-operator/pkg/data-handler/poolerclient" + cellcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/cell" + shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" + toposervercontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/toposerver" + "github.com/multigres/testkit/ctrltest" +) + +// OperatorNamespace stands in for the namespace the operator deploys into. The +// cache treats it specially (unfiltered), so the suite has to have one for the +// production cache config to mean anything. +const OperatorNamespace = "multigres-operator-system" + +// Suite is the suite-wide harness, booted once by TestMain. One envtest and one +// manager serve the whole package; isolate with Suite.Namespace(t). +var Suite *ctrltest.Suite + +// The operator's data plane doubles. Deliberately not suite members: every +// operator's data plane is different, so ctrltest has no place to put these and +// a test reaching for them is reaching for something about this operator. +var ( + rpc *rpcclient.FakeClient + topo *topoRegistry + poolers *poolerSim +) + +// Boot brings up envtest and the manager. The returned function tears both down +// and must run before any goroutine leak check, since the manager owns +// goroutines that only exit once its context is cancelled. +func Boot() (*ctrltest.Suite, func() error, error) { + scheme := runtime.NewScheme() + for _, add := range []func(*runtime.Scheme) error{ + clientgoscheme.AddToScheme, + multigresv1alpha1.AddToScheme, + appsv1.AddToScheme, + corev1.AddToScheme, + policyv1.AddToScheme, + networkingv1.AddToScheme, + storagev1.AddToScheme, + } { + if err := add(scheme); err != nil { + return nil, nil, fmt.Errorf("add to scheme: %w", err) + } + } + + return ctrltest.Boot(ctrltest.Options{ + Scheme: scheme, + CRDPaths: []string{filepath.Join("..", "..", "config", "crd", "bases")}, + WatchedKinds: watchedKinds(), + SimInterval: 250 * time.Millisecond, + Managers: []ctrltest.ManagerOptions{{ + Name: "multigres-operator", + CacheOptions: cacheopts.New(OperatorNamespace), + OperatorNamespace: OperatorNamespace, + // Matches main.go: the operator raises these to avoid client-side + // throttling once several controllers are reconciling at once. + QPS: 50, + Burst: 100, + Register: register, + }}, + }) +} + +// watchedKinds is every kind the operator writes in a test namespace. A kind +// missing here is a kind whose churn RequireQuiescent cannot see. +func watchedKinds() []client.ObjectList { + return []client.ObjectList{ + &multigresv1alpha1.MultigresClusterList{}, + &TopoServerList{}, + &multigresv1alpha1.CellList{}, + &TableGroupList{}, + &ShardList{}, + &corev1.PodList{}, + &appsv1.DeploymentList{}, + &appsv1.StatefulSetList{}, + &corev1.PersistentVolumeClaimList{}, + &corev1.ConfigMapList{}, + &corev1.ServiceList{}, + &policyv1.PodDisruptionBudgetList{}, + } +} + +// register wires every reconciler exactly as cmd/multigres-operator/main.go +// does, differing only in the seams a test has to fake: the topology store and +// the multipooler RPC client, and in the interceptor each one's reconcile +// boundary is wrapped in. +// +// Each reconciler is a named variable rather than the anonymous composite +// literal this used to be, and that is load bearing rather than tidying. +// SetupWithManagerReconciler substitutes the reconcile boundary and nothing +// else: the receiver stays live on the enqueue path, because map functions and +// predicates bind to it when the builder runs and call its client, and +// ShardReconciler keeps mutable state on itself. So the same pointer has to be +// both the receiver and what the interceptor delegates to. +func register(ctx context.Context, mgr manager.Manager, s *ctrltest.Suite) error { + base := mgr.GetClient() + + rpc = rpcclient.NewFakeClient() + topo = newTopoRegistry(ctx) + + // The pooler fake runs for the life of the suite, across every namespace, + // because the manager it feeds is also suite-wide. + poolers = &poolerSim{c: s.Client, rpc: rpc, topo: topo, interval: 500 * time.Millisecond} + go poolers.run(ctx) + + // Deliberately a bare option struct rather than each controller's + // production options. Every controller sets MaxConcurrentReconciles to 20 + // in its own SetupWithManager and then lets a caller-supplied + // controller.Options replace that wholesale, so this suite's options + // decide the value; it is set to 1 here rather than inherited by omission + // from the controller-runtime default, which is what used to happen. + // + // That divergence from production is load-bearing in both directions. + // It is what makes the recorder's per-controller ordering assertions + // meaningful: one reconcile goroutine per controller means writes are + // issued and recorded in program order. It is also what this suite + // therefore cannot catch, namely a controller racing itself across + // concurrent reconciles of different objects. Raising this to match + // production would silently turn every ordering assertion in the + // scenario tests into a race that fails a few times a week and reads as operator + // flakiness, and it would also break Interceptor.Ops, which identifies a + // pass's writes by an op log range that only one in-flight reconcile per + // controller can make unambiguous. + opts := controller.Options{ + SkipNameValidation: ptr.To(true), + MaxConcurrentReconciles: 1, + } + + cluster := &multigresclustercontroller.MultigresClusterReconciler{ + Client: s.Ops.For("multigrescluster", base), + Scheme: mgr.GetScheme(), + Recorder: mgr.GetEventRecorderFor("multigrescluster-controller"), + APIReader: mgr.GetAPIReader(), + CreateTopoStore: topo.ForClusterRef, + } + if err := cluster.SetupWithManagerReconciler( + mgr, s.Reconciles.Wrap("multigrescluster", cluster), opts, + ); err != nil { + return fmt.Errorf("setup multigrescluster: %w", err) + } + + tableGroup := &tablegroupcontroller.TableGroupReconciler{ + Client: s.Ops.For("tablegroup", base), + Scheme: mgr.GetScheme(), + Recorder: mgr.GetEventRecorderFor("tablegroup-controller"), + } + if err := tableGroup.SetupWithManagerReconciler( + mgr, s.Reconciles.Wrap("tablegroup", tableGroup), opts, + ); err != nil { + return fmt.Errorf("setup tablegroup: %w", err) + } + + cell := &cellcontroller.CellReconciler{ + Client: s.Ops.For("cell", base), + Scheme: mgr.GetScheme(), + Recorder: mgr.GetEventRecorderFor("cell-controller"), + } + if err := cell.SetupWithManagerReconciler( + mgr, s.Reconciles.Wrap("cell", cell), opts, + ); err != nil { + return fmt.Errorf("setup cell: %w", err) + } + + topoServer := &toposervercontroller.TopoServerReconciler{ + Client: s.Ops.For("toposerver", base), + Scheme: mgr.GetScheme(), + Recorder: mgr.GetEventRecorderFor("toposerver-controller"), + } + if err := topoServer.SetupWithManagerReconciler( + mgr, s.Reconciles.Wrap("toposerver", topoServer), opts, + ); err != nil { + return fmt.Errorf("setup toposerver: %w", err) + } + + shard := &shardcontroller.ShardReconciler{ + Client: s.Ops.For("shard", base), + Scheme: mgr.GetScheme(), + Recorder: mgr.GetEventRecorderFor("shard-controller"), + APIReader: mgr.GetAPIReader(), + PoolerClients: poolerclient.Static(rpc), + CreateTopoStore: topo.ForShard, + } + if err := shard.SetupWithManagerReconciler( + mgr, s.Reconciles.Wrap("shard", shard), opts, + ); err != nil { + return fmt.Errorf("setup shard: %w", err) + } + return nil +} diff --git a/test/suite/suite_test.go b/test/suite/suite_test.go new file mode 100644 index 00000000..2b8c8667 --- /dev/null +++ b/test/suite/suite_test.go @@ -0,0 +1,123 @@ +package suite + +import ( + "fmt" + "strings" + "testing" + "time" + + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" +) + +// TestClusterConvergesUnderAllControllers is the harness proof: one +// MultigresCluster, five live reconcilers, and a faked data plane, reaching a +// terminal healthy state. +// +// It asserts nothing about behaviour that single-controller tests already +// cover. Its job is to fail loudly if the harness itself stops working. +func TestClusterConvergesUnderAllControllers(t *testing.T) { + c := newCase(t) + ns := c.NS + cluster := c.MinimalCluster("minimal") + + c.Eventually(30*time.Second, "child CRs to be created", func() error { + topos := &TopoServerList{} + if err := c.List(topos); err != nil { + return err + } + cells := &multigresv1alpha1.CellList{} + if err := c.List(cells); err != nil { + return err + } + tgs := &TableGroupList{} + if err := c.List(tgs); err != nil { + return err + } + shards := &ShardList{} + if err := c.List(shards); err != nil { + return err + } + if len(topos.Items) == 0 || len(cells.Items) == 0 || + len(tgs.Items) == 0 || len(shards.Items) == 0 { + return fmt.Errorf("have %d TopoServer, %d Cell, %d TableGroup, %d Shard", + len(topos.Items), len(cells.Items), len(tgs.Items), len(shards.Items)) + } + return nil + }) + + c.WaitForClusterHealthy(cluster) + + // Attribution is the other half of the harness: a test that cannot say + // which controller wrote cannot assert a protocol between controllers. + wrote := map[string]bool{} + for _, op := range Suite.Ops.OpsInNamespace(ns) { + wrote[op.Controller] = true + } + for _, name := range []string{"multigrescluster", "cell", "toposerver", "tablegroup", "shard"} { + c.Check().True(wrote[name], "no recorded writes from the %s controller; "+ + "either it never ran or attribution is broken", name) + } +} + +// TestNamespacesAreIsolated runs two clusters at once to prove the isolation +// boundary holds, since every fake behind the suite (the topology store above +// all) is shared process-wide and keyed by namespace. +func TestNamespacesAreIsolated(t *testing.T) { + t.Parallel() + + for _, name := range []string{"iso-a", "iso-b"} { + t.Run(name, func(t *testing.T) { + t.Parallel() + c := newCase(t) + ns := c.NS + cluster := c.MinimalCluster(strings.TrimPrefix(name, "iso-")) + + c.WaitForClusterHealthy(cluster) + + shards := &ShardList{} + c.NoError(c.List(shards)) + c.Len(shards.Items, 1, "in %s", ns) + }) + } +} + +// TestProductionCacheConfigIsInEffect guards the fidelity trap: the manager +// must cache exactly what cmd/multigres-operator/main.go caches. +// +// Under the production config an unlabelled Secret outside the operator's own +// namespace is invisible to the cached client, which is why the reconcilers +// carry an APIReader at all. A manager built with default cache options makes +// every cached read behave differently from production and quietly voids the +// premise that these are the real controllers wired as in main.go. +// +// This assertion is only possible because the config lives in pkg/cacheopts +// rather than being copied out of package main, which cannot be imported. +func TestProductionCacheConfigIsInEffect(t *testing.T) { + c := newCase(t) + ns := c.NS + + secret := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{Name: "unlabelled", Namespace: ns}, + StringData: map[string]string{"password": "postgres"}, + } + c.NoError(c.Create(secret), "create secret") + key := client.ObjectKeyFromObject(secret) + + c.Eventually(30*time.Second, "the APIReader to see the secret", func() error { + return Suite.Manager("multigres-operator"). + Mgr.GetAPIReader(). + Get(c.Context(), key, &corev1.Secret{}) + }) + + err := Suite.Manager("multigres-operator"). + Mgr.GetClient(). + Get(c.Context(), key, &corev1.Secret{}) + c.True(apierrors.IsNotFound(err), + "cached client should not see an unlabelled Secret outside %s, got err=%v", + OperatorNamespace, err) +} diff --git a/test/suite/testdata/multigateway-deployment.golden.yaml b/test/suite/testdata/multigateway-deployment.golden.yaml new file mode 100644 index 00000000..c0c3211a --- /dev/null +++ b/test/suite/testdata/multigateway-deployment.golden.yaml @@ -0,0 +1,89 @@ +metadata: + labels: + app.kubernetes.io/component: multigateway + app.kubernetes.io/instance: golden-cluster + app.kubernetes.io/managed-by: multigres-operator + app.kubernetes.io/name: multigres + app.kubernetes.io/part-of: multigres + multigres.com/cell: zone1 + name: golden-cluster-zone1-multigateway-90bac2c9 + namespace: default + ownerReferences: + - apiVersion: multigres.com/v1alpha1 + blockOwnerDeletion: true + controller: true + kind: Cell + name: golden-cell + uid: golden-cell-uid +spec: + replicas: 1 + selector: + matchLabels: + app.kubernetes.io/component: multigateway + app.kubernetes.io/instance: golden-cluster + multigres.com/cell: zone1 + strategy: {} + template: + metadata: + annotations: + multigres.com/project-ref: golden-cluster + labels: + app.kubernetes.io/component: multigateway + app.kubernetes.io/instance: golden-cluster + app.kubernetes.io/managed-by: multigres-operator + app.kubernetes.io/name: multigres + app.kubernetes.io/part-of: multigres + multigres.com/cell: zone1 + spec: + containers: + - args: + - multigateway + - --http-port + - "15100" + - --grpc-port + - "15170" + - --pg-port + - "5432" + - --pg-replica-port + - "5433" + - --topo-global-server-addresses + - global-topo:2379 + - --topo-global-root + - /multigres/global + - --cell + - zone1 + - --log-level + - info + image: ghcr.io/multigres/multigres:golden-fixture + livenessProbe: + httpGet: + path: /live + port: 15100 + periodSeconds: 10 + name: multigateway + ports: + - containerPort: 15100 + name: http + protocol: TCP + - containerPort: 15170 + name: grpc + protocol: TCP + - containerPort: 5432 + name: postgres + protocol: TCP + - containerPort: 5433 + name: pg-replica + protocol: TCP + readinessProbe: + httpGet: + path: /ready + port: 15100 + periodSeconds: 5 + resources: {} + startupProbe: + failureThreshold: 30 + httpGet: + path: /ready + port: 15100 + periodSeconds: 5 +status: {} diff --git a/test/suite/types.go b/test/suite/types.go new file mode 100644 index 00000000..622c0d01 --- /dev/null +++ b/test/suite/types.go @@ -0,0 +1,48 @@ +package suite + +import multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + +// The API types this package builds most, without the package qualifier. +// +// multigresv1alpha1.MultigresCluster is thirty-three characters and says +// "multigres" twice, and a test body that constructs a dozen API objects +// spends more width on the qualifier than on what it is asserting. +// +// Two rules keep this from becoming a dot-import in disguise, which would +// trade that width for provenance nobody can recover: +// +// - Only types, never values or functions. multigresv1alpha1.PhaseHealthy +// and multigresv1alpha1.DeletePVCRetentionPolicy stay qualified, because +// a bare PhaseHealthy in an assertion genuinely does read as though it +// could be this package's own. A composite literal names its type on the +// line above, so &Shard{} does not have the same problem. +// - Only types used three or more times here. Aliasing a type used once +// saves eighteen characters and costs the next reader a lookup. +// +// Deliberately the same names as upstream, so there is nothing to learn and +// nothing to bikeshed: this drops the qualifier and changes nothing else. +// MultigresCluster in particular keeps its full name; a bare Cluster is far +// too overloaded in a workspace where that word also means a Kubernetes +// cluster and an EKS cluster. +// +// Orthogonal to any future rename of the multigresv1alpha1 alias itself, +// which is a repo-wide question across 217 files. These read the same either +// way. +type ( + MultigresCluster = multigresv1alpha1.MultigresCluster + Shard = multigresv1alpha1.Shard + PoolSpec = multigresv1alpha1.PoolSpec + ShardList = multigresv1alpha1.ShardList + PVCDeletionPolicy = multigresv1alpha1.PVCDeletionPolicy + CellConfig = multigresv1alpha1.CellConfig + MultigresClusterSpec = multigresv1alpha1.MultigresClusterSpec + PoolName = multigresv1alpha1.PoolName + PostgresPasswordSecretRef = multigresv1alpha1.PostgresPasswordSecretRef + TableGroupList = multigresv1alpha1.TableGroupList + CellName = multigresv1alpha1.CellName + DatabaseConfig = multigresv1alpha1.DatabaseConfig + TableGroupConfig = multigresv1alpha1.TableGroupConfig + ShardConfig = multigresv1alpha1.ShardConfig + ShardInlineSpec = multigresv1alpha1.ShardInlineSpec + TopoServerList = multigresv1alpha1.TopoServerList +) From db5a3dc9942a77aab622fae6568df71081b0f4e2 Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Sat, 19 Sep 2026 19:03:55 -0600 Subject: [PATCH 2/7] ci: add an advisory job for the multi-controller suite make test-suite runs it, and the other test targets exclude test/suite by path filter: the tier has no build tag, deliberately, so nothing can be hidden behind one. The job is advisory rather than required, because the suite pins live operator defects and a required check would block every PR on defects nobody is fixing in that PR. It runs verbosely, which is what makes a passing run readable: a pin that is doing its job logs and passes, and Go discards that output without -v, so a suite with seven live pins would otherwise print the same thing as a suite with none. Signed-off-by: Brent Graveland --- .github/workflows/build-and-release.yaml | 4 ++- .github/workflows/test-suite.yaml | 40 ++++++++++++++++++++++++ 2 files changed, 43 insertions(+), 1 deletion(-) create mode 100644 .github/workflows/test-suite.yaml diff --git a/.github/workflows/build-and-release.yaml b/.github/workflows/build-and-release.yaml index aedf85a0..9ceb2873 100644 --- a/.github/workflows/build-and-release.yaml +++ b/.github/workflows/build-and-release.yaml @@ -50,7 +50,9 @@ jobs: cache: false - name: Run tests - run: go test ./... -coverprofile=./cover.out -covermode=atomic -coverpkg=./... + run: | + go test $(go list ./... | grep -v /test/suite) \ + -coverprofile=./cover.out -covermode=atomic -coverpkg=./... - name: Run observer tests against current operator API working-directory: tools/observer diff --git a/.github/workflows/test-suite.yaml b/.github/workflows/test-suite.yaml new file mode 100644 index 00000000..ce840c22 --- /dev/null +++ b/.github/workflows/test-suite.yaml @@ -0,0 +1,40 @@ +# This job is advisory. It is deliberately not a required check and is not +# referenced by any other workflow's needs. A red run here means the suite +# caught a real problem, not that CI itself is broken: treat a failure as a +# finding to investigate, not as noise to wave off or restart until it goes +# green. +name: Test suite + +on: + pull_request: {} + # main only: pull_request already covers any branch with a PR open, so a + # wider push trigger would run the suite twice for it. main is here to catch + # a bad merge. + push: + branches: + - main + workflow_dispatch: {} + +permissions: + contents: read + +jobs: + test-suite: + runs-on: ubuntu-latest + timeout-minutes: 30 + permissions: + contents: read + steps: + - name: Check out code + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + with: + persist-credentials: false + + - name: Install Go + uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 + with: + go-version-file: go.mod + cache: false + + - name: Run test suite + run: make test-suite From 384491fa7267f987189e3674e82348047d303550 Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Sat, 19 Sep 2026 19:15:55 -0600 Subject: [PATCH 3/7] test: add a race-detector target for the multi-controller suite Nothing in this repo ran -race, which is an odd gap for the one suite where five controllers share a manager. Measured on the first run: 183s against a 166s baseline and zero data races. The 10% is cheaper than expected because this suite spends most of its wall clock waiting for controllers to converge, and the race detector does not slow down waiting. Separate from test-suite anyway. Certificate generation is the one CPU-bound step and has been measured swinging between 14 and 75 seconds under -race, which is enough to turn a wait sized against the normal run into a flake, so the timeout here is deliberately loose. The operator holds exactly one piece of state across reconcile goroutines, ShardReconciler.postureStrikes, and it is mutex-guarded with controller- runtime already serialising per object key. So this is a standing check that the answer has not changed rather than a hunt for a known race. Signed-off-by: Brent Graveland --- Makefile | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/Makefile b/Makefile index 9aef3644..a842b125 100644 --- a/Makefile +++ b/Makefile @@ -293,6 +293,33 @@ test-suite: manifests generate fmt vet setup-envtest ## Run the multi-controller KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ go test -v -p 1 -timeout 20m ./test/suite/... +# A separate target rather than a flag on the one above. Measured 2026-09-19: +# 183s against a 166s baseline, so about 10% rather than the roughly-double a +# CPU-bound suite would pay. This one spends most of its wall clock waiting for +# controllers to converge, and the race detector does not slow down waiting. +# +# Kept separate anyway, because the cost is not the same everywhere: certificate +# generation is the one CPU-bound step here and has been measured swinging +# between 14 and 75 seconds under -race, which is enough to turn a wait sized +# against the normal run into a flake. A budget that holds on both is looser +# than the default target should carry. +# +# Worth having at all because this suite is the only place five controllers +# share one manager, and the operator holds exactly one piece of state across +# reconcile goroutines: ShardReconciler.postureStrikes, a map guarded by a +# mutex. Nothing here exercises contention on it today, since the suite pins +# MaxConcurrentReconciles to 1 and controller-runtime already serialises +# reconciles per object key, so this is a standing check that the answer has +# not changed rather than a hunt for a known race. +# +# The timeout is generous rather than tight: the instrumented run is only +# slightly slower on average, but its slow tail is much fatter, and a timeout +# that fires on the tail reads as a hang rather than as the flake it is. +.PHONY: test-suite-race +test-suite-race: manifests generate fmt vet setup-envtest ## Run the multi-controller test suite under the race detector + KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ + go test -race -v -p 1 -timeout 40m ./test/suite/... + .PHONY: test test: manifests generate fmt vet ## Run tests (no integration testing) KUBEBUILDER_ASSETS="$(shell $(ENVTEST) use $(ENVTEST_K8S_VERSION) --bin-dir $(LOCALBIN) -p path)" \ From f41dbc220198d25928286f1c7096111f1ee50ea7 Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Fri, 25 Sep 2026 08:46:41 -0600 Subject: [PATCH 4/7] fix(shard): give every status write an explicit field owner A patch without client.FieldOwner does not opt out of field management. The API server derives one from the client's User-Agent, so the write gets an owner nobody named and nobody can see, and that owner then co-owns whatever fields the patch touched. Four Shard status writes did this, and the fields they claimed are also claimed by updateStatus on every reconcile, so two managers held them jointly. That is the same structure as the status hot loop: two managers, one object, overlapping fields. It was not looping, because these are merge patches that agree on values rather than applies that disagree, but it is one changed value away from the same outcome. The Pod status write in reconcile_readiness.go takes a distinct manager rather than the Shard's. It claims one condition on a Pod whose status otherwise belongs to kubelet, and a manager name is the only record of which concern took a field. The multi-controller suite pinned this with a KnownDefect on the field-ownership check in TestTwoControllersWriteOneShard. The pin becomes the positive assertion it was waiting to be: no field on the Shard is claimed by more than one manager. Signed-off-by: Brent Graveland --- .../controller/shard/reconcile_data_plane.go | 22 +++++++++++++++---- .../controller/shard/reconcile_deletion.go | 7 +++++- .../controller/shard/reconcile_readiness.go | 11 +++++++++- test/suite/scenario_race_test.go | 18 ++++----------- 4 files changed, 38 insertions(+), 20 deletions(-) diff --git a/pkg/resource-handler/controller/shard/reconcile_data_plane.go b/pkg/resource-handler/controller/shard/reconcile_data_plane.go index 57ac0301..0ce880e1 100644 --- a/pkg/resource-handler/controller/shard/reconcile_data_plane.go +++ b/pkg/resource-handler/controller/shard/reconcile_data_plane.go @@ -205,8 +205,12 @@ func (r *ShardReconciler) reconcileDataPlane( shard, fmt.Sprintf("Failed to check backup health: %v", err), ) - if patchErr := r.Status(). - Patch(ctx, shard, client.MergeFrom(backupBase)); patchErr != nil { + if patchErr := r.Status().Patch( + ctx, + shard, + client.MergeFrom(backupBase), + client.FieldOwner("multigres-resource-handler"), + ); patchErr != nil { return ctrl.Result{}, fmt.Errorf("update unavailable backup status: %w", patchErr) } } else if result != nil { @@ -222,7 +226,12 @@ func (r *ShardReconciler) reconcileDataPlane( r.Recorder.Event(shard, "Warning", "BackupStale", result.Message) } - if err := r.Status().Patch(ctx, shard, client.MergeFrom(backupBase)); err != nil { + if err := r.Status().Patch( + ctx, + shard, + client.MergeFrom(backupBase), + client.FieldOwner("multigres-resource-handler"), + ); err != nil { monitoring.RecordSpanError(childSpan, err) childSpan.End() logger.Error(err, "Failed to update shard backup status") @@ -317,7 +326,12 @@ func (r *ShardReconciler) reconcilePodRoles( } if rolesChanged { - if err := r.Status().Patch(ctx, shard, client.MergeFrom(statusBase)); err != nil { + if err := r.Status().Patch( + ctx, + shard, + client.MergeFrom(statusBase), + client.FieldOwner("multigres-resource-handler"), + ); err != nil { logger.Error(err, "Failed to update shard pod roles") } } diff --git a/pkg/resource-handler/controller/shard/reconcile_deletion.go b/pkg/resource-handler/controller/shard/reconcile_deletion.go index 4a949499..af17193c 100644 --- a/pkg/resource-handler/controller/shard/reconcile_deletion.go +++ b/pkg/resource-handler/controller/shard/reconcile_deletion.go @@ -342,7 +342,12 @@ func (r *ShardReconciler) handlePendingDeletion( ObservedGeneration: shard.Generation, LastTransitionTime: metav1.Now(), }) - if err := r.Status().Patch(ctx, shard, client.MergeFrom(statusBase)); err != nil { + if err := r.Status().Patch( + ctx, + shard, + client.MergeFrom(statusBase), + client.FieldOwner("multigres-resource-handler"), + ); err != nil { return ctrl.Result{}, fmt.Errorf("setting ReadyForDeletion condition: %w", err) } logger.Info("Set ReadyForDeletion condition") diff --git a/pkg/resource-handler/controller/shard/reconcile_readiness.go b/pkg/resource-handler/controller/shard/reconcile_readiness.go index fe6886f2..71234eae 100644 --- a/pkg/resource-handler/controller/shard/reconcile_readiness.go +++ b/pkg/resource-handler/controller/shard/reconcile_readiness.go @@ -71,7 +71,16 @@ func (r *ShardReconciler) reconcilePoolerReadiness( Reason: observation.Reason, Message: observation.Message, }) - if err := r.Status().Patch(ctx, pod, client.MergeFrom(base)); err != nil { + // Named apart from the Shard's own status manager: the claim here is over + // one condition on a Pod whose status otherwise belongs to kubelet, not + // over the Shard's status, and a manager name is the only record of which + // concern took a field. Same reasoning as the storage-class guard. + if err := r.Status().Patch( + ctx, + pod, + client.MergeFrom(base), + client.FieldOwner("multigres-resource-handler-readiness"), + ); err != nil { return fmt.Errorf("patch pooler readiness for pod %s: %w", pod.Name, err) } } diff --git a/test/suite/scenario_race_test.go b/test/suite/scenario_race_test.go index 8ac8a766..7b233ab8 100644 --- a/test/suite/scenario_race_test.go +++ b/test/suite/scenario_race_test.go @@ -64,21 +64,11 @@ func TestTwoControllersWriteOneShard(t *testing.T) { // managers. t.Run("field ownership on the Shard is disjoint", func(t *testing.T) { c := c.Sub(t) - // Several Shard status writes carry no field owner, so the API server - // attributes them to the manager that happens to be the process name, - // and that manager ends up co-owning fields the shard controller's own - // applier claims. Retire this pin with an explicit owner on every - // status write, and replace it with c.Empty on the conflicts. conflicts := c.fieldOwnershipConflicts(key) - c.KnownDefect("MGO-SHARD-STATUS-WRITES-NO-FIELD-OWNER", func() error { - if len(conflicts) == 0 { - return nil - } - return fmt.Errorf( - "Shard %s has fields claimed by more than one field manager, "+ - "the same shape of defect as the status hot loop:\n %s", - key.Name, strings.Join(conflicts, "\n ")) - }) + c.Empty(conflicts, + "Shard %s has fields claimed by more than one field manager, "+ + "the same shape of defect as the status hot loop:\n %s", + key.Name, strings.Join(conflicts, "\n ")) }) t.Run("tablegroup's patches to the Shard never change it", func(t *testing.T) { From 084d2bb0e9de58805b72c79f327799d296ecee33 Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Fri, 25 Sep 2026 09:48:42 -0600 Subject: [PATCH 5/7] chore: build with Go 1.27.1 testkit's assertions are generic methods, which need Go 1.27, so adding it as a test dependency raised this module's go directive from 1.26.6. Pin it to the exact release and move the builder image with it: official Go images set GOTOOLCHAIN=local, so a 1.26.6 builder refuses a module that asks for 1.27. golangci-lint goes to v2.13.2. v2.12.2's bundled staticcheck IR builder cannot parse Go 1.27 syntax and panics on the standard library (buildir: package "poll": unexpected expr: *ast.KeyValueExpr), which fails every lint run on Linux. Go 1.27 support landed in v2.13.0. The newer staticcheck names deprecated symbols by full package path, so four of the existing controller-runtime deprecation exclusions stopped matching. Their patterns now accept either form rather than adding new exclusions; the deprecations and the follow-up migration are unchanged. The golangci-lint binary's name now also carries the Go version it was built with. CI restores bin/ from older caches, and a name keyed only on the linter's version reused a binary built by go1.26 after the bump, which refused to load a go1.27 module. Any future Go bump would repeat that. tools/observer is a separate module that does not use testkit and stays on 1.26.6. Signed-off-by: Brent Graveland --- .golangci.toml | 13 +++++++++---- Dockerfile | 2 +- Makefile | 16 ++++++++++------ go.mod | 2 +- 4 files changed, 21 insertions(+), 12 deletions(-) diff --git a/.golangci.toml b/.golangci.toml index cf89d727..ce4a124a 100644 --- a/.golangci.toml +++ b/.golangci.toml @@ -10,11 +10,16 @@ enable = [ "gosec" ] # for a dedicated follow-up migration; excluded here so the dependency bump # that introduced the deprecations isn't blocked on an unrelated, # wide-reaching refactor. +# +# The optional package-path prefix is there because staticcheck names the +# symbol differently across versions: bare ("client.Apply") before +# golangci-lint v2.13, fully qualified ("sigs.k8s.io/.../pkg/client.Apply") +# from it. rules = [ - { linters = [ "staticcheck" ], text = "SA1019: client.Apply is deprecated" }, - { linters = [ "staticcheck" ], text = "SA1019: scheme.Builder is deprecated" }, - { linters = [ "staticcheck" ], text = "SA1019: webhook.CustomDefaulter is deprecated" }, - { linters = [ "staticcheck" ], text = "SA1019: webhook.CustomValidator is deprecated" }, + { linters = [ "staticcheck" ], text = 'SA1019: (\S+/)?client\.Apply is deprecated' }, + { linters = [ "staticcheck" ], text = 'SA1019: (\S+/)?scheme\.Builder is deprecated' }, + { linters = [ "staticcheck" ], text = 'SA1019: (\S+/)?webhook\.CustomDefaulter is deprecated' }, + { linters = [ "staticcheck" ], text = 'SA1019: (\S+/)?webhook\.CustomValidator is deprecated' }, { linters = [ "staticcheck" ], text = "WithCustomDefaulter is deprecated" }, { linters = [ "staticcheck" ], text = "WithCustomValidator is deprecated" }, { linters = [ "staticcheck" ], text = "GetEventRecorderFor is deprecated" }, diff --git a/Dockerfile b/Dockerfile index 3c215ac3..32a24b31 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,6 +1,6 @@ # Containerfile for multigres-operator -FROM --platform=$BUILDPLATFORM golang:1.26.6-alpine3.23 AS builder +FROM --platform=$BUILDPLATFORM golang:1.27.1-alpine3.23 AS builder ARG TARGETOS ARG TARGETARCH diff --git a/Makefile b/Makefile index a842b125..2c274622 100644 --- a/Makefile +++ b/Makefile @@ -108,7 +108,7 @@ KUSTOMIZE_VERSION ?= v5.6.0 # renovate: datasource=github-releases depName=kubernetes-sigs/controller-tools CONTROLLER_TOOLS_VERSION ?= v0.18.0 # renovate: datasource=github-releases depName=golangci/golangci-lint -GOLANGCI_LINT_VERSION ?= v2.12.2 +GOLANGCI_LINT_VERSION ?= v2.13.2 CERT_MANAGER_VERSION ?= v1.19.2 @@ -722,10 +722,13 @@ $(ENVTEST): $(LOCALBIN) golangci-lint: $(GOLANGCI_LINT) ## Download golangci-lint locally if necessary. # golangci-lint's own go.mod selects an older toolchain than this module # targets, and a linter built with a lower Go version refuses to run. Pin the -# build toolchain to the one resolved by this module's go.mod. +# build toolchain to the one resolved by this module's go.mod, and put that +# version in the binary's name: CI restores bin/ from older caches, and a +# name keyed only on the linter's version would reuse a binary built by the +# previous toolchain after a Go bump. $(GOLANGCI_LINT): export GOTOOLCHAIN = $(shell go env GOVERSION) $(GOLANGCI_LINT): $(LOCALBIN) - $(call go-install-tool,$(GOLANGCI_LINT),github.com/golangci/golangci-lint/v2/cmd/golangci-lint,$(GOLANGCI_LINT_VERSION)) + $(call go-install-tool,$(GOLANGCI_LINT),github.com/golangci/golangci-lint/v2/cmd/golangci-lint,$(GOLANGCI_LINT_VERSION),$(shell go env GOVERSION)) .PHONY: install-certmanager install-certmanager: ## Install Cert-Manager into the cluster @@ -738,16 +741,17 @@ install-certmanager: ## Install Cert-Manager into the cluster # $1 - target path with name of binary # $2 - package url which can be installed # $3 - specific version of package +# $4 - optional extra suffix for the binary's name, e.g. the Go version it was built with define go-install-tool -@[ -f "$(1)-$(3)" ] && [ "$$(readlink -- "$(1)" 2>/dev/null)" = "$(1)-$(3)" ] || { \ +@[ -f "$(1)-$(3)$(if $(4),-$(4))" ] && [ "$$(readlink -- "$(1)" 2>/dev/null)" = "$(1)-$(3)$(if $(4),-$(4))" ] || { \ set -e; \ package=$(2)@$(3) ;\ echo "Downloading $${package}" ;\ rm -f $(1) ;\ GOBIN=$(LOCALBIN) go install $${package} ;\ -mv $(1) $(1)-$(3) ;\ +mv $(1) $(1)-$(3)$(if $(4),-$(4)) ;\ } ;\ -ln -sf $$(realpath $(1)-$(3)) $(1) +ln -sf $$(realpath $(1)-$(3)$(if $(4),-$(4))) $(1) endef ##@ Backward Compatibility Aliases diff --git a/go.mod b/go.mod index 6b0d4ca9..57af9939 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/multigres/multigres-operator -go 1.27 +go 1.27.1 require ( github.com/go-logr/logr v1.4.4 From 225dddde65cb449e2f027d917235fdbc4d2f80bf Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Sat, 26 Sep 2026 12:34:50 -0600 Subject: [PATCH 6/7] fix(shard): requeue while a shard has not converged The commit that stopped the shard controller's status hot loop ("stop the status hot loop on a healthy shard") removed a driver that used to rerun every shard several times a second regardless of what reconcile asked for. That accidentally covered for a gap: once a posture observation settles (nothing inconsistent, nothing incomplete, or an unsettled observation accepted after its debounce), the shard requests no further reconcile at all, even when a managed pod has never reached posture readiness, or a shard mid-bootstrap has every pooler registered but no primary elected yet. Nothing in Kubernetes watches the topology store, so nothing else wakes the shard either: an accepted RPC blip or role mismatch leaves a pod NotReady or a shard Degraded until controller-runtime's 10h resync, and a fresh cluster's pool pods never go Ready at all. reconcilePosture now requests a requeue for any not-converged state, unsettled or merely not-ready, once the existing debounce for unsettled observations ends. The delay is clamped elapsed time since the shard was first observed not converged, five seconds to one minute, with up to 20% upward jitter (so five seconds to about seventy-two seconds including jitter), and resets the moment the shard converges. Elapsed time rather than a per-reconcile count, because pod status transitions, drain requeues and the operator's own status patches are each their own reconcile with no predicate filtering them, and a burst of those must not by itself run the backoff up to its ceiling. Scopes the posture pod list to pool pods: a shard's multiorch pod carries the same four identity labels and would otherwise count as a pod that never becomes ready. The pod-roles, drain, and pooler-prune lists are scoped the same way for consistency; pooler-prune's filter is the one that matters operationally, since it decides which topology poolers this reconcile marks LIFECYCLE_SHUTDOWN, and every pool pod has carried the component label since the pool controller was introduced, so this is not a behaviour change for existing clusters. test/suite's pool-scale-up pin is now a positive assertion: the role landing in status.podRoles after a scale-up whose pooler registers late used to be a coin flip, and is deterministic with this fix in place. Signed-off-by: Brent Graveland --- pkg/data-handler/posture/posture.go | 11 +- .../controller/shard/reconcile_data_plane.go | 197 +++- ...oncile_data_plane_posture_internal_test.go | 22 +- .../controller/shard/reconcile_deletion.go | 4 + .../shard/registration_requeue_test.go | 858 ++++++++++++++++++ .../controller/shard/shard_controller.go | 14 + test/suite/scenario_thrash_test.go | 87 +- 7 files changed, 1135 insertions(+), 58 deletions(-) create mode 100644 pkg/resource-handler/controller/shard/registration_requeue_test.go diff --git a/pkg/data-handler/posture/posture.go b/pkg/data-handler/posture/posture.go index 049b8ceb..8d0a83b6 100644 --- a/pkg/data-handler/posture/posture.go +++ b/pkg/data-handler/posture/posture.go @@ -47,6 +47,15 @@ type Readiness struct { Message string } +// reasonAwaitingRegistration is the readiness reason carried by a managed pod +// that has no corresponding pooler in the shard topology. +// +// Every managed pod is seeded with it and only overwritten once a topology +// entry matches. Note this is NOT what Result.Incomplete reports: that covers +// an unreachable cell, a topology entry with no matching pod, or an UNKNOWN +// posture, all of which are the opposite direction. +const reasonAwaitingRegistration = "AwaitingRegistration" + // Evaluate compares each managed pooler's observed postgres state with its // topology role. It returns nil when topology contains no active poolers, as // during bootstrap. @@ -61,7 +70,7 @@ func Evaluate( readiness := make(map[string]Readiness, len(managedPodNames)) for _, podName := range managedPodNames { readiness[podName] = Readiness{ - Reason: "AwaitingRegistration", + Reason: reasonAwaitingRegistration, Message: "pooler has not registered in the shard topology", } } diff --git a/pkg/resource-handler/controller/shard/reconcile_data_plane.go b/pkg/resource-handler/controller/shard/reconcile_data_plane.go index 0ce880e1..7bb1bd49 100644 --- a/pkg/resource-handler/controller/shard/reconcile_data_plane.go +++ b/pkg/resource-handler/controller/shard/reconcile_data_plane.go @@ -10,6 +10,7 @@ import ( corev1 "k8s.io/api/core/v1" meta "k8s.io/apimachinery/pkg/api/meta" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/util/wait" ctrl "sigs.k8s.io/controller-runtime" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/log" @@ -73,7 +74,12 @@ func (r *ShardReconciler) reconcileDataPlane( // A resolver error breaks the sequence of posture observations. It must // not let a strike from before the transport outage combine with the - // first unsettled observation after recovery. + // first unsettled observation after recovery. notConvergedSince needs no + // matching reset here: this path (like the nil-topology one below) + // returns its own fixed delay independent of that map, so a stale entry + // left over from before the outage only shortens the path to the + // backoff's ceiling once posture resumes, it never suppresses a requeue + // that is actually owed. That is intentional, not a gap. r.recordPostureObservation(shard, false) // A partial observation cannot clear a confirmed posture failure. Keep a @@ -279,12 +285,16 @@ func (r *ShardReconciler) reconcilePodRoles( ) { logger := log.FromContext(ctx) - // List managed pods for this shard (same pattern as reconcilePoolerPrune). + // List managed pool pods for this shard (same pattern as + // reconcilePoolerPrune). Scoped to pool pods: only they can ever match a + // topology pooler entry, and a shard's multiorch pod carries the same four + // identity labels but is never one. lbls := map[string]string{ metadata.LabelMultigresCluster: shard.Labels[metadata.LabelMultigresCluster], metadata.LabelMultigresDatabase: string(shard.Spec.DatabaseName), metadata.LabelMultigresTableGroup: string(shard.Spec.TableGroupName), metadata.LabelMultigresShard: string(shard.Spec.ShardName), + metadata.LabelAppComponent: PoolComponentName, } podList := &corev1.PodList{} if err := r.List(ctx, podList, @@ -345,6 +355,22 @@ const ( // client cannot be built yet (e.g. operator client cert not issued). poolerClientRetryDelay = 10 * time.Second + // A shard that is not converged, whatever the reason, is requeued on a + // backoff rather than left to the 10h resync. The delay is clamped elapsed + // time since the shard was first observed not-converged, not a per-pass + // count: a burst of unrelated pod/status events must not itself advance + // the backoff. The floor is the multipooler container's own readiness + // probe period, so the operator never polls topology more aggressively + // than the kubelet polls the pod, and the ceiling matches the + // empty-topology path. + readinessBackoffMinDelay = 5 * time.Second + readinessBackoffMaxDelay = time.Minute + // readinessBackoffJitter adds upward jitter. controller-runtime does not + // jitter RequeueAfter, and a fleet that lost convergence together + // (topology restart, mass scale-up) would otherwise retry in lockstep + // against the topology server that just came back. + readinessBackoffJitter = 0.2 + reasonPoolerClientUnavailable = "PoolerClientUnavailable" reasonAwaitingPoolerRegistration = "AwaitingPoolerRegistration" reasonObservationPending = "ObservationPending" @@ -357,11 +383,17 @@ func (r *ShardReconciler) reconcilePosture( shard *multigresv1alpha1.Shard, rpcClient rpcclient.MultipoolerClient, ) (time.Duration, error) { + // Pool pods only: a shard's multiorch pod carries the same four identity + // labels (buildMultiorchLabelsWithCell merges them in) but is never a + // pooler, so an unfiltered list seeds it AwaitingRegistration forever and + // anyPodNotReady never clears. reconcilePoolerReadiness already scopes + // this way; posture.Evaluate must match it. lbls := map[string]string{ metadata.LabelMultigresCluster: shard.Labels[metadata.LabelMultigresCluster], metadata.LabelMultigresDatabase: string(shard.Spec.DatabaseName), metadata.LabelMultigresTableGroup: string(shard.Spec.TableGroupName), metadata.LabelMultigresShard: string(shard.Spec.ShardName), + metadata.LabelAppComponent: PoolComponentName, } podList := &corev1.PodList{} if err := r.List(ctx, podList, @@ -388,6 +420,8 @@ func (r *ShardReconciler) reconcilePosture( // An empty topology is expected during bootstrap, but it is not a settled // posture observation. Keep polling until poolers register rather than // leaving a previous transport condition stuck until the periodic resync. + // This path always returns a fixed 1-minute delay of its own regardless of + // notConvergedSince, so it needs no matching reset there either. r.recordPostureObservation(shard, false) setPostureUnknownUnlessFalse( shard, @@ -450,12 +484,78 @@ func (r *ShardReconciler) reconcilePosture( r.Recorder.Event(shard, "Warning", reason, result.Message) } + // A shard is converged only once posture is settled AND every managed pod + // has reached posture readiness. Neither half implies the other: an + // accepted mismatch or accepted Incomplete observation (strikes past + // threshold) is settled but still has a pod sitting NotReady, and a shard + // mid-bootstrap with every pooler registered but no primary elected yet is + // neither inconsistent nor incomplete but has nothing actually ready. + // Registration itself, and any change multiorch makes in topology, are + // writes to the topology store with no Kubernetes event behind them, so + // nothing but this controller's own requeue will ever look again. Without + // it, a settled-but-not-ready shard sits until controller-runtime's 10h + // resync; a shard with a genuinely stuck pod (Pending, crash-looping, + // quarantined) sits forever, since a lost watch never recovers it. + // + // The debounce above stays authoritative for whether an unsettled + // observation is accepted into status: this never fires while + // `unsettled && strikes < postureStrikeThreshold`, so it cannot preempt or + // shorten that window, only pick up once it ends. + notConverged := unsettled || anyPodNotReady(result.Readiness) + elapsed := r.recordNotConverged(shard, notConverged) + if unsettled && strikes < postureStrikeThreshold { return postureDebounceRequeueDelay, nil } + if notConverged { + return readinessBackoffDelay(elapsed), nil + } return 0, nil } +// anyPodNotReady reports whether any managed pod has not yet reached posture +// readiness. +// +// Deliberately broader than a check for the seeded "AwaitingRegistration" +// reason alone. A pod whose pooler has never registered carries that reason, +// but every other reason poolerReadiness can return (NotInitialized, +// PostgresNotReady, CohortIneligible, NotCohortMember) means Ready is false +// too, and all of them are reachable by a shard that Evaluate reports as +// neither inconsistent nor incomplete, since that result only compares +// postures against topology roles and does not require a primary to exist. A +// shard mid-bootstrap with every pooler registered but no primary elected +// settles there, looking "consistent" while nothing is actually ready. +func anyPodNotReady(readiness map[string]posture.Readiness) bool { + for _, r := range readiness { + if !r.Ready { + return true + } + } + return false +} + +// readinessBackoffDelay clamps elapsed (time since the shard was first +// observed not-converged) to [readinessBackoffMinDelay, +// readinessBackoffMaxDelay], with upward jitter. +// +// Elapsed-time-based rather than a per-pass count: Owns(&corev1.Pod{}) and +// For(&Shard{}) carry no predicates, so a pod's Pending -> ContainerCreating -> +// per-container-Ready transitions, an operator gate patch, and a drain's 2s +// requeues are each their own reconcile. A scale-up alone is at least five of +// those before there is anything to register, and a wall-clock burst of +// events must not by itself run the delay up to the ceiling: two reconciles a +// second apart are five seconds not-converged either way, whether they were +// one pass or five. +func readinessBackoffDelay(elapsed time.Duration) time.Duration { + d := elapsed + if d < readinessBackoffMinDelay { + d = readinessBackoffMinDelay + } else if d > readinessBackoffMaxDelay { + d = readinessBackoffMaxDelay + } + return wait.Jitter(d, readinessBackoffJitter) +} + func setPostureUnknownUnlessFalse( shard *multigresv1alpha1.Shard, reason string, @@ -504,6 +604,18 @@ func withDataPlaneRequeue( return result } +// recordPostureObservation counts consecutive unsettled observations for one +// shard and returns the running total. +// +// A settled observation deletes the entry rather than writing zero. The two +// mean the same thing to every reader, since a missing key reads as zero, but +// they differ in what the map holds: writing zero keeps an entry for every +// shard this process has ever reconciled, while deleting keeps only the +// shards currently accumulating strikes, which in a healthy cluster is none. +// A Go map's table is sized by its peak simultaneous entries and does not +// shrink on delete, so this is what keeps that peak bounded by currently +// unsettled shards rather than by cumulative shards over the operator's +// lifetime. func (r *ShardReconciler) recordPostureObservation( shard *multigresv1alpha1.Shard, unsettled bool, @@ -512,17 +624,76 @@ func (r *ShardReconciler) recordPostureObservation( r.postureStrikesMu.Lock() defer r.postureStrikesMu.Unlock() + if !unsettled { + delete(r.postureStrikes, key) + return 0 + } if r.postureStrikes == nil { r.postureStrikes = make(map[string]int) } - if unsettled { - r.postureStrikes[key]++ - } else { - r.postureStrikes[key] = 0 - } + r.postureStrikes[key]++ return r.postureStrikes[key] } +// recordNotConverged tracks, per shard, the time a shard was first observed +// not converged (unsettled, or some managed pod not posture-ready), and +// returns how long that has been true. Same delete-on-settle pattern as +// recordPostureObservation, and a separate map for the same reason that one +// is not reused here: this one picks the backoff, and posture strikes must +// keep gating only posture.Apply. +// +// Storing the first-seen time rather than a per-pass count is what makes +// readinessBackoffDelay a function of elapsed wall-clock time instead of +// event count; see readinessBackoffDelay's own doc for why that matters. +func (r *ShardReconciler) recordNotConverged( + shard *multigresv1alpha1.Shard, + notConverged bool, +) time.Duration { + key := fmt.Sprintf("%s/%s", shard.Namespace, shard.Name) + now := r.now() + + r.notConvergedMu.Lock() + defer r.notConvergedMu.Unlock() + if !notConverged { + delete(r.notConvergedSince, key) + return 0 + } + if r.notConvergedSince == nil { + r.notConvergedSince = make(map[string]time.Time) + } + first, ok := r.notConvergedSince[key] + if !ok { + first = now + r.notConvergedSince[key] = first + } + return now.Sub(first) +} + +// now returns the reconciler's clock, defaulting to time.Now so production +// code never has to set Clock. Tests inject Clock to drive +// recordNotConverged's elapsed-time math deterministically. +func (r *ShardReconciler) now() time.Time { + if r.Clock != nil { + return r.Clock() + } + return time.Now() +} + +// forgetStrikes drops namespace/name's entry from both strike maps. Called on +// a Shard's deletion and not-found paths so a shard deleted mid-backoff, in +// either counter, does not leak its entry for the life of the process. +func (r *ShardReconciler) forgetStrikes(namespace, name string) { + key := fmt.Sprintf("%s/%s", namespace, name) + + r.postureStrikesMu.Lock() + delete(r.postureStrikes, key) + r.postureStrikesMu.Unlock() + + r.notConvergedMu.Lock() + delete(r.notConvergedSince, key) + r.notConvergedMu.Unlock() +} + // reconcileDrainState iterates pods with drain annotations and runs the // drain state machine for each one. func (r *ShardReconciler) reconcileDrainState( @@ -532,11 +703,17 @@ func (r *ShardReconciler) reconcileDrainState( ) (bool, error) { logger := log.FromContext(ctx) + // Pool pods only: the drain-requested annotation this loop acts on is only + // ever set by the pool scale-down/rolling-update path (reconcile_pool_pods.go), + // never on a shard's multiorch pod, so this is currently a no-op filter. + // Scoped anyway for the same reason reconcilePosture now is: relying on an + // annotation nothing else sets is a coincidence, not a guarantee. lbls := map[string]string{ metadata.LabelMultigresCluster: shard.Labels[metadata.LabelMultigresCluster], metadata.LabelMultigresDatabase: string(shard.Spec.DatabaseName), metadata.LabelMultigresTableGroup: string(shard.Spec.TableGroupName), metadata.LabelMultigresShard: string(shard.Spec.ShardName), + metadata.LabelAppComponent: PoolComponentName, } podList := &corev1.PodList{} if err := r.List( @@ -655,11 +832,17 @@ func (r *ShardReconciler) reconcilePoolerPrune( return } + // Pool pods only: topo.MarkDeadPoolers matches this set's names against + // topology pooler entries, and a multiorch pod's name never matches one + // (it is not a pooler), so including it here is currently a no-op. Scoped + // anyway to keep this list's meaning ("pods that can be poolers") aligned + // with what it is actually used for. lbls := map[string]string{ metadata.LabelMultigresCluster: shard.Labels[metadata.LabelMultigresCluster], metadata.LabelMultigresDatabase: string(shard.Spec.DatabaseName), metadata.LabelMultigresTableGroup: string(shard.Spec.TableGroupName), metadata.LabelMultigresShard: string(shard.Spec.ShardName), + metadata.LabelAppComponent: PoolComponentName, } podList := &corev1.PodList{} if err := r.List( diff --git a/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go b/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go index 28e2a918..cd83c8e0 100644 --- a/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go @@ -114,6 +114,10 @@ func postureTestReconciler( }, c } +// postureTestPod is a pool pod: reconcilePosture's own pod list is scoped to +// PoolComponentName (a shard's multiorch pod carries the same four identity +// labels but is never a pooler), so a pod fixture without this label is +// invisible to it regardless of what the test otherwise sets up. func postureTestPod() *corev1.Pod { return &corev1.Pod{ObjectMeta: metav1.ObjectMeta{ Name: "pooler-0", @@ -123,6 +127,7 @@ func postureTestPod() *corev1.Pod { metadata.LabelMultigresDatabase: "database", metadata.LabelMultigresTableGroup: "table-group", metadata.LabelMultigresShard: "0", + metadata.LabelAppComponent: PoolComponentName, }, }} } @@ -244,8 +249,14 @@ func TestReconcilePostureDebouncesFirstInconsistency(t *testing.T) { if err != nil { t.Fatalf("second reconcilePosture() error = %v", err) } - if retryAfter != 0 { - t.Error("second inconsistent posture observation requested another debounce requeue") + // The mismatch is now accepted into status (PostureConsistent=False), but + // this mock pooler never reports IsInitialized/PostgresReady, so it has + // also never reached posture readiness. An accepted-but-not-ready shard + // must keep requesting a requeue: nothing but this controller's own + // backoff will ever look again, since neither a status recovery in + // topology nor a role fix changes a Kubernetes object. + if retryAfter <= 0 { + t.Error("second inconsistent posture observation (still not ready) requested no requeue") } if !conditionIsFalse(shard.Status.Conditions, posture.ConditionConsistent) { t.Errorf( @@ -296,8 +307,11 @@ func TestReconcilePostureDebouncesFirstIncompleteObservation(t *testing.T) { if err != nil { t.Fatalf("second reconcilePosture() error = %v", err) } - if retryAfter != 0 { - t.Error("second incomplete posture observation requested another debounce requeue") + // Accepted into status (Unknown/ObservationIncomplete), but a pooler whose + // Status RPC errors has also never reached posture readiness, so this must + // still requeue rather than strand the pod until the 10h resync. + if retryAfter <= 0 { + t.Error("second incomplete posture observation (still not ready) requested no requeue") } for _, condition := range shard.Status.Conditions { if condition.Type != posture.ConditionConsistent { diff --git a/pkg/resource-handler/controller/shard/reconcile_deletion.go b/pkg/resource-handler/controller/shard/reconcile_deletion.go index af17193c..13dbaa6b 100644 --- a/pkg/resource-handler/controller/shard/reconcile_deletion.go +++ b/pkg/resource-handler/controller/shard/reconcile_deletion.go @@ -130,6 +130,10 @@ func (r *ShardReconciler) handleDeletion( return ctrl.Result{}, err } + // The shard is done reconciling; drop its strike entries so a shard + // deleted mid-backoff does not leak its counters. + r.forgetStrikes(shard.Namespace, shard.Name) + // Remove the finalizer last so Kubernetes can finish deletion now that the // PVC cleanup has run. if slices.Contains(shard.Finalizers, shardFinalizer) { diff --git a/pkg/resource-handler/controller/shard/registration_requeue_test.go b/pkg/resource-handler/controller/shard/registration_requeue_test.go new file mode 100644 index 00000000..b9a5ff99 --- /dev/null +++ b/pkg/resource-handler/controller/shard/registration_requeue_test.go @@ -0,0 +1,858 @@ +package shard + +import ( + "errors" + "fmt" + "testing" + "time" + + "github.com/multigres/multigres/go/common/rpcclient" + "github.com/multigres/multigres/go/common/topoclient" + "github.com/multigres/multigres/go/common/topoclient/memorytopo" + "github.com/multigres/multigres/go/pb/clustermetadata" + multipoolermanagerdatapb "github.com/multigres/multigres/go/pb/multipoolermanagerdata" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/types" + "k8s.io/client-go/tools/record" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + + multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + "github.com/multigres/multigres-operator/pkg/data-handler/poolerclient" + "github.com/multigres/multigres-operator/pkg/util/metadata" +) + +// fakeClock lets a test drive recordNotConverged's elapsed-time math without +// sleeping. Set r.Clock = clk.now. +type fakeClock struct { + t time.Time +} + +func (c *fakeClock) now() time.Time { return c.t } + +func (c *fakeClock) advance(d time.Duration) { c.t = c.t.Add(d) } + +// wantDelayRange mirrors readinessBackoffDelay's own clamp so a change to +// that clamp is caught here rather than only against itself. +func wantDelayRange(elapsed time.Duration) (min, max time.Duration) { + min = elapsed + if min < readinessBackoffMinDelay { + min = readinessBackoffMinDelay + } else if min > readinessBackoffMaxDelay { + min = readinessBackoffMaxDelay + } + max = time.Duration(float64(min) * 1.2) + return min, max +} + +// TestReadinessBackoffDelayClampsElapsedTime pins the clamp. The lower bound +// matters as much as the upper one: a delay that could return zero would +// reinstate the defect this exists to fix, since a zero RequeueAfter means no +// requeue at all rather than an immediate one. +func TestReadinessBackoffDelayClampsElapsedTime(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + elapsed time.Duration + min, max time.Duration + }{ + {elapsed: 0, min: 5 * time.Second, max: 6 * time.Second}, + {elapsed: 3 * time.Second, min: 5 * time.Second, max: 6 * time.Second}, + {elapsed: 5 * time.Second, min: 5 * time.Second, max: 6 * time.Second}, + {elapsed: 30 * time.Second, min: 30 * time.Second, max: 36 * time.Second}, + {elapsed: 59 * time.Second, min: 59 * time.Second, max: 70800 * time.Millisecond}, + {elapsed: time.Minute, min: time.Minute, max: 72 * time.Second}, + {elapsed: 70 * time.Second, min: time.Minute, max: 72 * time.Second}, + {elapsed: time.Hour, min: time.Minute, max: 72 * time.Second}, + } { + // Jitter is random, so the bound has to hold across repeats rather + // than on one draw. + for range 200 { + got := readinessBackoffDelay(tc.elapsed) + if got < tc.min || got > tc.max { + t.Fatalf("elapsed=%v: delay %v outside [%v, %v]", + tc.elapsed, got, tc.min, tc.max) + } + } + } +} + +// TestReadinessBackoffDelayIsNeverZero is the one that dies if the clamp is +// removed. A zero duration is not "retry immediately", it is "do not +// requeue", which is exactly how a shard ends up stranded not-converged. +func TestReadinessBackoffDelayIsNeverZero(t *testing.T) { + t.Parallel() + + for _, elapsed := range []time.Duration{ + -5 * time.Second, -1, 0, time.Second, 5 * time.Second, + 30 * time.Second, time.Minute, time.Hour, + } { + for range 50 { + if got := readinessBackoffDelay(elapsed); got <= 0 { + t.Fatalf("elapsed=%v produced a non-positive delay %v, "+ + "which controller-runtime reads as no requeue", elapsed, got) + } + } + } +} + +// TestReadinessBackoffDelayJitters guards the fleet-lockstep property: a +// constant delay would have every shard that lost convergence together retry +// in lockstep against the topology server that just came back. +func TestReadinessBackoffDelayJitters(t *testing.T) { + t.Parallel() + + seen := map[time.Duration]bool{} + for range 200 { + seen[readinessBackoffDelay(30*time.Second)] = true + } + if len(seen) < 10 { + t.Fatalf("only %d distinct delays across 200 draws; jitter is not applied", len(seen)) + } +} + +// shardNamed is the minimum a strike counter reads. +func shardNamed(ns, name string) *multigresv1alpha1.Shard { + return &multigresv1alpha1.Shard{ + ObjectMeta: metav1.ObjectMeta{Namespace: ns, Name: name}, + } +} + +// TestPostureStrikesLeaveNoEntryOnceSettled is what deleting on settle +// actually buys: a map's table is sized by its peak simultaneous entries and +// does not shrink on delete, so writing zero instead kept an entry for every +// shard the process had ever reconciled. +func TestPostureStrikesLeaveNoEntryOnceSettled(t *testing.T) { + t.Parallel() + + r := &ShardReconciler{} + for i := range 1000 { + s := shardNamed("ns", fmt.Sprintf("shard-%d", i)) + r.recordPostureObservation(s, true) + r.recordPostureObservation(s, false) + } + + if got := len(r.postureStrikes); got != 0 { + t.Fatalf("a thousand shards seen and settled left %d entries, want 0", got) + } +} + +// TestPostureStrikesDoNotSurviveRecreation pins the intended semantic: a +// recreated Shard (a different object at the same namespace/name) opens at +// strike 1, not at whatever count its predecessor left. Posture strikes gate +// posture.Apply (when an unsettled observation is accepted into status); they +// do not pick a requeue delay, that is notConvergedSince's job. +func TestPostureStrikesDoNotSurviveRecreation(t *testing.T) { + t.Parallel() + + r := &ShardReconciler{} + s := shardNamed("ns", "shard-0") + + for range 5 { + r.recordPostureObservation(s, true) + } + if got := r.recordPostureObservation(s, false); got != 0 { + t.Fatalf("a settled observation reported %d strikes, want 0", got) + } + + // The replacement is a different object at the same key, which is what + // the tablegroup controller creates after a Shard is deleted. + if got := r.recordPostureObservation(shardNamed("ns", "shard-0"), true); got != 1 { + t.Fatalf("a recreated shard opened at %d strikes, want 1", got) + } +} + +// TestPostureStrikesCountConsecutiveUnsettled pins what the counter is for, +// so settling on delete cannot be "fixed" into never counting at all. +func TestPostureStrikesCountConsecutiveUnsettled(t *testing.T) { + t.Parallel() + + r := &ShardReconciler{} + s := shardNamed("ns", "shard-0") + for want := 1; want <= 3; want++ { + if got := r.recordPostureObservation(s, true); got != want { + t.Fatalf("consecutive unsettled observation %d reported %d strikes", want, got) + } + } + // Shards are counted independently, which is the only reason the map has + // keys at all. + if got := r.recordPostureObservation(shardNamed("ns", "other"), true); got != 1 { + t.Fatalf("a second shard opened at %d strikes, want 1", got) + } + if got := r.recordPostureObservation(s, true); got != 4 { + t.Fatalf("the first shard reported %d strikes after a second shard, want 4", got) + } +} + +// TestNotConvergedSinceLeavesNoEntryOnceSettled mirrors +// TestPostureStrikesLeaveNoEntryOnceSettled for the not-converged-since map: a +// shard deleted mid-backoff must not leak its entry, and neither must one +// that simply converges. +func TestNotConvergedSinceLeavesNoEntryOnceSettled(t *testing.T) { + t.Parallel() + + r := &ShardReconciler{} + for i := range 1000 { + s := shardNamed("ns", fmt.Sprintf("shard-%d", i)) + r.recordNotConverged(s, true) + r.recordNotConverged(s, false) + } + + if got := len(r.notConvergedSince); got != 0 { + t.Fatalf("a thousand shards seen and settled left %d entries, want 0", got) + } +} + +// TestNotConvergedSinceTracksElapsedTime pins the elapsed-time semantics: the +// first not-converged observation opens the clock, later ones read the time +// since then rather than a per-call count, and a settled observation clears +// it so a later not-converged spell starts over rather than resuming. +func TestNotConvergedSinceTracksElapsedTime(t *testing.T) { + t.Parallel() + + clk := &fakeClock{t: time.Unix(1_700_000_000, 0)} + r := &ShardReconciler{Clock: clk.now} + s := shardNamed("ns", "shard-0") + + if got := r.recordNotConverged(s, true); got != 0 { + t.Fatalf("first not-converged observation reported elapsed %v, want 0", got) + } + clk.advance(37 * time.Second) + if got := r.recordNotConverged(s, true); got != 37*time.Second { + t.Fatalf("second not-converged observation reported elapsed %v, want 37s", got) + } + // A burst of same-instant calls (a wave of unrelated pod events) must not + // itself advance the elapsed time. + if got := r.recordNotConverged(s, true); got != 37*time.Second { + t.Fatalf( + "third not-converged observation (no time passed) reported elapsed %v, want 37s", got, + ) + } + + if got := r.recordNotConverged(s, false); got != 0 { + t.Fatalf("a settled observation reported elapsed %v, want 0", got) + } + if got := r.recordNotConverged(s, true); got != 0 { + t.Fatalf("a fresh not-converged spell reported elapsed %v, want 0 (not resumed)", got) + } +} + +// TestForgetStrikesDropsBothCounters pins that a Shard's strike entries do +// not survive forgetStrikes, in either counter. +func TestForgetStrikesDropsBothCounters(t *testing.T) { + t.Parallel() + + r := &ShardReconciler{} + s := shardNamed("ns", "shard-0") + r.recordPostureObservation(s, true) + r.recordNotConverged(s, true) + + r.forgetStrikes(s.Namespace, s.Name) + + if got := len(r.postureStrikes); got != 0 { + t.Fatalf("posture strikes: %d entries survived forgetStrikes, want 0", got) + } + if got := len(r.notConvergedSince); got != 0 { + t.Fatalf("not-converged-since: %d entries survived forgetStrikes, want 0", got) + } +} + +// TestHandleDeletionForgetsStrikes drives the deletion cleanup through the +// real controller path (handleDeletion) rather than calling forgetStrikes +// directly, so a regression that stops handleDeletion from reaching it is +// caught here rather than only in the helper's own unit test. +func TestHandleDeletionForgetsStrikes(t *testing.T) { + shard := postureTestShard() + shard.Finalizers = []string{shardFinalizer} + now := metav1.Now() + shard.DeletionTimestamp = &now + + scheme := postureTestScheme(t) + c := fake.NewClientBuilder(). + WithScheme(scheme). + WithObjects(shard). + WithStatusSubresource(&multigresv1alpha1.Shard{}). + Build() + r := &ShardReconciler{ + Client: c, + Scheme: scheme, + Recorder: record.NewFakeRecorder(20), + } + + key := fmt.Sprintf("%s/%s", shard.Namespace, shard.Name) + r.postureStrikes = map[string]int{key: 2} + r.notConvergedSince = map[string]time.Time{key: time.Now()} + + if _, err := r.handleDeletion(t.Context(), shard); err != nil { + t.Fatalf("handleDeletion() error = %v", err) + } + + if _, ok := r.postureStrikes[key]; ok { + t.Errorf("posture strikes entry for %s survived handleDeletion", key) + } + if _, ok := r.notConvergedSince[key]; ok { + t.Errorf("not-converged-since entry for %s survived handleDeletion", key) + } +} + +// TestReconcileForgetsStrikesOnNotFound drives the not-found cleanup through +// Reconcile itself: a Shard already gone from the API server (the common case +// once handleDeletion above has already run and removed the finalizer) must +// still have its strike entries dropped, as a backstop. +func TestReconcileForgetsStrikesOnNotFound(t *testing.T) { + scheme := postureTestScheme(t) + c := fake.NewClientBuilder().WithScheme(scheme).Build() + r := &ShardReconciler{ + Client: c, + Scheme: scheme, + Recorder: record.NewFakeRecorder(20), + } + + key := "default/gone-shard" + r.postureStrikes = map[string]int{key: 3} + r.notConvergedSince = map[string]time.Time{key: time.Now()} + + req := ctrl.Request{ + NamespacedName: types.NamespacedName{Namespace: "default", Name: "gone-shard"}, + } + if _, err := r.Reconcile(t.Context(), req); err != nil { + t.Fatalf("Reconcile() error = %v", err) + } + + if _, ok := r.postureStrikes[key]; ok { + t.Errorf("posture strikes entry for %s survived Reconcile on a missing Shard", key) + } + if _, ok := r.notConvergedSince[key]; ok { + t.Errorf("not-converged-since entry for %s survived Reconcile on a missing Shard", key) + } +} + +// gateTestReconciler is postureTestReconciler plus a Pod status subresource, +// for tests that read back the PoolerDataReady gate condition a real +// apiserver would only apply through Status().Patch. +func gateTestReconciler( + t *testing.T, + shard *multigresv1alpha1.Shard, + rpc rpcclient.MultipoolerClient, + objects ...client.Object, +) (*ShardReconciler, client.Client) { + t.Helper() + scheme := postureTestScheme(t) + allObjects := append([]client.Object{shard}, objects...) + c := fake.NewClientBuilder(). + WithScheme(scheme). + WithObjects(allObjects...). + WithStatusSubresource(&multigresv1alpha1.Shard{}, &corev1.Pod{}). + Build() + return &ShardReconciler{ + Client: c, + Scheme: scheme, + Recorder: record.NewFakeRecorder(20), + PoolerClients: poolerclient.Static(rpc), + CreateTopoStore: newMemoryTopoFactory(), + }, c +} + +// registeredReplica registers id in store as a non-primary member of shard, +// with an RPC status that reads Ready via poolerReadiness: initialized, +// accepting connections, cohort-eligible, and a committed member of its own +// single-member rule. Standing in for "this pooler has finished bootstrapping +// and has somewhere to belong," independent of whether anyone in the shard has +// been elected primary. +func registeredReplica( + t *testing.T, + store topoclient.Store, + rpc *rpcclient.FakeClient, + shard *multigresv1alpha1.Shard, + cell, name string, +) topoclient.ComponentID { + t.Helper() + id := &clustermetadata.ID{Cell: cell, Name: name} + if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + Id: id, + Hostname: name, + ShardKey: &clustermetadata.ShardKey{ + Database: string(shard.Spec.DatabaseName), + TableGroup: string(shard.Spec.TableGroupName), + Shard: string(shard.Spec.ShardName), + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + }, + }, false); err != nil { + t.Fatalf("register pooler %s: %v", name, err) + } + + componentID := topoclient.ComponentIDString(id) + rpc.SetStatusResponse(componentID, readyStatusResponse(id)) + return componentID +} + +// readyStatusResponse is a fully posture-ready StatusResponse for id: a +// replica, initialized, accepting connections, cohort-eligible, and a +// committed member of its own single-member rule. +func readyStatusResponse(id *clustermetadata.ID) *multipoolermanagerdatapb.StatusResponse { + return &multipoolermanagerdatapb.StatusResponse{ + Status: &multipoolermanagerdatapb.Status{ + IsInitialized: true, + PostgresReady: true, + PostgresStatus: multipoolermanagerdatapb.PostgresStatus_POSTGRES_STATUS_STANDBY, + }, + AvailabilityStatus: &clustermetadata.AvailabilityStatus{ + CohortEligibilityStatus: &clustermetadata.CohortEligibilityStatus{ + Signal: clustermetadata.CohortEligibilitySignal_COHORT_ELIGIBILITY_SIGNAL_ELIGIBLE, + }, + }, + ConsensusStatus: &clustermetadata.ConsensusStatus{ + Id: id, + CurrentPosition: &clustermetadata.PoolerPosition{ + Position: &clustermetadata.RulePosition{ + Decision: &clustermetadata.ShardRule{ + RuleNumber: &clustermetadata.RuleNumber{CoordinatorTerm: 1}, + LeaderId: id, + CohortMembers: []*clustermetadata.ID{id}, + DurabilityPolicy: topoclient.AtLeastN(1), + }, + }, + }, + }, + } +} + +// notYetSettledReplica registers id in store as a non-primary member of +// shard, with an RPC status that is initialized, accepting connections, and +// cohort-eligible, but has not yet committed a rule naming any cohort +// members. Standing in for a shard mid-bootstrap where every pooler has +// registered but multiorch has not yet elected a primary or committed a +// durability rule: nobody is a "postgres primary" (so nothing looks +// inconsistent) and nobody has an unmatched topology entry or an unreadable +// RPC (so nothing looks incomplete), but nobody is ready either. +func notYetSettledReplica( + t *testing.T, + store topoclient.Store, + rpc *rpcclient.FakeClient, + shard *multigresv1alpha1.Shard, + cell, name string, +) *clustermetadata.ID { + t.Helper() + id := &clustermetadata.ID{Cell: cell, Name: name} + if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + Id: id, + Hostname: name, + ShardKey: &clustermetadata.ShardKey{ + Database: string(shard.Spec.DatabaseName), + TableGroup: string(shard.Spec.TableGroupName), + Shard: string(shard.Spec.ShardName), + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + }, + }, false); err != nil { + t.Fatalf("register pooler %s: %v", name, err) + } + + componentID := topoclient.ComponentIDString(id) + rpc.SetStatusResponse(componentID, &multipoolermanagerdatapb.StatusResponse{ + Status: &multipoolermanagerdatapb.Status{ + IsInitialized: true, + PostgresReady: true, + PostgresStatus: multipoolermanagerdatapb.PostgresStatus_POSTGRES_STATUS_STANDBY, + }, + AvailabilityStatus: &clustermetadata.AvailabilityStatus{ + CohortEligibilityStatus: &clustermetadata.CohortEligibilityStatus{ + Signal: clustermetadata.CohortEligibilitySignal_COHORT_ELIGIBILITY_SIGNAL_ELIGIBLE, + }, + }, + // No ConsensusStatus: nothing has been committed yet, so + // committedCohortContains is false for everyone. + }) + return id +} + +// TestReconcilePostureConvergedShardReturnsZero pins the hot-loop fix's own +// invariant: a converged shard returns 0 and leaves no map entry. It includes +// a multiorch pod built from BuildMultiorchDeployment, since a shard's +// multiorch pod carries the same four identity labels as its pool pods, so a +// selector that forgets to scope to pool pods seeds it AwaitingRegistration +// forever and this shard never returns 0. +func TestReconcilePostureConvergedShardReturnsZero(t *testing.T) { + shard := postureTestShard() + shard.Labels[metadata.LabelMultigresDatabase] = "database" + shard.Labels[metadata.LabelMultigresTableGroup] = "table-group" + shard.Labels[metadata.LabelMultigresShard] = "0" + + dep, err := BuildMultiorchDeployment(shard, "cell1", postureTestScheme(t)) + if err != nil { + t.Fatalf("BuildMultiorchDeployment() error = %v", err) + } + orch := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{ + Name: "multiorch-abc", Namespace: shard.Namespace, Labels: dep.Spec.Template.Labels, + }} + + _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") + store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) + defer func() { _ = store.Close() }() + rpc := rpcclient.NewFakeClient() + registeredReplica(t, store, rpc, shard, "cell1", "pooler-0") + + pool := postureTestPod() + r, _ := postureTestReconciler(t, shard, rpc, pool, orch) + + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("reconcilePosture() error = %v", err) + } + if delay != 0 { + t.Errorf("delay = %v, want 0 for a converged shard", delay) + } + key := fmt.Sprintf("%s/%s", shard.Namespace, shard.Name) + if _, ok := r.postureStrikes[key]; ok { + t.Errorf("posture strikes entry left for a converged shard") + } + if _, ok := r.notConvergedSince[key]; ok { + t.Errorf("not-converged-since entry left for a converged shard") + } + if conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { + t.Errorf("conditions = %#v, want no failure for a converged shard", shard.Status.Conditions) + } +} + +// TestReconcilePostureAcceptedIncompleteObservationStillRequeues covers an +// accepted-but-not-ready observation: a Status RPC failure makes a pod +// UNKNOWN, which makes it Incomplete, which makes the shard unsettled. After +// the strike threshold the observation is accepted into status, but the +// failing pod has also never reached posture readiness, so this must keep +// requesting a requeue rather than stranding it until the 10h resync. +func TestReconcilePostureAcceptedIncompleteObservationStillRequeues(t *testing.T) { + shard := postureTestShard() + _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") + store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) + defer func() { _ = store.Close() }() + rpc := rpcclient.NewFakeClient() + registeredReplica(t, store, rpc, shard, "cell1", "pooler-0") + bad := registeredReplica(t, store, rpc, shard, "cell1", "pooler-1") + rpc.Errors[bad] = errors.New("dial: connection refused") + + p0 := postureTestPod() + p1 := postureTestPod() + p1.Name = "pooler-1" + r, _ := postureTestReconciler(t, shard, rpc, p0, p1) + + first, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("first reconcilePosture() error = %v", err) + } + if first != postureDebounceRequeueDelay { + t.Errorf("first delay = %v, want the %v debounce", first, postureDebounceRequeueDelay) + } + + for pass := 2; pass <= 4; pass++ { + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("pass %d reconcilePosture() error = %v", pass, err) + } + if delay <= 0 { + t.Errorf( + "pass %d: delay = %v, want non-zero (RPC failure still unresolved)", pass, delay, + ) + } + } +} + +// TestReconcilePostureAcceptedMismatchAndNotReadyStillRequeues covers a +// mismatch compounded with a pod that has never reached posture readiness at +// all (this mock pooler never reports IsInitialized/PostgresReady): a +// replica reporting postgres PRIMARY is a role mismatch. After the strike +// threshold it is accepted as PostureConsistent=False, and this must keep +// requesting a requeue. +func TestReconcilePostureAcceptedMismatchAndNotReadyStillRequeues(t *testing.T) { + shard := postureTestShard() + _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") + store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) + defer func() { _ = store.Close() }() + rpc := rpcclient.NewFakeClient() + id := &clustermetadata.ID{Cell: "cell1", Name: "pooler-0"} + if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + Id: id, + Hostname: "pooler-0", + ShardKey: &clustermetadata.ShardKey{ + Database: string(shard.Spec.DatabaseName), + TableGroup: string(shard.Spec.TableGroupName), + Shard: string(shard.Spec.ShardName), + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + }, + }, false); err != nil { + t.Fatalf("register pooler: %v", err) + } + componentID := topoclient.ComponentIDString(id) + rpc.SetStatusResponse(componentID, &multipoolermanagerdatapb.StatusResponse{ + Status: &multipoolermanagerdatapb.Status{ + PostgresStatus: multipoolermanagerdatapb.PostgresStatus_POSTGRES_STATUS_PRIMARY, + }, + }) + r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) + + first, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("first reconcilePosture() error = %v", err) + } + if first != postureDebounceRequeueDelay { + t.Errorf("first delay = %v, want the %v debounce", first, postureDebounceRequeueDelay) + } + + for pass := 2; pass <= 4; pass++ { + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("pass %d reconcilePosture() error = %v", pass, err) + } + if delay <= 0 { + t.Errorf("pass %d: delay = %v, want non-zero (mismatch still unresolved)", pass, delay) + } + if !conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { + t.Errorf("pass %d: conditions = %#v, want PostureConsistent=False once accepted", + pass, shard.Status.Conditions) + } + } +} + +// TestReconcilePostureAcceptedMismatchWithReadyPodsStillRequeues is the other +// half of the accepted-mismatch case: the pod itself is fully ready +// (registeredReplica: initialized, accepting connections, cohort-eligible, +// a committed cohort member), and the only thing wrong is that it reports +// postgres PRIMARY while topology still has it as REPLICA. anyPodNotReady is +// false throughout, so unsettled is the only thing driving this requeue: a +// shard where every pod is ready but multiorch and postgres disagree about +// who is primary must not be left Degraded until the 10h resync once that +// disagreement is accepted into status. +func TestReconcilePostureAcceptedMismatchWithReadyPodsStillRequeues(t *testing.T) { + shard := postureTestShard() + _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") + store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) + defer func() { _ = store.Close() }() + rpc := rpcclient.NewFakeClient() + id := registeredReplica(t, store, rpc, shard, "cell1", "pooler-0") + rpc.StatusResponses[id].Response.Status.PostgresStatus = multipoolermanagerdatapb.PostgresStatus_POSTGRES_STATUS_PRIMARY + r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) + + first, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("first reconcilePosture() error = %v", err) + } + if first != postureDebounceRequeueDelay { + t.Errorf("first delay = %v, want the %v debounce", first, postureDebounceRequeueDelay) + } + + for pass := 2; pass <= 4; pass++ { + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("pass %d reconcilePosture() error = %v", pass, err) + } + if delay <= 0 { + t.Errorf( + "pass %d: delay = %v, want non-zero (mismatch still unresolved, though the pod is ready)", + pass, + delay, + ) + } + if !conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { + t.Errorf("pass %d: conditions = %#v, want PostureConsistent=False once accepted", + pass, shard.Status.Conditions) + } + } +} + +// TestReconcilePostureRequeuesWhileAPodAwaitsItsPooler covers a shard with one +// settled, registered pooler and one managed pod that never registers: it +// must keep requesting a requeue, and the requested delay must grow with +// elapsed wall-clock time rather than sit at a fixed floor forever. +// +// Drives a fake clock directly so the growth assertion is exact rather than +// "second draw happened to be bigger": elapsed time is measured directly +// regardless of how many reconcile passes it took to get there, so a mutation +// that turns the backoff back into a per-pass count cannot pass by chance. +func TestReconcilePostureRequeuesWhileAPodAwaitsItsPooler(t *testing.T) { + shard := postureTestShard() + _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") + store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) + defer func() { _ = store.Close() }() + + rpc := rpcclient.NewFakeClient() + registeredReplica(t, store, rpc, shard, "cell1", "pooler-0") + + settledPod := postureTestPod() + awaitingPod := postureTestPod() + awaitingPod.Name = "pooler-1" + + r, _ := postureTestReconciler(t, shard, rpc, settledPod, awaitingPod) + clk := &fakeClock{t: time.Unix(1_700_000_000, 0)} + r.Clock = clk.now + + for _, tc := range []struct { + advance time.Duration + elapsed time.Duration + }{ + {advance: 0, elapsed: 0}, + {advance: 20 * time.Second, elapsed: 20 * time.Second}, + {advance: 50 * time.Second, elapsed: 70 * time.Second}, + } { + clk.advance(tc.advance) + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("reconcilePosture() error = %v", err) + } + min, max := wantDelayRange(tc.elapsed) + if delay < min || delay > max { + t.Errorf("at elapsed=%v: delay = %v, want in [%v, %v]", tc.elapsed, delay, min, max) + } + } +} + +// TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected covers the +// bring-up state a minimal cluster passes through before a primary is +// elected: every managed pod has a topology entry (nothing is "awaiting +// registration" in the topology-match sense), and nothing is inconsistent or +// incomplete (Evaluate only compares observed postgres primaries against +// topology roles and finds none of either), but nobody has committed a cohort +// membership because no primary has been elected yet. +// +// Drives a fake clock through several passes to pin the actual delay range, +// checks the PoolerDataReady gate stays False while waiting, then flips the +// fixture to a committed primary and checks the requeue stops, the gate goes +// True, and the not-converged-since entry is gone. +func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { + shard := postureTestShard() + _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") + store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) + defer func() { _ = store.Close() }() + + rpc := rpcclient.NewFakeClient() + id0 := notYetSettledReplica(t, store, rpc, shard, "cell1", "pooler-0") + id1 := notYetSettledReplica(t, store, rpc, shard, "cell1", "pooler-1") + + pod0 := postureTestPod() + pod1 := postureTestPod() + pod1.Name = "pooler-1" + + r, c := gateTestReconciler(t, shard, rpc, pod0, pod1) + clk := &fakeClock{t: time.Unix(1_700_000_000, 0)} + r.Clock = clk.now + + for _, tc := range []struct { + advance time.Duration + elapsed time.Duration + }{ + {advance: 0, elapsed: 0}, + {advance: 20 * time.Second, elapsed: 20 * time.Second}, + {advance: 50 * time.Second, elapsed: 70 * time.Second}, + } { + clk.advance(tc.advance) + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("reconcilePosture() error = %v", err) + } + min, max := wantDelayRange(tc.elapsed) + if delay < min || delay > max { + t.Errorf("at elapsed=%v: delay = %v, want in [%v, %v]", tc.elapsed, delay, min, max) + } + } + + for _, pod := range []*corev1.Pod{pod0, pod1} { + got := &corev1.Pod{} + if err := c.Get(t.Context(), client.ObjectKeyFromObject(pod), got); err != nil { + t.Fatalf("get pod %s: %v", pod.Name, err) + } + condition := readinessCondition(got.Status.Conditions) + if condition == nil || condition.Status != corev1.ConditionFalse { + t.Errorf("pod %s readiness condition = %#v, want False while waiting for a primary", + pod.Name, condition) + } + } + + // The fixture reaches a settled state: both poolers commit a rule naming + // pooler-0 as leader, so both are cohort-eligible members of the same + // durability rule and pooler-0's postgres reports PRIMARY, matching the + // topology role a leader-designate needs. + if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + Id: id0, + Hostname: "pooler-0", + ShardKey: &clustermetadata.ShardKey{ + Database: string(shard.Spec.DatabaseName), + TableGroup: string(shard.Spec.TableGroupName), + Shard: string(shard.Spec.ShardName), + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_PRIMARY, + }, + }, true); err != nil { + t.Fatalf("promote pooler-0 in topology: %v", err) + } + rule := &clustermetadata.ShardRule{ + RuleNumber: &clustermetadata.RuleNumber{CoordinatorTerm: 1}, + LeaderId: id0, + CohortMembers: []*clustermetadata.ID{id0, id1}, + DurabilityPolicy: topoclient.AtLeastN(1), + } + for _, elected := range []struct { + id *clustermetadata.ID + primary bool + }{{id0, true}, {id1, false}} { + status := multipoolermanagerdatapb.PostgresStatus_POSTGRES_STATUS_STANDBY + if elected.primary { + status = multipoolermanagerdatapb.PostgresStatus_POSTGRES_STATUS_PRIMARY + } + componentID := topoclient.ComponentIDString(elected.id) + rpc.SetStatusResponse(componentID, &multipoolermanagerdatapb.StatusResponse{ + Status: &multipoolermanagerdatapb.Status{ + IsInitialized: true, + PostgresReady: true, + PostgresStatus: status, + }, + AvailabilityStatus: &clustermetadata.AvailabilityStatus{ + CohortEligibilityStatus: &clustermetadata.CohortEligibilityStatus{ + Signal: clustermetadata.CohortEligibilitySignal_COHORT_ELIGIBILITY_SIGNAL_ELIGIBLE, + }, + }, + ConsensusStatus: &clustermetadata.ConsensusStatus{ + Id: elected.id, + CurrentPosition: &clustermetadata.PoolerPosition{ + Position: &clustermetadata.RulePosition{Decision: rule}, + }, + }, + }) + } + + clk.advance(time.Second) + delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) + if err != nil { + t.Fatalf("reconcilePosture() after election error = %v", err) + } + if delay != 0 { + t.Errorf("delay after a primary is elected = %v, want 0", delay) + } + key := fmt.Sprintf("%s/%s", shard.Namespace, shard.Name) + if _, ok := r.notConvergedSince[key]; ok { + t.Errorf("not-converged-since entry left after a primary is elected") + } + if conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { + t.Errorf( + "conditions = %#v, want no failure once a primary is elected", + shard.Status.Conditions, + ) + } + + for _, pod := range []*corev1.Pod{pod0, pod1} { + got := &corev1.Pod{} + if err := c.Get(t.Context(), client.ObjectKeyFromObject(pod), got); err != nil { + t.Fatalf("get pod %s: %v", pod.Name, err) + } + condition := readinessCondition(got.Status.Conditions) + if condition == nil || condition.Status != corev1.ConditionTrue { + t.Errorf("pod %s readiness condition = %#v, want True once a primary is elected", + pod.Name, condition) + } + } +} diff --git a/pkg/resource-handler/controller/shard/shard_controller.go b/pkg/resource-handler/controller/shard/shard_controller.go index d1229ef4..a36c2a7e 100644 --- a/pkg/resource-handler/controller/shard/shard_controller.go +++ b/pkg/resource-handler/controller/shard/shard_controller.go @@ -79,9 +79,22 @@ type ShardReconciler struct { APIReader client.Reader PoolerClients poolerclient.Resolver CreateTopoStore func(*multigresv1alpha1.Shard) (topoclient.Store, error) + // Clock overrides time.Now for recordNotConverged's elapsed-time math. + // Nil in production; tests inject it to drive the readiness backoff + // deterministically. + Clock func() time.Time postureStrikesMu sync.Mutex postureStrikes map[string]int + + // notConvergedMu and notConvergedSince record, per shard, the time it was + // first observed not converged: unsettled (posture debounce past + // threshold, i.e. an accepted mismatch or Incomplete observation), or some + // managed pod not yet posture-ready. Kept separate from postureStrikes, + // which gates posture.Apply: folding this into that counter would change + // when Apply fires. + notConvergedMu sync.Mutex + notConvergedSince map[string]time.Time } // Reconcile manages pool pods, PVCs, services, and data-plane topology for a Shard. @@ -109,6 +122,7 @@ func (r *ShardReconciler) Reconcile( if err := r.Get(ctx, req.NamespacedName, shard); err != nil { if errors.IsNotFound(err) { logger.Info("Shard resource not found, ignoring") + r.forgetStrikes(req.Namespace, req.Name) return ctrl.Result{}, nil } monitoring.RecordSpanError(span, err) diff --git a/test/suite/scenario_thrash_test.go b/test/suite/scenario_thrash_test.go index 71d14d63..f4718deb 100644 --- a/test/suite/scenario_thrash_test.go +++ b/test/suite/scenario_thrash_test.go @@ -43,10 +43,10 @@ const thrashPoolName = PoolName("default") // changes, then requires the namespace to go quiet and every pool PVC to be // bound to a live pod. // -// It does not assert on this thrashed shard's status.podRoles: whether the -// last scale-up's role lands there is a race against the defect that -// pinPoolScaleUpRoleStale constructs deterministically on a namespace of its -// own, so asserting it here would only sample that race. +// It does not assert on this thrashed shard's status.podRoles itself: +// requirePoolScaleUpLandsRole asserts that claim deterministically on a +// namespace of its own, and asserting it here too would only be a second, +// weaker sample of the same thing. func testPoolReplicaThrash(t *testing.T) { c := newCase(t) cluster := c.poolThrashCluster("pool-thrash", 1) @@ -60,15 +60,19 @@ func testPoolReplicaThrash(t *testing.T) { c.scalePoolTo(cluster, n) } - // MGO-POOL-SCALEUP-ROLE-STALE: once a shard reconciled to Healthy with N - // poolers, raising replicasPerCell to add an (N+1)th sometimes never gets - // that pooler's role into shard.Status.PodRoles, permanently. + // Fixed: scaling a pool from N to N+1 lands the new pooler's role in + // shard.Status.PodRoles even when the pooler registers in topology after + // the shard has already reconciled to Healthy. The shard controller's + // readiness requeue is what closes the gap, since nothing else notices a + // registration: it is a write to the topology store, with no Kubernetes + // event behind it. // - // Whether it bites depends on whether the pooler registers before or - // after the reconcile that declares the shard converged, so sampled - // naturally it reproduces about a third of the time. The pin constructs - // that ordering instead, which makes it deterministic. - pinPoolScaleUpRoleStale(t) + // The pin this replaces was statistical because the defect was: whether it + // bit depended on whether the pooler registered before or after the + // reconcile that declared the shard converged, so it reproduced about a + // third of the time. Fixed, it is deterministic, so a positive assertion + // replaces what used to be a KnownDefect pin. + requirePoolScaleUpLandsRole(t) c.RequireQuiescent(10*time.Second, 90*time.Second) @@ -84,20 +88,19 @@ func testPoolReplicaThrash(t *testing.T) { c.requireNoOrphanedPoolPVCs(live) } -// pinPoolScaleUpRoleStale pins that scaling a pool from one to two does not -// land the new pooler's role in status.podRoles when the pooler registers -// after the shard has already converged. +// requirePoolScaleUpLandsRole asserts that scaling a pool from one to two +// lands the new pooler's role in status.podRoles. // // On its own namespace and its own cluster, with no thrash, because the -// defect never needed one: a bare single scale-up reproduces it, and the -// thrash above only found it first. +// defect this replaces never needed one: a bare single scale-up reproduced +// it, and the thrash above only found it first. // // Nothing wakes the shard once it is Healthy: a registration is a write to -// the topology store, with no Kubernetes event behind it, and the shard does -// not requeue itself while a managed pod is still awaiting its pooler. The fix -// is that requeue; with it the role lands well inside a second here, because -// the suite compresses requeues, so the window below is generous. -func pinPoolScaleUpRoleStale(t *testing.T) { +// the topology store, with no Kubernetes event behind it. The shard +// controller now requeues on a backoff while any managed pod has not reached +// posture readiness, so the role lands well inside a second here, because the +// suite compresses requeues, so the window below is generous. +func requirePoolScaleUpLandsRole(t *testing.T) { t.Helper() c := newCase(t) @@ -162,31 +165,24 @@ func pinPoolScaleUpRoleStale(t *testing.T) { // Now the pooler appears, with no Kubernetes event to announce it: a // registration is a write to etcd. Only a requeue the operator asked for - // itself can notice, which is the thing this pins. + // itself can notice, which is the thing this asserts. release() - // Retire this pin by replacing it with c.Eventually on the same condition. - c.KnownDefect("MGO-POOL-SCALEUP-ROLE-STALE", func() error { - var members Members - deadline := time.Now().Add(20 * time.Second) - for { - var err error - members, err = MembersOf(c.Context(), c.Client(), key) + c.Eventually( + 60*time.Second, + "the scaled-up pooler's role to reach status.podRoles", + func() error { + members, err := MembersOf(c.Context(), c.Client(), key) // A read failure is the check's own setup failing, not the - // defect, so it fails the test rather than keeping the pin green. + // convergence it is waiting for, so it fails the test immediately + // rather than retrying it silently until the timeout. c.NoError(err, "read shard members") - if len(members.Replicas) == 1 && len(members.Quarantined) == 0 { - return nil - } - if time.Now().After(deadline) { - break + if len(members.Replicas) != 1 || len(members.Quarantined) != 0 { + return fmt.Errorf("got %+v", members) } - time.Sleep(250 * time.Millisecond) - } - return fmt.Errorf( - "the scaled-up pooler registered after the shard converged and its role "+ - "never reached status.podRoles within 20s: got %+v", members) - }) + return nil + }, + ) } func (c *C) poolThrashCluster( @@ -243,10 +239,9 @@ func (c *C) scalePoolTo(cluster *MultigresCluster, n int32) { // liveReadyPoolPodNames lists Ready pods belonging to the thrashed pool. // // This deliberately does not go through MembersOf/shard.Status.PodRoles: that -// path is exactly what MGO-POOL-SCALEUP-ROLE-STALE (see -// testPoolReplicaThrash) breaks, and a pod being live is a Kubernetes-level -// fact independent of whether the operator's own role bookkeeping has caught -// up to it. +// path is the one testPoolReplicaThrash's pool-scale-up assertion covers +// directly, and a pod being live is a Kubernetes-level fact independent of +// whether the operator's own role bookkeeping has caught up to it. func (c *C) liveReadyPoolPodNames() []string { c.Helper() pods := &corev1.PodList{} From ea8b5fb13763c9eb405053ac0f0c5582dea9d4e1 Mon Sep 17 00:00:00 2001 From: Brent Graveland Date: Sat, 26 Sep 2026 19:00:35 -0600 Subject: [PATCH 7/7] test: convert the operator's tests to testkit/assert Mechanical, produced entirely by the tool rather than by hand, after bumping github.com/multigres/testkit to v0.2.0 (the release that adds the methods the tool now converts to): go install github.com/multigres/testkit/tools/assertfix@v0.2.2 go fix -fixtool=assertfix ./... go fix -tags=integration,verbose -fixtool=assertfix ./... go fix -tags=e2e -fixtool=assertfix ./... gofmt -w golangci-lint fmt The three runs are needed because go fix only sees files matching the active build tags. The last step rewraps converted lines that exceed golines' limit; the same files needed no formatting before conversion. 167 test files, 91% of assertion sites (3,954 of 4,337 outside test/suite and tools/observer), and no test file outside test/suite imports testify any more. The rest is left untouched on purpose: the tool declines any site it cannot map with certainty, for example a compound condition whose failure message would panic or have side effects if evaluated on the passing path. No test changes behaviour: t.Errorf continues and t.Fatalf aborts, so a scope with both gets assert.NewCollecting and its aborting sites are written c.Require().X(...). Top-level test counts are identical under all three tag sets, and re-running the tool over the result changes nothing. Signed-off-by: Brent Graveland --- api/v1alpha1/etcd_maintenance_test.go | 11 +- api/v1alpha1/observability_helpers_test.go | 136 +- api/v1alpha1/topo_client_tls_helpers_test.go | 46 +- cmd/multigres-operator/main_test.go | 29 +- config/manager/manager_test.go | 24 +- go.mod | 4 +- go.sum | 4 +- pkg/cert/generator_test.go | 128 +- pkg/cert/manager_test.go | 281 ++-- .../multigrescluster/builders_cell_test.go | 81 +- .../multigrescluster/builders_global_test.go | 517 +++---- .../builders_tablegroup_test.go | 206 +-- .../multigrescluster/builders_test.go | 6 +- .../multigrescluster/certificate_test.go | 475 +++--- .../multigrescluster/certificate_topo_test.go | 234 ++- .../multigrescluster/images_test.go | 467 +++--- .../integration_adminweb_test.go | 144 +- .../integration_gateway_test.go | 152 +- .../integration_lifecycle_test.go | 77 +- ...integration_resolution_enforcement_test.go | 157 +- .../multigrescluster/integration_test.go | 253 ++- .../integration_validation_test.go | 101 +- .../multigrescluster_controller_test.go | 352 ++--- .../multigrescluster/reconcile_cells_test.go | 56 +- .../reconcile_database_test.go | 55 +- .../multigrescluster/reconcile_global_test.go | 261 ++-- .../reconcile_global_topo_tls_test.go | 67 +- .../reconcile_topology_test.go | 164 +- .../reconcile_topology_tls_test.go | 44 +- .../multigrescluster/status_test.go | 297 ++-- .../controller/tablegroup/builders_test.go | 172 +-- .../integration_consistency_test.go | 69 +- .../tablegroup/integration_lifecycle_test.go | 78 +- .../controller/tablegroup/integration_test.go | 30 +- .../tablegroup/reconcile_shards_test.go | 58 +- .../controller/tablegroup/status_test.go | 172 +-- .../tablegroup/tablegroup_controller_test.go | 393 ++--- .../backuphealth/backuphealth_test.go | 136 +- pkg/data-handler/drain/drain_test.go | 66 +- .../poolerclient/resolver_test.go | 318 ++-- pkg/data-handler/posture/disruption_test.go | 6 +- pkg/data-handler/posture/posture_test.go | 246 ++- pkg/data-handler/topo/cell_test.go | 180 +-- pkg/data-handler/topo/database_test.go | 243 ++- pkg/data-handler/topo/idempotency_test.go | 39 +- pkg/data-handler/topo/pooler_test.go | 235 ++- pkg/data-handler/topo/store_internal_test.go | 42 +- pkg/data-handler/topo/store_test.go | 42 +- pkg/data-handler/topo/topology_test.go | 377 ++--- pkg/gc/pvc/pvc_test.go | 47 +- pkg/images/images_test.go | 72 +- pkg/monitoring/metrics_test.go | 21 +- pkg/monitoring/recorder_test.go | 109 +- pkg/monitoring/tracing_test.go | 192 +-- pkg/postgresconfig/classify_test.go | 35 +- pkg/postgresconfig/hash_test.go | 98 +- pkg/postgresconfig/reload_marker_test.go | 51 +- pkg/postgresconfig/render_precedence_test.go | 19 +- pkg/postgresconfig/render_test.go | 119 +- pkg/postgresconfig/sizing_test.go | 140 +- pkg/postgresconfig/validate_test.go | 46 +- pkg/resolver/buffer_defaults_test.go | 38 +- pkg/resolver/cell_test.go | 91 +- pkg/resolver/cluster_test.go | 115 +- pkg/resolver/etcd_maintenance_test.go | 8 +- pkg/resolver/resolver_test.go | 384 ++--- pkg/resolver/shard_test.go | 329 ++-- pkg/resolver/validation_test.go | 46 +- .../cell/cell_controller_internal_test.go | 88 +- .../controller/cell/cell_controller_test.go | 138 +- .../controller/cell/integration_test.go | 308 +++- .../controller/cell/local_toposerver_test.go | 223 ++- .../controller/cell/multigateway_test.go | 362 ++--- .../controller/shard/configmap_test.go | 148 +- .../controller/shard/containers_test.go | 441 ++---- .../controller/shard/disruption_test.go | 145 +- .../controller/shard/integration_test.go | 650 +++++--- .../shard/maintenance_surge_test.go | 190 ++- .../controller/shard/multiorch_test.go | 68 +- .../controller/shard/pool_pod_test.go | 591 +++---- .../controller/shard/pool_pvc_test.go | 163 +- .../controller/shard/pool_service_test.go | 7 +- .../controller/shard/ports_test.go | 185 +-- .../controller/shard/postgres_config_test.go | 242 ++- ...oncile_data_plane_posture_internal_test.go | 336 ++-- .../shard/reconcile_deletion_internal_test.go | 23 +- .../shard/reconcile_deletion_test.go | 277 ++-- .../reconcile_quarantine_internal_test.go | 136 +- .../shard/reconcile_readiness_test.go | 41 +- .../shard/reconcile_shared_infra_test.go | 20 +- .../shard/registration_requeue_test.go | 331 ++-- .../controller/shard/reload_internal_test.go | 88 +- .../controller/shard/secret_test.go | 41 +- .../shard/shard_controller_internal_test.go | 1372 +++++++---------- .../controller/shard/shard_controller_test.go | 438 ++---- .../controller/shard/shard_pdb_test.go | 88 +- .../shard/storage_class_guard_test.go | 548 +++---- .../controller/shard/topo_client_tls_test.go | 74 +- .../controller/storage/pvc_test.go | 7 +- .../controller/toposerver/certificate_test.go | 246 ++- .../toposerver/container_env_test.go | 28 +- .../controller/toposerver/integration_test.go | 168 +- .../toposerver/maintenance_client_test.go | 26 +- .../toposerver/maintenance_etcd_test.go | 101 +- .../toposerver/maintenance_reconcile_test.go | 97 +- .../maintenance_status_integration_test.go | 38 +- .../controller/toposerver/maintenance_test.go | 149 +- .../controller/toposerver/pdb_test.go | 17 +- .../controller/toposerver/ports_test.go | 15 +- .../controller/toposerver/service_test.go | 11 +- .../controller/toposerver/statefulset_test.go | 37 +- .../toposerver/storage_class_guard_test.go | 118 +- .../toposerver_controller_internal_test.go | 71 +- .../toposerver/toposerver_controller_test.go | 241 ++- pkg/testutil/compare_internal_test.go | 17 +- pkg/testutil/compare_test.go | 20 +- pkg/testutil/e2e_test.go | 184 +-- pkg/testutil/envtest_internal_test.go | 114 +- pkg/testutil/envtest_test.go | 32 +- pkg/testutil/fake_client_test.go | 63 +- pkg/testutil/kind_test.go | 82 +- .../resource_watcher_cache_internal_test.go | 38 +- ...resource_watcher_deletion_internal_test.go | 15 +- .../resource_watcher_deletion_test.go | 87 +- .../resource_watcher_internal_test.go | 44 +- ...resource_watcher_listener_internal_test.go | 28 +- .../resource_watcher_match_edge_cases_test.go | 23 +- .../resource_watcher_match_internal_test.go | 61 +- pkg/testutil/resource_watcher_match_test.go | 67 +- pkg/testutil/resource_watcher_test.go | 146 +- pkg/topology/roots_test.go | 114 +- pkg/util/certs/certs_test.go | 260 ++-- pkg/util/metadata/labels_test.go | 47 +- pkg/util/metadata/project_ref_test.go | 6 +- pkg/util/name/name_test.go | 64 +- pkg/util/pvc/orphan_test.go | 51 +- pkg/util/pvc/retention_test.go | 11 +- pkg/util/status/conditions_test.go | 40 +- pkg/util/status/phase_test.go | 23 +- pkg/webhook/cel_validation_test.go | 385 +++-- .../etcd_maintenance_validation_test.go | 27 +- pkg/webhook/handlers/defaulter_test.go | 28 +- pkg/webhook/handlers/validator_test.go | 246 ++- pkg/webhook/integration_test.go | 286 ++-- pkg/webhook/pki_test.go | 198 ++- pkg/webhook/setup_test.go | 28 +- test/e2e/dedicated/deletion/deletion_test.go | 44 +- test/e2e/dedicated/inline/inline_test.go | 11 +- test/e2e/dedicated/minimal/minimal_test.go | 11 +- .../e2e/dedicated/templated/templated_test.go | 23 +- test/e2e/framework/diagnostics_test.go | 27 +- test/e2e/framework/fixtures_test.go | 37 +- test/e2e/framework/helpers_test.go | 73 +- test/e2e/framework/image_overrides_test.go | 37 +- test/e2e/framework/images_test.go | 22 +- test/e2e/shared/deletion/deletion_test.go | 166 +- test/e2e/shared/drain/drain_test.go | 331 ++-- test/e2e/shared/inline/inline_test.go | 11 +- test/e2e/shared/minimal/minimal_test.go | 11 +- .../postgresconfig/baseline_over_ref_test.go | 15 +- .../e2e/shared/postgresconfig/helpers_test.go | 15 +- .../postgresconfig/logconnections_test.go | 24 +- .../postgresconfig/postgresconfig_test.go | 77 +- test/e2e/shared/scaling/scaling_test.go | 38 +- test/e2e/shared/templated/templated_test.go | 23 +- test/e2e/shared/templates/templates_test.go | 127 +- test/e2e/shared/topotls/topotls_test.go | 35 +- .../shared/verification/verification_test.go | 216 ++- test/e2e/shared/webhook/webhook_test.go | 20 +- 169 files changed, 9942 insertions(+), 13110 deletions(-) diff --git a/api/v1alpha1/etcd_maintenance_test.go b/api/v1alpha1/etcd_maintenance_test.go index f4ddd985..214f518f 100644 --- a/api/v1alpha1/etcd_maintenance_test.go +++ b/api/v1alpha1/etcd_maintenance_test.go @@ -1,6 +1,10 @@ package v1alpha1 -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) func TestEffectiveCompaction(t *testing.T) { for _, tc := range []struct { @@ -23,9 +27,8 @@ func TestEffectiveCompaction(t *testing.T) { } { t.Run(tc.name, func(t *testing.T) { mode, retention, err := tc.config.EffectiveCompaction() - if (err != nil) != tc.wantErr || mode != tc.mode || retention != tc.retention { - t.Fatalf("got (%q,%q,%v)", mode, retention, err) - } + assert.NewAborting(t). + False((err != nil) != tc.wantErr || mode != tc.mode || retention != tc.retention, "got (%q,%q,%v)", mode, retention, err) }) } } diff --git a/api/v1alpha1/observability_helpers_test.go b/api/v1alpha1/observability_helpers_test.go index 4f598b2a..0522f963 100644 --- a/api/v1alpha1/observability_helpers_test.go +++ b/api/v1alpha1/observability_helpers_test.go @@ -20,6 +20,8 @@ import ( "testing" corev1 "k8s.io/api/core/v1" + + "github.com/multigres/testkit/assert" ) func TestBuildOTELEnvVars(t *testing.T) { @@ -115,15 +117,8 @@ func TestBuildOTELEnvVars(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := BuildOTELEnvVars(tc.cfg) - if len(got) != len(tc.want) { - t.Fatalf( - "len(BuildOTELEnvVars()) = %d, want %d\n got: %v\n want: %v", - len(got), - len(tc.want), - got, - tc.want, - ) - } + assert.NewAborting(t). + Len(got, len(tc.want), "len(BuildOTELEnvVars()) = %d, want %d\n got: %v\n want: %v", len(got), len(tc.want), got, tc.want) for i := range got { if got[i].Name != tc.want[i].Name || got[i].Value != tc.want[i].Value { t.Errorf("env[%d] = {%q, %q}, want {%q, %q}", @@ -165,15 +160,8 @@ func TestBuildOTELEnvVarsWithResourceAttributes(t *testing.T) { }, } - if len(got) != len(want) { - t.Fatalf( - "len(BuildOTELEnvVarsWithResourceAttributes()) = %d, want %d\n got: %v\n want: %v", - len(got), - len(want), - got, - want, - ) - } + assert.NewAborting(t). + Len(got, len(want), "len(BuildOTELEnvVarsWithResourceAttributes()) = %d, want %d\n got: %v\n want: %v", len(got), len(want), got, want) for i := range got { if got[i].Name != want[i].Name || got[i].Value != want[i].Value { t.Errorf("env[%d] = {%q, %q}, want {%q, %q}", @@ -186,40 +174,27 @@ func TestBuildOTELEnvVars_FallbackToEnv(t *testing.T) { t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://from-env:4318") t.Setenv("OTEL_EXPORTER_OTLP_METRICS_ENDPOINT", "http://metrics-from-env:4318/v1/metrics") t.Setenv("OTEL_EXPORTER_OTLP_PROTOCOL", "http/protobuf") + c := assert.NewCollecting(t) // nil config should fall back to env vars. got := BuildOTELEnvVars(nil) - if len(got) < 3 { - t.Fatalf("expected at least 3 env vars from fallback, got %d: %v", len(got), got) - } - if got[0].Value != "http://from-env:4318" { - t.Errorf("endpoint = %q, want %q", got[0].Value, "http://from-env:4318") - } - if got[1].Value != "http://metrics-from-env:4318/v1/metrics" { - t.Errorf( - "metrics endpoint = %q, want %q", - got[1].Value, - "http://metrics-from-env:4318/v1/metrics", - ) - } - if got[2].Value != "http/protobuf" { - t.Errorf("protocol = %q, want %q", got[2].Value, "http/protobuf") - } + c.Require(). + GreaterOrEqual(3, len(got), "expected at least 3 env vars from fallback, got %d: %v", len(got), got) + c.Eq("http://from-env:4318", got[0].Value, "endpoint") + c.Eq("http://metrics-from-env:4318/v1/metrics", got[1].Value, "metrics endpoint") + c.Eq("http/protobuf", got[2].Value, "protocol") } func TestBuildOTELEnvVars_CRDOverridesEnv(t *testing.T) { t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://from-env:4318") + c := assert.NewCollecting(t) cfg := &ObservabilityConfig{ OTLPEndpoint: "http://from-crd:4318", } got := BuildOTELEnvVars(cfg) - if len(got) == 0 { - t.Fatal("expected at least 1 env var") - } - if got[0].Value != "http://from-crd:4318" { - t.Errorf("endpoint = %q, want CRD value %q", got[0].Value, "http://from-crd:4318") - } + c.Require().NotEmpty(got, "expected at least 1 env var") + c.Eq("http://from-crd:4318", got[0].Value, "endpoint") } func TestEnvOrCRD(t *testing.T) { @@ -258,9 +233,7 @@ func TestEnvOrCRD(t *testing.T) { func(c *ObservabilityConfig) string { return c.OTLPEndpoint }, "TEST_ENVORCRD", ) - if got != tc.want { - t.Errorf("envOrCRD() = %q, want %q", got, tc.want) - } + assert.NewCollecting(t).Eq(tc.want, got, "envOrCRD()") }) } } @@ -414,70 +387,37 @@ func TestMergeBackupConfig(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) got := MergeBackupConfig(tc.child, tc.parent) if tc.want == nil { - if got != nil { - t.Errorf("MergeBackupConfig() = %+v, want nil", got) - } + c.Nil(got, "MergeBackupConfig()") return } - if got == nil { - t.Fatalf("MergeBackupConfig() = nil, want %+v", tc.want) - } - if got.Type != tc.want.Type { - t.Errorf("Type = %q, want %q", got.Type, tc.want.Type) - } + c.Require().NotNil(got, "MergeBackupConfig() = nil, want %+v", tc.want) + c.Eq(tc.want.Type, got.Type, "Type") if tc.want.Filesystem != nil { - if got.Filesystem == nil { - t.Fatal("Filesystem = nil, want non-nil") - } - if got.Filesystem.Path != tc.want.Filesystem.Path { - t.Errorf( - "Filesystem.Path = %q, want %q", - got.Filesystem.Path, - tc.want.Filesystem.Path, - ) - } - if got.Filesystem.Storage.Size != tc.want.Filesystem.Storage.Size { - t.Errorf( - "Filesystem.Storage.Size = %q, want %q", - got.Filesystem.Storage.Size, - tc.want.Filesystem.Storage.Size, - ) - } + c.Require().NotNil(got.Filesystem, "Filesystem = nil, want non-nil") + c.Eq(tc.want.Filesystem.Path, got.Filesystem.Path, "Filesystem.Path") + c.Eq( + tc.want.Filesystem.Storage.Size, + got.Filesystem.Storage.Size, + "Filesystem.Storage.Size", + ) } if tc.want.S3 != nil { - if got.S3 == nil { - t.Fatal("S3 = nil, want non-nil") - } - if got.S3.Bucket != tc.want.S3.Bucket { - t.Errorf("S3.Bucket = %q, want %q", got.S3.Bucket, tc.want.S3.Bucket) - } - if got.S3.Region != tc.want.S3.Region { - t.Errorf("S3.Region = %q, want %q", got.S3.Region, tc.want.S3.Region) - } - if got.S3.Endpoint != tc.want.S3.Endpoint { - t.Errorf("S3.Endpoint = %q, want %q", got.S3.Endpoint, tc.want.S3.Endpoint) - } - if got.S3.CredentialsSecret != tc.want.S3.CredentialsSecret { - t.Errorf( - "S3.CredentialsSecret = %q, want %q", - got.S3.CredentialsSecret, - tc.want.S3.CredentialsSecret, - ) - } + c.Require().NotNil(got.S3, "S3 = nil, want non-nil") + c.Eq(tc.want.S3.Bucket, got.S3.Bucket, "S3.Bucket") + c.Eq(tc.want.S3.Region, got.S3.Region, "S3.Region") + c.Eq(tc.want.S3.Endpoint, got.S3.Endpoint, "S3.Endpoint") + c.Eq(tc.want.S3.CredentialsSecret, got.S3.CredentialsSecret, "S3.CredentialsSecret") } if tc.want.PgBackRestTLS != nil { - if got.PgBackRestTLS == nil { - t.Fatal("PgBackRestTLS = nil, want non-nil") - } - if got.PgBackRestTLS.SecretName != tc.want.PgBackRestTLS.SecretName { - t.Errorf( - "PgBackRestTLS.SecretName = %q, want %q", - got.PgBackRestTLS.SecretName, - tc.want.PgBackRestTLS.SecretName, - ) - } + c.Require().NotNil(got.PgBackRestTLS, "PgBackRestTLS = nil, want non-nil") + c.Eq( + tc.want.PgBackRestTLS.SecretName, + got.PgBackRestTLS.SecretName, + "PgBackRestTLS.SecretName", + ) } }) } diff --git a/api/v1alpha1/topo_client_tls_helpers_test.go b/api/v1alpha1/topo_client_tls_helpers_test.go index b2025666..03ffb377 100644 --- a/api/v1alpha1/topo_client_tls_helpers_test.go +++ b/api/v1alpha1/topo_client_tls_helpers_test.go @@ -16,7 +16,11 @@ limitations under the License. package v1alpha1 -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) func TestTopoClientTLSConfigured(t *testing.T) { tests := map[string]struct { @@ -33,9 +37,8 @@ func TestTopoClientTLSConfigured(t *testing.T) { } for name, tc := range tests { t.Run(name, func(t *testing.T) { - if got := TopoClientTLSConfigured(tc.ref); got != tc.want { - t.Errorf("TopoClientTLSConfigured() = %v, want %v", got, tc.want) - } + assert.NewCollecting(t). + Eq(tc.want, TopoClientTLSConfigured(tc.ref), "TopoClientTLSConfigured()") }) } } @@ -44,57 +47,44 @@ func TestTopoClientTLSConfigured(t *testing.T) { // external topology may split them. The projection has to read the keypair from // ClientCertSecret and the CA from CASecret in either case. func TestBuildTopoClientTLSVolume_ProjectsFromBothSecrets(t *testing.T) { + c := assert.NewCollecting(t) ref := GlobalTopoServerRef{CASecret: "the-ca", ClientCertSecret: "the-client"} vol := BuildTopoClientTLSVolume(ref) - if vol.Name != TopoClientTLSVolumeName { - t.Errorf("volume name = %q, want %q", vol.Name, TopoClientTLSVolumeName) - } - if vol.Projected == nil { - t.Fatal("expected a projected volume source") - } + c.Eq(TopoClientTLSVolumeName, vol.Name, "volume name") + c.Require().NotNil(vol.Projected, "expected a projected volume source") sources := vol.Projected.Sources - if len(sources) != 2 { - t.Fatalf("got %d projection sources, want 2", len(sources)) - } + c.Require().Len(sources, 2, "got %d projection sources, want 2", len(sources)) keypair := sources[0].Secret - if keypair == nil || keypair.Name != "the-client" { - t.Fatalf("keypair source = %+v, want secret the-client", keypair) - } + c.Require(). + False(keypair == nil || keypair.Name != "the-client", "keypair source = %+v, want secret the-client", keypair) wantKeypairKeys := map[string]string{"tls.crt": "tls.crt", "tls.key": "tls.key"} gotKeypairKeys := map[string]string{} for _, item := range keypair.Items { gotKeypairKeys[item.Key] = item.Path } for k, v := range wantKeypairKeys { - if gotKeypairKeys[k] != v { - t.Errorf("keypair projects %q to %q, want %q", k, gotKeypairKeys[k], v) - } + c.Eq(v, gotKeypairKeys[k], "keypair projects %q to %q, want", k, gotKeypairKeys[k]) } ca := sources[1].Secret - if ca == nil || ca.Name != "the-ca" { - t.Fatalf("ca source = %+v, want secret the-ca", ca) - } + c.Require().False(ca == nil || ca.Name != "the-ca", "ca source = %+v, want secret the-ca", ca) if len(ca.Items) != 1 || ca.Items[0].Key != "ca.crt" || ca.Items[0].Path != "ca.crt" { t.Errorf("ca projection = %+v, want ca.crt to ca.crt", ca.Items) } } func TestTopoClientTLSArgs(t *testing.T) { + c := assert.NewCollecting(t) args := TopoClientTLSArgs() want := []string{ "--topo-etcd-tls-cert", TopoClientTLSCertFile, "--topo-etcd-tls-key", TopoClientTLSKeyFile, "--topo-etcd-tls-ca", TopoClientTLSCAFile, } - if len(args) != len(want) { - t.Fatalf("args = %v, want %v", args, want) - } + c.Require().Len(args, len(want), "args = %v, want %v", args, want) for i := range want { - if args[i] != want[i] { - t.Errorf("args[%d] = %q, want %q", i, args[i], want[i]) - } + c.Eq(want[i], args[i], "args[%d] = %q, want", i, args[i]) } } diff --git a/cmd/multigres-operator/main_test.go b/cmd/multigres-operator/main_test.go index ea56ffb3..d225e68e 100644 --- a/cmd/multigres-operator/main_test.go +++ b/cmd/multigres-operator/main_test.go @@ -1,6 +1,10 @@ package main -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) func TestGoMemLimitFromEnv(t *testing.T) { tests := map[string]struct { @@ -33,13 +37,10 @@ func TestGoMemLimitFromEnv(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewAborting(t) got, ok := goMemLimitFromEnv(tc.raw) - if ok != tc.wantOK { - t.Fatalf("goMemLimitFromEnv(%q) ok = %v, want %v", tc.raw, ok, tc.wantOK) - } - if got != tc.want { - t.Fatalf("goMemLimitFromEnv(%q) = %d, want %d", tc.raw, got, tc.want) - } + c.Eq(tc.wantOK, ok, "goMemLimitFromEnv(%q) ok = %v, want", tc.raw, ok) + c.Eq(tc.want, got, "goMemLimitFromEnv(%q) = %d, want", tc.raw, got) }) } } @@ -48,13 +49,15 @@ func TestGoMemLimitFromEnv(t *testing.T) { // container declares no memory limit, so the conversion has to stay in range // for values far larger than any limit we would set deliberately. func TestGoMemLimitFromEnvDoesNotOverflowAtNodeScale(t *testing.T) { + c := assert.NewAborting(t) const nodeAllocatable = 512 * 1024 * 1024 * 1024 // 512GiB got, ok := goMemLimitFromEnv("549755813888") - if !ok { - t.Fatal("goMemLimitFromEnv() rejected a node-sized limit") - } - if got <= 0 || got >= nodeAllocatable { - t.Fatalf("goMemLimitFromEnv() = %d, want a positive value below %d", got, nodeAllocatable) - } + c.True(ok, "goMemLimitFromEnv() rejected a node-sized limit") + c.False( + got <= 0 || got >= nodeAllocatable, + "goMemLimitFromEnv() = %d, want a positive value below %d", + got, + nodeAllocatable, + ) } diff --git a/config/manager/manager_test.go b/config/manager/manager_test.go index 19e23d9a..d2b7b2ba 100644 --- a/config/manager/manager_test.go +++ b/config/manager/manager_test.go @@ -7,17 +7,16 @@ import ( "testing" "k8s.io/apimachinery/pkg/util/yaml" + + "github.com/multigres/testkit/assert" ) func TestControllerManagerUsesNonOverlappingRollout(t *testing.T) { + c := assert.NewCollecting(t) f, err := os.Open("manager.yaml") - if err != nil { - t.Fatal(err) - } + c.Require().NoError(err) defer func() { - if err := f.Close(); err != nil { - t.Errorf("close manager manifest: %v", err) - } + c.NoError(f.Close(), "close manager manifest") }() type manifest struct { @@ -39,16 +38,11 @@ func TestControllerManagerUsesNonOverlappingRollout(t *testing.T) { if resource.Kind != "Deployment" { continue } - if got := resource.Spec.Strategy["type"]; got != "Recreate" { - t.Fatalf("controller-manager strategy = %q, want Recreate", got) - } + got := resource.Spec.Strategy["type"] + c.Require().False(got != "Recreate", "controller-manager strategy = %q, want Recreate", got) rollingUpdate, present := resource.Spec.Strategy["rollingUpdate"] - if !present { - t.Fatal("controller-manager strategy must explicitly clear rollingUpdate") - } - if rollingUpdate != nil { - t.Fatalf("controller-manager rollingUpdate = %#v, want null", rollingUpdate) - } + c.Require().True(present, "controller-manager strategy must explicitly clear rollingUpdate") + c.Require().Nil(rollingUpdate, "controller-manager rollingUpdate") return } diff --git a/go.mod b/go.mod index 57af9939..0515c135 100644 --- a/go.mod +++ b/go.mod @@ -6,10 +6,9 @@ require ( github.com/go-logr/logr v1.4.4 github.com/google/go-cmp v0.7.0 github.com/multigres/multigres v0.0.0-20260925193740-522b90425a83 - github.com/multigres/testkit v0.1.0 + github.com/multigres/testkit v0.2.0 github.com/prometheus/client_golang v1.24.1 github.com/prometheus/client_model v0.6.3 - github.com/stretchr/testify v1.12.1 go.etcd.io/etcd/api/v3 v3.7.1 go.etcd.io/etcd/client/v3 v3.7.1 go.opentelemetry.io/contrib/exporters/autoexport v0.71.0 @@ -101,6 +100,7 @@ require ( github.com/spf13/cobra v1.10.2 // indirect github.com/spf13/pflag v1.0.10 // indirect github.com/spf13/viper v1.21.0 // indirect + github.com/stretchr/testify v1.12.1 // indirect github.com/subosito/gotenv v1.6.0 // indirect github.com/tklauser/go-sysconf v0.3.16 // indirect github.com/tklauser/numcpus v0.11.0 // indirect diff --git a/go.sum b/go.sum index bdc302ba..85359d11 100644 --- a/go.sum +++ b/go.sum @@ -199,8 +199,8 @@ github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee h1:W5t00kpgFd github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= github.com/multigres/multigres v0.0.0-20260925193740-522b90425a83 h1:IdyFGtwc9pEZEzqs6h5JfavAnErWfuKwj3d42qkZUx8= github.com/multigres/multigres v0.0.0-20260925193740-522b90425a83/go.mod h1:Ov2hrkOguWSkCS2QIhAdguFeG5GlZ3v4WGqIdqkQ7Tg= -github.com/multigres/testkit v0.1.0 h1:i6DiCFZ9mVhEopBwakMmel5NTde1VErEPGtnoXn9CdU= -github.com/multigres/testkit v0.1.0/go.mod h1:3ONhsV/PNOUke7PID5HPlnxTLyQCcfOQ/JfLCRkSLOY= +github.com/multigres/testkit v0.2.0 h1:fy2GGA5kv8Nl8t9I04Q8u5Rct9hYUI431vhycUEKMFw= +github.com/multigres/testkit v0.2.0/go.mod h1:3ONhsV/PNOUke7PID5HPlnxTLyQCcfOQ/JfLCRkSLOY= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/onsi/ginkgo/v2 v2.27.4 h1:fcEcQW/A++6aZAZQNUmNjvA9PSOzefMJBerHJ4t8v8Y= diff --git a/pkg/cert/generator_test.go b/pkg/cert/generator_test.go index 3afda7f6..60a3d62f 100644 --- a/pkg/cert/generator_test.go +++ b/pkg/cert/generator_test.go @@ -11,7 +11,7 @@ import ( "strings" "testing" - "github.com/google/go-cmp/cmp" + "github.com/multigres/testkit/assert" ) func TestGenerator_Logic(t *testing.T) { @@ -26,17 +26,13 @@ func TestGenerator_Logic(t *testing.T) { return nil } cert, err := x509.ParseCertificate(block.Bytes) - if err != nil { - tb.Fatalf("failed to parse certificate: %v", err) - } + assert.NewAborting(tb).NoError(err, "failed to parse certificate") return cert } // Fixtures caArtifacts, err := GenerateCA("") - if err != nil { - t.Fatalf("setup failed: GenerateCA error = %v", err) - } + assert.NewAborting(t).NoError(err, "setup failed: GenerateCA error =") type input struct { ca *CAArtifacts @@ -52,13 +48,11 @@ func TestGenerator_Logic(t *testing.T) { }{ "Happy Path: Generate CA": { validate: func(tb testing.TB, _ *ServerArtifacts) { + c := assert.NewCollecting(tb) cert := decodeCert(tb, caArtifacts.CertPEM) - if !cert.IsCA { - tb.Error("Expected CA cert to have IsCA=true") - } - if got, want := cert.Subject.CommonName, "Multigres Operator CA"; got != want { - tb.Errorf("CommonName mismatch: got %q, want %q", got, want) - } + c.True(cert.IsCA, "Expected CA cert to have IsCA=true") + got, want := cert.Subject.CommonName, "Multigres Operator CA" + c.Eq(want, got, "CommonName mismatch: got") }, }, "Happy Path: Generate Server Cert": { @@ -68,23 +62,21 @@ func TestGenerator_Logic(t *testing.T) { dnsNames: []string{"test-svc", "test-svc.ns.svc"}, }, validate: func(tb testing.TB, arts *ServerArtifacts) { + c := assert.NewCollecting(tb) cert := decodeCert(tb, arts.CertPEM) - if cert.IsCA { - tb.Error("Expected server cert to NOT be CA") - } - if got, want := cert.Subject.CommonName, "test-svc.ns.svc"; got != want { - tb.Errorf("CN mismatch: got %q, want %q", got, want) - } - if diff := cmp.Diff( - cert.DNSNames, + c.False(cert.IsCA, "Expected server cert to NOT be CA") + got, want := cert.Subject.CommonName, "test-svc.ns.svc" + c.Eq(want, got, "CN mismatch: got") + c.EqDiff( []string{"test-svc", "test-svc.ns.svc"}, - ); diff != "" { - tb.Errorf("DNSNames mismatch (-got +want):\n%s", diff) - } + cert.DNSNames, + "DNSNames mismatch", + ) // Verify chain - if err := cert.CheckSignatureFrom(caArtifacts.Cert); err != nil { - tb.Errorf("Signature verification failed: %v", err) - } + c.NoError( + cert.CheckSignatureFrom(caArtifacts.Cert), + "Signature verification failed", + ) // Verify default ExtKeyUsage is ServerAuth only if len(cert.ExtKeyUsage) != 1 || cert.ExtKeyUsage[0] != x509.ExtKeyUsageServerAuth { tb.Errorf("Expected default ExtKeyUsage [ServerAuth], got %v", cert.ExtKeyUsage) @@ -115,22 +107,20 @@ func TestGenerator_Logic(t *testing.T) { }, }, validate: func(tb testing.TB, arts *ServerArtifacts) { + c := assert.NewCollecting(tb) cert := decodeCert(tb, arts.CertPEM) - if len(cert.ExtKeyUsage) != 2 { - tb.Fatalf("Expected 2 ExtKeyUsages, got %d", len(cert.ExtKeyUsage)) - } - if cert.ExtKeyUsage[0] != x509.ExtKeyUsageServerAuth { - tb.Errorf( - "Expected first ExtKeyUsage to be ServerAuth, got %v", - cert.ExtKeyUsage[0], - ) - } - if cert.ExtKeyUsage[1] != x509.ExtKeyUsageClientAuth { - tb.Errorf( - "Expected second ExtKeyUsage to be ClientAuth, got %v", - cert.ExtKeyUsage[1], - ) - } + c.Require(). + Len(cert.ExtKeyUsage, 2, "Expected 2 ExtKeyUsages, got %d", len(cert.ExtKeyUsage)) + c.Eq( + x509.ExtKeyUsageServerAuth, + cert.ExtKeyUsage[0], + "Expected first ExtKeyUsage to be ServerAuth, got", + ) + c.Eq( + x509.ExtKeyUsageClientAuth, + cert.ExtKeyUsage[1], + "Expected second ExtKeyUsage to be ClientAuth, got", + ) }, }, "Happy Path: Single ExtKeyUsage Override": { @@ -160,6 +150,7 @@ func TestGenerator_Logic(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) // Skip execution for CA-only test case if name == "Happy Path: Generate CA" { @@ -173,14 +164,10 @@ func TestGenerator_Logic(t *testing.T) { tc.input.dnsNames, tc.input.opts...) if tc.wantErr { - if err == nil { - t.Error("Expected error, got nil") - } + c.Error(err, "Expected error, got nil") return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.Require().NoError(err, "Unexpected error") if tc.validate != nil { tc.validate(t, arts) @@ -249,22 +236,15 @@ func TestParseCA_Logic(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) got, err := ParseCA(tc.certBytes, tc.keyBytes) if tc.wantErr != "" { - if err == nil { - t.Fatal("Expected error, got nil") - } - if !strings.Contains(err.Error(), tc.wantErr) { - t.Errorf("Error mismatch. Got %q, want substring %q", err.Error(), tc.wantErr) - } + c.Require().Error(err, "Expected error, got nil") + c.StrContains(err.Error(), tc.wantErr, "Error mismatch. Got") } else { - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } - if got == nil { - t.Fatal("Expected artifacts, got nil") - } + c.Require().NoError(err, "Unexpected error") + c.Require().NotNil(got, "Expected artifacts, got nil") } }) } @@ -287,17 +267,18 @@ func TestGenerator_EntropyFailures(t *testing.T) { defer func() { randReader = oldReader }() t.Run("GenerateServerCert: serial number failure", func(t *testing.T) { + c := assert.NewCollecting(t) ca, err := GenerateCA("") - if err != nil { - t.Fatalf("GenerateCA failed: %v", err) - } + c.Require().NoError(err, "GenerateCA failed") randReader = errorReader{} defer func() { randReader = oldReader }() _, err = GenerateServerCert(ca, "foo", nil) - if err == nil || !strings.Contains(err.Error(), "failed to generate serial number") { - t.Errorf("Expected serial number error, got %v", err) - } + c.False( + err == nil || !strings.Contains(err.Error(), "failed to generate serial number"), + "Expected serial number error, got %v", + err, + ) }) } @@ -314,9 +295,8 @@ func TestGenerator_MockFailures(t *testing.T) { } _, err := GenerateCA("") - if err == nil || !strings.Contains(err.Error(), "failed to parse generated CA") { - t.Errorf("Expected parse error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to parse generated CA"), "Expected parse error, got %v", err) }) t.Run("GenerateCA: Marshal Key Failure", func(t *testing.T) { @@ -326,9 +306,8 @@ func TestGenerator_MockFailures(t *testing.T) { } _, err := GenerateCA("") - if err == nil || !strings.Contains(err.Error(), "failed to marshal CA key") { - t.Errorf("Expected marshal error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to marshal CA key"), "Expected marshal error, got %v", err) }) } @@ -347,8 +326,7 @@ func TestGenerator_MockFailures_ServerCert(t *testing.T) { } _, err := GenerateServerCert(ca, "foo", nil) - if err == nil || !strings.Contains(err.Error(), "failed to marshal server key") { - t.Errorf("Expected marshal error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to marshal server key"), "Expected marshal error, got %v", err) }) } diff --git a/pkg/cert/manager_test.go b/pkg/cert/manager_test.go index ae566218..660169d3 100644 --- a/pkg/cert/manager_test.go +++ b/pkg/cert/manager_test.go @@ -30,6 +30,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) const ( @@ -525,6 +527,7 @@ func TestManager_EnsureCerts(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) fakeClient := fake.NewClientBuilder(). WithScheme(s). @@ -578,21 +581,16 @@ func TestManager_EnsureCerts(t *testing.T) { err := mgr.Bootstrap(t.Context()) if tc.wantErr { - if err == nil { - t.Fatal("Expected error, got nil") - } - if tc.errContains != "" && !strings.Contains(err.Error(), tc.errContains) { - t.Errorf( - "Error message mismatch. Got: %v, Want substring: %s", - err, - tc.errContains, - ) - } + c.Require().Error(err, "Expected error, got nil") + c.False( + tc.errContains != "" && !strings.Contains(err.Error(), tc.errContains), + "Error message mismatch. Got: %v, Want substring: %s", + err, + tc.errContains, + ) return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.Require().NoError(err, "Unexpected error") if tc.checkFiles { if _, err := os.Stat(filepath.Join(certDir, CertFileName)); os.IsNotExist(err) { @@ -615,9 +613,10 @@ func TestManager_EnsureCerts(t *testing.T) { break } } - if len(original) > 0 && bytes.Equal(secret.Data["tls.crt"], original) { - t.Error("Expected rotation, but cert did not change") - } + c.False( + len(original) > 0 && bytes.Equal(secret.Data["tls.crt"], original), + "Expected rotation, but cert did not change", + ) } }) } @@ -686,9 +685,7 @@ func (b *badSchemeClient) Scheme() *runtime.Scheme { func generateCAPEM(tb testing.TB) ([]byte, []byte) { tb.Helper() ca, err := GenerateCA("") - if err != nil { - tb.Fatal(err) - } + assert.NewAborting(tb).NoError(err) return ca.CertPEM, ca.KeyPEM } @@ -699,10 +696,9 @@ func generateSignedCertPEM( dnsNames []string, ) []byte { tb.Helper() + c := assert.NewAborting(tb) priv, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - tb.Fatal(err) - } + c.NoError(err) tmpl := x509.Certificate{ SerialNumber: big.NewInt(2), @@ -713,9 +709,7 @@ func generateSignedCertPEM( } der, err := x509.CreateCertificate(rand.Reader, &tmpl, ca.Cert, &priv.PublicKey, ca.Key) - if err != nil { - tb.Fatal(err) - } + c.NoError(err) return pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) } @@ -728,6 +722,7 @@ func TestManager_PostReconcileHook(t *testing.T) { t.Run("Hook Called with CA Bundle", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) var hookCalled atomic.Bool var receivedCABundle []byte @@ -746,22 +741,14 @@ func TestManager_PostReconcileHook(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.Require().NoError(mgr.Bootstrap(t.Context()), "Unexpected error") - if !hookCalled.Load() { - t.Error("PostReconcileHook was not called") - } - if len(receivedCABundle) == 0 { - t.Error("PostReconcileHook received empty CA bundle") - } + c.True(hookCalled.Load(), "PostReconcileHook was not called") + c.NotEmpty(receivedCABundle, "PostReconcileHook received empty CA bundle") // Verify the CA bundle is valid PEM block, _ := pem.Decode(receivedCABundle) - if block == nil { - t.Error("CA bundle is not valid PEM") - } + c.NotNil(block, "CA bundle is not valid PEM") }) t.Run("Hook Error Propagates", func(t *testing.T) { @@ -780,9 +767,8 @@ func TestManager_PostReconcileHook(t *testing.T) { mgr := NewManager(cl, record.NewFakeRecorder(10), opts) err := mgr.Bootstrap(t.Context()) - if err == nil || !strings.Contains(err.Error(), "post-reconcile hook failed") { - t.Errorf("Expected hook error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "post-reconcile hook failed"), "Expected hook error, got %v", err) }) t.Run("No Hook (nil) Succeeds", func(t *testing.T) { @@ -798,9 +784,7 @@ func TestManager_PostReconcileHook(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + assert.NewAborting(t).NoError(mgr.Bootstrap(t.Context()), "Unexpected error") }) } @@ -812,6 +796,7 @@ func TestManager_OwnerRef(t *testing.T) { t.Run("Owner Set: Secrets Get Owner Reference", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) owner := &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ @@ -831,37 +816,28 @@ func TestManager_OwnerRef(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.Require().NoError(mgr.Bootstrap(t.Context()), "Unexpected error") // Verify CA secret has owner reference caSecret := &corev1.Secret{} - if err := cl.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(cl.Get(t.Context(), types.NamespacedName{ Name: testCASecretName, Namespace: "test-ns", - }, caSecret); err != nil { - t.Fatalf("Failed to get CA secret: %v", err) - } - if len(caSecret.OwnerReferences) == 0 { - t.Error("Expected CA secret to have owner reference") - } + }, caSecret), "Failed to get CA secret") + c.NotEmpty(caSecret.OwnerReferences, "Expected CA secret to have owner reference") // Verify server secret has owner reference srvSecret := &corev1.Secret{} - if err := cl.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(cl.Get(t.Context(), types.NamespacedName{ Name: testServerSecretName, Namespace: "test-ns", - }, srvSecret); err != nil { - t.Fatalf("Failed to get server secret: %v", err) - } - if len(srvSecret.OwnerReferences) == 0 { - t.Error("Expected server secret to have owner reference") - } + }, srvSecret), "Failed to get server secret") + c.NotEmpty(srvSecret.OwnerReferences, "Expected server secret to have owner reference") }) t.Run("No Owner: Secrets Created Without Owner Reference", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cl := fake.NewClientBuilder().WithScheme(s).Build() opts := Options{ @@ -873,20 +849,14 @@ func TestManager_OwnerRef(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.Require().NoError(mgr.Bootstrap(t.Context()), "Unexpected error") caSecret := &corev1.Secret{} - if err := cl.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(cl.Get(t.Context(), types.NamespacedName{ Name: testCASecretName, Namespace: "test-ns", - }, caSecret); err != nil { - t.Fatalf("Failed to get CA secret: %v", err) - } - if len(caSecret.OwnerReferences) != 0 { - t.Error("Expected CA secret to have no owner references") - } + }, caSecret), "Failed to get CA secret") + c.Empty(caSecret.OwnerReferences, "Expected CA secret to have no owner references") }) t.Run("Owner with Bad Scheme: SetControllerReference Fails", func(t *testing.T) { @@ -915,9 +885,8 @@ func TestManager_OwnerRef(t *testing.T) { mgr := NewManager(cl, record.NewFakeRecorder(10), opts) err := mgr.Bootstrap(t.Context()) - if err == nil || !strings.Contains(err.Error(), "failed to set controller reference") { - t.Errorf("Expected controller ref error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to set controller reference"), "Expected controller ref error, got %v", err) }) t.Run("Owner with Bad Scheme: SetControllerReference Fails on Server Cert", func(t *testing.T) { @@ -951,10 +920,11 @@ func TestManager_OwnerRef(t *testing.T) { mgr := NewManager(cl, record.NewFakeRecorder(10), opts) err := mgr.Bootstrap(t.Context()) - if err == nil || - !strings.Contains(err.Error(), "failed to set owner for server cert secret") { - t.Errorf("Expected server cert owner error, got %v", err) - } + assert.NewCollecting(t).False(err == nil || + !strings.Contains( + err.Error(), + "failed to set owner for server cert secret", + ), "Expected server cert owner error, got %v", err) }) } @@ -978,9 +948,7 @@ func TestManager_WaitForProjection(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + assert.NewAborting(t).NoError(mgr.Bootstrap(t.Context()), "Unexpected error") }) t.Run("Enabled: Mismatch Timeout", func(t *testing.T) { @@ -1024,9 +992,7 @@ func TestManager_WaitForProjection(t *testing.T) { defer cancel() err := mgr.waitForProjection(ctx, []byte("expected")) - if err == nil { - t.Error("Expected timeout error for missing file") - } + assert.NewCollecting(t).Error(err, "Expected timeout error for missing file") }) t.Run("Enabled: File Matches Immediately", func(t *testing.T) { @@ -1045,9 +1011,8 @@ func TestManager_WaitForProjection(t *testing.T) { WaitForProjection: true, }) - if err := mgr.waitForProjection(t.Context(), expected); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(mgr.waitForProjection(t.Context(), expected), "Unexpected error") }) } @@ -1059,6 +1024,7 @@ func TestManager_ExtKeyUsages(t *testing.T) { t.Run("Custom ExtKeyUsages Flow Through to Cert", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cl := fake.NewClientBuilder().WithScheme(s).Build() opts := Options{ @@ -1073,41 +1039,29 @@ func TestManager_ExtKeyUsages(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.Require().NoError(mgr.Bootstrap(t.Context()), "Unexpected error") // Read the generated server cert and verify ExtKeyUsages secret := &corev1.Secret{} - if err := cl.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(cl.Get(t.Context(), types.NamespacedName{ Name: testServerSecretName, Namespace: "test-ns", - }, secret); err != nil { - t.Fatalf("Failed to get server secret: %v", err) - } + }, secret), "Failed to get server secret") block, _ := pem.Decode(secret.Data["tls.crt"]) - if block == nil { - t.Fatal("Failed to decode server cert PEM") - } + c.Require().NotNil(block, "Failed to decode server cert PEM") cert, err := x509.ParseCertificate(block.Bytes) - if err != nil { - t.Fatalf("Failed to parse server cert: %v", err) - } + c.Require().NoError(err, "Failed to parse server cert") - if len(cert.ExtKeyUsage) != 2 { - t.Fatalf("Expected 2 ExtKeyUsages, got %d", len(cert.ExtKeyUsage)) - } - if cert.ExtKeyUsage[0] != x509.ExtKeyUsageServerAuth { - t.Errorf("Expected ServerAuth, got %v", cert.ExtKeyUsage[0]) - } - if cert.ExtKeyUsage[1] != x509.ExtKeyUsageClientAuth { - t.Errorf("Expected ClientAuth, got %v", cert.ExtKeyUsage[1]) - } + c.Require(). + Len(cert.ExtKeyUsage, 2, "Expected 2 ExtKeyUsages, got %d", len(cert.ExtKeyUsage)) + c.Eq(x509.ExtKeyUsageServerAuth, cert.ExtKeyUsage[0], "Expected ServerAuth, got") + c.Eq(x509.ExtKeyUsageClientAuth, cert.ExtKeyUsage[1], "Expected ClientAuth, got") }) t.Run("Default ExtKeyUsages (ServerAuth Only)", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) cl := fake.NewClientBuilder().WithScheme(s).Build() opts := Options{ @@ -1119,17 +1073,13 @@ func TestManager_ExtKeyUsages(t *testing.T) { } mgr := NewManager(cl, record.NewFakeRecorder(10), opts) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Unexpected error: %v", err) - } + c.NoError(mgr.Bootstrap(t.Context()), "Unexpected error") secret := &corev1.Secret{} - if err := cl.Get(t.Context(), types.NamespacedName{ + c.NoError(cl.Get(t.Context(), types.NamespacedName{ Name: testServerSecretName, Namespace: "test-ns", - }, secret); err != nil { - t.Fatalf("Failed to get server secret: %v", err) - } + }, secret), "Failed to get server secret") block, _ := pem.Decode(secret.Data["tls.crt"]) cert, _ := x509.ParseCertificate(block.Bytes) @@ -1158,9 +1108,9 @@ func TestDNSNameSetsEqual(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() - if got := dnsNameSetsEqual(tc.a, tc.b); got != tc.want { - t.Errorf("dnsNameSetsEqual(%v, %v) = %v, want %v", tc.a, tc.b, got, tc.want) - } + got := dnsNameSetsEqual(tc.a, tc.b) + assert.NewCollecting(t). + Eq(tc.want, got, "dnsNameSetsEqual(%v, %v) = %v, want", tc.a, tc.b, got) }) } } @@ -1230,50 +1180,43 @@ func TestManager_Misc(t *testing.T) { }) // Start should not return an error — it logs reconcile failures - if err := mgr.Start(timeoutCtx); err != nil { - t.Fatalf("Start should only return nil, got %v", err) - } + assert.NewAborting(t).NoError(mgr.Start(timeoutCtx), "Start should only return nil, got") }) t.Run("ComponentName Default", func(t *testing.T) { t.Parallel() opts := Options{} - if got := opts.componentName(); got != "cert" { - t.Errorf("Expected default componentName 'cert', got %q", got) - } + assert.NewCollecting(t). + Eq("cert", opts.componentName(), "Expected default componentName 'cert', got") }) t.Run("ComponentName Custom", func(t *testing.T) { t.Parallel() opts := Options{ComponentName: "webhook"} - if got := opts.componentName(); got != "webhook" { - t.Errorf("Expected componentName 'webhook', got %q", got) - } + assert.NewCollecting(t). + Eq("webhook", opts.componentName(), "Expected componentName 'webhook', got") }) t.Run("ExtKeyUsages Default", func(t *testing.T) { t.Parallel() opts := Options{} usages := opts.extKeyUsages() - if len(usages) != 1 || usages[0] != x509.ExtKeyUsageServerAuth { - t.Errorf("Expected default [ServerAuth], got %v", usages) - } + assert.NewCollecting(t). + False(len(usages) != 1 || usages[0] != x509.ExtKeyUsageServerAuth, "Expected default [ServerAuth], got %v", usages) }) t.Run("Organization Default", func(t *testing.T) { t.Parallel() opts := Options{} - if got := opts.organization(); got != Organization { - t.Errorf("Expected default organization %q, got %q", Organization, got) - } + assert.NewCollecting(t). + Eq(Organization, opts.organization(), "Expected default organization") }) t.Run("Organization Custom", func(t *testing.T) { t.Parallel() opts := Options{Organization: "Acme Corp"} - if got := opts.organization(); got != "Acme Corp" { - t.Errorf("Expected organization 'Acme Corp', got %q", got) - } + assert.NewCollecting(t). + Eq("Acme Corp", opts.organization(), "Expected organization 'Acme Corp', got") }) t.Run("RecorderEvent with Nil Recorder", func(t *testing.T) { @@ -1323,9 +1266,8 @@ func TestManager_EntropyFailures(t *testing.T) { randReader = errorReader{} err := mgr.reconcilePKI(t.Context()) - if err == nil || !strings.Contains(err.Error(), "failed to generate server cert") { - t.Errorf("Expected server cert gen error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to generate server cert"), "Expected server cert gen error, got %v", err) }) t.Run("ensureServerCert: GenerateServerCert Failure (Rotation)", func(t *testing.T) { @@ -1363,9 +1305,8 @@ func TestManager_EntropyFailures(t *testing.T) { randReader = errorReader{} err := mgr.reconcilePKI(t.Context()) - if err == nil || !strings.Contains(err.Error(), "failed to generate new server cert") { - t.Errorf("Expected server cert rotation error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to generate new server cert"), "Expected server cert rotation error, got %v", err) }) } @@ -1405,6 +1346,7 @@ func TestManager_CacheRaceConditions(t *testing.T) { t.Run("ensureCA: MaxRecursionDepth", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cl := &alreadyExistsOnCreateClient{ Client: fake.NewClientBuilder().WithScheme(s).Build(), @@ -1420,15 +1362,19 @@ func TestManager_CacheRaceConditions(t *testing.T) { }) err := mgr.Bootstrap(t.Context()) - if err == nil { - t.Fatal("Expected error, got nil") - } - if !strings.Contains(err.Error(), "failed to ensure CA secret") { - t.Errorf("Expected max recursion error, got: %v", err) - } - if !strings.Contains(err.Error(), "informer cache") { - t.Errorf("Expected cache label hint in error, got: %v", err) - } + c.Require().Error(err, "Expected error, got nil") + c.StrContains( + err.Error(), + "failed to ensure CA secret", + "Expected max recursion error, got: %v", + err, + ) + c.StrContains( + err.Error(), + "informer cache", + "Expected cache label hint in error, got: %v", + err, + ) }) t.Run("ensureCA: AlreadyExists Retry Succeeds", func(t *testing.T) { @@ -1447,13 +1393,13 @@ func TestManager_CacheRaceConditions(t *testing.T) { ServiceName: "test-svc", }) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Expected success after retry, got: %v", err) - } + assert.NewAborting(t). + NoError(mgr.Bootstrap(t.Context()), "Expected success after retry, got") }) t.Run("ensureServerCert: MaxRecursionDepth", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) validCABytes, validCAKeyBytes := generateCAPEM(t) caSecret := &corev1.Secret{ @@ -1475,15 +1421,19 @@ func TestManager_CacheRaceConditions(t *testing.T) { }) err := mgr.Bootstrap(t.Context()) - if err == nil { - t.Fatal("Expected error, got nil") - } - if !strings.Contains(err.Error(), "failed to ensure server cert secret") { - t.Errorf("Expected max recursion error, got: %v", err) - } - if !strings.Contains(err.Error(), "informer cache") { - t.Errorf("Expected cache label hint in error, got: %v", err) - } + c.Require().Error(err, "Expected error, got nil") + c.StrContains( + err.Error(), + "failed to ensure server cert secret", + "Expected max recursion error, got: %v", + err, + ) + c.StrContains( + err.Error(), + "informer cache", + "Expected cache label hint in error, got: %v", + err, + ) }) t.Run("ensureServerCert: AlreadyExists Retry Succeeds", func(t *testing.T) { @@ -1508,8 +1458,7 @@ func TestManager_CacheRaceConditions(t *testing.T) { ServiceName: "test-svc", }) - if err := mgr.Bootstrap(t.Context()); err != nil { - t.Fatalf("Expected success after retry, got: %v", err) - } + assert.NewAborting(t). + NoError(mgr.Bootstrap(t.Context()), "Expected success after retry, got") }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/builders_cell_test.go b/pkg/cluster-handler/controller/multigrescluster/builders_cell_test.go index 612d2d45..08105d0b 100644 --- a/pkg/cluster-handler/controller/multigrescluster/builders_cell_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/builders_cell_test.go @@ -3,7 +3,6 @@ package multigrescluster import ( "testing" - "github.com/google/go-cmp/cmp" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" "k8s.io/utils/ptr" @@ -11,6 +10,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestBuildCell(t *testing.T) { @@ -48,6 +49,7 @@ func TestBuildCell(t *testing.T) { allCells := []multigresv1alpha1.CellName{"zone-a", "zone-b"} t.Run("Success", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildCell( cluster, cellCfg, @@ -58,32 +60,26 @@ func TestBuildCell(t *testing.T) { allCells, scheme, ) - if err != nil { - t.Fatalf("BuildCell() error = %v", err) - } + c.Require().NoError(err, "BuildCell() error =") // Calculate expected hash: md5("my-cluster", "zone-a") -> "6b6f7386" expectedName := name.JoinWithConstraints(name.DefaultConstraints, "my-cluster", "zone-a") - if got.Name != expectedName { - t.Errorf("Name = %v, want %v", got.Name, expectedName) - } - if got.Spec.ZoneID != "use1-az1" { - t.Errorf("ZoneID = %v, want %v", got.Spec.ZoneID, "use1-az1") - } - if got.Spec.Images.Multigateway != "gateway:latest" { - t.Errorf("Gateway Image = %v, want %v", got.Spec.Images.Multigateway, "gateway:latest") - } - if diff := cmp.Diff(allCells, got.Spec.AllCells); diff != "" { - t.Errorf("AllCells mismatch (-want +got):\n%s", diff) - } + c.Eq(expectedName, got.Name, "Name") + c.Eq("use1-az1", got.Spec.ZoneID, "ZoneID") + c.Eq("gateway:latest", got.Spec.Images.Multigateway, "Gateway Image") + c.EqDiff(allCells, got.Spec.AllCells, "AllCells mismatch") // Verify OwnerReference - if len(got.OwnerReferences) != 1 { - t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) - } + c.Len( + got.OwnerReferences, + 1, + "OwnerReferences count = %v, want 1", + len(got.OwnerReferences), + ) }) t.Run("Propagates ZoneID", func(t *testing.T) { + c := assert.NewCollecting(t) cellCfgWithZoneID := &multigresv1alpha1.CellConfig{ Name: "zone-a", ZoneID: "use1-az1", @@ -98,15 +94,12 @@ func TestBuildCell(t *testing.T) { allCells, scheme, ) - if err != nil { - t.Fatalf("BuildCell() error = %v", err) - } - if got.Spec.ZoneID != "use1-az1" { - t.Errorf("ZoneID = %v, want use1-az1", got.Spec.ZoneID) - } + c.Require().NoError(err, "BuildCell() error =") + c.Eq("use1-az1", got.Spec.ZoneID, "ZoneID") }) t.Run("Propagates InternalTLS", func(t *testing.T) { + c := assert.NewAborting(t) clusterWithInternalTLS := cluster.DeepCopy() clusterWithInternalTLS.Spec.InternalTLS = &multigresv1alpha1.InternalTLSConfig{ Enabled: ptr.To(true), @@ -122,16 +115,8 @@ func TestBuildCell(t *testing.T) { allCells, scheme, ) - if err != nil { - t.Fatalf("BuildCell() error = %v", err) - } - if got.Spec.InternalTLS != clusterWithInternalTLS.Spec.InternalTLS { - t.Fatalf( - "InternalTLS = %#v, want propagated pointer %#v", - got.Spec.InternalTLS, - clusterWithInternalTLS.Spec.InternalTLS, - ) - } + c.NoError(err, "BuildCell() error =") + c.Eq(clusterWithInternalTLS.Spec.InternalTLS, got.Spec.InternalTLS, "InternalTLS") }) t.Run("ControllerRefError", func(t *testing.T) { @@ -146,12 +131,11 @@ func TestBuildCell(t *testing.T) { allCells, emptyScheme, ) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) t.Run("Propagates explicit project ref annotation", func(t *testing.T) { + c := assert.NewAborting(t) clusterWithProjectRef := cluster.DeepCopy() clusterWithProjectRef.Annotations = map[string]string{ metadata.AnnotationProjectRef: "proj_123", @@ -167,17 +151,14 @@ func TestBuildCell(t *testing.T) { allCells, scheme, ) - if err != nil { - t.Fatalf("BuildCell() error = %v", err) - } - - if got.Annotations[metadata.AnnotationProjectRef] != "proj_123" { - t.Fatalf( - "annotation %q = %q, want %q", - metadata.AnnotationProjectRef, - got.Annotations[metadata.AnnotationProjectRef], - "proj_123", - ) - } + c.NoError(err, "BuildCell() error =") + + c.Eq( + "proj_123", + got.Annotations[metadata.AnnotationProjectRef], + "annotation %q = %q, want", + metadata.AnnotationProjectRef, + got.Annotations[metadata.AnnotationProjectRef], + ) }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/builders_global_test.go b/pkg/cluster-handler/controller/multigrescluster/builders_global_test.go index dd29b654..6bf633e5 100644 --- a/pkg/cluster-handler/controller/multigrescluster/builders_global_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/builders_global_test.go @@ -4,10 +4,6 @@ import ( "fmt" "testing" - "github.com/google/go-cmp/cmp" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - corev1 "k8s.io/api/core/v1" networkingv1 "k8s.io/api/networking/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -16,6 +12,8 @@ import ( "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildGlobalTopoServer(t *testing.T) { @@ -31,6 +29,7 @@ func TestBuildGlobalTopoServer(t *testing.T) { } t.Run("Etcd Enabled", func(t *testing.T) { + c := assert.NewCollecting(t) spec := &multigresv1alpha1.GlobalTopoServerSpec{ Etcd: &multigresv1alpha1.EtcdSpec{ Image: "etcd:latest", @@ -39,19 +38,11 @@ func TestBuildGlobalTopoServer(t *testing.T) { } got, err := BuildGlobalTopoServer(cluster, spec, scheme) - if err != nil { - t.Fatalf("BuildGlobalTopoServer() error = %v", err) - } + c.Require().NoError(err, "BuildGlobalTopoServer() error =") - if got == nil { - t.Fatal("Expected TopoServer, got nil") - } - if got.Name != "my-cluster-global-topo" { - t.Errorf("Name = %v, want %v", got.Name, "my-cluster-global-topo") - } - if got.Spec.Etcd.Image != "etcd:latest" { - t.Errorf("Image = %v, want %v", got.Spec.Etcd.Image, "etcd:latest") - } + c.Require().NotNil(got, "Expected TopoServer, got nil") + c.Eq("my-cluster-global-topo", got.Name, "Name") + c.Eq("etcd:latest", got.Spec.Etcd.Image, "Image") // Verify OwnerReference if len(got.OwnerReferences) != 1 { t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) @@ -61,6 +52,7 @@ func TestBuildGlobalTopoServer(t *testing.T) { }) t.Run("Etcd Enabled with placement", func(t *testing.T) { + c := assert.NewCollecting(t) spec := &multigresv1alpha1.GlobalTopoServerSpec{ Etcd: &multigresv1alpha1.EtcdSpec{Image: "etcd:latest"}, Placement: &multigresv1alpha1.TopoServerPlacementSpec{ @@ -76,26 +68,19 @@ func TestBuildGlobalTopoServer(t *testing.T) { } got, err := BuildGlobalTopoServer(cluster, spec, scheme) - if err != nil { - t.Fatalf("BuildGlobalTopoServer() error = %v", err) - } - if diff := cmp.Diff(spec.Placement, got.Spec.Placement); diff != "" { - t.Errorf("Placement diff (-want +got):\n%s", diff) - } + c.Require().NoError(err, "BuildGlobalTopoServer() error =") + c.EqDiff(spec.Placement, got.Spec.Placement, "Placement diff") }) t.Run("Etcd Disabled (External)", func(t *testing.T) { + c := assert.NewCollecting(t) spec := &multigresv1alpha1.GlobalTopoServerSpec{ Etcd: nil, // Simulating external mode where Etcd spec is nil } got, err := BuildGlobalTopoServer(cluster, spec, scheme) - if err != nil { - t.Fatalf("BuildGlobalTopoServer() error = %v", err) - } - if got != nil { - t.Errorf("Expected nil when Etcd spec is nil, got %v", got) - } + c.Require().NoError(err, "BuildGlobalTopoServer() error =") + c.Nil(got, "Expected nil when Etcd spec is nil, got") }) t.Run("ControllerRefError", func(t *testing.T) { @@ -104,9 +89,7 @@ func TestBuildGlobalTopoServer(t *testing.T) { Etcd: &multigresv1alpha1.EtcdSpec{Image: "img"}, } _, err := BuildGlobalTopoServer(cluster, spec, emptyScheme) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) } @@ -138,36 +121,26 @@ func TestBuildMultiadminDeployment(t *testing.T) { } t.Run("Success", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildMultiadminDeployment(cluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } - - if got.Name != "my-cluster-multiadmin" { - t.Errorf("Name = %v, want %v", got.Name, "my-cluster-multiadmin") - } - if *got.Spec.Replicas != 2 { - t.Errorf("Replicas = %v, want 2", *got.Spec.Replicas) - } - if got.Spec.Template.Labels["custom"] != "label" { - t.Errorf("PodLabels missing custom label") - } - if got.Spec.Template.Annotations["anno"] != "tation" { - t.Errorf("PodAnnotations missing annotation") - } - assert.Contains(t, got.Spec.Template.Spec.Containers[0].Args, - "--topo-global-server-addresses=shared-etcd.default.svc:2379") - assert.Contains(t, got.Spec.Template.Spec.Containers[0].Args, - "--topo-global-root=/multigres/default/my-cluster/global") + c.Require().NoError(err, "BuildMultiadminDeployment() error =") + + c.Eq("my-cluster-multiadmin", got.Name, "Name") + c.Eq(2, *got.Spec.Replicas, "Replicas") + c.Eq("label", got.Spec.Template.Labels["custom"], "PodLabels missing custom label") + c.Eq("tation", got.Spec.Template.Annotations["anno"], "PodAnnotations missing annotation") + c.Contains( + got.Spec.Template.Spec.Containers[0].Args, + "--topo-global-server-addresses=shared-etcd.default.svc:2379", + ) + c.Contains( + got.Spec.Template.Spec.Containers[0].Args, + "--topo-global-root=/multigres/default/my-cluster/global", + ) // Verify container image from cluster spec if len(got.Spec.Template.Spec.Containers) > 0 { - if got.Spec.Template.Spec.Containers[0].Image != "multiadmin:latest" { - t.Errorf( - "Container Image = %v, want multiadmin:latest", - got.Spec.Template.Spec.Containers[0].Image, - ) - } + c.Eq("multiadmin:latest", got.Spec.Template.Spec.Containers[0].Image, "Container Image") } // Verify Selector does NOT contain mutable labels @@ -178,17 +151,20 @@ func TestBuildMultiadminDeployment(t *testing.T) { if _, ok := selector["app.kubernetes.io/managed-by"]; ok { t.Error("Selector should not contain app.kubernetes.io/managed-by") } - if _, ok := selector["app.kubernetes.io/component"]; !ok { - t.Error("Selector MUST contain app.kubernetes.io/component") - } + _, ok := selector["app.kubernetes.io/component"] + c.True(ok, "Selector MUST contain app.kubernetes.io/component") // Verify OwnerReference - if len(got.OwnerReferences) != 1 { - t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) - } + c.Len( + got.OwnerReferences, + 1, + "OwnerReferences count = %v, want 1", + len(got.OwnerReferences), + ) }) t.Run("Success with Observability", func(t *testing.T) { + c := assert.NewCollecting(t) obsCluster := cluster.DeepCopy() obsCluster.Spec.Observability = &multigresv1alpha1.ObservabilityConfig{ TracesSampler: "multigres_custom", @@ -198,44 +174,31 @@ func TestBuildMultiadminDeployment(t *testing.T) { }, } got, err := BuildMultiadminDeployment(obsCluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } - if len(got.Spec.Template.Spec.Volumes) == 0 { - t.Errorf("Expected OTEL volume to be added") - } - if len(got.Spec.Template.Spec.Containers[0].VolumeMounts) == 0 { - t.Errorf("Expected OTEL volume mount to be added") - } + c.Require().NoError(err, "BuildMultiadminDeployment() error =") + c.NotEmpty(got.Spec.Template.Spec.Volumes, "Expected OTEL volume to be added") + c.NotEmpty( + got.Spec.Template.Spec.Containers[0].VolumeMounts, + "Expected OTEL volume mount to be added", + ) }) t.Run("Success with internal mTLS and empty CertCommonName", func(t *testing.T) { + c := assert.NewCollecting(t) tlsCluster := cluster.DeepCopy() tlsCluster.Spec.InternalTLS = &multigresv1alpha1.InternalTLSConfig{Enabled: ptr.To(true)} - if tlsCluster.Spec.CertCommonName != "" { - t.Fatalf("test requires empty CertCommonName, got %q", tlsCluster.Spec.CertCommonName) - } + c.Require(). + Eq("", tlsCluster.Spec.CertCommonName, "test requires empty CertCommonName, got") got, err := BuildMultiadminDeployment(tlsCluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } + c.Require().NoError(err, "BuildMultiadminDeployment() error =") var foundVol bool for _, v := range got.Spec.Template.Spec.Volumes { if v.Name == multiAdminTLSVolumeName { foundVol = true - if v.Secret == nil { - t.Fatal("TLS volume should use Secret source") - } + c.Require().NotNil(v.Secret, "TLS volume should use Secret source") wantSecretName := "multiadmin.my-cluster.default.multigres.internal" //nolint:gosec // test constant - if v.Secret.SecretName != wantSecretName { - t.Errorf( - "TLS secretName = %q, want %q", - v.Secret.SecretName, - wantSecretName, - ) - } + c.Eq(wantSecretName, v.Secret.SecretName, "TLS secretName") if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf( "TLS secret defaultMode = %v, want 0444", @@ -244,27 +207,17 @@ func TestBuildMultiadminDeployment(t *testing.T) { } } } - if !foundVol { - t.Fatalf("expected TLS volume %q", multiAdminTLSVolumeName) - } + c.Require().True(foundVol, "expected TLS volume %q", multiAdminTLSVolumeName) container := got.Spec.Template.Spec.Containers[0] var foundMount bool for _, m := range container.VolumeMounts { if m.Name == multiAdminTLSVolumeName { foundMount = true - if m.MountPath != multiAdminTLSMountPath { - t.Errorf( - "TLS mount path = %q, want %q", - m.MountPath, - multiAdminTLSMountPath, - ) - } + c.Eq(multiAdminTLSMountPath, m.MountPath, "TLS mount path") } } - if !foundMount { - t.Fatalf("expected TLS volume mount %q", multiAdminTLSVolumeName) - } + c.Require().True(foundMount, "expected TLS volume mount %q", multiAdminTLSVolumeName) wantArgs := []string{ "--grpc-cert", multiAdminTLSCertFile, @@ -279,9 +232,7 @@ func TestBuildMultiadminDeployment(t *testing.T) { "--multipooler-grpc-require-tls", } tailArgs := container.Args[len(container.Args)-len(wantArgs):] - if diff := cmp.Diff(wantArgs, tailArgs); diff != "" { - t.Errorf("mTLS args mismatch (-want +got):\n%s", diff) - } + c.EqDiff(wantArgs, tailArgs, "mTLS args mismatch") }) for name, mutateCluster := range map[string]func(*multigresv1alpha1.MultigresCluster){ @@ -294,23 +245,18 @@ func TestBuildMultiadminDeployment(t *testing.T) { }, } { t.Run("No internal mTLS when "+name, func(t *testing.T) { + c := assert.NewCollecting(t) disabledCluster := cluster.DeepCopy() mutateCluster(disabledCluster) got, err := BuildMultiadminDeployment(disabledCluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } + c.Require().NoError(err, "BuildMultiadminDeployment() error =") for _, volume := range got.Spec.Template.Spec.Volumes { - if volume.Name == multiAdminTLSVolumeName { - t.Errorf("unexpected TLS volume %q", multiAdminTLSVolumeName) - } + c.NotEq(multiAdminTLSVolumeName, volume.Name, "unexpected TLS volume") } container := got.Spec.Template.Spec.Containers[0] for _, mount := range container.VolumeMounts { - if mount.Name == multiAdminTLSVolumeName { - t.Errorf("unexpected TLS volume mount %q", multiAdminTLSVolumeName) - } + c.NotEq(multiAdminTLSVolumeName, mount.Name, "unexpected TLS volume mount") } internalTLSArgs := map[string]struct{}{ "--grpc-cert": {}, @@ -324,14 +270,14 @@ func TestBuildMultiadminDeployment(t *testing.T) { "--multipooler-grpc-require-tls": {}, } for _, arg := range container.Args { - if _, found := internalTLSArgs[arg]; found { - t.Errorf("unexpected internal TLS argument %q", arg) - } + _, found := internalTLSArgs[arg] + c.False(found, "unexpected internal TLS argument %q", arg) } }) } t.Run("Success with tolerations", func(t *testing.T) { + c := assert.NewCollecting(t) placement := &multigresv1alpha1.PodPlacementSpec{ Tolerations: []corev1.Toleration{ { @@ -343,30 +289,21 @@ func TestBuildMultiadminDeployment(t *testing.T) { }, } got, err := BuildMultiadminDeployment(cluster, spec, placement, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } - if diff := cmp.Diff(placement.Tolerations, got.Spec.Template.Spec.Tolerations); diff != "" { - t.Errorf("Tolerations diff (-want +got):\n%s", diff) - } + c.Require().NoError(err, "BuildMultiadminDeployment() error =") + c.EqDiff(placement.Tolerations, got.Spec.Template.Spec.Tolerations, "Tolerations diff") }) t.Run("Success with nil placement", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildMultiadminDeployment(cluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } - if len(got.Spec.Template.Spec.Tolerations) != 0 { - t.Errorf("Tolerations = %v, want none", got.Spec.Template.Spec.Tolerations) - } + c.Require().NoError(err, "BuildMultiadminDeployment() error =") + c.Empty(got.Spec.Template.Spec.Tolerations, "Tolerations") }) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() _, err := BuildMultiadminDeployment(cluster, spec, nil, globalTopo, emptyScheme) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) } @@ -394,32 +331,22 @@ func TestBuildMultiadminWebDeployment(t *testing.T) { } t.Run("Success", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildMultiadminWebDeployment(cluster, spec, scheme) - if err != nil { - t.Fatalf("BuildMultiadminWebDeployment() error = %v", err) - } + c.Require().NoError(err, "BuildMultiadminWebDeployment() error =") - if got.Name != "my-cluster-multiadmin-web" { - t.Errorf("Name = %v, want %v", got.Name, "my-cluster-multiadmin-web") - } - if *got.Spec.Replicas != 2 { - t.Errorf("Replicas = %v, want 2", *got.Spec.Replicas) - } - if got.Spec.Template.Labels["custom"] != "label" { - t.Errorf("PodLabels missing custom label") - } - if got.Spec.Template.Annotations["anno"] != "tation" { - t.Errorf("PodAnnotations missing annotation") - } + c.Eq("my-cluster-multiadmin-web", got.Name, "Name") + c.Eq(2, *got.Spec.Replicas, "Replicas") + c.Eq("label", got.Spec.Template.Labels["custom"], "PodLabels missing custom label") + c.Eq("tation", got.Spec.Template.Annotations["anno"], "PodAnnotations missing annotation") // Verify container image from cluster spec if len(got.Spec.Template.Spec.Containers) > 0 { - if got.Spec.Template.Spec.Containers[0].Image != "multiadmin-web:latest" { - t.Errorf( - "Container Image = %v, want multiadmin-web:latest", - got.Spec.Template.Spec.Containers[0].Image, - ) - } + c.Eq( + "multiadmin-web:latest", + got.Spec.Template.Spec.Containers[0].Image, + "Container Image", + ) } // Verify env vars @@ -436,15 +363,11 @@ func TestBuildMultiadminWebDeployment(t *testing.T) { for _, ev := range envVars { if ev.Name == wantName { found = true - if ev.Value != wantValue { - t.Errorf("Env %s = %q, want %q", wantName, ev.Value, wantValue) - } + c.Eq(wantValue, ev.Value, "Env %s = %q, want", wantName, ev.Value) break } } - if !found { - t.Errorf("Missing env var %s", wantName) - } + c.True(found, "Missing env var %s", wantName) } // Verify Selector does NOT contain mutable labels @@ -455,44 +378,39 @@ func TestBuildMultiadminWebDeployment(t *testing.T) { if _, ok := selector["app.kubernetes.io/managed-by"]; ok { t.Error("Selector should not contain app.kubernetes.io/managed-by") } - if _, ok := selector["app.kubernetes.io/component"]; !ok { - t.Error("Selector MUST contain app.kubernetes.io/component") - } + _, ok := selector["app.kubernetes.io/component"] + c.True(ok, "Selector MUST contain app.kubernetes.io/component") // Verify OwnerReference - if len(got.OwnerReferences) != 1 { - t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) - } + c.Len( + got.OwnerReferences, + 1, + "OwnerReferences count = %v, want 1", + len(got.OwnerReferences), + ) }) t.Run("CustomPostgresSuperuser", func(t *testing.T) { + ck := assert.NewCollecting(t) c := *cluster c.Spec.PostgresSuperuser = "admin" got, err := BuildMultiadminWebDeployment(&c, spec, scheme) - if err != nil { - t.Fatalf("BuildMultiadminWebDeployment() error = %v", err) - } + ck.Require().NoError(err, "BuildMultiadminWebDeployment() error =") found := false for _, ev := range got.Spec.Template.Spec.Containers[0].Env { if ev.Name == "POSTGRES_USER" { found = true - if ev.Value != "admin" { - t.Errorf("POSTGRES_USER = %q, want %q", ev.Value, "admin") - } + ck.Eq("admin", ev.Value, "POSTGRES_USER") break } } - if !found { - t.Fatal("Missing env var POSTGRES_USER") - } + ck.Require().True(found, "Missing env var POSTGRES_USER") }) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() _, err := BuildMultiadminWebDeployment(cluster, spec, emptyScheme) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) } @@ -589,33 +507,35 @@ func TestBuildMultiadminWebService(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildMultiadminWebService(cluster, tc.extAW, scheme) - require.NoError(t, err) + c.Require().NoError(err) - assert.Equal(t, "my-cluster-multiadmin-web", got.Name) - assert.Equal(t, "default", got.Namespace) - assert.Equal(t, tc.wantType, got.Spec.Type) - assert.Equal(t, tc.wantExternalIPs, got.Spec.ExternalIPs) + c.EqDeep("my-cluster-multiadmin-web", got.Name) + c.EqDeep("default", got.Namespace) + c.EqDeep(tc.wantType, got.Spec.Type) + c.EqDeep(tc.wantExternalIPs, got.Spec.ExternalIPs) - require.Len(t, got.Spec.Ports, 1) - assert.Equal(t, wantPort, got.Spec.Ports[0]) + c.Require().Len(got.Spec.Ports, 1) + c.EqDeep(wantPort, got.Spec.Ports[0]) - assert.Equal(t, wantLabels, got.Labels) + c.EqDeep(wantLabels, got.Labels) if tc.wantAnnotations != nil { for k, v := range tc.wantAnnotations { - assert.Equal(t, v, got.Annotations[k], "annotation %s", k) + c.EqDeep(v, got.Annotations[k], "annotation %s", k) } } else { - assert.Empty(t, got.Annotations) + c.Empty(got.Annotations) } - require.Len(t, got.OwnerReferences, 1) - assert.Equal(t, "my-cluster", got.OwnerReferences[0].Name) + c.Require().Len(got.OwnerReferences, 1) + c.EqDeep("my-cluster", got.OwnerReferences[0].Name) }) } t.Run("Annotation removal on disable", func(t *testing.T) { + c := assert.NewCollecting(t) enabledCfg := &multigresv1alpha1.ExternalAdminWebConfig{ Enabled: true, Annotations: map[string]string{ @@ -623,25 +543,21 @@ func TestBuildMultiadminWebService(t *testing.T) { }, } enabled, err := BuildMultiadminWebService(cluster, enabledCfg, scheme) - require.NoError(t, err) - assert.Equal(t, corev1.ServiceTypeClusterIP, enabled.Spec.Type) - assert.Equal( - t, - "platform-engineering", - enabled.Annotations["team.example.com/owner"], - ) + c.Require().NoError(err) + c.EqDeep(corev1.ServiceTypeClusterIP, enabled.Spec.Type) + c.EqDeep("platform-engineering", enabled.Annotations["team.example.com/owner"]) disabledCfg := &multigresv1alpha1.ExternalAdminWebConfig{Enabled: false} disabled, err := BuildMultiadminWebService(cluster, disabledCfg, scheme) - require.NoError(t, err) - assert.Equal(t, corev1.ServiceTypeClusterIP, disabled.Spec.Type) - assert.Empty(t, disabled.Annotations) + c.Require().NoError(err) + c.EqDeep(corev1.ServiceTypeClusterIP, disabled.Spec.Type) + c.Empty(disabled.Annotations) }) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() _, err := BuildMultiadminWebService(cluster, nil, emptyScheme) - assert.Error(t, err) + assert.NewCollecting(t).Error(err) }) } @@ -659,27 +575,25 @@ func TestBuildMultiadminService(t *testing.T) { } t.Run("Success", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildMultiadminService(cluster, scheme) - if err != nil { - t.Fatalf("BuildMultiadminService() error = %v", err) - } + c.Require().NoError(err, "BuildMultiadminService() error =") - if got.Name != "my-cluster-multiadmin" { - t.Errorf("Name = %v, want %v", got.Name, "my-cluster-multiadmin") - } + c.Eq("my-cluster-multiadmin", got.Name, "Name") // Verify OwnerReference - if len(got.OwnerReferences) != 1 { - t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) - } + c.Len( + got.OwnerReferences, + 1, + "OwnerReferences count = %v, want 1", + len(got.OwnerReferences), + ) }) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() _, err := BuildMultiadminService(cluster, emptyScheme) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) } @@ -783,46 +697,48 @@ func TestBuildMultigatewayGlobalService(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildMultigatewayGlobalService(cluster, tc.extGw, scheme) - require.NoError(t, err) + c.Require().NoError(err) // Name and namespace - assert.Equal(t, "my-cluster-multigateway", got.Name) - assert.Equal(t, "default", got.Namespace) + c.EqDeep("my-cluster-multigateway", got.Name) + c.EqDeep("default", got.Namespace) // Service type - assert.Equal(t, tc.wantType, got.Spec.Type) - assert.Equal(t, tc.wantExternalIPs, got.Spec.ExternalIPs) + c.EqDeep(tc.wantType, got.Spec.Type) + c.EqDeep(tc.wantExternalIPs, got.Spec.ExternalIPs) // Port 5432 invariant - require.Len(t, got.Spec.Ports, 1) - assert.Equal(t, wantPort, got.Spec.Ports[0]) + c.Require().Len(got.Spec.Ports, 1) + c.EqDeep(wantPort, got.Spec.Ports[0]) // Labels preserved - assert.Equal(t, wantLabels, got.Labels) + c.EqDeep(wantLabels, got.Labels) // Annotations if tc.wantAnnotations != nil { for k, v := range tc.wantAnnotations { - assert.Equal(t, v, got.Annotations[k], "annotation %s", k) + c.EqDeep(v, got.Annotations[k], "annotation %s", k) } } else { // No gateway annotations expected; annotations should be nil or empty - assert.Empty(t, got.Annotations) + c.Empty(got.Annotations) } // Selector: component + instance, no cell label - assert.Equal(t, "multigateway", got.Spec.Selector["app.kubernetes.io/component"]) - assert.Equal(t, "my-cluster", got.Spec.Selector["app.kubernetes.io/instance"]) - assert.NotContains(t, got.Spec.Selector, "multigres.com/cell") + c.EqDeep("multigateway", got.Spec.Selector["app.kubernetes.io/component"]) + c.EqDeep("my-cluster", got.Spec.Selector["app.kubernetes.io/instance"]) + c.NotHasKey(got.Spec.Selector, "multigres.com/cell") // Owner reference - require.Len(t, got.OwnerReferences, 1) - assert.Equal(t, "my-cluster", got.OwnerReferences[0].Name) + c.Require().Len(got.OwnerReferences, 1) + c.EqDeep("my-cluster", got.OwnerReferences[0].Name) }) } t.Run("Annotation removal on disable", func(t *testing.T) { + c := assert.NewCollecting(t) // Build with annotations enabled enabledCfg := &multigresv1alpha1.ExternalGatewayConfig{ Enabled: true, @@ -831,30 +747,27 @@ func TestBuildMultigatewayGlobalService(t *testing.T) { }, } enabled, err := BuildMultigatewayGlobalService(cluster, enabledCfg, scheme) - require.NoError(t, err) - assert.Equal(t, corev1.ServiceTypeClusterIP, enabled.Spec.Type) - assert.Equal( - t, - "platform-engineering", - enabled.Annotations["team.example.com/owner"], - ) + c.Require().NoError(err) + c.EqDeep(corev1.ServiceTypeClusterIP, enabled.Spec.Type) + c.EqDeep("platform-engineering", enabled.Annotations["team.example.com/owner"]) // Build with disabled config; previously-set gateway annotations absent disabledCfg := &multigresv1alpha1.ExternalGatewayConfig{Enabled: false} disabled, err := BuildMultigatewayGlobalService(cluster, disabledCfg, scheme) - require.NoError(t, err) - assert.Equal(t, corev1.ServiceTypeClusterIP, disabled.Spec.Type) - assert.Empty(t, disabled.Annotations) + c.Require().NoError(err) + c.EqDeep(corev1.ServiceTypeClusterIP, disabled.Spec.Type) + c.Empty(disabled.Annotations) }) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() _, err := BuildMultigatewayGlobalService(cluster, nil, emptyScheme) - assert.Error(t, err) + assert.NewCollecting(t).Error(err) }) } func TestBuildMultigatewayGlobalReplicaService(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -868,32 +781,32 @@ func TestBuildMultigatewayGlobalReplicaService(t *testing.T) { } got, err := BuildMultigatewayGlobalReplicaService(cluster, scheme) - require.NoError(t, err) + c.Require().NoError(err) - assert.Equal(t, "my-cluster-multigateway-replica", got.Name) - assert.Equal(t, "default", got.Namespace) - assert.Equal(t, corev1.ServiceTypeClusterIP, got.Spec.Type) + c.EqDeep("my-cluster-multigateway-replica", got.Name) + c.EqDeep("default", got.Namespace) + c.EqDeep(corev1.ServiceTypeClusterIP, got.Spec.Type) - require.Len(t, got.Spec.Ports, 1) - assert.Equal(t, corev1.ServicePort{ + c.Require().Len(got.Spec.Ports, 1) + c.EqDeep(corev1.ServicePort{ Name: "pg-replica", Port: 5433, TargetPort: intstr.FromString("pg-replica"), Protocol: corev1.ProtocolTCP, }, got.Spec.Ports[0]) - assert.Equal(t, map[string]string{ + c.EqDeep(map[string]string{ "app.kubernetes.io/component": "multigateway", "app.kubernetes.io/instance": "my-cluster", }, got.Spec.Selector) - require.Len(t, got.OwnerReferences, 1) - assert.Equal(t, "my-cluster", got.OwnerReferences[0].Name) + c.Require().Len(got.OwnerReferences, 1) + c.EqDeep("my-cluster", got.OwnerReferences[0].Name) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() _, err := BuildMultigatewayGlobalReplicaService(cluster, emptyScheme) - assert.Error(t, err) + assert.NewCollecting(t).Error(err) }) } @@ -916,21 +829,24 @@ func TestBuildAdminNetworkPolicies(t *testing.T) { } t.Run("NilConfig", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildAdminNetworkPolicies(baseCluster(nil), scheme) - require.NoError(t, err) - assert.Nil(t, got) + c.Require().NoError(err) + c.Nil(got) }) t.Run("Disabled", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildAdminNetworkPolicies( baseCluster(&multigresv1alpha1.NetworkPolicyConfig{Enabled: false}), scheme, ) - require.NoError(t, err) - assert.Nil(t, got) + c.Require().NoError(err) + c.Nil(got) }) t.Run("EnabledWithAllowedNamespaces", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildAdminNetworkPolicies( baseCluster(&multigresv1alpha1.NetworkPolicyConfig{ Enabled: true, @@ -938,8 +854,8 @@ func TestBuildAdminNetworkPolicies(t *testing.T) { }), scheme, ) - require.NoError(t, err) - require.Len(t, got, 2) + c.Require().NoError(err) + c.Require().Len(got, 2) wantNames := []string{ "my-cluster-multiadmin-restrict-ingress", @@ -948,50 +864,48 @@ func TestBuildAdminNetworkPolicies(t *testing.T) { wantComponents := []string{"multiadmin", "multiadmin-web"} for i, policy := range got { - assert.Equal(t, wantNames[i], policy.Name) - assert.Equal(t, "default", policy.Namespace) - require.Len(t, policy.OwnerReferences, 1) - - assert.Equal(t, - map[string]string{ - "app.kubernetes.io/component": wantComponents[i], - "app.kubernetes.io/instance": "my-cluster", - "multigres.com/cluster": "my-cluster", - }, - policy.Spec.PodSelector.MatchLabels, - ) - assert.Equal(t, + c.EqDeep(wantNames[i], policy.Name) + c.EqDeep("default", policy.Namespace) + c.Require().Len(policy.OwnerReferences, 1) + + c.EqDeep(map[string]string{ + "app.kubernetes.io/component": wantComponents[i], + "app.kubernetes.io/instance": "my-cluster", + "multigres.com/cluster": "my-cluster", + }, policy.Spec.PodSelector.MatchLabels) + c.EqDeep( []networkingv1.PolicyType{networkingv1.PolicyTypeIngress}, policy.Spec.PolicyTypes, ) - require.Len(t, policy.Spec.Ingress, 1) + c.Require().Len(policy.Spec.Ingress, 1) peers := policy.Spec.Ingress[0].From - require.Len(t, peers, 3) + c.Require().Len(peers, 3) // An empty pod selector without a namespace selector matches local pods. - assert.NotNil(t, peers[0].PodSelector) - assert.Nil(t, peers[0].NamespaceSelector) + c.NotNil(peers[0].PodSelector) + c.Nil(peers[0].NamespaceSelector) for j, ns := range []string{"envoy-gateway-system", "multigres-operator"} { peer := peers[j+1] - require.NotNil(t, peer.NamespaceSelector) - assert.Equal(t, + c.Require().NotNil(peer.NamespaceSelector) + c.EqDeep( map[string]string{"kubernetes.io/metadata.name": ns}, peer.NamespaceSelector.MatchLabels, ) - assert.Nil(t, peer.PodSelector) + c.Nil(peer.PodSelector) } } }) t.Run("EnabledWithoutAllowedNamespaces", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildAdminNetworkPolicies( baseCluster(&multigresv1alpha1.NetworkPolicyConfig{Enabled: true}), scheme, ) - require.NoError(t, err) - require.Len(t, got, 2) - require.Len(t, got[0].Spec.Ingress, 1) - assert.Len(t, got[0].Spec.Ingress[0].From, 1) + c.Require().NoError(err) + c.Require().Len(got, 2) + c.Require().Len(got[0].Spec.Ingress, 1) + c.Len(got[0].Spec.Ingress[0].From, 1) }) t.Run("ControllerRefError", func(t *testing.T) { @@ -1000,7 +914,7 @@ func TestBuildAdminNetworkPolicies(t *testing.T) { baseCluster(&multigresv1alpha1.NetworkPolicyConfig{Enabled: true}), emptyScheme, ) - assert.Error(t, err) + assert.NewCollecting(t).Error(err) }) } @@ -1021,6 +935,7 @@ func TestBuildMultiadminDeployment_TopoClientTLS(t *testing.T) { spec := &multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))} t.Run("presents the client certificate when the reference carries it", func(t *testing.T) { + ck := assert.NewCollecting(t) secret := multigresv1alpha1.TopoClientCertSecretName("my-cluster") globalTopo := multigresv1alpha1.GlobalTopoServerRef{ Address: "my-cluster-global-topo.default.svc:2379", @@ -1029,40 +944,44 @@ func TestBuildMultiadminDeployment_TopoClientTLS(t *testing.T) { ClientCertSecret: secret, } got, err := BuildMultiadminDeployment(cluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } + ck.Require().NoError(err, "BuildMultiadminDeployment() error =") c := got.Spec.Template.Spec.Containers[0] - if !hasArgValue(c.Args, "--topo-etcd-tls-cert", multigresv1alpha1.TopoClientTLSCertFile) { - t.Errorf("missing --topo-etcd-tls-cert flag: %v", c.Args) - } - if !containerMountsVolume(c, multigresv1alpha1.TopoClientTLSVolumeName) { - t.Error("multiadmin does not mount the topo client certificate") - } + ck.True( + hasArgValue(c.Args, "--topo-etcd-tls-cert", multigresv1alpha1.TopoClientTLSCertFile), + "missing --topo-etcd-tls-cert flag: %v", + c.Args, + ) + ck.True( + containerMountsVolume(c, multigresv1alpha1.TopoClientTLSVolumeName), + "multiadmin does not mount the topo client certificate", + ) volumes := got.Spec.Template.Spec.Volumes - if !podHasVolume(volumes, multigresv1alpha1.TopoClientTLSVolumeName) { - t.Error("topo client volume missing from pod spec") - } + ck.True( + podHasVolume(volumes, multigresv1alpha1.TopoClientTLSVolumeName), + "topo client volume missing from pod spec", + ) }) t.Run("renders unchanged when the reference carries no credential", func(t *testing.T) { + ck := assert.NewCollecting(t) globalTopo := multigresv1alpha1.GlobalTopoServerRef{ Address: "my-cluster-global-topo.default.svc:2379", RootPath: "/multigres/global", } got, err := BuildMultiadminDeployment(cluster, spec, nil, globalTopo, scheme) - if err != nil { - t.Fatalf("BuildMultiadminDeployment() error = %v", err) - } + ck.Require().NoError(err, "BuildMultiadminDeployment() error =") c := got.Spec.Template.Spec.Containers[0] for _, a := range c.Args { - if a == "--topo-etcd-tls-cert" { - t.Error("topo TLS flag present with no credential on the reference") - } - } - if podHasVolume(got.Spec.Template.Spec.Volumes, multigresv1alpha1.TopoClientTLSVolumeName) { - t.Error("topo client volume present with no credential on the reference") + ck.NotEq( + "--topo-etcd-tls-cert", + a, + "topo TLS flag present with no credential on the reference", + ) } + ck.False( + podHasVolume(got.Spec.Template.Spec.Volumes, multigresv1alpha1.TopoClientTLSVolumeName), + "topo client volume present with no credential on the reference", + ) }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/builders_tablegroup_test.go b/pkg/cluster-handler/controller/multigrescluster/builders_tablegroup_test.go index d9908026..f0b8d4fa 100644 --- a/pkg/cluster-handler/controller/multigrescluster/builders_tablegroup_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/builders_tablegroup_test.go @@ -11,6 +11,8 @@ import ( "k8s.io/utils/ptr" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestBuildTableGroup(t *testing.T) { @@ -31,6 +33,7 @@ func TestBuildTableGroup(t *testing.T) { globalTopoRef := multigresv1alpha1.GlobalTopoServerRef{} t.Run("Success", func(t *testing.T) { + c := assert.NewCollecting(t) tgCfg := &multigresv1alpha1.TableGroupConfig{ Name: "tg-1", } @@ -39,9 +42,7 @@ func TestBuildTableGroup(t *testing.T) { } got, err := BuildTableGroup(cluster, dbCfg, tgCfg, resolvedShards, globalTopoRef, scheme) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } + c.Require().NoError(err, "BuildTableGroup() error =") // Calculate expected hash: md5("my-cluster", "my-db", "tg-1") -> "d5708433" expectedName := name.JoinWithConstraints( @@ -50,55 +51,41 @@ func TestBuildTableGroup(t *testing.T) { "my-db", "tg-1", ) - if got.Name != expectedName { - t.Errorf("Name = %v, want %v", got.Name, expectedName) - } - if got.Labels["multigres.com/database"] != "my-db" { - t.Errorf("Label[database] = %v, want my-db", got.Labels["multigres.com/database"]) - } + c.Eq(expectedName, got.Name, "Name") + c.Eq("my-db", got.Labels["multigres.com/database"], "Label[database]") // Verify OwnerReference - if len(got.OwnerReferences) != 1 { - t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) - } + c.Len( + got.OwnerReferences, + 1, + "OwnerReferences count = %v, want 1", + len(got.OwnerReferences), + ) }) t.Run("CustomPostgresSuperuser", func(t *testing.T) { + ck := assert.NewCollecting(t) c := *cluster c.Spec.PostgresSuperuser = "admin" tgCfg := &multigresv1alpha1.TableGroupConfig{Name: "tg-superuser"} got, err := BuildTableGroup(&c, dbCfg, tgCfg, nil, globalTopoRef, scheme) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } - if got.Spec.PostgresSuperuser != "admin" { - t.Errorf( - "PostgresSuperuser = %q, want %q", - got.Spec.PostgresSuperuser, - "admin", - ) - } + ck.Require().NoError(err, "BuildTableGroup() error =") + ck.Eq("admin", got.Spec.PostgresSuperuser, "PostgresSuperuser") }) t.Run("Propagates InternalTLS", func(t *testing.T) { + ck := assert.NewAborting(t) c := cluster.DeepCopy() c.Spec.InternalTLS = &multigresv1alpha1.InternalTLSConfig{Enabled: ptr.To(true)} tgCfg := &multigresv1alpha1.TableGroupConfig{Name: "tg-internal-tls"} got, err := BuildTableGroup(c, dbCfg, tgCfg, nil, globalTopoRef, scheme) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } - if got.Spec.InternalTLS != c.Spec.InternalTLS { - t.Fatalf( - "InternalTLS = %#v, want propagated pointer %#v", - got.Spec.InternalTLS, - c.Spec.InternalTLS, - ) - } + ck.NoError(err, "BuildTableGroup() error =") + ck.Eq(c.Spec.InternalTLS, got.Spec.InternalTLS, "InternalTLS") }) t.Run("PostgresPasswordSecretRef", func(t *testing.T) { + ck := assert.NewCollecting(t) c := *cluster c.Spec.PostgresPasswordSecretRef = multigresv1alpha1.PostgresPasswordSecretRef{ Name: "multigres-admin-password", @@ -106,26 +93,17 @@ func TestBuildTableGroup(t *testing.T) { } tgCfg := &multigresv1alpha1.TableGroupConfig{Name: "tg-password"} got, err := BuildTableGroup(&c, dbCfg, tgCfg, nil, globalTopoRef, scheme) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } - if got.Spec.PostgresPasswordSecretRef.Name != "multigres-admin-password" { - t.Errorf( - "PostgresPasswordSecretRef.Name = %q, want %q", - got.Spec.PostgresPasswordSecretRef.Name, - "multigres-admin-password", - ) - } - if got.Spec.PostgresPasswordSecretRef.Key != "current" { - t.Errorf( - "PostgresPasswordSecretRef.Key = %q, want %q", - got.Spec.PostgresPasswordSecretRef.Key, - "current", - ) - } + ck.Require().NoError(err, "BuildTableGroup() error =") + ck.Eq( + "multigres-admin-password", + got.Spec.PostgresPasswordSecretRef.Name, + "PostgresPasswordSecretRef.Name", + ) + ck.Eq("current", got.Spec.PostgresPasswordSecretRef.Key, "PostgresPasswordSecretRef.Key") }) t.Run("PostgresInitSecretsRef propagated when set", func(t *testing.T) { + ck := assert.NewCollecting(t) c := *cluster c.Spec.PostgresInitSecretsRef = &multigresv1alpha1.PostgresInitSecretsRef{ Name: "multigres-init-secrets", @@ -133,42 +111,33 @@ func TestBuildTableGroup(t *testing.T) { } tgCfg := &multigresv1alpha1.TableGroupConfig{Name: "tg-init-secrets"} got, err := BuildTableGroup(&c, dbCfg, tgCfg, nil, globalTopoRef, scheme) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } - if got.Spec.PostgresInitSecretsRef == nil { - t.Fatal("PostgresInitSecretsRef = nil, want propagated ref") - } - if got.Spec.PostgresInitSecretsRef.Name != "multigres-init-secrets" { - t.Errorf( - "PostgresInitSecretsRef.Name = %q, want %q", - got.Spec.PostgresInitSecretsRef.Name, - "multigres-init-secrets", - ) - } - if got.Spec.PostgresInitSecretsRef.Key != "init-secrets.json" { - t.Errorf( - "PostgresInitSecretsRef.Key = %q, want %q", - got.Spec.PostgresInitSecretsRef.Key, - "init-secrets.json", - ) - } + ck.Require().NoError(err, "BuildTableGroup() error =") + ck.Require(). + NotNil(got.Spec.PostgresInitSecretsRef, "PostgresInitSecretsRef = nil, want propagated ref") + ck.Eq( + "multigres-init-secrets", + got.Spec.PostgresInitSecretsRef.Name, + "PostgresInitSecretsRef.Name", + ) + ck.Eq( + "init-secrets.json", + got.Spec.PostgresInitSecretsRef.Key, + "PostgresInitSecretsRef.Key", + ) }) t.Run("PostgresInitSecretsRef nil when unset", func(t *testing.T) { + ck := assert.NewCollecting(t) c := *cluster c.Spec.PostgresInitSecretsRef = nil tgCfg := &multigresv1alpha1.TableGroupConfig{Name: "tg-no-init-secrets"} got, err := BuildTableGroup(&c, dbCfg, tgCfg, nil, globalTopoRef, scheme) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } - if got.Spec.PostgresInitSecretsRef != nil { - t.Errorf("PostgresInitSecretsRef = %+v, want nil", got.Spec.PostgresInitSecretsRef) - } + ck.Require().NoError(err, "BuildTableGroup() error =") + ck.Nil(got.Spec.PostgresInitSecretsRef, "PostgresInitSecretsRef") }) t.Run("Name Truncation", func(t *testing.T) { + c := assert.NewCollecting(t) longName := strings.Repeat("a", 250) // Very long name tgCfg := &multigresv1alpha1.TableGroupConfig{ Name: multigresv1alpha1.TableGroupName(longName), @@ -176,21 +145,16 @@ func TestBuildTableGroup(t *testing.T) { resolvedShards := []multigresv1alpha1.ShardResolvedSpec{} got, err := BuildTableGroup(cluster, dbCfg, tgCfg, resolvedShards, globalTopoRef, scheme) - if err != nil { - t.Errorf("BuildTableGroup() error = %v, want nil", err) - } + c.NoError(err, "BuildTableGroup() error") // Should be truncated to 253 chars - if len(got.Name) > 253 { - t.Errorf("Expected name length <= 253, got %d", len(got.Name)) - } + c.LessOrEqual(253, len(got.Name), "Expected name length <= 253, got") // Confirm it ends with a hash (8 chars) // and has the truncation mark "---" - if !strings.Contains(got.Name, "---") { - t.Errorf("Expected truncation mark '---', got %s", got.Name) - } + c.StrContains(got.Name, "---", "Expected truncation mark '---', got") }) t.Run("CellTopologyLabels ZoneID", func(t *testing.T) { + ck := assert.NewCollecting(t) c := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "my-cluster", @@ -209,18 +173,17 @@ func TestBuildTableGroup(t *testing.T) { globalTopoRef, scheme, ) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } - if got.Spec.CellTopologyLabels["az-cell"]["topology.k8s.aws/zone-id"] != "use1-az1" { - t.Errorf( - "expected topology.k8s.aws/zone-id=use1-az1, got %v", - got.Spec.CellTopologyLabels["az-cell"], - ) - } + ck.Require().NoError(err, "BuildTableGroup() error =") + ck.Eq( + "use1-az1", + got.Spec.CellTopologyLabels["az-cell"]["topology.k8s.aws/zone-id"], + "expected topology.k8s.aws/zone-id=use1-az1, got %v", + got.Spec.CellTopologyLabels["az-cell"], + ) }) t.Run("CellTopologyLabels ZoneID only", func(t *testing.T) { + ck := assert.NewCollecting(t) c := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "my-cluster", @@ -241,16 +204,18 @@ func TestBuildTableGroup(t *testing.T) { globalTopoRef, scheme, ) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } + ck.Require().NoError(err, "BuildTableGroup() error =") labels := got.Spec.CellTopologyLabels["both-cell"] - if labels["topology.k8s.aws/zone-id"] != "use1-az1" { - t.Errorf("expected topology.k8s.aws/zone-id=use1-az1, got %v", labels) - } + ck.Eq( + "use1-az1", + labels["topology.k8s.aws/zone-id"], + "expected topology.k8s.aws/zone-id=use1-az1, got %v", + labels, + ) }) t.Run("CellTopologyLabels Region", func(t *testing.T) { + c := assert.NewCollecting(t) regionCluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "my-cluster", @@ -274,19 +239,14 @@ func TestBuildTableGroup(t *testing.T) { globalTopoRef, scheme, ) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } + c.Require().NoError(err, "BuildTableGroup() error =") labels, ok := got.Spec.CellTopologyLabels["region-cell"] - if !ok { - t.Fatal("Expected CellTopologyLabels to contain region-cell") - } - if labels["topology.kubernetes.io/region"] != "us-east-1" { - t.Errorf( - "Expected region label us-east-1, got %s", - labels["topology.kubernetes.io/region"], - ) - } + c.Require().True(ok, "Expected CellTopologyLabels to contain region-cell") + c.Eq( + "us-east-1", + labels["topology.kubernetes.io/region"], + "Expected region label us-east-1, got", + ) }) t.Run("ControllerRefError", func(t *testing.T) { @@ -302,12 +262,11 @@ func TestBuildTableGroup(t *testing.T) { globalTopoRef, emptyScheme, ) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) t.Run("Propagates explicit project ref annotation", func(t *testing.T) { + c := assert.NewAborting(t) clusterWithProjectRef := cluster.DeepCopy() clusterWithProjectRef.Annotations = map[string]string{ metadata.AnnotationProjectRef: "proj_123", @@ -323,17 +282,14 @@ func TestBuildTableGroup(t *testing.T) { globalTopoRef, scheme, ) - if err != nil { - t.Fatalf("BuildTableGroup() error = %v", err) - } + c.NoError(err, "BuildTableGroup() error =") - if got.Annotations[metadata.AnnotationProjectRef] != "proj_123" { - t.Fatalf( - "annotation %q = %q, want %q", - metadata.AnnotationProjectRef, - got.Annotations[metadata.AnnotationProjectRef], - "proj_123", - ) - } + c.Eq( + "proj_123", + got.Annotations[metadata.AnnotationProjectRef], + "annotation %q = %q, want", + metadata.AnnotationProjectRef, + got.Annotations[metadata.AnnotationProjectRef], + ) }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/builders_test.go b/pkg/cluster-handler/controller/multigrescluster/builders_test.go index f5cc5063..8a71f94d 100644 --- a/pkg/cluster-handler/controller/multigrescluster/builders_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/builders_test.go @@ -6,6 +6,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" + + "github.com/multigres/testkit/assert" ) func TestBuildGlobalTopoServer_Errors(t *testing.T) { @@ -22,7 +24,5 @@ func TestBuildGlobalTopoServer_Errors(t *testing.T) { } _, err := BuildGlobalTopoServer(cluster, cluster.Spec.GlobalTopoServer, scheme) - if err != nil { - t.Errorf("Expected nil error for nil GlobalTopoServerSpec, got %v", err) - } + assert.NewCollecting(t).NoError(err, "Expected nil error for nil GlobalTopoServerSpec, got") } diff --git a/pkg/cluster-handler/controller/multigrescluster/certificate_test.go b/pkg/cluster-handler/controller/multigrescluster/certificate_test.go index 5aed7a7d..431cc7fa 100644 --- a/pkg/cluster-handler/controller/multigrescluster/certificate_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/certificate_test.go @@ -19,6 +19,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/interceptor" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) // registerCertManagerTypes registers cert-manager Certificate as an @@ -111,54 +113,41 @@ func TestBuildCertificate(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) got, err := buildCertificate(tc.cluster, scheme) - if err != nil { - t.Fatalf("buildCertificate() error: %v", err) - } + c.Require().NoError(err, "buildCertificate() error") wantGVK := schema.GroupVersionKind{ Group: "cert-manager.io", Version: "v1", Kind: "Certificate", } - if diff := cmp.Diff(wantGVK, got.GroupVersionKind()); diff != "" { - t.Errorf("GVK mismatch (-want +got):\n%s", diff) - } - if got.GetName() != tc.wantName { - t.Errorf("Name = %q, want %q", got.GetName(), tc.wantName) - } + c.EqDiff(wantGVK, got.GroupVersionKind(), "GVK mismatch") + c.Eq(tc.wantName, got.GetName(), "Name") // Verify owner reference points to the cluster ownerRefs := got.GetOwnerReferences() - if len(ownerRefs) != 1 { - t.Fatalf("expected 1 ownerReference, got %d", len(ownerRefs)) - } - if ownerRefs[0].Name != tc.cluster.Name { - t.Errorf( - "ownerRef.Name = %q, want %q", - ownerRefs[0].Name, tc.cluster.Name, - ) - } - if ownerRefs[0].Kind != "MultigresCluster" { - t.Errorf( - "ownerRef.Kind = %q, want MultigresCluster", - ownerRefs[0].Kind, - ) - } + c.Require().Len(ownerRefs, 1, "expected 1 ownerReference, got %d", len(ownerRefs)) + c.Eq(tc.cluster.Name, ownerRefs[0].Name, "ownerRef.Name") + c.Eq("MultigresCluster", ownerRefs[0].Kind, "ownerRef.Kind") spec, ok := got.Object["spec"].(map[string]any) - if !ok { - t.Fatal("spec is not a map") - } - if diff := cmp.Diff(tc.wantDNSNames, spec["dnsNames"]); diff != "" { - t.Errorf("dnsNames mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantSubject, spec["literalSubject"]); diff != "" { - t.Errorf("literalSubject mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantSecretName, spec["secretName"]); diff != "" { - t.Errorf("secretName mismatch (-want +got):\n%s", diff) - } + c.Require().True(ok, "spec is not a map") + c.Eq( + "", + cmp.Diff(tc.wantDNSNames, spec["dnsNames"]), + "dnsNames mismatch (-want +got):\n", + ) + c.Eq( + "", + cmp.Diff(tc.wantSubject, spec["literalSubject"]), + "literalSubject mismatch (-want +got):\n", + ) + c.Eq( + "", + cmp.Diff(tc.wantSecretName, spec["secretName"]), + "secretName mismatch (-want +got):\n", + ) wantIssuer := tc.wantIssuerName if wantIssuer == "" { wantIssuer = CertIssuerName @@ -168,22 +157,23 @@ func TestBuildCertificate(t *testing.T) { "kind": "ClusterIssuer", "group": "cert-manager.io", } - if diff := cmp.Diff(wantIssuerRef, spec["issuerRef"]); diff != "" { - t.Errorf("issuerRef mismatch (-want +got):\n%s", diff) - } + c.Eq( + "", + cmp.Diff(wantIssuerRef, spec["issuerRef"]), + "issuerRef mismatch (-want +got):\n", + ) wantUsages := []any{ "digital signature", "key encipherment", "server auth", } - if diff := cmp.Diff(wantUsages, spec["usages"]); diff != "" { - t.Errorf("usages mismatch (-want +got):\n%s", diff) - } + c.Eq("", cmp.Diff(wantUsages, spec["usages"]), "usages mismatch (-want +got):\n") }) } } func TestBuildInternalCertificates(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ TypeMeta: metav1.TypeMeta{ @@ -202,9 +192,7 @@ func TestBuildInternalCertificates(t *testing.T) { } got, err := buildInternalCertificates(cluster, scheme) - if err != nil { - t.Fatalf("buildInternalCertificates() error: %v", err) - } + c.Require().NoError(err, "buildInternalCertificates() error") want := map[string]string{ "multiadmin.test-cluster.supabase.multigres.internal": multigresv1alpha1.ComponentCertSecretName( @@ -233,26 +221,28 @@ func TestBuildInternalCertificates(t *testing.T) { cluster.Namespace, ), } - if len(got) != len(want) { - t.Fatalf("got %d certs, want %d", len(got), len(want)) - } + c.Require().Len(got, len(want), "got %d certs, want", len(got)) for _, cert := range got { secretName, ok := want[cert.GetName()] if !ok { t.Fatalf("unexpected Certificate %q", cert.GetName()) } spec, ok := cert.Object["spec"].(map[string]any) - if !ok { - t.Fatal("spec is not a map") - } - if diff := cmp.Diff(secretName, spec["secretName"]); diff != "" { - t.Errorf("secretName mismatch for %s (-want +got):\n%s", cert.GetName(), diff) - } + c.Require().True(ok, "spec is not a map") + c.Eq( + "", + cmp.Diff(secretName, spec["secretName"]), + "secretName mismatch for %s (-want +got):\n", + cert.GetName(), + ) wantCommonName := cert.GetName() wantSubject := "C=US, ST=Delware, L=New Castle,O=Supabase Inc, CN=" + wantCommonName - if diff := cmp.Diff(wantSubject, spec["literalSubject"]); diff != "" { - t.Errorf("literalSubject mismatch for %s (-want +got):\n%s", cert.GetName(), diff) - } + c.Eq( + "", + cmp.Diff(wantSubject, spec["literalSubject"]), + "literalSubject mismatch for %s (-want +got):\n", + cert.GetName(), + ) wantUsages := []any{ "digital signature", "key encipherment", @@ -266,9 +256,12 @@ func TestBuildInternalCertificates(t *testing.T) { "client auth", } } - if diff := cmp.Diff(wantUsages, spec["usages"]); diff != "" { - t.Errorf("usages mismatch for %s (-want +got):\n%s", cert.GetName(), diff) - } + c.Eq( + "", + cmp.Diff(wantUsages, spec["usages"]), + "usages mismatch for %s (-want +got):\n", + cert.GetName(), + ) wantDNSNames := []any{cert.GetName()} if cert.GetName() == "multigres-operator.test-cluster.supabase.multigres.internal" { wantDNSNames = []any{} @@ -281,24 +274,20 @@ func TestBuildInternalCertificates(t *testing.T) { "multipooler.test-cluster.supabase.multigres.internal", ) } - if diff := cmp.Diff(wantDNSNames, spec["dnsNames"]); diff != "" { - t.Errorf("dnsNames mismatch for %s (-want +got):\n%s", cert.GetName(), diff) - } + c.Eq( + "", + cmp.Diff(wantDNSNames, spec["dnsNames"]), + "dnsNames mismatch for %s (-want +got):\n", + cert.GetName(), + ) } changedExternalName := cluster.DeepCopy() changedExternalName.Spec.CertCommonName = "db.changed.supabase.red" gotAfterExternalNameChange, err := buildInternalCertificates(changedExternalName, scheme) - if err != nil { - t.Fatalf("buildInternalCertificates() after external name change: %v", err) - } - if len(gotAfterExternalNameChange) != len(got) { - t.Fatalf( - "got %d certs after external name change, want %d", - len(gotAfterExternalNameChange), - len(got), - ) - } + c.Require().NoError(err, "buildInternalCertificates() after external name change") + c.Require(). + Len(gotAfterExternalNameChange, len(got), "got %d certs after external name change, want", len(gotAfterExternalNameChange)) changedByName := make(map[string]*unstructured.Unstructured, len(gotAfterExternalNameChange)) for _, cert := range gotAfterExternalNameChange { @@ -310,13 +299,12 @@ func TestBuildInternalCertificates(t *testing.T) { if !ok { t.Fatalf("Certificate %q missing after external name change", originalCert.GetName()) } - if diff := cmp.Diff(originalCert, changedCert); diff != "" { - t.Errorf( - "internal Certificate %q changed with external CertCommonName (-want +got):\n%s", - originalCert.GetName(), - diff, - ) - } + c.EqDiff( + originalCert, + changedCert, + "internal Certificate %q changed with external CertCommonName", + originalCert.GetName(), + ) } } @@ -326,51 +314,37 @@ func TestTruncateCommonName(t *testing.T) { const longCNVariant = "multiadmin.mgc-iaogrkvrpaubkinljowm.ha-project-iaogrkvrpaubkinljowl.multigres.internal" t.Run("short CN is returned unchanged", func(t *testing.T) { - if len(shortCN) > maxCommonNameBytes { - t.Fatalf( - "test fixture shortCN is %d bytes, want <= %d", - len(shortCN), - maxCommonNameBytes, - ) - } + c := assert.NewCollecting(t) + c.Require().LessOrEqual(maxCommonNameBytes, len(shortCN), "test fixture shortCN is") got := truncateCommonName(shortCN) - if diff := cmp.Diff(shortCN, got); diff != "" { - t.Errorf("truncateCommonName() mismatch (-want +got):\n%s", diff) - } + c.EqDiff(shortCN, got, "truncateCommonName() mismatch") }) t.Run("long CN is truncated to the X.509 limit", func(t *testing.T) { - if len(longCN) <= maxCommonNameBytes { - t.Fatalf("test fixture longCN is %d bytes, want > %d", len(longCN), maxCommonNameBytes) - } + c := assert.NewCollecting(t) + c.Require().Greater(maxCommonNameBytes, len(longCN), "test fixture longCN is") got := truncateCommonName(longCN) - if len(got) > maxCommonNameBytes { - t.Errorf( - "truncateCommonName() = %q (%d bytes), want <= %d bytes", - got, - len(got), - maxCommonNameBytes, - ) - } + c.LessOrEqual( + maxCommonNameBytes, + len(got), + "truncateCommonName() = %q (%d bytes), want <= %d bytes", + got, + len(got), + maxCommonNameBytes, + ) }) t.Run("truncation is deterministic", func(t *testing.T) { first := truncateCommonName(longCN) second := truncateCommonName(longCN) - if diff := cmp.Diff(first, second); diff != "" { - t.Errorf("truncateCommonName() not deterministic (-first +second):\n%s", diff) - } + assert.NewCollecting(t). + Eq("", cmp.Diff(first, second), "truncateCommonName() not deterministic (-first +second):\n") }) t.Run("different long inputs produce different outputs", func(t *testing.T) { got := truncateCommonName(longCN) gotVariant := truncateCommonName(longCNVariant) - if got == gotVariant { - t.Errorf( - "truncateCommonName() collided: %q == %q for different inputs", - got, gotVariant, - ) - } + assert.NewCollecting(t).NotEq(gotVariant, got, "truncateCommonName() collided") }) } @@ -407,18 +381,15 @@ func TestReconcileCertificate(t *testing.T) { want map[string]string, ) { t.Helper() + c := assert.NewAborting(t) certificates := &unstructured.UnstructuredList{} certificates.SetGroupVersionKind(certGVK) - if err := fc.List( + c.NoError(fc.List( t.Context(), certificates, client.InNamespace(cluster.Namespace), - ); err != nil { - t.Fatalf("list Certificates: %v", err) - } - if len(certificates.Items) != len(want) { - t.Fatalf("got %d Certificates, want %d", len(certificates.Items), len(want)) - } + ), "list Certificates") + c.Len(certificates.Items, len(want), "got %d Certificates, want", len(certificates.Items)) seen := make(map[string]struct{}, len(certificates.Items)) for _, certificate := range certificates.Items { @@ -479,9 +450,8 @@ func TestReconcileCertificate(t *testing.T) { InternalTLS: &multigresv1alpha1.InternalTLSConfig{Enabled: ptr.To(true)}, }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcileCertificate(t.Context(), cluster), "unexpected error") assertCertificates(t, fc, cluster, wantInternalCertificates(cluster)) }) @@ -504,9 +474,8 @@ func TestReconcileCertificate(t *testing.T) { CertCommonName: "db.abc123.supabase.red", }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcileCertificate(t.Context(), cluster), "unexpected error") assertCertificates(t, fc, cluster, map[string]string{ cluster.Spec.CertCommonName: multigresv1alpha1.CertSecretName, @@ -529,9 +498,8 @@ func TestReconcileCertificate(t *testing.T) { InternalTLS: &multigresv1alpha1.InternalTLSConfig{Enabled: ptr.To(false)}, }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcileCertificate(t.Context(), cluster), "unexpected error") assertCertificates(t, fc, cluster, map[string]string{}) }, ) @@ -555,9 +523,8 @@ func TestReconcileCertificate(t *testing.T) { CertCommonName: "db.both.supabase.red", }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcileCertificate(t.Context(), cluster), "unexpected error") want := wantInternalCertificates(cluster) want[cluster.Spec.CertCommonName] = multigresv1alpha1.CertSecretName @@ -567,6 +534,7 @@ func TestReconcileCertificate(t *testing.T) { t.Run( "disabling internal TLS deletes internal Certificates but keeps public", func(t *testing.T) { + c := assert.NewAborting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -586,17 +554,16 @@ func TestReconcileCertificate(t *testing.T) { CertCommonName: "db.toggle.supabase.red", }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("create internal and public Certificates: %v", err) - } + c.NoError( + r.reconcileCertificate(t.Context(), cluster), + "create internal and public Certificates", + ) want := wantInternalCertificates(cluster) want[cluster.Spec.CertCommonName] = multigresv1alpha1.CertSecretName assertCertificates(t, fc, cluster, want) cluster.Spec.InternalTLS.Enabled = ptr.To(false) - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("disable internal TLS: %v", err) - } + c.NoError(r.reconcileCertificate(t.Context(), cluster), "disable internal TLS") assertCertificates(t, fc, cluster, map[string]string{ cluster.Spec.CertCommonName: multigresv1alpha1.CertSecretName, }) @@ -604,6 +571,7 @@ func TestReconcileCertificate(t *testing.T) { ) t.Run("idempotent on repeated calls", func(t *testing.T) { + c := assert.NewAborting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -622,15 +590,12 @@ func TestReconcileCertificate(t *testing.T) { CertCommonName: "db.xyz.supabase.red", }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("first call: %v", err) - } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("second call: %v", err) - } + c.NoError(r.reconcileCertificate(t.Context(), cluster), "first call") + c.NoError(r.reconcileCertificate(t.Context(), cluster), "second call") }) t.Run("no Patch when nothing changed", func(t *testing.T) { + c := assert.NewCollecting(t) var patchCount int fc := fake.NewClientBuilder(). WithScheme(scheme). @@ -666,27 +631,17 @@ func TestReconcileCertificate(t *testing.T) { }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("first reconcile: %v", err) - } - if patchCount != 6 { - t.Fatalf("patchCount after first reconcile = %d, want 6", patchCount) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "first reconcile") + c.Require().Eq(6, patchCount, "patchCount after first reconcile") // Reconciling again with the same spec should not re-patch any // Certificate since the live specs already match desired. - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("second reconcile: %v", err) - } - if patchCount != 6 { - t.Errorf( - "patchCount after second reconcile = %d, want 6 (no new patches)", - patchCount, - ) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "second reconcile") + c.Eq(6, patchCount, "patchCount after second reconcile") }) t.Run("CN change updates Certificate", func(t *testing.T) { + c := assert.NewCollecting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -707,36 +662,29 @@ func TestReconcileCertificate(t *testing.T) { } // Create with old CN - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("create old: %v", err) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "create old") // Change CN cluster.Spec.CertCommonName = "db.new.supabase.red" - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("create new: %v", err) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "create new") // New cert should exist got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(fc.Get(t.Context(), types.NamespacedName{ Name: "db.new.supabase.red", Namespace: "default", - }, got); err != nil { - t.Fatalf("new Certificate should exist: %v", err) - } + }, got), "new Certificate should exist") // Old cert should be deleted by reconcileCertificate old := &unstructured.Unstructured{} old.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Error(fc.Get(t.Context(), types.NamespacedName{ Name: "db.old.supabase.red", Namespace: "default", - }, old); err == nil { - t.Error("old Certificate should be deleted on CN change") - } + }, old), "old Certificate should be deleted on CN change") }) t.Run("CN unset cleans up Certificate", func(t *testing.T) { + c := assert.NewCollecting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -757,37 +705,30 @@ func TestReconcileCertificate(t *testing.T) { } // Create cert - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("create: %v", err) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "create") // Verify it exists got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(fc.Get(t.Context(), types.NamespacedName{ Name: "db.cleanup.supabase.red", Namespace: "default", - }, got); err != nil { - t.Fatalf("Certificate should exist before cleanup: %v", err) - } + }, got), "Certificate should exist before cleanup") // Unset CN and reconcile cluster.Spec.CertCommonName = "" - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("cleanup: %v", err) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "cleanup") // Cert should be deleted err := fc.Get(t.Context(), types.NamespacedName{ Name: "db.cleanup.supabase.red", Namespace: "default", }, got) - if err == nil { - t.Error("Certificate should be deleted after unsetting CN") - } + c.Error(err, "Certificate should be deleted after unsetting CN") }) t.Run( "cleanup ignores certs not owned by this cluster", func(t *testing.T) { + c := assert.NewAborting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -811,9 +752,7 @@ func TestReconcileCertificate(t *testing.T) { other.Object["spec"] = map[string]any{ "secretName": multigresv1alpha1.CertSecretName, } - if err := fc.Create(t.Context(), other); err != nil { - t.Fatalf("failed to create other cert: %v", err) - } + c.NoError(fc.Create(t.Context(), other), "failed to create other cert") // Our cluster has no CN — should not delete the other cert cluster := &multigresv1alpha1.MultigresCluster{ @@ -821,24 +760,21 @@ func TestReconcileCertificate(t *testing.T) { Name: "c6", Namespace: "default", UID: "uid-6", }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("cleanup: %v", err) - } + c.NoError(r.reconcileCertificate(t.Context(), cluster), "cleanup") // Other cert should still exist got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.NoError(fc.Get(t.Context(), types.NamespacedName{ Name: "db.other.supabase.red", Namespace: "default", - }, got); err != nil { - t.Fatal("Certificate owned by another cluster should survive") - } + }, got), "Certificate owned by another cluster should survive") }, ) t.Run( "legacy internal identity retries cleanup after Secret deletion fails", func(t *testing.T) { + c := assert.NewCollecting(t) oldCertificateName := "multipooler.db.secretold.supabase.red" oldSecretName := oldCertificateName deleteFailure := errors.New("injected Secret deletion failure") @@ -885,58 +821,42 @@ func TestReconcileCertificate(t *testing.T) { dnsNames: []any{oldCertificateName}, usages: []any{"server auth", "client auth"}, }) - if err != nil { - t.Fatalf("build legacy Certificate: %v", err) - } - if err := fc.Create(t.Context(), legacyCertificate); err != nil { - t.Fatalf("create legacy Certificate: %v", err) - } + c.Require().NoError(err, "build legacy Certificate") + c.Require(). + NoError(fc.Create(t.Context(), legacyCertificate), "create legacy Certificate") oldSecret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{ Name: oldSecretName, Namespace: "default", }, } - if err := fc.Create(t.Context(), oldSecret); err != nil { - t.Fatalf("failed to create old secret: %v", err) - } + c.Require().NoError(fc.Create(t.Context(), oldSecret), "failed to create old secret") failSecretDelete = true err = r.reconcileCertificate(t.Context(), cluster) - if !errors.Is(err, deleteFailure) { - t.Fatalf("first cleanup error = %v, want wrapped deletion failure", err) - } + c.Require().ErrorIs(err, deleteFailure, "first cleanup error") staleCertificate := &unstructured.Unstructured{} staleCertificate.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.NoError(fc.Get(t.Context(), types.NamespacedName{ Name: oldCertificateName, Namespace: "default", - }, staleCertificate); err != nil { - t.Errorf("stale Certificate should remain after Secret deletion failure: %v", err) - } - if err := fc.Get(t.Context(), types.NamespacedName{ + }, staleCertificate), "stale Certificate should remain after Secret deletion failure") + c.NoError(fc.Get(t.Context(), types.NamespacedName{ Name: oldSecretName, Namespace: "default", - }, &corev1.Secret{}); err != nil { - t.Errorf("stale Secret should remain after deletion failure: %v", err) - } + }, &corev1.Secret{}), "stale Secret should remain after deletion failure") - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("cleanup retry: %v", err) - } - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "cleanup retry") + c.Error(fc.Get(t.Context(), types.NamespacedName{ Name: oldCertificateName, Namespace: "default", - }, staleCertificate); err == nil { - t.Error("stale Certificate should be deleted on cleanup retry") - } - if err := fc.Get(t.Context(), types.NamespacedName{ + }, staleCertificate), "stale Certificate should be deleted on cleanup retry") + c.Error(fc.Get(t.Context(), types.NamespacedName{ Name: oldSecretName, Namespace: "default", - }, &corev1.Secret{}); err == nil { - t.Error("stale internal Secret should be deleted on cleanup retry") - } + }, &corev1.Secret{}), "stale internal Secret should be deleted on cleanup retry") }, ) t.Run("CN change keeps the shared external Secret", func(t *testing.T) { + c := assert.NewCollecting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -955,9 +875,7 @@ func TestReconcileCertificate(t *testing.T) { CertCommonName: "db.sharedold.supabase.red", }, } - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("create old: %v", err) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "create old") // The external gateway cert always uses the fixed secret name, so it // must survive the CN rotation for the new Certificate to reuse it. @@ -967,25 +885,20 @@ func TestReconcileCertificate(t *testing.T) { Namespace: "default", }, } - if err := fc.Create(t.Context(), sharedSecret); err != nil { - t.Fatalf("failed to create shared secret: %v", err) - } + c.Require().NoError(fc.Create(t.Context(), sharedSecret), "failed to create shared secret") cluster.Spec.CertCommonName = "db.sharednew.supabase.red" - if err := r.reconcileCertificate(t.Context(), cluster); err != nil { - t.Fatalf("create new: %v", err) - } + c.Require().NoError(r.reconcileCertificate(t.Context(), cluster), "create new") - if err := fc.Get(t.Context(), types.NamespacedName{ + c.NoError(fc.Get(t.Context(), types.NamespacedName{ Name: multigresv1alpha1.CertSecretName, Namespace: "default", - }, &corev1.Secret{}); err != nil { - t.Errorf("shared external Secret should survive CN rotation: %v", err) - } + }, &corev1.Secret{}), "shared external Secret should survive CN rotation") }) t.Run( "reports a collision for a Certificate with matching spec but no ownerRef", func(t *testing.T) { + c := assert.NewCollecting(t) fc := fake.NewClientBuilder().WithScheme(scheme).Build() r := &MultigresClusterReconciler{ Client: fc, @@ -1006,9 +919,7 @@ func TestReconcileCertificate(t *testing.T) { } desired, err := buildCertificate(cluster, scheme) - if err != nil { - t.Fatalf("buildCertificate: %v", err) - } + c.Require().NoError(err, "buildCertificate") // Pre-create a Certificate with a matching spec but no // ownerRef, simulating one left unmanaged by a prior bug. unowned := &unstructured.Unstructured{} @@ -1016,30 +927,24 @@ func TestReconcileCertificate(t *testing.T) { unowned.SetName(desired.GetName()) unowned.SetNamespace("default") unowned.Object["spec"] = desired.Object["spec"] - if err := fc.Create(t.Context(), unowned); err != nil { - t.Fatalf("failed to create unowned cert: %v", err) - } + c.Require().NoError(fc.Create(t.Context(), unowned), "failed to create unowned cert") - if err := r.reconcileCertificate(t.Context(), cluster); err == nil { - t.Fatal("reconcile: got nil error, want collision error") - } + c.Require(). + Error(r.reconcileCertificate(t.Context(), cluster), "reconcile: got nil error, want collision error") got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(fc.Get(t.Context(), types.NamespacedName{ Name: desired.GetName(), Namespace: "default", - }, got); err != nil { - t.Fatalf("Certificate should exist: %v", err) - } - if len(got.GetOwnerReferences()) != 0 { - t.Error("foreign Certificate should not be adopted") - } + }, got), "Certificate should exist") + c.Empty(got.GetOwnerReferences(), "foreign Certificate should not be adopted") }, ) t.Run( "two clusters in same namespace get independent certs", func(t *testing.T) { + c := assert.NewCollecting(t) // Regression: the original cell-level architecture had // multiple cells fighting over the same Certificate with // ownerRef flipping. Moving to the cluster controller @@ -1082,16 +987,12 @@ func TestReconcileCertificate(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := rA.reconcileCertificate( + c.Require().NoError(rA.reconcileCertificate( t.Context(), clusterA, - ); err != nil { - t.Fatalf("cluster-a reconcile: %v", err) - } - if err := rB.reconcileCertificate( + ), "cluster-a reconcile") + c.Require().NoError(rB.reconcileCertificate( t.Context(), clusterB, - ); err != nil { - t.Fatalf("cluster-b reconcile: %v", err) - } + ), "cluster-b reconcile") // Both certs exist for _, name := range []string{ @@ -1100,68 +1001,44 @@ func TestReconcileCertificate(t *testing.T) { } { got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.NoError(fc.Get(t.Context(), types.NamespacedName{ Name: name, Namespace: "default", - }, got); err != nil { - t.Errorf("Certificate %q should exist: %v", name, err) - } + }, got), "Certificate %q should exist", name) } // Each cert is owned by the correct cluster certA := &unstructured.Unstructured{} certA.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(fc.Get(t.Context(), types.NamespacedName{ Name: "db.projA.supabase.red", Namespace: "default", - }, certA); err != nil { - t.Fatalf("failed to get certA: %v", err) - } - if certA.GetOwnerReferences()[0].UID != "uid-a" { - t.Errorf( - "certA owner UID = %q, want uid-a", - certA.GetOwnerReferences()[0].UID, - ) - } + }, certA), "failed to get certA") + c.Eq("uid-a", certA.GetOwnerReferences()[0].UID, "certA owner UID") certB := &unstructured.Unstructured{} certB.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Require().NoError(fc.Get(t.Context(), types.NamespacedName{ Name: "db.projB.supabase.red", Namespace: "default", - }, certB); err != nil { - t.Fatalf("failed to get certB: %v", err) - } - if certB.GetOwnerReferences()[0].UID != "uid-b" { - t.Errorf( - "certB owner UID = %q, want uid-b", - certB.GetOwnerReferences()[0].UID, - ) - } + }, certB), "failed to get certB") + c.Eq("uid-b", certB.GetOwnerReferences()[0].UID, "certB owner UID") // Unsetting CN on cluster-a only deletes its cert clusterA.Spec.CertCommonName = "" - if err := rA.reconcileCertificate( + c.Require().NoError(rA.reconcileCertificate( t.Context(), clusterA, - ); err != nil { - t.Fatalf("cluster-a cleanup: %v", err) - } + ), "cluster-a cleanup") gone := &unstructured.Unstructured{} gone.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.Error(fc.Get(t.Context(), types.NamespacedName{ Name: "db.projA.supabase.red", Namespace: "default", - }, gone); err == nil { - t.Error("cluster-a cert should be deleted") - } + }, gone), "cluster-a cert should be deleted") // cluster-b cert is untouched still := &unstructured.Unstructured{} still.SetGroupVersionKind(certGVK) - if err := fc.Get(t.Context(), types.NamespacedName{ + c.NoError(fc.Get(t.Context(), types.NamespacedName{ Name: "db.projB.supabase.red", Namespace: "default", - }, still); err != nil { - t.Errorf( - "cluster-b cert should survive: %v", err, - ) - } + }, still), "cluster-b cert should survive") }, ) } diff --git a/pkg/cluster-handler/controller/multigrescluster/certificate_topo_test.go b/pkg/cluster-handler/controller/multigrescluster/certificate_topo_test.go index e22e57dc..03db82c8 100644 --- a/pkg/cluster-handler/controller/multigrescluster/certificate_topo_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/certificate_topo_test.go @@ -8,8 +8,6 @@ import ( "github.com/google/go-cmp/cmp" "github.com/multigres/multigres/go/common/topoclient" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" "k8s.io/apimachinery/pkg/api/meta" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/apis/meta/v1/unstructured" @@ -22,9 +20,12 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/resolver" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestTopologyTLSLongFallbackUsesSameRoot(t *testing.T) { + ck := assert.NewCollecting(t) cluster := topoTLSCluster("cluster-abcdefghijklmnop", "namespace-abcdefghijklmnopqrstu", nil) cluster.Spec.Cells = []multigresv1alpha1.CellConfig{{Name: "zone-a"}} scheme := setupScheme() @@ -33,33 +34,32 @@ func TestTopologyTLSLongFallbackUsesSameRoot(t *testing.T) { r := &MultigresClusterReconciler{ Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10), CreateTopoStore: func(ref multigresv1alpha1.GlobalTopoServerRef) (topoclient.Store, error) { - assert.Equal( - t, + ck.EqDeep( "/multigres-fallback/b-Tmo_r9oWzDWEuz_6f6LFOAmYO7ve1i4ksIy7qa9ac/global", ref.RootPath, ) return noCloseStore{Store: store}, nil }, } - require.NoError(t, r.reconcileCertificate(t.Context(), cluster)) + ck.Require().NoError(r.reconcileCertificate(t.Context(), cluster)) cert := &unstructured.Unstructured{} cert.SetGroupVersionKind(certGVK) - require.NoError(t, c.Get(t.Context(), client.ObjectKey{ + ck.Require().NoError(c.Get(t.Context(), client.ObjectKey{ Namespace: cluster.Namespace, Name: multigresv1alpha1.TopoClientCertName(cluster.Name), }, cert)) subject, _, err := unstructured.NestedString(cert.Object, "spec", "literalSubject") - require.NoError(t, err) + ck.Require().NoError(err) cn, ok := parseCommonName(subject) - require.True(t, ok) - assert.Equal(t, "/multigres-fallback/b-Tmo_r9oWzDWEuz_6f6LFOAmYO7ve1i4ksIy7qa9ac", cn) + ck.Require().True(ok) + ck.EqDeep("/multigres-fallback/b-Tmo_r9oWzDWEuz_6f6LFOAmYO7ve1i4ksIy7qa9ac", cn) res := resolver.NewResolver(c, cluster.Namespace) result, err := r.reconcileTopology(t.Context(), cluster, res) - require.NoError(t, err) - assert.Zero(t, result.RequeueAfter) + ck.Require().NoError(err) + ck.Zero(result.RequeueAfter) cell, err := store.GetCell(t.Context(), "zone-a") - require.NoError(t, err) - assert.Equal(t, cn+"/global", cell.Root) + ck.Require().NoError(err) + ck.EqDeep(cn+"/global", cell.Root) cluster.Spec.Cells[0].Spec = &multigresv1alpha1.CellInlineSpec{ LocalTopoServer: &multigresv1alpha1.LocalTopoServerSpec{ @@ -67,13 +67,14 @@ func TestTopologyTLSLongFallbackUsesSameRoot(t *testing.T) { }, } _, _, local, err := res.ResolveCell(t.Context(), cluster, &cluster.Spec.Cells[0]) - require.NoError(t, err) - assert.Equal(t, cn+"/zone-a", local.Etcd.RootPath) + ck.Require().NoError(err) + ck.EqDeep(cn+"/zone-a", local.Etcd.RootPath) } func TestReconcileOversizedProjectRefStatusAndRecovery(t *testing.T) { for _, ref := range []string{strings.Repeat("p", 55), strings.Repeat("/", 19)} { t.Run(ref, func(t *testing.T) { + ck := assert.NewCollecting(t) cluster := topoTLSCluster( "cluster", "default", @@ -93,35 +94,35 @@ func TestReconcileOversizedProjectRefStatusAndRecovery(t *testing.T) { } key := client.ObjectKeyFromObject(cluster) _, err := r.Reconcile(t.Context(), ctrl.Request{NamespacedName: key}) - require.ErrorContains(t, err, "64 byte certificate common name limit") - require.NoError(t, c.Get(t.Context(), key, cluster)) - assert.Equal(t, multigresv1alpha1.PhaseDegraded, cluster.Status.Phase) + ck.Require().ErrorContains(err, "64 byte certificate common name limit") + ck.Require().NoError(c.Get(t.Context(), key, cluster)) + ck.EqDeep(multigresv1alpha1.PhaseDegraded, cluster.Status.Phase) condition := meta.FindStatusCondition(cluster.Status.Conditions, conditionTopologyReady) - require.NotNil(t, condition) - assert.Equal(t, metav1.ConditionFalse, condition.Status) - assert.Equal(t, "TopoCertificateFailed", condition.Reason) - assert.Equal(t, cluster.Generation, condition.ObservedGeneration) - assert.Contains( - t, + ck.Require().NotNil(condition) + ck.EqDeep(metav1.ConditionFalse, condition.Status) + ck.EqDeep("TopoCertificateFailed", condition.Reason) + ck.EqDeep(cluster.Generation, condition.ObservedGeneration) + ck.StrContains( condition.Message, "shorten annotation multigres.com/project-ref to at most 53 bytes after path escaping", ) cluster.Annotations[metadata.AnnotationProjectRef] = strings.Repeat("p", 53) - require.NoError(t, c.Update(t.Context(), cluster)) + ck.Require().NoError(c.Update(t.Context(), cluster)) _, err = r.Reconcile(t.Context(), ctrl.Request{NamespacedName: key}) - require.NoError(t, err) - require.NoError(t, c.Get(t.Context(), key, cluster)) + ck.Require().NoError(err) + ck.Require().NoError(c.Get(t.Context(), key, cluster)) condition = meta.FindStatusCondition(cluster.Status.Conditions, conditionTopologyReady) - require.NotNil(t, condition) - assert.Equal(t, metav1.ConditionTrue, condition.Status) - assert.Equal(t, "TopoConnected", condition.Reason) - assert.NotContains(t, cluster.Status.Message, "certificate common name limit") + ck.Require().NotNil(condition) + ck.EqDeep(metav1.ConditionTrue, condition.Status) + ck.EqDeep("TopoConnected", condition.Reason) + ck.NotStrContains(cluster.Status.Message, "certificate common name limit") }) } } func TestReconcileCertificatePreservesExplicitLongRoot(t *testing.T) { + ck := assert.NewCollecting(t) cluster := topoTLSCluster("cluster-abcdefghijklmnop", "namespace-abcdefghijklmnopqrstu", nil) const explicitRoot = "/multigres/namespace-abcdefghijklmnopqrstu/cluster-abcdefghijklmnop/global" cluster.Spec.GlobalTopoServer = &multigresv1alpha1.GlobalTopoServerSpec{ @@ -138,17 +139,14 @@ func TestReconcileCertificatePreservesExplicitLongRoot(t *testing.T) { Scheme: scheme, Recorder: record.NewFakeRecorder(10), } - require.ErrorContains( - t, - r.reconcileCertificate(t.Context(), cluster), - "outside certificate identity", - ) - require.NoError(t, c.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) - assert.Equal(t, explicitRoot, cluster.Spec.GlobalTopoServer.Etcd.RootPath) + ck.Require(). + ErrorContains(r.reconcileCertificate(t.Context(), cluster), "outside certificate identity") + ck.Require().NoError(c.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) + ck.EqDeep(explicitRoot, cluster.Spec.GlobalTopoServer.Etcd.RootPath) condition := meta.FindStatusCondition(cluster.Status.Conditions, conditionTopologyReady) - require.NotNil(t, condition) - assert.Equal(t, "TopoCertificateFailed", condition.Reason) - assert.Contains(t, condition.Message, "migrate any existing topology data") + ck.Require().NotNil(condition) + ck.EqDeep("TopoCertificateFailed", condition.Reason) + ck.StrContains(condition.Message, "migrate any existing topology data") cluster.Spec.GlobalTopoServer.Etcd.RootPath = "" cluster.Spec.Cells = []multigresv1alpha1.CellConfig{{ @@ -159,12 +157,9 @@ func TestReconcileCertificatePreservesExplicitLongRoot(t *testing.T) { }, }, }} - require.NoError(t, c.Update(t.Context(), cluster)) - require.ErrorContains( - t, - r.reconcileCertificate(t.Context(), cluster), - `cell "zone-a" topology root`, - ) + ck.Require().NoError(c.Update(t.Context(), cluster)) + ck.Require(). + ErrorContains(r.reconcileCertificate(t.Context(), cluster), `cell "zone-a" topology root`) } func topoTLSCluster( @@ -225,48 +220,36 @@ func TestBuildTopoClientCertificate(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() got, err := buildTopoClientCertificate(tc.cluster, scheme) - if err != nil { - t.Fatalf("buildTopoClientCertificate() error = %v", err) - } + c.Require().NoError(err, "buildTopoClientCertificate() error =") wantName := tc.cluster.Name + "-topo-client-tls" - if got.GetName() != wantName { - t.Errorf("name = %q, want %q", got.GetName(), wantName) - } - if got.GetNamespace() != tc.cluster.Namespace { - t.Errorf("namespace = %q, want %q", got.GetNamespace(), tc.cluster.Namespace) - } + c.Eq(wantName, got.GetName(), "name") + c.Eq(tc.cluster.Namespace, got.GetNamespace(), "namespace") ownerRefs := got.GetOwnerReferences() - if len(ownerRefs) != 1 || ownerRefs[0].Kind != "MultigresCluster" { - t.Fatalf("ownerReferences = %+v, want one MultigresCluster ref", ownerRefs) - } + c.Require(). + False(len(ownerRefs) != 1 || ownerRefs[0].Kind != "MultigresCluster", "ownerReferences = %+v, want one MultigresCluster ref", ownerRefs) spec, ok := got.Object["spec"].(map[string]any) - if !ok { - t.Fatal("spec is not a map") - } + c.Require().True(ok, "spec is not a map") wantSubject := fmt.Sprintf(CertLiteralSubjectTemplate, tc.wantSubject) - if diff := cmp.Diff(wantSubject, spec["literalSubject"]); diff != "" { - t.Errorf("literalSubject mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(wantName, spec["secretName"]); diff != "" { - t.Errorf("secretName mismatch (-want +got):\n%s", diff) - } + c.Eq( + "", + cmp.Diff(wantSubject, spec["literalSubject"]), + "literalSubject mismatch (-want +got):\n", + ) + c.Eq("", cmp.Diff(wantName, spec["secretName"]), "secretName mismatch (-want +got):\n") // A client credential is verified by subject, so it carries no SANs. - if diff := cmp.Diff([]any{}, spec["dnsNames"]); diff != "" { - t.Errorf("dnsNames mismatch (-want +got):\n%s", diff) - } + c.Eq("", cmp.Diff([]any{}, spec["dnsNames"]), "dnsNames mismatch (-want +got):\n") wantUsages := []any{ "digital signature", "key encipherment", "client auth", } - if diff := cmp.Diff(wantUsages, spec["usages"]); diff != "" { - t.Errorf("usages mismatch (-want +got):\n%s", diff) - } + c.Eq("", cmp.Diff(wantUsages, spec["usages"]), "usages mismatch (-want +got):\n") // The credential is only useful if it chains to the CA the // topology server trusts, never to the cluster's own issuer. wantIssuerRef := map[string]any{ @@ -274,9 +257,11 @@ func TestBuildTopoClientCertificate(t *testing.T) { "kind": "ClusterIssuer", "group": "cert-manager.io", } - if diff := cmp.Diff(wantIssuerRef, spec["issuerRef"]); diff != "" { - t.Errorf("issuerRef mismatch (-want +got):\n%s", diff) - } + c.Eq( + "", + cmp.Diff(wantIssuerRef, spec["issuerRef"]), + "issuerRef mismatch (-want +got):\n", + ) }) } } @@ -284,25 +269,21 @@ func TestBuildTopoClientCertificate(t *testing.T) { // The CN is the identity a topology server authorizes against, so it has to // be the string ClusterRoot() produces and not a re-derivation of it. func TestBuildTopoClientCertificateCommonNameMatchesTopologyRoot(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() cluster := topoTLSCluster("test-cluster", "supabase", map[string]string{ metadata.AnnotationProjectRef: "proj_123", }) cert, err := buildTopoClientCertificate(cluster, scheme) - if err != nil { - t.Fatalf("buildTopoClientCertificate() error = %v", err) - } + c.Require().NoError(err, "buildTopoClientCertificate() error =") subject, _, _ := unstructured.NestedString(cert.Object, "spec", "literalSubject") globalRoot := "/multigres/proj_123/global" cn, ok := parseCommonName(subject) - if !ok { - t.Fatalf("no CN in literalSubject %q", subject) - } - if got := cn + "/global"; got != globalRoot { - t.Errorf("CN %q does not prefix the global root: got %q, want %q", cn, got, globalRoot) - } + c.Require().True(ok, "no CN in literalSubject %q", subject) + got := cn + "/global" + c.Eq(globalRoot, got, "CN %q does not prefix the global root: got %q, want", cn, got) } func parseCommonName(subject string) (string, bool) { @@ -318,6 +299,7 @@ func TestReconcileCertificateTopoTLS(t *testing.T) { certName := "test-cluster-topo-client-tls" t.Run("issues the client certificate when topology TLS is enabled", func(t *testing.T) { + ck := assert.NewAborting(t) scheme := setupScheme() cluster := topoTLSCluster("test-cluster", "supabase", nil) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(cluster).Build() @@ -327,19 +309,22 @@ func TestReconcileCertificateTopoTLS(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), cluster); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.NoError( + r.reconcileCertificate(context.Background(), cluster), + "reconcileCertificate() error =", + ) got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) key := client.ObjectKey{Namespace: "supabase", Name: certName} - if err := c.Get(context.Background(), key, got); err != nil { - t.Fatalf("expected topology client Certificate, got error %v", err) - } + ck.NoError( + c.Get(context.Background(), key, got), + "expected topology client Certificate, got error", + ) }) t.Run("prunes the client certificate when topology TLS is disabled", func(t *testing.T) { + ck := assert.NewAborting(t) scheme := setupScheme() cluster := topoTLSCluster("test-cluster", "supabase", nil) cluster.Spec.TopoTLS.Enabled = ptr.To(false) @@ -347,9 +332,7 @@ func TestReconcileCertificateTopoTLS(t *testing.T) { existing, err := buildTopoClientCertificate( topoTLSCluster("test-cluster", "supabase", nil), scheme, ) - if err != nil { - t.Fatalf("buildTopoClientCertificate() error = %v", err) - } + ck.NoError(err, "buildTopoClientCertificate() error =") c := fake.NewClientBuilder(). WithScheme(scheme). WithObjects(cluster, existing). @@ -360,20 +343,20 @@ func TestReconcileCertificateTopoTLS(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), cluster); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.NoError( + r.reconcileCertificate(context.Background(), cluster), + "reconcileCertificate() error =", + ) got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) key := client.ObjectKey{Namespace: "supabase", Name: certName} err = c.Get(context.Background(), key, got) - if err == nil { - t.Fatal("expected topology client Certificate to be deleted") - } + ck.Error(err, "expected topology client Certificate to be deleted") }) t.Run("issues nothing when the topology server is external", func(t *testing.T) { + ck := assert.NewAborting(t) scheme := setupScheme() cluster := topoTLSCluster("test-cluster", "supabase", nil) // An external topology server brings its own CA and client Secrets, so @@ -390,19 +373,22 @@ func TestReconcileCertificateTopoTLS(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), cluster); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.NoError( + r.reconcileCertificate(context.Background(), cluster), + "reconcileCertificate() error =", + ) got := &unstructured.Unstructured{} got.SetGroupVersionKind(certGVK) key := client.ObjectKey{Namespace: "supabase", Name: certName} - if err := c.Get(context.Background(), key, got); err == nil { - t.Fatal("expected no topology client Certificate for an external topology server") - } + ck.Error( + c.Get(context.Background(), key, got), + "expected no topology client Certificate for an external topology server", + ) }) t.Run("issues nothing when topology TLS is unset", func(t *testing.T) { + ck := assert.NewCollecting(t) scheme := setupScheme() cluster := topoTLSCluster("test-cluster", "supabase", nil) cluster.Spec.TopoTLS = nil @@ -413,18 +399,13 @@ func TestReconcileCertificateTopoTLS(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), cluster); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.Require(). + NoError(r.reconcileCertificate(context.Background(), cluster), "reconcileCertificate() error =") list := &unstructured.UnstructuredList{} list.SetGroupVersionKind(certGVK) - if err := c.List(context.Background(), list); err != nil { - t.Fatalf("List() error = %v", err) - } - if len(list.Items) != 0 { - t.Errorf("got %d Certificates, want 0", len(list.Items)) - } + ck.Require().NoError(c.List(context.Background(), list), "List() error =") + ck.Empty(list.Items, "got %d Certificates, want 0", len(list.Items)) }) } @@ -436,26 +417,29 @@ func TestBuildTopoClientCertificateRejectsOverLongRoot(t *testing.T) { metadata.AnnotationProjectRef: strings.Repeat("p", 64), }) - if _, err := buildTopoClientCertificate(cluster, scheme); err == nil { - t.Fatal("buildTopoClientCertificate() = nil error, want a common name limit error") - } + _, err := buildTopoClientCertificate(cluster, scheme) + assert.NewAborting(t). + Error(err, "buildTopoClientCertificate() = nil error, want a common name limit error") } // Internal component certificates keep using the cluster's own issuer; only // the topology credential follows the topology CA. func TestInternalCertificatesKeepClusterIssuer(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() cluster := topoTLSCluster("test-cluster", "supabase", nil) cluster.Spec.InternalTLS = &multigresv1alpha1.InternalTLSConfig{Enabled: ptr.To(true)} built, err := buildInternalCertificates(cluster, scheme) - if err != nil { - t.Fatalf("buildInternalCertificates() error = %v", err) - } + c.Require().NoError(err, "buildInternalCertificates() error =") for _, cert := range built { issuer, _, _ := unstructured.NestedString(cert.Object, "spec", "issuerRef", "name") - if issuer != "cluster-issuer" { - t.Errorf("%s issuerRef.name = %q, want cluster-issuer", cert.GetName(), issuer) - } + c.Eq( + "cluster-issuer", + issuer, + "%s issuerRef.name = %q, want cluster-issuer", + cert.GetName(), + issuer, + ) } } diff --git a/pkg/cluster-handler/controller/multigrescluster/images_test.go b/pkg/cluster-handler/controller/multigrescluster/images_test.go index 663f75bf..6b5193bb 100644 --- a/pkg/cluster-handler/controller/multigrescluster/images_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/images_test.go @@ -20,6 +20,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/images" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) type patchCountingClient struct { @@ -81,9 +83,7 @@ func olderApplied() multigresv1alpha1.ComponentImages { func mustJSON(t *testing.T, set multigresv1alpha1.ComponentImages) string { t.Helper() raw, err := json.Marshal(set) - if err != nil { - t.Fatal(err) - } + assert.NewAborting(t).NoError(err) return string(raw) } @@ -155,18 +155,15 @@ func TestResolveImages(t *testing.T) { } t.Run("immediate strategy adopts current defaults", func(t *testing.T) { + c := assert.NewCollecting(t) old := olderApplied() cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: mustJSON(t, old), }, nil) r := newHarness(t, images.UpdateImmediate, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v2" { - t.Errorf("expected current default, got %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v2", cluster.Spec.Images.Postgres, "expected current default, got") if cluster.Status.Images.Source != multigresv1alpha1.ImageSourceDefaults || cluster.Status.Images.Effective != testImagesConfig(images.UpdateImmediate).Defaults { t.Errorf("unexpected effective image status: %+v", cluster.Status.Images) @@ -177,21 +174,20 @@ func TestResolveImages(t *testing.T) { }) t.Run("lazy strategy holds recorded set without acknowledgement", func(t *testing.T) { + c := assert.NewCollecting(t) old := olderApplied() cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: mustJSON(t, old), }, nil) r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v1" { - t.Errorf("expected held image, got %s", cluster.Spec.Images.Postgres) - } - if cluster.Status.Images.AppliedRevision == cluster.Status.Images.AvailableRevision { - t.Error("expected a pending rollout (applied != available)") - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v1", cluster.Spec.Images.Postgres, "expected held image, got") + c.NotEq( + cluster.Status.Images.AvailableRevision, + cluster.Status.Images.AppliedRevision, + "expected a pending rollout (applied != available)", + ) if cluster.Status.Images.Applied != old || cluster.Status.Images.Available != testImagesConfig(images.UpdateLazy).Defaults { t.Errorf( @@ -199,18 +195,18 @@ func TestResolveImages(t *testing.T) { cluster.Status.Images, ) } - if cluster.Status.Images.UpdateStrategy != string(images.UpdateLazy) { - t.Errorf("status must report the running strategy, got %q", - cluster.Status.Images.UpdateStrategy) - } + c.Eq( + string(images.UpdateLazy), + cluster.Status.Images.UpdateStrategy, + "status must report the running strategy, got", + ) cond := pendingCond(cluster) - if cond == nil || cond.Status != metav1.ConditionTrue || - cond.Reason != multigresv1alpha1.ReasonAwaitingAcknowledgement { - t.Errorf("expected ImageRolloutPending=True/AwaitingAcknowledgement, got %+v", cond) - } + c.False(cond == nil || cond.Status != metav1.ConditionTrue || + cond.Reason != multigresv1alpha1.ReasonAwaitingAcknowledgement, "expected ImageRolloutPending=True/AwaitingAcknowledgement, got %+v", cond) }) t.Run("lazy strategy survives status loss via annotation", func(t *testing.T) { + c := assert.NewCollecting(t) // Status wiped (backup/restore, kubectl replace); only the // annotation remains. The cluster must NOT roll to new defaults. old := olderApplied() @@ -219,15 +215,12 @@ func TestResolveImages(t *testing.T) { }, nil) r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v1" { - t.Errorf("status loss rolled the cluster: %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v1", cluster.Spec.Images.Postgres, "status loss rolled the cluster") }) t.Run("lazy strategy adopts when acknowledged revision matches", func(t *testing.T) { + c := assert.NewCollecting(t) old := olderApplied() cfg := testImagesConfig(images.UpdateLazy) cluster := newCluster(map[string]string{ @@ -238,27 +231,23 @@ func TestResolveImages(t *testing.T) { } r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v2" { - t.Errorf("expected adopted image, got %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v2", cluster.Spec.Images.Postgres, "expected adopted image, got") if cond := pendingCond(cluster); cond == nil || cond.Status != metav1.ConditionFalse { t.Errorf("expected ImageRolloutPending=False after adoption, got %+v", cond) } // The durable record must move with the adoption. stored := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(context.Background(), key, stored); err != nil { - t.Fatal(err) - } - if stored.Annotations[metadata.AnnotationAppliedImages] != mustJSON(t, cfg.Defaults) { - t.Errorf("applied-images annotation not updated: %s", - stored.Annotations[metadata.AnnotationAppliedImages]) - } + c.Require().NoError(r.Get(context.Background(), key, stored)) + c.Eq( + mustJSON(t, cfg.Defaults), + stored.Annotations[metadata.AnnotationAppliedImages], + "applied-images annotation not updated", + ) }) t.Run("consumed acknowledgement is awaiting, not a mismatch", func(t *testing.T) { + c := assert.NewCollecting(t) // The normal state between rollouts: the previous acknowledgment was // adopted and is still in the spec when a newer set becomes available. old := olderApplied() @@ -271,23 +260,26 @@ func TestResolveImages(t *testing.T) { r := newHarness(t, images.UpdateLazy, cluster) rec := r.Recorder.(*record.FakeRecorder) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v1" { - t.Errorf("consumed acknowledgement rolled the cluster: %s", - cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq( + "test/pgctld:v1", + cluster.Spec.Images.Postgres, + "consumed acknowledgement rolled the cluster", + ) cond := pendingCond(cluster) - if cond == nil || cond.Reason != multigresv1alpha1.ReasonAwaitingAcknowledgement { - t.Errorf("expected AwaitingAcknowledgement, got %+v", cond) - } - if hasEvent(t, rec, "ImagesRevisionMismatch") { - t.Error("consumed acknowledgement must not warn as a mismatch") - } + c.False( + cond == nil || cond.Reason != multigresv1alpha1.ReasonAwaitingAcknowledgement, + "expected AwaitingAcknowledgement, got %+v", + cond, + ) + c.False( + hasEvent(t, rec, "ImagesRevisionMismatch"), + "consumed acknowledgement must not warn as a mismatch", + ) }) t.Run("mismatched acknowledgement holds and reports RevisionMismatch", func(t *testing.T) { + c := assert.NewCollecting(t) old := olderApplied() cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: mustJSON(t, old), @@ -297,52 +289,45 @@ func TestResolveImages(t *testing.T) { } r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v1" { - t.Errorf("mismatched acknowledgement rolled the cluster: %s", - cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq( + "test/pgctld:v1", + cluster.Spec.Images.Postgres, + "mismatched acknowledgement rolled the cluster", + ) cond := pendingCond(cluster) - if cond == nil || cond.Reason != multigresv1alpha1.ReasonRevisionMismatch { - t.Errorf("expected RevisionMismatch reason, got %+v", cond) - } + c.False( + cond == nil || cond.Reason != multigresv1alpha1.ReasonRevisionMismatch, + "expected RevisionMismatch reason, got %+v", + cond, + ) }) t.Run("new cluster adopts immediately and records the set", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(nil, nil) r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v2" { - t.Errorf("expected current default, got %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v2", cluster.Spec.Images.Postgres, "expected current default, got") stored := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(context.Background(), key, stored); err != nil { - t.Fatal(err) - } - if stored.Annotations[metadata.AnnotationAppliedImages] == "" { - t.Error("applied-images annotation not recorded for new cluster") - } + c.Require().NoError(r.Get(context.Background(), key, stored)) + c.NotEq( + "", + stored.Annotations[metadata.AnnotationAppliedImages], + "applied-images annotation not recorded for new cluster", + ) }) t.Run("explicit spec pins always win", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(nil, nil) cluster.Spec.Images.Postgres = "pinned/pgctld:v0" r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "pinned/pgctld:v0" { - t.Errorf("explicit pin overwritten: %s", cluster.Spec.Images.Postgres) - } - if cluster.Spec.Images.Multiorch != "test/multigres:v2" { - t.Errorf("unset field not resolved: %s", cluster.Spec.Images.Multiorch) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("pinned/pgctld:v0", cluster.Spec.Images.Postgres, "explicit pin overwritten") + c.Eq("test/multigres:v2", cluster.Spec.Images.Multiorch, "unset field not resolved") if cluster.Status.Images.Source != multigresv1alpha1.ImageSourceMixed || cluster.Status.Images.Effective.Postgres != "pinned/pgctld:v0" { t.Errorf("unexpected mixed image status: %+v", cluster.Status.Images) @@ -350,18 +335,15 @@ func TestResolveImages(t *testing.T) { }) t.Run("partial unpin after fully pinned adopts current defaults", func(t *testing.T) { + c := assert.NewAborting(t) cluster := newCluster(nil, nil) r := newHarness(t, images.UpdateLazy, cluster) counter := r.Client.(*patchCountingClient) old := olderApplied() r.Images.Defaults = old - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if counter.patches != 1 { - t.Fatalf("initial default record patches = %d, want 1", counter.patches) - } + c.NoError(r.resolveImages(context.Background(), cluster)) + c.Eq(1, counter.patches, "initial default record patches") pinned := multigresv1alpha1.ComponentImages{ Postgres: "pinned/pgctld:v3", @@ -377,15 +359,13 @@ func TestResolveImages(t *testing.T) { pinnedCtx := logr.NewContext(context.Background(), logr.New(sink)) rec := r.Recorder.(*record.FakeRecorder) - if err := r.resolveImages(pinnedCtx, cluster); err != nil { - t.Fatal(err) - } - if counter.patches != 2 { - t.Fatalf("pin transition patches = %d, want exactly 2 total", counter.patches) - } - if got := cluster.Annotations[metadata.AnnotationAppliedImages]; got != appliedImagesFullyPinned { - t.Fatalf("applied-images annotation = %q, want fully-pinned tombstone", got) - } + c.NoError(r.resolveImages(pinnedCtx, cluster)) + c.Eq(2, counter.patches, "pin transition patches") + c.Eq( + appliedImagesFullyPinned, + cluster.Annotations[metadata.AnnotationAppliedImages], + "applied-images annotation", + ) if sink.entries != 0 || hasEvent(t, rec, "") { t.Fatalf("pin transition emitted image activity: logs=%d", sink.entries) } @@ -396,43 +376,35 @@ func TestResolveImages(t *testing.T) { t.Fatalf("unexpected fully pinned status: %+v", got) } cond := pendingCond(cluster) - if cond == nil || cond.Status != metav1.ConditionFalse || - cond.Reason != multigresv1alpha1.ReasonFullyPinned { - t.Fatalf("expected FullyPinned condition, got %+v", cond) - } + c.False(cond == nil || cond.Status != metav1.ConditionFalse || + cond.Reason != multigresv1alpha1.ReasonFullyPinned, "expected FullyPinned condition, got %+v", cond) - if err := r.resolveImages(pinnedCtx, cluster); err != nil { - t.Fatal(err) - } + c.NoError(r.resolveImages(pinnedCtx, cluster)) if counter.patches != 2 || sink.entries != 0 || hasEvent(t, rec, "") { t.Fatalf("steady pinned reconcile was not quiet: patches=%d logs=%d", counter.patches, sink.entries) } cluster.Spec.Images.Postgres = "" - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } + c.NoError(r.resolveImages(context.Background(), cluster)) current := testImagesConfig(images.UpdateLazy).Defaults - if cluster.Spec.Images.Postgres != current.Postgres { - t.Fatalf("partial unpin restored %q, want current default %q", - cluster.Spec.Images.Postgres, current.Postgres) - } + c.Eq(current.Postgres, cluster.Spec.Images.Postgres, "partial unpin restored") if cluster.Status.Images.Source != multigresv1alpha1.ImageSourceMixed || cluster.Status.Images.Effective.Postgres != current.Postgres { t.Fatalf("unexpected partial-unpin status: %+v", cluster.Status.Images) } stored := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(context.Background(), key, stored); err != nil { - t.Fatal(err) - } + c.NoError(r.Get(context.Background(), key, stored)) wantAnnotation := mustJSON(t, current) - if got := stored.Annotations[metadata.AnnotationAppliedImages]; got != wantAnnotation { - t.Fatalf("partial unpin recorded %q, want current defaults", got) - } + c.Eq( + wantAnnotation, + stored.Annotations[metadata.AnnotationAppliedImages], + "partial unpin recorded", + ) }) t.Run("fully pinned tombstone survives interrupted status update", func(t *testing.T) { + c := assert.NewAborting(t) old := olderApplied() pinned := multigresv1alpha1.ComponentImages{ Postgres: "pinned/pgctld:v3", @@ -453,33 +425,29 @@ func TestResolveImages(t *testing.T) { // Simulate reconciliation stopping after image resolution but before // the in-memory fully pinned status is persisted. - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } + c.NoError(r.resolveImages(context.Background(), cluster)) stored := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(context.Background(), key, stored); err != nil { - t.Fatal(err) - } - if stored.Status.Images.Applied != old { - t.Fatalf("test setup lost stale status: %+v", stored.Status.Images) - } + c.NoError(r.Get(context.Background(), key, stored)) + c.Eq( + old, + stored.Status.Images.Applied, + "test setup lost stale status: %+v", + stored.Status.Images, + ) stored.Spec.Images.Postgres = "" - if err := r.Update(context.Background(), stored); err != nil { - t.Fatal(err) - } + c.NoError(r.Update(context.Background(), stored)) fresh := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(context.Background(), key, fresh); err != nil { - t.Fatal(err) - } - if err := r.resolveImages(context.Background(), fresh); err != nil { - t.Fatal(err) - } - if fresh.Spec.Images.Postgres != testImagesConfig(images.UpdateLazy).Defaults.Postgres { - t.Fatalf("unpin restored stale image %q", fresh.Spec.Images.Postgres) - } + c.NoError(r.Get(context.Background(), key, fresh)) + c.NoError(r.resolveImages(context.Background(), fresh)) + c.Eq( + testImagesConfig(images.UpdateLazy).Defaults.Postgres, + fresh.Spec.Images.Postgres, + "unpin restored stale image", + ) }) t.Run("fully pinned to fully unpinned adopts current defaults", func(t *testing.T) { + c := assert.NewAborting(t) pinned := multigresv1alpha1.ComponentImages{ Postgres: "pinned/pgctld:v3", Multiadmin: "pinned/multigres:v3", @@ -496,36 +464,29 @@ func TestResolveImages(t *testing.T) { r := newHarness(t, images.UpdateLazy, cluster) counter := r.Client.(*patchCountingClient) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if counter.patches != 1 { - t.Fatalf("pin transition patches = %d, want 1", counter.patches) - } + c.NoError(r.resolveImages(context.Background(), cluster)) + c.Eq(1, counter.patches, "pin transition patches") cluster.Spec.Images = multigresv1alpha1.ClusterImages{} - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } + c.NoError(r.resolveImages(context.Background(), cluster)) current := testImagesConfig(images.UpdateLazy).Defaults - if got := images.FromSpec(cluster.Spec.Images); got != current { - t.Fatalf("full unpin resolved %+v, want current defaults %+v", got, current) - } + c.Eq(current, images.FromSpec(cluster.Spec.Images), "full unpin resolved") if cluster.Status.Images.Source != multigresv1alpha1.ImageSourceDefaults || cluster.Status.Images.Effective != current { t.Fatalf("unexpected full-unpin status: %+v", cluster.Status.Images) } stored := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(context.Background(), key, stored); err != nil { - t.Fatal(err) - } + c.NoError(r.Get(context.Background(), key, stored)) wantAnnotation := mustJSON(t, current) - if got := stored.Annotations[metadata.AnnotationAppliedImages]; got != wantAnnotation { - t.Fatalf("full unpin recorded %q, want current defaults", got) - } + c.Eq( + wantAnnotation, + stored.Annotations[metadata.AnnotationAppliedImages], + "full unpin recorded", + ) }) t.Run("corrupt state fails closed under lazy strategy", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: "{not json", }, nil) @@ -534,36 +495,34 @@ func TestResolveImages(t *testing.T) { counter := r.Client.(*patchCountingClient) err := r.resolveImages(context.Background(), cluster) - if err == nil || !strings.Contains(err.Error(), "lazy update strategy") { - t.Fatalf("resolveImages() error = %v, want lazy state error", err) - } - if cluster.Spec.Images.Postgres != "" { - t.Errorf("corrupt lazy state resolved image %q", cluster.Spec.Images.Postgres) - } - if counter.patches != 0 { - t.Errorf("corrupt lazy state performed %d patches, want 0", counter.patches) - } - if !hasEvent(t, rec, "ImagesRecordInvalid") { - t.Error("expected ImagesRecordInvalid warning event") - } + c.Require(). + False(err == nil || !strings.Contains(err.Error(), "lazy update strategy"), "resolveImages() error = %v, want lazy state error", err) + c.Eq("", cluster.Spec.Images.Postgres, "corrupt lazy state resolved image") + c.Eq(0, counter.patches, "corrupt lazy state performed") + c.True( + hasEvent(t, rec, "ImagesRecordInvalid"), + "expected ImagesRecordInvalid warning event", + ) }) t.Run("empty annotation fails closed under lazy strategy", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: "", }, nil) r := newHarness(t, images.UpdateLazy, cluster) rec := r.Recorder.(*record.FakeRecorder) - if err := r.resolveImages(context.Background(), cluster); err == nil { - t.Fatal("resolveImages() error = nil, want invalid lazy state error") - } - if !hasEvent(t, rec, "ImagesRecordInvalid") { - t.Error("expected ImagesRecordInvalid warning event") - } + c.Require(). + Error(r.resolveImages(context.Background(), cluster), "resolveImages() error = nil, want invalid lazy state error") + c.True( + hasEvent(t, rec, "ImagesRecordInvalid"), + "expected ImagesRecordInvalid warning event", + ) }) t.Run("corrupted annotation falls back to status and holds", func(t *testing.T) { + c := assert.NewCollecting(t) old := olderApplied() cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: "{not json", @@ -573,15 +532,16 @@ func TestResolveImages(t *testing.T) { }) r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v1" { - t.Errorf("expected hold via status fallback, got %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq( + "test/pgctld:v1", + cluster.Spec.Images.Postgres, + "expected hold via status fallback, got", + ) }) t.Run("partial annotation fails closed under lazy strategy", func(t *testing.T) { + c := assert.NewCollecting(t) // A recorded set missing components must not resolve some components // from the record and drop the rest to compiled-in fallbacks. cluster := newCluster(map[string]string{ @@ -590,36 +550,35 @@ func TestResolveImages(t *testing.T) { r := newHarness(t, images.UpdateLazy, cluster) rec := r.Recorder.(*record.FakeRecorder) - if err := r.resolveImages(context.Background(), cluster); err == nil { - t.Fatal("resolveImages() error = nil, want invalid lazy state error") - } + c.Require(). + Error(r.resolveImages(context.Background(), cluster), "resolveImages() error = nil, want invalid lazy state error") if cluster.Spec.Images.Postgres != "" || cluster.Spec.Images.Multiorch != "" { t.Errorf("partial lazy record resolved images: %+v", cluster.Spec.Images) } - if !hasEvent(t, rec, "ImagesRecordInvalid") { - t.Error("expected ImagesRecordInvalid warning event") - } + c.True( + hasEvent(t, rec, "ImagesRecordInvalid"), + "expected ImagesRecordInvalid warning event", + ) }) t.Run("corrupt state recovers under immediate strategy", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: "{not json", }, nil) r := newHarness(t, images.UpdateImmediate, cluster) rec := r.Recorder.(*record.FakeRecorder) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v2" { - t.Errorf("expected current defaults, got %s", cluster.Spec.Images.Postgres) - } - if !hasEvent(t, rec, "ImagesRecordInvalid") { - t.Error("expected ImagesRecordInvalid warning event") - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v2", cluster.Spec.Images.Postgres, "expected current defaults, got") + c.True( + hasEvent(t, rec, "ImagesRecordInvalid"), + "expected ImagesRecordInvalid warning event", + ) }) t.Run("acknowledgement with immediate strategy warns that it is ignored", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(nil, nil) cluster.Spec.ImageUpdatePolicy = &multigresv1alpha1.ImageUpdatePolicy{ AcknowledgedRevision: "abcdef123456", @@ -627,15 +586,15 @@ func TestResolveImages(t *testing.T) { r := newHarness(t, images.UpdateImmediate, cluster) rec := r.Recorder.(*record.FakeRecorder) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if !hasEvent(t, rec, "ImagesAcknowledgementIgnored") { - t.Error("expected ImagesAcknowledgementIgnored warning event") - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.True( + hasEvent(t, rec, "ImagesAcknowledgementIgnored"), + "expected ImagesAcknowledgementIgnored warning event", + ) }) t.Run("per-cluster lazy overrides an immediate operator", func(t *testing.T) { + c := assert.NewCollecting(t) // A single cluster can be frozen while the fleet follows the operator. old := olderApplied() cluster := newCluster(map[string]string{ @@ -646,19 +605,17 @@ func TestResolveImages(t *testing.T) { } r := newHarness(t, images.UpdateImmediate, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v1" { - t.Errorf("per-cluster lazy did not hold: %s", cluster.Spec.Images.Postgres) - } - if cluster.Status.Images.UpdateStrategy != string(images.UpdateLazy) { - t.Errorf("status must report the effective strategy, got %q", - cluster.Status.Images.UpdateStrategy) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v1", cluster.Spec.Images.Postgres, "per-cluster lazy did not hold") + c.Eq( + string(images.UpdateLazy), + cluster.Status.Images.UpdateStrategy, + "status must report the effective strategy, got", + ) }) t.Run("per-cluster immediate overrides a lazy operator", func(t *testing.T) { + c := assert.NewCollecting(t) old := olderApplied() cluster := newCluster(map[string]string{ metadata.AnnotationAppliedImages: mustJSON(t, old), @@ -668,25 +625,22 @@ func TestResolveImages(t *testing.T) { } r := newHarness(t, images.UpdateLazy, cluster) - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != "test/pgctld:v2" { - t.Errorf("per-cluster immediate did not adopt: %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq("test/pgctld:v2", cluster.Spec.Images.Postgres, "per-cluster immediate did not adopt") }) t.Run("zero-value config falls back to compiled defaults", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := newCluster(nil, nil) r := newHarness(t, "", cluster) r.Images = images.Config{} - if err := r.resolveImages(context.Background(), cluster); err != nil { - t.Fatal(err) - } - if cluster.Spec.Images.Postgres != multigresv1alpha1.DefaultPostgresImage { - t.Errorf("expected compiled default, got %s", cluster.Spec.Images.Postgres) - } + c.Require().NoError(r.resolveImages(context.Background(), cluster)) + c.Eq( + multigresv1alpha1.DefaultPostgresImage, + cluster.Spec.Images.Postgres, + "expected compiled default, got", + ) }) } @@ -695,6 +649,7 @@ func TestResolveImages(t *testing.T) { // image set until spec.imageUpdatePolicy.acknowledgedRevision names the new // revision, and follow it once it does. func TestReconcile_LazyImageRollout(t *testing.T) { + ck := assert.NewAborting(t) coreTpl, cellTpl, shardTpl, baseCluster, clusterName, namespace := setupFixtures(t) // The fixture pins spec.images; this test is about operator defaults. baseCluster.Spec.Images = multigresv1alpha1.ClusterImages{} @@ -732,31 +687,30 @@ func TestReconcile_LazyImageRollout(t *testing.T) { assertChildImages := func(t *testing.T, wantGateway, wantOrch multigresv1alpha1.ImageRef) { t.Helper() + ck := assert.NewCollecting(t) cells := &multigresv1alpha1.CellList{} - if err := c.List(t.Context(), cells); err != nil { - t.Fatal(err) - } - if len(cells.Items) == 0 { - t.Fatal("no Cell children created") - } + ck.Require().NoError(c.List(t.Context(), cells)) + ck.Require().NotEmpty(cells.Items, "no Cell children created") for _, cell := range cells.Items { - if cell.Spec.Images.Multigateway != wantGateway { - t.Errorf("cell %s multigateway = %s, want %s", - cell.Name, cell.Spec.Images.Multigateway, wantGateway) - } + ck.Eq( + wantGateway, + cell.Spec.Images.Multigateway, + "cell %s multigateway = %s, want", + cell.Name, + cell.Spec.Images.Multigateway, + ) } tgs := &multigresv1alpha1.TableGroupList{} - if err := c.List(t.Context(), tgs); err != nil { - t.Fatal(err) - } - if len(tgs.Items) == 0 { - t.Fatal("no TableGroup children created") - } + ck.Require().NoError(c.List(t.Context(), tgs)) + ck.Require().NotEmpty(tgs.Items, "no TableGroup children created") for _, tg := range tgs.Items { - if tg.Spec.Images.Multiorch != wantOrch { - t.Errorf("tablegroup %s multiorch = %s, want %s", - tg.Name, tg.Spec.Images.Multiorch, wantOrch) - } + ck.Eq( + wantOrch, + tg.Spec.Images.Multiorch, + "tablegroup %s multiorch = %s, want", + tg.Name, + tg.Spec.Images.Multiorch, + ) } } @@ -769,17 +723,12 @@ func TestReconcile_LazyImageRollout(t *testing.T) { // Pass 2: available revision acknowledged in the spec — children must follow. cluster := &multigresv1alpha1.MultigresCluster{} - if err := c.Get(t.Context(), req.NamespacedName, cluster); err != nil { - t.Fatal(err) - } + ck.NoError(c.Get(t.Context(), req.NamespacedName, cluster)) cluster.Spec.ImageUpdatePolicy = &multigresv1alpha1.ImageUpdatePolicy{ AcknowledgedRevision: images.Revision(testImagesConfig(images.UpdateLazy).Defaults), } - if err := c.Update(t.Context(), cluster); err != nil { - t.Fatal(err) - } - if _, err := r.Reconcile(t.Context(), req); err != nil { - t.Fatalf("second reconcile: %v", err) - } + ck.NoError(c.Update(t.Context(), cluster)) + _, err := r.Reconcile(t.Context(), req) + ck.NoError(err, "second reconcile") assertChildImages(t, "test/multigres:v2", "test/multigres:v2") } diff --git a/pkg/cluster-handler/controller/multigrescluster/integration_adminweb_test.go b/pkg/cluster-handler/controller/multigrescluster/integration_adminweb_test.go index 7ca51c94..ff03bcf3 100644 --- a/pkg/cluster-handler/controller/multigrescluster/integration_adminweb_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/integration_adminweb_test.go @@ -8,8 +8,6 @@ import ( "testing" "github.com/google/go-cmp/cmp/cmpopts" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" @@ -21,10 +19,13 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestExternalAdminWeb_EnableDisableLifecycle(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) const clusterName = "aw-lifecycle" @@ -57,7 +58,7 @@ func TestExternalAdminWeb_EnableDisableLifecycle(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - require.NoError(t, k8sClient.Create(t.Context(), cluster)) + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) watcher.SetCmpOpts( testutil.IgnoreMetaRuntimeFields(), @@ -93,67 +94,86 @@ func TestExternalAdminWeb_EnableDisableLifecycle(t *testing.T) { }, } - require.NoError(t, watcher.WaitForMatch(expectedAWSvc), - "multiadmin-web Service should be ClusterIP with externalIPs and annotations") + c.Require(). + NoError(watcher.WaitForMatch(expectedAWSvc), "multiadmin-web Service should be ClusterIP with externalIPs and annotations") // Step 3: Verify initial condition is NoReadyAdminWeb (endpoint assigned via externalIP, 0 ready pods). - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionFalse && cond.Reason == multigresv1alpha1.ReasonNoReadyAdminWeb && mgc.Status.AdminWeb != nil && mgc.Status.AdminWeb.ExternalEndpoint == "2001:db8::200" - }, testTimeout, pollInterval, "condition should be False/NoReadyAdminWeb before pods are ready") + }, "condition should be False/NoReadyAdminWeb before pods are ready") // Step 4: Simulate the admin-web Deployment reporting ready replicas. var awDeploy appsv1.Deployment - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { err := k8sClient.Get(t.Context(), client.ObjectKey{ Name: clusterName + "-multiadmin-web", Namespace: testNamespace, }, &awDeploy) return err == nil - }, testTimeout, pollInterval, "admin-web Deployment should exist") + }, "admin-web Deployment should exist") awDeploy.Status.ReadyReplicas = 1 awDeploy.Status.Replicas = 1 - require.NoError(t, k8sClient.Status().Update(t.Context(), &awDeploy), - "simulating admin-web Deployment reporting ready replicas") + c.Require(). + NoError(k8sClient.Status().Update(t.Context(), &awDeploy), "simulating admin-web Deployment reporting ready replicas") // Step 5: Verify condition transitions to EndpointReady. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionTrue && cond.Reason == multigresv1alpha1.ReasonEndpointReady && mgc.Status.AdminWeb != nil && mgc.Status.AdminWeb.ExternalEndpoint == "2001:db8::200" - }, testTimeout, pollInterval, "condition should be True/EndpointReady with endpoint in status") + }, "condition should be True/EndpointReady with endpoint in status") // Verify observedGeneration is set on the condition. var mgcCheck multigresv1alpha1.MultigresCluster - require.NoError(t, k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) - cond := meta.FindStatusCondition(mgcCheck.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) - require.NotNil(t, cond) - assert.Equal(t, mgcCheck.Generation, cond.ObservedGeneration, - "observedGeneration should match cluster generation") + c.Require().NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) + cond := meta.FindStatusCondition( + mgcCheck.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) + c.Require().NotNil(cond) + c.EqDeep( + mgcCheck.Generation, + cond.ObservedGeneration, + "observedGeneration should match cluster generation", + ) // Step 6: Disable external admin web and verify reversion. - require.NoError(t, k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) + c.Require().NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) cluster.Spec.ExternalAdminWeb = &multigresv1alpha1.ExternalAdminWebConfig{ Enabled: false, } - require.NoError(t, k8sClient.Update(t.Context(), cluster), - "disabling external admin web") + c.Require().NoError(k8sClient.Update(t.Context(), cluster), "disabling external admin web") // Step 7: Verify Service reverts to ClusterIP with no annotations. expectedClusterIPSvc := &corev1.Service{ @@ -177,22 +197,30 @@ func TestExternalAdminWeb_EnableDisableLifecycle(t *testing.T) { }, } - require.NoError(t, watcher.WaitForMatch(expectedClusterIPSvc), - "multiadmin-web Service should revert to ClusterIP after disabling") + c.Require(). + NoError(watcher.WaitForMatch(expectedClusterIPSvc), "multiadmin-web Service should revert to ClusterIP after disabling") // Step 8: Verify admin web status is nil and condition is removed. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) return mgc.Status.AdminWeb == nil && cond == nil - }, testTimeout, pollInterval, "admin web status should be nil and condition removed after disabling") + }, "admin web status should be nil and condition removed after disabling") } func TestExternalAdminWeb_NoReadyAdminWebTransition(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) const clusterName = "aw-no-ready" @@ -222,57 +250,77 @@ func TestExternalAdminWeb_NoReadyAdminWebTransition(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - require.NoError(t, k8sClient.Create(t.Context(), cluster)) + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) // Step 2: Wait for initial NoReadyAdminWeb condition. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionFalse && cond.Reason == multigresv1alpha1.ReasonNoReadyAdminWeb && strings.Contains(cond.Message, "no multiadmin-web pods are ready") && mgc.Status.AdminWeb != nil && mgc.Status.AdminWeb.ExternalEndpoint == "2001:db8::201" - }, testTimeout, pollInterval, "condition should be False/NoReadyAdminWeb with endpoint populated") + }, "condition should be False/NoReadyAdminWeb with endpoint populated") // Step 3: Simulate the admin-web Deployment reporting ready replicas. var awDeploy appsv1.Deployment - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { err := k8sClient.Get(t.Context(), client.ObjectKey{ Name: clusterName + "-multiadmin-web", Namespace: testNamespace, }, &awDeploy) return err == nil - }, testTimeout, pollInterval, "admin-web Deployment should exist") + }, "admin-web Deployment should exist") awDeploy.Status.ReadyReplicas = 2 awDeploy.Status.Replicas = 2 - require.NoError(t, k8sClient.Status().Update(t.Context(), &awDeploy)) + c.Require().NoError(k8sClient.Status().Update(t.Context(), &awDeploy)) // Step 4: Verify condition transitions to True/EndpointReady. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionTrue && cond.Reason == multigresv1alpha1.ReasonEndpointReady && strings.Contains(cond.Message, "is serving traffic") && mgc.Status.AdminWeb != nil && mgc.Status.AdminWeb.ExternalEndpoint == "2001:db8::201" - }, testTimeout, pollInterval, "condition should transition to True/EndpointReady after Deployment reports ready") + }, "condition should transition to True/EndpointReady after Deployment reports ready") // Verify observedGeneration is set correctly. var mgcCheck multigresv1alpha1.MultigresCluster - require.NoError(t, k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) - cond := meta.FindStatusCondition(mgcCheck.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady) - require.NotNil(t, cond) - assert.Equal(t, mgcCheck.Generation, cond.ObservedGeneration, - "observedGeneration should match cluster generation") + c.Require().NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) + cond := meta.FindStatusCondition( + mgcCheck.Status.Conditions, + multigresv1alpha1.ConditionAdminWebExternalReady, + ) + c.Require().NotNil(cond) + c.EqDeep( + mgcCheck.Generation, + cond.ObservedGeneration, + "observedGeneration should match cluster generation", + ) } diff --git a/pkg/cluster-handler/controller/multigrescluster/integration_gateway_test.go b/pkg/cluster-handler/controller/multigrescluster/integration_gateway_test.go index b439a655..3a737dec 100644 --- a/pkg/cluster-handler/controller/multigrescluster/integration_gateway_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/integration_gateway_test.go @@ -9,8 +9,6 @@ import ( "time" "github.com/google/go-cmp/cmp/cmpopts" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/meta" @@ -21,12 +19,15 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) const pollInterval = 200 * time.Millisecond func TestExternalGateway_EnableDisableLifecycle(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) const clusterName = "gw-lifecycle" @@ -59,7 +60,7 @@ func TestExternalGateway_EnableDisableLifecycle(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - require.NoError(t, k8sClient.Create(t.Context(), cluster)) + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) // Add extra comparison options for Service runtime fields. watcher.SetCmpOpts( @@ -123,69 +124,88 @@ func TestExternalGateway_EnableDisableLifecycle(t *testing.T) { }, } - require.NoError(t, watcher.WaitForMatch(expectedGwSvc), - "global multigateway Service should be ClusterIP with externalIPs and annotations") - require.NoError(t, watcher.WaitForMatch(expectedReplicaGwSvc), - "global multigateway replica Service should be ClusterIP") + c.Require(). + NoError(watcher.WaitForMatch(expectedGwSvc), "global multigateway Service should be ClusterIP with externalIPs and annotations") + c.Require(). + NoError(watcher.WaitForMatch(expectedReplicaGwSvc), "global multigateway replica Service should be ClusterIP") // Step 3: Verify initial condition is NoReadyGateways (endpoint assigned via externalIP, 0 ready gateways). - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionFalse && cond.Reason == multigresv1alpha1.ReasonNoReadyGateways && mgc.Status.Gateway != nil && mgc.Status.Gateway.ExternalEndpoint == "2001:db8::100" - }, testTimeout, pollInterval, "condition should be False/NoReadyGateways before gateways are ready") + }, "condition should be False/NoReadyGateways before gateways are ready") // Step 4: Simulate Cell reporting ready gateways by updating Cell status. var cellList multigresv1alpha1.CellList - require.NoError(t, k8sClient.List(t.Context(), &cellList, + c.Require().NoError(k8sClient.List(t.Context(), &cellList, client.InNamespace(testNamespace), client.MatchingLabels{"multigres.com/cluster": clusterName}, )) - require.NotEmpty(t, cellList.Items, "expected at least one Cell CR") + c.Require().NotEmpty(cellList.Items, "expected at least one Cell CR") cell := &cellList.Items[0] cell.Status.GatewayReadyReplicas = 1 cell.Status.GatewayReplicas = 1 cell.Status.ObservedGeneration = cell.Generation - require.NoError(t, k8sClient.Status().Update(t.Context(), cell), - "simulating Cell reporting ready gateway replicas") + c.Require(). + NoError(k8sClient.Status().Update(t.Context(), cell), "simulating Cell reporting ready gateway replicas") // Step 5: Verify condition transitions to EndpointReady. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionTrue && cond.Reason == multigresv1alpha1.ReasonEndpointReady && mgc.Status.Gateway != nil && mgc.Status.Gateway.ExternalEndpoint == "2001:db8::100" - }, testTimeout, pollInterval, "condition should be True/EndpointReady with endpoint in status") + }, "condition should be True/EndpointReady with endpoint in status") // Verify observedGeneration is set on the condition. var mgcCheck multigresv1alpha1.MultigresCluster - require.NoError(t, k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) - cond := meta.FindStatusCondition(mgcCheck.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) - require.NotNil(t, cond) - assert.Equal(t, mgcCheck.Generation, cond.ObservedGeneration, - "observedGeneration should match cluster generation") + c.Require().NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) + cond := meta.FindStatusCondition( + mgcCheck.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) + c.Require().NotNil(cond) + c.EqDeep( + mgcCheck.Generation, + cond.ObservedGeneration, + "observedGeneration should match cluster generation", + ) // Step 6: Disable external gateway and verify reversion. - require.NoError(t, k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) + c.Require().NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) cluster.Spec.ExternalGateway = &multigresv1alpha1.ExternalGatewayConfig{ Enabled: false, } - require.NoError(t, k8sClient.Update(t.Context(), cluster), - "disabling external gateway") + c.Require().NoError(k8sClient.Update(t.Context(), cluster), "disabling external gateway") // Step 7: Verify Service reverts to ClusterIP with no gateway annotations. expectedClusterIPSvc := &corev1.Service{ @@ -236,24 +256,32 @@ func TestExternalGateway_EnableDisableLifecycle(t *testing.T) { }, } - require.NoError(t, watcher.WaitForMatch(expectedClusterIPSvc), - "global multigateway Service should revert to ClusterIP after disabling") - require.NoError(t, watcher.WaitForMatch(expectedReplicaClusterIPSvc), - "global multigateway replica Service should remain ClusterIP") + c.Require(). + NoError(watcher.WaitForMatch(expectedClusterIPSvc), "global multigateway Service should revert to ClusterIP after disabling") + c.Require(). + NoError(watcher.WaitForMatch(expectedReplicaClusterIPSvc), "global multigateway replica Service should remain ClusterIP") // Step 8: Verify gateway status is nil and condition is removed. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) return mgc.Status.Gateway == nil && cond == nil - }, testTimeout, pollInterval, "gateway status should be nil and condition removed after disabling") + }, "gateway status should be nil and condition removed after disabling") } func TestExternalGateway_NoReadyGatewaysTransition(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) const clusterName = "gw-no-ready" @@ -283,57 +311,77 @@ func TestExternalGateway_NoReadyGatewaysTransition(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - require.NoError(t, k8sClient.Create(t.Context(), cluster)) + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) // Step 2: Wait for initial NoReadyGateways condition (external endpoint assigned, 0 ready gateways). - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionFalse && cond.Reason == multigresv1alpha1.ReasonNoReadyGateways && strings.Contains(cond.Message, "no multigateway pods are ready") && mgc.Status.Gateway != nil && mgc.Status.Gateway.ExternalEndpoint == "2001:db8::101" - }, testTimeout, pollInterval, "condition should be False/NoReadyGateways with endpoint populated") + }, "condition should be False/NoReadyGateways with endpoint populated") // Step 3: Simulate Cell reporting gatewayReadyReplicas > 0. var cellList multigresv1alpha1.CellList - require.NoError(t, k8sClient.List(t.Context(), &cellList, + c.Require().NoError(k8sClient.List(t.Context(), &cellList, client.InNamespace(testNamespace), client.MatchingLabels{"multigres.com/cluster": clusterName}, )) - require.NotEmpty(t, cellList.Items, "expected at least one Cell CR") + c.Require().NotEmpty(cellList.Items, "expected at least one Cell CR") cell := &cellList.Items[0] cell.Status.GatewayReadyReplicas = 2 cell.Status.GatewayReplicas = 2 cell.Status.ObservedGeneration = cell.Generation - require.NoError(t, k8sClient.Status().Update(t.Context(), cell)) + c.Require().NoError(k8sClient.Status().Update(t.Context(), cell)) // Step 4: Verify condition transitions to True/EndpointReady. - assert.Eventually(t, func() bool { + c.EventuallyTrue(testTimeout, pollInterval, func() bool { var mgc multigresv1alpha1.MultigresCluster - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgc); err != nil { + if err := k8sClient.Get( + t.Context(), + client.ObjectKeyFromObject(cluster), + &mgc, + ); err != nil { return false } - cond := meta.FindStatusCondition(mgc.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) + cond := meta.FindStatusCondition( + mgc.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) return cond != nil && cond.Status == metav1.ConditionTrue && cond.Reason == multigresv1alpha1.ReasonEndpointReady && strings.Contains(cond.Message, "is serving traffic") && mgc.Status.Gateway != nil && mgc.Status.Gateway.ExternalEndpoint == "2001:db8::101" - }, testTimeout, pollInterval, "condition should transition to True/EndpointReady after Cell reports ready gateways") + }, "condition should transition to True/EndpointReady after Cell reports ready gateways") // Verify observedGeneration is set correctly. var mgcCheck multigresv1alpha1.MultigresCluster - require.NoError(t, k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) - cond := meta.FindStatusCondition(mgcCheck.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady) - require.NotNil(t, cond) - assert.Equal(t, mgcCheck.Generation, cond.ObservedGeneration, - "observedGeneration should match cluster generation") + c.Require().NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), &mgcCheck)) + cond := meta.FindStatusCondition( + mgcCheck.Status.Conditions, + multigresv1alpha1.ConditionGatewayExternalReady, + ) + c.Require().NotNil(cond) + c.EqDeep( + mgcCheck.Generation, + cond.ObservedGeneration, + "observedGeneration should match cluster generation", + ) } diff --git a/pkg/cluster-handler/controller/multigrescluster/integration_lifecycle_test.go b/pkg/cluster-handler/controller/multigrescluster/integration_lifecycle_test.go index 073fc335..5c87191a 100644 --- a/pkg/cluster-handler/controller/multigrescluster/integration_lifecycle_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/integration_lifecycle_test.go @@ -20,6 +20,8 @@ import ( "github.com/multigres/multigres-operator/pkg/resolver" "github.com/multigres/multigres-operator/pkg/util/metadata" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestMultigresCluster_Lifecycle(t *testing.T) { @@ -27,6 +29,7 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { t.Run("TableGroup Long Name Hashing", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, _ := setupIntegration(t) // Name length math: 25 (cluster) + 8 (db) + 25 (tg) + 2 (hyphens) = 60 chars. longClusterName := "valid-cluster-name-123456" // 25 chars @@ -59,9 +62,7 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster), "Failed to create cluster") // Verify the TableGroup DOES exist (hashed) tgName := nameutil.JoinWithConstraints( @@ -75,7 +76,11 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { // We just want to wait for it to exist found := false for i := 0; i < 20; i++ { - err := k8sClient.Get(t.Context(), client.ObjectKey{Name: tgName, Namespace: testNamespace}, expectedTG) + err := k8sClient.Get( + t.Context(), + client.ObjectKey{Name: tgName, Namespace: testNamespace}, + expectedTG, + ) if err == nil { found = true break @@ -85,19 +90,23 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { time.Sleep(200 * time.Millisecond) } - if !found { - t.Errorf("Expected TableGroup %s to be created using hashing, but it was not found after timeout", tgName) - } + c.True( + found, + "Expected TableGroup %s to be created using hashing, but it was not found after timeout", + tgName, + ) // Ensure Cluster exists fetchedCluster := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), fetchedCluster); err != nil { - t.Error("Cluster should exist") - } + c.NoError( + k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), fetchedCluster), + "Cluster should exist", + ) }) t.Run("Annotation Limit (Bombing)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) // 250 chars is near limit (256). If controller appends to this value, it might fail. longAnnotation := strings.Repeat("a", 250) @@ -112,9 +121,7 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster), "Failed to create cluster") // Verify Multiadmin Deployment created successfully WITH annotation deploy := &appsv1.Deployment{ @@ -127,7 +134,9 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultAdminReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "short-annot-bomb", "multiadmin", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "short-annot-bomb", "multiadmin", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -212,13 +221,15 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { }, }, } - if err := watcher.WaitForMatch(deploy); err != nil { - t.Errorf("Multiadmin deployment failed to create with massive annotation: %v", err) - } + c.NoError( + watcher.WaitForMatch(deploy), + "Multiadmin deployment failed to create with massive annotation", + ) }) t.Run("Mutability (Image Update)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "mut-test", Namespace: testNamespace}, @@ -227,14 +238,11 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster), "Failed to create cluster") // Wait for v1 deploy := &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "mut-test-multiadmin", Namespace: testNamespace, Labels: clusterLabels(t, "mut-test", "multiadmin", ""), @@ -243,7 +251,9 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultAdminReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "mut-test", "multiadmin", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "mut-test", "multiadmin", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -325,18 +335,14 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { }, }, } - if err := watcher.WaitForMatch(deploy); err != nil { - t.Fatalf("Failed to wait for initial deployment v1: %v", err) - } + c.Require(). + NoError(watcher.WaitForMatch(deploy), "Failed to wait for initial deployment v1") // Update Image - if err := k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster); err != nil { - t.Fatal(err) - } + c.Require(). + NoError(k8sClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) cluster.Spec.Images.Multiadmin = "admin:v2" - if err := k8sClient.Update(t.Context(), cluster); err != nil { - t.Fatalf("Failed to update cluster: %v", err) - } + c.Require().NoError(k8sClient.Update(t.Context(), cluster), "Failed to update cluster") // Verify v2 deployV2 := &appsv1.Deployment{ @@ -349,7 +355,9 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultAdminReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "mut-test", "multiadmin", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "mut-test", "multiadmin", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -431,9 +439,6 @@ func TestMultigresCluster_Lifecycle(t *testing.T) { }, }, } - if err := watcher.WaitForMatch(deployV2); err != nil { - t.Errorf("Deployment failed to update to v2: %v", err) - } + c.NoError(watcher.WaitForMatch(deployV2), "Deployment failed to update to v2") }) - } diff --git a/pkg/cluster-handler/controller/multigrescluster/integration_resolution_enforcement_test.go b/pkg/cluster-handler/controller/multigrescluster/integration_resolution_enforcement_test.go index 0330b5d3..ccdbacdd 100644 --- a/pkg/cluster-handler/controller/multigrescluster/integration_resolution_enforcement_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/integration_resolution_enforcement_test.go @@ -14,6 +14,8 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) // TestMultigresCluster_ResolutionLogic validates the "4-Level Override Chain" @@ -23,6 +25,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { t.Run("4-Level Override Precedence", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) // 1. Setup Templates @@ -30,23 +33,23 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { smallTpl := &multigresv1alpha1.CellTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "small", Namespace: testNamespace}, Spec: multigresv1alpha1.CellTemplateSpec{ - Multigateway: &multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + Multigateway: &multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, + }, }, } // "Large" Template -> Replicas: 5 largeTpl := &multigresv1alpha1.CellTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "large", Namespace: testNamespace}, Spec: multigresv1alpha1.CellTemplateSpec{ - Multigateway: &multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(5))}}, + Multigateway: &multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(5))}, + }, }, } - if err := k8sClient.Create(t.Context(), smallTpl); err != nil { - t.Fatal(err) - } - if err := k8sClient.Create(t.Context(), largeTpl); err != nil { - t.Fatal(err) - } + ck.Require().NoError(k8sClient.Create(t.Context(), smallTpl)) + ck.Require().NoError(k8sClient.Create(t.Context(), largeTpl)) // 2. Create Cluster with various levels of overrides clusterName := "precedence-test" @@ -74,7 +77,11 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { ZoneID: "use1-az3", CellTemplate: "large", Overrides: &multigresv1alpha1.CellOverrides{ - Multigateway: &multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(3))}}, + Multigateway: &multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(3)), + }, + }, }, }, @@ -83,7 +90,11 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { Name: "zone-d", ZoneID: "use1-az4", Spec: &multigresv1alpha1.CellInlineSpec{ - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(9))}}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(9)), + }, + }, }, }, }, @@ -92,9 +103,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + ck.Require().NoError(k8sClient.Create(t.Context(), cluster), "Failed to create cluster") // 3. Verify Results // Use CompareSpecOnly to avoid needing to construct OwnerRefs/UIDs manually @@ -165,14 +174,17 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { } for _, tc := range cases { - if err := watcher.WaitForMatch(makeExpected(tc.zone, tc.wantReplicas, allCells)); err != nil { - t.Errorf("Precedence failed for %s: %v", tc.zone, err) - } + ck.NoError( + watcher.WaitForMatch(makeExpected(tc.zone, tc.wantReplicas, allCells)), + "Precedence failed for %s", + tc.zone, + ) } }) t.Run("Implicit Namespace Defaulting", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) clusterName := "implicit-default-test" @@ -188,9 +200,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) // We expect the "default" template (created in setupIntegration) to be used. // That template has Replicas: 1. @@ -238,13 +248,15 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { }, } - if err := watcher.WaitForMatch(wantCell); err != nil { - t.Error("Failed to implicitly resolve namespace 'default' template") - } + c.NoError( + watcher.WaitForMatch(wantCell), + "Failed to implicitly resolve namespace 'default' template", + ) }) t.Run("List Replacement Logic (Cells)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) // Setup ShardTemplate @@ -257,9 +269,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { }, }, } - if err := k8sClient.Create(t.Context(), tpl); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(t.Context(), tpl)) // Create cluster clusterName := "list-replace-test" @@ -297,9 +307,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) // Verify by checking the TableGroup (since Shard controller isn't running) // We expect the resolved Spec in TableGroup to have the correct list. @@ -348,7 +356,9 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { // VERIFICATION: Only zone-c should be present Cells: []multigresv1alpha1.CellName{"zone-c"}, StatelessSpec: multigresv1alpha1.StatelessSpec{ - Replicas: ptr.To(int32(1)), // From implicit defaults + Replicas: ptr.To( + int32(1), + ), // From implicit defaults Resources: resolver.DefaultResourcesOrch(), // FIX: Expect defaults }, }, @@ -370,8 +380,13 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { }, PVCDeletionPolicy: nil, // Shard-level is nil Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, }, }, }, @@ -380,8 +395,13 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, }, TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{Enabled: ptr.To(true)}, DurabilityPolicy: "AT_LEAST_2", @@ -390,9 +410,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { } setTestTableGroupPostgresPasswordSecretRef(wantTG) - if err := watcher.WaitForMatch(wantTG); err != nil { - t.Errorf("List replacement failed: %v", err) - } + c.NoError(watcher.WaitForMatch(wantTG), "List replacement failed") }) } @@ -400,6 +418,7 @@ func TestMultigresCluster_ResolutionLogic(t *testing.T) { // enforces the desired state, including reverting manual changes (immutability). func TestMultigresCluster_EnforcementLogic(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) clusterName := "enforcement-test" @@ -410,7 +429,11 @@ func TestMultigresCluster_EnforcementLogic(t *testing.T) { { Name: "zone-a", ZoneID: "use1-az1", Spec: &multigresv1alpha1.CellInlineSpec{ - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(2))}}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(2)), + }, + }, }, }, }, @@ -419,9 +442,7 @@ func TestMultigresCluster_EnforcementLogic(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) watcher.SetCmpOpts(testutil.CompareSpecOnly()...) @@ -469,28 +490,20 @@ func TestMultigresCluster_EnforcementLogic(t *testing.T) { }, } - if err := watcher.WaitForMatch(wantCell); err != nil { - t.Fatal("Initial cell creation failed") - } + c.Require().NoError(watcher.WaitForMatch(wantCell), "Initial cell creation failed") // 2. Tamper (Scale up manually) cellKey := client.ObjectKey{Name: wantCell.Name, Namespace: wantCell.Namespace} cell := &multigresv1alpha1.Cell{} - if err := k8sClient.Get(t.Context(), cellKey, cell); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Get(t.Context(), cellKey, cell)) cell.Spec.Multigateway.Replicas = ptr.To(int32(100)) - if err := k8sClient.Update(t.Context(), cell); err != nil { - t.Fatal("Failed to tamper with cell") - } + c.Require().NoError(k8sClient.Update(t.Context(), cell), "Failed to tamper with cell") // 3. Verify Reversion // We wait for the object to match 'wantCell' again. // Since client.Update succeeded, the object *was* changed. The fact that it matches // wantCell (2 replicas) afterwards proves the controller reverted it. - if err := watcher.WaitForMatch(wantCell); err != nil { - t.Errorf("Controller failed to revert manual change: %v", err) - } + c.NoError(watcher.WaitForMatch(wantCell), "Controller failed to revert manual change") } // TestMultigresCluster_V1Alpha1Constraints verifies strict v1alpha1 validations @@ -540,8 +553,11 @@ func TestMultigresCluster_V1Alpha1Constraints(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { cluster := &multigresv1alpha1.MultigresCluster{ - ObjectMeta: metav1.ObjectMeta{Name: strings.ToLower(strings.ReplaceAll(tc.name, " ", "-")), Namespace: testNamespace}, - Spec: tc.clusterSpec, + ObjectMeta: metav1.ObjectMeta{ + Name: strings.ToLower(strings.ReplaceAll(tc.name, " ", "-")), + Namespace: testNamespace, + }, + Spec: tc.clusterSpec, } setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(t.Context(), cluster) @@ -558,6 +574,7 @@ func TestMultigresCluster_V1Alpha1Constraints(t *testing.T) { // ShardTemplates (specifically PVCDeletionPolicy) are correctly applied. func TestMultigresCluster_TemplateOverrides(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) // 1. Create a ShardTemplate with specific PVC Policy @@ -571,9 +588,7 @@ func TestMultigresCluster_TemplateOverrides(t *testing.T) { }, }, } - if err := k8sClient.Create(t.Context(), tpl); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(t.Context(), tpl)) // 2. Create Cluster using this template in Defaults clusterName := "template-policy-test" @@ -604,9 +619,7 @@ func TestMultigresCluster_TemplateOverrides(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(t.Context(), cluster)) // 3. Verify TableGroup has the correct Resolved Shard Spec watcher.SetCmpOpts(testutil.CompareSpecOnly()...) @@ -652,7 +665,9 @@ func TestMultigresCluster_TemplateOverrides(t *testing.T) { Multiorch: multigresv1alpha1.MultiorchSpec{ Cells: []multigresv1alpha1.CellName{"zone-a"}, StatelessSpec: multigresv1alpha1.StatelessSpec{ - Replicas: ptr.To(int32(1)), // From implicit defaults in resolver + Replicas: ptr.To( + int32(1), + ), // From implicit defaults in resolver Resources: resolver.DefaultResourcesOrch(), // Defaults }, }, @@ -678,8 +693,13 @@ func TestMultigresCluster_TemplateOverrides(t *testing.T) { WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, }, }, }, @@ -689,8 +709,11 @@ func TestMultigresCluster_TemplateOverrides(t *testing.T) { WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}, + }, }, TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{Enabled: ptr.To(true)}, DurabilityPolicy: "AT_LEAST_2", @@ -699,7 +722,5 @@ func TestMultigresCluster_TemplateOverrides(t *testing.T) { } setTestTableGroupPostgresPasswordSecretRef(wantTG) - if err := watcher.WaitForMatch(wantTG); err != nil { - t.Errorf("Template override validation failed: %v", err) - } + c.NoError(watcher.WaitForMatch(wantTG), "Template override validation failed") } diff --git a/pkg/cluster-handler/controller/multigrescluster/integration_test.go b/pkg/cluster-handler/controller/multigrescluster/integration_test.go index f27eeecc..2ef5bf02 100644 --- a/pkg/cluster-handler/controller/multigrescluster/integration_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/integration_test.go @@ -27,6 +27,8 @@ import ( "github.com/multigres/multigres/go/common/topoclient" "github.com/multigres/multigres/go/common/topoclient/memorytopo" + + "github.com/multigres/testkit/assert" ) // ============================================================================ @@ -46,6 +48,7 @@ func canonicalGlobalRoot(clusterName string) string { // It returns a ready-to-use K8s Client and a ResourceWatcher. func setupIntegration(t *testing.T) (client.Client, *testutil.ResourceWatcher) { t.Helper() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -93,11 +96,9 @@ func setupIntegration(t *testing.T) (client.Client, *testutil.ResourceWatcher) { }, } - if err := reconciler.SetupWithManager(mgr, controller.Options{ + c.NoError(reconciler.SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller: %v", err) - } + }), "Failed to create controller") k8sClient := mgr.GetClient() @@ -115,7 +116,9 @@ func setupIntegration(t *testing.T) (client.Client, *testutil.ResourceWatcher) { &multigresv1alpha1.CellTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "default", Namespace: testNamespace}, Spec: multigresv1alpha1.CellTemplateSpec{ - Multigateway: &multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + Multigateway: &multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, + }, }, }, &multigresv1alpha1.ShardTemplate{ @@ -125,9 +128,13 @@ func setupIntegration(t *testing.T) (client.Client, *testutil.ResourceWatcher) { } for _, obj := range defaults { - if err := k8sClient.Create(t.Context(), obj); client.IgnoreAlreadyExists(err) != nil { - t.Fatalf("Failed to create default template %s: %v", obj.GetName(), err) - } + err := k8sClient.Create(t.Context(), obj) + c.NoError( + client.IgnoreAlreadyExists(err), + "Failed to create default template %s: %v", + obj.GetName(), + err, + ) } return k8sClient, watcher @@ -230,9 +237,17 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Spec: &multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, }, Cells: []multigresv1alpha1.CellConfig{ - {Name: "zone-a", ZoneID: "use1-az1", Spec: &multigresv1alpha1.CellInlineSpec{ - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, - }}, + { + Name: "zone-a", + ZoneID: "use1-az1", + Spec: &multigresv1alpha1.CellInlineSpec{ + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(1)), + }, + }, + }, + }, }, Databases: []multigresv1alpha1.DatabaseConfig{ { @@ -245,12 +260,18 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Shards: []multigresv1alpha1.ShardConfig{{ Name: "0-inf", Spec: &multigresv1alpha1.ShardInlineSpec{ - Multiorch: multigresv1alpha1.MultiorchSpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + Multiorch: multigresv1alpha1.MultiorchSpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(1)), + }, + }, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ "primary": { ReplicasPerCell: ptr.To(int32(3)), Type: "readWrite", - Cells: []multigresv1alpha1.CellName{"zone-a"}, + Cells: []multigresv1alpha1.CellName{ + "zone-a", + }, }, }, }, @@ -272,10 +293,12 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, Spec: multigresv1alpha1.TopoServerSpec{ Etcd: &multigresv1alpha1.EtcdSpec{ - Image: "etcd:latest", - Replicas: ptr.To(resolver.DefaultEtcdReplicas), - RootPath: "/multigres/default/test-cluster/global", - Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultEtcdStorageSize}, + Image: "etcd:latest", + Replicas: ptr.To(resolver.DefaultEtcdReplicas), + RootPath: "/multigres/default/test-cluster/global", + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultEtcdStorageSize, + }, Resources: resolver.DefaultResourcesEtcd(), }, PVCDeletionPolicy: &multigresv1alpha1.PVCDeletionPolicy{ @@ -293,9 +316,13 @@ func TestMultigresCluster_HappyPath(t *testing.T) { OwnerReferences: clusterOwnerRefs(t, clusterName), }, Spec: appsv1.DeploymentSpec{ - Replicas: ptr.To(resolver.DefaultAdminReplicas), // Matches default in test input + Replicas: ptr.To( + resolver.DefaultAdminReplicas, + ), // Matches default in test input Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, clusterName, "multiadmin", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, clusterName, "multiadmin", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -305,7 +332,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, }, Spec: corev1.PodSpec{ - ImagePullSecrets: []corev1.LocalObjectReference{{Name: "pull-secret"}}, + ImagePullSecrets: []corev1.LocalObjectReference{ + {Name: "pull-secret"}, + }, Containers: []corev1.Container{ { Name: "multiadmin", @@ -391,7 +420,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultMultiadminWebReplicas), // Defaults Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, clusterName, "multiadmin-web", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, clusterName, "multiadmin-web", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -401,7 +432,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, }, Spec: corev1.PodSpec{ - ImagePullSecrets: []corev1.LocalObjectReference{{Name: "pull-secret"}}, + ImagePullSecrets: []corev1.LocalObjectReference{ + {Name: "pull-secret"}, + }, Containers: []corev1.Container{ { Name: "multiadmin-web", @@ -416,8 +449,14 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Resources: resolver.DefaultResourcesAdminWeb(), Env: []corev1.EnvVar{ {Name: "HOSTNAME", Value: "::"}, - {Name: "MULTIADMIN_API_URL", Value: "http://" + clusterName + "-multiadmin:18000"}, - {Name: "POSTGRES_HOST", Value: clusterName + "-multigateway"}, + { + Name: "MULTIADMIN_API_URL", + Value: "http://" + clusterName + "-multiadmin:18000", + }, + { + Name: "POSTGRES_HOST", + Value: clusterName + "-multigateway", + }, {Name: "POSTGRES_PORT", Value: "5432"}, {Name: "POSTGRES_DATABASE", Value: "postgres"}, {Name: "POSTGRES_USER", Value: "postgres"}, @@ -583,15 +622,26 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Type: "readWrite", Cells: []multigresv1alpha1.CellName{"zone-a"}, // FIX: Expect defaults for pool resources - Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultEtcdStorageSize}, - Postgres: multigresv1alpha1.ContainerConfig{Resources: resolver.DefaultResourcesPostgres()}, - Multipooler: multigresv1alpha1.ContainerConfig{Resources: resolver.DefaultResourcesPooler()}, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultEtcdStorageSize, + }, + Postgres: multigresv1alpha1.ContainerConfig{ + Resources: resolver.DefaultResourcesPostgres(), + }, + Multipooler: multigresv1alpha1.ContainerConfig{ + Resources: resolver.DefaultResourcesPooler(), + }, }, }, PVCDeletionPolicy: nil, // Shard-level policy is nil (inherited) Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, }, }, }, @@ -600,10 +650,17 @@ func TestMultigresCluster_HappyPath(t *testing.T) { WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, + }, + TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{ + Enabled: ptr.To(true), }, - TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{Enabled: ptr.To(true)}, DurabilityPolicy: "AT_LEAST_2", PostgresSuperuser: "postgres", }, @@ -644,10 +701,12 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, Spec: multigresv1alpha1.TopoServerSpec{ Etcd: &multigresv1alpha1.EtcdSpec{ - Image: "etcd:default", - Replicas: ptr.To(resolver.DefaultEtcdReplicas), - RootPath: "/multigres/default/minimal-cluster/global", - Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultEtcdStorageSize}, + Image: "etcd:default", + Replicas: ptr.To(resolver.DefaultEtcdReplicas), + RootPath: "/multigres/default/minimal-cluster/global", + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultEtcdStorageSize, + }, Resources: resolver.DefaultResourcesEtcd(), }, PVCDeletionPolicy: &multigresv1alpha1.PVCDeletionPolicy{ @@ -667,7 +726,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultAdminReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "minimal-cluster", "multiadmin", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "minimal-cluster", "multiadmin", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -762,7 +823,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultMultiadminWebReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "minimal-cluster", "multiadmin-web", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "minimal-cluster", "multiadmin-web", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -786,8 +849,14 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Resources: resolver.DefaultResourcesAdminWeb(), Env: []corev1.EnvVar{ {Name: "HOSTNAME", Value: "::"}, - {Name: "MULTIADMIN_API_URL", Value: "http://minimal-cluster-multiadmin:18000"}, - {Name: "POSTGRES_HOST", Value: "minimal-cluster-multigateway"}, + { + Name: "MULTIADMIN_API_URL", + Value: "http://minimal-cluster-multiadmin:18000", + }, + { + Name: "POSTGRES_HOST", + Value: "minimal-cluster-multigateway", + }, {Name: "POSTGRES_PORT", Value: "5432"}, {Name: "POSTGRES_DATABASE", Value: "postgres"}, {Name: "POSTGRES_USER", Value: "postgres"}, @@ -952,7 +1021,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Cells: []multigresv1alpha1.CellName{"zone-a"}, // Single-cell pool defaults to 2 (AT_LEAST_2 minimum). ReplicasPerCell: ptr.To(int32(2)), - Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultEtcdStorageSize}, // "1Gi" + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultEtcdStorageSize, + }, // "1Gi" Postgres: multigresv1alpha1.ContainerConfig{ Resources: resolver.DefaultResourcesPostgres(), }, @@ -963,8 +1034,13 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, PVCDeletionPolicy: nil, // Shard-level is nil Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, }, }, }, @@ -973,10 +1049,17 @@ func TestMultigresCluster_HappyPath(t *testing.T) { WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, + }, + TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{ + Enabled: ptr.To(true), }, - TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{Enabled: ptr.To(true)}, DurabilityPolicy: "AT_LEAST_2", PostgresSuperuser: "postgres", }, @@ -1012,10 +1095,12 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, Spec: multigresv1alpha1.TopoServerSpec{ Etcd: &multigresv1alpha1.EtcdSpec{ - Image: "etcd:default", - Replicas: ptr.To(resolver.DefaultEtcdReplicas), - RootPath: "/multigres/default/lazy-cluster/global", - Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultEtcdStorageSize}, + Image: "etcd:default", + Replicas: ptr.To(resolver.DefaultEtcdReplicas), + RootPath: "/multigres/default/lazy-cluster/global", + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultEtcdStorageSize, + }, Resources: resolver.DefaultResourcesEtcd(), }, PVCDeletionPolicy: &multigresv1alpha1.PVCDeletionPolicy{ @@ -1035,7 +1120,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultAdminReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "lazy-cluster", "multiadmin", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "lazy-cluster", "multiadmin", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -1130,7 +1217,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(resolver.DefaultMultiadminWebReplicas), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(clusterLabels(t, "lazy-cluster", "multiadmin-web", "")), + MatchLabels: metadata.GetSelectorLabels( + clusterLabels(t, "lazy-cluster", "multiadmin-web", ""), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ @@ -1154,8 +1243,14 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Resources: resolver.DefaultResourcesAdminWeb(), Env: []corev1.EnvVar{ {Name: "HOSTNAME", Value: "::"}, - {Name: "MULTIADMIN_API_URL", Value: "http://lazy-cluster-multiadmin:18000"}, - {Name: "POSTGRES_HOST", Value: "lazy-cluster-multigateway"}, + { + Name: "MULTIADMIN_API_URL", + Value: "http://lazy-cluster-multiadmin:18000", + }, + { + Name: "POSTGRES_HOST", + Value: "lazy-cluster-multigateway", + }, {Name: "POSTGRES_PORT", Value: "5432"}, {Name: "POSTGRES_DATABASE", Value: "postgres"}, {Name: "POSTGRES_USER", Value: "postgres"}, @@ -1320,7 +1415,9 @@ func TestMultigresCluster_HappyPath(t *testing.T) { Cells: []multigresv1alpha1.CellName{"zone-a"}, // Single-cell pool defaults to 2 (AT_LEAST_2 minimum). ReplicasPerCell: ptr.To(int32(2)), - Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultEtcdStorageSize}, // "1Gi" + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultEtcdStorageSize, + }, // "1Gi" Postgres: multigresv1alpha1.ContainerConfig{ Resources: resolver.DefaultResourcesPostgres(), }, @@ -1331,8 +1428,13 @@ func TestMultigresCluster_HappyPath(t *testing.T) { }, PVCDeletionPolicy: nil, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, }, }, }, @@ -1341,10 +1443,17 @@ func TestMultigresCluster_HappyPath(t *testing.T) { WhenScaled: multigresv1alpha1.DeletePVCRetentionPolicy, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: resolver.DefaultBackupPath, Storage: multigresv1alpha1.StorageSpec{Size: resolver.DefaultBackupStorageSize}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: resolver.DefaultBackupPath, + Storage: multigresv1alpha1.StorageSpec{ + Size: resolver.DefaultBackupStorageSize, + }, + }, + }, + TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{ + Enabled: ptr.To(true), }, - TopologyPruning: &multigresv1alpha1.TopologyPruningConfig{Enabled: ptr.To(true)}, DurabilityPolicy: "AT_LEAST_2", PostgresSuperuser: "postgres", }, @@ -1356,16 +1465,26 @@ func TestMultigresCluster_HappyPath(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setupIntegration(t) // Patch wantResources with hashed names for _, obj := range tc.wantResources { if cell, ok := obj.(*multigresv1alpha1.Cell); ok { - hashedName := nameutil.JoinWithConstraints(nameutil.DefaultConstraints, tc.cluster.Name, string(cell.Spec.Name)) + hashedName := nameutil.JoinWithConstraints( + nameutil.DefaultConstraints, + tc.cluster.Name, + string(cell.Spec.Name), + ) cell.Name = hashedName } if tg, ok := obj.(*multigresv1alpha1.TableGroup); ok { - hashedName := nameutil.JoinWithConstraints(nameutil.DefaultConstraints, tc.cluster.Name, string(tg.Spec.DatabaseName), string(tg.Spec.TableGroupName)) + hashedName := nameutil.JoinWithConstraints( + nameutil.DefaultConstraints, + tc.cluster.Name, + string(tg.Spec.DatabaseName), + string(tg.Spec.TableGroupName), + ) tg.Name = hashedName setTestTableGroupPostgresPasswordSecretRef(tg) } @@ -1373,15 +1492,11 @@ func TestMultigresCluster_HappyPath(t *testing.T) { // Create Cluster setTestPostgresPasswordSecretRef(tc.cluster) - if err := k8sClient.Create(t.Context(), tc.cluster); err != nil { - t.Fatalf("Failed to create the initial cluster, %v", err) - } + c.Require(). + NoError(k8sClient.Create(t.Context(), tc.cluster), "Failed to create the initial cluster") // Assert Resources - if err := watcher.WaitForMatch(tc.wantResources...); err != nil { - t.Errorf("Resources mismatch:\n%v", err) - } - + c.NoError(watcher.WaitForMatch(tc.wantResources...), "Resources mismatch:\n") }) } } diff --git a/pkg/cluster-handler/controller/multigrescluster/integration_validation_test.go b/pkg/cluster-handler/controller/multigrescluster/integration_validation_test.go index a95a8b23..2462a109 100644 --- a/pkg/cluster-handler/controller/multigrescluster/integration_validation_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/integration_validation_test.go @@ -10,6 +10,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/utils/ptr" + + "github.com/multigres/testkit/assert" ) func TestMultigresCluster_Validation(t *testing.T) { @@ -17,6 +19,7 @@ func TestMultigresCluster_Validation(t *testing.T) { t.Run("Explicit Empty GlobalTopoServer (Should Fail)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, _ := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "fail-empty-struct", Namespace: testNamespace}, @@ -26,17 +29,20 @@ func TestMultigresCluster_Validation(t *testing.T) { } setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(t.Context(), cluster) - if err == nil { - t.Fatal("Expected error creating cluster with empty GlobalTopoServer struct, got nil") - } + c.Require(). + Error(err, "Expected error creating cluster with empty GlobalTopoServer struct, got nil") // We expect CEL validation error here - if !strings.Contains(err.Error(), "must specify exactly one of") { - t.Errorf("Expected CEL validation error, got: %v", err) - } + c.StrContains( + err.Error(), + "must specify exactly one of", + "Expected CEL validation error, got: %v", + err, + ) }) t.Run("Multiadmin XOR Violation (Should Fail)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, _ := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "fail-xor-admin", Namespace: testNamespace}, @@ -49,16 +55,19 @@ func TestMultigresCluster_Validation(t *testing.T) { } setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(t.Context(), cluster) - if err == nil { - t.Fatal("Expected error creating cluster with Multiadmin XOR violation, got nil") - } - if !strings.Contains(err.Error(), "cannot specify both") { - t.Errorf("Expected CEL validation error, got: %v", err) - } + c.Require(). + Error(err, "Expected error creating cluster with Multiadmin XOR violation, got nil") + c.StrContains( + err.Error(), + "cannot specify both", + "Expected CEL validation error, got: %v", + err, + ) }) t.Run("Multiple Databases (Should Fail)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, _ := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "fail-multi-db", Namespace: testNamespace}, @@ -71,19 +80,21 @@ func TestMultigresCluster_Validation(t *testing.T) { } setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(t.Context(), cluster) - if err == nil { - t.Fatal("Expected error creating cluster with multiple databases, got nil") - } + c.Require().Error(err, "Expected error creating cluster with multiple databases, got nil") // Expect MaxItems=1 or system database rule violation - if !strings.Contains(err.Error(), "Invalid value") && !strings.Contains(err.Error(), "only the single system database") { - t.Errorf("Expected validation error regarding DB count/rules, got: %v", err) - } + c.False( + !strings.Contains(err.Error(), "Invalid value") && + !strings.Contains(err.Error(), "only the single system database"), + "Expected validation error regarding DB count/rules, got: %v", + err, + ) }) // Topology access is only as narrow as the CA that backs it, so an enabled // configuration has to name its issuer rather than inherit one. t.Run("Topology TLS Without Issuer (Should Fail)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, _ := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "fail-topo-tls", Namespace: testNamespace}, @@ -93,12 +104,14 @@ func TestMultigresCluster_Validation(t *testing.T) { } setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(t.Context(), cluster) - if err == nil { - t.Fatal("Expected error creating cluster with topology TLS and no issuer, got nil") - } - if !strings.Contains(err.Error(), "issuerName is required when topology TLS is enabled") { - t.Errorf("Expected CEL validation error, got: %v", err) - } + c.Require(). + Error(err, "Expected error creating cluster with topology TLS and no issuer, got nil") + c.StrContains( + err.Error(), + "issuerName is required when topology TLS is enabled", + "Expected CEL validation error, got: %v", + err, + ) }) t.Run("Topology TLS With Issuer (Should Pass)", func(t *testing.T) { @@ -114,9 +127,8 @@ func TestMultigresCluster_Validation(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Expected cluster with topology TLS and an issuer to be accepted, got: %v", err) - } + assert.NewAborting(t). + NoError(k8sClient.Create(t.Context(), cluster), "Expected cluster with topology TLS and an issuer to be accepted, got") }) // Disabled and absent configurations stay valid without an issuer, so the @@ -131,15 +143,15 @@ func TestMultigresCluster_Validation(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Expected cluster with topology TLS disabled to be accepted, got: %v", err) - } + assert.NewAborting(t). + NoError(k8sClient.Create(t.Context(), cluster), "Expected cluster with topology TLS disabled to be accepted, got") }) // Enablement is fixed at creation: flipping it on a running cluster cannot be // rolled out safely, so the API rejects the transition. t.Run("Enabling Topology TLS After Creation (Should Fail)", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) k8sClient, _ := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "mutate-topo-tls", Namespace: testNamespace}, @@ -148,17 +160,19 @@ func TestMultigresCluster_Validation(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Expected initial cluster to be accepted, got: %v", err) - } + c.NoError( + k8sClient.Create(t.Context(), cluster), + "Expected initial cluster to be accepted, got", + ) cluster.Spec.TopoTLS = &multigresv1alpha1.TopoTLSConfig{ Enabled: ptr.To(true), IssuerName: "multigres-infra-issuer", } - if err := k8sClient.Update(t.Context(), cluster); err == nil { - t.Fatal("Expected enabling topology TLS after creation to be rejected") - } + c.Error( + k8sClient.Update(t.Context(), cluster), + "Expected enabling topology TLS after creation to be rejected", + ) }) // Rotating the issuer rotates the CA every certificate chains to, which the @@ -166,6 +180,7 @@ func TestMultigresCluster_Validation(t *testing.T) { // too. t.Run("Changing Topology TLS Issuer After Creation (Should Fail)", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) k8sClient, _ := setupIntegration(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "rotate-topo-issuer", Namespace: testNamespace}, @@ -177,13 +192,15 @@ func TestMultigresCluster_Validation(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("Expected initial cluster to be accepted, got: %v", err) - } + c.NoError( + k8sClient.Create(t.Context(), cluster), + "Expected initial cluster to be accepted, got", + ) cluster.Spec.TopoTLS.IssuerName = "issuer-b" - if err := k8sClient.Update(t.Context(), cluster); err == nil { - t.Fatal("Expected changing the topology TLS issuer after creation to be rejected") - } + c.Error( + k8sClient.Update(t.Context(), cluster), + "Expected changing the topology TLS issuer after creation to be rejected", + ) }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/multigrescluster_controller_test.go b/pkg/cluster-handler/controller/multigrescluster/multigrescluster_controller_test.go index 2b70fe5e..ea581330 100644 --- a/pkg/cluster-handler/controller/multigrescluster/multigrescluster_controller_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/multigrescluster_controller_test.go @@ -4,7 +4,6 @@ import ( "context" "errors" "fmt" - "slices" "strconv" "strings" "testing" @@ -34,6 +33,8 @@ import ( "github.com/multigres/multigres-operator/pkg/util/name" "github.com/multigres/multigres/go/common/topoclient" "github.com/multigres/multigres/go/common/topoclient/memorytopo" + + "github.com/multigres/testkit/assert" ) // ============================================================================ @@ -67,6 +68,7 @@ func runReconcileTest(t *testing.T, tests map[string]reconcileTestCase) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) // Default to all standard templates if existingObjects is nil objects := tc.existingObjects @@ -108,29 +110,23 @@ func runReconcileTest(t *testing.T, tests map[string]reconcileTestCase) { check, ) if apierrors.IsNotFound(err) { - if err := baseClient.Create(t.Context(), cluster); err != nil { - t.Fatalf("failed to create initial cluster: %v", err) - } + c.Require(). + NoError(baseClient.Create(t.Context(), cluster), "failed to create initial cluster") if shouldDelete { // Simulate deletion workflow - if err := baseClient.Get( + c.Require().NoError(baseClient.Get( t.Context(), types.NamespacedName{Name: cluster.Name, Namespace: cluster.Namespace}, cluster, - ); err != nil { - t.Fatalf("failed to refresh cluster before delete: %v", err) - } - if err := baseClient.Delete(t.Context(), cluster); err != nil { - t.Fatalf("failed to set deletion timestamp: %v", err) - } - if err := baseClient.Get( + ), "failed to refresh cluster before delete") + c.Require(). + NoError(baseClient.Delete(t.Context(), cluster), "failed to set deletion timestamp") + c.Require().NoError(baseClient.Get( t.Context(), types.NamespacedName{Name: cluster.Name, Namespace: cluster.Namespace}, cluster, - ); err != nil { - t.Fatalf("failed to refresh cluster after deletion: %v", err) - } + ), "failed to refresh cluster after deletion") } } } @@ -197,13 +193,12 @@ func runReconcileTest(t *testing.T, tests map[string]reconcileTestCase) { break } } - if !found { - t.Errorf( - "Expected event containing %q not found. Got events: %v", - want, - gotEvents, - ) - } + c.True( + found, + "Expected event containing %q not found. Got events: %v", + want, + gotEvents, + ) } } @@ -364,6 +359,7 @@ func TestHandleDeletionReleasesPoolerClient(t *testing.T) { } t.Run("after successful cleanup", func(t *testing.T) { + c := assert.NewAborting(t) forgotten := make(chan types.NamespacedName, 1) r := &MultigresClusterReconciler{ Client: fake.NewClientBuilder().WithScheme(setupScheme()).Build(), @@ -373,20 +369,18 @@ func TestHandleDeletionReleasesPoolerClient(t *testing.T) { }), } - if _, err := r.handleDeletion(t.Context(), cluster.DeepCopy()); err != nil { - t.Fatalf("handleDeletion() error = %v", err) - } + _, err := r.handleDeletion(t.Context(), cluster.DeepCopy()) + c.NoError(err, "handleDeletion() error =") select { case got := <-forgotten: - if got != clusterKey { - t.Fatalf("forgot cluster %v, want %v", got, clusterKey) - } + c.Eq(clusterKey, got, "forgot cluster") default: t.Fatal("pooler client cache was not notified") } }) t.Run("not when cleanup fails", func(t *testing.T) { + c := assert.NewAborting(t) baseClient := fake.NewClientBuilder().WithScheme(setupScheme()).Build() failingClient := testutil.NewFakeClientWithFailures(baseClient, &testutil.FailureConfig{ OnList: testutil.FailObjListAfterNCalls(0, errors.New("list failed")), @@ -400,12 +394,9 @@ func TestHandleDeletionReleasesPoolerClient(t *testing.T) { }), } - if _, err := r.handleDeletion(t.Context(), cluster.DeepCopy()); err == nil { - t.Fatal("handleDeletion() error = nil, want cleanup error") - } - if forgotten { - t.Fatal("pooler client cache notified after failed cleanup") - } + _, err := r.handleDeletion(t.Context(), cluster.DeepCopy()) + c.Error(err, "handleDeletion() error = nil, want cleanup error") + c.False(forgotten, "pooler client cache notified after failed cleanup") }) } @@ -465,6 +456,7 @@ func TestHandleDeletionWaitsForChildren(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewAborting(t) cluster := newCluster() forgotten := false r := &MultigresClusterReconciler{ @@ -479,31 +471,22 @@ func TestHandleDeletionWaitsForChildren(t *testing.T) { } result, err := r.handleDeletion(t.Context(), cluster.DeepCopy()) - if err != nil { - t.Fatalf("handleDeletion() error = %v", err) - } - if result.RequeueAfter != childDeletionRequeueDelay { - t.Fatalf( - "RequeueAfter = %v, want %v", - result.RequeueAfter, - childDeletionRequeueDelay, - ) - } - if forgotten { - t.Fatal("pooler client cache notified while children remain") - } + c.NoError(err, "handleDeletion() error =") + c.Eq(childDeletionRequeueDelay, result.RequeueAfter, "RequeueAfter") + c.False(forgotten, "pooler client cache notified while children remain") got := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(t.Context(), clusterKey, got); err != nil { - t.Fatalf("failed to get cluster: %v", err) - } - if !slices.Contains(got.Finalizers, multigresv1alpha1.FinalizerClusterCleanup) { - t.Fatal("cluster cleanup finalizer removed while children remain") - } + c.NoError(r.Get(t.Context(), clusterKey, got), "failed to get cluster") + c.Contains( + got.Finalizers, + multigresv1alpha1.FinalizerClusterCleanup, + "cluster cleanup finalizer removed while children remain", + ) }) } t.Run("no children", func(t *testing.T) { + c := assert.NewAborting(t) cluster := newCluster() r := &MultigresClusterReconciler{ Client: fake.NewClientBuilder(). @@ -518,22 +501,14 @@ func TestHandleDeletionWaitsForChildren(t *testing.T) { }) fetched := &multigresv1alpha1.MultigresCluster{} - if err := r.Get(t.Context(), clusterKey, fetched); err != nil { - t.Fatalf("failed to get cluster: %v", err) - } + c.NoError(r.Get(t.Context(), clusterKey, fetched), "failed to get cluster") result, err := r.handleDeletion(t.Context(), fetched) - if err != nil { - t.Fatalf("handleDeletion() error = %v", err) - } - if result.RequeueAfter != 0 { - t.Fatalf("RequeueAfter = %v, want 0", result.RequeueAfter) - } + c.NoError(err, "handleDeletion() error =") + c.Eq(0, result.RequeueAfter, "RequeueAfter") select { case got := <-forgotten: - if got != clusterKey { - t.Fatalf("forgot cluster %v, want %v", got, clusterKey) - } + c.Eq(clusterKey, got, "forgot cluster") default: t.Fatal("pooler client cache was not notified") } @@ -542,9 +517,11 @@ func TestHandleDeletionWaitsForChildren(t *testing.T) { err = r.Get(t.Context(), clusterKey, got) switch { case err == nil: - if slices.Contains(got.Finalizers, multigresv1alpha1.FinalizerClusterCleanup) { - t.Fatal("cluster cleanup finalizer not removed") - } + c.NotContains( + got.Finalizers, + multigresv1alpha1.FinalizerClusterCleanup, + "cluster cleanup finalizer not removed", + ) case apierrors.IsNotFound(err): default: t.Fatalf("failed to get cluster: %v", err) @@ -553,6 +530,7 @@ func TestHandleDeletionWaitsForChildren(t *testing.T) { } func TestReconcileNotFoundReleasesPoolerClient(t *testing.T) { + c := assert.NewAborting(t) key := types.NamespacedName{Name: "missing", Namespace: "test-ns"} forgotten := make(chan types.NamespacedName, 1) r := &MultigresClusterReconciler{ @@ -562,14 +540,11 @@ func TestReconcileNotFoundReleasesPoolerClient(t *testing.T) { }), } - if _, err := r.Reconcile(t.Context(), ctrl.Request{NamespacedName: key}); err != nil { - t.Fatalf("Reconcile() error = %v", err) - } + _, err := r.Reconcile(t.Context(), ctrl.Request{NamespacedName: key}) + c.NoError(err, "Reconcile() error =") select { case got := <-forgotten: - if got != key { - t.Fatalf("forgot cluster %v, want %v", got, key) - } + c.Eq(key, got, "forgot cluster") default: t.Fatal("pooler client cache was not notified for missing cluster") } @@ -590,6 +565,7 @@ func TestMultigresClusterReconciler_Lifecycle(t *testing.T) { "Create: Full Cluster Creation - Verify Images and Wiring": { expectedEvents: []string{"Normal Synced Successfully reconciled MultigresCluster"}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() // Verify Cell (Basic wiring check) cell := &multigresv1alpha1.Cell{} @@ -598,18 +574,15 @@ func TestMultigresClusterReconciler_Lifecycle(t *testing.T) { clusterName, "zone-a", ) - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: cellName, Namespace: namespace}, cell, - ); err != nil { - t.Fatalf("Expected Cell %s to exist: %v", cellName, err) - } - if got, want := cell.Spec.Images.Multigateway, multigresv1alpha1.ImageRef( + ), "Expected Cell %s to exist", cellName) + got, want := cell.Spec.Images.Multigateway, multigresv1alpha1.ImageRef( "gateway:latest", - ); got != want { - t.Errorf("Cell image mismatch got %q, want %q", got, want) - } + ) + ck.Eq(want, got, "Cell image mismatch got") }, }, @@ -706,18 +679,20 @@ func TestMultigresClusterReconciler_Lifecycle(t *testing.T) { "Normal PendingDeletion Marked TableGroup", }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) tg := &multigresv1alpha1.TableGroup{} err := c.Get( t.Context(), types.NamespacedName{Name: clusterName + "-orphan-tg", Namespace: namespace}, tg, ) - if err != nil { - t.Fatalf("Expected orphan TG to still exist with PendingDeletion, got: %v", err) - } - if tg.Annotations[multigresv1alpha1.AnnotationPendingDeletion] == "" { - t.Error("Expected orphan TableGroup to have PendingDeletion annotation") - } + ck.Require(). + NoError(err, "Expected orphan TG to still exist with PendingDeletion, got") + ck.NotEq( + "", + tg.Annotations[multigresv1alpha1.AnnotationPendingDeletion], + "Expected orphan TableGroup to have PendingDeletion annotation", + ) }, }, "Object Not Found (Clean Exit)": { @@ -805,31 +780,28 @@ func TestMultigresClusterReconciler_Lifecycle(t *testing.T) { existingObjects: []client.Object{coreTpl, cellTpl, shardTpl}, expectedEvents: []string{"Normal Synced Successfully reconciled MultigresCluster"}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) cell := &multigresv1alpha1.Cell{} cellName := name.JoinWithConstraints( name.DefaultConstraints, clusterName, "zone-a", ) - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: cellName, Namespace: namespace}, cell, - ); err != nil { - t.Fatalf("Expected Cell %s to exist: %v", cellName, err) - } - if cell.Spec.GlobalTopoServer.Address != "http://external:2379" { - t.Errorf( - "Expected external address http://external:2379, got %s", - cell.Spec.GlobalTopoServer.Address, - ) - } - if cell.Spec.GlobalTopoServer.RootPath != "/custom/root" { - t.Errorf( - "Expected external root path /custom/root, got %s", - cell.Spec.GlobalTopoServer.RootPath, - ) - } + ), "Expected Cell %s to exist", cellName) + ck.Eq( + "http://external:2379", + cell.Spec.GlobalTopoServer.Address, + "Expected external address http://external:2379, got", + ) + ck.Eq( + "/custom/root", + cell.Spec.GlobalTopoServer.RootPath, + "Expected external root path /custom/root, got", + ) }, }, "Success: Early Return on Deletion": { @@ -1098,6 +1070,7 @@ func TestEnqueueRequestsFromTemplate(t *testing.T) { } t.Run("CoreTemplate matches only referencing and nil-status clusters", func(t *testing.T) { + c := assert.NewCollecting(t) tpl := &multigresv1alpha1.CoreTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "prod-core", Namespace: "default"}, } @@ -1107,24 +1080,24 @@ func TestEnqueueRequestsFromTemplate(t *testing.T) { for _, req := range requests { names[req.Name] = true } - if !names["cluster-core"] { - t.Error("Expected cluster-core (references prod-core) to be enqueued") - } - if !names["cluster-nil"] { - t.Error("Expected cluster-nil (nil status) to be enqueued") - } - if names["cluster-shard"] { - t.Error("cluster-shard should not be enqueued for CoreTemplate change") - } - if names["cluster-other"] { - t.Error("cluster-other (different namespace) should not be enqueued") - } - if len(requests) != 2 { - t.Errorf("Expected 2 requests, got %d: %v", len(requests), names) - } + c.False( + !names["cluster-core"], + "Expected cluster-core (references prod-core) to be enqueued", + ) + c.False(!names["cluster-nil"], "Expected cluster-nil (nil status) to be enqueued") + c.False( + names["cluster-shard"], + "cluster-shard should not be enqueued for CoreTemplate change", + ) + c.False( + names["cluster-other"], + "cluster-other (different namespace) should not be enqueued", + ) + c.Len(requests, 2, "Expected 2 requests, got %d: %v", len(requests), names) }) t.Run("ShardTemplate matches only referencing and nil-status clusters", func(t *testing.T) { + c := assert.NewCollecting(t) tpl := &multigresv1alpha1.ShardTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "prod-shard", Namespace: "default"}, } @@ -1134,18 +1107,13 @@ func TestEnqueueRequestsFromTemplate(t *testing.T) { for _, req := range requests { names[req.Name] = true } - if !names["cluster-shard"] { - t.Error("Expected cluster-shard to be enqueued") - } - if !names["cluster-nil"] { - t.Error("Expected cluster-nil (nil status) to be enqueued") - } - if names["cluster-core"] { - t.Error("cluster-core should not be enqueued for ShardTemplate change") - } - if len(requests) != 2 { - t.Errorf("Expected 2 requests, got %d: %v", len(requests), names) - } + c.False(!names["cluster-shard"], "Expected cluster-shard to be enqueued") + c.False(!names["cluster-nil"], "Expected cluster-nil (nil status) to be enqueued") + c.False( + names["cluster-core"], + "cluster-core should not be enqueued for ShardTemplate change", + ) + c.Len(requests, 2, "Expected 2 requests, got %d: %v", len(requests), names) }) t.Run("Unmatched template enqueues only nil-status clusters", func(t *testing.T) { @@ -1154,9 +1122,8 @@ func TestEnqueueRequestsFromTemplate(t *testing.T) { } requests := r.enqueueRequestsFromTemplate(context.Background(), tpl) - if len(requests) != 1 { - t.Errorf("Expected 1 request (nil-status cluster only), got %d", len(requests)) - } + assert.NewCollecting(t). + Len(requests, 1, "Expected 1 request (nil-status cluster only), got %d", len(requests)) if len(requests) == 1 && requests[0].Name != "cluster-nil" { t.Errorf("Expected cluster-nil, got %s", requests[0].Name) } @@ -1167,9 +1134,7 @@ func TestEnqueueRequestsFromTemplate(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "not-a-template", Namespace: "default"}, } requests := r.enqueueRequestsFromTemplate(context.Background(), unknown) - if requests != nil { - t.Errorf("Expected nil for unknown object type, got %v", requests) - } + assert.NewCollecting(t).Nil(requests, "Expected nil for unknown object type, got") }) t.Run("List error returns empty", func(t *testing.T) { @@ -1183,9 +1148,8 @@ func TestEnqueueRequestsFromTemplate(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "prod-core", Namespace: "default"}, } requests := r.enqueueRequestsFromTemplate(context.Background(), tpl) - if len(requests) != 0 { - t.Errorf("Expected 0 requests on list error, got %d", len(requests)) - } + assert.NewCollecting(t). + Empty(requests, "Expected 0 requests on list error, got %d", len(requests)) }) } @@ -1202,9 +1166,8 @@ func TestTemplateKindFromObject(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := templateKindFromObject(tt.obj); got != tt.want { - t.Errorf("templateKindFromObject() = %q, want %q", got, tt.want) - } + assert.NewCollecting(t). + Eq(tt.want, templateKindFromObject(tt.obj), "templateKindFromObject()") }) } } @@ -1236,15 +1199,15 @@ func TestReferencesTemplate(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := referencesTemplate(tt.rt, tt.kind, tt.tpl); got != tt.want { - t.Errorf("referencesTemplate() = %v, want %v", got, tt.want) - } + assert.NewCollecting(t). + Eq(tt.want, referencesTemplate(tt.rt, tt.kind, tt.tpl), "referencesTemplate()") }) } } func TestCollectResolvedTemplates(t *testing.T) { t.Run("All template refs populated", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ Spec: multigresv1alpha1.MultigresClusterSpec{ TemplateDefaults: multigresv1alpha1.TemplateDefaults{ @@ -1279,17 +1242,11 @@ func TestCollectResolvedTemplates(t *testing.T) { rt := collectResolvedTemplates(cluster) wantCore := []multigresv1alpha1.TemplateRef{"admin-core", "default-core", "gts-core"} - if !slices.Equal(rt.CoreTemplates, wantCore) { - t.Errorf("CoreTemplates = %v, want %v", rt.CoreTemplates, wantCore) - } + c.EqDiff(wantCore, rt.CoreTemplates, "CoreTemplates") wantCell := []multigresv1alpha1.TemplateRef{"cell-ha", "cell-std", "default-cell"} - if !slices.Equal(rt.CellTemplates, wantCell) { - t.Errorf("CellTemplates = %v, want %v", rt.CellTemplates, wantCell) - } + c.EqDiff(wantCell, rt.CellTemplates, "CellTemplates") wantShard := []multigresv1alpha1.TemplateRef{"default-shard", "shard-prod"} - if !slices.Equal(rt.ShardTemplates, wantShard) { - t.Errorf("ShardTemplates = %v, want %v", rt.ShardTemplates, wantShard) - } + c.EqDiff(wantShard, rt.ShardTemplates, "ShardTemplates") }) t.Run("Duplicates are deduplicated", func(t *testing.T) { @@ -1338,6 +1295,7 @@ func TestCollectResolvedTemplates(t *testing.T) { }) t.Run("No templates (pure inline)", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ Spec: multigresv1alpha1.MultigresClusterSpec{ Cells: []multigresv1alpha1.CellConfig{ @@ -1348,15 +1306,9 @@ func TestCollectResolvedTemplates(t *testing.T) { rt := collectResolvedTemplates(cluster) - if len(rt.CoreTemplates) != 0 { - t.Errorf("CoreTemplates should be empty, got %v", rt.CoreTemplates) - } - if len(rt.CellTemplates) != 0 { - t.Errorf("CellTemplates should be empty, got %v", rt.CellTemplates) - } - if len(rt.ShardTemplates) != 0 { - t.Errorf("ShardTemplates should be empty, got %v", rt.ShardTemplates) - } + c.Empty(rt.CoreTemplates, "CoreTemplates should be empty, got") + c.Empty(rt.CellTemplates, "CellTemplates should be empty, got") + c.Empty(rt.ShardTemplates, "ShardTemplates should be empty, got") }) t.Run("MultiadminWeb templateRef included", func(t *testing.T) { @@ -1378,6 +1330,7 @@ func TestCollectResolvedTemplates(t *testing.T) { func TestCollectTrackingLabels(t *testing.T) { t.Run("All template kinds referenced", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ Spec: multigresv1alpha1.MultigresClusterSpec{ TemplateDefaults: multigresv1alpha1.TemplateDefaults{ @@ -1390,15 +1343,9 @@ func TestCollectTrackingLabels(t *testing.T) { labels := collectTrackingLabels(cluster) - if labels[metadata.LabelUsesCoreTemplate] != "true" { - t.Error("Expected uses-core-template=true") - } - if labels[metadata.LabelUsesCellTemplate] != "true" { - t.Error("Expected uses-cell-template=true") - } - if labels[metadata.LabelUsesShardTemplate] != "true" { - t.Error("Expected uses-shard-template=true") - } + c.Eq("true", labels[metadata.LabelUsesCoreTemplate], "Expected uses-core-template=true") + c.Eq("true", labels[metadata.LabelUsesCellTemplate], "Expected uses-cell-template=true") + c.Eq("true", labels[metadata.LabelUsesShardTemplate], "Expected uses-shard-template=true") }) t.Run("No templates (pure inline)", func(t *testing.T) { @@ -1412,12 +1359,12 @@ func TestCollectTrackingLabels(t *testing.T) { labels := collectTrackingLabels(cluster) - if len(labels) != 0 { - t.Errorf("Expected no tracking labels for inline-only cluster, got %v", labels) - } + assert.NewCollecting(t). + Empty(labels, "Expected no tracking labels for inline-only cluster, got") }) t.Run("Only cell template from per-cell ref", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ Spec: multigresv1alpha1.MultigresClusterSpec{ Cells: []multigresv1alpha1.CellConfig{ @@ -1431,15 +1378,13 @@ func TestCollectTrackingLabels(t *testing.T) { if _, ok := labels[metadata.LabelUsesCoreTemplate]; ok { t.Error("Unexpected uses-core-template label") } - if labels[metadata.LabelUsesCellTemplate] != "true" { - t.Error("Expected uses-cell-template=true") - } - if _, ok := labels[metadata.LabelUsesShardTemplate]; ok { - t.Error("Unexpected uses-shard-template label") - } + c.Eq("true", labels[metadata.LabelUsesCellTemplate], "Expected uses-cell-template=true") + _, ok := labels[metadata.LabelUsesShardTemplate] + c.False(ok, "Unexpected uses-shard-template label") }) t.Run("Only shard template from per-shard ref", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ Spec: multigresv1alpha1.MultigresClusterSpec{ Databases: []multigresv1alpha1.DatabaseConfig{ @@ -1461,12 +1406,9 @@ func TestCollectTrackingLabels(t *testing.T) { if _, ok := labels[metadata.LabelUsesCoreTemplate]; ok { t.Error("Unexpected uses-core-template label") } - if _, ok := labels[metadata.LabelUsesCellTemplate]; ok { - t.Error("Unexpected uses-cell-template label") - } - if labels[metadata.LabelUsesShardTemplate] != "true" { - t.Error("Expected uses-shard-template=true") - } + _, ok := labels[metadata.LabelUsesCellTemplate] + c.False(ok, "Unexpected uses-cell-template label") + c.Eq("true", labels[metadata.LabelUsesShardTemplate], "Expected uses-shard-template=true") }) t.Run("Core from GlobalTopoServer templateRef", func(t *testing.T) { @@ -1480,9 +1422,8 @@ func TestCollectTrackingLabels(t *testing.T) { labels := collectTrackingLabels(cluster) - if labels[metadata.LabelUsesCoreTemplate] != "true" { - t.Error("Expected uses-core-template=true from GlobalTopoServer ref") - } + assert.NewCollecting(t). + Eq("true", labels[metadata.LabelUsesCoreTemplate], "Expected uses-core-template=true from GlobalTopoServer ref") }) t.Run("Core from MultiadminWeb templateRef", func(t *testing.T) { @@ -1496,9 +1437,8 @@ func TestCollectTrackingLabels(t *testing.T) { labels := collectTrackingLabels(cluster) - if labels[metadata.LabelUsesCoreTemplate] != "true" { - t.Error("Expected uses-core-template=true from MultiadminWeb ref") - } + assert.NewCollecting(t). + Eq("true", labels[metadata.LabelUsesCoreTemplate], "Expected uses-core-template=true from MultiadminWeb ref") }) } @@ -1525,9 +1465,8 @@ func TestReconciler_PatchTrackingLabelsError(t *testing.T) { context.Background(), ctrl.Request{NamespacedName: types.NamespacedName{Name: clusterName, Namespace: namespace}}, ) - if err == nil || !strings.Contains(err.Error(), "failed to patch tracking labels") { - t.Errorf("Expected patch error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to patch tracking labels"), "Expected patch error, got %v", err) } func TestReconciler_TopologyFailure(t *testing.T) { @@ -1552,12 +1491,12 @@ func TestReconciler_TopologyFailure(t *testing.T) { context.Background(), ctrl.Request{NamespacedName: types.NamespacedName{Name: clusterName, Namespace: namespace}}, ) - if err == nil || !strings.Contains(err.Error(), "failed to open topology store") { - t.Errorf("Expected topology reconcile error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to open topology store"), "Expected topology reconcile error, got %v", err) } func TestReconciler_TopologyRequeueWhenUnavailable(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() coreTpl, cellTpl, shardTpl, baseCluster, clusterName, namespace := setupFixtures(t) baseCluster.CreationTimestamp = metav1.Now() // Trigger grace period @@ -1580,12 +1519,8 @@ func TestReconciler_TopologyRequeueWhenUnavailable(t *testing.T) { context.Background(), ctrl.Request{NamespacedName: types.NamespacedName{Name: clusterName, Namespace: namespace}}, ) - if err != nil { - t.Errorf("Expected nil error, got %v", err) - } - if res.RequeueAfter == 0 { - t.Errorf("Expected RequeueAfter > 0") - } + c.NoError(err, "Expected nil error, got") + c.NotEq(0, res.RequeueAfter, "Expected RequeueAfter > 0") } func TestReconciler_TopologyErrorWhenUnavailableExpired(t *testing.T) { @@ -1610,7 +1545,6 @@ func TestReconciler_TopologyErrorWhenUnavailableExpired(t *testing.T) { context.Background(), ctrl.Request{NamespacedName: types.NamespacedName{Name: clusterName, Namespace: namespace}}, ) - if err == nil || !strings.Contains(err.Error(), "topology server unavailable") { - t.Errorf("Expected topology unavailable error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "topology server unavailable"), "Expected topology unavailable error, got %v", err) } diff --git a/pkg/cluster-handler/controller/multigrescluster/reconcile_cells_test.go b/pkg/cluster-handler/controller/multigrescluster/reconcile_cells_test.go index c906d2f2..45908097 100644 --- a/pkg/cluster-handler/controller/multigrescluster/reconcile_cells_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/reconcile_cells_test.go @@ -13,6 +13,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + "github.com/multigres/testkit/assert" ) func TestReconcileCells_ErrorPaths(t *testing.T) { @@ -42,9 +44,7 @@ func TestReconcileCells_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to missing global topo, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing global topo, got nil") }) t.Run("Error: Resolve Cell Failed", func(t *testing.T) { @@ -75,9 +75,7 @@ func TestReconcileCells_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to missing cell template, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing cell template, got nil") }) t.Run("Error: List Existing Cells Failed", func(t *testing.T) { @@ -101,9 +99,8 @@ func TestReconcileCells_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || err.Error() != "failed to list existing cells: list error" { - t.Errorf("Expected 'list error', got %v", err) - } + assert.NewCollecting(t). + False(err == nil || err.Error() != "failed to list existing cells: list error", "Expected 'list error', got %v", err) }) t.Run("Error: Patch Cell Failed", func(t *testing.T) { @@ -135,9 +132,8 @@ func TestReconcileCells_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || err.Error() != "failed to apply cell 'zone-a': patch error" { - t.Errorf("Expected 'patch error', got %v", err) - } + assert.NewCollecting(t). + False(err == nil || err.Error() != "failed to apply cell 'zone-a': patch error", "Expected 'patch error', got %v", err) }) t.Run("Error: Set PendingDeletion on Orphaned Cell Failed", func(t *testing.T) { @@ -175,10 +171,8 @@ func TestReconcileCells_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || - err.Error() != "failed to set PendingDeletion on cell 'test-zone-orphan': patch error" { - t.Errorf("Expected PendingDeletion patch error, got %v", err) - } + assert.NewCollecting(t).False(err == nil || + err.Error() != "failed to set PendingDeletion on cell 'test-zone-orphan': patch error", "Expected PendingDeletion patch error, got %v", err) }) t.Run("Error: Build Cell Failed", func(t *testing.T) { @@ -207,9 +201,8 @@ func TestReconcileCells_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to build failure (scheme mismatch), got nil") - } + assert.NewCollecting(t). + Error(err, "Expected error due to build failure (scheme mismatch), got nil") }) } @@ -217,6 +210,7 @@ func TestReconcileCells_HappyPath(t *testing.T) { scheme := setupScheme() t.Run("Happy Path: Create Cells and Mark Orphan PendingDeletion", func(t *testing.T) { + ck := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "test", Namespace: "default"}, Spec: multigresv1alpha1.MultigresClusterSpec{ @@ -250,31 +244,25 @@ func TestReconcileCells_HappyPath(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err != nil { - t.Fatalf("Expected happy path success, got %v", err) - } - if !pending { - t.Error("Expected pending=true for orphan cell pending deletion") - } + ck.Require().NoError(err, "Expected happy path success, got") + ck.True(pending, "Expected pending=true for orphan cell pending deletion") // Verify "zone-new" created and orphan has PendingDeletion annotation cells := &multigresv1alpha1.CellList{} - if err := c.List(context.Background(), cells); err != nil { - t.Fatal(err) - } + ck.Require().NoError(c.List(context.Background(), cells)) foundNew := false for _, cell := range cells.Items { if cell.Spec.Name == "zone-new" { foundNew = true } if cell.Spec.Name == "zone-old" { - if cell.Annotations[multigresv1alpha1.AnnotationPendingDeletion] == "" { - t.Error("Expected orphan cell 'zone-old' to have PendingDeletion annotation") - } + ck.NotEq( + "", + cell.Annotations[multigresv1alpha1.AnnotationPendingDeletion], + "Expected orphan cell 'zone-old' to have PendingDeletion annotation", + ) } } - if !foundNew { - t.Error("Expected new cell 'zone-new' to be created") - } + ck.True(foundNew, "Expected new cell 'zone-new' to be created") }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/reconcile_database_test.go b/pkg/cluster-handler/controller/multigrescluster/reconcile_database_test.go index 7a4c02e5..4f196c84 100644 --- a/pkg/cluster-handler/controller/multigrescluster/reconcile_database_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/reconcile_database_test.go @@ -17,6 +17,8 @@ import ( "k8s.io/apimachinery/pkg/types" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestReconcile_Databases(t *testing.T) { @@ -31,6 +33,7 @@ func TestReconcile_Databases(t *testing.T) { }, existingObjects: []client.Object{coreTpl, cellTpl, shardTpl}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() // System catalog is always "postgres" db, "default" tablegroup @@ -60,16 +63,14 @@ func TestReconcile_Databases(t *testing.T) { t.Fatalf("System Catalog TableGroup not found: %v", err) } - if len(tg.Spec.Shards) != 1 { - t.Fatalf("Expected 1 shard (injected '0'), got %d", len(tg.Spec.Shards)) - } + ck.Require(). + Len(tg.Spec.Shards, 1, "Expected 1 shard (injected '0'), got %d", len(tg.Spec.Shards)) // Verify defaults applied. // NOTE: We expect 3 replicas here because 'shardTpl' (the default template in fixtures) // defines replicas: 3. The resolver correctly prioritizes the Namespace Default (Level 3) // over the Operator Default (Level 4, which is 1). - if got, want := *tg.Spec.Shards[0].Multiorch.Replicas, int32(3); got != want { - t.Errorf("Injected shard replicas mismatch. Replicas: %d, Want: %d", got, want) - } + got, want := *tg.Spec.Shards[0].Multiorch.Replicas, int32(3) + ck.Eq(want, got, "Injected shard replicas mismatch. Replicas") if len(tg.Spec.Shards[0].Multiorch.Cells) != 1 || tg.Spec.Shards[0].Multiorch.Cells[0] != "zone-a" { t.Errorf( @@ -89,6 +90,7 @@ func TestReconcile_Databases(t *testing.T) { }, existingObjects: []client.Object{coreTpl, cellTpl, shardTpl}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) tg := &multigresv1alpha1.TableGroup{} tgName := name.JoinWithConstraints( name.DefaultConstraints, @@ -96,16 +98,16 @@ func TestReconcile_Databases(t *testing.T) { "db1", "tg1", ) - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, tg, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if got := tg.Spec.Shards[0].Multiorch.Cells[0]; got != "zone-custom" { - t.Errorf("Expected explicit cell 'zone-custom', got %s", got) - } + ), "failed to get tablegroup") + ck.Eq( + "zone-custom", + tg.Spec.Shards[0].Multiorch.Cells[0], + "Expected explicit cell 'zone-custom', got", + ) }, }, "Reconcile: Implicit Cell Sorting": { @@ -134,6 +136,7 @@ func TestReconcile_Databases(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() tg := &multigresv1alpha1.TableGroup{} tgName := name.JoinWithConstraints( @@ -142,20 +145,18 @@ func TestReconcile_Databases(t *testing.T) { "db1", "tg1", ) - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: tgName, Namespace: namespace}, tg, - ); err != nil { - t.Fatal(err) - } + )) cells := tg.Spec.Shards[0].Multiorch.Cells - if len(cells) != 2 { - t.Fatalf("Expected 2 cells, got %d", len(cells)) - } - if cells[0] != "zone-a" || cells[1] != "zone-b" { - t.Errorf("Cells not sorted: %v", cells) - } + ck.Require().Len(cells, 2, "Expected 2 cells, got %d", len(cells)) + ck.False( + cells[0] != "zone-a" || cells[1] != "zone-b", + "Cells not sorted: %v", + cells, + ) }, }, "Error: Explicit Shard Template Missing": { @@ -322,9 +323,7 @@ func TestReconcileDatabases_BuildError_SchemeMismatch(t *testing.T) { // So we should add ShardTemplate to the scheme too. scheme.AddKnownTypes(multigresv1alpha1.GroupVersion, &multigresv1alpha1.ShardTemplate{}) - if err := cl.Create(t.Context(), shardTpl); err != nil { - t.Fatal(err) - } + assert.NewAborting(t).NoError(cl.Create(t.Context(), shardTpl)) cluster.Spec.TemplateDefaults.ShardTemplate = "default-shard" // Execution @@ -391,8 +390,6 @@ func TestReconcileDatabases_Direct_Error_GlobalTopoRef(t *testing.T) { t.Error("Expected error from reconcileDatabases, got nil") } else { expectedMsg := "failed to get global topo ref" - if !strings.Contains(err.Error(), expectedMsg) { - t.Errorf("Expected error containing %q, got %q", expectedMsg, err.Error()) - } + assert.NewCollecting(t).StrContains(err.Error(), expectedMsg, "Expected error containing") } } diff --git a/pkg/cluster-handler/controller/multigrescluster/reconcile_global_test.go b/pkg/cluster-handler/controller/multigrescluster/reconcile_global_test.go index 6d0764c6..d9591130 100644 --- a/pkg/cluster-handler/controller/multigrescluster/reconcile_global_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/reconcile_global_test.go @@ -25,6 +25,8 @@ import ( "github.com/multigres/multigres-operator/pkg/resolver" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestReconcileGlobal_ErrorPaths(t *testing.T) { @@ -52,9 +54,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to missing global topo spec, got nil") - } + assert.NewCollecting(t). + Error(err, "Expected error due to missing global topo spec, got nil") }) t.Run("Error: Resolve Multiadmin Failed", func(t *testing.T) { @@ -79,9 +80,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to missing multi admin spec, got nil") - } + assert.NewCollecting(t). + Error(err, "Expected error due to missing multi admin spec, got nil") }) t.Run("Error: Resolve MultiadminWeb Failed", func(t *testing.T) { @@ -106,9 +106,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to missing multi admin web spec, got nil") - } + assert.NewCollecting(t). + Error(err, "Expected error due to missing multi admin web spec, got nil") }) t.Run("Error: Patch Global Topo Failed", func(t *testing.T) { @@ -137,9 +136,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || err.Error() != "failed to apply global topo server: patch error" { - t.Errorf("Expected 'patch error', got %v", err) - } + assert.NewCollecting(t). + False(err == nil || err.Error() != "failed to apply global topo server: patch error", "Expected 'patch error', got %v", err) }) t.Run("Error: Patch Multiadmin Failed", func(t *testing.T) { @@ -164,9 +162,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || err.Error() != "failed to apply multiadmin deployment: patch error" { - t.Errorf("Expected 'patch error', got %v", err) - } + assert.NewCollecting(t). + False(err == nil || err.Error() != "failed to apply multiadmin deployment: patch error", "Expected 'patch error', got %v", err) }) t.Run("Error: Patch MultiadminWeb Failed", func(t *testing.T) { @@ -191,9 +188,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || err.Error() != "failed to apply multiadmin-web deployment: patch error" { - t.Errorf("Expected 'patch error', got %v", err) - } + assert.NewCollecting(t). + False(err == nil || err.Error() != "failed to apply multiadmin-web deployment: patch error", "Expected 'patch error', got %v", err) }) t.Run("Error: Build Global Topo Failed", func(t *testing.T) { @@ -219,9 +215,7 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to build failure, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to build failure, got nil") }) t.Run("Error: Build Multiadmin Failed", func(t *testing.T) { @@ -243,9 +237,7 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to build failure, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to build failure, got nil") }) t.Run("Error: Build MultiadminWeb Failed", func(t *testing.T) { @@ -267,9 +259,7 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to build failure, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to build failure, got nil") }) t.Run("Error: Build Multiadmin Service Failed", func(t *testing.T) { @@ -298,9 +288,7 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to build failure, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to build failure, got nil") }) t.Run("Error: Patch Multiadmin Service Failed", func(t *testing.T) { @@ -328,9 +316,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || err.Error() != "failed to apply multiadmin service: service patch error" { - t.Errorf("Expected 'service patch error', got %v", err) - } + assert.NewCollecting(t). + False(err == nil || err.Error() != "failed to apply multiadmin service: service patch error", "Expected 'service patch error', got %v", err) }) t.Run("Error: Build MultiadminWeb Service Failed", func(t *testing.T) { @@ -354,9 +341,7 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil { - t.Error("Expected error due to build failure, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to build failure, got nil") }) t.Run("Error: Patch MultiadminWeb Service Failed", func(t *testing.T) { @@ -384,10 +369,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || - err.Error() != "failed to apply multiadmin-web service: service patch error" { - t.Errorf("Expected 'service patch error', got %v", err) - } + assert.NewCollecting(t).False(err == nil || + err.Error() != "failed to apply multiadmin-web service: service patch error", "Expected 'service patch error', got %v", err) }) t.Run("Error: Patch Multigateway Global Service Failed", func(t *testing.T) { @@ -420,10 +403,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || - err.Error() != "failed to apply global multigateway service: gw service patch error" { - t.Errorf("Expected 'gw service patch error', got %v", err) - } + assert.NewCollecting(t).False(err == nil || + err.Error() != "failed to apply global multigateway service: gw service patch error", "Expected 'gw service patch error', got %v", err) }) t.Run("Error: Patch Multigateway Global Replica Service Failed", func(t *testing.T) { @@ -452,10 +433,8 @@ func TestReconcileGlobal_ErrorPaths(t *testing.T) { cluster, resolver.NewResolver(c, "default"), ) - if err == nil || - err.Error() != "failed to apply global replica multigateway service: gw replica service patch error" { - t.Errorf("Expected 'gw replica service patch error', got %v", err) - } + assert.NewCollecting(t).False(err == nil || + err.Error() != "failed to apply global replica multigateway service: gw replica service patch error", "Expected 'gw replica service patch error', got %v", err) }) } @@ -492,16 +471,15 @@ func TestReconcile_Global(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() ts := &multigresv1alpha1.TopoServer{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: clusterName + "-global-topo", Namespace: namespace}, ts, - ); err != nil { - t.Fatal(err) - } + )) if got, want := ts.Spec.Etcd.Image, multigresv1alpha1.ImageRef( "etcd:topo", ); got != want { @@ -509,61 +487,51 @@ func TestReconcile_Global(t *testing.T) { } deploy := &appsv1.Deployment{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: clusterName + "-multiadmin", Namespace: namespace}, deploy, - ); err != nil { - t.Fatal(err) - } + )) if got, want := *deploy.Spec.Replicas, int32(5); got != want { t.Errorf("Multiadmin replicas mismatch got %d, want %d", got, want) } webDeploy := &appsv1.Deployment{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{ Name: clusterName + "-multiadmin-web", Namespace: namespace, }, webDeploy, - ); err != nil { - t.Fatal(err) - } + )) // Default replicas is 1 - if got, want := *webDeploy.Spec.Replicas, int32(1); got != want { - t.Errorf("MultiadminWeb replicas mismatch got %d, want %d", got, want) - } + got, want := *webDeploy.Spec.Replicas, int32(1) + ck.Eq(want, got, "MultiadminWeb replicas mismatch got") // Verify global multigateway Service exists gwSvc := &corev1.Service{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: clusterName + "-multigateway", Namespace: namespace}, gwSvc, - ); err != nil { - t.Fatalf("Expected global multigateway Service to exist: %v", err) - } - if gwSvc.Spec.Selector["app.kubernetes.io/component"] != "multigateway" { - t.Errorf( - "Global multigateway Service selector component = %v, want multigateway", - gwSvc.Spec.Selector["app.kubernetes.io/component"], - ) - } + ), "Expected global multigateway Service to exist") + ck.Eq( + "multigateway", + gwSvc.Spec.Selector["app.kubernetes.io/component"], + "Global multigateway Service selector component", + ) // Verify global replica multigateway Service exists gwReplicaSvc := &corev1.Service{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{ Name: clusterName + "-multigateway-replica", Namespace: namespace, }, gwReplicaSvc, - ); err != nil { - t.Fatalf("Expected global multigateway replica Service to exist: %v", err) - } + ), "Expected global multigateway replica Service to exist") if len(gwReplicaSvc.Spec.Ports) != 1 || gwReplicaSvc.Spec.Ports[0].Port != 5433 || gwReplicaSvc.Spec.Ports[0].TargetPort != intstr.FromString("pg-replica") { @@ -586,18 +554,16 @@ func TestReconcile_Global(t *testing.T) { }, existingObjects: []client.Object{coreTpl, cellTpl, shardTpl}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() ts := &multigresv1alpha1.TopoServer{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: clusterName + "-global-topo", Namespace: namespace}, ts, - ); err != nil { - t.Fatal(err) - } - if got, want := ts.Spec.Etcd.RootPath, "/custom/root"; got != want { - t.Errorf("RootPath mismatch got %q, want %q", got, want) - } + )) + got, want := ts.Spec.Etcd.RootPath, "/custom/root" + ck.Eq(want, got, "RootPath mismatch got") }, }, @@ -611,17 +577,17 @@ func TestReconcile_Global(t *testing.T) { }, existingObjects: []client.Object{coreTpl, cellTpl, shardTpl}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() ts := &multigresv1alpha1.TopoServer{} - if err := c.Get( + err := c.Get( ctx, types.NamespacedName{Name: clusterName + "-global-topo", Namespace: namespace}, ts, - ); !apierrors.IsNotFound( + ) + assert.NewAborting(t).True(apierrors.IsNotFound( err, - ) { - t.Fatal("Global TopoServer should NOT be created for External mode") - } + ), "Global TopoServer should NOT be created for External mode") cell := &multigresv1alpha1.Cell{} // Use hashed name for Cell cellName := name.JoinWithConstraints( @@ -629,16 +595,13 @@ func TestReconcile_Global(t *testing.T) { clusterName, "zone-a", ) - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: cellName, Namespace: namespace}, cell, - ); err != nil { - t.Fatalf("Expected Cell %s to exist: %v", cellName, err) - } - if got, want := cell.Spec.GlobalTopoServer.Address, "http://external-etcd:2379"; got != want { - t.Errorf("External address mismatch got %q, want %q", got, want) - } + ), "Expected Cell %s to exist", cellName) + got, want := cell.Spec.GlobalTopoServer.Address, "http://external-etcd:2379" + ck.Eq(want, got, "External address mismatch got") }, }, "Error: Explicit Core Template Missing (Should Fail)": { @@ -727,43 +690,32 @@ func TestReconcile_Global(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ts := &multigresv1alpha1.TopoServer{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: clusterName + "-global-topo", Namespace: namespace}, ts, - ); err != nil { - t.Fatal(err) - } - if ts.Spec.Etcd.Image != "new-etcd" { - t.Errorf("TopoServer not updated") - } + )) + ck.Eq("new-etcd", ts.Spec.Etcd.Image, "TopoServer not updated") deploy := &appsv1.Deployment{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: clusterName + "-multiadmin", Namespace: namespace}, deploy, - ); err != nil { - t.Fatal(err) - } - if *deploy.Spec.Replicas != 3 { - t.Errorf("Multiadmin not updated") - } + )) + ck.Eq(3, *deploy.Spec.Replicas, "Multiadmin not updated") webDeploy := &appsv1.Deployment{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{ Name: clusterName + "-multiadmin-web", Namespace: namespace, }, webDeploy, - ); err != nil { - t.Fatal(err) - } - if *webDeploy.Spec.Replicas != 2 { - t.Errorf("MultiadminWeb not updated") - } + )) + ck.Eq(2, *webDeploy.Spec.Replicas, "MultiadminWeb not updated") }, }, "Idempotency: No changes needed": { @@ -923,9 +875,8 @@ func TestReconcile_Global_BuilderErrors(t *testing.T) { c := clientBuilder.Build() // Manually create template to ensure it exists - if err := c.Create(t.Context(), coreTpl.DeepCopy()); err != nil { - t.Fatalf("Failed to create CoreTemplate: %v", err) - } + assert.NewAborting(t). + NoError(c.Create(t.Context(), coreTpl.DeepCopy()), "Failed to create CoreTemplate") reconciler := &MultigresClusterReconciler{ Client: c, @@ -967,9 +918,8 @@ func TestReconcile_Global_BuilderErrors(t *testing.T) { c := clientBuilder.Build() // Manually create template to ensure it exists - if err := c.Create(t.Context(), coreTpl.DeepCopy()); err != nil { - t.Fatalf("Failed to create CoreTemplate: %v", err) - } + assert.NewAborting(t). + NoError(c.Create(t.Context(), coreTpl.DeepCopy()), "Failed to create CoreTemplate") reconciler := &MultigresClusterReconciler{ Client: c, @@ -1003,9 +953,8 @@ func TestReconcile_Global_BuilderErrors(t *testing.T) { c := clientBuilder.Build() // Manually create template to ensure it exists - if err := c.Create(t.Context(), coreTpl.DeepCopy()); err != nil { - t.Fatalf("Failed to create CoreTemplate: %v", err) - } + assert.NewAborting(t). + NoError(c.Create(t.Context(), coreTpl.DeepCopy()), "Failed to create CoreTemplate") reconciler := &MultigresClusterReconciler{ Client: c, @@ -1036,9 +985,8 @@ func TestReconcile_Global_BuilderErrors(t *testing.T) { WithStatusSubresource(&multigresv1alpha1.MultigresCluster{}) c := clientBuilder.Build() - if err := c.Create(t.Context(), coreTpl.DeepCopy()); err != nil { - t.Fatalf("Failed to create CoreTemplate: %v", err) - } + assert.NewAborting(t). + NoError(c.Create(t.Context(), coreTpl.DeepCopy()), "Failed to create CoreTemplate") reconciler := &MultigresClusterReconciler{ Client: c, @@ -1069,9 +1017,8 @@ func TestReconcile_Global_BuilderErrors(t *testing.T) { WithStatusSubresource(&multigresv1alpha1.MultigresCluster{}) c := clientBuilder.Build() - if err := c.Create(t.Context(), coreTpl.DeepCopy()); err != nil { - t.Fatalf("Failed to create CoreTemplate: %v", err) - } + assert.NewAborting(t). + NoError(c.Create(t.Context(), coreTpl.DeepCopy()), "Failed to create CoreTemplate") reconciler := &MultigresClusterReconciler{ Client: c, @@ -1103,13 +1050,13 @@ func TestReconcileAdminNetworkPolicies(t *testing.T) { listPolicies := func(t *testing.T, c client.Client) []networkingv1.NetworkPolicy { t.Helper() list := &networkingv1.NetworkPolicyList{} - if err := c.List(context.Background(), list, client.InNamespace("default")); err != nil { - t.Fatalf("failed to list network policies: %v", err) - } + assert.NewAborting(t). + NoError(c.List(context.Background(), list, client.InNamespace("default")), "failed to list network policies") return list.Items } t.Run("Enabled creates policies for multiadmin and multiadmin-web", func(t *testing.T) { + ck := assert.NewAborting(t) cluster := newCluster(&multigresv1alpha1.NetworkPolicyConfig{ Enabled: true, AllowedIngressNamespaces: []string{"envoy-gateway-system"}, @@ -1126,17 +1073,17 @@ func TestReconcileAdminNetworkPolicies(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileAdminNetworkPolicies(context.Background(), cluster); err != nil { - t.Fatalf("reconcileAdminNetworkPolicies() error = %v", err) - } + ck.NoError( + r.reconcileAdminNetworkPolicies(context.Background(), cluster), + "reconcileAdminNetworkPolicies() error =", + ) policies := listPolicies(t, c) - if len(policies) != 2 { - t.Fatalf("expected 2 network policies, got %d", len(policies)) - } + ck.Len(policies, 2, "expected 2 network policies, got %d", len(policies)) }) t.Run("Disabled deletes previously created policies", func(t *testing.T) { + ck := assert.NewAborting(t) cluster := newCluster(&multigresv1alpha1.NetworkPolicyConfig{Enabled: true}) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(cluster). WithTypeConverters(managedfields.NewDeducedTypeConverter()).Build() @@ -1146,20 +1093,18 @@ func TestReconcileAdminNetworkPolicies(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileAdminNetworkPolicies(context.Background(), cluster); err != nil { - t.Fatalf("reconcileAdminNetworkPolicies() error = %v", err) - } - if got := len(listPolicies(t, c)); got != 2 { - t.Fatalf("expected 2 network policies before disable, got %d", got) - } + ck.NoError( + r.reconcileAdminNetworkPolicies(context.Background(), cluster), + "reconcileAdminNetworkPolicies() error =", + ) + ck.Eq(2, len(listPolicies(t, c)), "expected 2 network policies before disable, got") cluster.Spec.NetworkPolicy = nil - if err := r.reconcileAdminNetworkPolicies(context.Background(), cluster); err != nil { - t.Fatalf("reconcileAdminNetworkPolicies() error = %v", err) - } - if got := len(listPolicies(t, c)); got != 0 { - t.Fatalf("expected 0 network policies after disable, got %d", got) - } + ck.NoError( + r.reconcileAdminNetworkPolicies(context.Background(), cluster), + "reconcileAdminNetworkPolicies() error =", + ) + ck.Eq(0, len(listPolicies(t, c)), "expected 0 network policies after disable, got") }) t.Run("Disabled with nothing to delete is a no-op", func(t *testing.T) { @@ -1171,9 +1116,8 @@ func TestReconcileAdminNetworkPolicies(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileAdminNetworkPolicies(context.Background(), cluster); err != nil { - t.Fatalf("reconcileAdminNetworkPolicies() error = %v", err) - } + assert.NewAborting(t). + NoError(r.reconcileAdminNetworkPolicies(context.Background(), cluster), "reconcileAdminNetworkPolicies() error =") }) t.Run("Apply error is propagated", func(t *testing.T) { @@ -1194,8 +1138,7 @@ func TestReconcileAdminNetworkPolicies(t *testing.T) { } err := r.reconcileAdminNetworkPolicies(context.Background(), cluster) - if err == nil || !strings.Contains(err.Error(), "failed to apply network policy") { - t.Fatalf("expected apply error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to apply network policy"), "expected apply error, got %v", err) }) } diff --git a/pkg/cluster-handler/controller/multigrescluster/reconcile_global_topo_tls_test.go b/pkg/cluster-handler/controller/multigrescluster/reconcile_global_topo_tls_test.go index dab8c95e..dcd95c30 100644 --- a/pkg/cluster-handler/controller/multigrescluster/reconcile_global_topo_tls_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/reconcile_global_topo_tls_test.go @@ -11,6 +11,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/resolver" + + "github.com/multigres/testkit/assert" ) func topoRefCluster( @@ -49,9 +51,7 @@ func resolveTopoRef( ref, err := r.globalTopoRef( context.Background(), cluster, resolver.NewResolver(c, cluster.Namespace), ) - if err != nil { - t.Fatalf("globalTopoRef() error = %v", err) - } + assert.NewAborting(t).NoError(err, "globalTopoRef() error =") return ref } @@ -61,29 +61,23 @@ func TestGlobalTopoRefResolvesManagedSecrets(t *testing.T) { } t.Run("topology TLS enabled resolves the operator-issued secret", func(t *testing.T) { + c := assert.NewCollecting(t) ref := resolveTopoRef(t, topoRefCluster( &multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}, managed, )) want := "test-cluster-topo-client-tls" - if ref.ClientCertSecret != want { - t.Errorf("ClientCertSecret = %q, want %q", ref.ClientCertSecret, want) - } + c.Eq(want, ref.ClientCertSecret, "ClientCertSecret") // cert-manager writes ca.crt into the same Secret as the keypair. - if ref.CASecret != want { - t.Errorf("CASecret = %q, want %q", ref.CASecret, want) - } + c.Eq(want, ref.CASecret, "CASecret") }) t.Run("topology TLS unset leaves both references empty", func(t *testing.T) { + c := assert.NewCollecting(t) ref := resolveTopoRef(t, topoRefCluster(nil, managed)) - if ref.ClientCertSecret != "" { - t.Errorf("ClientCertSecret = %q, want empty", ref.ClientCertSecret) - } - if ref.CASecret != "" { - t.Errorf("CASecret = %q, want empty", ref.CASecret) - } + c.Eq("", ref.ClientCertSecret, "ClientCertSecret") + c.Eq("", ref.CASecret, "CASecret") }) t.Run("topology TLS disabled leaves both references empty", func(t *testing.T) { @@ -91,12 +85,8 @@ func TestGlobalTopoRefResolvesManagedSecrets(t *testing.T) { &multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(false)}, managed, )) - if ref.ClientCertSecret != "" || ref.CASecret != "" { - t.Errorf( - "want empty references, got CASecret=%q ClientCertSecret=%q", - ref.CASecret, ref.ClientCertSecret, - ) - } + assert.NewCollecting(t). + False(ref.ClientCertSecret != "" || ref.CASecret != "", "want empty references, got CASecret=%q ClientCertSecret=%q", ref.CASecret, ref.ClientCertSecret) }) } @@ -104,6 +94,7 @@ func TestGlobalTopoRefResolvesManagedSecrets(t *testing.T) { // spec's secret names are what reach the ref the cell and shard controllers // read. func TestGlobalTopoRefResolvesExternalSecrets(t *testing.T) { + c := assert.NewCollecting(t) ref := resolveTopoRef(t, topoRefCluster(nil, &multigresv1alpha1.GlobalTopoServerSpec{ //nolint:gosec // K8s resource names, not credentials External: &multigresv1alpha1.ExternalTopoServerSpec{ @@ -115,15 +106,12 @@ func TestGlobalTopoRefResolvesExternalSecrets(t *testing.T) { }, })) - if ref.CASecret != "infra-etcd-ca" { - t.Errorf("CASecret = %q, want infra-etcd-ca", ref.CASecret) - } - if ref.ClientCertSecret != "proj-123-topo-client" { - t.Errorf("ClientCertSecret = %q, want proj-123-topo-client", ref.ClientCertSecret) - } + c.Eq("infra-etcd-ca", ref.CASecret, "CASecret") + c.Eq("proj-123-topo-client", ref.ClientCertSecret, "ClientCertSecret") } func TestBuildGlobalTopoServerPropagatesTopoTLS(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() tls := &multigresv1alpha1.TopoTLSConfig{ Enabled: ptr.To(true), @@ -135,21 +123,16 @@ func TestBuildGlobalTopoServerPropagatesTopoTLS(t *testing.T) { } ts, err := BuildGlobalTopoServer(cluster, spec, scheme) - if err != nil { - t.Fatalf("BuildGlobalTopoServer() error = %v", err) - } - if ts.Spec.TLS == nil { - t.Fatal("TopoServer.Spec.TLS = nil, want the cluster's topology TLS config") - } - if !ts.Spec.TLS.IsEnabled() { - t.Error("TopoServer.Spec.TLS is not enabled") - } - if ts.Spec.TLS.IssuerName != "multigres-infra-issuer" { - t.Errorf("IssuerName = %q, want multigres-infra-issuer", ts.Spec.TLS.IssuerName) - } + c.Require().NoError(err, "BuildGlobalTopoServer() error =") + c.Require(). + NotNil(ts.Spec.TLS, "TopoServer.Spec.TLS = nil, want the cluster's topology TLS config") + c.True(ts.Spec.TLS.IsEnabled(), "TopoServer.Spec.TLS is not enabled") + c.Eq("multigres-infra-issuer", ts.Spec.TLS.IssuerName, "IssuerName") // A deep copy, so mutating the child spec cannot reach back into the cluster. ts.Spec.TLS.IssuerName = "mutated" - if cluster.Spec.TopoTLS.IssuerName != "multigres-infra-issuer" { - t.Error("mutating the child TLS config changed the cluster spec") - } + c.Eq( + "multigres-infra-issuer", + cluster.Spec.TopoTLS.IssuerName, + "mutating the child TLS config changed the cluster spec", + ) } diff --git a/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_test.go b/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_test.go index 51cdbc4c..50d6fc95 100644 --- a/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_test.go @@ -3,7 +3,6 @@ package multigrescluster import ( "context" "path" - "reflect" "testing" "github.com/multigres/multigres/go/common/topoclient" @@ -18,10 +17,13 @@ import ( "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/pkg/util/name" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + + "github.com/multigres/testkit/assert" ) func TestConvergedTopologyReconcileDoesNotWrite(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ctx := t.Context() cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "cluster", Namespace: "default"}, @@ -50,15 +52,11 @@ func TestConvergedTopologyReconcileDoesNotWrite(t *testing.T) { t.Fatal(err) } conn, err := store.ConnForCell(ctx, topoclient.GlobalCell) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) versions := map[string]string{} for _, file := range []string{path.Join(topoclient.CellsPath, "cell1", topoclient.CellFile), path.Join(topoclient.DatabasesPath, "db", topoclient.DatabaseFile)} { _, v, err := conn.Get(ctx, file) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) versions[file] = v.String() } for range 5 { @@ -67,18 +65,15 @@ func TestConvergedTopologyReconcileDoesNotWrite(t *testing.T) { } for file, want := range versions { _, v, err := conn.Get(ctx, file) - if err != nil { - t.Fatal(err) - } - if v.String() != want { - t.Fatalf("converged reconcile rewrote %s", file) - } + ck.NoError(err) + ck.Eq(want, v.String(), "converged reconcile rewrote %s", file) } } } func TestReconcileTopologySharedTopo(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ @@ -126,42 +121,35 @@ func TestReconcileTopologySharedTopo(t *testing.T) { cluster, resolver.NewResolver(client, "default"), ) - if err != nil { - t.Fatalf("reconcileTopology() error = %v", err) - } - if result.RequeueAfter != 0 { - t.Fatalf("RequeueAfter = %v, want 0", result.RequeueAfter) - } + c.Require().NoError(err, "reconcileTopology() error =") + c.Require().Eq(0, result.RequeueAfter, "RequeueAfter") - if openedRef.Address != "http://global-etcd:2379" { - t.Fatalf("expected topology store to open global address, got %q", openedRef.Address) - } - if openedRef.RootPath != "/multigres/clusters/cluster/global" { - t.Fatalf("expected topology store to open global root, got %q", openedRef.RootPath) - } + c.Require(). + Eq("http://global-etcd:2379", openedRef.Address, "expected topology store to open global address, got") + c.Require(). + Eq("/multigres/clusters/cluster/global", openedRef.RootPath, "expected topology store to open global root, got") cell, err := store.GetCell(context.Background(), "cell1") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if !reflect.DeepEqual(cell.ServerAddresses, []string{"http://cell1-local-etcd:2379"}) { - t.Errorf("expected cell record to point at local topology, got %v", cell.ServerAddresses) - } - if cell.Root != "/multigres/clusters/cluster/cells/cell1" { - t.Errorf("expected cell record local root, got %q", cell.Root) - } + c.Require().NoError(err, "cell not found") + c.EqDiff( + []string{"http://cell1-local-etcd:2379"}, + cell.ServerAddresses, + "expected cell record to point at local topology, got", + ) + c.Eq( + "/multigres/clusters/cluster/cells/cell1", + cell.Root, + "expected cell record local root, got", + ) db, err := store.GetDatabase(context.Background(), "commerce") - if err != nil { - t.Fatalf("database not found in global topology store: %v", err) - } - if db.Name != "commerce" { - t.Errorf("expected database commerce, got %q", db.Name) - } + c.Require().NoError(err, "database not found in global topology store") + c.Eq("commerce", db.Name, "expected database commerce, got") } func TestReconcileTopologyManagedLocalTopo(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ @@ -194,9 +182,8 @@ func TestReconcileTopologyManagedLocalTopo(t *testing.T) { WithObjects(cluster, k8sCell, ts). Build() ts.Status.ObservedGeneration = ts.Generation - if err := client.Status().Update(context.Background(), ts); err != nil { - t.Fatalf("failed to update TopoServer status: %v", err) - } + c.Require(). + NoError(client.Status().Update(context.Background(), ts), "failed to update TopoServer status") reconciler := &MultigresClusterReconciler{ Client: client, Scheme: scheme, @@ -210,31 +197,26 @@ func TestReconcileTopologyManagedLocalTopo(t *testing.T) { cluster, resolver.NewResolver(client, "default"), ) - if err != nil { - t.Fatalf("reconcileTopology() error = %v", err) - } - if result.RequeueAfter != 0 { - t.Fatalf("RequeueAfter = %v, want 0", result.RequeueAfter) - } + c.Require().NoError(err, "reconcileTopology() error =") + c.Require().Eq(0, result.RequeueAfter, "RequeueAfter") cell, err := store.GetCell(context.Background(), "cell1") - if err != nil { - t.Fatalf("cell not found: %v", err) - } + c.Require().NoError(err, "cell not found") wantAddress := topo.ManagedLocalTopoServerAddress( name.JoinWithConstraints(name.DefaultConstraints, "cluster", "cell1"), "default", ) - if !reflect.DeepEqual(cell.ServerAddresses, []string{wantAddress}) { - t.Errorf("expected managed local topology service, got %v", cell.ServerAddresses) - } - if cell.Root != "/multigres/default/cluster/cell1" { - t.Errorf("expected default local root, got %q", cell.Root) - } + c.EqDiff( + []string{wantAddress}, + cell.ServerAddresses, + "expected managed local topology service, got", + ) + c.Eq("/multigres/default/cluster/cell1", cell.Root, "expected default local root, got") } func TestReconcileTopologyWaitsForManagedLocalTopoNotOwnedByCell(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ @@ -266,9 +248,10 @@ func TestReconcileTopologyWaitsForManagedLocalTopoNotOwnedByCell(t *testing.T) { WithObjects(cluster, cell, ts). Build() ts.Status.ObservedGeneration = ts.Generation - if err := client.Status().Update(context.Background(), ts); err != nil { - t.Fatalf("failed to update TopoServer status: %v", err) - } + c.NoError( + client.Status().Update(context.Background(), ts), + "failed to update TopoServer status", + ) reconciler := &MultigresClusterReconciler{ Client: client, Scheme: scheme, @@ -284,16 +267,13 @@ func TestReconcileTopologyWaitsForManagedLocalTopoNotOwnedByCell(t *testing.T) { cluster, resolver.NewResolver(client, "default"), ) - if err != nil { - t.Fatalf("reconcileTopology() error = %v", err) - } - if result.RequeueAfter != localTopoServerRequeueDelay { - t.Fatalf("RequeueAfter = %v, want %v", result.RequeueAfter, localTopoServerRequeueDelay) - } + c.NoError(err, "reconcileTopology() error =") + c.Eq(localTopoServerRequeueDelay, result.RequeueAfter, "RequeueAfter") } func TestReconcileTopologyWaitsForManagedLocalTopo(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ @@ -331,12 +311,8 @@ func TestReconcileTopologyWaitsForManagedLocalTopo(t *testing.T) { cluster, resolver.NewResolver(client, "default"), ) - if err != nil { - t.Fatalf("reconcileTopology() error = %v", err) - } - if result.RequeueAfter != localTopoServerRequeueDelay { - t.Fatalf("RequeueAfter = %v, want %v", result.RequeueAfter, localTopoServerRequeueDelay) - } + c.NoError(err, "reconcileTopology() error =") + c.Eq(localTopoServerRequeueDelay, result.RequeueAfter, "RequeueAfter") if _, err := store.GetCell(context.Background(), "cell1"); err == nil { t.Fatal("cell should not be registered before managed local TopoServer is ready") } @@ -344,6 +320,7 @@ func TestReconcileTopologyWaitsForManagedLocalTopo(t *testing.T) { func TestReconcileTopologyKeepsExistingCellRecordWhileManagedLocalTopoWaits(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ @@ -367,7 +344,7 @@ func TestReconcileTopologyKeepsExistingCellRecordWhileManagedLocalTopoWaits(t *t } store := newClusterTopologyMemoryStore(t) - if err := topo.RegisterCellFromSpec( + c.Require().NoError(topo.RegisterCellFromSpec( context.Background(), store, record.NewFakeRecorder(10), @@ -378,9 +355,7 @@ func TestReconcileTopologyKeepsExistingCellRecordWhileManagedLocalTopoWaits(t *t Address: "http://global-etcd:2379", RootPath: "/multigres/global", }, - ); err != nil { - t.Fatalf("failed to seed existing cell topology: %v", err) - } + ), "failed to seed existing cell topology") client := fake.NewClientBuilder().WithScheme(scheme).WithObjects(cluster).Build() reconciler := &MultigresClusterReconciler{ @@ -397,27 +372,18 @@ func TestReconcileTopologyKeepsExistingCellRecordWhileManagedLocalTopoWaits(t *t cluster, resolver.NewResolver(client, "default"), ) - if err != nil { - t.Fatalf("reconcileTopology() error = %v", err) - } - if result.RequeueAfter != localTopoServerRequeueDelay { - t.Fatalf("RequeueAfter = %v, want %v", result.RequeueAfter, localTopoServerRequeueDelay) - } + c.Require().NoError(err, "reconcileTopology() error =") + c.Require().Eq(localTopoServerRequeueDelay, result.RequeueAfter, "RequeueAfter") cell, err := store.GetCell(context.Background(), "cell1") - if err != nil { - t.Fatalf("existing cell record should remain available: %v", err) - } - if !reflect.DeepEqual(cell.ServerAddresses, []string{"http://global-etcd:2379"}) { - t.Errorf("existing cell address = %v, want global topology address", cell.ServerAddresses) - } - if cell.Root != "/multigres/global" { - t.Errorf("existing cell root = %q, want global topology root", cell.Root) - } + c.Require().NoError(err, "existing cell record should remain available") + c.EqDiff([]string{"http://global-etcd:2379"}, cell.ServerAddresses, "existing cell address") + c.Eq("/multigres/global", cell.Root, "existing cell root") } func TestReconcileTopologyKeepsPendingDeletionCellRecord(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ @@ -448,7 +414,7 @@ func TestReconcileTopologyKeepsPendingDeletionCellRecord(t *testing.T) { store := newClusterTopologyMemoryStore(t) for _, cellName := range []multigresv1alpha1.CellName{"cell1", "cell2"} { - if err := topo.RegisterCellFromSpec( + c.NoError(topo.RegisterCellFromSpec( context.Background(), store, record.NewFakeRecorder(10), @@ -459,9 +425,7 @@ func TestReconcileTopologyKeepsPendingDeletionCellRecord(t *testing.T) { Address: "http://global-etcd:2379", RootPath: "/multigres/global", }, - ); err != nil { - t.Fatalf("failed to seed cell %s topology: %v", cellName, err) - } + ), "failed to seed cell %s topology", cellName) } client := fake.NewClientBuilder().WithScheme(scheme).WithObjects(cluster, pendingCell).Build() @@ -480,12 +444,8 @@ func TestReconcileTopologyKeepsPendingDeletionCellRecord(t *testing.T) { resolver.NewResolver(client, "default"), true, ) - if err != nil { - t.Fatalf("reconcileTopology() error = %v", err) - } - if result.RequeueAfter != 0 { - t.Fatalf("RequeueAfter = %v, want 0", result.RequeueAfter) - } + c.NoError(err, "reconcileTopology() error =") + c.Eq(0, result.RequeueAfter, "RequeueAfter") if _, err := store.GetCell(context.Background(), "cell1"); err != nil { t.Fatalf("active cell record should remain: %v", err) diff --git a/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_tls_test.go b/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_tls_test.go index d2b543db..e3c36864 100644 --- a/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_tls_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/reconcile_topology_tls_test.go @@ -3,7 +3,6 @@ package multigrescluster import ( "context" "errors" - "strings" "testing" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -12,12 +11,15 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) // A topology connection failure has to be visible in the cluster status, not // only in the logs, so an operator can see why the cluster is not becoming // ready and which Secret is missing. func TestMarkTopologyConnectFailed_SurfacesInStatus(t *testing.T) { + ck := assert.NewCollecting(t) scheme := setupScheme() cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ @@ -43,35 +45,37 @@ func TestMarkTopologyConnectFailed_SurfacesInStatus(t *testing.T) { r.markTopologyFailed(context.Background(), cluster, "TopoConnectFailed", cause, testLogger{}) got := &multigresv1alpha1.MultigresCluster{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "test-cluster"}, got, - ); err != nil { - t.Fatalf("Get() error = %v", err) - } + ), "Get() error =") - if got.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("phase = %q, want %q", got.Status.Phase, multigresv1alpha1.PhaseDegraded) - } - if !strings.Contains(got.Status.Message, "test-cluster-topo-client-tls") { - t.Errorf("status message does not name the Secret: %q", got.Status.Message) - } + ck.Eq(multigresv1alpha1.PhaseDegraded, got.Status.Phase, "phase") + ck.StrContains( + got.Status.Message, + "test-cluster-topo-client-tls", + "status message does not name the Secret", + ) var found bool for _, cond := range got.Status.Conditions { if cond.Type == conditionTopologyReady { found = true - if cond.Status != metav1.ConditionFalse { - t.Errorf("%s status = %q, want False", conditionTopologyReady, cond.Status) - } - if !strings.Contains(cond.Message, "test-cluster-topo-client-tls") { - t.Errorf("condition message does not name the Secret: %q", cond.Message) - } + ck.Eq( + metav1.ConditionFalse, + cond.Status, + "%s status = %q, want False", + conditionTopologyReady, + cond.Status, + ) + ck.StrContains( + cond.Message, + "test-cluster-topo-client-tls", + "condition message does not name the Secret", + ) } } - if !found { - t.Errorf("no %s condition set", conditionTopologyReady) - } + ck.True(found, "no %s condition set", conditionTopologyReady) } type testLogger struct{} diff --git a/pkg/cluster-handler/controller/multigrescluster/status_test.go b/pkg/cluster-handler/controller/multigrescluster/status_test.go index 997e15f3..6a95e7cd 100644 --- a/pkg/cluster-handler/controller/multigrescluster/status_test.go +++ b/pkg/cluster-handler/controller/multigrescluster/status_test.go @@ -6,9 +6,6 @@ import ( "fmt" "testing" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/meta" @@ -20,6 +17,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) func TestReconcile_Status(t *testing.T) { @@ -73,6 +72,7 @@ func TestReconcile_Status(t *testing.T) { } func TestUpdateStatus_Coverage(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -127,30 +127,22 @@ func TestUpdateStatus_Coverage(t *testing.T) { Recorder: record.NewFakeRecorder(100), } - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + ck.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + ck.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") found := false for _, c := range cluster.Status.Conditions { if c.Type == "Available" { found = true - if c.Status != metav1.ConditionTrue { - t.Errorf("Expected Available=True, got %s", c.Status) - } + ck.Eq(metav1.ConditionTrue, c.Status, "Expected Available=True, got") } } - if !found { - t.Error("Available condition not found") - } + ck.True(found, "Available condition not found") if s, ok := cluster.Status.Cells["cell-1"]; !ok || !s.Ready { t.Errorf("Expected cell-1 to be ready in status summary, got %v", s) @@ -179,21 +171,15 @@ func TestUpdateStatus_Coverage(t *testing.T) { Build() r.Client = fakeClient - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + ck.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + ck.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") - if cluster.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("Expected PhaseDegraded, got %s", cluster.Status.Phase) - } + ck.Eq(multigresv1alpha1.PhaseDegraded, cluster.Status.Phase, "Expected PhaseDegraded, got") cProgressing := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ @@ -214,19 +200,17 @@ func TestUpdateStatus_Coverage(t *testing.T) { Build() r.Client = fakeClient - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } - if err := fakeClient.Get( + ck.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") + ck.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } - if cluster.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf("Expected PhaseProgressing, got %s", cluster.Status.Phase) - } + ), "Failed to refresh cluster") + ck.Eq( + multigresv1alpha1.PhaseProgressing, + cluster.Status.Phase, + "Expected PhaseProgressing, got", + ) tgDegraded := &multigresv1alpha1.TableGroup{ ObjectMeta: metav1.ObjectMeta{ @@ -250,21 +234,15 @@ func TestUpdateStatus_Coverage(t *testing.T) { Build() r.Client = fakeClient - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + ck.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + ck.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") - if cluster.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("Expected PhaseDegraded, got %s", cluster.Status.Phase) - } + ck.Eq(multigresv1alpha1.PhaseDegraded, cluster.Status.Phase, "Expected PhaseDegraded, got") tgInit := &multigresv1alpha1.TableGroup{ ObjectMeta: metav1.ObjectMeta{ @@ -288,21 +266,19 @@ func TestUpdateStatus_Coverage(t *testing.T) { Build() r.Client = fakeClient - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + ck.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + ck.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") - if cluster.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf("Expected PhaseProgressing, got %s", cluster.Status.Phase) - } + ck.Eq( + multigresv1alpha1.PhaseProgressing, + cluster.Status.Phase, + "Expected PhaseProgressing, got", + ) tsDegraded := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ @@ -322,24 +298,19 @@ func TestUpdateStatus_Coverage(t *testing.T) { Build() r.Client = fakeClient - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + ck.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + ck.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") - if cluster.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("Expected PhaseDegraded, got %s", cluster.Status.Phase) - } + ck.Eq(multigresv1alpha1.PhaseDegraded, cluster.Status.Phase, "Expected PhaseDegraded, got") } func TestUpdateStatus_ZeroResources(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -378,33 +349,25 @@ func TestUpdateStatus_ZeroResources(t *testing.T) { Recorder: record.NewFakeRecorder(100), } - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + c.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") - assert.Equal(t, multigresv1alpha1.PhaseProgressing, cluster.Status.Phase) - assert.Nil(t, cluster.Status.InitializedAt) + c.EqDeep(multigresv1alpha1.PhaseProgressing, cluster.Status.Phase) + c.Nil(cluster.Status.InitializedAt) cond := meta.FindStatusCondition(cluster.Status.Conditions, "Available") - if cond == nil { - t.Fatal("Available condition missing") - } - if cond.Status != metav1.ConditionFalse { - t.Errorf("Expected Available=False (no cells), got %s", cond.Status) - } + c.Require().NotNil(cond, "Available condition missing") + c.Eq(metav1.ConditionFalse, cond.Status, "Expected Available=False (no cells), got") } func TestUpdateStatus_ExpectedChildren(t *testing.T) { scheme := runtime.NewScheme() - require.NoError(t, multigresv1alpha1.AddToScheme(scheme)) + assert.NewAborting(t).NoError(multigresv1alpha1.AddToScheme(scheme)) externalTopo := &multigresv1alpha1.GlobalTopoServerSpec{ External: &multigresv1alpha1.ExternalTopoServerSpec{ @@ -529,6 +492,7 @@ func TestUpdateStatus_ExpectedChildren(t *testing.T) { for name, tt := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) objects := append([]client.Object{tt.cluster}, tt.children...) fakeClient := fake.NewClientBuilder(). WithScheme(scheme). @@ -541,18 +505,17 @@ func TestUpdateStatus_ExpectedChildren(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - require.NoError(t, r.updateStatus(t.Context(), tt.cluster)) - require.NoError( - t, - fakeClient.Get(t.Context(), client.ObjectKeyFromObject(tt.cluster), tt.cluster), - ) - assert.Equal(t, tt.wantPhase, tt.cluster.Status.Phase) - assert.Equal(t, tt.initialized, tt.cluster.Status.InitializedAt != nil) + c.Require().NoError(r.updateStatus(t.Context(), tt.cluster)) + c.Require(). + NoError(fakeClient.Get(t.Context(), client.ObjectKeyFromObject(tt.cluster), tt.cluster)) + c.EqDeep(tt.wantPhase, tt.cluster.Status.Phase) + c.EqDeep(tt.initialized, tt.cluster.Status.InitializedAt != nil) }) } } func TestUpdateStatus_InitializedAtSticky(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -616,11 +579,11 @@ func TestUpdateStatus_InitializedAtSticky(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - require.NoError(t, r.updateStatus(t.Context(), cluster)) - require.NoError(t, fakeClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) + c.Require().NoError(r.updateStatus(t.Context(), cluster)) + c.Require().NoError(fakeClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) - require.Equal(t, multigresv1alpha1.PhaseHealthy, cluster.Status.Phase) - require.NotNil(t, cluster.Status.InitializedAt) + c.Require().EqDeep(multigresv1alpha1.PhaseHealthy, cluster.Status.Phase) + c.Require().NotNil(cluster.Status.InitializedAt) initializedAt := *cluster.Status.InitializedAt degradedCell := &multigresv1alpha1.Cell{ @@ -642,15 +605,16 @@ func TestUpdateStatus_InitializedAtSticky(t *testing.T) { Build() r.Client = fakeClient - require.NoError(t, r.updateStatus(t.Context(), cluster)) - require.NoError(t, fakeClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) + c.Require().NoError(r.updateStatus(t.Context(), cluster)) + c.Require().NoError(fakeClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) - assert.Equal(t, multigresv1alpha1.PhaseDegraded, cluster.Status.Phase) - require.NotNil(t, cluster.Status.InitializedAt) - assert.Equal(t, initializedAt, *cluster.Status.InitializedAt) + c.EqDeep(multigresv1alpha1.PhaseDegraded, cluster.Status.Phase) + c.Require().NotNil(cluster.Status.InitializedAt) + c.EqDeep(initializedAt, *cluster.Status.InitializedAt) } func TestUpdateStatus_GenerationMismatch(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -719,29 +683,25 @@ func TestUpdateStatus_GenerationMismatch(t *testing.T) { Recorder: record.NewFakeRecorder(100), } - if err := r.updateStatus(context.Background(), cluster); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + c.Require().NoError(r.updateStatus(context.Background(), cluster), "updateStatus failed") - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(cluster), cluster, - ); err != nil { - t.Fatalf("Failed to refresh cluster: %v", err) - } + ), "Failed to refresh cluster") - if cluster.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf( - "Expected PhaseProgressing due to generation mismatch, got %s", - cluster.Status.Phase, - ) - } + c.Eq( + multigresv1alpha1.PhaseProgressing, + cluster.Status.Phase, + "Expected PhaseProgressing due to generation mismatch, got", + ) } func TestUpdateStatus_UsesAPIReaderForChildHealth(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() - require.NoError(t, multigresv1alpha1.AddToScheme(scheme)) + c.Require().NoError(multigresv1alpha1.AddToScheme(scheme)) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ @@ -811,10 +771,10 @@ func TestUpdateStatus_UsesAPIReaderForChildHealth(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - require.NoError(t, r.updateStatus(t.Context(), cluster)) - require.NoError(t, cachedClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) - assert.Equal(t, multigresv1alpha1.PhaseProgressing, cluster.Status.Phase) - assert.Nil(t, cluster.Status.InitializedAt) + c.Require().NoError(r.updateStatus(t.Context(), cluster)) + c.Require().NoError(cachedClient.Get(t.Context(), client.ObjectKeyFromObject(cluster), cluster)) + c.EqDeep(multigresv1alpha1.PhaseProgressing, cluster.Status.Phase) + c.Nil(cluster.Status.InitializedAt) } func TestExtractExternalEndpoint(t *testing.T) { @@ -902,7 +862,7 @@ func TestExtractExternalEndpoint(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { got := extractExternalEndpoint(tc.svc) - assert.Equal(t, tc.want, got) + assert.NewCollecting(t).EqDeep(tc.want, got) }) } } @@ -968,6 +928,7 @@ func TestComputeGatewayCondition(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := computeGatewayCondition( tc.externalGatewayEnabled, tc.externalEndpoint, @@ -976,16 +937,16 @@ func TestComputeGatewayCondition(t *testing.T) { ) if tc.wantNil { - assert.Nil(t, got) + c.Nil(got) return } - require.NotNil(t, got) - assert.Equal(t, multigresv1alpha1.ConditionGatewayExternalReady, got.Type) - assert.Equal(t, tc.wantStatus, got.Status) - assert.Equal(t, tc.wantReason, got.Reason) - assert.Contains(t, got.Message, tc.wantMessageContains) - assert.Equal(t, tc.clusterGeneration, got.ObservedGeneration) + c.Require().NotNil(got) + c.EqDeep(multigresv1alpha1.ConditionGatewayExternalReady, got.Type) + c.EqDeep(tc.wantStatus, got.Status) + c.EqDeep(tc.wantReason, got.Reason) + c.StrContains(got.Message, tc.wantMessageContains) + c.EqDeep(tc.clusterGeneration, got.ObservedGeneration) }) } } @@ -1009,6 +970,7 @@ func TestUpdateStatus_GatewayServiceErrors(t *testing.T) { } t.Run("non-NotFound error fetching global service returns error", func(t *testing.T) { + c := assert.NewCollecting(t) injectedErr := fmt.Errorf("simulated API server error") cl := fake.NewClientBuilder(). @@ -1029,11 +991,12 @@ func TestUpdateStatus_GatewayServiceErrors(t *testing.T) { cluster := baseCluster.DeepCopy() err := r.updateStatus(t.Context(), cluster) - require.Error(t, err) - assert.Contains(t, err.Error(), "failed to get global multigateway service for status") + c.Require().Error(err) + c.StrContains(err.Error(), "failed to get global multigateway service for status") }) t.Run("NotFound global service sets AwaitingEndpoint when enabled", func(t *testing.T) { + c := assert.NewCollecting(t) // No Service object created — Get will return NotFound. cl := fake.NewClientBuilder(). WithScheme(scheme). @@ -1049,22 +1012,23 @@ func TestUpdateStatus_GatewayServiceErrors(t *testing.T) { cluster := baseCluster.DeepCopy() err := r.updateStatus(t.Context(), cluster) - require.NoError(t, err) + c.Require().NoError(err) - require.NotNil(t, cluster.Status.Gateway) - assert.Empty(t, cluster.Status.Gateway.ExternalEndpoint) + c.Require().NotNil(cluster.Status.Gateway) + c.Empty(cluster.Status.Gateway.ExternalEndpoint) cond := meta.FindStatusCondition( cluster.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady, ) - require.NotNil(t, cond) - assert.Equal(t, metav1.ConditionFalse, cond.Status) - assert.Equal(t, multigresv1alpha1.ReasonAwaitingEndpoint, cond.Reason) + c.Require().NotNil(cond) + c.EqDeep(metav1.ConditionFalse, cond.Status) + c.EqDeep(multigresv1alpha1.ReasonAwaitingEndpoint, cond.Reason) }) } func TestUpdateStatus_StaleCellGenerationIgnored(t *testing.T) { + c := assert.NewCollecting(t) scheme := setupScheme() clusterName := "gw-stale" @@ -1142,7 +1106,7 @@ func TestUpdateStatus_StaleCellGenerationIgnored(t *testing.T) { clusterCopy := cluster.DeepCopy() err := r.updateStatus(t.Context(), clusterCopy) - require.NoError(t, err) + c.Require().NoError(err) // Only the fresh cell's 2 ready gateways should count. // With endpoint present and aggregateReadyGateways=2, condition should be EndpointReady. @@ -1150,10 +1114,10 @@ func TestUpdateStatus_StaleCellGenerationIgnored(t *testing.T) { clusterCopy.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady, ) - require.NotNil(t, cond) - assert.Equal(t, metav1.ConditionTrue, cond.Status) - assert.Equal(t, multigresv1alpha1.ReasonEndpointReady, cond.Reason) - assert.Contains(t, cond.Message, "lb.example.com") + c.Require().NotNil(cond) + c.EqDeep(metav1.ConditionTrue, cond.Status) + c.EqDeep(multigresv1alpha1.ReasonEndpointReady, cond.Reason) + c.StrContains(cond.Message, "lb.example.com") // Now verify that if we remove the fresh cell (only stale remains), // the aggregate is 0 and condition becomes NoReadyGateways. @@ -1171,15 +1135,15 @@ func TestUpdateStatus_StaleCellGenerationIgnored(t *testing.T) { clusterCopy2 := cluster.DeepCopy() err = r2.updateStatus(t.Context(), clusterCopy2) - require.NoError(t, err) + c.Require().NoError(err) cond2 := meta.FindStatusCondition( clusterCopy2.Status.Conditions, multigresv1alpha1.ConditionGatewayExternalReady, ) - require.NotNil(t, cond2) - assert.Equal(t, metav1.ConditionFalse, cond2.Status) - assert.Equal(t, multigresv1alpha1.ReasonNoReadyGateways, cond2.Reason) + c.Require().NotNil(cond2) + c.EqDeep(metav1.ConditionFalse, cond2.Status) + c.EqDeep(multigresv1alpha1.ReasonNoReadyGateways, cond2.Reason) } func TestComputeAdminWebCondition(t *testing.T) { @@ -1235,6 +1199,7 @@ func TestComputeAdminWebCondition(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := computeAdminWebCondition( tc.enabled, tc.externalEndpoint, @@ -1243,16 +1208,16 @@ func TestComputeAdminWebCondition(t *testing.T) { ) if tc.wantNil { - assert.Nil(t, got) + c.Nil(got) return } - require.NotNil(t, got) - assert.Equal(t, multigresv1alpha1.ConditionAdminWebExternalReady, got.Type) - assert.Equal(t, tc.wantStatus, got.Status) - assert.Equal(t, tc.wantReason, got.Reason) - assert.Contains(t, got.Message, tc.wantMessageContains) - assert.Equal(t, tc.clusterGeneration, got.ObservedGeneration) + c.Require().NotNil(got) + c.EqDeep(multigresv1alpha1.ConditionAdminWebExternalReady, got.Type) + c.EqDeep(tc.wantStatus, got.Status) + c.EqDeep(tc.wantReason, got.Reason) + c.StrContains(got.Message, tc.wantMessageContains) + c.EqDeep(tc.clusterGeneration, got.ObservedGeneration) }) } } @@ -1276,6 +1241,7 @@ func TestUpdateStatus_AdminWebServiceErrors(t *testing.T) { } t.Run("non-NotFound error fetching admin-web service returns error", func(t *testing.T) { + c := assert.NewCollecting(t) injectedErr := fmt.Errorf("simulated API server error") cl := fake.NewClientBuilder(). @@ -1296,11 +1262,12 @@ func TestUpdateStatus_AdminWebServiceErrors(t *testing.T) { cluster := baseCluster.DeepCopy() err := r.updateStatus(t.Context(), cluster) - require.Error(t, err) - assert.Contains(t, err.Error(), "failed to get multiadmin-web service for status") + c.Require().Error(err) + c.StrContains(err.Error(), "failed to get multiadmin-web service for status") }) t.Run("NotFound admin-web service sets AwaitingEndpoint when enabled", func(t *testing.T) { + c := assert.NewCollecting(t) cl := fake.NewClientBuilder(). WithScheme(scheme). WithObjects(baseCluster.DeepCopy()). @@ -1315,21 +1282,22 @@ func TestUpdateStatus_AdminWebServiceErrors(t *testing.T) { cluster := baseCluster.DeepCopy() err := r.updateStatus(t.Context(), cluster) - require.NoError(t, err) + c.Require().NoError(err) - require.NotNil(t, cluster.Status.AdminWeb) - assert.Empty(t, cluster.Status.AdminWeb.ExternalEndpoint) + c.Require().NotNil(cluster.Status.AdminWeb) + c.Empty(cluster.Status.AdminWeb.ExternalEndpoint) cond := meta.FindStatusCondition( cluster.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady, ) - require.NotNil(t, cond) - assert.Equal(t, metav1.ConditionFalse, cond.Status) - assert.Equal(t, multigresv1alpha1.ReasonAwaitingEndpoint, cond.Reason) + c.Require().NotNil(cond) + c.EqDeep(metav1.ConditionFalse, cond.Status) + c.EqDeep(multigresv1alpha1.ReasonAwaitingEndpoint, cond.Reason) }) t.Run("non-NotFound error fetching admin-web deployment returns error", func(t *testing.T) { + c := assert.NewCollecting(t) injectedErr := fmt.Errorf("simulated API server error") awSvc := &corev1.Service{ @@ -1360,12 +1328,13 @@ func TestUpdateStatus_AdminWebServiceErrors(t *testing.T) { cluster := baseCluster.DeepCopy() err := r.updateStatus(t.Context(), cluster) - require.Error(t, err) + c.Require().Error(err) // The Service Get will fail first since both have the same name - assert.Contains(t, err.Error(), "failed to get multiadmin-web") + c.StrContains(err.Error(), "failed to get multiadmin-web") }) t.Run("admin-web deployment ready replicas drives condition", func(t *testing.T) { + c := assert.NewCollecting(t) awSvc := &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ Name: clusterName + "-multiadmin-web", @@ -1400,17 +1369,17 @@ func TestUpdateStatus_AdminWebServiceErrors(t *testing.T) { cluster := baseCluster.DeepCopy() err := r.updateStatus(t.Context(), cluster) - require.NoError(t, err) + c.Require().NoError(err) - require.NotNil(t, cluster.Status.AdminWeb) - assert.Equal(t, "10.0.0.1", cluster.Status.AdminWeb.ExternalEndpoint) + c.Require().NotNil(cluster.Status.AdminWeb) + c.EqDeep("10.0.0.1", cluster.Status.AdminWeb.ExternalEndpoint) cond := meta.FindStatusCondition( cluster.Status.Conditions, multigresv1alpha1.ConditionAdminWebExternalReady, ) - require.NotNil(t, cond) - assert.Equal(t, metav1.ConditionTrue, cond.Status) - assert.Equal(t, multigresv1alpha1.ReasonEndpointReady, cond.Reason) + c.Require().NotNil(cond) + c.EqDeep(metav1.ConditionTrue, cond.Status) + c.EqDeep(multigresv1alpha1.ReasonEndpointReady, cond.Reason) }) } diff --git a/pkg/cluster-handler/controller/tablegroup/builders_test.go b/pkg/cluster-handler/controller/tablegroup/builders_test.go index 99ba1d92..c3d9e931 100644 --- a/pkg/cluster-handler/controller/tablegroup/builders_test.go +++ b/pkg/cluster-handler/controller/tablegroup/builders_test.go @@ -3,7 +3,6 @@ package tablegroup import ( "testing" - "github.com/google/go-cmp/cmp" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" "k8s.io/utils/ptr" @@ -11,6 +10,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestBuildShard(t *testing.T) { @@ -41,10 +42,9 @@ func TestBuildShard(t *testing.T) { } t.Run("Success", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildShard(tg, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } + c.Require().NoError(err, "BuildShard() error =") // Calculate expected hash: md5("my-cluster", "my-db", "my-tg", "shard-0") -> "a068d59f" expectedName := name.JoinWithConstraints( @@ -54,79 +54,47 @@ func TestBuildShard(t *testing.T) { "my-tg", "shard-0", ) - if got.Name != expectedName { - t.Errorf("Name = %v, want %v", got.Name, expectedName) - } - if got.Namespace != "default" { - t.Errorf("Namespace = %v, want %v", got.Namespace, "default") - } - if got.Labels["multigres.com/cluster"] != "my-cluster" { - t.Errorf( - "Labels[cluster] = %v, want %v", - got.Labels["multigres.com/cluster"], - "my-cluster", - ) - } + c.Eq(expectedName, got.Name, "Name") + c.Eq("default", got.Namespace, "Namespace") + c.Eq("my-cluster", got.Labels["multigres.com/cluster"], "Labels[cluster]") // Verify OwnerReference pointing to TableGroup if len(got.OwnerReferences) != 1 { t.Errorf("OwnerReferences count = %v, want 1", len(got.OwnerReferences)) } else { - if got.OwnerReferences[0].Name != "my-tg" { - t.Errorf("OwnerReference Name = %v, want %v", got.OwnerReferences[0].Name, "my-tg") - } - if got.OwnerReferences[0].UID != "tg-uid" { - t.Errorf("OwnerReference UID = %v, want %v", got.OwnerReferences[0].UID, "tg-uid") - } + c.Eq("my-tg", got.OwnerReferences[0].Name, "OwnerReference Name") + c.Eq("tg-uid", got.OwnerReferences[0].UID, "OwnerReference UID") } // Verify Spec fields are copied - if got.Spec.ShardName != "shard-0" { - t.Errorf("Spec.ShardName = %v, want %v", got.Spec.ShardName, "shard-0") - } - if got.Spec.DatabaseName != "my-db" { - t.Errorf("Spec.DatabaseName = %v, want %v", got.Spec.DatabaseName, "my-db") - } - if diff := cmp.Diff(shardSpec.Pools, got.Spec.Pools); diff != "" { - t.Errorf("Spec.Pools mismatch (-want +got):\n%s", diff) - } + c.Eq("shard-0", got.Spec.ShardName, "Spec.ShardName") + c.Eq("my-db", got.Spec.DatabaseName, "Spec.DatabaseName") + c.EqDiff(shardSpec.Pools, got.Spec.Pools, "Spec.Pools mismatch") }) t.Run("DurabilityPolicy propagates from TableGroup to Shard", func(t *testing.T) { + c := assert.NewCollecting(t) tgWithPolicy := tg.DeepCopy() tgWithPolicy.Spec.DurabilityPolicy = "MULTI_CELL_AT_LEAST_2" got, err := BuildShard(tgWithPolicy, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - if got.Spec.DurabilityPolicy != "MULTI_CELL_AT_LEAST_2" { - t.Errorf( - "Spec.DurabilityPolicy = %v, want MULTI_CELL_AT_LEAST_2", - got.Spec.DurabilityPolicy, - ) - } + c.Require().NoError(err, "BuildShard() error =") + c.Eq("MULTI_CELL_AT_LEAST_2", got.Spec.DurabilityPolicy, "Spec.DurabilityPolicy") }) t.Run("InternalTLS propagates from TableGroup to Shard", func(t *testing.T) { + c := assert.NewAborting(t) tgWithInternalTLS := tg.DeepCopy() tgWithInternalTLS.Spec.InternalTLS = &multigresv1alpha1.InternalTLSConfig{ Enabled: ptr.To(true), } got, err := BuildShard(tgWithInternalTLS, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - if got.Spec.InternalTLS != tgWithInternalTLS.Spec.InternalTLS { - t.Fatalf( - "Spec.InternalTLS = %#v, want propagated pointer %#v", - got.Spec.InternalTLS, - tgWithInternalTLS.Spec.InternalTLS, - ) - } + c.NoError(err, "BuildShard() error =") + c.Eq(tgWithInternalTLS.Spec.InternalTLS, got.Spec.InternalTLS, "Spec.InternalTLS") }) t.Run("PostgresPasswordSecretRef propagates from TableGroup to Shard", func(t *testing.T) { + c := assert.NewCollecting(t) tgWithRef := tg.DeepCopy() tgWithRef.Spec.PostgresPasswordSecretRef = multigresv1alpha1.PostgresPasswordSecretRef{ Name: "multigres-admin-password", @@ -134,26 +102,23 @@ func TestBuildShard(t *testing.T) { } got, err := BuildShard(tgWithRef, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - if got.Spec.PostgresPasswordSecretRef.Name != "multigres-admin-password" { - t.Errorf( - "Spec.PostgresPasswordSecretRef.Name = %v, want multigres-admin-password", - got.Spec.PostgresPasswordSecretRef.Name, - ) - } - if got.Spec.PostgresPasswordSecretRef.Key != "current" { - t.Errorf( - "Spec.PostgresPasswordSecretRef.Key = %v, want current", - got.Spec.PostgresPasswordSecretRef.Key, - ) - } + c.Require().NoError(err, "BuildShard() error =") + c.Eq( + "multigres-admin-password", + got.Spec.PostgresPasswordSecretRef.Name, + "Spec.PostgresPasswordSecretRef.Name", + ) + c.Eq( + "current", + got.Spec.PostgresPasswordSecretRef.Key, + "Spec.PostgresPasswordSecretRef.Key", + ) }) t.Run( "PostgresInitSecretsRef propagates from TableGroup to Shard when set", func(t *testing.T) { + c := assert.NewCollecting(t) tgWithRef := tg.DeepCopy() tgWithRef.Spec.PostgresInitSecretsRef = &multigresv1alpha1.PostgresInitSecretsRef{ Name: "multigres-init-secrets", @@ -161,77 +126,62 @@ func TestBuildShard(t *testing.T) { } got, err := BuildShard(tgWithRef, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - if got.Spec.PostgresInitSecretsRef == nil { - t.Fatal("Spec.PostgresInitSecretsRef = nil, want propagated ref") - } - if got.Spec.PostgresInitSecretsRef.Name != "multigres-init-secrets" { - t.Errorf( - "Spec.PostgresInitSecretsRef.Name = %v, want multigres-init-secrets", - got.Spec.PostgresInitSecretsRef.Name, - ) - } - if got.Spec.PostgresInitSecretsRef.Key != "init-secrets.json" { - t.Errorf( - "Spec.PostgresInitSecretsRef.Key = %v, want init-secrets.json", - got.Spec.PostgresInitSecretsRef.Key, - ) - } + c.Require().NoError(err, "BuildShard() error =") + c.Require(). + NotNil(got.Spec.PostgresInitSecretsRef, "Spec.PostgresInitSecretsRef = nil, want propagated ref") + c.Eq( + "multigres-init-secrets", + got.Spec.PostgresInitSecretsRef.Name, + "Spec.PostgresInitSecretsRef.Name", + ) + c.Eq( + "init-secrets.json", + got.Spec.PostgresInitSecretsRef.Key, + "Spec.PostgresInitSecretsRef.Key", + ) }, ) t.Run("PostgresInitSecretsRef nil on Shard when unset on TableGroup", func(t *testing.T) { + c := assert.NewCollecting(t) tgWithoutRef := tg.DeepCopy() tgWithoutRef.Spec.PostgresInitSecretsRef = nil got, err := BuildShard(tgWithoutRef, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - if got.Spec.PostgresInitSecretsRef != nil { - t.Errorf("Spec.PostgresInitSecretsRef = %+v, want nil", got.Spec.PostgresInitSecretsRef) - } + c.Require().NoError(err, "BuildShard() error =") + c.Nil(got.Spec.PostgresInitSecretsRef, "Spec.PostgresInitSecretsRef") }) t.Run("DurabilityPolicy empty when not set on TableGroup", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildShard(tg, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - if got.Spec.DurabilityPolicy != "" { - t.Errorf("Spec.DurabilityPolicy = %v, want empty", got.Spec.DurabilityPolicy) - } + c.Require().NoError(err, "BuildShard() error =") + c.Eq("", got.Spec.DurabilityPolicy, "Spec.DurabilityPolicy") }) t.Run("ControllerRefError", func(t *testing.T) { emptyScheme := runtime.NewScheme() // Intentionally missing scheme registrations to force SetControllerReference failure _, err := BuildShard(tg, shardSpec, emptyScheme) - if err == nil { - t.Error("Expected error due to missing scheme types, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error due to missing scheme types, got nil") }) t.Run("Propagates explicit project ref annotation", func(t *testing.T) { + c := assert.NewAborting(t) tgWithProjectRef := tg.DeepCopy() tgWithProjectRef.Annotations = map[string]string{ metadata.AnnotationProjectRef: "proj_123", } got, err := BuildShard(tgWithProjectRef, shardSpec, scheme) - if err != nil { - t.Fatalf("BuildShard() error = %v", err) - } - - if got.Annotations[metadata.AnnotationProjectRef] != "proj_123" { - t.Fatalf( - "annotation %q = %q, want %q", - metadata.AnnotationProjectRef, - got.Annotations[metadata.AnnotationProjectRef], - "proj_123", - ) - } + c.NoError(err, "BuildShard() error =") + + c.Eq( + "proj_123", + got.Annotations[metadata.AnnotationProjectRef], + "annotation %q = %q, want", + metadata.AnnotationProjectRef, + got.Annotations[metadata.AnnotationProjectRef], + ) }) } diff --git a/pkg/cluster-handler/controller/tablegroup/integration_consistency_test.go b/pkg/cluster-handler/controller/tablegroup/integration_consistency_test.go index 797c0549..a3caa169 100644 --- a/pkg/cluster-handler/controller/tablegroup/integration_consistency_test.go +++ b/pkg/cluster-handler/controller/tablegroup/integration_consistency_test.go @@ -21,6 +21,8 @@ import ( "github.com/multigres/multigres-operator/pkg/cluster-handler/controller/tablegroup" "github.com/multigres/multigres-operator/pkg/testutil" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) // TestTableGroup_ConsistencyConvergence runs the controller against a real @@ -31,6 +33,7 @@ import ( // tests. func TestTableGroup_ConsistencyConvergence(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) globalTopo := multigresv1alpha1.GlobalTopoServerRef{ Address: "etcd-client:2379", @@ -47,13 +50,11 @@ func TestTableGroup_ConsistencyConvergence(t *testing.T) { testutil.WithCRDPaths("../../../../config/crd/bases"), ) - if err := (&tablegroup.TableGroupReconciler{ + c.NoError((&tablegroup.TableGroupReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("tablegroup-controller"), - }).SetupWithManager(mgr, controller.Options{SkipNameValidation: ptr.To(true)}); err != nil { - t.Fatalf("Failed to set up controller: %v", err) - } + }).SetupWithManager(mgr, controller.Options{SkipNameValidation: ptr.To(true)}), "Failed to set up controller") watcher := testutil.NewResourceWatcher(t, t.Context(), mgr, testutil.WithCmpOpts(testutil.IgnoreMetaRuntimeFields()), @@ -123,9 +124,7 @@ func TestTableGroup_ConsistencyConvergence(t *testing.T) { // The initial three-shard spec establishes a healthy terminal baseline under // a real apiserver and cache before the test introduces an add/remove change. - if err := k8sClient.Create(ctx, tg); err != nil { - t.Fatalf("Failed to create TableGroup: %v", err) - } + c.NoError(k8sClient.Create(ctx, tg), "Failed to create TableGroup") waitForChildShards(t, ctx, k8sClient, namespace, clusterName, dbName, tgName, []string{childName("s1"), childName("s2"), childName("s3")}) @@ -159,7 +158,7 @@ func TestTableGroup_ConsistencyConvergence(t *testing.T) { // The add/remove update forces the controller to create one desired child, // retire one orphan through PendingDeletion, and keep status truthful while // both old and new children may briefly exist. - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + c.NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { latest := &multigresv1alpha1.TableGroup{} if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(tg), latest); err != nil { return err @@ -170,9 +169,7 @@ func TestTableGroup_ConsistencyConvergence(t *testing.T) { shardSpec("s4", "zone-d"), } return k8sClient.Update(ctx, latest) - }); err != nil { - t.Fatalf("Failed to update TableGroup spec: %v", err) - } + }), "Failed to update TableGroup spec") // The old child can still exist while it drains, so the useful assertion here // is that the new desired child appears, not that the child set is already @@ -185,7 +182,7 @@ func TestTableGroup_ConsistencyConvergence(t *testing.T) { waitForShardAnnotation(t, ctx, k8sClient, s3Name, namespace, multigresv1alpha1.AnnotationPendingDeletion) - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + c.NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { latest := &multigresv1alpha1.Shard{} if err := k8sClient.Get( ctx, @@ -203,15 +200,11 @@ func TestTableGroup_ConsistencyConvergence(t *testing.T) { } latest.Status.Conditions = append(latest.Status.Conditions, cond) return k8sClient.Status().Update(ctx, latest) - }); err != nil { - t.Fatalf("Failed to set ReadyForDeletion on orphan shard: %v", err) - } + }), "Failed to set ReadyForDeletion on orphan shard") - if err := watcher.WaitForDeletion(&multigresv1alpha1.Shard{ + c.NoError(watcher.WaitForDeletion(&multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: s3Name, Namespace: namespace}, - }); err != nil { - t.Fatalf("Orphan shard s3 was not pruned: %v", err) - } + }), "Orphan shard s3 was not pruned") // Once the replacement child reports Healthy, the parent should converge to a // stable terminal status for the new generation. @@ -275,9 +268,8 @@ func waitForChildShards( } } - if time.Now().After(deadline) { - t.Fatalf("timed out waiting for child shards %v", wantNames) - } + assert.NewAborting(t). + False(time.Now().After(deadline), "timed out waiting for child shards %v", wantNames) time.Sleep(200 * time.Millisecond) } } @@ -293,7 +285,7 @@ func driveShardHealthy( ) { t.Helper() - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + assert.NewAborting(t).NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { latest := &multigresv1alpha1.Shard{} if err := k8sClient.Get( ctx, @@ -306,9 +298,7 @@ func driveShardHealthy( latest.Status.Message = "Ready" latest.Status.ObservedGeneration = latest.Generation return k8sClient.Status().Update(ctx, latest) - }); err != nil { - t.Fatalf("Failed to drive shard %s healthy: %v", name, err) - } + }), "Failed to drive shard %s healthy", name) } // waitForShardExists polls until the named Shard exists. @@ -330,9 +320,8 @@ func waitForShardExists( ); err == nil { return } - if time.Now().After(deadline) { - t.Fatalf("timed out waiting for shard %s to exist", name) - } + assert.NewAborting(t). + False(time.Now().After(deadline), "timed out waiting for shard %s to exist", name) time.Sleep(200 * time.Millisecond) } } @@ -359,9 +348,8 @@ func waitForShardAnnotation( return } } - if time.Now().After(deadline) { - t.Fatalf("timed out waiting for shard %s annotation %q", name, annotation) - } + assert.NewAborting(t). + False(time.Now().After(deadline), "timed out waiting for shard %s annotation %q", name, annotation) time.Sleep(200 * time.Millisecond) } } @@ -388,16 +376,8 @@ func waitForTableGroup( return g } } - if time.Now().After(deadline) { - t.Fatalf( - "timed out waiting for TableGroup to converge; last observed phase=%q ready=%d total=%d observedGen=%d gen=%d", - last.Status.Phase, - last.Status.ReadyShards, - last.Status.TotalShards, - last.Status.ObservedGeneration, - last.Generation, - ) - } + assert.NewAborting(t). + False(time.Now().After(deadline), "timed out waiting for TableGroup to converge; last observed phase=%q ready=%d total=%d observedGen=%d gen=%d", last.Status.Phase, last.Status.ReadyShards, last.Status.TotalShards, last.Status.ObservedGeneration, last.Generation) time.Sleep(200 * time.Millisecond) } } @@ -418,9 +398,8 @@ func assertTableGroupStaysStable( deadline := time.Now().Add(window) for { g := &multigresv1alpha1.TableGroup{} - if err := k8sClient.Get(ctx, key, g); err != nil { - t.Fatalf("Failed to get TableGroup during stability check: %v", err) - } + assert.NewAborting(t). + NoError(k8sClient.Get(ctx, key, g), "Failed to get TableGroup during stability check") if !predicate(g) { t.Fatalf( "TableGroup left its stable terminal state: phase=%q observedGen=%d gen=%d", diff --git a/pkg/cluster-handler/controller/tablegroup/integration_lifecycle_test.go b/pkg/cluster-handler/controller/tablegroup/integration_lifecycle_test.go index 95146a4c..cfbaf6ed 100644 --- a/pkg/cluster-handler/controller/tablegroup/integration_lifecycle_test.go +++ b/pkg/cluster-handler/controller/tablegroup/integration_lifecycle_test.go @@ -21,6 +21,8 @@ import ( "github.com/multigres/multigres-operator/pkg/cluster-handler/controller/tablegroup" "github.com/multigres/multigres-operator/pkg/testutil" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestTableGroup_Lifecycle(t *testing.T) { @@ -47,13 +49,11 @@ func TestTableGroup_Lifecycle(t *testing.T) { testutil.WithCRDPaths("../../../../config/crd/bases"), ) - if err := (&tablegroup.TableGroupReconciler{ + assert.NewAborting(t).NoError((&tablegroup.TableGroupReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("tablegroup-controller"), - }).SetupWithManager(mgr, controller.Options{SkipNameValidation: ptr.To(true)}); err != nil { - t.Fatal(err) - } + }).SetupWithManager(mgr, controller.Options{SkipNameValidation: ptr.To(true)})) watcher := testutil.NewResourceWatcher(t, t.Context(), mgr, testutil.WithCmpOpts(testutil.IgnoreMetaRuntimeFields()), @@ -65,6 +65,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { t.Run("Pruning and Same-Name Replacement", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setup(t) ctx := t.Context() @@ -110,9 +111,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { // Begin with two desired children so removing one from the spec later has // to go through the same graceful-deletion path used in production. setTestPostgresPasswordSecretRef(tg) - if err := k8sClient.Create(ctx, tg); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(ctx, tg)) // The expected Shards describe the child specs TableGroup owns. Metadata // and status are intentionally ignored because those are filled by the @@ -178,15 +177,13 @@ func TestTableGroup_Lifecycle(t *testing.T) { // Compare only spec fields because the controller owns metadata and status. watcher.SetCmpOpts(testutil.CompareSpecOnly()...) - if err := watcher.WaitForMatch(shard1, shard2); err != nil { - t.Fatalf("Failed to create initial shards: %v", err) - } + c.Require().NoError(watcher.WaitForMatch(shard1, shard2), "Failed to create initial shards") // Removing a child from the desired spec should not delete the Shard // immediately. The parent first asks the child to drain by preserving the // object and adding PendingDeletion. // Retry because the controller may update status or finalizers concurrently. - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + c.Require().NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(tg), tg); err != nil { return err } @@ -200,9 +197,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { }, } return k8sClient.Update(ctx, tg) - }); err != nil { - t.Fatal(err) - } + })) // The shard controller is not running in this test. Once the parent has // asked for a drain, the test simulates the child controller's @@ -242,7 +237,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { // Re-adding the same logical shard while the old object is draining must // not resurrect or update that object. Cleanup remains the active // lifecycle until the child reports ReadyForDeletion and is deleted. - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + c.Require().NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(tg), tg); err != nil { return err } @@ -263,9 +258,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { }, } return k8sClient.Update(ctx, tg) - }); err != nil { - t.Fatal(err) - } + })) waitForTableGroup(t, ctx, k8sClient, client.ObjectKeyFromObject(tg), 15*time.Second, func(g *multigresv1alpha1.TableGroup) bool { @@ -274,25 +267,20 @@ func TestTableGroup_Lifecycle(t *testing.T) { g.Status.Message == "Waiting for shard cleanup to finish" }) - if err := k8sClient.Get( + c.Require().NoError(k8sClient.Get( ctx, client.ObjectKey{Name: deleteMeName, Namespace: "default"}, &deleteMe, - ); err != nil { - t.Fatalf("Failed to get draining shard after re-add: %v", err) - } - if got := string(deleteMe.UID); got != oldDeleteMeUID { - t.Fatalf("Draining shard was replaced before cleanup: got UID %s, want %s", - got, oldDeleteMeUID) - } - if deleteMe.Annotations[multigresv1alpha1.AnnotationPendingDeletion] == "" { - t.Fatal("Draining shard lost PendingDeletion after same-name re-add") - } + ), "Failed to get draining shard after re-add") + c.Require(). + Eq(oldDeleteMeUID, string(deleteMe.UID), "Draining shard was replaced before cleanup: got UID") + c.Require(). + NotEq("", deleteMe.Annotations[multigresv1alpha1.AnnotationPendingDeletion], "Draining shard lost PendingDeletion after same-name re-add") if got := deleteMe.Spec.Multiorch.StatelessSpec.Replicas; got == nil || *got != 1 { t.Fatalf("Draining shard was updated during cleanup: replicas = %v, want 1", got) } - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + c.Require().NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { latest := &multigresv1alpha1.Shard{} if err := k8sClient.Get( ctx, @@ -309,14 +297,10 @@ func TestTableGroup_Lifecycle(t *testing.T) { LastTransitionTime: metav1.Now(), }) return k8sClient.Status().Update(ctx, latest) - }); err != nil { - t.Fatalf("Failed to set ReadyForDeletion condition: %v", err) - } + }), "Failed to set ReadyForDeletion condition") // Deletion is only valid after the child has acknowledged the drain. - if err := watcher.WaitForDeletion(shard2); err != nil { - t.Errorf("Shard 'delete-me' was not pruned: %v", err) - } + c.NoError(watcher.WaitForDeletion(shard2), "Shard 'delete-me' was not pruned") waitCtx, waitCancel = context.WithTimeout(ctx, 15*time.Second) defer waitCancel() @@ -354,6 +338,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { t.Run("Enforcement (Revert Manual Changes)", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) k8sClient, watcher := setup(t) ctx := t.Context() @@ -393,9 +378,7 @@ func TestTableGroup_Lifecycle(t *testing.T) { // desired shape. setTestPostgresPasswordSecretRef(tg) - if err := k8sClient.Create(ctx, tg); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(ctx, tg)) // Compare only spec fields because the controller owns metadata and status. watcher.SetCmpOpts(testutil.CompareSpecOnly()...) @@ -429,14 +412,12 @@ func TestTableGroup_Lifecycle(t *testing.T) { }, } setTestShardPostgresPasswordSecretRef(goodShard) - if err := watcher.WaitForMatch(goodShard); err != nil { - t.Fatalf("Initial shard creation failed: %v", err) - } + c.Require().NoError(watcher.WaitForMatch(goodShard), "Initial shard creation failed") // Mutate the child directly, as an out-of-band actor might. The retry is // for resourceVersion conflicts with the controller's own writes. latestShard := &multigresv1alpha1.Shard{} - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + c.Require().NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { if err := k8sClient.Get( ctx, client.ObjectKeyFromObject(goodShard), @@ -446,14 +427,13 @@ func TestTableGroup_Lifecycle(t *testing.T) { } latestShard.Spec.Multiorch.Replicas = ptr.To(int32(99)) return k8sClient.Update(ctx, latestShard) - }); err != nil { - t.Fatal(err) - } + })) // Assert the desired spec again rather than trying to observe the bad // intermediate state; the controller may repair it before the watch sees it. - if err := watcher.WaitForMatch(goodShard); err != nil { - t.Errorf("Controller failed to revert manual shard change: %v", err) - } + c.NoError( + watcher.WaitForMatch(goodShard), + "Controller failed to revert manual shard change", + ) }) } diff --git a/pkg/cluster-handler/controller/tablegroup/integration_test.go b/pkg/cluster-handler/controller/tablegroup/integration_test.go index 30fa8762..48d0acd5 100644 --- a/pkg/cluster-handler/controller/tablegroup/integration_test.go +++ b/pkg/cluster-handler/controller/tablegroup/integration_test.go @@ -21,6 +21,8 @@ import ( "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestSetupWithManager(t *testing.T) { @@ -37,15 +39,13 @@ func TestSetupWithManager(t *testing.T) { ), ) - if err := (&tablegroup.TableGroupReconciler{ + assert.NewAborting(t).NoError((&tablegroup.TableGroupReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("tablegroup-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") } func TestSetupWithManager_Failure(t *testing.T) { @@ -63,15 +63,13 @@ func TestSetupWithManager_Failure(t *testing.T) { ), ) - if err := (&tablegroup.TableGroupReconciler{ + assert.NewAborting(t).Error((&tablegroup.TableGroupReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("tablegroup-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err == nil { - t.Fatal("Expected SetupWithManager to fail due to missing type in scheme, got nil") - } + }), "Expected SetupWithManager to fail due to missing type in scheme, got nil") } func setTestPostgresPasswordSecretRef(tableGroup *multigresv1alpha1.TableGroup) { @@ -248,6 +246,7 @@ func TestTableGroupReconciliation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) ctx := t.Context() mgr := testutil.SetUpEnvtestManager(t, scheme, @@ -274,16 +273,13 @@ func TestTableGroupReconciliation(t *testing.T) { Recorder: mgr.GetEventRecorderFor("tablegroup-controller"), } - if err := reconciler.SetupWithManager(mgr, controller.Options{ + c.Require().NoError(reconciler.SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") setTestPostgresPasswordSecretRef(tc.tableGroup) - if err := k8sClient.Create(ctx, tc.tableGroup); err != nil { - t.Fatalf("Failed to create the initial tablegroup, %v", err) - } + c.Require(). + NoError(k8sClient.Create(ctx, tc.tableGroup), "Failed to create the initial tablegroup") // Expected Shard names use the same name constraints as the controller. for _, obj := range tc.wantResources { @@ -301,9 +297,7 @@ func TestTableGroupReconciliation(t *testing.T) { } } - if err := watcher.WaitForMatch(tc.wantResources...); err != nil { - t.Errorf("Resources mismatch:\n%v", err) - } + c.NoError(watcher.WaitForMatch(tc.wantResources...), "Resources mismatch:\n") }) } } diff --git a/pkg/cluster-handler/controller/tablegroup/reconcile_shards_test.go b/pkg/cluster-handler/controller/tablegroup/reconcile_shards_test.go index 55156c34..be39982b 100644 --- a/pkg/cluster-handler/controller/tablegroup/reconcile_shards_test.go +++ b/pkg/cluster-handler/controller/tablegroup/reconcile_shards_test.go @@ -19,6 +19,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) // TestShardStepsPrune covers the retained PendingDeletion handshake for child @@ -186,6 +188,7 @@ func TestShardStepsPrune(t *testing.T) { for tn, tc := range tests { t.Run(tn, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) var shardPatchCount atomic.Int64 var shardDeleteCount atomic.Int64 @@ -251,26 +254,18 @@ func TestShardStepsPrune(t *testing.T) { } } - if got := shardPatchCount.Load(); got != tc.wantPatches { - t.Errorf("Patch count mismatch: got %d, want %d", got, tc.wantPatches) - } - if got := shardDeleteCount.Load(); got != tc.wantDeletes { - t.Errorf("Delete count mismatch: got %d, want %d", got, tc.wantDeletes) - } - if rc.pendingDeletion != tc.wantPending { - t.Errorf("pending mismatch: got %t, want %t", rc.pendingDeletion, tc.wantPending) - } - if got := rc.activeShardNames[shardName]; got != tc.wantActive { - t.Errorf("active child mismatch: got %t, want %t", got, tc.wantActive) - } + ck.Eq(tc.wantPatches, shardPatchCount.Load(), "Patch count mismatch: got") + ck.Eq(tc.wantDeletes, shardDeleteCount.Load(), "Delete count mismatch: got") + ck.Eq(tc.wantPending, rc.pendingDeletion, "pending mismatch: got") + ck.Eq(tc.wantActive, rc.activeShardNames[shardName], "active child mismatch: got") res, err := stepRequeueIfPending(t.Context(), rc) - if err != nil { - t.Fatalf("stepRequeueIfPending returned error: %v", err) - } + ck.Require().NoError(err, "stepRequeueIfPending returned error") if tc.wantPending { - if !res.done || res.result.RequeueAfter != 5*time.Second { - t.Errorf("requeue mismatch: got %+v, want 5s requeue", res.result) - } + ck.False( + !res.done || res.result.RequeueAfter != 5*time.Second, + "requeue mismatch: got %+v, want 5s requeue", + res.result, + ) } else if res.done { t.Errorf("unexpected requeue result: %+v", res.result) } @@ -283,9 +278,7 @@ func TestShardStepsPrune(t *testing.T) { found = true } } - if !found { - t.Errorf("expected an event containing %q", tc.wantEvent) - } + ck.True(found, "expected an event containing %q", tc.wantEvent) } fetched := &multigresv1alpha1.Shard{} @@ -296,23 +289,22 @@ func TestShardStepsPrune(t *testing.T) { ) if tc.wantDeleted { - if !errors.IsNotFound(getErr) { - t.Errorf( - "expected child Shard to be deleted, but it still exists (err: %v)", - getErr, - ) - } + ck.True( + errors.IsNotFound(getErr), + "expected child Shard to be deleted, but it still exists (err: %v)", + getErr, + ) return } - if getErr != nil { - t.Fatalf("expected child Shard to still exist: %v", getErr) - } + ck.Require().NoError(getErr, "expected child Shard to still exist") got := fetched.Annotations[multigresv1alpha1.AnnotationPendingDeletion] if tc.wantAnnotationNonEmpty { - if got == "" { - t.Error("expected a freshly stamped PendingDeletion annotation, got empty") - } + ck.NotEq( + "", + got, + "expected a freshly stamped PendingDeletion annotation, got empty", + ) } else if got != tc.wantAnnotation { t.Errorf( "PendingDeletion annotation mismatch: got %q, want %q", diff --git a/pkg/cluster-handler/controller/tablegroup/status_test.go b/pkg/cluster-handler/controller/tablegroup/status_test.go index a15a7a66..210ef997 100644 --- a/pkg/cluster-handler/controller/tablegroup/status_test.go +++ b/pkg/cluster-handler/controller/tablegroup/status_test.go @@ -13,6 +13,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) // TestStepComputeStatus pins the status aggregation branches. Status must come @@ -137,6 +139,7 @@ func TestStepComputeStatus(t *testing.T) { for tn, tc := range tests { t.Run(tn, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) tg := &multigresv1alpha1.TableGroup{ ObjectMeta: metav1.ObjectMeta{ @@ -189,68 +192,32 @@ func TestStepComputeStatus(t *testing.T) { t.Fatalf("stepComputeStatus returned error: %v", err) } res, err := stepPatchStatus(t.Context(), rc) - if err != nil { - t.Fatalf("stepPatchStatus returned error: %v", err) - } + ck.Require().NoError(err, "stepPatchStatus returned error") // A successful status update never requeues on its own. - if res.result.RequeueAfter != 0 { - t.Errorf("expected no requeue, got %+v", res.result) - } + ck.Eq(0, res.result.RequeueAfter, "expected no requeue, got %+v", res.result) updated := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updated, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } + ), "failed to get tablegroup") - if got := updated.Status.Phase; got != tc.wantPhase { - t.Errorf("Phase mismatch: got %q, want %q", got, tc.wantPhase) - } - if got := updated.Status.TotalShards; got != tc.wantTotal { - t.Errorf("TotalShards mismatch: got %d, want %d", got, tc.wantTotal) - } - if got := updated.Status.ReadyShards; got != tc.wantReady { - t.Errorf("ReadyShards mismatch: got %d, want %d", got, tc.wantReady) - } - if got := updated.Status.Message; got != tc.wantMessage { - t.Errorf("Message mismatch: got %q, want %q", got, tc.wantMessage) - } + ck.Eq(tc.wantPhase, updated.Status.Phase, "Phase mismatch: got") + ck.Eq(tc.wantTotal, updated.Status.TotalShards, "TotalShards mismatch: got") + ck.Eq(tc.wantReady, updated.Status.ReadyShards, "ReadyShards mismatch: got") + ck.Eq(tc.wantMessage, updated.Status.Message, "Message mismatch: got") cond := meta.FindStatusCondition(updated.Status.Conditions, "Available") - if cond == nil { - t.Fatal("expected an Available condition to be set") - } - if cond.Status != tc.wantAvailable { - t.Errorf( - "Available condition status mismatch: got %q, want %q", - cond.Status, - tc.wantAvailable, - ) - } - if cond.Reason != tc.wantReason { - t.Errorf( - "Available condition reason mismatch: got %q, want %q", - cond.Reason, - tc.wantReason, - ) - } - if cond.Message != tc.wantCondMsg { - t.Errorf( - "Available condition message mismatch: got %q, want %q", - cond.Message, - tc.wantCondMsg, - ) - } - if cond.ObservedGeneration != tg.Generation { - t.Errorf( - "Available condition observedGeneration mismatch: got %d, want %d", - cond.ObservedGeneration, - tg.Generation, - ) - } + ck.Require().NotNil(cond, "expected an Available condition to be set") + ck.Eq(tc.wantAvailable, cond.Status, "Available condition status mismatch: got") + ck.Eq(tc.wantReason, cond.Reason, "Available condition reason mismatch: got") + ck.Eq(tc.wantCondMsg, cond.Message, "Available condition message mismatch: got") + ck.Eq( + tg.Generation, + cond.ObservedGeneration, + "Available condition observedGeneration mismatch: got", + ) }) } } @@ -260,6 +227,7 @@ func TestStepComputeStatus(t *testing.T) { // status. func TestStepComputeStatus_IgnoresUndesiredOrphans(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) tg := &multigresv1alpha1.TableGroup{ ObjectMeta: metav1.ObjectMeta{Name: "tg", Namespace: "default", Generation: 1}, @@ -292,29 +260,19 @@ func TestStepComputeStatus_IgnoresUndesiredOrphans(t *testing.T) { appliedGeneration: map[string]int64{"desired": 1}, } - if _, err := stepComputeStatus(t.Context(), rc); err != nil { - t.Fatalf("stepComputeStatus returned error: %v", err) - } + _, err := stepComputeStatus(t.Context(), rc) + c.Require().NoError(err, "stepComputeStatus returned error") - if got := tg.Status.Phase; got != multigresv1alpha1.PhaseHealthy { - t.Errorf( - "Phase mismatch: got %q, want %q (orphan must not count)", - got, - multigresv1alpha1.PhaseHealthy, - ) - } - if got := tg.Status.ReadyShards; got != 1 { - t.Errorf("ReadyShards mismatch: got %d, want 1", got) - } - if got := tg.Status.TotalShards; got != 1 { - t.Errorf("TotalShards mismatch: got %d, want 1", got) - } + c.Eq(multigresv1alpha1.PhaseHealthy, tg.Status.Phase, "Phase mismatch: got") + c.Eq(1, tg.Status.ReadyShards, "ReadyShards mismatch: got") + c.Eq(1, tg.Status.TotalShards, "TotalShards mismatch: got") } // TestStepComputeStatus_PendingDeletionPreventsHealthy verifies that pending // cleanup keeps the parent Progressing even when all desired children are ready. func TestStepComputeStatus_PendingDeletionPreventsHealthy(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) tg := &multigresv1alpha1.TableGroup{ ObjectMeta: metav1.ObjectMeta{Name: "tg", Namespace: "default", Generation: 2}, @@ -348,43 +306,26 @@ func TestStepComputeStatus_PendingDeletionPreventsHealthy(t *testing.T) { pendingDeletion: true, } - if _, err := stepComputeStatus(t.Context(), rc); err != nil { - t.Fatalf("stepComputeStatus returned error: %v", err) - } + _, err := stepComputeStatus(t.Context(), rc) + c.Require().NoError(err, "stepComputeStatus returned error") - if got := tg.Status.Phase; got != multigresv1alpha1.PhaseProgressing { - t.Errorf("Phase mismatch: got %q, want %q", got, multigresv1alpha1.PhaseProgressing) - } - if got, want := tg.Status.Message, "Waiting for shard cleanup to finish"; got != want { - t.Errorf("Message mismatch: got %q, want %q", got, want) - } - if got := tg.Status.ReadyShards; got != 1 { - t.Errorf("ReadyShards mismatch: got %d, want 1", got) - } + c.Eq(multigresv1alpha1.PhaseProgressing, tg.Status.Phase, "Phase mismatch: got") + got, want := tg.Status.Message, "Waiting for shard cleanup to finish" + assert.NewCollecting(t).Eq(want, got, "Message mismatch: got") + c.Eq(1, tg.Status.ReadyShards, "ReadyShards mismatch: got") cond := meta.FindStatusCondition(tg.Status.Conditions, "Available") - if cond == nil { - t.Fatal("expected an Available condition to be set") - } - if cond.Status != metav1.ConditionFalse { - t.Errorf("Available status mismatch: got %q, want %q", cond.Status, metav1.ConditionFalse) - } - if cond.Reason != "CleanupPending" { - t.Errorf("Available reason mismatch: got %q, want CleanupPending", cond.Reason) - } - if cond.ObservedGeneration != tg.Generation { - t.Errorf( - "Available observedGeneration mismatch: got %d, want %d", - cond.ObservedGeneration, - tg.Generation, - ) - } + c.Require().NotNil(cond, "expected an Available condition to be set") + c.Eq(metav1.ConditionFalse, cond.Status, "Available status mismatch: got") + c.Eq("CleanupPending", cond.Reason, "Available reason mismatch: got") + c.Eq(tg.Generation, cond.ObservedGeneration, "Available observedGeneration mismatch: got") } // TestStepComputeStatus_IgnoresChildUntilSpecChangeObserved verifies that a // just-changed child is not ready until its observed generation catches up. func TestStepComputeStatus_IgnoresChildUntilSpecChangeObserved(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) newTG := func() *multigresv1alpha1.TableGroup { return &multigresv1alpha1.TableGroup{ @@ -417,16 +358,12 @@ func TestStepComputeStatus_IgnoresChildUntilSpecChangeObserved(t *testing.T) { if _, err := stepComputeStatus(t.Context(), changed); err != nil { t.Fatalf("stepComputeStatus returned error: %v", err) } - if got := changed.tg.Status.Phase; got != multigresv1alpha1.PhaseProgressing { - t.Errorf( - "Phase mismatch after spec change: got %q, want %q", - got, - multigresv1alpha1.PhaseProgressing, - ) - } - if got := changed.tg.Status.ReadyShards; got != 0 { - t.Errorf("ReadyShards mismatch after spec change: got %d, want 0", got) - } + c.Eq( + multigresv1alpha1.PhaseProgressing, + changed.tg.Status.Phase, + "Phase mismatch after spec change: got", + ) + c.Eq(0, changed.tg.Status.ReadyShards, "ReadyShards mismatch after spec change: got") // Once applied and observed generations match, the child counts as ready. caughtUp := &reconcileContext{ @@ -435,17 +372,12 @@ func TestStepComputeStatus_IgnoresChildUntilSpecChangeObserved(t *testing.T) { activeShardNames: map[string]bool{"desired": true}, appliedGeneration: map[string]int64{"desired": 1}, } - if _, err := stepComputeStatus(t.Context(), caughtUp); err != nil { - t.Fatalf("stepComputeStatus returned error: %v", err) - } - if got := caughtUp.tg.Status.Phase; got != multigresv1alpha1.PhaseHealthy { - t.Errorf( - "Phase mismatch once observed: got %q, want %q", - got, - multigresv1alpha1.PhaseHealthy, - ) - } - if got := caughtUp.tg.Status.ReadyShards; got != 1 { - t.Errorf("ReadyShards mismatch once observed: got %d, want 1", got) - } + _, err := stepComputeStatus(t.Context(), caughtUp) + c.Require().NoError(err, "stepComputeStatus returned error") + c.Eq( + multigresv1alpha1.PhaseHealthy, + caughtUp.tg.Status.Phase, + "Phase mismatch once observed: got", + ) + c.Eq(1, caughtUp.tg.Status.ReadyShards, "ReadyShards mismatch once observed: got") } diff --git a/pkg/cluster-handler/controller/tablegroup/tablegroup_controller_test.go b/pkg/cluster-handler/controller/tablegroup/tablegroup_controller_test.go index 75fef08e..cd0ddbe9 100644 --- a/pkg/cluster-handler/controller/tablegroup/tablegroup_controller_test.go +++ b/pkg/cluster-handler/controller/tablegroup/tablegroup_controller_test.go @@ -25,6 +25,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func setupFixtures( @@ -103,6 +105,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "Normal Synced Successfully reconciled TableGroup", }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) ctx := t.Context() // The name is the md5 hash of test-cluster, db1, tg1, shard-0. shardNameFull := name.JoinWithConstraints( @@ -113,18 +116,15 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "shard-0", ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + ck.Require().NoError(c.Get( ctx, types.NamespacedName{Name: shardNameFull, Namespace: namespace}, shard, - ); err != nil { - t.Fatalf("Shard %s not created: %v", shardNameFull, err) - } - if got, want := shard.Spec.DatabaseName, multigresv1alpha1.DatabaseName( + ), "Shard %s not created", shardNameFull) + got, want := shard.Spec.DatabaseName, multigresv1alpha1.DatabaseName( dbName, - ); got != want { - t.Errorf("Shard DB name mismatch got %q, want %q", got, want) - } + ) + ck.Eq(want, got, "Shard DB name mismatch got") }, }, @@ -165,17 +165,15 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "Normal Synced Successfully reconciled TableGroup", }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if got, want := updatedTG.Status.ReadyShards, int32(1); got != want { - t.Errorf("ReadyShards mismatch got %d, want %d", got, want) - } + ), "failed to get tablegroup") + got, want := updatedTG.Status.ReadyShards, int32(1) + ck.Eq(want, got, "ReadyShards mismatch got") }, }, "Status: Partial Ready (Not all shards ready)": { @@ -196,20 +194,19 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if got, want := updatedTG.Status.ReadyShards, int32(0); got != want { - t.Errorf("ReadyShards mismatch got %d, want %d", got, want) - } - if meta.IsStatusConditionTrue(updatedTG.Status.Conditions, "Available") { - t.Error("TableGroup should NOT be Available") - } + ), "failed to get tablegroup") + got, want := updatedTG.Status.ReadyShards, int32(0) + ck.Eq(want, got, "ReadyShards mismatch got") + ck.False( + meta.IsStatusConditionTrue(updatedTG.Status.Conditions, "Available"), + "TableGroup should NOT be Available", + ) }, }, "Status: Shard Not Ready (False Condition)": { @@ -241,17 +238,15 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if got, want := updatedTG.Status.ReadyShards, int32(0); got != want { - t.Errorf("ReadyShards mismatch got %d, want %d", got, want) - } + ), "failed to get tablegroup") + got, want := updatedTG.Status.ReadyShards, int32(0) + ck.Eq(want, got, "ReadyShards mismatch got") }, }, "Status: Degraded Shard": { @@ -280,17 +275,15 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if got, want := updatedTG.Status.Phase, multigresv1alpha1.PhaseDegraded; got != want { - t.Errorf("Phase mismatch got %q, want %q", got, want) - } + ), "failed to get tablegroup") + got, want := updatedTG.Status.Phase, multigresv1alpha1.PhaseDegraded + ck.Eq(want, got, "Phase mismatch got") }, }, "Status: Zero Shards (Vacuously True)": { @@ -300,17 +293,17 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { }, existingObjects: []client.Object{}, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if !meta.IsStatusConditionTrue(updatedTG.Status.Conditions, "Available") { - t.Error("Zero shard TableGroup should be Available") - } + ), "failed to get tablegroup") + ck.True( + meta.IsStatusConditionTrue(updatedTG.Status.Conditions, "Available"), + "Zero shard TableGroup should be Available", + ) }, }, @@ -345,6 +338,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { }, }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} shardName := name.JoinWithConstraints( name.DefaultConstraints, @@ -353,16 +347,12 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { tgLabelName, "shard-0", ) - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: shardName, Namespace: namespace}, shard, - ); err != nil { - t.Fatal(err) - } - if *shard.Spec.Multiorch.Replicas != 5 { - t.Errorf("Shard replicas not updated") - } + )) + ck.Eq(5, *shard.Spec.Multiorch.Replicas, "Shard replicas not updated") }, }, "Success: Early Return on Deletion": { @@ -378,15 +368,14 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { validate: func(t testing.TB, c client.Client) { shardNameFull := fmt.Sprintf("%s-%s", tgName, "shard-0") shard := &multigresv1alpha1.Shard{} - if err := c.Get( + err := c.Get( t.Context(), types.NamespacedName{Name: shardNameFull, Namespace: namespace}, shard, - ); !apierrors.IsNotFound( + ) + assert.NewCollecting(t).True(apierrors.IsNotFound( err, - ) { - t.Errorf("Expected Shard %s to NOT be created", shardNameFull) - } + ), "Expected Shard %s to NOT be created", shardNameFull) }, }, "Success: Prune Orphan Shard (Sets PendingDeletion)": { @@ -418,6 +407,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "Normal PendingDeletion Marked Shard", }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) shardName := name.JoinWithConstraints( name.DefaultConstraints, clusterName, @@ -426,16 +416,16 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "shard-0", ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: shardName, Namespace: namespace}, shard, - ); err != nil { - t.Fatalf("Shard should still exist: %v", err) - } - if shard.Annotations[multigresv1alpha1.AnnotationPendingDeletion] == "" { - t.Error("Expected PendingDeletion annotation to be set") - } + ), "Shard should still exist") + ck.NotEq( + "", + shard.Annotations[multigresv1alpha1.AnnotationPendingDeletion], + "Expected PendingDeletion annotation to be set", + ) }, }, @@ -478,13 +468,11 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "shard-0", ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + assert.NewAborting(t).NoError(c.Get( t.Context(), types.NamespacedName{Name: shardName, Namespace: namespace}, shard, - ); err != nil { - t.Fatalf("Shard should still exist while draining: %v", err) - } + ), "Shard should still exist while draining") }, }, @@ -538,13 +526,13 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "shard-0", ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + err := c.Get( t.Context(), types.NamespacedName{Name: shardName, Namespace: namespace}, shard, - ); !apierrors.IsNotFound(err) { - t.Errorf("Expected Shard %s to be deleted", shardName) - } + ) + assert.NewCollecting(t). + True(apierrors.IsNotFound(err), "Expected Shard %s to be deleted", shardName) }, }, @@ -587,6 +575,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "Normal PendingDeletion Marked Shard", }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) shardName := name.JoinWithConstraints( name.DefaultConstraints, clusterName, @@ -595,16 +584,16 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "shard-0", ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: shardName, Namespace: namespace}, shard, - ); err != nil { - t.Fatalf("Shard should exist: %v", err) - } - if shard.Annotations[multigresv1alpha1.AnnotationPendingDeletion] == "" { - t.Error("Expected PendingDeletion annotation to be set on Shard") - } + ), "Shard should exist") + ck.NotEq( + "", + shard.Annotations[multigresv1alpha1.AnnotationPendingDeletion], + "Expected PendingDeletion annotation to be set on Shard", + ) }, }, "Success: Handle Pending Deletion (Child Already Terminating)": { @@ -641,6 +630,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { }, expectedResult: ptr.To(ctrl.Result{RequeueAfter: 5 * time.Second}), validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) shardName := name.JoinWithConstraints( name.DefaultConstraints, clusterName, @@ -649,31 +639,27 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "shard-0", ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: shardName, Namespace: namespace}, shard, - ); err != nil { - t.Fatalf("terminating Shard should still be visible: %v", err) - } - if got := shard.Annotations[multigresv1alpha1.AnnotationPendingDeletion]; got != "" { - t.Errorf("terminating Shard should not be restamped, got %q", got) - } + ), "terminating Shard should still be visible") + ck.Eq( + "", + shard.Annotations[multigresv1alpha1.AnnotationPendingDeletion], + "terminating Shard should not be restamped, got", + ) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get TableGroup: %v", err) - } - if meta.IsStatusConditionTrue( + ), "failed to get TableGroup") + ck.False(meta.IsStatusConditionTrue( updatedTG.Status.Conditions, multigresv1alpha1.ConditionReadyForDeletion, - ) { - t.Error("TableGroup should wait for terminating child Shard to disappear") - } + ), "TableGroup should wait for terminating child Shard to disappear") }, }, "Success: Handle Pending Deletion (All Shards Ready For Deletion)": { @@ -721,20 +707,17 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { "Normal ReadyForDeletion TableGroup test-tg marked ready for deletion", }, validate: func(t testing.TB, c client.Client) { + ck := assert.NewCollecting(t) updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } - if !meta.IsStatusConditionTrue( + ), "failed to get tablegroup") + ck.True(meta.IsStatusConditionTrue( updatedTG.Status.Conditions, multigresv1alpha1.ConditionReadyForDeletion, - ) { - t.Error("TableGroup should be ReadyForDeletion") - } + ), "TableGroup should be ReadyForDeletion") }, }, } @@ -742,6 +725,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) if tc.preReconcileUpdate != nil { tc.preReconcileUpdate(t, tc.tableGroup) @@ -778,9 +762,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { } result, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Errorf("Unexpected error from Reconcile: %v", err) - } + c.NoError(err, "Unexpected error from Reconcile") if tc.expectedResult != nil && result != *tc.expectedResult { t.Errorf( "Unexpected result from Reconcile: got %+v, want %+v", @@ -804,13 +786,12 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { break } } - if !found { - t.Errorf( - "Expected event containing %q not found. Got events: %v", - want, - gotEvents, - ) - } + c.True( + found, + "Expected event containing %q not found. Got events: %v", + want, + gotEvents, + ) } } @@ -823,6 +804,7 @@ func TestTableGroupReconciler_Reconcile_Success(t *testing.T) { func TestTableGroupReconciler_DefaultMVPShape(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -912,13 +894,11 @@ func TestTableGroupReconciler_DefaultMVPShape(t *testing.T) { shardName, ) shard := &multigresv1alpha1.Shard{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: fullShardName, Namespace: namespace}, shard, - ); err != nil { - t.Fatalf("default 0-inf Shard was not created: %v", err) - } + ), "default 0-inf Shard was not created") if got, want := shard.Spec.DatabaseName, multigresv1alpha1.DatabaseName(dbName); got != want { t.Errorf("Shard DatabaseName mismatch: got %q, want %q", got, want) @@ -956,13 +936,11 @@ func TestTableGroupReconciler_DefaultMVPShape(t *testing.T) { } updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get TableGroup: %v", err) - } + ), "failed to get TableGroup") if got, want := updatedTG.Status.Phase, multigresv1alpha1.PhaseHealthy; got != want { t.Errorf("TableGroup phase mismatch: got %q, want %q", got, want) @@ -970,12 +948,12 @@ func TestTableGroupReconciler_DefaultMVPShape(t *testing.T) { if got, want := updatedTG.Status.ReadyShards, int32(1); got != want { t.Errorf("ReadyShards mismatch: got %d, want %d", got, want) } - if got, want := updatedTG.Status.TotalShards, int32(1); got != want { - t.Errorf("TotalShards mismatch: got %d, want %d", got, want) - } - if !meta.IsStatusConditionTrue(updatedTG.Status.Conditions, "Available") { - t.Error("default 0-inf TableGroup should be Available after its shard is healthy") - } + got, want := updatedTG.Status.TotalShards, int32(1) + ck.Eq(want, got, "TotalShards mismatch: got") + ck.True( + meta.IsStatusConditionTrue(updatedTG.Status.Conditions, "Available"), + "default 0-inf TableGroup should be Available after its shard is healthy", + ) } func TestTableGroupReconciler_Reconcile_Failure(t *testing.T) { @@ -1256,6 +1234,7 @@ func TestTableGroupReconciler_Reconcile_Failure(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) if tc.preReconcileUpdate != nil { tc.preReconcileUpdate(t, tc.tableGroup) @@ -1291,9 +1270,7 @@ func TestTableGroupReconciler_Reconcile_Failure(t *testing.T) { } _, err := reconciler.Reconcile(t.Context(), req) - if err == nil { - t.Error("Expected error from Reconcile, got nil") - } + c.Error(err, "Expected error from Reconcile, got nil") if len(tc.expectedEvents) > 0 { close(fakeRecorder.Events) @@ -1310,13 +1287,12 @@ func TestTableGroupReconciler_Reconcile_Failure(t *testing.T) { break } } - if !found { - t.Errorf( - "Expected event containing %q not found. Got events: %v", - want, - gotEvents, - ) - } + c.True( + found, + "Expected event containing %q not found. Got events: %v", + want, + gotEvents, + ) } } }) @@ -1342,9 +1318,7 @@ func TestTableGroupReconciler_Reconcile_BuildFailure(t *testing.T) { NamespacedName: types.NamespacedName{Name: baseTG.Name, Namespace: baseTG.Namespace}, } _, err := r.Reconcile(t.Context(), req) - if err == nil { - t.Fatal("Expected Reconcile to fail due to Build error") - } + assert.NewAborting(t).Error(err, "Expected Reconcile to fail due to Build error") if err.Error() != "failed to build shard: no kind is registered for the type v1alpha1.TableGroup" { t.Logf("Got error: %v", err) } @@ -1444,6 +1418,7 @@ func TestTableGroupReconciler_ReadsChildShardsOnce(t *testing.T) { for tn, tc := range tests { t.Run(tn, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) if tc.preReconcileUpdate != nil { tc.preReconcileUpdate(t, tc.tableGroup) @@ -1490,16 +1465,14 @@ func TestTableGroupReconciler_ReadsChildShardsOnce(t *testing.T) { }, } - if _, err := reconciler.Reconcile(t.Context(), req); err != nil { - t.Fatalf("unexpected error from Reconcile: %v", err) - } + _, err := reconciler.Reconcile(t.Context(), req) + ck.Require().NoError(err, "unexpected error from Reconcile") - if got := shardListCount.Load(); got != 1 { - t.Errorf( - "expected exactly one List of child Shards per reconcile, got %d", - got, - ) - } + ck.Eq( + 1, + shardListCount.Load(), + "expected exactly one List of child Shards per reconcile, got", + ) }) } } @@ -1582,6 +1555,7 @@ func TestTableGroupReconciler_PendingDeletionUpgradeCompatibility(t *testing.T) for tn, tc := range tests { t.Run(tn, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) // No desired Shards, so the seeded child is an orphan. tg := baseTG.DeepCopy() @@ -1641,22 +1615,19 @@ func TestTableGroupReconciler_PendingDeletionUpgradeCompatibility(t *testing.T) } got, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Fatalf("unexpected error from Reconcile: %v", err) - } + ck.Require().NoError(err, "unexpected error from Reconcile") - if got != tc.wantResult { - t.Errorf("unexpected reconcile result: got %+v, want %+v", got, tc.wantResult) - } + ck.Eq(tc.wantResult, got, "unexpected reconcile result: got") // The annotation must never be re-stamped. - if patches := shardPatchCount.Load(); patches != 0 { - t.Errorf( - "handshake double-driven: expected zero Patch calls on the child Shard "+ - "(annotation already set by a prior version), got %d", - patches, - ) - } + patches := shardPatchCount.Load() + ck.Eq( + 0, + patches, + "handshake double-driven: expected zero Patch calls on the child Shard "+ + "(annotation already set by a prior version), got %d", + patches, + ) fetched := &multigresv1alpha1.Shard{} getErr := c.Get( @@ -1674,27 +1645,26 @@ func TestTableGroupReconciler_PendingDeletionUpgradeCompatibility(t *testing.T) getErr, ) } - if deletes := shardDeleteCount.Load(); deletes != 1 { - t.Errorf( - "expected exactly one Delete of the child Shard, got %d", - deletes, - ) - } + ck.Eq( + 1, + shardDeleteCount.Load(), + "expected exactly one Delete of the child Shard, got", + ) return } // While it's still draining the child should stay put and keep its // original annotation. - if getErr != nil { - t.Fatalf("expected orphan child Shard to still exist while draining: %v", getErr) - } - if deletes := shardDeleteCount.Load(); deletes != 0 { - t.Errorf( - "handshake double-driven: expected no Delete of the child Shard while it "+ - "is still draining, got %d", - deletes, - ) - } + ck.Require(). + NoError(getErr, "expected orphan child Shard to still exist while draining") + deletes := shardDeleteCount.Load() + ck.Eq( + 0, + deletes, + "handshake double-driven: expected no Delete of the child Shard while it "+ + "is still draining, got %d", + deletes, + ) if got := fetched.Annotations[multigresv1alpha1.AnnotationPendingDeletion]; got != priorVersionTimestamp { t.Errorf( @@ -1713,6 +1683,7 @@ func TestTableGroupReconciler_PendingDeletionUpgradeCompatibility(t *testing.T) // the pending cleanup requeue. func TestTableGroupReconciler_PendingDeletionPublishesProgressingStatus(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1784,21 +1755,16 @@ func TestTableGroupReconciler_PendingDeletionPublishesProgressingStatus(t *testi got, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: tgName, Namespace: namespace}, }) - if err != nil { - t.Fatalf("unexpected error from Reconcile: %v", err) - } - if want := (ctrl.Result{RequeueAfter: 5 * time.Second}); got != want { - t.Fatalf("unexpected reconcile result: got %+v, want %+v", got, want) - } + ck.Require().NoError(err, "unexpected error from Reconcile") + ck.Require(). + Eq((ctrl.Result{RequeueAfter: 5 * time.Second}), got, "unexpected reconcile result: got") updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } + ), "failed to get tablegroup") if got, want := updatedTG.Status.Phase, multigresv1alpha1.PhaseProgressing; got != want { t.Errorf("Phase mismatch while cleanup is pending: got %q, want %q", got, want) @@ -1817,31 +1783,15 @@ func TestTableGroupReconciler_PendingDeletionPublishesProgressingStatus(t *testi } cond := meta.FindStatusCondition(updatedTG.Status.Conditions, "Available") - if cond == nil { - t.Fatal("expected an Available condition to be set") - } - if cond.Status != metav1.ConditionFalse { - t.Errorf("Available status mismatch: got %q, want %q", cond.Status, metav1.ConditionFalse) - } - if cond.Reason != "CleanupPending" { - t.Errorf("Available reason mismatch: got %q, want CleanupPending", cond.Reason) - } - if cond.Message != "Waiting for shard cleanup to finish" { - t.Errorf("Available message mismatch: got %q", cond.Message) - } - if cond.ObservedGeneration != tg.Generation { - t.Errorf( - "Available observedGeneration mismatch: got %d, want %d", - cond.ObservedGeneration, - tg.Generation, - ) - } + ck.Require().NotNil(cond, "expected an Available condition to be set") + ck.Eq(metav1.ConditionFalse, cond.Status, "Available status mismatch: got") + ck.Eq("CleanupPending", cond.Reason, "Available reason mismatch: got") + ck.Eq("Waiting for shard cleanup to finish", cond.Message, "Available message mismatch: got") + ck.Eq(tg.Generation, cond.ObservedGeneration, "Available observedGeneration mismatch: got") close(recorder.Events) for evt := range recorder.Events { - if strings.Contains(evt, "Synced") { - t.Errorf("unexpected Synced event while cleanup is still pending: %q", evt) - } + ck.NotStrContains(evt, "Synced", "unexpected Synced event while cleanup is still pending") } } @@ -1851,6 +1801,7 @@ func TestTableGroupReconciler_PendingDeletionPublishesProgressingStatus(t *testi // full-reconcile assertion, since the PhaseChange gating spans the whole chain. func TestTableGroupReconciler_SteadyStateDoesNotFlap(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1919,25 +1870,19 @@ func TestTableGroupReconciler_SteadyStateDoesNotFlap(t *testing.T) { } got, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Fatalf("unexpected error from Reconcile: %v", err) - } + ck.Require().NoError(err, "unexpected error from Reconcile") // Nothing changed, so we shouldn't get a requeue. - if got != (ctrl.Result{}) { - t.Errorf("unexpected reconcile result: got %+v, want empty ctrl.Result{}", got) - } + ck.Eq((ctrl.Result{}), got, "unexpected reconcile result: got") // The status stays Healthy because it comes from the child's observed // .Status, not the empty-status desired object. updatedTG := &multigresv1alpha1.TableGroup{} - if err := c.Get( + ck.Require().NoError(c.Get( t.Context(), types.NamespacedName{Name: tgName, Namespace: namespace}, updatedTG, - ); err != nil { - t.Fatalf("failed to get tablegroup: %v", err) - } + ), "failed to get tablegroup") if got, want := updatedTG.Status.Phase, multigresv1alpha1.PhaseHealthy; got != want { t.Errorf( @@ -1955,8 +1900,10 @@ func TestTableGroupReconciler_SteadyStateDoesNotFlap(t *testing.T) { // Phase didn't change, so there should be no PhaseChange event. close(recorder.Events) for evt := range recorder.Events { - if strings.Contains(evt, "PhaseChange") { - t.Errorf("unexpected spurious PhaseChange event emitted in steady state: %q", evt) - } + ck.NotStrContains( + evt, + "PhaseChange", + "unexpected spurious PhaseChange event emitted in steady state", + ) } } diff --git a/pkg/data-handler/backuphealth/backuphealth_test.go b/pkg/data-handler/backuphealth/backuphealth_test.go index 82547ba1..df295bba 100644 --- a/pkg/data-handler/backuphealth/backuphealth_test.go +++ b/pkg/data-handler/backuphealth/backuphealth_test.go @@ -14,10 +14,13 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/backuphealth" + + "github.com/multigres/testkit/assert" ) func TestEvaluateBackups_Healthy(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -40,16 +43,13 @@ func TestEvaluateBackups_Healthy(t *testing.T) { if !result.Healthy { t.Errorf("expected healthy, got message: %s", result.Message) } - if result.LastBackupType != "full" { - t.Errorf("expected type=full, got %s", result.LastBackupType) - } - if result.LastBackupTime == nil { - t.Error("expected LastBackupTime to be set") - } + c.Eq("full", result.LastBackupType, "expected type=full, got") + c.NotNil(result.LastBackupTime, "expected LastBackupTime to be set") } func TestEvaluateBackups_Stale(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -69,29 +69,23 @@ func TestEvaluateBackups_Stale(t *testing.T) { } result := backuphealth.EvaluateBackups(shard, backups) - if result.Healthy { - t.Error("expected unhealthy for 48h-old backup") - } - if result.LastBackupType != "diff" { - t.Errorf("expected type=diff, got %s", result.LastBackupType) - } + c.False(result.Healthy, "expected unhealthy for 48h-old backup") + c.Eq("diff", result.LastBackupType, "expected type=diff, got") } func TestEvaluateBackups_NoBackups(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} result := backuphealth.EvaluateBackups(shard, nil) - if result.Healthy { - t.Error("expected unhealthy when no backups") - } - if result.Message != "No backups found" { - t.Errorf("unexpected message: %s", result.Message) - } + c.False(result.Healthy, "expected unhealthy when no backups") + c.Eq("No backups found", result.Message, "unexpected message") } func TestEvaluateBackups_NoCompleted(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} backups := []*multipoolermanagerdata.BackupMetadata{ @@ -103,12 +97,8 @@ func TestEvaluateBackups_NoCompleted(t *testing.T) { } result := backuphealth.EvaluateBackups(shard, backups) - if result.Healthy { - t.Error("expected unhealthy when no completed backups") - } - if result.Message != "No completed backups found" { - t.Errorf("unexpected message: %s", result.Message) - } + c.False(result.Healthy, "expected unhealthy when no completed backups") + c.Eq("No completed backups found", result.Message, "unexpected message") } func TestEvaluateBackups_SelectsMostRecent(t *testing.T) { @@ -138,9 +128,8 @@ func TestEvaluateBackups_SelectsMostRecent(t *testing.T) { } result := backuphealth.EvaluateBackups(shard, backups) - if result.LastBackupType != "incr" { - t.Errorf("expected most recent backup type=incr, got %s", result.LastBackupType) - } + assert.NewCollecting(t). + Eq("incr", result.LastBackupType, "expected most recent backup type=incr, got") } func TestApply(t *testing.T) { @@ -148,6 +137,7 @@ func TestApply(t *testing.T) { t.Run("sets healthy condition", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Generation: 5}, @@ -162,30 +152,18 @@ func TestApply(t *testing.T) { backuphealth.Apply(shard, result) - if len(shard.Status.Conditions) != 1 { - t.Fatalf("expected 1 condition, got %d", len(shard.Status.Conditions)) - } + ck.Require(). + Len(shard.Status.Conditions, 1, "expected 1 condition, got %d", len(shard.Status.Conditions)) c := shard.Status.Conditions[0] - if c.Type != backuphealth.ConditionHealthy { - t.Errorf( - "expected condition type %s, got %s", - backuphealth.ConditionHealthy, - c.Type, - ) - } - if c.Status != metav1.ConditionTrue { - t.Errorf("expected True, got %s", c.Status) - } - if c.Reason != "BackupRecent" { - t.Errorf("expected reason BackupRecent, got %s", c.Reason) - } - if shard.Status.LastBackupType != "full" { - t.Errorf("expected LastBackupType=full, got %s", shard.Status.LastBackupType) - } + ck.Eq(backuphealth.ConditionHealthy, c.Type, "expected condition type") + ck.Eq(metav1.ConditionTrue, c.Status, "expected True, got") + ck.Eq("BackupRecent", c.Reason, "expected reason BackupRecent, got") + ck.Eq("full", shard.Status.LastBackupType, "expected LastBackupType=full, got") }) t.Run("sets unhealthy condition", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} result := &backuphealth.Result{ @@ -195,16 +173,11 @@ func TestApply(t *testing.T) { backuphealth.Apply(shard, result) - if len(shard.Status.Conditions) != 1 { - t.Fatalf("expected 1 condition, got %d", len(shard.Status.Conditions)) - } + ck.Require(). + Len(shard.Status.Conditions, 1, "expected 1 condition, got %d", len(shard.Status.Conditions)) c := shard.Status.Conditions[0] - if c.Status != metav1.ConditionFalse { - t.Errorf("expected False, got %s", c.Status) - } - if c.Reason != "BackupStale" { - t.Errorf("expected reason BackupStale, got %s", c.Reason) - } + ck.Eq(metav1.ConditionFalse, c.Status, "expected False, got") + ck.Eq("BackupStale", c.Reason, "expected reason BackupStale, got") }) t.Run("nil result is no-op", func(t *testing.T) { @@ -212,9 +185,8 @@ func TestApply(t *testing.T) { shard := &multigresv1alpha1.Shard{} backuphealth.Apply(shard, nil) - if len(shard.Status.Conditions) != 0 { - t.Error("expected no conditions for nil result") - } + assert.NewCollecting(t). + Empty(shard.Status.Conditions, "expected no conditions for nil result") }) } @@ -241,9 +213,8 @@ func TestParseTime(t *testing.T) { t.Run(name, func(t *testing.T) { t.Parallel() got := backuphealth.ParseTime(tc.input) - if !got.Equal(tc.want) { - t.Errorf("ParseTime(%q) = %v, want %v", tc.input, got, tc.want) - } + assert.NewCollecting(t). + True(got.Equal(tc.want), "ParseTime(%q) = %v, want %v", tc.input, got, tc.want) }) } } @@ -251,13 +222,12 @@ func TestParseTime(t *testing.T) { func TestParseTime_InvalidFormat(t *testing.T) { t.Parallel() got := backuphealth.ParseTime("ABCDEFG-HIJKLMN") - if !got.IsZero() { - t.Errorf("expected zero time for invalid format, got %v", got) - } + assert.NewCollecting(t).True(got.IsZero(), "expected zero time for invalid format, got %v", got) } func TestEvaluateBackups_MalformedBackupID(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -276,12 +246,8 @@ func TestEvaluateBackups_MalformedBackupID(t *testing.T) { } result := backuphealth.EvaluateBackups(shard, backups) - if result.Healthy { - t.Error("expected unhealthy for malformed backup ID") - } - if result.LastBackupTime != nil { - t.Errorf("expected nil LastBackupTime, got %v", result.LastBackupTime) - } + c.False(result.Healthy, "expected unhealthy for malformed backup ID") + c.Nil(result.LastBackupTime, "expected nil LastBackupTime, got") } type mockTopoStore struct { @@ -335,6 +301,7 @@ func TestEvaluate(t *testing.T) { } t.Run("No primary found", func(t *testing.T) { + c := assert.NewCollecting(t) store := &mockTopoStore{ getMultipoolersByCellFunc: func(ctx context.Context, cellName string, opt *topoclient.GetMultipoolersByCellOptions) ([]*topoclient.MultipoolerInfo, error) { return nil, nil @@ -343,12 +310,8 @@ func TestEvaluate(t *testing.T) { rpc := &mockMultipoolerClient{} res, err := backuphealth.Evaluate(ctx, store, rpc, shard) - if err != nil { - t.Errorf("unexpected error: %v", err) - } - if res != nil { - t.Errorf("expected nil result, got %v", res) - } + c.NoError(err, "unexpected error") + c.Nil(res, "expected nil result, got") }) t.Run("Primary found but GetBackups fails", func(t *testing.T) { @@ -372,12 +335,11 @@ func TestEvaluate(t *testing.T) { } _, err := backuphealth.Evaluate(ctx, store, rpc, shard) - if err == nil { - t.Errorf("expected error, got nil") - } + assert.NewCollecting(t).Error(err, "expected error, got nil") }) t.Run("Primary found and EvaluateBackups runs", func(t *testing.T) { + c := assert.NewCollecting(t) primaryInfo := &topoclient.MultipoolerInfo{ Multipooler: &clustermetadata.Multipooler{ Id: &clustermetadata.ID{Name: "primary-1"}, @@ -408,15 +370,9 @@ func TestEvaluate(t *testing.T) { } res, err := backuphealth.Evaluate(ctx, store, rpc, shard) - if err != nil { - t.Errorf("unexpected error: %v", err) - } - if res == nil { - t.Fatal("expected result, got nil") - } - if !res.Healthy { - t.Errorf("expected healthy true, got false") - } + c.NoError(err, "unexpected error") + c.Require().NotNil(res, "expected result, got nil") + c.True(res.Healthy, "expected healthy true, got false") }) t.Run("FindPrimaryPooler error", func(t *testing.T) { @@ -428,8 +384,6 @@ func TestEvaluate(t *testing.T) { rpc := &mockMultipoolerClient{} _, err := backuphealth.Evaluate(ctx, store, rpc, shard) - if err == nil { - t.Errorf("expected find pooler error") - } + assert.NewCollecting(t).Error(err, "expected find pooler error") }) } diff --git a/pkg/data-handler/drain/drain_test.go b/pkg/data-handler/drain/drain_test.go index 3f399fc0..9e75e50b 100644 --- a/pkg/data-handler/drain/drain_test.go +++ b/pkg/data-handler/drain/drain_test.go @@ -15,6 +15,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/drain" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestExecuteDrainStateMachine(t *testing.T) { @@ -41,66 +43,55 @@ func TestExecuteDrainStateMachine(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) shard, pod, k8sClient := testObjects(t, tt.from) requeue, err := drain.ExecuteDrainStateMachine( context.Background(), k8sClient, record.NewFakeRecorder(1), shard, pod, ) - if err != nil { - t.Fatalf("execute drain state machine: %v", err) - } - if !requeue { - t.Fatal("expected a requeue after a state transition") - } + c.NoError(err, "execute drain state machine") + c.True(requeue, "expected a requeue after a state transition") updated := &corev1.Pod{} - if err := k8sClient.Get( + c.NoError(k8sClient.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("get updated pod: %v", err) - } - if got := updated.Annotations[metadata.AnnotationDrainState]; got != tt.want { - t.Fatalf("drain state = %q, want %q", got, tt.want) - } + ), "get updated pod") + c.Eq(tt.want, updated.Annotations[metadata.AnnotationDrainState], "drain state") }) } } func TestExecuteDrainStateMachineTimeout(t *testing.T) { + c := assert.NewAborting(t) shard, pod, k8sClient := testObjects(t, metadata.DrainStateDraining) pod.Annotations[metadata.AnnotationDrainRequestedAt] = time.Now(). Add(-drain.DrainTimeout - time.Second). Format(time.RFC3339) - if err := k8sClient.Update(context.Background(), pod); err != nil { - t.Fatalf("update pod: %v", err) - } + c.NoError(k8sClient.Update(context.Background(), pod), "update pod") requeue, err := drain.ExecuteDrainStateMachine(context.Background(), k8sClient, nil, shard, pod) - if err != nil { - t.Fatalf("execute timed out drain: %v", err) - } - if !requeue { - t.Fatal("expected a requeue after the timeout transition") - } + c.NoError(err, "execute timed out drain") + c.True(requeue, "expected a requeue after the timeout transition") updated := &corev1.Pod{} - if err := k8sClient.Get( + c.NoError(k8sClient.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("get updated pod: %v", err) - } - if got := updated.Annotations[metadata.AnnotationDrainState]; got != metadata.DrainStateReadyForDeletion { - t.Fatalf("drain state = %q, want %q", got, metadata.DrainStateReadyForDeletion) - } + ), "get updated pod") + c.Eq( + metadata.DrainStateReadyForDeletion, + updated.Annotations[metadata.AnnotationDrainState], + "drain state", + ) } func TestExecuteDrainStateMachineNoop(t *testing.T) { for _, state := range []string{"", metadata.DrainStateReadyForDeletion} { t.Run(state, func(t *testing.T) { + c := assert.NewAborting(t) shard, pod, k8sClient := testObjects(t, state) requeue, err := drain.ExecuteDrainStateMachine( context.Background(), @@ -109,12 +100,8 @@ func TestExecuteDrainStateMachineNoop(t *testing.T) { shard, pod, ) - if err != nil { - t.Fatalf("execute drain state machine: %v", err) - } - if requeue { - t.Fatal("did not expect a requeue") - } + c.NoError(err, "execute drain state machine") + c.False(requeue, "did not expect a requeue") }) } } @@ -124,13 +111,10 @@ func testObjects( state string, ) (*multigresv1alpha1.Shard, *corev1.Pod, client.Client) { t.Helper() + c := assert.NewAborting(t) scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("add Multigres scheme: %v", err) - } - if err := corev1.AddToScheme(scheme); err != nil { - t.Fatalf("add core scheme: %v", err) - } + c.NoError(multigresv1alpha1.AddToScheme(scheme), "add Multigres scheme") + c.NoError(corev1.AddToScheme(scheme), "add core scheme") shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ diff --git a/pkg/data-handler/poolerclient/resolver_test.go b/pkg/data-handler/poolerclient/resolver_test.go index c3cd7c2e..77131c7c 100644 --- a/pkg/data-handler/poolerclient/resolver_test.go +++ b/pkg/data-handler/poolerclient/resolver_test.go @@ -26,17 +26,16 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func testScheme(t *testing.T) *runtime.Scheme { t.Helper() + c := assert.NewAborting(t) s := runtime.NewScheme() - if err := clientgoscheme.AddToScheme(s); err != nil { - t.Fatalf("add scheme: %v", err) - } - if err := multigresv1alpha1.AddToScheme(s); err != nil { - t.Fatalf("add Multigres scheme: %v", err) - } + c.NoError(clientgoscheme.AddToScheme(s), "add scheme") + c.NoError(multigresv1alpha1.AddToScheme(s), "add Multigres scheme") return s } @@ -73,10 +72,9 @@ type testCA struct { func newTestCA(t *testing.T) *testCA { t.Helper() + c := assert.NewAborting(t) key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("generate CA key: %v", err) - } + c.NoError(err, "generate CA key") template := &x509.Certificate{ SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "test-ca"}, @@ -87,13 +85,9 @@ func newTestCA(t *testing.T) *testCA { BasicConstraintsValid: true, } der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) - if err != nil { - t.Fatalf("create CA: %v", err) - } + c.NoError(err, "create CA") cert, err := x509.ParseCertificate(der) - if err != nil { - t.Fatalf("parse CA: %v", err) - } + c.NoError(err, "parse CA") return &testCA{ cert: cert, key: key, @@ -107,10 +101,9 @@ func (ca *testCA) issue( usages ...x509.ExtKeyUsage, ) (certPEM, keyPEM []byte, leaf *x509.Certificate) { t.Helper() + c := assert.NewAborting(t) key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("generate leaf key: %v", err) - } + c.NoError(err, "generate leaf key") template := &x509.Certificate{ SerialNumber: big.NewInt(time.Now().UnixNano()), Subject: pkix.Name{CommonName: "leaf"}, @@ -121,17 +114,11 @@ func (ca *testCA) issue( DNSNames: dnsNames, } der, err := x509.CreateCertificate(rand.Reader, template, ca.cert, &key.PublicKey, ca.key) - if err != nil { - t.Fatalf("create leaf: %v", err) - } + c.NoError(err, "create leaf") leaf, err = x509.ParseCertificate(der) - if err != nil { - t.Fatalf("parse leaf: %v", err) - } + c.NoError(err, "parse leaf") keyDER, err := x509.MarshalECPrivateKey(key) - if err != nil { - t.Fatalf("marshal leaf key: %v", err) - } + c.NoError(err, "marshal leaf key") return pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}), leaf } @@ -171,9 +158,7 @@ func newResolver( c := fake.NewClientBuilder().WithScheme(testScheme(t)).WithObjects(objects...).Build() insecure := rpcclient.NewFakeClient() r, err := NewOperatorCertResolver(c, Options{Capacity: 10, Insecure: insecure}) - if err != nil { - t.Fatalf("NewOperatorCertResolver() error = %v", err) - } + assert.NewAborting(t).NoError(err, "NewOperatorCertResolver() error =") t.Cleanup(r.Close) return r, insecure, c } @@ -201,41 +186,33 @@ func TestNewOperatorCertResolverValidatesDependencies(t *testing.T) { } func TestStaticResolver(t *testing.T) { + c := assert.NewCollecting(t) want := rpcclient.NewFakeClient() got, err := Static(want).ClientFor(t.Context(), testShard("s", "ns", "c", true)) - if err != nil { - t.Fatalf("ClientFor() error = %v", err) - } - if got != want { - t.Error("Static resolver returned a different client") - } + c.Require().NoError(err, "ClientFor() error =") + c.False(got != want, "Static resolver returned a different client") } func TestClientForTLSDisabledReturnsInsecure(t *testing.T) { + c := assert.NewCollecting(t) r, insecure, _ := newResolver(t) got, err := r.ClientFor(t.Context(), testShard("s", "ns", "c", false)) - if err != nil { - t.Fatalf("ClientFor() error = %v", err) - } - if got != insecure { - t.Error("TLS-disabled shard did not get the insecure client") - } + c.Require().NoError(err, "ClientFor() error =") + c.False(got != insecure, "TLS-disabled shard did not get the insecure client") } func TestClientForRequiresClusterLabel(t *testing.T) { r, _, _ := newResolver(t) _, err := r.ClientFor(t.Context(), testShard("s", "ns", "", true)) - if err == nil || !strings.Contains(err.Error(), metadata.LabelMultigresCluster) { - t.Fatalf("ClientFor() error = %v, want missing label error", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), metadata.LabelMultigresCluster), "ClientFor() error = %v, want missing label error", err) } func TestClientForSecretNotIssued(t *testing.T) { r, _, _ := newResolver(t) _, err := r.ClientFor(t.Context(), activeTLSShard(r, "s", "ns", "c")) - if err == nil || !strings.Contains(err.Error(), "reading operator internal TLS secret") { - t.Fatalf("ClientFor() error = %v, want missing Secret error", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "reading operator internal TLS secret"), "ClientFor() error = %v, want missing Secret error", err) } type countingBlockingReader struct { @@ -275,6 +252,7 @@ func (r *countingBlockingReader) getCount() int { } func TestClientForCoalescesAndCachesInitialFailure(t *testing.T) { + c := assert.NewAborting(t) r, _, reader := newResolver(t) blocking := &countingBlockingReader{ Reader: reader, @@ -300,9 +278,7 @@ func TestClientForCoalescesAndCachesInitialFailure(t *testing.T) { } // Give the other callers a chance to join the in-flight initialization. time.Sleep(20 * time.Millisecond) - if got := blocking.getCount(); got != 1 { - t.Fatalf("concurrent initial Secret reads = %d, want 1", got) - } + c.Eq(1, blocking.getCount(), "concurrent initial Secret reads") close(blocking.release) for range callers { if err := <-errs; err == nil || @@ -311,12 +287,9 @@ func TestClientForCoalescesAndCachesInitialFailure(t *testing.T) { } } - if _, err := r.ClientFor(t.Context(), shard); err == nil { - t.Fatal("ClientFor() expected cached missing Secret error") - } - if got := blocking.getCount(); got != 1 { - t.Fatalf("Secret reads during failure cache = %d, want 1", got) - } + _, err := r.ClientFor(t.Context(), shard) + c.Error(err, "ClientFor() expected cached missing Secret error") + c.Eq(1, blocking.getCount(), "Secret reads during failure cache") } type cancelFirstReader struct { @@ -352,6 +325,7 @@ func (r *cancelFirstReader) getCount() int { } func TestClientForCanceledColdLeaderDoesNotPoisonJoiner(t *testing.T) { + c := assert.NewAborting(t) ca := newTestCA(t) r, _, reader := newResolver(t, operatorSecret(t, ca, "ns", "c", "1")) cancelReader := &cancelFirstReader{ @@ -389,18 +363,20 @@ func TestClientForCanceledColdLeaderDoesNotPoisonJoiner(t *testing.T) { } select { case result := <-joinerResult: - if result.err != nil || result.client == nil { - t.Fatalf("joiner ClientFor() = (%v, %v), want a client", result.client, result.err) - } + c.False( + result.err != nil || result.client == nil, + "joiner ClientFor() = (%v, %v), want a client", + result.client, + result.err, + ) case <-time.After(time.Second): t.Fatal("healthy joiner did not retry canceled initialization") } - if got := cancelReader.getCount(); got != 2 { - t.Fatalf("Secret reads = %d, want canceled read plus healthy retry", got) - } + c.Eq(2, cancelReader.getCount(), "Secret reads") } func TestClientForCachesPerClusterAndBindsServerName(t *testing.T) { + c := assert.NewCollecting(t) ca := newTestCA(t) r, insecure, _ := newResolver(t, operatorSecret(t, ca, "ns-a", "cluster-a", "1"), @@ -415,23 +391,16 @@ func TestClientForCachesPerClusterAndBindsServerName(t *testing.T) { r.ActivateCluster(types.NamespacedName{Namespace: "ns-b", Name: "cluster-b"}) a1, err := r.ClientFor(t.Context(), testShard("s1", "ns-a", "cluster-a", true)) - if err != nil { - t.Fatalf("first cluster A ClientFor() error = %v", err) - } + c.Require().NoError(err, "first cluster A ClientFor() error =") a2, err := r.ClientFor(t.Context(), testShard("s2", "ns-a", "cluster-a", true)) - if err != nil { - t.Fatalf("second cluster A ClientFor() error = %v", err) - } + c.Require().NoError(err, "second cluster A ClientFor() error =") b, err := r.ClientFor(t.Context(), testShard("s", "ns-b", "cluster-b", true)) - if err != nil { - t.Fatalf("cluster B ClientFor() error = %v", err) - } - if a1 == insecure || b == insecure || a1 != a2 || a1 == b { - t.Error("clients were not cached independently per cluster") - } - if len(configs) != 2 { - t.Fatalf("created %d TLS configs, want 2", len(configs)) - } + c.Require().NoError(err, "cluster B ClientFor() error =") + c.False( + a1 == insecure || b == insecure || a1 != a2 || a1 == b, + "clients were not cached independently per cluster", + ) + c.Require().Len(configs, 2, "created %d TLS configs, want 2", len(configs)) wantA := "multipooler.cluster-a.ns-a.multigres.internal" wantB := "multipooler.cluster-b.ns-b.multigres.internal" if configs[0].ServerName != wantA || configs[1].ServerName != wantB { @@ -443,12 +412,14 @@ func TestClientForCachesPerClusterAndBindsServerName(t *testing.T) { wantB, ) } - if configs[0].InsecureSkipVerify || configs[0].VerifyConnection != nil { - t.Error("cluster client bypasses normal TLS hostname verification") - } + c.False( + configs[0].InsecureSkipVerify || configs[0].VerifyConnection != nil, + "cluster client bypasses normal TLS hostname verification", + ) } func TestClientForKeepsClientOnRefreshFailure(t *testing.T) { + ck := assert.NewCollecting(t) ca := newTestCA(t) secret := operatorSecret(t, ca, "ns", "c", "1") r, _, c := newResolver(t, secret) @@ -456,19 +427,11 @@ func TestClientForKeepsClientOnRefreshFailure(t *testing.T) { shard := activeTLSShard(r, "s", "ns", "c") want, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } - if err := c.Delete(t.Context(), secret); err != nil { - t.Fatalf("delete secret: %v", err) - } + ck.Require().NoError(err, "first ClientFor() error =") + ck.Require().NoError(c.Delete(t.Context(), secret), "delete secret") got, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("refresh ClientFor() error = %v", err) - } - if got != want { - t.Error("refresh failure replaced the working client") - } + ck.Require().NoError(err, "refresh ClientFor() error =") + ck.False(got != want, "refresh failure replaced the working client") } type countingReader struct { @@ -496,6 +459,7 @@ func (r *countingReader) getCount() int { } func TestClientForRetriesWarmRefreshFailureOnFailureInterval(t *testing.T) { + ck := assert.NewAborting(t) ca := newTestCA(t) secret := operatorSecret(t, ca, "ns", "c", "1") r, _, c := newResolver(t, secret) @@ -506,54 +470,41 @@ func TestClientForRetriesWarmRefreshFailureOnFailureInterval(t *testing.T) { shard := activeTLSShard(r, "s", "ns", "c") want, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } + ck.NoError(err, "first ClientFor() error =") state := r.states[types.NamespacedName{Namespace: "ns", Name: "c"}] state.mu.Lock() state.fetchedAt = time.Now().Add(-2 * r.refreshInterval) state.mu.Unlock() - if err := c.Delete(t.Context(), secret); err != nil { - t.Fatalf("delete Secret: %v", err) - } + ck.NoError(c.Delete(t.Context(), secret), "delete Secret") if got, err := r.ClientFor(t.Context(), shard); err != nil || got != want { t.Fatalf("failed refresh ClientFor() = (%v, %v), want existing client", got, err) } - if got := reader.getCount(); got != 2 { - t.Fatalf("Secret reads after failed refresh = %d, want 2", got) - } + ck.Eq(2, reader.getCount(), "Secret reads after failed refresh") if _, err := r.ClientFor(t.Context(), shard); err != nil { t.Fatalf("ClientFor() during failure retry window error = %v", err) } - if got := reader.getCount(); got != 2 { - t.Fatalf("Secret reads during failure retry window = %d, want 2", got) - } + ck.Eq(2, reader.getCount(), "Secret reads during failure retry window") replacement := operatorSecret(t, ca, "ns", "c", "") - if err := c.Create(t.Context(), replacement); err != nil { - t.Fatalf("recreate Secret: %v", err) - } + ck.NoError(c.Create(t.Context(), replacement), "recreate Secret") state.mu.Lock() state.fetchedAt = time.Now().Add(-2 * r.failureRetry) state.mu.Unlock() if _, err := r.ClientFor(t.Context(), shard); err != nil { t.Fatalf("ClientFor() after failure retry interval error = %v", err) } - if got := reader.getCount(); got != 3 { - t.Fatalf("Secret reads after failure retry interval = %d, want 3", got) - } + ck.Eq(3, reader.getCount(), "Secret reads after failure retry interval") } func TestClientForCanceledWarmRefreshDoesNotDelayRetry(t *testing.T) { + ck := assert.NewAborting(t) ca := newTestCA(t) secret := operatorSecret(t, ca, "ns", "c", "1") r, _, c := newResolver(t, secret) r.refreshInterval = time.Hour shard := activeTLSShard(r, "s", "ns", "c") want, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } + ck.NoError(err, "first ClientFor() error =") state := r.states[types.NamespacedName{Namespace: "ns", Name: "c"}] state.mu.Lock() state.fetchedAt = time.Now().Add(-2 * r.refreshInterval) @@ -575,25 +526,20 @@ func TestClientForCanceledWarmRefreshDoesNotDelayRetry(t *testing.T) { <-cancelReader.started cancelRefresh() result := <-refreshResult - if result.err != nil || result.client != want { - t.Fatalf( - "canceled refresh ClientFor() = (%v, %v), want existing client", - result.client, - result.err, - ) - } + ck.False( + result.err != nil || result.client != want, + "canceled refresh ClientFor() = (%v, %v), want existing client", + result.client, + result.err, + ) state.mu.Lock() lastErr := state.lastErr state.mu.Unlock() - if lastErr != nil { - t.Fatalf("canceled refresh cached error %v", lastErr) - } + ck.NoError(lastErr, "canceled refresh cached error") if got, err := r.ClientFor(t.Context(), shard); err != nil || got != want { t.Fatalf("healthy retry ClientFor() = (%v, %v), want existing client", got, err) } - if got := cancelReader.getCount(); got != 2 { - t.Fatalf("Secret reads = %d, want immediate retry after cancellation", got) - } + ck.Eq(2, cancelReader.getCount(), "Secret reads") } type blockingReader struct { @@ -618,6 +564,7 @@ func (r *blockingReader) Get( } func TestClientForDoesNotBlockOnConcurrentRefresh(t *testing.T) { + ck := assert.NewAborting(t) ca := newTestCA(t) secret := operatorSecret(t, ca, "ns", "c", "1") r, _, c := newResolver(t, secret) @@ -625,9 +572,7 @@ func TestClientForDoesNotBlockOnConcurrentRefresh(t *testing.T) { shard := activeTLSShard(r, "s", "ns", "c") want, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } + ck.NoError(err, "first ClientFor() error =") started := make(chan struct{}, 1) release := make(chan struct{}) r.reader = &blockingReader{Reader: c, started: started, release: release} @@ -663,12 +608,11 @@ func TestClientForDoesNotBlockOnConcurrentRefresh(t *testing.T) { t.Fatal("concurrent reconcile blocked behind Secret refresh") } close(release) - if err := <-refreshDone; err != nil { - t.Fatalf("background refresh error = %v", err) - } + ck.NoError(<-refreshDone, "background refresh error =") } func TestClientForKeepsClientOnMalformedRotation(t *testing.T) { + ck := assert.NewCollecting(t) ca := newTestCA(t) secret := operatorSecret(t, ca, "ns", "c", "1") r, _, c := newResolver(t, secret) @@ -676,28 +620,19 @@ func TestClientForKeepsClientOnMalformedRotation(t *testing.T) { shard := activeTLSShard(r, "s", "ns", "c") want, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } + ck.Require().NoError(err, "first ClientFor() error =") current := &corev1.Secret{} key := types.NamespacedName{Namespace: secret.Namespace, Name: secret.Name} - if err := c.Get(t.Context(), key, current); err != nil { - t.Fatalf("get secret: %v", err) - } + ck.Require().NoError(c.Get(t.Context(), key, current), "get secret") current.Data[corev1.TLSCertKey] = []byte("not a cert") - if err := c.Update(t.Context(), current); err != nil { - t.Fatalf("update secret: %v", err) - } + ck.Require().NoError(c.Update(t.Context(), current), "update secret") got, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("refresh ClientFor() error = %v", err) - } - if got != want { - t.Error("malformed rotation replaced the working client") - } + ck.Require().NoError(err, "refresh ClientFor() error =") + ck.False(got != want, "malformed rotation replaced the working client") } func TestClientForRebuildsOnSecretRotation(t *testing.T) { + ck := assert.NewCollecting(t) ca := newTestCA(t) secret := operatorSecret(t, ca, "ns", "c", "1") r, _, c := newResolver(t, secret) @@ -712,37 +647,26 @@ func TestClientForRebuildsOnSecretRotation(t *testing.T) { } shard := activeTLSShard(r, "s", "ns", "c") first, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } + ck.Require().NoError(err, "first ClientFor() error =") current := &corev1.Secret{} key := types.NamespacedName{Namespace: secret.Namespace, Name: secret.Name} - if err := c.Get(t.Context(), key, current); err != nil { - t.Fatalf("get secret: %v", err) - } + ck.Require().NoError(c.Get(t.Context(), key, current), "get secret") certPEM, keyPEM, _ := ca.issue(t, nil, x509.ExtKeyUsageClientAuth) current.Data[corev1.TLSCertKey] = certPEM current.Data[corev1.TLSPrivateKeyKey] = keyPEM - if err := c.Update(t.Context(), current); err != nil { - t.Fatalf("rotate secret: %v", err) - } + ck.Require().NoError(c.Update(t.Context(), current), "rotate secret") after, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("post-rotation ClientFor() error = %v", err) - } - if after == first { - t.Error("rotated secret did not produce a new client") - } + ck.Require().NoError(err, "post-rotation ClientFor() error =") + ck.False(after == first, "rotated secret did not produce a new client") state := r.states[types.NamespacedName{Namespace: "ns", Name: "c"}] - if len(state.retiredClients) != 1 || state.retiredClients[0].client != first { - t.Error("old client was not retained for in-flight RPCs") - } + ck.False( + len(state.retiredClients) != 1 || state.retiredClients[0].client != first, + "old client was not retained for in-flight RPCs", + ) r.Close() - if closeCount != 2 { - t.Errorf("client close count after resolver shutdown = %d, want 2", closeCount) - } + ck.Eq(2, closeCount, "client close count after resolver shutdown") } type closeTrackingClient struct { @@ -765,6 +689,7 @@ func (c *closeSignalClient) Close() { } func TestForgetClusterBlocksStaleShardUntilReactivated(t *testing.T) { + c := assert.NewAborting(t) ca := newTestCA(t) r, _, _ := newResolver(t, operatorSecret(t, ca, "ns", "c", "1")) r.retirementGrace = 10 * time.Millisecond @@ -780,9 +705,7 @@ func TestForgetClusterBlocksStaleShardUntilReactivated(t *testing.T) { key := types.NamespacedName{Namespace: "ns", Name: "c"} shard := activeTLSShard(r, "s", key.Namespace, key.Name) first, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("first ClientFor() error = %v", err) - } + c.NoError(err, "first ClientFor() error =") r.ForgetCluster(key) if _, err := r.ClientFor(t.Context(), shard); err == nil || @@ -793,13 +716,12 @@ func TestForgetClusterBlocksStaleShardUntilReactivated(t *testing.T) { _, stateRecreated := r.states[key] _, stillActive := r.active[key] r.mu.Unlock() - if stateRecreated || stillActive { - t.Fatalf( - "forgotten cluster state/active = %v/%v, want false/false", - stateRecreated, - stillActive, - ) - } + c.False( + stateRecreated || stillActive, + "forgotten cluster state/active = %v/%v, want false/false", + stateRecreated, + stillActive, + ) select { case <-clients[0].closed: case <-time.After(time.Second): @@ -808,15 +730,12 @@ func TestForgetClusterBlocksStaleShardUntilReactivated(t *testing.T) { r.ActivateCluster(key) after, err := r.ClientFor(t.Context(), shard) - if err != nil { - t.Fatalf("reactivated ClientFor() error = %v", err) - } - if after == first { - t.Fatal("reactivated cluster reused the forgotten client") - } + c.NoError(err, "reactivated ClientFor() error =") + c.False(after == first, "reactivated cluster reused the forgotten client") } func TestClientForValidatesLiveClusterBeforeClusterControllerReconciles(t *testing.T) { + c := assert.NewAborting(t) ca := newTestCA(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "c", Namespace: "ns"}, @@ -824,27 +743,23 @@ func TestClientForValidatesLiveClusterBeforeClusterControllerReconciles(t *testi r, insecure, _ := newResolver(t, cluster, operatorSecret(t, ca, "ns", "c", "1")) got, err := r.ClientFor(t.Context(), testShard("s", "ns", "c", true)) - if err != nil { - t.Fatalf("ClientFor() before cluster reconcile error = %v", err) - } - if got == nil || got == insecure { - t.Fatal("ClientFor() before cluster reconcile did not build an mTLS client") - } + c.NoError(err, "ClientFor() before cluster reconcile error =") + c.False( + got == nil || got == insecure, + "ClientFor() before cluster reconcile did not build an mTLS client", + ) r.mu.Lock() _, active := r.active[types.NamespacedName{Namespace: "ns", Name: "c"}] r.mu.Unlock() - if !active { - t.Fatal("live cluster was not added to lifecycle registry") - } + c.True(active, "live cluster was not added to lifecycle registry") } func TestClientForMalformedInitialSecret(t *testing.T) { secret := operatorSecret(t, newTestCA(t), "ns", "c", "1") delete(secret.Data, corev1.TLSPrivateKeyKey) r, _, _ := newResolver(t, secret) - if _, err := r.ClientFor(t.Context(), activeTLSShard(r, "s", "ns", "c")); err == nil { - t.Fatal("ClientFor() expected malformed Secret error") - } + _, err := r.ClientFor(t.Context(), activeTLSShard(r, "s", "ns", "c")) + assert.NewAborting(t).Error(err, "ClientFor() expected malformed Secret error") } func TestBuildTLSConfigVerifiesExactClusterIdentity(t *testing.T) { @@ -852,9 +767,7 @@ func TestBuildTLSConfigVerifiesExactClusterIdentity(t *testing.T) { secret := operatorSecret(t, ca, "ns-a", "cluster-a", "1") serverName := "multipooler.cluster-a.ns-a.multigres.internal" config, err := buildTLSConfig(secret, serverName) - if err != nil { - t.Fatalf("buildTLSConfig() error = %v", err) - } + assert.NewAborting(t).NoError(err, "buildTLSConfig() error =") if config.ServerName != serverName || config.InsecureSkipVerify { t.Fatalf( "TLS config ServerName/InsecureSkipVerify = %q/%v", @@ -917,9 +830,8 @@ func TestBuildTLSConfigVerifiesExactClusterIdentity(t *testing.T) { Roots: config.RootCAs, DNSName: config.ServerName, }) - if (err == nil) != tt.wantOK { - t.Errorf("Verify() error = %v, want success %v", err, tt.wantOK) - } + assert.NewCollecting(t). + Eq(tt.wantOK, (err == nil), "Verify() error = %v, want success", err) }) } } diff --git a/pkg/data-handler/posture/disruption_test.go b/pkg/data-handler/posture/disruption_test.go index 1de4ed8b..ae351201 100644 --- a/pkg/data-handler/posture/disruption_test.go +++ b/pkg/data-handler/posture/disruption_test.go @@ -13,6 +13,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/posture" + + "github.com/multigres/testkit/assert" ) func TestCheckDisruption(t *testing.T) { @@ -235,9 +237,7 @@ func TestCheckDisruption(t *testing.T) { Name: target, Unscheduled: tc.unscheduled, }, ) - if (err != nil) != tc.wantError { - t.Fatalf("CheckDisruption() = %v, wantError %v", err, tc.wantError) - } + assert.NewAborting(t).ErrorWhen(tc.wantError, err, "CheckDisruption()") }) } } diff --git a/pkg/data-handler/posture/posture_test.go b/pkg/data-handler/posture/posture_test.go index 59423e45..9800fe52 100644 --- a/pkg/data-handler/posture/posture_test.go +++ b/pkg/data-handler/posture/posture_test.go @@ -13,6 +13,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/posture" + + "github.com/multigres/testkit/assert" ) type mockTopoStore struct { @@ -102,6 +104,7 @@ func TestEvaluate(t *testing.T) { t.Run("consistent postures", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() primary := poolerInfo( @@ -127,29 +130,17 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpc, shard, []string{"primary-pod", "replica-pod"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result == nil { - t.Fatal("expected result, got nil") - } - if result.MultiplePrimaries { - t.Error("expected MultiplePrimaries=false") - } - if len(result.Mismatches) != 0 { - t.Errorf("expected no mismatches, got %v", result.Mismatches) - } - if result.PrimaryCount != 1 { - t.Errorf("expected PrimaryCount=1, got %d", result.PrimaryCount) - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(result, "expected result, got nil") + c.False(result.MultiplePrimaries, "expected MultiplePrimaries=false") + c.Empty(result.Mismatches, "expected no mismatches, got") + c.Eq(1, result.PrimaryCount, "expected PrimaryCount=1, got") wantPrimary := result.Postures["primary-pod"] != "PRIMARY" wantReplica := result.Postures["replica-pod"] != "STANDBY" if wantPrimary || wantReplica { t.Errorf("unexpected postures: %v", result.Postures) } - if result.Message != "postures consistent with topology roles" { - t.Errorf("unexpected message: %s", result.Message) - } + c.Eq("postures consistent with topology roles", result.Message, "unexpected message") for _, podName := range []string{"primary-pod", "replica-pod"} { if got := result.Readiness[podName]; !got.Ready || got.Reason != "DataPlaneReady" { t.Errorf("readiness[%s] = %#v, want data-plane ready", podName, got) @@ -178,9 +169,7 @@ func TestEvaluate(t *testing.T) { rpc.SetStatusResponse(topoclient.ComponentIDString(replica.Id), response) result, err := posture.Evaluate(t.Context(), store, rpc, shard, []string{"replica-pod"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") if got := result.Readiness["replica-pod"]; got.Ready || got.Reason != "PostgresNotReady" { t.Errorf("readiness = %#v, want PostgresNotReady", got) } @@ -207,9 +196,7 @@ func TestEvaluate(t *testing.T) { rpc.SetStatusResponse(topoclient.ComponentIDString(replica.Id), response) result, err := posture.Evaluate(t.Context(), store, rpc, shard, []string{"replica-pod"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") if got := result.Readiness["replica-pod"]; got.Ready || got.Reason != "NotCohortMember" { t.Errorf("readiness = %#v, want NotCohortMember", got) } @@ -217,6 +204,7 @@ func TestEvaluate(t *testing.T) { t.Run("multiple primaries detected", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() primaryA := poolerInfo( @@ -240,18 +228,10 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpc, shard, []string{"pod-a", "pod-b"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result == nil { - t.Fatal("expected result, got nil") - } - if !result.MultiplePrimaries { - t.Error("expected MultiplePrimaries=true") - } - if result.PrimaryCount != 2 { - t.Errorf("expected PrimaryCount=2, got %d", result.PrimaryCount) - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(result, "expected result, got nil") + c.True(result.MultiplePrimaries, "expected MultiplePrimaries=true") + c.Eq(2, result.PrimaryCount, "expected PrimaryCount=2, got") if len(result.Mismatches) != 1 || result.Mismatches[0] != "pod-b" { t.Errorf("expected mismatch [pod-b], got %v", result.Mismatches) } @@ -259,6 +239,7 @@ func TestEvaluate(t *testing.T) { t.Run("replica reporting primary posture is a mismatch", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() replica := poolerInfo( @@ -277,26 +258,22 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpc, shard, []string{"replica-pod"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result == nil { - t.Fatal("expected result, got nil") - } - if result.MultiplePrimaries { - t.Error("expected MultiplePrimaries=false with only one observed primary") - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(result, "expected result, got nil") + c.False( + result.MultiplePrimaries, + "expected MultiplePrimaries=false with only one observed primary", + ) if len(result.Mismatches) != 1 || result.Mismatches[0] != "replica-pod" { t.Errorf("expected mismatch [replica-pod], got %v", result.Mismatches) } wantMsg := "pod replica-pod reports postgres primary but topology role is REPLICA" - if result.Message != wantMsg { - t.Errorf("unexpected message: got %q, want %q", result.Message, wantMsg) - } + c.Eq(wantMsg, result.Message, "unexpected message: got") }) t.Run("promoting is not a mismatch", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() replica := poolerInfo( @@ -315,25 +292,15 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpc, shard, []string{"replica-pod"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result == nil { - t.Fatal("expected result, got nil") - } - if len(result.Mismatches) != 0 { - t.Errorf( - "expected no mismatches during promotion transition, got %v", - result.Mismatches, - ) - } - if result.Postures["replica-pod"] != "PROMOTING" { - t.Errorf("expected posture PROMOTING, got %s", result.Postures["replica-pod"]) - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(result, "expected result, got nil") + c.Empty(result.Mismatches, "expected no mismatches during promotion transition, got") + c.Eq("PROMOTING", result.Postures["replica-pod"], "expected posture PROMOTING, got") }) t.Run("RPC error records UNKNOWN without false positive", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() replica := poolerInfo( @@ -352,31 +319,21 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpc, shard, []string{"replica-pod"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result == nil { - t.Fatal("expected result, got nil") - } - if result.Postures["replica-pod"] != "UNKNOWN" { - t.Errorf( - "expected UNKNOWN posture on RPC error, got %s", - result.Postures["replica-pod"], - ) - } - if len(result.Mismatches) != 0 { - t.Errorf("expected no mismatches, got %v", result.Mismatches) - } - if result.MultiplePrimaries { - t.Error("expected MultiplePrimaries=false") - } - if !result.Incomplete { - t.Error("expected RPC failure to mark observation incomplete") - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(result, "expected result, got nil") + c.Eq( + "UNKNOWN", + result.Postures["replica-pod"], + "expected UNKNOWN posture on RPC error, got", + ) + c.Empty(result.Mismatches, "expected no mismatches, got") + c.False(result.MultiplePrimaries, "expected MultiplePrimaries=false") + c.True(result.Incomplete, "expected RPC failure to mark observation incomplete") }) t.Run("unavailable topology cell returns incomplete observation", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() store := &mockTopoStore{ getMultipoolersByCellFunc: func(ctx context.Context, cellName string, opt *topoclient.GetMultipoolersByCellOptions) ([]*topoclient.MultipoolerInfo, error) { @@ -387,19 +344,15 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpcclient.NewFakeClient(), shard, nil, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result == nil || !result.Incomplete { - t.Fatalf("expected incomplete result, got %#v", result) - } - if result.Message != "posture observation incomplete" { - t.Errorf("unexpected message: %q", result.Message) - } + c.Require().NoError(err, "unexpected error") + c.Require(). + False(result == nil || !result.Incomplete, "expected incomplete result, got %#v", result) + c.Eq("posture observation incomplete", result.Message, "unexpected message") }) t.Run("shutdown pooler is skipped", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() dead := poolerInfo( @@ -418,16 +371,13 @@ func TestEvaluate(t *testing.T) { result, err := posture.Evaluate( context.Background(), store, rpc, shard, []string{"dead-pod"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result != nil { - t.Errorf("expected nil result when only pooler is shut down, got %v", result) - } + c.Require().NoError(err, "unexpected error") + c.Nil(result, "expected nil result when only pooler is shut down, got") }) t.Run("no poolers matched returns nil result", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := testShard() store := &mockTopoStore{ getMultipoolersByCellFunc: func(ctx context.Context, cellName string, opt *topoclient.GetMultipoolersByCellOptions) ([]*topoclient.MultipoolerInfo, error) { @@ -437,12 +387,8 @@ func TestEvaluate(t *testing.T) { rpc := rpcclient.NewFakeClient() result, err := posture.Evaluate(context.Background(), store, rpc, shard, nil) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result != nil { - t.Errorf("expected nil result, got %v", result) - } + c.Require().NoError(err, "unexpected error") + c.Nil(result, "expected nil result, got") }) t.Run("non-unavailable topo error is returned", func(t *testing.T) { @@ -456,9 +402,7 @@ func TestEvaluate(t *testing.T) { rpc := rpcclient.NewFakeClient() _, err := posture.Evaluate(context.Background(), store, rpc, shard, nil) - if err == nil { - t.Error("expected error, got nil") - } + assert.NewCollecting(t).Error(err, "expected error, got nil") }) } @@ -467,6 +411,7 @@ func TestApply(t *testing.T) { t.Run("sets consistent condition", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ObjectMeta: metav1.ObjectMeta{Generation: 3}} result := &posture.Result{ Postures: map[string]string{"pod-a": "PRIMARY"}, @@ -475,26 +420,23 @@ func TestApply(t *testing.T) { posture.Apply(shard, result) - if len(shard.Status.Conditions) != 1 { - t.Fatalf("expected 1 condition, got %d", len(shard.Status.Conditions)) - } + ck.Require(). + Len(shard.Status.Conditions, 1, "expected 1 condition, got %d", len(shard.Status.Conditions)) c := shard.Status.Conditions[0] - if c.Type != posture.ConditionConsistent { - t.Errorf("expected type %s, got %s", posture.ConditionConsistent, c.Type) - } - if c.Status != metav1.ConditionTrue { - t.Errorf("expected True, got %s", c.Status) - } - if c.Reason != "Consistent" { - t.Errorf("expected reason Consistent, got %s", c.Reason) - } - if shard.Status.PodPostures["pod-a"] != "PRIMARY" { - t.Errorf("expected PodPostures to be set, got %v", shard.Status.PodPostures) - } + ck.Eq(posture.ConditionConsistent, c.Type, "expected type") + ck.Eq(metav1.ConditionTrue, c.Status, "expected True, got") + ck.Eq("Consistent", c.Reason, "expected reason Consistent, got") + ck.Eq( + "PRIMARY", + shard.Status.PodPostures["pod-a"], + "expected PodPostures to be set, got %v", + shard.Status.PodPostures, + ) }) t.Run("sets MultiplePrimaries condition, takes priority over mismatches", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} result := &posture.Result{ Postures: map[string]string{"pod-a": "PRIMARY", "pod-b": "PRIMARY"}, @@ -506,16 +448,13 @@ func TestApply(t *testing.T) { posture.Apply(shard, result) c := shard.Status.Conditions[0] - if c.Status != metav1.ConditionFalse { - t.Errorf("expected False, got %s", c.Status) - } - if c.Reason != "MultiplePrimaries" { - t.Errorf("expected reason MultiplePrimaries, got %s", c.Reason) - } + ck.Eq(metav1.ConditionFalse, c.Status, "expected False, got") + ck.Eq("MultiplePrimaries", c.Reason, "expected reason MultiplePrimaries, got") }) t.Run("sets RoleMismatch condition", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} result := &posture.Result{ Postures: map[string]string{"pod-a": "PRIMARY"}, @@ -526,16 +465,13 @@ func TestApply(t *testing.T) { posture.Apply(shard, result) c := shard.Status.Conditions[0] - if c.Status != metav1.ConditionFalse { - t.Errorf("expected False, got %s", c.Status) - } - if c.Reason != "RoleMismatch" { - t.Errorf("expected reason RoleMismatch, got %s", c.Reason) - } + ck.Eq(metav1.ConditionFalse, c.Status, "expected False, got") + ck.Eq("RoleMismatch", c.Reason, "expected reason RoleMismatch, got") }) t.Run("incomplete observation sets Unknown instead of consistent", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} posture.Apply(shard, &posture.Result{ Postures: map[string]string{"pod-a": "UNKNOWN"}, @@ -544,16 +480,13 @@ func TestApply(t *testing.T) { }) condition := shard.Status.Conditions[0] - if condition.Status != metav1.ConditionUnknown { - t.Errorf("expected Unknown, got %s", condition.Status) - } - if condition.Reason != "ObservationIncomplete" { - t.Errorf("expected ObservationIncomplete, got %s", condition.Reason) - } + c.Eq(metav1.ConditionUnknown, condition.Status, "expected Unknown, got") + c.Eq("ObservationIncomplete", condition.Reason, "expected ObservationIncomplete, got") }) t.Run("incomplete observation preserves confirmed failure", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{Status: multigresv1alpha1.ShardStatus{ Conditions: []metav1.Condition{{ Type: posture.ConditionConsistent, @@ -570,12 +503,17 @@ func TestApply(t *testing.T) { }) condition := shard.Status.Conditions[0] - if condition.Status != metav1.ConditionFalse || condition.Reason != "MultiplePrimaries" { - t.Errorf("expected existing failure to remain, got %#v", condition) - } - if shard.Status.PodPostures["pod-a"] != "UNKNOWN" { - t.Errorf("expected latest posture visibility, got %v", shard.Status.PodPostures) - } + c.False( + condition.Status != metav1.ConditionFalse || condition.Reason != "MultiplePrimaries", + "expected existing failure to remain, got %#v", + condition, + ) + c.Eq( + "UNKNOWN", + shard.Status.PodPostures["pod-a"], + "expected latest posture visibility, got %v", + shard.Status.PodPostures, + ) }) t.Run( @@ -591,21 +529,17 @@ func TestApply(t *testing.T) { }) condition := shard.Status.Conditions[0] - if condition.Status != metav1.ConditionFalse || condition.Reason != "RoleMismatch" { - t.Errorf("expected definite mismatch failure, got %#v", condition) - } + assert.NewCollecting(t). + False(condition.Status != metav1.ConditionFalse || condition.Reason != "RoleMismatch", "expected definite mismatch failure, got %#v", condition) }, ) t.Run("nil result is no-op", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} posture.Apply(shard, nil) - if len(shard.Status.Conditions) != 0 { - t.Error("expected no conditions for nil result") - } - if shard.Status.PodPostures != nil { - t.Error("expected PodPostures to remain nil for nil result") - } + c.Empty(shard.Status.Conditions, "expected no conditions for nil result") + c.Nil(shard.Status.PodPostures, "expected PodPostures to remain nil for nil result") }) } diff --git a/pkg/data-handler/topo/cell_test.go b/pkg/data-handler/topo/cell_test.go index 18878789..b1f4392d 100644 --- a/pkg/data-handler/topo/cell_test.go +++ b/pkg/data-handler/topo/cell_test.go @@ -3,7 +3,6 @@ package topo_test import ( "context" "fmt" - "reflect" "testing" "github.com/multigres/multigres/go/common/topoclient" @@ -15,6 +14,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/topo" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func newTestCell(name string) *multigresv1alpha1.Cell { @@ -84,6 +85,7 @@ func TestRegisterCell(t *testing.T) { t.Run("creates new cell in topology", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -94,30 +96,23 @@ func TestRegisterCell(t *testing.T) { // Register a different cell name to ensure it's not already in topo cell := newTestCell("cell2") - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(t.Context(), store, recorder, cell, false), "unexpected error") got, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found in topo after registration: %v", err) - } - if got.Name != "cell2" { - t.Errorf("expected cell name cell2, got %s", got.Name) - } - if !reflect.DeepEqual( - got.ServerAddresses, + c.Require().NoError(err, "cell not found in topo after registration") + c.Eq("cell2", got.Name, "expected cell name cell2, got") + c.EqDiff( []string{"http://local-etcd-1:2379", "http://local-etcd-2:2379"}, - ) { - t.Errorf("expected local topo addresses, got %v", got.ServerAddresses) - } - if got.Root != "/multigres/cells/cell2" { - t.Errorf("expected local topo root, got %s", got.Root) - } + got.ServerAddresses, + "expected local topo addresses, got", + ) + c.Eq("/multigres/cells/cell2", got.Root, "expected local topo root, got") }) t.Run("copies metadata verbatim into the topo record", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -128,21 +123,21 @@ func TestRegisterCell(t *testing.T) { cell := newTestCell("cell2") cell.Spec.Metadata = `{"zoneId":"use1-az1","custom":"value"}` - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(t.Context(), store, recorder, cell, false), "unexpected error") got, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if got.Metadata != `{"zoneId":"use1-az1","custom":"value"}` { - t.Errorf("expected metadata copied verbatim, got %q", got.Metadata) - } + c.Require().NoError(err, "cell not found") + c.Eq( + `{"zoneId":"use1-az1","custom":"value"}`, + got.Metadata, + "expected metadata copied verbatim, got", + ) }) t.Run("updates metadata on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -153,22 +148,16 @@ func TestRegisterCell(t *testing.T) { cell := newTestCell("cell1") cell.Spec.Metadata = `{"zoneId":"use1-az1"}` - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("first registration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(t.Context(), store, recorder, cell, false), "first registration failed") cell.Spec.Metadata = `{"zoneId":"use1-az2"}` - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("re-registration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(t.Context(), store, recorder, cell, false), "re-registration failed") got, err := store.GetCell(context.Background(), "cell1") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if got.Metadata != `{"zoneId":"use1-az2"}` { - t.Errorf("expected updated metadata, got %q", got.Metadata) - } + c.Require().NoError(err, "cell not found") + c.Eq(`{"zoneId":"use1-az2"}`, got.Metadata, "expected updated metadata, got") }) t.Run("returns error on failure", func(t *testing.T) { @@ -183,13 +172,12 @@ func TestRegisterCell(t *testing.T) { cell := newTestCell("cell1") err := topo.RegisterCell(t.Context(), store, recorder, cell, false) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("idempotent when cell already exists", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -199,16 +187,19 @@ func TestRegisterCell(t *testing.T) { recorder := record.NewFakeRecorder(10) cell := newTestCell("cell1") - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("first registration failed: %v", err) - } - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("second registration should succeed (idempotent), got: %v", err) - } + c.NoError( + topo.RegisterCell(t.Context(), store, recorder, cell, false), + "first registration failed", + ) + c.NoError( + topo.RegisterCell(t.Context(), store, recorder, cell, false), + "second registration should succeed (idempotent), got", + ) }) t.Run("updates stale cell topology on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -217,7 +208,7 @@ func TestRegisterCell(t *testing.T) { recorder := record.NewFakeRecorder(10) cell := newTestCell("cell1") - if err := store.UpdateCellFields( + c.Require().NoError(store.UpdateCellFields( context.Background(), "cell1", func(existing *clustermetadata.Cell) error { @@ -225,31 +216,24 @@ func TestRegisterCell(t *testing.T) { existing.Root = "/stale/root" return nil }, - ); err != nil { - t.Fatalf("seeding stale cell: %v", err) - } + ), "seeding stale cell") - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("re-registration should update stale cell, got: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(t.Context(), store, recorder, cell, false), "re-registration should update stale cell, got") got, err := store.GetCell(context.Background(), "cell1") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if !reflect.DeepEqual( - got.ServerAddresses, + c.Require().NoError(err, "cell not found") + c.EqDiff( []string{"http://local-etcd-1:2379", "http://local-etcd-2:2379"}, - ) { - t.Errorf("expected local topo addresses, got %v", got.ServerAddresses) - } - if got.Root != "/multigres/cells/cell1" { - t.Errorf("expected local topo root, got %s", got.Root) - } + got.ServerAddresses, + "expected local topo addresses, got", + ) + c.Eq("/multigres/cells/cell1", got.Root, "expected local topo root, got") }) t.Run("falls back to global topology when no local topology is configured", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -260,24 +244,22 @@ func TestRegisterCell(t *testing.T) { cell := newTestCell("cell2") cell.Spec.TopoServer = nil - if err := topo.RegisterCell(t.Context(), store, recorder, cell, false); err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(t.Context(), store, recorder, cell, false), "unexpected error") got, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if !reflect.DeepEqual(got.ServerAddresses, []string{"localhost:2379"}) { - t.Errorf("expected global topo address fallback, got %v", got.ServerAddresses) - } - if got.Root != "/test" { - t.Errorf("expected global topology root fallback, got %s", got.Root) - } + c.Require().NoError(err, "cell not found") + c.EqDiff( + []string{"localhost:2379"}, + got.ServerAddresses, + "expected global topo address fallback, got", + ) + c.Eq("/test", got.Root, "expected global topology root fallback, got") }) t.Run("uses project identity for a defaulted local topology root", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -288,19 +270,13 @@ func TestRegisterCell(t *testing.T) { cell.Annotations = map[string]string{metadata.AnnotationProjectRef: "proj_123"} cell.Spec.TopoServer.External.RootPath = "" - if err := topo.RegisterCell( + c.Require().NoError(topo.RegisterCell( context.Background(), store, record.NewFakeRecorder(10), cell, false, - ); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ), "unexpected error") got, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if got.Root != "/multigres/proj_123/cell2" { - t.Errorf("expected project-scoped cell root, got %s", got.Root) - } + c.Require().NoError(err, "cell not found") + c.Eq("/multigres/proj_123/cell2", got.Root, "expected project-scoped cell root, got") }) } @@ -309,6 +285,7 @@ func TestUnregisterCell(t *testing.T) { t.Run("removes existing cell", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -319,17 +296,13 @@ func TestUnregisterCell(t *testing.T) { cell := newTestCell("cell1") ctx := context.Background() - if err := topo.RegisterCell(ctx, store, recorder, cell, false); err != nil { - t.Fatalf("registration failed: %v", err) - } - if err := topo.UnregisterCell(ctx, store, recorder, cell); err != nil { - t.Fatalf("unregistration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterCell(ctx, store, recorder, cell, false), "registration failed") + c.Require(). + NoError(topo.UnregisterCell(ctx, store, recorder, cell), "unregistration failed") _, err := store.GetCell(ctx, "cell1") - if err == nil { - t.Error("expected cell to be gone from topo after unregistration") - } + c.Error(err, "expected cell to be gone from topo after unregistration") }) t.Run("idempotent when cell does not exist", func(t *testing.T) { @@ -343,9 +316,8 @@ func TestUnregisterCell(t *testing.T) { recorder := record.NewFakeRecorder(10) cell := newTestCell("nonexistent") - if err := topo.UnregisterCell(context.Background(), store, recorder, cell); err != nil { - t.Fatalf("unregistering nonexistent cell should succeed (idempotent), got: %v", err) - } + assert.NewAborting(t). + NoError(topo.UnregisterCell(context.Background(), store, recorder, cell), "unregistering nonexistent cell should succeed (idempotent), got") }) t.Run("returns error on failure other than TopoUnavailable", func(t *testing.T) { @@ -360,9 +332,7 @@ func TestUnregisterCell(t *testing.T) { cell := newTestCell("cell1") err := topo.UnregisterCell(context.Background(), store, recorder, cell) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("returns error on TopoUnavailable", func(t *testing.T) { @@ -377,8 +347,6 @@ func TestUnregisterCell(t *testing.T) { cell := newTestCell("cell1") err := topo.UnregisterCell(context.Background(), store, recorder, cell) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) } diff --git a/pkg/data-handler/topo/database_test.go b/pkg/data-handler/topo/database_test.go index 46244585..106b4e9a 100644 --- a/pkg/data-handler/topo/database_test.go +++ b/pkg/data-handler/topo/database_test.go @@ -3,7 +3,6 @@ package topo_test import ( "context" "fmt" - "strings" "testing" "github.com/multigres/multigres/go/common/topoclient" @@ -14,6 +13,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/topo" + + "github.com/multigres/testkit/assert" ) func newTestShard(name string) *multigresv1alpha1.Shard { @@ -79,6 +80,7 @@ func TestRegisterDatabase(t *testing.T) { t.Run("creates new database in topology", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -88,21 +90,17 @@ func TestRegisterDatabase(t *testing.T) { recorder := record.NewFakeRecorder(10) shard := newTestShard("test-shard") - if err := topo.RegisterDatabase(context.Background(), store, recorder, shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require(). + NoError(topo.RegisterDatabase(context.Background(), store, recorder, shard), "unexpected error") db, err := store.GetDatabase(context.Background(), "test-db") - if err != nil { - t.Fatalf("database not found in topo after registration: %v", err) - } - if db.Name != "test-db" { - t.Errorf("expected database name test-db, got %s", db.Name) - } + c.Require().NoError(err, "database not found in topo after registration") + c.Eq("test-db", db.Name, "expected database name test-db, got") }) t.Run("updates existing database on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1", "cell2") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -113,29 +111,24 @@ func TestRegisterDatabase(t *testing.T) { shard := newTestShard("test-shard") ctx := context.Background() - if err := topo.RegisterDatabase(ctx, store, recorder, shard); err != nil { - t.Fatalf("first registration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterDatabase(ctx, store, recorder, shard), "first registration failed") // Modify shard to add a second cell, re-register should update. shard.Spec.Pools["pool2"] = multigresv1alpha1.PoolSpec{ Cells: []multigresv1alpha1.CellName{"cell2"}, } - if err := topo.RegisterDatabase(ctx, store, recorder, shard); err != nil { - t.Fatalf("second registration (update) failed: %v", err) - } + c.Require(). + NoError(topo.RegisterDatabase(ctx, store, recorder, shard), "second registration (update) failed") db, err := store.GetDatabase(ctx, "test-db") - if err != nil { - t.Fatalf("database not found after update: %v", err) - } - if len(db.Cells) != 2 { - t.Errorf("expected 2 cells after update, got %d: %v", len(db.Cells), db.Cells) - } + c.Require().NoError(err, "database not found after update") + c.Len(db.Cells, 2, "expected 2 cells after update, got %d", len(db.Cells)) }) t.Run("syncs durability policy on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -147,35 +140,27 @@ func TestRegisterDatabase(t *testing.T) { ctx := context.Background() // First registration with default AT_LEAST_2. - if err := topo.RegisterDatabase(ctx, store, recorder, shard); err != nil { - t.Fatalf("first registration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterDatabase(ctx, store, recorder, shard), "first registration failed") db, err := store.GetDatabase(ctx, "test-db") - if err != nil { - t.Fatalf("database not found: %v", err) - } - if db.BootstrapDurabilityPolicy.GetPolicyName() != "AT_LEAST_2" { - t.Errorf( - "expected AT_LEAST_2 after first registration, got %s", - db.BootstrapDurabilityPolicy.GetPolicyName(), - ) - } + c.Require().NoError(err, "database not found") + c.Eq( + "AT_LEAST_2", + db.BootstrapDurabilityPolicy.GetPolicyName(), + "expected AT_LEAST_2 after first registration, got", + ) // Change to MULTI_CELL_AT_LEAST_2 and re-register. shard.Spec.DurabilityPolicy = "MULTI_CELL_AT_LEAST_2" - if err := topo.RegisterDatabase(ctx, store, recorder, shard); err != nil { - t.Fatalf("second registration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterDatabase(ctx, store, recorder, shard), "second registration failed") db, err = store.GetDatabase(ctx, "test-db") - if err != nil { - t.Fatalf("database not found after update: %v", err) - } - if db.BootstrapDurabilityPolicy.GetPolicyName() != "MULTI_CELL_AT_LEAST_2" { - t.Errorf( - "expected MULTI_CELL_AT_LEAST_2 after update, got %s", - db.BootstrapDurabilityPolicy.GetPolicyName(), - ) - } + c.Require().NoError(err, "database not found after update") + c.Eq( + "MULTI_CELL_AT_LEAST_2", + db.BootstrapDurabilityPolicy.GetPolicyName(), + "expected MULTI_CELL_AT_LEAST_2 after update, got", + ) }) t.Run("returns error on creation failure", func(t *testing.T) { @@ -190,9 +175,7 @@ func TestRegisterDatabase(t *testing.T) { shard := newTestShard("test-shard") err := topo.RegisterDatabase(context.Background(), store, recorder, shard) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("returns error on UpdateDatabaseFields failure", func(t *testing.T) { @@ -210,9 +193,7 @@ func TestRegisterDatabase(t *testing.T) { shard := newTestShard("test-shard") err := topo.RegisterDatabase(context.Background(), store, recorder, shard) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) } @@ -221,6 +202,7 @@ func TestUnregisterDatabase(t *testing.T) { t.Run("removes existing database", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -231,17 +213,13 @@ func TestUnregisterDatabase(t *testing.T) { shard := newTestShard("test-shard") ctx := context.Background() - if err := topo.RegisterDatabase(ctx, store, recorder, shard); err != nil { - t.Fatalf("registration failed: %v", err) - } - if err := topo.UnregisterDatabase(ctx, store, recorder, shard); err != nil { - t.Fatalf("unregistration failed: %v", err) - } + c.Require(). + NoError(topo.RegisterDatabase(ctx, store, recorder, shard), "registration failed") + c.Require(). + NoError(topo.UnregisterDatabase(ctx, store, recorder, shard), "unregistration failed") _, err := store.GetDatabase(ctx, "test-db") - if err == nil { - t.Error("expected database to be gone after unregistration") - } + c.Error(err, "expected database to be gone after unregistration") }) t.Run("idempotent when database does not exist", func(t *testing.T) { @@ -255,14 +233,12 @@ func TestUnregisterDatabase(t *testing.T) { recorder := record.NewFakeRecorder(10) shard := newTestShard("test-shard") - if err := topo.UnregisterDatabase( + assert.NewAborting(t).NoError(topo.UnregisterDatabase( context.Background(), store, recorder, shard, - ); err != nil { - t.Fatalf("unregistering nonexistent database should succeed, got: %v", err) - } + ), "unregistering nonexistent database should succeed, got") }) t.Run("returns error on failure other than TopoUnavailable", func(t *testing.T) { @@ -277,9 +253,7 @@ func TestUnregisterDatabase(t *testing.T) { shard := newTestShard("test-shard") err := topo.UnregisterDatabase(context.Background(), store, recorder, shard) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("returns error on TopoUnavailable", func(t *testing.T) { @@ -294,9 +268,7 @@ func TestUnregisterDatabase(t *testing.T) { shard := newTestShard("test-shard") err := topo.UnregisterDatabase(context.Background(), store, recorder, shard) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) } @@ -305,76 +277,62 @@ func TestGetDurabilityPolicy(t *testing.T) { t.Run("defaults to AT_LEAST_2 when empty", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := newTestShard("test-shard") got, err := topo.GetDurabilityPolicy(shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got.GetPolicyName() != "AT_LEAST_2" { - t.Errorf("expected AT_LEAST_2, got %s", got.GetPolicyName()) - } - if got.GetQuorumType() != clustermetadatapb.QuorumType_QUORUM_TYPE_AT_LEAST_N { - t.Errorf("expected QUORUM_TYPE_AT_LEAST_N, got %s", got.GetQuorumType()) - } - if got.GetRequiredCount() != 2 { - t.Errorf("expected RequiredCount 2, got %d", got.GetRequiredCount()) - } + c.Require().NoError(err, "unexpected error") + c.Eq("AT_LEAST_2", got.GetPolicyName(), "expected AT_LEAST_2, got") + c.Eq( + clustermetadatapb.QuorumType_QUORUM_TYPE_AT_LEAST_N, + got.GetQuorumType(), + "expected QUORUM_TYPE_AT_LEAST_N, got", + ) + c.Eq(2, got.GetRequiredCount(), "expected RequiredCount 2, got") }) t.Run("returns explicit AT_LEAST_2", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := newTestShard("test-shard") shard.Spec.DurabilityPolicy = "AT_LEAST_2" got, err := topo.GetDurabilityPolicy(shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got.GetPolicyName() != "AT_LEAST_2" { - t.Errorf("expected AT_LEAST_2, got %s", got.GetPolicyName()) - } - if got.GetQuorumType() != clustermetadatapb.QuorumType_QUORUM_TYPE_AT_LEAST_N { - t.Errorf("expected QUORUM_TYPE_AT_LEAST_N, got %s", got.GetQuorumType()) - } - if got.GetRequiredCount() != 2 { - t.Errorf("expected RequiredCount 2, got %d", got.GetRequiredCount()) - } + c.Require().NoError(err, "unexpected error") + c.Eq("AT_LEAST_2", got.GetPolicyName(), "expected AT_LEAST_2, got") + c.Eq( + clustermetadatapb.QuorumType_QUORUM_TYPE_AT_LEAST_N, + got.GetQuorumType(), + "expected QUORUM_TYPE_AT_LEAST_N, got", + ) + c.Eq(2, got.GetRequiredCount(), "expected RequiredCount 2, got") }) t.Run("returns MULTI_CELL_AT_LEAST_2 when set", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := newTestShard("test-shard") shard.Spec.DurabilityPolicy = "MULTI_CELL_AT_LEAST_2" got, err := topo.GetDurabilityPolicy(shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got.GetPolicyName() != "MULTI_CELL_AT_LEAST_2" { - t.Errorf("expected MULTI_CELL_AT_LEAST_2, got %s", got.GetPolicyName()) - } - if got.GetQuorumType() != clustermetadatapb.QuorumType_QUORUM_TYPE_MULTI_CELL_AT_LEAST_N { - t.Errorf("expected QUORUM_TYPE_MULTI_CELL_AT_LEAST_N, got %s", got.GetQuorumType()) - } - if got.GetRequiredCount() != 2 { - t.Errorf("expected RequiredCount 2, got %d", got.GetRequiredCount()) - } + c.Require().NoError(err, "unexpected error") + c.Eq("MULTI_CELL_AT_LEAST_2", got.GetPolicyName(), "expected MULTI_CELL_AT_LEAST_2, got") + c.Eq( + clustermetadatapb.QuorumType_QUORUM_TYPE_MULTI_CELL_AT_LEAST_N, + got.GetQuorumType(), + "expected QUORUM_TYPE_MULTI_CELL_AT_LEAST_N, got", + ) + c.Eq(2, got.GetRequiredCount(), "expected RequiredCount 2, got") }) t.Run("unknown policy returns error", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := newTestShard("test-shard") shard.Spec.DurabilityPolicy = "CUSTOM" got, err := topo.GetDurabilityPolicy(shard) - if err == nil { - t.Fatal("expected error, got nil") - } - if got != nil { - t.Errorf("expected nil policy on error, got %+v", got) - } + c.Require().Error(err, "expected error, got nil") + c.Nil(got, "expected nil policy on error, got") msg := err.Error() for _, want := range []string{"CUSTOM", "AT_LEAST_2", "MULTI_CELL_AT_LEAST_2"} { - if !strings.Contains(msg, want) { - t.Errorf("expected error message to contain %q, got %q", want, msg) - } + c.StrContains(msg, want, "expected error message to contain") } }) } @@ -384,6 +342,7 @@ func TestGetBackupLocation(t *testing.T) { t.Run("S3 backup", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ Spec: multigresv1alpha1.ShardSpec{ Backup: &multigresv1alpha1.BackupConfig{ @@ -400,28 +359,17 @@ func TestGetBackupLocation(t *testing.T) { } loc := topo.GetBackupLocation(shard) s3 := loc.GetS3() - if s3 == nil { - t.Fatal("expected S3 backup location") - } - if s3.Bucket != "my-bucket" { - t.Errorf("expected bucket my-bucket, got %s", s3.Bucket) - } - if s3.Region != "us-west-2" { - t.Errorf("expected region us-west-2, got %s", s3.Region) - } - if s3.Endpoint != "https://s3.example.com" { - t.Errorf("expected endpoint, got %s", s3.Endpoint) - } - if s3.KeyPrefix != "prefix/" { - t.Errorf("expected key prefix, got %s", s3.KeyPrefix) - } - if !s3.UseEnvCredentials { - t.Error("expected UseEnvCredentials=true") - } + c.Require().NotNil(s3, "expected S3 backup location") + c.Eq("my-bucket", s3.Bucket, "expected bucket my-bucket, got") + c.Eq("us-west-2", s3.Region, "expected region us-west-2, got") + c.Eq("https://s3.example.com", s3.Endpoint, "expected endpoint, got") + c.Eq("prefix/", s3.KeyPrefix, "expected key prefix, got") + c.True(s3.UseEnvCredentials, "expected UseEnvCredentials=true") }) t.Run("filesystem with custom path", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ Spec: multigresv1alpha1.ShardSpec{ Backup: &multigresv1alpha1.BackupConfig{ @@ -434,25 +382,18 @@ func TestGetBackupLocation(t *testing.T) { } loc := topo.GetBackupLocation(shard) fs := loc.GetFilesystem() - if fs == nil { - t.Fatal("expected filesystem backup location") - } - if fs.Path != "/custom/backups" { - t.Errorf("expected path /custom/backups, got %s", fs.Path) - } + c.Require().NotNil(fs, "expected filesystem backup location") + c.Eq("/custom/backups", fs.Path, "expected path /custom/backups, got") }) t.Run("default filesystem", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} loc := topo.GetBackupLocation(shard) fs := loc.GetFilesystem() - if fs == nil { - t.Fatal("expected filesystem backup location") - } - if fs.Path != "/backups" { - t.Errorf("expected default path /backups, got %s", fs.Path) - } + c.Require().NotNil(fs, "expected filesystem backup location") + c.Eq("/backups", fs.Path, "expected default path /backups, got") }) t.Run("encryption enabled", func(t *testing.T) { @@ -471,17 +412,15 @@ func TestGetBackupLocation(t *testing.T) { }, } loc := topo.GetBackupLocation(shard) - if !loc.GetRequireInitialRepoEncryption() { - t.Error("expected RequireInitialRepoEncryption=true") - } + assert.NewCollecting(t). + True(loc.GetRequireInitialRepoEncryption(), "expected RequireInitialRepoEncryption=true") }) t.Run("encryption not set", func(t *testing.T) { t.Parallel() shard := &multigresv1alpha1.Shard{} loc := topo.GetBackupLocation(shard) - if loc.GetRequireInitialRepoEncryption() { - t.Error("expected RequireInitialRepoEncryption=false") - } + assert.NewCollecting(t). + False(loc.GetRequireInitialRepoEncryption(), "expected RequireInitialRepoEncryption=false") }) } diff --git a/pkg/data-handler/topo/idempotency_test.go b/pkg/data-handler/topo/idempotency_test.go index 0c91b364..efd540fe 100644 --- a/pkg/data-handler/topo/idempotency_test.go +++ b/pkg/data-handler/topo/idempotency_test.go @@ -10,6 +10,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/topo" + + "github.com/multigres/testkit/assert" ) // Inspect the backing record version, not just the returned value: rewriting @@ -19,6 +21,7 @@ func TestRegistrationDoesNotRewriteUnchangedRecords(t *testing.T) { for _, kind := range []string{"cell", "database", "shard-database"} { t.Run(kind, func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) ctx := t.Context() store := newMemoryStore(t) recorder := record.NewFakeRecorder(100) @@ -68,28 +71,18 @@ func TestRegistrationDoesNotRewriteUnchangedRecords(t *testing.T) { file = path.Join(topoclient.CellsPath, "cell1", topoclient.CellFile) } conn, err := store.ConnForCell(ctx, topoclient.GlobalCell) - if err != nil { - t.Fatal(err) - } + c.NoError(err) version := func() string { t.Helper() _, v, err := conn.Get(ctx, file) - if err != nil { - t.Fatal(err) - } + c.NoError(err) return v.String() } - if err := register(); err != nil { - t.Fatal(err) - } + c.NoError(register()) initial := version() for range 5 { - if err := register(); err != nil { - t.Fatal(err) - } - if got := version(); got != initial { - t.Fatalf("unchanged registration rewrote record: %s -> %s", initial, got) - } + c.NoError(register()) + c.Eq(initial, version(), "unchanged registration rewrote record") } switch kind { case "cell": @@ -101,19 +94,11 @@ func TestRegistrationDoesNotRewriteUnchangedRecords(t *testing.T) { Cells: []multigresv1alpha1.CellName{"cell2"}, } } - if err := register(); err != nil { - t.Fatal(err) - } + c.NoError(register()) changed := version() - if changed == initial { - t.Fatal("changed registration did not update record") - } - if err := register(); err != nil { - t.Fatal(err) - } - if version() != changed { - t.Fatal("registration did not converge after change") - } + c.NotEq(initial, changed, "changed registration did not update record") + c.NoError(register()) + c.Eq(changed, version(), "registration did not converge after change") }) } } diff --git a/pkg/data-handler/topo/pooler_test.go b/pkg/data-handler/topo/pooler_test.go index cb95ee85..8edf4000 100644 --- a/pkg/data-handler/topo/pooler_test.go +++ b/pkg/data-handler/topo/pooler_test.go @@ -12,6 +12,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/topo" + + "github.com/multigres/testkit/assert" ) func routingState(role clustermetadata.RoutingRole) *clustermetadata.RoutingState { @@ -23,6 +25,7 @@ func TestFindPrimaryPooler(t *testing.T) { t.Run("returns nil when no primary exists", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(context.Background(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -43,12 +46,8 @@ func TestFindPrimaryPooler(t *testing.T) { shard, []string{"cell1"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if primary != nil { - t.Error("expected nil primary when none registered") - } + c.Require().NoError(err, "unexpected error") + c.Nil(primary, "expected nil primary when none registered") }) t.Run("returns error for non-unavailable topo errors", func(t *testing.T) { @@ -68,13 +67,12 @@ func TestFindPrimaryPooler(t *testing.T) { _, err := topo.FindPrimaryPooler( t.Context(), store, shard, []string{"nonexistent-cell"}, ) - if err == nil { - t.Error("expected error for non-unavailable topo error") - } + assert.NewCollecting(t).Error(err, "expected error for non-unavailable topo error") }) t.Run("returns primary from second cell", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1", "cell2") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -102,15 +100,9 @@ func TestFindPrimaryPooler(t *testing.T) { } primary, err := topo.FindPrimaryPooler(ctx, store, shard, []string{"cell1", "cell2"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if primary == nil { - t.Fatal("expected primary to be found") - } - if primary.Id.Name != "primary-pod" { - t.Errorf("expected primary-pod, got %s", primary.Id.Name) - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(primary, "expected primary to be found") + c.Eq("primary-pod", primary.Id.Name, "expected primary-pod, got") }) t.Run("skips a shut-down primary", func(t *testing.T) { @@ -141,9 +133,7 @@ func TestFindPrimaryPooler(t *testing.T) { } primary, err := topo.FindPrimaryPooler(ctx, store, shard, []string{"cell1"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") if primary != nil { t.Errorf("expected nil primary (dead primary skipped), got %s", primary.Id.Name) } @@ -175,9 +165,7 @@ func TestFindPrimaryPooler(t *testing.T) { } primary, err := topo.FindPrimaryPooler(ctx, store, shard, []string{"cell1"}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") if primary != nil { t.Errorf("expected nil primary (quarantined primary skipped), got %s", primary.Id.Name) } @@ -234,6 +222,7 @@ func TestFindPrimaryPooler_TopoUnavailableSkip(t *testing.T) { t.Run("skips unavailable cell and finds primary in next cell", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := &multiCellStore{ cells: map[string][]*topoclient.MultipoolerInfo{ "cell2": {{ @@ -263,19 +252,14 @@ func TestFindPrimaryPooler_TopoUnavailableSkip(t *testing.T) { primary, err := topo.FindPrimaryPooler( t.Context(), store, shard, []string{"cell1", "cell2"}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if primary == nil { - t.Fatal("expected primary from cell2 after skipping unavailable cell1") - } - if primary.Id.Name != "primary-pod" { - t.Errorf("expected primary-pod, got %s", primary.Id.Name) - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(primary, "expected primary from cell2 after skipping unavailable cell1") + c.Eq("primary-pod", primary.Id.Name, "expected primary-pod, got") }) t.Run("returns error when all cells are unavailable", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := &multiCellStore{ errorCells: map[string]error{ "cell1": errors.New("Code: UNAVAILABLE"), @@ -292,12 +276,8 @@ func TestFindPrimaryPooler_TopoUnavailableSkip(t *testing.T) { primary, err := topo.FindPrimaryPooler( t.Context(), store, shard, []string{"cell1", "cell2"}, ) - if err == nil { - t.Fatal("expected error when all cells are unavailable") - } - if primary != nil { - t.Error("expected nil primary when all cells are unavailable") - } + c.Require().Error(err, "expected error when all cells are unavailable") + c.Nil(primary, "expected nil primary when all cells are unavailable") }) } @@ -311,6 +291,7 @@ func TestMarkDeadPoolers(t *testing.T) { t.Run("marks dead poolers shut down without deleting them", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -348,18 +329,13 @@ func TestMarkDeadPoolers(t *testing.T) { activePods := map[string]bool{"active-pod": true} marked, err := topo.MarkDeadPoolers(ctx, store, shard, activePods) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if marked != 2 { - t.Errorf("expected 2 marked, got %d", marked) - } + c.Require().NoError(err, "unexpected error") + c.Eq(2, marked, "expected 2 marked, got") // Entries are left in place (tombstones), not deleted. remaining, _ := store.GetMultipoolersByCell(ctx, "cell1", nil) - if len(remaining) != 3 { - t.Fatalf("expected 3 remaining poolers (none deleted), got %d", len(remaining)) - } + c.Require(). + Len(remaining, 3, "expected 3 remaining poolers (none deleted), got %d", len(remaining)) byName := make(map[string]*clustermetadata.Multipooler, len(remaining)) for _, p := range remaining { @@ -375,21 +351,23 @@ func TestMarkDeadPoolers(t *testing.T) { } } mp := byName["stale-pod"] - if topo.PoolerRoutingRole(mp) != clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA { - t.Errorf( - "expected stale-pod routing role left untouched (REPLICA), got %v", - topo.PoolerRoutingRole(mp), - ) - } + c.Eq( + clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + topo.PoolerRoutingRole(mp), + "expected stale-pod routing role left untouched (REPLICA), got", + ) // Active pooler left untouched. - if active := byName["active-pod"]; isShutdown(active) { - t.Error("expected active-pod to be left untouched, but it was marked shut down") - } + active := byName["active-pod"] + c.False( + isShutdown(active), + "expected active-pod to be left untouched, but it was marked shut down", + ) }) t.Run("is idempotent for already-shutdown poolers", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -417,16 +395,13 @@ func TestMarkDeadPoolers(t *testing.T) { } marked, err := topo.MarkDeadPoolers(ctx, store, shard, map[string]bool{}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if marked != 0 { - t.Errorf("expected 0 marked for already-shutdown pooler, got %d", marked) - } + c.Require().NoError(err, "unexpected error") + c.Eq(0, marked, "expected 0 marked for already-shutdown pooler, got") }) t.Run("noop when all poolers are active", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -452,16 +427,13 @@ func TestMarkDeadPoolers(t *testing.T) { activePods := map[string]bool{"pod-1": true} marked, err := topo.MarkDeadPoolers(ctx, store, shard, activePods) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if marked != 0 { - t.Errorf("expected 0 marked, got %d", marked) - } + c.Require().NoError(err, "unexpected error") + c.Eq(0, marked, "expected 0 marked, got") }) t.Run("skips unavailable cells gracefully", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := &multiCellStore{ errorCells: map[string]error{ "cell1": errors.New("Code: UNAVAILABLE"), @@ -480,16 +452,13 @@ func TestMarkDeadPoolers(t *testing.T) { marked, err := topo.MarkDeadPoolers( t.Context(), store, shard, map[string]bool{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if marked != 0 { - t.Errorf("expected 0 marked for unavailable cell, got %d", marked) - } + c.Require().NoError(err, "unexpected error") + c.Eq(0, marked, "expected 0 marked for unavailable cell, got") }) t.Run("does not mark active poolers with FQDN hostnames", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -522,24 +491,21 @@ func TestMarkDeadPoolers(t *testing.T) { // active-pod is in the active set; stale-pod is NOT. activePods := map[string]bool{"active-pod": true} marked, err := topo.MarkDeadPoolers(ctx, store, shard, activePods) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if marked != 1 { - t.Errorf("expected 1 marked (stale-pod), got %d", marked) - } + c.Require().NoError(err, "unexpected error") + c.Eq(1, marked, "expected 1 marked (stale-pod), got") remaining, _ := store.GetMultipoolersByCell(ctx, "cell1", nil) - if len(remaining) != 2 { - t.Fatalf("expected 2 remaining poolers (none deleted), got %d", len(remaining)) - } + c.Require(). + Len(remaining, 2, "expected 2 remaining poolers (none deleted), got %d", len(remaining)) for _, p := range remaining { - if p.Id.Name == "active-pod" && isShutdown(p.Multipooler) { - t.Error("expected active-pod to be left untouched") - } - if p.Id.Name == "stale-pod" && !isShutdown(p.Multipooler) { - t.Error("expected stale-pod to be marked LIFECYCLE_SHUTDOWN") - } + c.False( + p.Id.Name == "active-pod" && isShutdown(p.Multipooler), + "expected active-pod to be left untouched", + ) + c.False( + p.Id.Name == "stale-pod" && !isShutdown(p.Multipooler), + "expected stale-pod to be marked LIFECYCLE_SHUTDOWN", + ) } }) @@ -557,13 +523,12 @@ func TestMarkDeadPoolers(t *testing.T) { } _, err := topo.MarkDeadPoolers(t.Context(), store, shard, map[string]bool{}) - if err == nil { - t.Error("expected error when GetMultipoolersByCell fails") - } + assert.NewCollecting(t).Error(err, "expected error when GetMultipoolersByCell fails") }) t.Run("continues and logs error on UpdateMultipoolerFields failure", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := &mockPoolerTopoStore{ getMultipoolersByCellFunc: func(ctx context.Context, cell string, opts *topoclient.GetMultipoolersByCellOptions) ([]*topoclient.MultipoolerInfo, error) { p := &topoclient.MultipoolerInfo{ @@ -597,16 +562,13 @@ func TestMarkDeadPoolers(t *testing.T) { } marked, err := topo.MarkDeadPoolers(t.Context(), store, shard, map[string]bool{}) - if err != nil { - t.Fatalf("expected nil error (caught and logged), got %v", err) - } - if marked != 0 { - t.Errorf("expected 0 marked due to error, got %d", marked) - } + c.Require().NoError(err, "expected nil error (caught and logged), got") + c.Eq(0, marked, "expected 0 marked due to error, got") }) t.Run("uses Id.Name when hostname is empty", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := &mockPoolerTopoStore{ getMultipoolersByCellFunc: func(ctx context.Context, cell string, opts *topoclient.GetMultipoolersByCellOptions) ([]*topoclient.MultipoolerInfo, error) { p := &topoclient.MultipoolerInfo{ @@ -643,12 +605,8 @@ func TestMarkDeadPoolers(t *testing.T) { } marked, err := topo.MarkDeadPoolers(t.Context(), store, shard, map[string]bool{}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if marked != 1 { - t.Errorf("expected 1 marked, got %d", marked) - } + c.Require().NoError(err, "unexpected error") + c.Eq(1, marked, "expected 1 marked, got") }) } @@ -667,6 +625,7 @@ func (s *errorGetPoolersStore) Close() error { return nil } func TestCollectCells(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ Spec: multigresv1alpha1.ShardSpec{ @@ -678,18 +637,14 @@ func TestCollectCells(t *testing.T) { } cells := topo.CollectCells(shard) - if len(cells) != 3 { - t.Errorf("expected 3 unique cells, got %d: %v", len(cells), cells) - } + ck.Len(cells, 3, "expected 3 unique cells, got %d", len(cells)) cellSet := make(map[string]bool) for _, c := range cells { cellSet[c] = true } for _, want := range []string{"zone-a", "zone-b", "zone-c"} { - if !cellSet[want] { - t.Errorf("expected cell %q in result", want) - } + ck.False(!cellSet[want], "expected cell %q in result", want) } } @@ -698,6 +653,7 @@ func TestGetPoolerStatus(t *testing.T) { t.Run("returns advertised routing roles", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -748,27 +704,18 @@ func TestGetPoolerStatus(t *testing.T) { shard, []string{"primary", "replica", "unknown", "quarantined"}, ) - if !result.QuerySuccess { - t.Error("expected QuerySuccess=true") - } - if result.Roles["primary"] != "PRIMARY" { - t.Errorf("expected PRIMARY, got %s", result.Roles["primary"]) - } - if result.Roles["replica"] != "REPLICA" { - t.Errorf("expected REPLICA, got %s", result.Roles["replica"]) - } - if result.Roles["unknown"] != "REPLICA" { - t.Errorf("expected REPLICA fallback, got %s", result.Roles["unknown"]) - } + c.True(result.QuerySuccess, "expected QuerySuccess=true") + c.Eq("PRIMARY", result.Roles["primary"], "expected PRIMARY, got") + c.Eq("REPLICA", result.Roles["replica"], "expected REPLICA, got") + c.Eq("REPLICA", result.Roles["unknown"], "expected REPLICA fallback, got") // Quarantined poolers get a distinct QUARANTINED role (visible in status) // but are handled by quarantine remediation, not routed. - if result.Roles["quarantined"] != "QUARANTINED" { - t.Errorf("expected QUARANTINED, got %s", result.Roles["quarantined"]) - } + c.Eq("QUARANTINED", result.Roles["quarantined"], "expected QUARANTINED, got") }) t.Run("skips shut-down poolers even if a pod name matches", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -797,16 +744,13 @@ func TestGetPoolerStatus(t *testing.T) { } result := topo.GetPoolerStatus(ctx, store, shard, []string{"shutdown-pod"}) - if !result.QuerySuccess { - t.Error("expected QuerySuccess=true") - } - if len(result.Roles) != 0 { - t.Errorf("expected no roles for shut-down pooler, got %v", result.Roles) - } + c.True(result.QuerySuccess, "expected QuerySuccess=true") + c.Empty(result.Roles, "expected no roles for shut-down pooler, got") }) t.Run("skips orphaned poolers gracefully", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -832,12 +776,8 @@ func TestGetPoolerStatus(t *testing.T) { } result := topo.GetPoolerStatus(ctx, store, shard, []string{"other-pod"}) - if !result.QuerySuccess { - t.Error("expected QuerySuccess=true") - } - if len(result.Roles) != 0 { - t.Errorf("expected no roles mapped for orphaned pod, got %v", result.Roles) - } + c.True(result.QuerySuccess, "expected QuerySuccess=true") + c.Empty(result.Roles, "expected no roles mapped for orphaned pod, got") }) t.Run("sets QuerySuccess false on error", func(t *testing.T) { @@ -855,13 +795,13 @@ func TestGetPoolerStatus(t *testing.T) { } result := topo.GetPoolerStatus(t.Context(), store, shard, nil) - if result.QuerySuccess { - t.Error("expected QuerySuccess=false when store errors") - } + assert.NewCollecting(t). + False(result.QuerySuccess, "expected QuerySuccess=false when store errors") }) t.Run("uses Id.Name when hostname is empty", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory( factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), @@ -887,11 +827,12 @@ func TestGetPoolerStatus(t *testing.T) { } result := topo.GetPoolerStatus(ctx, store, shard, []string{"my-pod-0"}) - if !result.QuerySuccess { - t.Error("expected QuerySuccess=true") - } - if result.Roles["my-pod-0"] != "PRIMARY" { - t.Errorf("expected key 'my-pod-0' with PRIMARY, got roles: %v", result.Roles) - } + c.True(result.QuerySuccess, "expected QuerySuccess=true") + c.Eq( + "PRIMARY", + result.Roles["my-pod-0"], + "expected key 'my-pod-0' with PRIMARY, got roles: %v", + result.Roles, + ) }) } diff --git a/pkg/data-handler/topo/store_internal_test.go b/pkg/data-handler/topo/store_internal_test.go index f0a9a132..151d4d22 100644 --- a/pkg/data-handler/topo/store_internal_test.go +++ b/pkg/data-handler/topo/store_internal_test.go @@ -12,6 +12,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func tlsTestClient(objs ...client.Object) client.Client { @@ -21,19 +23,17 @@ func tlsTestClient(objs ...client.Object) client.Client { } func TestLoadClientTLS_PlaintextWhenNoSecrets(t *testing.T) { + c := assert.NewCollecting(t) ref := multigresv1alpha1.GlobalTopoServerRef{Address: "localhost:2379"} opts, err := loadClientTLS(context.Background(), tlsTestClient(), "team-a", ref) - if err != nil { - t.Fatalf("loadClientTLS() error = %v", err) - } - if opts != nil { - t.Errorf("expected nil TLS options for a plaintext reference, got %+v", opts) - } + c.Require().NoError(err, "loadClientTLS() error =") + c.Nil(opts, "expected nil TLS options for a plaintext reference, got") } // The managed path names one Secret for both the keypair and the CA. The loaded // options have to carry the tls.crt, tls.key, and ca.crt bytes verbatim. func TestLoadClientTLS_ManagedSecret(t *testing.T) { + c := assert.NewCollecting(t) secret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{Name: "topo-tls", Namespace: "team-a"}, Data: map[string][]byte{ @@ -49,22 +49,20 @@ func TestLoadClientTLS_ManagedSecret(t *testing.T) { } opts, err := loadClientTLS(context.Background(), tlsTestClient(secret), "team-a", ref) - if err != nil { - t.Fatalf("loadClientTLS() error = %v", err) - } - if opts == nil { - t.Fatal("expected TLS options, got nil") - } - if !bytes.Equal(opts.CertPEM, []byte("CERT")) || + c.Require().NoError(err, "loadClientTLS() error =") + c.Require().NotNil(opts, "expected TLS options, got nil") + c.False(!bytes.Equal(opts.CertPEM, []byte("CERT")) || !bytes.Equal(opts.KeyPEM, []byte("KEY")) || - !bytes.Equal(opts.CAPEM, []byte("CA")) { - t.Errorf("TLS options do not carry the Secret material verbatim: %+v", opts) - } + !bytes.Equal( + opts.CAPEM, + []byte("CA"), + ), "TLS options do not carry the Secret material verbatim: %+v", opts) } // An external topology can split the CA and the client keypair across two // Secrets; both are read. func TestLoadClientTLS_SeparateCAAndClientSecrets(t *testing.T) { + c := assert.NewCollecting(t) client := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{Name: "client", Namespace: "team-a"}, Data: map[string][]byte{"tls.crt": []byte("CERT"), "tls.key": []byte("KEY")}, @@ -80,10 +78,10 @@ func TestLoadClientTLS_SeparateCAAndClientSecrets(t *testing.T) { } opts, err := loadClientTLS(context.Background(), tlsTestClient(client, ca), "team-a", ref) - if err != nil { - t.Fatalf("loadClientTLS() error = %v", err) - } - if !bytes.Equal(opts.CertPEM, []byte("CERT")) || !bytes.Equal(opts.CAPEM, []byte("CA")) { - t.Errorf("expected material from both Secrets: %+v", opts) - } + c.Require().NoError(err, "loadClientTLS() error =") + c.False( + !bytes.Equal(opts.CertPEM, []byte("CERT")) || !bytes.Equal(opts.CAPEM, []byte("CA")), + "expected material from both Secrets: %+v", + opts, + ) } diff --git a/pkg/data-handler/topo/store_test.go b/pkg/data-handler/topo/store_test.go index c7bf2d64..158959a2 100644 --- a/pkg/data-handler/topo/store_test.go +++ b/pkg/data-handler/topo/store_test.go @@ -15,6 +15,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/topo" + + "github.com/multigres/testkit/assert" ) func topoTestClient(objs ...client.Object) client.Client { @@ -79,6 +81,7 @@ func TestNewStoreFromShard_InvalidImplementation(t *testing.T) { // fail with the Secret name, so a misconfigured cluster reports the missing // material instead of silently connecting without a certificate. func TestNewStoreFromShard_MissingSecretIsLoud(t *testing.T) { + c := assert.NewCollecting(t) tlsName := "cluster-topo-client-tls" shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Namespace: "team-a"}, @@ -93,17 +96,14 @@ func TestNewStoreFromShard_MissingSecretIsLoud(t *testing.T) { } _, err := topo.NewStoreFromShard(context.Background(), topoTestClient(), shard) - if err == nil { - t.Fatal("expected an error for a missing client credential Secret") - } - if !strings.Contains(err.Error(), tlsName) { - t.Errorf("error does not name the missing Secret: %v", err) - } + c.Require().Error(err, "expected an error for a missing client credential Secret") + c.StrContains(err.Error(), tlsName, "error does not name the missing Secret: %v", err) } // A Secret that exists but is missing a required key also fails loudly, naming // the key. func TestNewStoreFromRef_MissingKeyIsLoud(t *testing.T) { + c := assert.NewCollecting(t) secret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{Name: "topo-tls", Namespace: "team-a"}, Data: map[string][]byte{ @@ -120,12 +120,12 @@ func TestNewStoreFromRef_MissingKeyIsLoud(t *testing.T) { } _, err := topo.NewStoreFromRef(context.Background(), topoTestClient(secret), "team-a", ref) - if err == nil { - t.Fatal("expected an error for a Secret missing tls.key") - } - if !strings.Contains(err.Error(), "tls.key") || !strings.Contains(err.Error(), "topo-tls") { - t.Errorf("error does not name the Secret and missing key: %v", err) - } + c.Require().Error(err, "expected an error for a Secret missing tls.key") + c.False( + !strings.Contains(err.Error(), "tls.key") || !strings.Contains(err.Error(), "topo-tls"), + "error does not name the Secret and missing key: %v", + err, + ) } func TestNewStoreFromRef_WithClientCert(t *testing.T) { @@ -145,9 +145,7 @@ func TestNewStoreFromRef_WithClientCert(t *testing.T) { } store, err := topo.NewStoreFromRef(context.Background(), topoTestClient(secret), "team-a", ref) - if err != nil { - t.Fatalf("NewStoreFromRef() unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "NewStoreFromRef() unexpected error") if store != nil { _ = store.Close() } @@ -166,9 +164,9 @@ func TestIsTopoUnavailable(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { - if got := topo.IsTopoUnavailable(tc.err); got != tc.want { - t.Errorf("IsTopoUnavailable(%v) = %v, want %v", tc.err, got, tc.want) - } + got := topo.IsTopoUnavailable(tc.err) + assert.NewCollecting(t). + Eq(tc.want, got, "IsTopoUnavailable(%v) = %v, want", tc.err, got) }) } } @@ -184,9 +182,7 @@ func TestNewStoreFromCell(t *testing.T) { } store, err := topo.NewStoreFromCell(context.Background(), topoTestClient(), cell) - if err != nil { - t.Fatalf("NewStoreFromCell() unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "NewStoreFromCell() unexpected error") if store != nil { _ = store.Close() } @@ -219,9 +215,7 @@ func TestNewStoreFromRef(t *testing.T) { } store, err := topo.NewStoreFromRef(context.Background(), topoTestClient(), "team-a", ref) - if err != nil { - t.Fatalf("NewStoreFromRef() unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "NewStoreFromRef() unexpected error") if store != nil { _ = store.Close() } diff --git a/pkg/data-handler/topo/topology_test.go b/pkg/data-handler/topo/topology_test.go index 2a27ce95..b3aefb59 100644 --- a/pkg/data-handler/topo/topology_test.go +++ b/pkg/data-handler/topo/topology_test.go @@ -3,7 +3,6 @@ package topo_test import ( "context" "errors" - "reflect" "testing" "github.com/multigres/multigres/go/common/topoclient" @@ -14,6 +13,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/topo" + + "github.com/multigres/testkit/assert" ) func newMemoryStore(t *testing.T, cells ...string) topoclient.Store { @@ -44,6 +45,7 @@ func (f rootedMemoryFactory) Create( func TestSharedTopologyRootsIsolateClusters(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) const ( address = "shared-topo:2379" @@ -86,19 +88,15 @@ func TestSharedTopologyRootsIsolateClusters(t *testing.T) { {store: clusterA, cellRoot: cellRootA, host: "pooler-a"}, {store: clusterB, cellRoot: cellRootB, host: "pooler-b"}, } { - if err := cluster.store.CreateCell(ctx, cellName, &clustermetadatapb.Cell{ + c.NoError(cluster.store.CreateCell(ctx, cellName, &clustermetadatapb.Cell{ Name: cellName, ServerAddresses: []string{address}, Root: cluster.cellRoot, - }); err != nil { - t.Fatalf("CreateCell(%s): %v", cluster.cellRoot, err) - } - if err := cluster.store.CreateMultipooler( + }), "CreateCell(%s)", cluster.cellRoot) + c.NoError(cluster.store.CreateMultipooler( ctx, topoclient.NewMultipooler("pooler", cellName, cluster.host), - ); err != nil { - t.Fatalf("CreateMultipooler(%s): %v", cluster.cellRoot, err) - } + ), "CreateMultipooler(%s)", cluster.cellRoot) } for _, cluster := range []struct { @@ -109,15 +107,9 @@ func TestSharedTopologyRootsIsolateClusters(t *testing.T) { {store: clusterB, wantHost: "pooler-b"}, } { poolers, err := cluster.store.GetMultipoolersByCell(ctx, cellName, nil) - if err != nil { - t.Fatal(err) - } - if len(poolers) != 1 { - t.Fatalf("got %d multipoolers, want 1", len(poolers)) - } - if got := poolers[0].GetHostname(); got != cluster.wantHost { - t.Fatalf("multipooler host = %q, want %q", got, cluster.wantHost) - } + c.NoError(err) + c.Len(poolers, 1, "got %d multipoolers, want 1", len(poolers)) + c.Eq(cluster.wantHost, poolers[0].GetHostname(), "multipooler host") } } @@ -215,6 +207,7 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { t.Run("creates database with filesystem backup", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -226,20 +219,15 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { context.Background(), store, recorder, owner, dbConfig, []string{"cell1"}, nil, "", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") db, err := store.GetDatabase(context.Background(), "mydb") - if err != nil { - t.Fatalf("database not found: %v", err) - } - if db.BootstrapDurabilityPolicy.GetPolicyName() != "AT_LEAST_2" { - t.Errorf( - "expected default durability AT_LEAST_2, got %s", - db.BootstrapDurabilityPolicy.GetPolicyName(), - ) - } + c.Require().NoError(err, "database not found") + c.Eq( + "AT_LEAST_2", + db.BootstrapDurabilityPolicy.GetPolicyName(), + "expected default durability AT_LEAST_2, got", + ) fs := db.BackupLocation.GetFilesystem() if fs == nil || fs.Path != "/backups" { t.Errorf("expected filesystem backup at /backups, got %v", db.BackupLocation) @@ -248,6 +236,7 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { t.Run("creates database with S3 backup", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -264,14 +253,10 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { multigresv1alpha1.DatabaseConfig{Name: "s3db"}, []string{"cell1"}, backup, "", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") db, err := store.GetDatabase(context.Background(), "s3db") - if err != nil { - t.Fatalf("database not found: %v", err) - } + c.NoError(err, "database not found") s3 := db.BackupLocation.GetS3() if s3 == nil || s3.Bucket != "my-bucket" { t.Errorf("expected S3 backup with bucket my-bucket, got %v", db.BackupLocation) @@ -280,6 +265,7 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { t.Run("creates database with custom filesystem path", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -295,14 +281,10 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { multigresv1alpha1.DatabaseConfig{Name: "fsdb"}, []string{"cell1"}, backup, "", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") db, err := store.GetDatabase(context.Background(), "fsdb") - if err != nil { - t.Fatalf("database not found: %v", err) - } + c.NoError(err, "database not found") fs := db.BackupLocation.GetFilesystem() if fs == nil || fs.Path != "/custom/path" { t.Errorf("expected filesystem backup at /custom/path, got %v", db.BackupLocation) @@ -311,43 +293,36 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { t.Run("updates existing database on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1", "cell2") recorder := record.NewFakeRecorder(10) ctx := context.Background() dbConfig := multigresv1alpha1.DatabaseConfig{Name: "upddb"} - if err := topo.RegisterDatabaseFromSpec( + c.Require().NoError(topo.RegisterDatabaseFromSpec( ctx, store, recorder, owner, dbConfig, []string{"cell1"}, nil, "", - ); err != nil { - t.Fatalf("first registration: %v", err) - } + ), "first registration") // Re-register with different cells. - if err := topo.RegisterDatabaseFromSpec( + c.Require().NoError(topo.RegisterDatabaseFromSpec( ctx, store, recorder, owner, dbConfig, []string{"cell1", "cell2"}, nil, "MULTI_CELL_AT_LEAST_2", - ); err != nil { - t.Fatalf("re-registration: %v", err) - } + ), "re-registration") db, err := store.GetDatabase(ctx, "upddb") - if err != nil { - t.Fatalf("database not found: %v", err) - } - if len(db.Cells) != 2 { - t.Errorf("expected 2 cells, got %d", len(db.Cells)) - } - if db.BootstrapDurabilityPolicy.GetPolicyName() != "MULTI_CELL_AT_LEAST_2" { - t.Errorf( - "expected MULTI_CELL_AT_LEAST_2, got %s", - db.BootstrapDurabilityPolicy.GetPolicyName(), - ) - } + c.Require().NoError(err, "database not found") + c.Len(db.Cells, 2, "expected 2 cells, got %d", len(db.Cells)) + c.Eq( + "MULTI_CELL_AT_LEAST_2", + db.BootstrapDurabilityPolicy.GetPolicyName(), + "expected MULTI_CELL_AT_LEAST_2, got", + ) }) t.Run("creates database with encryption enabled", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -366,33 +341,29 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { multigresv1alpha1.DatabaseConfig{Name: "encdb"}, []string{"cell1"}, backup, "", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") db, err := store.GetDatabase(context.Background(), "encdb") - if err != nil { - t.Fatalf("database not found: %v", err) - } - if !db.BackupLocation.GetRequireInitialRepoEncryption() { - t.Error("expected RequireInitialRepoEncryption=true") - } + c.Require().NoError(err, "database not found") + c.True( + db.BackupLocation.GetRequireInitialRepoEncryption(), + "expected RequireInitialRepoEncryption=true", + ) }) t.Run("update path propagates encryption flag", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) ctx := context.Background() dbConfig := multigresv1alpha1.DatabaseConfig{Name: "encupddb"} - if err := topo.RegisterDatabaseFromSpec( + c.Require().NoError(topo.RegisterDatabaseFromSpec( ctx, store, recorder, owner, dbConfig, []string{"cell1"}, nil, "", - ); err != nil { - t.Fatalf("first registration: %v", err) - } + ), "first registration") backup := &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, @@ -404,24 +375,22 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { }, } - if err := topo.RegisterDatabaseFromSpec( + c.Require().NoError(topo.RegisterDatabaseFromSpec( ctx, store, recorder, owner, dbConfig, []string{"cell1"}, backup, "", - ); err != nil { - t.Fatalf("re-registration: %v", err) - } + ), "re-registration") db, err := store.GetDatabase(ctx, "encupddb") - if err != nil { - t.Fatalf("database not found: %v", err) - } - if !db.BackupLocation.GetRequireInitialRepoEncryption() { - t.Error("expected RequireInitialRepoEncryption=true after update") - } + c.Require().NoError(err, "database not found") + c.True( + db.BackupLocation.GetRequireInitialRepoEncryption(), + "expected RequireInitialRepoEncryption=true after update", + ) }) t.Run("uses database-level durability policy", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -434,20 +403,15 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { context.Background(), store, recorder, owner, dbConfig, []string{"cell1"}, nil, "AT_LEAST_2", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") db, err := store.GetDatabase(context.Background(), "durdb") - if err != nil { - t.Fatalf("database not found: %v", err) - } - if db.BootstrapDurabilityPolicy.GetPolicyName() != "MULTI_CELL_AT_LEAST_2" { - t.Errorf( - "expected database-level policy MULTI_CELL_AT_LEAST_2, got %s", - db.BootstrapDurabilityPolicy.GetPolicyName(), - ) - } + c.Require().NoError(err, "database not found") + c.Eq( + "MULTI_CELL_AT_LEAST_2", + db.BootstrapDurabilityPolicy.GetPolicyName(), + "expected database-level policy MULTI_CELL_AT_LEAST_2, got", + ) }) t.Run("returns error on CreateDatabase failure", func(t *testing.T) { @@ -462,9 +426,7 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { context.Background(), store, record.NewFakeRecorder(10), owner, multigresv1alpha1.DatabaseConfig{Name: "errdb"}, []string{"cell1"}, nil, "", ) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("returns error on UpdateDatabaseFields failure", func(t *testing.T) { @@ -482,9 +444,7 @@ func TestRegisterDatabaseFromSpec(t *testing.T) { context.Background(), store, record.NewFakeRecorder(10), owner, multigresv1alpha1.DatabaseConfig{Name: "errdb"}, []string{"cell1"}, nil, "", ) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) } @@ -501,6 +461,7 @@ func TestRegisterCellFromSpec(t *testing.T) { t.Run("creates new cell", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -514,27 +475,22 @@ func TestRegisterCellFromSpec(t *testing.T) { err := topo.RegisterCellFromSpec( context.Background(), store, recorder, owner, cellCfg, localTopo, topoRef, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") cell, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if cell.Name != "cell2" { - t.Errorf("expected cell2, got %s", cell.Name) - } - if !reflect.DeepEqual(cell.ServerAddresses, []string{"http://cell2-local:2379"}) { - t.Errorf("expected local topo addresses, got %v", cell.ServerAddresses) - } - if cell.Root != "/multigres/cells/cell2" { - t.Errorf("expected local topo root, got %s", cell.Root) - } + c.Require().NoError(err, "cell not found") + c.Eq("cell2", cell.Name, "expected cell2, got") + c.EqDiff( + []string{"http://cell2-local:2379"}, + cell.ServerAddresses, + "expected local topo addresses, got", + ) + c.Eq("/multigres/cells/cell2", cell.Root, "expected local topo root, got") }) t.Run("idempotent on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) ctx := context.Background() @@ -546,7 +502,7 @@ func TestRegisterCellFromSpec(t *testing.T) { RootPath: "/multigres/cells/cell1", }, } - if err := topo.RegisterCellFromSpec( + c.NoError(topo.RegisterCellFromSpec( ctx, store, recorder, @@ -554,10 +510,8 @@ func TestRegisterCellFromSpec(t *testing.T) { cellCfg, localTopo, topoRef, - ); err != nil { - t.Fatalf("first: %v", err) - } - if err := topo.RegisterCellFromSpec( + ), "first") + c.NoError(topo.RegisterCellFromSpec( ctx, store, recorder, @@ -565,18 +519,17 @@ func TestRegisterCellFromSpec(t *testing.T) { cellCfg, localTopo, topoRef, - ); err != nil { - t.Fatalf("second: %v", err) - } + ), "second") }) t.Run("updates stale cell topology on re-registration", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) ctx := context.Background() - if err := store.UpdateCellFields( + c.Require().NoError(store.UpdateCellFields( ctx, "cell1", func(existing *clustermetadatapb.Cell) error { @@ -584,9 +537,7 @@ func TestRegisterCellFromSpec(t *testing.T) { existing.Root = "/stale/cell1" return nil }, - ); err != nil { - t.Fatalf("seeding stale cell: %v", err) - } + ), "seeding stale cell") localTopo := &multigresv1alpha1.LocalTopoServerSpec{ External: &multigresv1alpha1.ExternalTopoServerSpec{ @@ -597,7 +548,7 @@ func TestRegisterCellFromSpec(t *testing.T) { RootPath: "/multigres/cells/cell1", }, } - if err := topo.RegisterCellFromSpec( + c.Require().NoError(topo.RegisterCellFromSpec( ctx, store, recorder, @@ -605,27 +556,21 @@ func TestRegisterCellFromSpec(t *testing.T) { multigresv1alpha1.CellConfig{Name: "cell1"}, localTopo, topoRef, - ); err != nil { - t.Fatalf("re-registration should update stale cell, got: %v", err) - } + ), "re-registration should update stale cell, got") cell, err := store.GetCell(ctx, "cell1") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if !reflect.DeepEqual( - cell.ServerAddresses, + c.Require().NoError(err, "cell not found") + c.EqDiff( []string{"http://cell1-local-a:2379", "http://cell1-local-b:2379"}, - ) { - t.Errorf("expected local topo addresses, got %v", cell.ServerAddresses) - } - if cell.Root != "/multigres/cells/cell1" { - t.Errorf("expected local topo root, got %s", cell.Root) - } + cell.ServerAddresses, + "expected local topo addresses, got", + ) + c.Eq("/multigres/cells/cell1", cell.Root, "expected local topo root, got") }) t.Run("falls back to global topology when no local topology is configured", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -638,24 +583,21 @@ func TestRegisterCellFromSpec(t *testing.T) { nil, topoRef, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") cell, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if !reflect.DeepEqual(cell.ServerAddresses, []string{"global:2379"}) { - t.Errorf("expected global topo address fallback, got %v", cell.ServerAddresses) - } - if cell.Root != "/multigres/global" { - t.Errorf("expected global topology root fallback, got %s", cell.Root) - } + c.Require().NoError(err, "cell not found") + c.EqDiff( + []string{"global:2379"}, + cell.ServerAddresses, + "expected global topo address fallback, got", + ) + c.Eq("/multigres/global", cell.Root, "expected global topology root fallback, got") }) t.Run("uses managed local topology address for etcd local topology", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -675,20 +617,16 @@ func TestRegisterCellFromSpec(t *testing.T) { topoRef, managedAddress, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") cell, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if !reflect.DeepEqual(cell.ServerAddresses, []string{managedAddress}) { - t.Errorf("expected managed local topo address, got %v", cell.ServerAddresses) - } - if cell.Root != "/multigres/cells/cell2" { - t.Errorf("expected managed local topo root, got %s", cell.Root) - } + c.Require().NoError(err, "cell not found") + c.EqDiff( + []string{managedAddress}, + cell.ServerAddresses, + "expected managed local topo address, got", + ) + c.Eq("/multigres/cells/cell2", cell.Root, "expected managed local topo root, got") }) t.Run("returns error on CreateCell failure", func(t *testing.T) { @@ -703,13 +641,12 @@ func TestRegisterCellFromSpec(t *testing.T) { context.Background(), store, record.NewFakeRecorder(10), owner, multigresv1alpha1.CellConfig{Name: "cell1"}, nil, topoRef, ) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("copies metadata from CellConfig into topo record", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) @@ -717,19 +654,13 @@ func TestRegisterCellFromSpec(t *testing.T) { Name: "cell2", Metadata: `{"zoneId":"use1-az1"}`, } - if err := topo.RegisterCellFromSpec( + c.Require().NoError(topo.RegisterCellFromSpec( context.Background(), store, recorder, owner, cellCfg, nil, topoRef, - ); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ), "unexpected error") cell, err := store.GetCell(context.Background(), "cell2") - if err != nil { - t.Fatalf("cell not found: %v", err) - } - if cell.Metadata != `{"zoneId":"use1-az1"}` { - t.Errorf("expected metadata copied verbatim, got %q", cell.Metadata) - } + c.Require().NoError(err, "cell not found") + c.Eq(`{"zoneId":"use1-az1"}`, cell.Metadata, "expected metadata copied verbatim, got") }) } @@ -742,6 +673,7 @@ func TestPruneDatabases(t *testing.T) { t.Run("removes stale database", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) ctx := context.Background() @@ -753,26 +685,23 @@ func TestPruneDatabases(t *testing.T) { multigresv1alpha1.DatabaseConfig{Name: multigresv1alpha1.DatabaseName(name)}, []string{"cell1"}, nil, "", ) - if err != nil { - t.Fatalf("registering %s: %v", name, err) - } + c.Require().NoError(err, "registering %s", name) } // Prune, keeping only db1. - if err := topo.PruneDatabases(ctx, store, recorder, owner, []string{"db1"}); err != nil { - t.Fatalf("prune: %v", err) - } + c.Require(). + NoError(topo.PruneDatabases(ctx, store, recorder, owner, []string{"db1"}), "prune") if _, err := store.GetDatabase(ctx, "db1"); err != nil { t.Error("db1 should still exist") } - if _, err := store.GetDatabase(ctx, "db2"); err == nil { - t.Error("db2 should have been pruned") - } + _, err := store.GetDatabase(ctx, "db2") + c.Error(err, "db2 should have been pruned") }) t.Run("no-op when all databases active", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) ctx := context.Background() @@ -782,13 +711,9 @@ func TestPruneDatabases(t *testing.T) { multigresv1alpha1.DatabaseConfig{Name: "db1"}, []string{"cell1"}, nil, "", ) - if err != nil { - t.Fatalf("registering: %v", err) - } + c.NoError(err, "registering") - if err := topo.PruneDatabases(ctx, store, recorder, owner, []string{"db1"}); err != nil { - t.Fatalf("prune: %v", err) - } + c.NoError(topo.PruneDatabases(ctx, store, recorder, owner, []string{"db1"}), "prune") if _, err := store.GetDatabase(ctx, "db1"); err != nil { t.Error("db1 should still exist") @@ -810,9 +735,7 @@ func TestPruneDatabases(t *testing.T) { owner, nil, ) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("skips if database delete returns NoNode", func(t *testing.T) { @@ -833,9 +756,7 @@ func TestPruneDatabases(t *testing.T) { owner, nil, ) - if err != nil { - t.Fatalf("expected nil as NoNode is skipped, got %v", err) - } + assert.NewAborting(t).NoError(err, "expected nil as NoNode is skipped, got") }) t.Run("returns error on deleting database", func(t *testing.T) { @@ -856,9 +777,7 @@ func TestPruneDatabases(t *testing.T) { owner, nil, ) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) } @@ -875,52 +794,46 @@ func TestPruneCells(t *testing.T) { t.Run("removes stale cell", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) store := newMemoryStore(t, "cell1", "cell2") recorder := record.NewFakeRecorder(10) ctx := context.Background() for _, c := range []string{"cell1", "cell2"} { - if err := topo.RegisterCellFromSpec( + ck.Require().NoError(topo.RegisterCellFromSpec( ctx, store, recorder, owner, multigresv1alpha1.CellConfig{Name: multigresv1alpha1.CellName(c)}, nil, topoRef, - ); err != nil { - t.Fatalf("registering %s: %v", c, err) - } + ), "registering %s", c) } // Prune, keeping only cell1. - if err := topo.PruneCells(ctx, store, recorder, owner, []string{"cell1"}); err != nil { - t.Fatalf("prune: %v", err) - } + ck.Require(). + NoError(topo.PruneCells(ctx, store, recorder, owner, []string{"cell1"}), "prune") if _, err := store.GetCell(ctx, "cell1"); err != nil { t.Error("cell1 should still exist") } - if _, err := store.GetCell(ctx, "cell2"); err == nil { - t.Error("cell2 should have been pruned") - } + _, err := store.GetCell(ctx, "cell2") + ck.Error(err, "cell2 should have been pruned") }) t.Run("no-op when all cells active", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) store := newMemoryStore(t, "cell1") recorder := record.NewFakeRecorder(10) ctx := context.Background() - if err := topo.RegisterCellFromSpec( + c.Require().NoError(topo.RegisterCellFromSpec( ctx, store, recorder, owner, multigresv1alpha1.CellConfig{Name: "cell1"}, nil, topoRef, - ); err != nil { - t.Fatalf("registering: %v", err) - } + ), "registering") - if err := topo.PruneCells(ctx, store, recorder, owner, []string{"cell1"}); err != nil { - t.Fatalf("prune: %v", err) - } + c.Require(). + NoError(topo.PruneCells(ctx, store, recorder, owner, []string{"cell1"}), "prune") - if _, err := store.GetCell(ctx, "cell1"); err != nil { - t.Error("cell1 should still exist") - } + _, err := store.GetCell(ctx, "cell1") + c.NoError(err, "cell1 should still exist") }) t.Run("returns error on getting cell names", func(t *testing.T) { @@ -932,9 +845,7 @@ func TestPruneCells(t *testing.T) { } err := topo.PruneCells(context.Background(), store, record.NewFakeRecorder(10), owner, nil) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) t.Run("skips if cell delete returns NoNode", func(t *testing.T) { @@ -949,9 +860,7 @@ func TestPruneCells(t *testing.T) { } err := topo.PruneCells(context.Background(), store, record.NewFakeRecorder(10), owner, nil) - if err != nil { - t.Fatalf("expected nil as NoNode is skipped, got %v", err) - } + assert.NewAborting(t).NoError(err, "expected nil as NoNode is skipped, got") }) t.Run("returns error on deleting cell", func(t *testing.T) { @@ -966,8 +875,6 @@ func TestPruneCells(t *testing.T) { } err := topo.PruneCells(context.Background(), store, record.NewFakeRecorder(10), owner, nil) - if err == nil { - t.Fatal("expected error, got nil") - } + assert.NewAborting(t).Error(err, "expected error, got nil") }) } diff --git a/pkg/gc/pvc/pvc_test.go b/pkg/gc/pvc/pvc_test.go index 0b199181..e851dcd5 100644 --- a/pkg/gc/pvc/pvc_test.go +++ b/pkg/gc/pvc/pvc_test.go @@ -7,7 +7,6 @@ import ( "time" "github.com/go-logr/logr/testr" - "github.com/stretchr/testify/require" corev1 "k8s.io/api/core/v1" apierrors "k8s.io/apimachinery/pkg/api/errors" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -18,6 +17,8 @@ import ( "github.com/multigres/multigres-operator/pkg/gc" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) const ( @@ -60,19 +61,20 @@ func run( opts.Retention = 30 * 24 * time.Hour } res, err := New(cl, testr.New(t), opts).Clean(context.Background()) - require.NoError(t, err) + assert.NewAborting(t).NoError(err) return res, cl } func parseNow(t *testing.T, s string) func() time.Time { t.Helper() ts, err := time.Parse(metadata.OrphanTimestampFormat, s) - require.NoError(t, err) + assert.NewAborting(t).NoError(err) return func() time.Time { return ts } } func TestClean_DeletesExpiredOnly(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) res, cl := run( t, @@ -87,13 +89,13 @@ func TestClean_DeletesExpiredOnly(t *testing.T) { ), ) - require.Equal(t, ObjectKind, res.Kind) - require.Equal(t, 2, res.Scanned) - require.Equal(t, 1, res.Deleted) - require.Equal(t, 1, res.Skipped) - require.Zero(t, res.Errors+res.Malformed+res.WouldDelete) + c.EqDeep(ObjectKind, res.Kind) + c.EqDeep(2, res.Scanned) + c.EqDeep(1, res.Deleted) + c.EqDeep(1, res.Skipped) + c.Zero(res.Errors + res.Malformed + res.WouldDelete) - require.True(t, apierrors.IsNotFound( + c.True(apierrors.IsNotFound( cl.Get( context.Background(), client.ObjectKey{Namespace: "ns1", Name: "old"}, @@ -104,28 +106,31 @@ func TestClean_DeletesExpiredOnly(t *testing.T) { func TestClean_DryRun(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) res, _ := run(t, gc.Options{DryRun: true}, interceptor.Funcs{}, pvc("old", orphanLabels(tsOld)), ) - require.Equal(t, 1, res.WouldDelete) - require.Zero(t, res.Deleted) + c.EqDeep(1, res.WouldDelete) + c.Zero(res.Deleted) } func TestClean_MalformedTimestamp(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) res, _ := run(t, gc.Options{}, interceptor.Funcs{}, pvc("bad", orphanLabels("not-a-timestamp")), ) - require.Equal(t, 1, res.Malformed) - require.Zero(t, res.Deleted) + c.EqDeep(1, res.Malformed) + c.Zero(res.Deleted) } func TestClean_NamespaceScope(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) a := pvc("a", orphanLabels(tsOld)) b := pvc("b", orphanLabels(tsOld)) @@ -133,12 +138,13 @@ func TestClean_NamespaceScope(t *testing.T) { res, _ := run(t, gc.Options{Namespace: "ns1"}, interceptor.Funcs{}, a, b) - require.Equal(t, 1, res.Scanned) - require.Equal(t, 1, res.Deleted) + c.EqDeep(1, res.Scanned) + c.EqDeep(1, res.Deleted) } func TestClean_DeleteErrorCounted(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) boom := errors.New("api server unavailable") res, _ := run(t, gc.Options{}, interceptor.Funcs{ @@ -147,8 +153,8 @@ func TestClean_DeleteErrorCounted(t *testing.T) { }, }, pvc("old", orphanLabels(tsOld))) - require.Equal(t, 1, res.Errors) - require.Zero(t, res.Deleted) + ck.EqDeep(1, res.Errors) + ck.Zero(res.Deleted) } func TestClean_NotFoundIsNotError(t *testing.T) { @@ -160,21 +166,22 @@ func TestClean_NotFoundIsNotError(t *testing.T) { }, }, pvc("old", orphanLabels(tsOld))) - require.Zero(t, res.Errors) + assert.NewAborting(t).Zero(res.Errors) } // Pinpoint regression: a PVC orphaned later in the day must NOT be deleted // until a full retention has elapsed from that exact timestamp. func TestClean_SameDayPrecision(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) res, _ := run(t, gc.Options{ Retention: 30 * 24 * time.Hour, Now: parseNow(t, "2026-05-15T12-00-00Z"), }, interceptor.Funcs{}, pvc("borderline", orphanLabels("2026-04-15T18-00-00Z"))) - require.Equal(t, 1, res.Skipped) - require.Zero(t, res.Deleted) + c.EqDeep(1, res.Skipped) + c.Zero(res.Deleted) } // Static check that *Cleankeeper satisfies gc.Cleankeeper. diff --git a/pkg/images/images_test.go b/pkg/images/images_test.go index 2d6fb994..d318cf35 100644 --- a/pkg/images/images_test.go +++ b/pkg/images/images_test.go @@ -4,68 +4,54 @@ import ( "testing" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestDefaultsFromEnv(t *testing.T) { t.Run("no overrides returns compiled defaults", func(t *testing.T) { + c := assert.NewCollecting(t) set, overrides := DefaultsFromEnv() - if set != CompiledDefaults() { - t.Errorf("expected compiled defaults, got %+v", set) - } - if len(overrides) != 0 { - t.Errorf("expected no overrides, got %v", overrides) - } + c.Eq(CompiledDefaults(), set, "expected compiled defaults, got") + c.Empty(overrides, "expected no overrides, got") }) t.Run("env override replaces single component", func(t *testing.T) { t.Setenv(EnvPostgresImage, "custom/pgctld:v9") t.Setenv(EnvMultigatewayImage, " custom/gateway:v9 ") + c := assert.NewCollecting(t) set, overrides := DefaultsFromEnv() - if set.Postgres != "custom/pgctld:v9" { - t.Errorf("postgres override not applied: %s", set.Postgres) - } - if set.Multigateway != "custom/gateway:v9" { - t.Errorf("multigateway override not trimmed/applied: %s", set.Multigateway) - } - if set.Multiadmin != CompiledDefaults().Multiadmin { - t.Errorf("multiadmin should keep compiled default, got %s", set.Multiadmin) - } - if len(overrides) != 2 { - t.Errorf("expected 2 active overrides, got %v", overrides) - } + c.Eq("custom/pgctld:v9", set.Postgres, "postgres override not applied") + c.Eq("custom/gateway:v9", set.Multigateway, "multigateway override not trimmed/applied") + c.Eq( + CompiledDefaults().Multiadmin, + set.Multiadmin, + "multiadmin should keep compiled default, got", + ) + c.Len(overrides, 2, "expected 2 active overrides, got") }) } func TestRevision(t *testing.T) { + c := assert.NewCollecting(t) base := CompiledDefaults() rev := Revision(base) - if len(rev) != 12 { - t.Fatalf("expected 12-char revision, got %q", rev) - } - if Revision(base) != rev { - t.Error("revision is not deterministic") - } + c.Require().Len(rev, 12, "expected 12-char revision, got") + c.Eq(rev, Revision(base), "revision is not deterministic") changed := base changed.Postgres = "other/pgctld:v1" - if Revision(changed) == rev { - t.Error("revision did not change when an image changed") - } + c.NotEq(rev, Revision(changed), "revision did not change when an image changed") } func TestIsComplete(t *testing.T) { - if !IsComplete(CompiledDefaults()) { - t.Error("compiled defaults must always form a complete set") - } + c := assert.NewCollecting(t) + c.True(IsComplete(CompiledDefaults()), "compiled defaults must always form a complete set") partial := CompiledDefaults() partial.Multipooler = "" - if IsComplete(partial) { - t.Error("a set with an empty component must not be complete") - } - if IsComplete(multigresv1alpha1.ComponentImages{}) { - t.Error("the zero set must not be complete") - } + c.False(IsComplete(partial), "a set with an empty component must not be complete") + c.False(IsComplete(multigresv1alpha1.ComponentImages{}), "the zero set must not be complete") } func TestComplete(t *testing.T) { @@ -74,19 +60,15 @@ func TestComplete(t *testing.T) { t.Run("fills unset fields", func(t *testing.T) { spec := multigresv1alpha1.ClusterImages{} Complete(&spec, defaults) - if spec.Postgres != defaults.Postgres || spec.Multigateway != defaults.Multigateway { - t.Errorf("unset fields not filled: %+v", spec) - } + assert.NewCollecting(t). + False(spec.Postgres != defaults.Postgres || spec.Multigateway != defaults.Multigateway, "unset fields not filled: %+v", spec) }) t.Run("explicit values win", func(t *testing.T) { + c := assert.NewCollecting(t) spec := multigresv1alpha1.ClusterImages{Postgres: "pinned/pgctld:v1"} Complete(&spec, defaults) - if spec.Postgres != "pinned/pgctld:v1" { - t.Errorf("explicit value overwritten: %s", spec.Postgres) - } - if spec.Multiorch != defaults.Multiorch { - t.Errorf("unset field not filled: %s", spec.Multiorch) - } + c.Eq("pinned/pgctld:v1", spec.Postgres, "explicit value overwritten") + c.Eq(defaults.Multiorch, spec.Multiorch, "unset field not filled") }) } diff --git a/pkg/monitoring/metrics_test.go b/pkg/monitoring/metrics_test.go index 071faa46..ce38ecf9 100644 --- a/pkg/monitoring/metrics_test.go +++ b/pkg/monitoring/metrics_test.go @@ -6,13 +6,13 @@ import ( "github.com/prometheus/client_golang/prometheus" dto "github.com/prometheus/client_model/go" + + "github.com/multigres/testkit/assert" ) func TestCollectorsRegistered(t *testing.T) { collectors := Collectors() - if len(collectors) == 0 { - t.Fatal("expected at least one collector, got 0") - } + assert.NewAborting(t).NotEmpty(collectors, "expected at least one collector, got 0") } func TestMetricNamingConvention(t *testing.T) { @@ -38,9 +38,8 @@ func TestMetricHelpNonEmpty(t *testing.T) { for desc := range ch { help := extractHelp(desc) - if help == "" { - t.Errorf("metric %q has empty help string", desc.String()) - } + assert.NewCollecting(t). + NotEq("", help, "metric %q has empty help string", desc.String()) } } } @@ -87,14 +86,8 @@ func TestGaugeLabels(t *testing.T) { descStr := desc.String() for _, label := range tt.wantLabels { - if !strings.Contains(descStr, label) { - t.Errorf( - "metric %s missing label %q in descriptor: %s", - tt.name, - label, - descStr, - ) - } + assert.NewCollecting(t). + StrContains(descStr, label, "metric %s missing label %q in descriptor", tt.name, label) } }) } diff --git a/pkg/monitoring/recorder_test.go b/pkg/monitoring/recorder_test.go index 36c33a10..ab16dc87 100644 --- a/pkg/monitoring/recorder_test.go +++ b/pkg/monitoring/recorder_test.go @@ -7,34 +7,32 @@ import ( "github.com/prometheus/client_golang/prometheus" dto "github.com/prometheus/client_model/go" + + "github.com/multigres/testkit/assert" ) func TestSetClusterInfo(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { clusterInfo.Reset() }) SetClusterInfo("test-cluster", "default", "Healthy", true) val := gaugeValue(t, clusterInfo, "test-cluster", "default", "Healthy", "true") - if val != 1 { - t.Errorf("expected clusterInfo gauge to be 1, got %f", val) - } + c.Eq(1, val, "expected clusterInfo gauge to be 1, got") // Phase change should clean up old label set, initialized stays true SetClusterInfo("test-cluster", "default", "Degraded", true) val = gaugeValue(t, clusterInfo, "test-cluster", "default", "Degraded", "true") - if val != 1 { - t.Errorf("expected clusterInfo gauge for Degraded to be 1, got %f", val) - } + c.Eq(1, val, "expected clusterInfo gauge for Degraded to be 1, got") // Old phase must have been cleaned up (value 0) oldVal := gaugeValue(t, clusterInfo, "test-cluster", "default", "Healthy", "true") - if oldVal != 0 { - t.Error("old phase label set should have been cleaned up") - } + c.Eq(0, oldVal, "old phase label set should have been cleaned up") } func TestSetClusterTopology(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { clusterCellsTotal.Reset() clusterShardsTotal.Reset() @@ -43,31 +41,25 @@ func TestSetClusterTopology(t *testing.T) { SetClusterTopology("test-cluster", "default", 3, 6) cells := gaugeValue(t, clusterCellsTotal, "test-cluster", "default") - if cells != 3 { - t.Errorf("expected cells=3, got %f", cells) - } + c.Eq(3, cells, "expected cells=3, got") shards := gaugeValue(t, clusterShardsTotal, "test-cluster", "default") - if shards != 6 { - t.Errorf("expected shards=6, got %f", shards) - } + c.Eq(6, shards, "expected shards=6, got") } func TestSetCellGatewayReplicas(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { cellGatewayReplicas.Reset() }) SetCellGatewayReplicas("cell-1", "default", 3, 2) desired := gaugeValue(t, cellGatewayReplicas, "cell-1", "default", "desired") - if desired != 3 { - t.Errorf("expected desired=3, got %f", desired) - } + c.Eq(3, desired, "expected desired=3, got") ready := gaugeValue(t, cellGatewayReplicas, "cell-1", "default", "ready") - if ready != 2 { - t.Errorf("expected ready=2, got %f", ready) - } + c.Eq(2, ready, "expected ready=2, got") } func TestSetShardPoolReplicas(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { shardPoolReplicas.Reset() }) SetShardPoolReplicas("test-cluster", "shard-1", "primary", "z1", "default", 3, 3) @@ -82,9 +74,7 @@ func TestSetShardPoolReplicas(t *testing.T) { "default", "desired", ) - if desired != 3 { - t.Errorf("expected desired=3, got %f", desired) - } + c.Eq(3, desired, "expected desired=3, got") ready := gaugeValue( t, shardPoolReplicas, @@ -95,27 +85,23 @@ func TestSetShardPoolReplicas(t *testing.T) { "default", "ready", ) - if ready != 3 { - t.Errorf("expected ready=3, got %f", ready) - } + c.Eq(3, ready, "expected ready=3, got") } func TestSetTopoServerReplicas(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { toposerverReplicas.Reset() }) SetTopoServerReplicas("topo-1", "default", 3, 1) desired := gaugeValue(t, toposerverReplicas, "topo-1", "default", "desired") - if desired != 3 { - t.Errorf("expected desired=3, got %f", desired) - } + c.Eq(3, desired, "expected desired=3, got") ready := gaugeValue(t, toposerverReplicas, "topo-1", "default", "ready") - if ready != 1 { - t.Errorf("expected ready=1, got %f", ready) - } + c.Eq(1, ready, "expected ready=1, got") } func TestRecordWebhookRequest(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { webhookRequestTotal.Reset() webhookRequestDuration.Reset() @@ -130,31 +116,24 @@ func TestRecordWebhookRequest(t *testing.T) { ) successVal := counterValue(t, webhookRequestTotal, "CREATE", "MultigresCluster", "success") - if successVal != 1 { - t.Errorf("expected success counter=1, got %f", successVal) - } + c.Eq(1, successVal, "expected success counter=1, got") errorVal := counterValue(t, webhookRequestTotal, "UPDATE", "MultigresCluster", "error") - if errorVal != 1 { - t.Errorf("expected error counter=1, got %f", errorVal) - } + c.Eq(1, errorVal, "expected error counter=1, got") } func TestSetPoolPodsDrifted(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { poolPodsDrifted.Reset() }) SetPoolPodsDrifted("cluster-1", "shard-1", "primary", "zone-a", "default", 3) val := gaugeValue(t, poolPodsDrifted, "cluster-1", "shard-1", "primary", "zone-a", "default") - if val != 3 { - t.Errorf("expected poolPodsDrifted gauge to be 3, got %f", val) - } + c.Eq(3, val, "expected poolPodsDrifted gauge to be 3, got") SetPoolPodsDrifted("cluster-1", "shard-1", "primary", "zone-a", "default", 0) val = gaugeValue(t, poolPodsDrifted, "cluster-1", "shard-1", "primary", "zone-a", "default") - if val != 0 { - t.Errorf("expected poolPodsDrifted gauge to be 0, got %f", val) - } + c.Eq(0, val, "expected poolPodsDrifted gauge to be 0, got") } func TestSetLastBackupAge(t *testing.T) { @@ -164,12 +143,11 @@ func TestSetLastBackupAge(t *testing.T) { SetLastBackupAge("cluster-1", "shard-1", "default", age) val := gaugeValue(t, lastBackupAgeSeconds, "cluster-1", "shard-1", "default") - if val != age.Seconds() { - t.Errorf("expected lastBackupAgeSeconds gauge to be %f, got %f", age.Seconds(), val) - } + assert.NewCollecting(t).Eq(age.Seconds(), val, "expected lastBackupAgeSeconds gauge to be") } func TestIncrementDrainOperations(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { drainOperationsTotal.Reset() }) IncrementDrainOperations("cluster-1", "shard-1", "success") @@ -177,16 +155,13 @@ func TestIncrementDrainOperations(t *testing.T) { IncrementDrainOperations("cluster-1", "shard-1", "error") successVal := counterValue(t, drainOperationsTotal, "cluster-1", "shard-1", "success") - if successVal != 2 { - t.Errorf("expected drain success counter=2, got %f", successVal) - } + c.Eq(2, successVal, "expected drain success counter=2, got") errorVal := counterValue(t, drainOperationsTotal, "cluster-1", "shard-1", "error") - if errorVal != 1 { - t.Errorf("expected drain error counter=1, got %f", errorVal) - } + c.Eq(1, errorVal, "expected drain error counter=1, got") } func TestSetRollingUpdateInProgress(t *testing.T) { + c := assert.NewCollecting(t) t.Cleanup(func() { rollingUpdateInProgress.Reset() }) SetRollingUpdateInProgress("cluster-1", "shard-1", "primary", "zone-a", "default", true) @@ -199,9 +174,7 @@ func TestSetRollingUpdateInProgress(t *testing.T) { "zone-a", "default", ) - if val != 1 { - t.Errorf("expected rollingUpdateInProgress=1 when true, got %f", val) - } + c.Eq(1, val, "expected rollingUpdateInProgress=1 when true, got") SetRollingUpdateInProgress("cluster-1", "shard-1", "primary", "zone-a", "default", false) val = gaugeValue( @@ -213,35 +186,27 @@ func TestSetRollingUpdateInProgress(t *testing.T) { "zone-a", "default", ) - if val != 0 { - t.Errorf("expected rollingUpdateInProgress=0 when false, got %f", val) - } + c.Eq(0, val, "expected rollingUpdateInProgress=0 when false, got") } // --- helpers --- func gaugeValue(t *testing.T, vec *prometheus.GaugeVec, labels ...string) float64 { t.Helper() + c := assert.NewAborting(t) g, err := vec.GetMetricWithLabelValues(labels...) - if err != nil { - t.Fatalf("GetMetricWithLabelValues(%v): %v", labels, err) - } + c.NoError(err, "GetMetricWithLabelValues(%v)", labels) m := &dto.Metric{} - if err := g.Write(m); err != nil { - t.Fatalf("Write: %v", err) - } + c.NoError(g.Write(m), "Write") return m.GetGauge().GetValue() } func counterValue(t *testing.T, vec *prometheus.CounterVec, labels ...string) float64 { t.Helper() + ck := assert.NewAborting(t) c, err := vec.GetMetricWithLabelValues(labels...) - if err != nil { - t.Fatalf("GetMetricWithLabelValues(%v): %v", labels, err) - } + ck.NoError(err, "GetMetricWithLabelValues(%v)", labels) m := &dto.Metric{} - if err := c.Write(m); err != nil { - t.Fatalf("Write: %v", err) - } + ck.NoError(c.Write(m), "Write") return m.GetCounter().GetValue() } diff --git a/pkg/monitoring/tracing_test.go b/pkg/monitoring/tracing_test.go index edbce9f1..95572c26 100644 --- a/pkg/monitoring/tracing_test.go +++ b/pkg/monitoring/tracing_test.go @@ -4,7 +4,6 @@ import ( "context" "errors" "strconv" - "strings" "testing" "time" @@ -18,9 +17,12 @@ import ( "go.opentelemetry.io/otel/sdk/trace/tracetest" "go.opentelemetry.io/otel/trace" "sigs.k8s.io/controller-runtime/pkg/log" + + "github.com/multigres/testkit/assert" ) func TestStartReconcileSpan(t *testing.T) { + c := assert.NewCollecting(t) exporter := tracetest.NewInMemoryExporter() tp := sdktrace.NewTracerProvider(sdktrace.WithSyncer(exporter)) t.Cleanup(func() { _ = tp.Shutdown(context.Background()) }) @@ -39,14 +41,10 @@ func TestStartReconcileSpan(t *testing.T) { span.End() spans := exporter.GetSpans() - if len(spans) != 1 { - t.Fatalf("expected 1 span, got %d", len(spans)) - } + c.Require().Len(spans, 1, "expected 1 span, got %d", len(spans)) s := spans[0] - if s.Name != "MultigresCluster.Reconcile" { - t.Errorf("span name = %q, want %q", s.Name, "MultigresCluster.Reconcile") - } + c.Eq("MultigresCluster.Reconcile", s.Name, "span name") wantAttrs := map[string]string{ "k8s.resource.name": "my-cluster", @@ -58,23 +56,24 @@ func TestStartReconcileSpan(t *testing.T) { for _, attr := range s.Attributes { if string(attr.Key) == key { found = true - if attr.Value.AsString() != want { - t.Errorf("attribute %q = %q, want %q", key, attr.Value.AsString(), want) - } + c.Eq( + want, + attr.Value.AsString(), + "attribute %q = %q, want", + key, + attr.Value.AsString(), + ) } } - if !found { - t.Errorf("attribute %q not found on span", key) - } + c.True(found, "attribute %q not found on span", key) } // Verify the context carries the span. - if ctx == context.Background() { - t.Error("expected context to carry span") - } + c.False(ctx == context.Background(), "expected context to carry span") } func TestStartChildSpan(t *testing.T) { + c := assert.NewCollecting(t) exporter := tracetest.NewInMemoryExporter() tp := sdktrace.NewTracerProvider(sdktrace.WithSyncer(exporter)) t.Cleanup(func() { _ = tp.Shutdown(context.Background()) }) @@ -88,23 +87,13 @@ func TestStartChildSpan(t *testing.T) { parent.End() spans := exporter.GetSpans() - if len(spans) != 2 { - t.Fatalf("expected 2 spans, got %d", len(spans)) - } + c.Require().Len(spans, 2, "expected 2 spans, got %d", len(spans)) // Child span should reference the parent's span context. childSpan := spans[0] parentSpan := spans[1] - if childSpan.Parent.SpanID() != parentSpan.SpanContext.SpanID() { - t.Errorf( - "child parent span ID = %s, want %s", - childSpan.Parent.SpanID(), - parentSpan.SpanContext.SpanID(), - ) - } - if childSpan.Name != "ChildOperation" { - t.Errorf("child span name = %q, want %q", childSpan.Name, "ChildOperation") - } + c.Eq(parentSpan.SpanContext.SpanID(), childSpan.Parent.SpanID(), "child parent span ID") + c.Eq("ChildOperation", childSpan.Name, "child span name") } func TestRecordSpanError(t *testing.T) { @@ -115,6 +104,7 @@ func TestRecordSpanError(t *testing.T) { Tracer = tp.Tracer(tracerName) t.Run("records error on span", func(t *testing.T) { + c := assert.NewCollecting(t) exporter.Reset() _, span := StartReconcileSpan(context.Background(), "Op", "n", "ns", "K") testErr := errors.New("something failed") @@ -122,21 +112,11 @@ func TestRecordSpanError(t *testing.T) { span.End() spans := exporter.GetSpans() - if len(spans) != 1 { - t.Fatalf("expected 1 span, got %d", len(spans)) - } + c.Require().Len(spans, 1, "expected 1 span, got %d", len(spans)) s := spans[0] - if s.Status.Code != codes.Error { - t.Errorf("span status = %v, want Error", s.Status.Code) - } - if s.Status.Description != "something failed" { - t.Errorf( - "span status description = %q, want %q", - s.Status.Description, - "something failed", - ) - } + c.Eq(codes.Error, s.Status.Code, "span status") + c.Eq("something failed", s.Status.Description, "span status description") // Check that an error event was recorded. foundErrorEvent := false @@ -151,38 +131,30 @@ func TestRecordSpanError(t *testing.T) { } } } - if !foundErrorEvent { - t.Error("expected an exception event on the span") - } + c.True(foundErrorEvent, "expected an exception event on the span") }) t.Run("nil error is no-op", func(t *testing.T) { + c := assert.NewCollecting(t) exporter.Reset() _, span := StartReconcileSpan(context.Background(), "Op", "n", "ns", "K") RecordSpanError(span, nil) span.End() spans := exporter.GetSpans() - if len(spans) != 1 { - t.Fatalf("expected 1 span, got %d", len(spans)) - } - if spans[0].Status.Code == codes.Error { - t.Error("nil error should not set error status") - } + c.Require().Len(spans, 1, "expected 1 span, got %d", len(spans)) + c.NotEq(codes.Error, spans[0].Status.Code, "nil error should not set error status") }) } func TestInitTracing_NoopWhenEndpointUnset(t *testing.T) { // Ensure OTEL_EXPORTER_OTLP_ENDPOINT is unset. t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "") + c := assert.NewAborting(t) shutdown, err := InitTracing(context.Background(), "test-svc", "v0.0.1") - if err != nil { - t.Fatalf("InitTracing() returned error: %v", err) - } - if err := shutdown(context.Background()); err != nil { - t.Fatalf("shutdown() returned error: %v", err) - } + c.NoError(err, "InitTracing() returned error") + c.NoError(shutdown(context.Background()), "shutdown() returned error") } func TestInjectAndExtractTraceContext(t *testing.T) { @@ -194,6 +166,7 @@ func TestInjectAndExtractTraceContext(t *testing.T) { otel.SetTextMapPropagator(propagation.TraceContext{}) t.Run("round-trips trace context through annotations", func(t *testing.T) { + c := assert.NewCollecting(t) ctx, span := Tracer.Start(context.Background(), "webhook") originalTraceID := span.SpanContext().TraceID() @@ -204,18 +177,13 @@ func TestInjectAndExtractTraceContext(t *testing.T) { if _, ok := annotations[annotationTraceparent]; !ok { t.Fatal("expected traceparent annotation to be set") } - if _, ok := annotations[annotationTraceparentTS]; !ok { - t.Fatal("expected traceparent-ts annotation to be set") - } + _, ok := annotations[annotationTraceparentTS] + c.Require().True(ok, "expected traceparent-ts annotation to be set") parentCtx, isStale := ExtractTraceContext(annotations) - if isStale { - t.Error("fresh annotation should not be stale") - } + c.False(isStale, "fresh annotation should not be stale") sc := trace.SpanFromContext(parentCtx).SpanContext() - if sc.TraceID() != originalTraceID { - t.Errorf("extracted trace ID = %s, want %s", sc.TraceID(), originalTraceID) - } + c.Eq(originalTraceID, sc.TraceID(), "extracted trace ID") }) t.Run("stale annotation", func(t *testing.T) { @@ -230,20 +198,15 @@ func TestInjectAndExtractTraceContext(t *testing.T) { annotations[annotationTraceparentTS] = strconv.FormatInt(staleTS, 10) _, isStale := ExtractTraceContext(annotations) - if !isStale { - t.Error("expected stale annotation to be detected") - } + assert.NewCollecting(t).True(isStale, "expected stale annotation to be detected") }) t.Run("missing annotation returns background context", func(t *testing.T) { + c := assert.NewCollecting(t) parentCtx, isStale := ExtractTraceContext(map[string]string{}) - if isStale { - t.Error("empty annotations should not be stale") - } + c.False(isStale, "empty annotations should not be stale") sc := trace.SpanFromContext(parentCtx).SpanContext() - if sc.IsValid() { - t.Error("expected invalid span context from empty annotations") - } + c.False(sc.IsValid(), "expected invalid span context from empty annotations") }) t.Run("missing timestamp treated as stale", func(t *testing.T) { @@ -255,9 +218,7 @@ func TestInjectAndExtractTraceContext(t *testing.T) { delete(annotations, annotationTraceparentTS) _, isStale := ExtractTraceContext(annotations) - if !isStale { - t.Error("missing timestamp should be treated as stale") - } + assert.NewCollecting(t).True(isStale, "missing timestamp should be treated as stale") }) } @@ -280,9 +241,8 @@ func TestEnrichLoggerWithTrace(t *testing.T) { logger := log.FromContext(enrichedCtx) // We can't easily inspect logr values, but we can verify the function // doesn't panic and returns a different context. - if enrichedCtx == ctx { - t.Error("expected enriched context to differ from original") - } + assert.NewCollecting(t). + False(enrichedCtx == ctx, "expected enriched context to differ from original") _ = logger }) @@ -290,9 +250,7 @@ func TestEnrichLoggerWithTrace(t *testing.T) { ctx := logr.NewContext(context.Background(), logr.Discard()) result := EnrichLoggerWithTrace(ctx) // With no valid span, the context should be returned unchanged. - if result != ctx { - t.Error("expected unchanged context for invalid span") - } + assert.NewCollecting(t).False(result != ctx, "expected unchanged context for invalid span") }) } @@ -303,17 +261,12 @@ func TestInitTracing_WithEndpoint(t *testing.T) { // tracer re-acquisition. t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://localhost:4318") t.Setenv("OTEL_TRACES_EXPORTER", "none") + c := assert.NewAborting(t) shutdown, err := InitTracing(context.Background(), "test-svc", "v0.0.1") - if err != nil { - t.Fatalf("InitTracing() returned error: %v", err) - } - if shutdown == nil { - t.Fatal("expected non-nil shutdown function") - } - if err := shutdown(context.Background()); err != nil { - t.Fatalf("shutdown() returned error: %v", err) - } + c.NoError(err, "InitTracing() returned error") + c.NotNil(shutdown, "expected non-nil shutdown function") + c.NoError(shutdown(context.Background()), "shutdown() returned error") } func TestInitTracing_ExporterError(t *testing.T) { @@ -321,23 +274,19 @@ func TestInitTracing_ExporterError(t *testing.T) { t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://localhost:4318") // Set invalid exporter type to trigger error in autoexport.NewSpanExporter t.Setenv("OTEL_TRACES_EXPORTER", "invalid-exporter-type") + c := assert.NewCollecting(t) // InitTracing should fail shutdown, err := InitTracing(context.Background(), "test-svc", "v0.0.1") - if err == nil { - t.Fatal("InitTracing() should have failed with invalid exporter type") - } - if shutdown != nil { - t.Fatal("shutdown function should be nil on error") - } - if !strings.Contains(err.Error(), "creating OTLP exporter") { - t.Errorf("unexpected error message: %v", err) - } + c.Require().Error(err, "InitTracing() should have failed with invalid exporter type") + c.Require().Nil(shutdown, "shutdown function should be nil on error") + c.StrContains(err.Error(), "creating OTLP exporter", "unexpected error message: %v", err) } func TestInitTracing_ResourceError(t *testing.T) { t.Setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://localhost:4318") t.Setenv("OTEL_TRACES_EXPORTER", "none") + c := assert.NewCollecting(t) injectedErr := errors.New("synthetic resource failure") original := newResource @@ -347,18 +296,10 @@ func TestInitTracing_ResourceError(t *testing.T) { t.Cleanup(func() { newResource = original }) shutdown, err := InitTracing(context.Background(), "test-svc", "v0.0.1") - if err == nil { - t.Fatal("expected error from InitTracing when resource creation fails") - } - if !strings.Contains(err.Error(), "creating OTel resource") { - t.Errorf("unexpected error message: %v", err) - } - if !errors.Is(err, injectedErr) { - t.Errorf("expected wrapped injectedErr, got: %v", err) - } - if shutdown != nil { - t.Fatal("shutdown function should be nil on error") - } + c.Require().Error(err, "expected error from InitTracing when resource creation fails") + c.StrContains(err.Error(), "creating OTel resource", "unexpected error message: %v", err) + c.ErrorIs(err, injectedErr, "expected wrapped injectedErr, got") + c.Require().Nil(shutdown, "shutdown function should be nil on error") } func TestInjectTraceContext_TracestateRename(t *testing.T) { @@ -386,9 +327,8 @@ func TestInjectTraceContext_TracestateRename(t *testing.T) { if _, ok := annotations["tracestate"]; ok { t.Error("standard 'tracestate' key should be renamed") } - if _, ok := annotations["multigres.com/tracestate"]; !ok { - t.Error("expected 'multigres.com/tracestate' annotation to be set") - } + _, ok := annotations["multigres.com/tracestate"] + assert.NewCollecting(t).True(ok, "expected 'multigres.com/tracestate' annotation to be set") } func TestExtractTraceContext_InvalidTimestamp(t *testing.T) { @@ -408,21 +348,18 @@ func TestExtractTraceContext_InvalidTimestamp(t *testing.T) { annotations[annotationTraceparentTS] = "not-a-number" _, isStale := ExtractTraceContext(annotations) - if !isStale { - t.Error("invalid timestamp should be treated as stale") - } + assert.NewCollecting(t).True(isStale, "invalid timestamp should be treated as stale") } func TestInjectTraceContext_InvalidSpanContext(t *testing.T) { annotations := make(map[string]string) InjectTraceContext(context.Background(), annotations) - if len(annotations) != 0 { - t.Errorf("expected no annotations for invalid span, got %v", annotations) - } + assert.NewCollecting(t).Empty(annotations, "expected no annotations for invalid span, got") } func TestExtractTraceContext_WithTracestate(t *testing.T) { + c := assert.NewAborting(t) exporter := tracetest.NewInMemoryExporter() tp := sdktrace.NewTracerProvider(sdktrace.WithSyncer(exporter)) t.Cleanup(func() { _ = tp.Shutdown(context.Background()) }) @@ -441,14 +378,11 @@ func TestExtractTraceContext_WithTracestate(t *testing.T) { span.End() // Verify the tracestate was injected under our custom key. - if _, ok := annotations["multigres.com/tracestate"]; !ok { - t.Fatal("expected multigres.com/tracestate annotation") - } + _, ok := annotations["multigres.com/tracestate"] + c.True(ok, "expected multigres.com/tracestate annotation") // Now extract and verify the tracestate is restored. extractedCtx, _ := ExtractTraceContext(annotations) sc := trace.SpanFromContext(extractedCtx).SpanContext() - if !sc.IsValid() { - t.Fatal("expected valid span context after extraction") - } + c.True(sc.IsValid(), "expected valid span context after extraction") } diff --git a/pkg/postgresconfig/classify_test.go b/pkg/postgresconfig/classify_test.go index e5d1c7c0..d93087b1 100644 --- a/pkg/postgresconfig/classify_test.go +++ b/pkg/postgresconfig/classify_test.go @@ -1,6 +1,10 @@ package postgresconfig -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) func TestRequiresRestart(t *testing.T) { tests := []struct { @@ -41,9 +45,8 @@ func TestRequiresRestart(t *testing.T) { {"auto_explain.log_min_duration", true}, } for _, tt := range tests { - if got := RequiresRestart(tt.name); got != tt.want { - t.Errorf("RequiresRestart(%q) = %v, want %v", tt.name, got, tt.want) - } + got := RequiresRestart(tt.name) + assert.NewCollecting(t).Eq(tt.want, got, "RequiresRestart(%q) = %v, want", tt.name, got) } } @@ -68,9 +71,8 @@ func TestRequiresRestartKnownRestartKeys(t *testing.T) { "cron.log_statement", // pg_cron PGC_POSTMASTER; namespaced → conservative restart } for _, k := range restartKeys { - if !RequiresRestart(k) { - t.Errorf("RequiresRestart(%q) = false, want true (restart-required)", k) - } + assert.NewCollecting(t). + True(RequiresRestart(k), "RequiresRestart(%q) = false, want true (restart-required)", k) } } @@ -79,6 +81,7 @@ func TestRequiresRestartKnownRestartKeys(t *testing.T) { // empty context, everything would (conservatively) require a restart, silently // disabling the reload path. Assert a healthy split exists. func TestRequiresRestartMatchesCatalogContext(t *testing.T) { + c := assert.NewCollecting(t) var reloadable, total int for name := range catalog { total++ @@ -86,16 +89,14 @@ func TestRequiresRestartMatchesCatalogContext(t *testing.T) { reloadable++ } } - if total < 300 { - t.Fatalf("catalog has %d entries, expected the full PG17 set", total) - } + c.Require().GreaterOrEqual(300, total, "catalog has") // The majority of PG17 GUCs are reload-safe (sighup/user/superuser); if the // context column were missing this would be ~0. - if reloadable < 200 { - t.Errorf( - "only %d/%d params classified reloadable; context column likely missing from catalog", - reloadable, - total, - ) - } + c.GreaterOrEqual( + 200, + reloadable, + "only %d/%d params classified reloadable; context column likely missing from catalog", + reloadable, + total, + ) } diff --git a/pkg/postgresconfig/hash_test.go b/pkg/postgresconfig/hash_test.go index 77293616..b2c41155 100644 --- a/pkg/postgresconfig/hash_test.go +++ b/pkg/postgresconfig/hash_test.go @@ -1,6 +1,10 @@ package postgresconfig -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) // baseConf is a minimal rendered postgresql.conf covering a restart param // (shared_buffers, postmaster), a couple of reload params (work_mem/user, @@ -12,6 +16,7 @@ max_wal_size = '1GB' # sighup ` func TestSplitHashesReloadOnlyChangeLeavesRestartHash(t *testing.T) { + c := assert.NewCollecting(t) _, base := StampAndSplit(baseConf) // Change only work_mem (a reload-safe/user param). @@ -21,19 +26,16 @@ work_mem = '8MB' max_wal_size = '1GB' # sighup `) - if changed.RestartHash != base.RestartHash { - t.Errorf( - "restart-hash moved on a reload-only change: %s -> %s", - base.RestartHash, - changed.RestartHash, - ) - } - if changed.ReloadHash == base.ReloadHash { - t.Errorf("reload-hash did not move on a work_mem change (still %s)", base.ReloadHash) - } + c.Eq(base.RestartHash, changed.RestartHash, "restart-hash moved on a reload-only change") + c.NotEq( + base.ReloadHash, + changed.ReloadHash, + "reload-hash did not move on a work_mem change (still", + ) } func TestSplitHashesRestartChangeMovesRestartHash(t *testing.T) { + c := assert.NewCollecting(t) _, base := StampAndSplit(baseConf) // Change shared_buffers (a postmaster/restart param). @@ -43,22 +45,16 @@ work_mem = '4MB' max_wal_size = '1GB' # sighup `) - if changed.RestartHash == base.RestartHash { - t.Errorf( - "restart-hash did not move on a shared_buffers change (still %s)", - base.RestartHash, - ) - } - if changed.ReloadHash != base.ReloadHash { - t.Errorf( - "reload-hash moved on a restart-only change: %s -> %s", - base.ReloadHash, - changed.ReloadHash, - ) - } + c.NotEq( + base.RestartHash, + changed.RestartHash, + "restart-hash did not move on a shared_buffers change (still", + ) + c.Eq(base.ReloadHash, changed.ReloadHash, "reload-hash moved on a restart-only change") } func TestSplitHashesCosmeticEditsMoveNeither(t *testing.T) { + c := assert.NewCollecting(t) _, base := StampAndSplit(baseConf) // Reordered, differently commented, extra blank lines, different inline @@ -72,20 +68,8 @@ work_mem = '4MB' shared_buffers = '128MB' `) - if cosmetic.RestartHash != base.RestartHash { - t.Errorf( - "restart-hash moved on a cosmetic-only edit: %s -> %s", - base.RestartHash, - cosmetic.RestartHash, - ) - } - if cosmetic.ReloadHash != base.ReloadHash { - t.Errorf( - "reload-hash moved on a cosmetic-only edit: %s -> %s", - base.ReloadHash, - cosmetic.ReloadHash, - ) - } + c.Eq(base.RestartHash, cosmetic.RestartHash, "restart-hash moved on a cosmetic-only edit") + c.Eq(base.ReloadHash, cosmetic.ReloadHash, "reload-hash moved on a cosmetic-only edit") } func TestSplitHashesLastWins(t *testing.T) { @@ -96,9 +80,7 @@ work_mem = '8MB' `) _, single := StampAndSplit(`work_mem = '8MB' `) - if dup.ReloadHash != single.ReloadHash { - t.Errorf("last-wins not honored: dup=%s single=%s", dup.ReloadHash, single.ReloadHash) - } + assert.NewCollecting(t).Eq(single.ReloadHash, dup.ReloadHash, "last-wins not honored: dup") } // TestSplitHashesValueFormatting documents how the reload-hash treats value @@ -132,9 +114,9 @@ func TestSplitHashesValueFormatting(t *testing.T) { // Internal whitespace: token-based hashing treats '4 MB' as distinct from // '4MB'. This documents the known (harmless) redundant-reload behavior. - if s, q := reloadHashOf("work_mem = '4 MB'\n"), reloadHashOf("work_mem = '4MB'\n"); s == q { - t.Errorf("expected '4 MB' to hash differently from '4MB' under token-based comparison") - } + s, q := reloadHashOf("work_mem = '4 MB'\n"), reloadHashOf("work_mem = '4MB'\n") + assert.NewCollecting(t). + NotEq(q, s, "expected '4 MB' to hash differently from '4MB' under token-based comparison") } func TestStripInlineComment(t *testing.T) { @@ -162,9 +144,8 @@ func TestStripInlineComment(t *testing.T) { {"'just a test ''", "'just a test ''"}, } for _, tt := range tests { - if got := stripInlineComment(tt.in); got != tt.want { - t.Errorf("stripInlineComment(%q) = %q, want %q", tt.in, got, tt.want) - } + got := stripInlineComment(tt.in) + assert.NewCollecting(t).Eq(tt.want, got, "stripInlineComment(%q) = %q, want", tt.in, got) } } @@ -183,12 +164,12 @@ log_line_prefix = '%h %m [%p] ' # ok, for contrast // Deterministic: the same malformed input always hashes the same way. r1, s1 := StampAndSplit(malformed) r2, s2 := StampAndSplit(malformed) - if s1.ReloadHash != s2.ReloadHash || s1.RestartHash != s2.RestartHash || r1 != r2 { - t.Errorf("split of malformed config is not deterministic") - } + assert.NewCollecting(t). + False(s1.ReloadHash != s2.ReloadHash || s1.RestartHash != s2.RestartHash || r1 != r2, "split of malformed config is not deterministic") } func TestReloadSettings(t *testing.T) { + c := assert.NewCollecting(t) rendered := `# rendered shared_buffers = '128MB' # postmaster → restart, excluded work_mem = '32MB' # user → reload @@ -199,19 +180,12 @@ cron.database_name = 'postgres' # namespaced → restart, excluded _, split := StampAndSplit(rendered) got := split.ReloadSettings - if got["work_mem"] != "32MB" { - t.Errorf("work_mem = %q, want 32MB (unquoted)", got["work_mem"]) - } - if got["max_wal_size"] != "1024MB" { - t.Errorf("max_wal_size = %q, want 1024MB", got["max_wal_size"]) - } - if got["log_line_prefix"] != "%h %m [%p] " { - t.Errorf("log_line_prefix = %q, want unquoted verbatim", got["log_line_prefix"]) - } + c.Eq("32MB", got["work_mem"], "work_mem") + c.Eq("1024MB", got["max_wal_size"], "max_wal_size") + c.Eq("%h %m [%p] ", got["log_line_prefix"], "log_line_prefix") if _, ok := got["shared_buffers"]; ok { t.Error("shared_buffers (postmaster) must be excluded from reload settings") } - if _, ok := got["cron.database_name"]; ok { - t.Error("cron.database_name (namespaced→restart) must be excluded") - } + _, ok := got["cron.database_name"] + c.False(ok, "cron.database_name (namespaced→restart) must be excluded") } diff --git a/pkg/postgresconfig/reload_marker_test.go b/pkg/postgresconfig/reload_marker_test.go index 77008895..1d6f8ff6 100644 --- a/pkg/postgresconfig/reload_marker_test.go +++ b/pkg/postgresconfig/reload_marker_test.go @@ -1,11 +1,13 @@ package postgresconfig import ( - "strings" "testing" + + "github.com/multigres/testkit/assert" ) func TestStampAndSplitStampsMarker(t *testing.T) { + c := assert.NewCollecting(t) const conf = `work_mem = '4MB' max_wal_size = '1GB' ` @@ -13,39 +15,24 @@ max_wal_size = '1GB' // The marker line is appended to the returned file with value == reload-hash. wantLine := ReloadMarkerGUC + " = '" + split.ReloadHash + "'" - if !strings.Contains(rendered, wantLine) { - t.Errorf("rendered config missing marker line %q:\n%s", wantLine, rendered) - } + c.StrContains(rendered, wantLine, "rendered config missing marker line") // The marker is present in the expected-settings map, unquoted, == reload-hash. - if got := split.ReloadSettings[ReloadMarkerGUC]; got != split.ReloadHash { - t.Errorf( - "ReloadSettings[%q] = %q, want reload-hash %q", - ReloadMarkerGUC, - got, - split.ReloadHash, - ) - } + got := split.ReloadSettings[ReloadMarkerGUC] + c.Eq(split.ReloadHash, got, "ReloadSettings[%q] = %q, want reload-hash", ReloadMarkerGUC, got) // The reload-hash covers the real settings only, not the marker: hashing the // reload settings with the marker removed must reproduce the reload-hash. delete(split.ReloadSettings, ReloadMarkerGUC) - if h := hashSettings(split.ReloadSettings); h != split.ReloadHash { - t.Errorf( - "reload-hash includes the marker: %s over settings-without-marker != %s", - h, - split.ReloadHash, - ) - } + c.Eq(split.ReloadHash, hashSettings(split.ReloadSettings), "reload-hash includes the marker") } func TestStampAndSplitMarkerOnAllRestartConfig(t *testing.T) { // A config with only restart-only settings still gets a marker, so even an // all-restart render carries a version marker in its reload settings. _, split := StampAndSplit("shared_buffers = '128MB'\n") - if split.ReloadSettings[ReloadMarkerGUC] != split.ReloadHash { - t.Errorf("marker not stamped for an all-restart config: %v", split.ReloadSettings) - } + assert.NewCollecting(t). + Eq(split.ReloadHash, split.ReloadSettings[ReloadMarkerGUC], "marker not stamped for an all-restart config: %v", split.ReloadSettings) } // TestReloadMarkerDetectsRemoval is the core guard for the removal-only fix. When @@ -57,6 +44,7 @@ func TestStampAndSplitMarkerOnAllRestartConfig(t *testing.T) { // (carrying the previous marker) fails the gate and the reload is retried until // the kubelet syncs. func TestReloadMarkerDetectsRemoval(t *testing.T) { + c := assert.NewCollecting(t) const withParam = `work_mem = '4MB' random_page_cost = '1.1' ` @@ -69,22 +57,15 @@ random_page_cost = '1.1' // Neither config has a restart-only setting, so the restart-hash is unchanged // (this stays a reload, not a pod recreation). - if before.RestartHash != after.RestartHash { - t.Errorf("restart-hash moved on a reload-only removal: %s -> %s", - before.RestartHash, after.RestartHash) - } + c.Eq(after.RestartHash, before.RestartHash, "restart-hash moved on a reload-only removal") // The removal changed the reload-safe partition, so the reload-hash — and thus // the marker value — must move. - if before.ReloadHash == after.ReloadHash { - t.Fatalf("reload-hash did not move when a reload-safe setting was removed (still %s)", - before.ReloadHash) - } + c.Require(). + NotEq(after.ReloadHash, before.ReloadHash, "reload-hash did not move when a reload-safe setting was removed (still") beforeMarker := before.ReloadSettings[ReloadMarkerGUC] afterMarker := after.ReloadSettings[ReloadMarkerGUC] - if beforeMarker == afterMarker { - t.Fatalf("marker did not move on removal: %q", afterMarker) - } + c.Require().NotEq(afterMarker, beforeMarker, "marker did not move on removal") // The crux: every non-marker expected setting of the post-removal config is // still satisfied by the pre-removal (stale) file — so without the marker the @@ -94,9 +75,7 @@ random_page_cost = '1.1' for name, want := range after.ReloadSettings { staleVal, present := staleFile[name] if name == ReloadMarkerGUC { - if present { - t.Errorf("stale file unexpectedly already carries the marker") - } + c.False(present, "stale file unexpectedly already carries the marker") continue // the marker is what the stale file cannot satisfy } if !present || unquoteValue(staleVal) != want { diff --git a/pkg/postgresconfig/render_precedence_test.go b/pkg/postgresconfig/render_precedence_test.go index e2ade030..043a501c 100644 --- a/pkg/postgresconfig/render_precedence_test.go +++ b/pkg/postgresconfig/render_precedence_test.go @@ -1,6 +1,10 @@ package postgresconfig -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) // TestBaselineWinsOverRefForResourceDerivedKeys reproduces a precedence problem // in Render(): the operator's own resource-derived baseline must not be @@ -23,6 +27,7 @@ import "testing" // wins). After swapping the order to ref -> baseline -> inline it must be "192MB" // (baseline wins), since no inline override was given. func TestBaselineWinsOverRefForResourceDerivedKeys(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() // effective_cache_size baseline default is "192MB" // A ref value clearly different from the baseline so a precedence bug is @@ -31,17 +36,9 @@ func TestBaselineWinsOverRefForResourceDerivedKeys(t *testing.T) { const refContent = "effective_cache_size = '999MB'" rendered, err := Render(cfg, refContent, nil) - if err != nil { - t.Fatalf("Render: %v", err) - } + c.Require().NoError(err, "Render") _, split := StampAndSplit(rendered) got := split.ReloadSettings["effective_cache_size"] - if want := "192MB"; got != want { - t.Errorf( - "effective_cache_size = %q, want %q: the operator's resource-derived baseline must win over the deprecated PostgresConfigRef (only inline spec.postgresConfig should override it)", - got, - want, - ) - } + c.Eq("192MB", got, "effective_cache_size") } diff --git a/pkg/postgresconfig/render_test.go b/pkg/postgresconfig/render_test.go index 50118645..e83f0077 100644 --- a/pkg/postgresconfig/render_test.go +++ b/pkg/postgresconfig/render_test.go @@ -3,14 +3,15 @@ package postgresconfig import ( "strings" "testing" + + "github.com/multigres/testkit/assert" ) func TestRender(t *testing.T) { t.Run("renders the baseline from Config", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := Render(Defaults(), "", nil) - if err != nil { - t.Fatalf("Render() error = %v", err) - } + c.Require().NoError(err, "Render() error =") // A few representative baseline lines must be present with default values. for _, want := range []string{ "max_connections = 60", @@ -19,114 +20,98 @@ func TestRender(t *testing.T) { "wal_level = logical", "cluster_name = 'default'", } { - if !strings.Contains(got, want) { - t.Errorf("rendered baseline missing %q, got:\n%s", want, got) - } + c.StrContains(got, want, "rendered baseline missing") } }) t.Run("Config values flow into the template", func(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() cfg.SharedBuffers = "2GB" cfg.MaxConnections = 200 got, err := Render(cfg, "", nil) - if err != nil { - t.Fatalf("Render() error = %v", err) - } - if !strings.Contains(got, "shared_buffers = 2GB") { - t.Errorf("shared_buffers override missing, got:\n%s", got) - } - if !strings.Contains(got, "max_connections = 200") { - t.Errorf("max_connections override missing, got:\n%s", got) - } + c.Require().NoError(err, "Render() error =") + c.StrContains(got, "shared_buffers = 2GB", "shared_buffers override missing, got:\n") + c.StrContains(got, "max_connections = 200", "max_connections override missing, got:\n") }) t.Run("ref content emitted verbatim before the baseline", func(t *testing.T) { + c := assert.NewCollecting(t) ref := "shared_buffers = '8GB'\n# a comment" got, err := Render(Defaults(), ref, nil) - if err != nil { - t.Fatalf("Render() error = %v", err) - } - if !strings.Contains(got, ref) { - t.Errorf("ref content not emitted verbatim, got:\n%s", got) - } + c.Require().NoError(err, "Render() error =") + c.StrContains(got, ref, "ref content not emitted verbatim, got:\n") // Ref must come BEFORE the baseline so the operator's resource-derived // baseline wins last-write-wins: here the baseline's shared_buffers = 64MB // must override the ref's 8GB. The deprecated ref may not override the // operator's sizing math — only inline spec.postgresConfig can. - if strings.Index(got, ref) > strings.Index(got, "shared_buffers = 64MB") { - t.Errorf("ref content should precede the baseline, got:\n%s", got) - } + c.LessOrEqual( + strings.Index(got, "shared_buffers = 64MB"), + strings.Index(got, ref), + "ref content should precede the baseline, got:\n%s", + got, + ) }) t.Run("inline map appended last as sorted quoted lines", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := Render(Defaults(), "ref = 'x'", map[string]string{ "work_mem": "16MB", "max_connections": "200", }) - if err != nil { - t.Fatalf("Render() error = %v", err) - } + c.Require().NoError(err, "Render() error =") wantMax := "max_connections = '200'" wantWork := "work_mem = '16MB'" - if !strings.Contains(got, wantMax) || !strings.Contains(got, wantWork) { - t.Fatalf("inline lines missing, got:\n%s", got) - } + c.Require(). + False(!strings.Contains(got, wantMax) || !strings.Contains(got, wantWork), "inline lines missing, got:\n%s", got) // Sorted keys, and the whole map block comes after the ref block. - if strings.Index(got, wantMax) > strings.Index(got, wantWork) { - t.Errorf("inline keys not sorted, got:\n%s", got) - } - if strings.Index(got, wantMax) < strings.Index(got, "ref = 'x'") { - t.Errorf("inline map should follow the ref, got:\n%s", got) - } + c.LessOrEqual( + strings.Index(got, wantWork), + strings.Index(got, wantMax), + "inline keys not sorted, got:\n%s", + got, + ) + c.GreaterOrEqual( + strings.Index(got, "ref = 'x'"), + strings.Index(got, wantMax), + "inline map should follow the ref, got:\n%s", + got, + ) }) t.Run("single quotes in inline values are escaped", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := Render(Defaults(), "", map[string]string{"log_line_prefix": "it's %m"}) - if err != nil { - t.Fatalf("Render() error = %v", err) - } - if !strings.Contains(got, "log_line_prefix = 'it''s %m'") { - t.Errorf("single quote not escaped, got:\n%s", got) - } + c.Require().NoError(err, "Render() error =") + c.StrContains(got, "log_line_prefix = 'it''s %m'", "single quote not escaped, got:\n") }) t.Run("empty ref and nil map render only the baseline", func(t *testing.T) { + c := assert.NewCollecting(t) got, err := Render(Defaults(), "\n\n", nil) - if err != nil { - t.Fatalf("Render() error = %v", err) - } - if strings.Contains(got, "# postgresConfigRef") || - strings.Contains(got, "# spec.postgresConfig") { - t.Errorf("unexpected override section for empty inputs, got:\n%s", got) - } + c.Require().NoError(err, "Render() error =") + c.False(strings.Contains(got, "# postgresConfigRef") || + strings.Contains( + got, + "# spec.postgresConfig", + ), "unexpected override section for empty inputs, got:\n%s", got) }) t.Run("deterministic across calls", func(t *testing.T) { + c := assert.NewCollecting(t) in := map[string]string{"a": "1", "b": "2", "c": "3"} first, err := Render(Defaults(), "x = 'y'", in) - if err != nil { - t.Fatalf("Render() error = %v", err) - } + c.Require().NoError(err, "Render() error =") second, err := Render(Defaults(), "x = 'y'", in) - if err != nil { - t.Fatalf("Render() error = %v", err) - } - if first != second { - t.Error("Render is not deterministic") - } + c.Require().NoError(err, "Render() error =") + c.Eq(second, first, "Render is not deterministic") }) } func TestDefaults(t *testing.T) { + c := assert.NewCollecting(t) d := Defaults() - if d.MaxConnections != 60 { - t.Errorf("MaxConnections = %d, want 60", d.MaxConnections) - } - if d.SharedBuffers != "64MB" { - t.Errorf("SharedBuffers = %q, want 64MB", d.SharedBuffers) - } - if d.ClusterName != "default" { - t.Errorf("ClusterName = %q, want default", d.ClusterName) - } + c.Eq(60, d.MaxConnections, "MaxConnections") + c.Eq("64MB", d.SharedBuffers, "SharedBuffers") + c.Eq("default", d.ClusterName, "ClusterName") } diff --git a/pkg/postgresconfig/sizing_test.go b/pkg/postgresconfig/sizing_test.go index 780aab57..8ca89b8a 100644 --- a/pkg/postgresconfig/sizing_test.go +++ b/pkg/postgresconfig/sizing_test.go @@ -1,13 +1,16 @@ package postgresconfig -import "testing" +import ( + "testing" + + "github.com/multigres/testkit/assert" +) func TestApplyResourceSizing_Memory(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() // 512Mi memory, no CPU, no disk. - if err := ApplyResourceSizing(&cfg, 512*mib, 0, 0); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } + c.Require().NoError(ApplyResourceSizing(&cfg, 512*mib, 0, 0), "ApplyResourceSizing() error =") checks := map[string]string{ "SharedBuffers": "128MB", // 512Mi / 4 "EffectiveCacheSize": "384MB", // 512Mi * 3/4 @@ -22,29 +25,24 @@ func TestApplyResourceSizing_Memory(t *testing.T) { "WorkMem": cfg.WorkMem, "WalBuffers": cfg.WalBuffers, } - if cfg.MaxConnections != 60 { - t.Errorf("MaxConnections = %d, want 60 (below 2GiB)", cfg.MaxConnections) - } - if cfg.MaxWalSenders != 5 || cfg.MaxReplicationSlots != 5 { - t.Errorf("MaxWalSenders/Slots = %d/%d, want 5/5", - cfg.MaxWalSenders, cfg.MaxReplicationSlots) - } + c.Eq(60, cfg.MaxConnections, "MaxConnections") + c.False( + cfg.MaxWalSenders != 5 || cfg.MaxReplicationSlots != 5, + "MaxWalSenders/Slots = %d/%d, want 5/5", + cfg.MaxWalSenders, + cfg.MaxReplicationSlots, + ) for k, want := range checks { - if got[k] != want { - t.Errorf("%s = %q, want %q", k, got[k], want) - } + c.Eq(want, got[k], "%s = %q, want", k, got[k]) } } func TestApplyResourceSizing_MaintenanceWorkMemCap(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() // 64Gi / 16 = 4Gi, which must be capped at 2GB. - if err := ApplyResourceSizing(&cfg, 64*gib, 0, 0); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } - if cfg.MaintenanceWorkMem != "2GB" { - t.Errorf("MaintenanceWorkMem = %q, want 2GB (capped)", cfg.MaintenanceWorkMem) - } + c.Require().NoError(ApplyResourceSizing(&cfg, 64*gib, 0, 0), "ApplyResourceSizing() error =") + c.Eq("2GB", cfg.MaintenanceWorkMem, "MaintenanceWorkMem") } func TestApplyResourceSizing_CPU(t *testing.T) { @@ -81,27 +79,13 @@ func TestApplyResourceSizing_CPU(t *testing.T) { } for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() - if err := ApplyResourceSizing(&cfg, 0, tc.millicores, 0); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } - if cfg.MaxWorkerProcesses != tc.wantWorker { - t.Errorf("MaxWorkerProcesses = %d, want %d", cfg.MaxWorkerProcesses, tc.wantWorker) - } - if cfg.MaxParallelWorkersPerGather != tc.wantGather { - t.Errorf( - "MaxParallelWorkersPerGather = %d, want %d", - cfg.MaxParallelWorkersPerGather, - tc.wantGather, - ) - } - if cfg.MaxParallelMaintenanceWorkers != tc.wantMaint { - t.Errorf( - "MaxParallelMaintenanceWorkers = %d, want %d", - cfg.MaxParallelMaintenanceWorkers, - tc.wantMaint, - ) - } + c.Require(). + NoError(ApplyResourceSizing(&cfg, 0, tc.millicores, 0), "ApplyResourceSizing() error =") + c.Eq(tc.wantWorker, cfg.MaxWorkerProcesses, "MaxWorkerProcesses") + c.Eq(tc.wantGather, cfg.MaxParallelWorkersPerGather, "MaxParallelWorkersPerGather") + c.Eq(tc.wantMaint, cfg.MaxParallelMaintenanceWorkers, "MaxParallelMaintenanceWorkers") }) } } @@ -113,16 +97,14 @@ func TestApplyResourceSizing_WorkMemUsesParallelWorkers(t *testing.T) { _ = ApplyResourceSizing(&single, 512*mib, 0, 0) parallel := Defaults() _ = ApplyResourceSizing(¶llel, 512*mib, 8000, 0) - if single.WorkMem == parallel.WorkMem { - t.Errorf("work_mem should shrink with more parallel workers: both %q", single.WorkMem) - } + assert.NewCollecting(t). + NotEq(parallel.WorkMem, single.WorkMem, "work_mem should shrink with more parallel workers: both") } func TestApplyResourceSizing_WAL(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() - if err := ApplyResourceSizing(&cfg, 0, 0, 1*gib); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } + c.Require().NoError(ApplyResourceSizing(&cfg, 0, 0, 1*gib), "ApplyResourceSizing() error =") checks := map[string]string{ "MinWalSize": "64MB", "MaxWalSize": "256MB", @@ -136,9 +118,7 @@ func TestApplyResourceSizing_WAL(t *testing.T) { "MaxSlotWalKeepSize": cfg.MaxSlotWalKeepSize, } for k, want := range checks { - if got[k] != want { - t.Errorf("%s = %q, want %q", k, got[k], want) - } + c.Eq(want, got[k], "%s = %q, want", k, got[k]) } } @@ -147,20 +127,15 @@ func TestApplyResourceSizing_WALScalesDownOnSmallVolume(t *testing.T) { _ = ApplyResourceSizing(&small, 0, 0, 256*mib) // 256Mi volume: max_wal_size = clamp(256/4=64, floor 64, cap) = 64MB, well // below the 1Gi volume's 256MB. - if small.MaxWalSize != "64MB" { - t.Errorf("MaxWalSize for 256Mi volume = %q, want 64MB", small.MaxWalSize) - } + assert.NewCollecting(t).Eq("64MB", small.MaxWalSize, "MaxWalSize for 256Mi volume") } func TestApplyResourceSizing_ZeroInputsLeaveBaseline(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() base := Defaults() - if err := ApplyResourceSizing(&cfg, 0, 0, 0); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } - if cfg != base { - t.Errorf("zero inputs mutated the config: %+v != %+v", cfg, base) - } + c.Require().NoError(ApplyResourceSizing(&cfg, 0, 0, 0), "ApplyResourceSizing() error =") + c.Eq(base, cfg, "zero inputs mutated the config") } func TestFormatBytes(t *testing.T) { @@ -172,9 +147,8 @@ func TestFormatBytes(t *testing.T) { 0: "0kB", } for in, want := range tests { - if got := formatBytes(in); got != want { - t.Errorf("formatBytes(%d) = %q, want %q", in, got, want) - } + got := formatBytes(in) + assert.NewCollecting(t).Eq(want, got, "formatBytes(%d) = %q, want", in, got) } } @@ -186,9 +160,9 @@ func TestDeriveWalSettings_Errors(t *testing.T) { t.Error("expected error for non-MB-aligned WAL segment size") } // A WAL segment large enough that its floor exceeds the max_wal_size cap. - if _, err := deriveWalSettings(1*uint64(gib), 2048*megabyte); err == nil { - t.Error("expected error when segment size forces max_wal_size above the cap") - } + _, err := deriveWalSettings(1*uint64(gib), 2048*megabyte) + assert.NewCollecting(t). + Error(err, "expected error when segment size forces max_wal_size above the cap") } func TestApplyResourceSizing_MaxConnections(t *testing.T) { @@ -214,13 +188,11 @@ func TestApplyResourceSizing_MaxConnections(t *testing.T) { } for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() - if err := ApplyResourceSizing(&cfg, tc.memBytes, 0, 0); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } - if cfg.MaxConnections != tc.want { - t.Errorf("MaxConnections = %d, want %d", cfg.MaxConnections, tc.want) - } + c.Require(). + NoError(ApplyResourceSizing(&cfg, tc.memBytes, 0, 0), "ApplyResourceSizing() error =") + c.Eq(tc.want, cfg.MaxConnections, "MaxConnections") }) } } @@ -239,39 +211,31 @@ func TestApplyResourceSizing_MaxWalSenders(t *testing.T) { } for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) cfg := Defaults() - if err := ApplyResourceSizing(&cfg, tc.memBytes, 0, 0); err != nil { - t.Fatalf("ApplyResourceSizing() error = %v", err) - } - if cfg.MaxWalSenders != tc.want { - t.Errorf("MaxWalSenders = %d, want %d", cfg.MaxWalSenders, tc.want) - } - if cfg.MaxReplicationSlots != tc.want { - t.Errorf("MaxReplicationSlots = %d, want %d", cfg.MaxReplicationSlots, tc.want) - } + c.Require(). + NoError(ApplyResourceSizing(&cfg, tc.memBytes, 0, 0), "ApplyResourceSizing() error =") + c.Eq(tc.want, cfg.MaxWalSenders, "MaxWalSenders") + c.Eq(tc.want, cfg.MaxReplicationSlots, "MaxReplicationSlots") }) } } func TestApplyResourceSizing_MaxConnectionsFeedsWorkMem(t *testing.T) { + c := assert.NewCollecting(t) // work_mem = (mem - shared) / (conns * 3) / parallel. With the derived // MaxConnections=160 at 8GiB, work_mem must be ~1/2.66× what it would be // under the old hardcoded 60. Two different memory anchors sanity-check // that the divisor tracks the derived value rather than the baseline. cfg8 := Defaults() _ = ApplyResourceSizing(&cfg8, 8*gib, 0, 0) - if cfg8.MaxConnections != 160 { - t.Fatalf("precondition: MaxConnections at 8GiB = %d, want 160", cfg8.MaxConnections) - } + c.Require().Eq(160, cfg8.MaxConnections, "precondition: MaxConnections at 8GiB") // (8*gib - 2*gib) / (160*3) = 6GiB/480 = 12.8 MiB, formatted to kB. - if cfg8.WorkMem != "13107kB" { - t.Errorf("WorkMem at 8GiB, 160 conns = %q, want 13107kB", cfg8.WorkMem) - } + c.Eq("13107kB", cfg8.WorkMem, "WorkMem at 8GiB, 160 conns") } func TestDefaults_EffectiveIoConcurrency(t *testing.T) { // SSDs everywhere; PgTune's OLTP/DW profile for SSD storage. - if got := Defaults().EffectiveIoConcurrency; got != 200 { - t.Errorf("Defaults().EffectiveIoConcurrency = %d, want 200", got) - } + assert.NewCollecting(t). + Eq(200, Defaults().EffectiveIoConcurrency, "Defaults().EffectiveIoConcurrency") } diff --git a/pkg/postgresconfig/validate_test.go b/pkg/postgresconfig/validate_test.go index a038002d..cd094358 100644 --- a/pkg/postgresconfig/validate_test.go +++ b/pkg/postgresconfig/validate_test.go @@ -3,6 +3,8 @@ package postgresconfig import ( "strings" "testing" + + "github.com/multigres/testkit/assert" ) func TestValidate_Accepts(t *testing.T) { @@ -27,9 +29,8 @@ func TestValidate_Accepts(t *testing.T) { } for name, cfg := range tests { t.Run(name, func(t *testing.T) { - if err := Validate(cfg); err != nil { - t.Errorf("Validate(%v) = %v, want nil", cfg, err) - } + err := Validate(cfg) + assert.NewCollecting(t).NoError(err, "Validate(%v) = %v, want nil", cfg, err) }) } } @@ -74,42 +75,38 @@ func TestValidate_Rejects(t *testing.T) { } for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) err := Validate(tc.cfg) - if err == nil { - t.Fatalf("Validate(%v) = nil, want error", tc.cfg) - } + c.Require().Error(err, "Validate(%v) = nil, want error", tc.cfg) for _, sub := range tc.wantSubs { - if !strings.Contains(err.Error(), sub) { - t.Errorf("error %q missing %q", err.Error(), sub) - } + c.StrContains(err.Error(), sub, "error") } }) } } func TestValidate_AggregatesAllProblems(t *testing.T) { + c := assert.NewCollecting(t) err := Validate(map[string]string{ "maxx_connections": "200", // unknown "fsync": "maybe", // bad bool "max_connections": "200", // valid — should not appear }) - if err == nil { - t.Fatal("expected error") - } - if !strings.Contains(err.Error(), "maxx_connections") || - !strings.Contains(err.Error(), "fsync") { - t.Errorf("error should mention both problems: %v", err) - } - if strings.Contains(err.Error(), `"max_connections"`) { - t.Errorf("error should not flag the valid parameter: %v", err) - } + c.Require().Error(err, "expected error") + c.False(!strings.Contains(err.Error(), "maxx_connections") || + !strings.Contains(err.Error(), "fsync"), "error should mention both problems: %v", err) + c.NotStrContains( + err.Error(), + `"max_connections"`, + "error should not flag the valid parameter: %v", + err, + ) } func TestCatalogLoaded(t *testing.T) { + c := assert.NewCollecting(t) // The embedded catalog must be non-trivial and contain well-known params. - if len(catalog) < 300 { - t.Errorf("catalog has %d entries, expected the full PG17 set", len(catalog)) - } + c.GreaterOrEqual(300, len(catalog), "catalog has") for name, want := range map[string]gucType{ "max_connections": gucInteger, "shared_buffers": gucInteger, @@ -118,8 +115,7 @@ func TestCatalogLoaded(t *testing.T) { "wal_level": gucEnum, "log_line_prefix": gucString, } { - if got := catalog[name].typ; got != want { - t.Errorf("catalog[%q].typ = %q, want %q", name, got, want) - } + got := catalog[name].typ + c.Eq(want, got, "catalog[%q].typ = %q, want", name, got) } } diff --git a/pkg/resolver/buffer_defaults_test.go b/pkg/resolver/buffer_defaults_test.go index b1e20fc5..b1dc5a5b 100644 --- a/pkg/resolver/buffer_defaults_test.go +++ b/pkg/resolver/buffer_defaults_test.go @@ -6,6 +6,8 @@ import ( "github.com/multigres/multigres/go/services/multigateway/buffer" "github.com/multigres/multigres/go/tools/viperutil" + + "github.com/multigres/testkit/assert" ) // TestBufferDefaultsMatchBinary pins the hardcoded admission constants — and @@ -15,25 +17,25 @@ import ( // silently desynchronizing webhook verdicts (or documentation) from binary // startup behavior. func TestBufferDefaultsMatchBinary(t *testing.T) { + c := assert.NewCollecting(t) cfg := buffer.NewConfig(viperutil.NewRegistry()) - if got := cfg.Window.Default(); got != defaultBufferWindow { - t.Errorf("defaultBufferWindow = %s, binary default = %s", defaultBufferWindow, got) - } - if got := cfg.MaxFailoverDuration.Default(); got != defaultBufferMaxFailoverDuration { - t.Errorf( - "defaultBufferMaxFailoverDuration = %s, binary default = %s", - defaultBufferMaxFailoverDuration, got, - ) - } + c.Eq(defaultBufferWindow, cfg.Window.Default(), "defaultBufferWindow") + c.Eq( + defaultBufferMaxFailoverDuration, + cfg.MaxFailoverDuration.Default(), + "defaultBufferMaxFailoverDuration", + ) // Documented (not validated) defaults: update api/v1alpha1/cell_types.go // and config/samples/README.md if any of these fail. - if got := cfg.MinTimeBetweenFailovers.Default(); got != time.Minute { - t.Errorf("documented minTimeBetweenFailovers default 1m, binary default = %s", got) - } - if got := cfg.Size.Default(); got != 1000 { - t.Errorf("documented size default 1000, binary default = %d", got) - } - if got := cfg.DrainConcurrency.Default(); got != 1 { - t.Errorf("documented drainConcurrency default 1, binary default = %d", got) - } + c.Eq( + time.Minute, + cfg.MinTimeBetweenFailovers.Default(), + "documented minTimeBetweenFailovers default 1m, binary default =", + ) + c.Eq(1000, cfg.Size.Default(), "documented size default 1000, binary default =") + c.Eq( + 1, + cfg.DrainConcurrency.Default(), + "documented drainConcurrency default 1, binary default =", + ) } diff --git a/pkg/resolver/cell_test.go b/pkg/resolver/cell_test.go index 48f4ddcd..ee78f495 100644 --- a/pkg/resolver/cell_test.go +++ b/pkg/resolver/cell_test.go @@ -17,6 +17,8 @@ import ( "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestResolver_ResolveCell(t *testing.T) { @@ -147,6 +149,7 @@ func TestResolver_ResolveCell(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) var c client.Client if name == "Client Error" { base := fake.NewClientBuilder().WithScheme(scheme).Build() @@ -163,34 +166,29 @@ func TestResolver_ResolveCell(t *testing.T) { gw, placement, topo, err := r.ResolveCell(t.Context(), cluster, tc.config) if tc.wantErr { - if err == nil { - t.Error("Expected error") - } + ck.Error(err, "Expected error") return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } + ck.Require().NoError(err, "Unexpected error") - if diff := cmp.Diff( + ck.EqDiffOpts( tc.wantGw, gw, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Gateway Diff (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantPlacement, placement, cmpopts.EquateEmpty()); diff != "" { - t.Errorf("Placement Diff (-want +got):\n%s", diff) - } - if diff := cmp.Diff( + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Gateway Diff", + ) + ck.EqDiffOpts( + tc.wantPlacement, + placement, + []cmp.Option{cmpopts.EquateEmpty()}, + "Placement Diff", + ) + ck.EqDiffOpts( tc.wantTopo, topo, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Topo Diff (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Topo Diff", + ) }) } } @@ -245,6 +243,7 @@ func TestResolver_ResolveCellTemplate(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder(). WithScheme(scheme). WithObjects(tc.existingObjects...). @@ -253,9 +252,7 @@ func TestResolver_ResolveCellTemplate(t *testing.T) { res, err := r.ResolveCellTemplate(t.Context(), tc.reqName) if tc.wantErr { - if err == nil { - t.Fatal("Expected error, got nil") - } + ck.Require().Error(err, "Expected error, got nil") if tc.errContains != "" && !strings.Contains(err.Error(), tc.errContains) { t.Errorf( "Error message mismatch: got %q, want substring %q", @@ -269,20 +266,14 @@ func TestResolver_ResolveCellTemplate(t *testing.T) { } if !tc.wantFound { - if res == nil { - t.Fatal( - "Expected non-nil result structure even for not-found implicit fallback", - ) - } - if res.GetName() != "" { - t.Errorf("Expected empty result, got object with name %q", res.GetName()) - } + ck.Require(). + NotNil(res, "Expected non-nil result structure even for not-found implicit fallback") + ck.Eq("", res.GetName(), "Expected empty result, got object with name") return } - if got, want := res.GetName(), tc.wantResName; got != want { - t.Errorf("Result name mismatch: got %q, want %q", got, want) - } + got, want := res.GetName(), tc.wantResName + ck.Eq(want, got, "Result name mismatch: got") }) } } @@ -548,22 +539,22 @@ func TestMergeCellConfig(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) gw, placement, topo := mergeCellConfig(tc.tpl, tc.overrides, tc.inline) - if diff := cmp.Diff( + c.EqDiffOpts( tc.wantGw, gw, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Gateway mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantPlacement, placement, cmpopts.EquateEmpty()); diff != "" { - t.Errorf("Placement mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantTopo, topo); diff != "" { - t.Errorf("Topo mismatch (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Gateway mismatch", + ) + c.EqDiffOpts( + tc.wantPlacement, + placement, + []cmp.Option{cmpopts.EquateEmpty()}, + "Placement mismatch", + ) + c.EqDiff(tc.wantTopo, topo, "Topo mismatch") }) } } @@ -582,8 +573,6 @@ func TestResolver_ClientErrors_Cell(t *testing.T) { r := NewResolver(mc, "default") _, err := r.ResolveCellTemplate(t.Context(), "any") - if err == nil || - err.Error() != "failed to get CellTemplate: simulated database connection error" { - t.Errorf("Error mismatch: got %v, want simulated error", err) - } + assert.NewCollecting(t).False(err == nil || + err.Error() != "failed to get CellTemplate: simulated database connection error", "Error mismatch: got %v, want simulated error", err) } diff --git a/pkg/resolver/cluster_test.go b/pkg/resolver/cluster_test.go index 78ca228a..6c127727 100644 --- a/pkg/resolver/cluster_test.go +++ b/pkg/resolver/cluster_test.go @@ -15,6 +15,8 @@ import ( "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestResolver_PopulateClusterDefaults(t *testing.T) { @@ -489,24 +491,22 @@ func TestResolver_PopulateClusterDefaults(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) r := NewResolver( fake.NewClientBuilder().WithScheme(scheme).WithObjects(tc.objects...).Build(), "default", ) got := tc.input.DeepCopy() - if _, err := r.PopulateClusterDefaults(t.Context(), got); err != nil { - t.Fatalf("PopulateClusterDefaults failed: %v", err) - } + _, err := r.PopulateClusterDefaults(t.Context(), got) + c.Require().NoError(err, "PopulateClusterDefaults failed") - if diff := cmp.Diff( + c.EqDiffOpts( tc.want, got, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Diff (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Diff", + ) }) } } @@ -531,9 +531,8 @@ func TestResolver_PopulateClusterDefaults_ClientError(t *testing.T) { } _, err := r.PopulateClusterDefaults(t.Context(), input) - if err == nil || !errors.Is(err, errSim) { - t.Errorf("Expected simulated error, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !errors.Is(err, errSim), "Expected simulated error, got %v", err) } func TestResolver_ResolveGlobalTopo(t *testing.T) { @@ -922,6 +921,7 @@ func TestResolver_ResolveGlobalTopo(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) if tc.cluster.Name == "" { tc.cluster.Name = "test-cluster" } @@ -933,22 +933,16 @@ func TestResolver_ResolveGlobalTopo(t *testing.T) { got, err := r.ResolveGlobalTopo(t.Context(), tc.cluster) if tc.wantErr { - if err == nil { - t.Error("Expected error, got nil") - } + ck.Error(err, "Expected error, got nil") return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } - if diff := cmp.Diff( + ck.Require().NoError(err, "Unexpected error") + ck.EqDiffOpts( tc.want, got, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Diff (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Diff", + ) }) } } @@ -1116,30 +1110,28 @@ func TestResolver_ResolveMultiadmin(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(tc.objects...).Build() r := NewResolver(c, ns) got, gotPlacement, err := r.ResolveMultiadmin(t.Context(), tc.cluster) if tc.wantErr { - if err == nil { - t.Error("Expected error") - } + ck.Error(err, "Expected error") return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } - if diff := cmp.Diff( + ck.Require().NoError(err, "Unexpected error") + ck.EqDiffOpts( tc.want, got, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Diff (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantPlacement, gotPlacement, cmpopts.EquateEmpty()); diff != "" { - t.Errorf("Placement diff (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Diff", + ) + ck.EqDiffOpts( + tc.wantPlacement, + gotPlacement, + []cmp.Option{cmpopts.EquateEmpty()}, + "Placement diff", + ) }) } } @@ -1158,22 +1150,17 @@ func TestResolver_ResolveCoreTemplate(t *testing.T) { // 1. Implicit Fallback ("default" or "") -> Not Found -> Returns nil, nil (No Error) // This covers: "if isImplicitFallback { return ... nil }" t.Run("Implicit Fallback Missing", func(t *testing.T) { + c := assert.NewCollecting(t) tpl, err := r.ResolveCoreTemplate(t.Context(), "") - if err != nil { - t.Errorf("Expected nil error for implicit missing, got %v", err) - } - if tpl.Name != "" { // Empty struct - t.Errorf("Expected empty template, got %v", tpl) - } + c.NoError(err, "Expected nil error for implicit missing, got") + c.Eq("", tpl.Name, "Expected empty template, got %v", tpl) // Empty struct }) // 2. Explicit Template ("custom") -> Not Found -> Returns Error // This covers: "return nil, fmt.Errorf(...)" t.Run("Explicit Template Missing", func(t *testing.T) { _, err := r.ResolveCoreTemplate(t.Context(), "missing-custom") - if err == nil { - t.Error("Expected error for explicit missing template") - } + assert.NewCollecting(t).Error(err, "Expected error for explicit missing template") }) } @@ -1192,9 +1179,8 @@ func TestResolver_ClientErrors_Core(t *testing.T) { r := NewResolver(c, "default") _, err := r.ResolveCoreTemplate(t.Context(), "any") - if err == nil || !errors.Is(err, errSim) { - t.Errorf("Error mismatch: got %v, want %v", err, errSim) - } + assert.NewCollecting(t). + False(err == nil || !errors.Is(err, errSim), "Error mismatch: got %v, want %v", err, errSim) } func TestResolver_ResolveMultiadminWeb(t *testing.T) { @@ -1294,27 +1280,22 @@ func TestResolver_ResolveMultiadminWeb(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(tc.objects...).Build() r := NewResolver(c, ns) got, err := r.ResolveMultiadminWeb(t.Context(), tc.cluster) if tc.wantErr { - if err == nil { - t.Error("Expected error") - } + ck.Error(err, "Expected error") return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } - if diff := cmp.Diff( + ck.Require().NoError(err, "Unexpected error") + ck.EqDiffOpts( tc.want, got, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Diff (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Diff", + ) }) } } @@ -1339,9 +1320,7 @@ func TestResolveGlobalTopo_PVCDeletionPolicy(t *testing.T) { }, } got, err := r.ResolveGlobalTopo(t.Context(), cluster) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") if got.PVCDeletionPolicy == nil || got.PVCDeletionPolicy.WhenDeleted != multigresv1alpha1.DeletePVCRetentionPolicy { t.Errorf("Expected GlobalTopo PVCDeletionPolicy=Delete, got %v", got.PVCDeletionPolicy) @@ -1362,9 +1341,7 @@ func TestResolveGlobalTopo_PVCDeletionPolicy(t *testing.T) { }, } got, err := r.ResolveGlobalTopo(t.Context(), cluster) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") // mergeEtcdSpec should have applied this to the base if got.Etcd.PVCDeletionPolicy == nil || got.Etcd.PVCDeletionPolicy.WhenDeleted != multigresv1alpha1.RetainPVCRetentionPolicy { diff --git a/pkg/resolver/etcd_maintenance_test.go b/pkg/resolver/etcd_maintenance_test.go index f7b9164b..b193b17e 100644 --- a/pkg/resolver/etcd_maintenance_test.go +++ b/pkg/resolver/etcd_maintenance_test.go @@ -5,6 +5,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "k8s.io/utils/ptr" + + "github.com/multigres/testkit/assert" ) func TestMergeEtcdMaintenance(t *testing.T) { @@ -28,8 +30,6 @@ func TestMergeEtcdMaintenance(t *testing.T) { } *override.Maintenance.DefragmentationEnabled = true *override.Maintenance.QuotaBackendBytes = 2 << 30 - if base.Maintenance.DefragmentationIsEnabled() || - base.Maintenance.EffectiveQuotaBackendBytes() != 512<<20 { - t.Fatal("merged maintenance aliases override") - } + assert.NewAborting(t).False(base.Maintenance.DefragmentationIsEnabled() || + base.Maintenance.EffectiveQuotaBackendBytes() != 512<<20, "merged maintenance aliases override") } diff --git a/pkg/resolver/resolver_test.go b/pkg/resolver/resolver_test.go index a6595fdf..7fed0a45 100644 --- a/pkg/resolver/resolver_test.go +++ b/pkg/resolver/resolver_test.go @@ -4,7 +4,6 @@ import ( "errors" "testing" - "github.com/google/go-cmp/cmp" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" corev1 "k8s.io/api/core/v1" @@ -15,6 +14,8 @@ import ( "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) // setupFixtures helper returns a fresh set of test objects. @@ -79,9 +80,8 @@ func TestNewResolver(t *testing.T) { if got, want := r.Client, c; got != want { t.Errorf("Client mismatch: got %v, want %v", got, want) } - if got, want := r.Namespace, "ns"; got != want { - t.Errorf("Namespace mismatch: got %q, want %q", got, want) - } + got, want := r.Namespace, "ns" + assert.NewCollecting(t).Eq(want, got, "Namespace mismatch: got") } // setupScheme creates a new scheme with all required types registered @@ -195,80 +195,99 @@ func TestResolver_ValidateReference(t *testing.T) { // Case 1: Core Template t.Run("Core", func(t *testing.T) { + c := assert.NewCollecting(t) r := NewResolver(cEmpty, ns) // Empty name -> Valid (no explicit reference) - if err := r.ValidateCoreTemplateReference(t.Context(), ""); err != nil { - t.Errorf("Empty name should be valid, got %v", err) - } + c.NoError( + r.ValidateCoreTemplateReference(t.Context(), ""), + "Empty name should be valid, got", + ) // Explicit "default" reference with missing template -> Invalid - if err := r.ValidateCoreTemplateReference(t.Context(), FallbackCoreTemplate); err == nil { - t.Error("Explicit 'default' reference should error when template is missing") - } + c.Error( + r.ValidateCoreTemplateReference(t.Context(), FallbackCoreTemplate), + "Explicit 'default' reference should error when template is missing", + ) // Random missing -> Invalid - if err := r.ValidateCoreTemplateReference(t.Context(), "missing"); err == nil { - t.Error("Missing template should error") - } + c.Error( + r.ValidateCoreTemplateReference(t.Context(), "missing"), + "Missing template should error", + ) // Real existence (Implicit Default) rExists := NewResolver(cWithDefaults, ns) - if err := rExists.ValidateCoreTemplateReference(t.Context(), "default"); err != nil { - t.Errorf("Existing template should be valid, got %v", err) - } + c.NoError( + rExists.ValidateCoreTemplateReference(t.Context(), "default"), + "Existing template should be valid, got", + ) // Real existence (Explicit Custom) - Hits "exists" branch - if err := rExists.ValidateCoreTemplateReference(t.Context(), "custom"); err != nil { - t.Errorf("Custom template should be valid, got %v", err) - } + c.NoError( + rExists.ValidateCoreTemplateReference(t.Context(), "custom"), + "Custom template should be valid, got", + ) }) // Case 2: Cell Template t.Run("Cell", func(t *testing.T) { + c := assert.NewCollecting(t) r := NewResolver(cEmpty, ns) - if err := r.ValidateCellTemplateReference(t.Context(), ""); err != nil { - t.Errorf("Empty name should be valid, got %v", err) - } + c.NoError( + r.ValidateCellTemplateReference(t.Context(), ""), + "Empty name should be valid, got", + ) // Explicit "default" reference with missing template -> Invalid - if err := r.ValidateCellTemplateReference(t.Context(), FallbackCellTemplate); err == nil { - t.Error("Explicit 'default' reference should error when template is missing") - } - if err := r.ValidateCellTemplateReference(t.Context(), "missing"); err == nil { - t.Error("Missing template should error") - } + c.Error( + r.ValidateCellTemplateReference(t.Context(), FallbackCellTemplate), + "Explicit 'default' reference should error when template is missing", + ) + c.Error( + r.ValidateCellTemplateReference(t.Context(), "missing"), + "Missing template should error", + ) rExists := NewResolver(cWithDefaults, ns) - if err := rExists.ValidateCellTemplateReference(t.Context(), "default"); err != nil { - t.Errorf("Existing template should be valid, got %v", err) - } - if err := rExists.ValidateCellTemplateReference(t.Context(), "custom"); err != nil { - t.Errorf("Custom template should be valid, got %v", err) - } + c.NoError( + rExists.ValidateCellTemplateReference(t.Context(), "default"), + "Existing template should be valid, got", + ) + c.NoError( + rExists.ValidateCellTemplateReference(t.Context(), "custom"), + "Custom template should be valid, got", + ) }) // Case 3: Shard Template t.Run("Shard", func(t *testing.T) { + c := assert.NewCollecting(t) r := NewResolver(cEmpty, ns) - if err := r.ValidateShardTemplateReference(t.Context(), ""); err != nil { - t.Errorf("Empty name should be valid, got %v", err) - } + c.NoError( + r.ValidateShardTemplateReference(t.Context(), ""), + "Empty name should be valid, got", + ) // Explicit "default" reference with missing template -> Invalid - if err := r.ValidateShardTemplateReference(t.Context(), FallbackShardTemplate); err == nil { - t.Error("Explicit 'default' reference should error when template is missing") - } - if err := r.ValidateShardTemplateReference(t.Context(), "missing"); err == nil { - t.Error("Missing template should error") - } + c.Error( + r.ValidateShardTemplateReference(t.Context(), FallbackShardTemplate), + "Explicit 'default' reference should error when template is missing", + ) + c.Error( + r.ValidateShardTemplateReference(t.Context(), "missing"), + "Missing template should error", + ) rExists := NewResolver(cWithDefaults, ns) - if err := rExists.ValidateShardTemplateReference(t.Context(), "default"); err != nil { - t.Errorf("Existing template should be valid, got %v", err) - } - if err := rExists.ValidateShardTemplateReference(t.Context(), "custom"); err != nil { - t.Errorf("Custom template should be valid, got %v", err) - } + c.NoError( + rExists.ValidateShardTemplateReference(t.Context(), "default"), + "Existing template should be valid, got", + ) + c.NoError( + rExists.ValidateShardTemplateReference(t.Context(), "custom"), + "Custom template should be valid, got", + ) }) // Case 4: Client Failure t.Run("ClientFailure", func(t *testing.T) { + c := assert.NewCollecting(t) errSim := testutil.ErrInjected failClient := testutil.NewFakeClientWithFailures( fake.NewClientBuilder().Build(), @@ -279,18 +298,21 @@ func TestResolver_ValidateReference(t *testing.T) { rFail := NewResolver(failClient, ns) // Should propagate error - if err := rFail.ValidateCoreTemplateReference(t.Context(), "any"); !errors.Is(err, errSim) { - t.Errorf("Expected error propagation for Core, got %v", err) - } - if err := rFail.ValidateCellTemplateReference(t.Context(), "any"); !errors.Is(err, errSim) { - t.Errorf("Expected error propagation for Cell, got %v", err) - } - if err := rFail.ValidateShardTemplateReference(t.Context(), "any"); !errors.Is( - err, + c.ErrorIs( + rFail.ValidateCoreTemplateReference(t.Context(), "any"), errSim, - ) { - t.Errorf("Expected error propagation for Shard, got %v", err) - } + "Expected error propagation for Core, got", + ) + c.ErrorIs( + rFail.ValidateCellTemplateReference(t.Context(), "any"), + errSim, + "Expected error propagation for Cell, got", + ) + c.ErrorIs( + rFail.ValidateShardTemplateReference(t.Context(), "any"), + errSim, + "Expected error propagation for Shard, got", + ) }) } @@ -305,6 +327,7 @@ func TestResolver_Caching(t *testing.T) { baseClient := fake.NewClientBuilder().WithScheme(scheme).WithObjects(objs...).Build() t.Run("CoreTemplate", func(t *testing.T) { + c := assert.NewCollecting(t) // Use a counter to track Get calls var getCalls int clientWithCounter := testutil.NewFakeClientWithFailures(baseClient, &testutil.FailureConfig{ @@ -318,34 +341,24 @@ func TestResolver_Caching(t *testing.T) { // First call - should hit the API tpl1, err := r.ResolveCoreTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("First ResolveCoreTemplate failed: %v", err) - } - if tpl1 == nil { - t.Fatal("Expected non-nil template") - } - if getCalls != 1 { - t.Errorf("Expected 1 Get call after first resolve, got %d", getCalls) - } + c.Require().NoError(err, "First ResolveCoreTemplate failed") + c.Require().NotNil(tpl1, "Expected non-nil template") + c.Eq(1, getCalls, "Expected 1 Get call after first resolve, got") // Second call - should use cache tpl2, err := r.ResolveCoreTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("Second ResolveCoreTemplate failed: %v", err) - } - if tpl2 == nil { - t.Fatal("Expected non-nil template") - } - if getCalls != 1 { - t.Errorf("Expected 1 Get call after second resolve (cached), got %d", getCalls) - } + c.Require().NoError(err, "Second ResolveCoreTemplate failed") + c.Require().NotNil(tpl2, "Expected non-nil template") + c.Eq(1, getCalls, "Expected 1 Get call after second resolve (cached), got") // Verify DeepCopy (forward): modifying first result shouldn't affect second if tpl1.Spec.GlobalTopoServer != nil && tpl1.Spec.GlobalTopoServer.Etcd != nil { tpl1.Spec.GlobalTopoServer.Etcd.Image = "modified" - if tpl2.Spec.GlobalTopoServer.Etcd.Image == "modified" { - t.Error("DeepCopy failed - modifications leaked from tpl1 to tpl2") - } + c.NotEq( + "modified", + tpl2.Spec.GlobalTopoServer.Etcd.Image, + "DeepCopy failed - modifications leaked from tpl1 to tpl2", + ) } // Verify DeepCopy (reverse): modifying a cached result must not @@ -354,16 +367,17 @@ func TestResolver_Caching(t *testing.T) { if tpl2.Spec.GlobalTopoServer != nil && tpl2.Spec.GlobalTopoServer.Etcd != nil { tpl2.Spec.GlobalTopoServer.Etcd.Image = "corrupted" tpl3, err := r.ResolveCoreTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("Third ResolveCoreTemplate failed: %v", err) - } - if tpl3.Spec.GlobalTopoServer.Etcd.Image == "corrupted" { - t.Error("Cache corruption - mutating a cached result polluted subsequent resolves") - } + c.Require().NoError(err, "Third ResolveCoreTemplate failed") + c.NotEq( + "corrupted", + tpl3.Spec.GlobalTopoServer.Etcd.Image, + "Cache corruption - mutating a cached result polluted subsequent resolves", + ) } }) t.Run("CellTemplate", func(t *testing.T) { + c := assert.NewCollecting(t) var getCalls int clientWithCounter := testutil.NewFakeClientWithFailures(baseClient, &testutil.FailureConfig{ OnGet: func(_ client.ObjectKey) error { @@ -376,34 +390,23 @@ func TestResolver_Caching(t *testing.T) { // First call tpl1, err := r.ResolveCellTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("First ResolveCellTemplate failed: %v", err) - } - if tpl1 == nil { - t.Fatal("Expected non-nil template") - } - if getCalls != 1 { - t.Errorf("Expected 1 Get call after first resolve, got %d", getCalls) - } + c.Require().NoError(err, "First ResolveCellTemplate failed") + c.Require().NotNil(tpl1, "Expected non-nil template") + c.Eq(1, getCalls, "Expected 1 Get call after first resolve, got") // Second call - should use cache tpl2, err := r.ResolveCellTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("Second ResolveCellTemplate failed: %v", err) - } - if tpl2 == nil { - t.Fatal("Expected non-nil template") - } - if getCalls != 1 { - t.Errorf("Expected 1 Get call after second resolve (cached), got %d", getCalls) - } + c.Require().NoError(err, "Second ResolveCellTemplate failed") + c.Require().NotNil(tpl2, "Expected non-nil template") + c.Eq(1, getCalls, "Expected 1 Get call after second resolve (cached), got") // Verify DeepCopy (forward) if tpl1.Spec.Multigateway != nil { tpl1.Spec.Multigateway.Replicas = ptr.To(int32(999)) - if tpl2.Spec.Multigateway.Replicas != nil && *tpl2.Spec.Multigateway.Replicas == 999 { - t.Error("DeepCopy failed - modifications leaked from tpl1 to tpl2") - } + c.False( + tpl2.Spec.Multigateway.Replicas != nil && *tpl2.Spec.Multigateway.Replicas == 999, + "DeepCopy failed - modifications leaked from tpl1 to tpl2", + ) } // Verify DeepCopy (reverse): mutating a cached result must not @@ -411,16 +414,16 @@ func TestResolver_Caching(t *testing.T) { if tpl2.Spec.Multigateway != nil { tpl2.Spec.Multigateway.Replicas = ptr.To(int32(777)) tpl3, err := r.ResolveCellTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("Third ResolveCellTemplate failed: %v", err) - } - if tpl3.Spec.Multigateway.Replicas != nil && *tpl3.Spec.Multigateway.Replicas == 777 { - t.Error("Cache corruption - mutating a cached result polluted subsequent resolves") - } + c.Require().NoError(err, "Third ResolveCellTemplate failed") + c.False( + tpl3.Spec.Multigateway.Replicas != nil && *tpl3.Spec.Multigateway.Replicas == 777, + "Cache corruption - mutating a cached result polluted subsequent resolves", + ) } }) t.Run("ShardTemplate", func(t *testing.T) { + c := assert.NewCollecting(t) var getCalls int clientWithCounter := testutil.NewFakeClientWithFailures(baseClient, &testutil.FailureConfig{ OnGet: func(_ client.ObjectKey) error { @@ -433,34 +436,23 @@ func TestResolver_Caching(t *testing.T) { // First call tpl1, err := r.ResolveShardTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("First ResolveShardTemplate failed: %v", err) - } - if tpl1 == nil { - t.Fatal("Expected non-nil template") - } - if getCalls != 1 { - t.Errorf("Expected 1 Get call after first resolve, got %d", getCalls) - } + c.Require().NoError(err, "First ResolveShardTemplate failed") + c.Require().NotNil(tpl1, "Expected non-nil template") + c.Eq(1, getCalls, "Expected 1 Get call after first resolve, got") // Second call - should use cache tpl2, err := r.ResolveShardTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("Second ResolveShardTemplate failed: %v", err) - } - if tpl2 == nil { - t.Fatal("Expected non-nil template") - } - if getCalls != 1 { - t.Errorf("Expected 1 Get call after second resolve (cached), got %d", getCalls) - } + c.Require().NoError(err, "Second ResolveShardTemplate failed") + c.Require().NotNil(tpl2, "Expected non-nil template") + c.Eq(1, getCalls, "Expected 1 Get call after second resolve (cached), got") // Verify DeepCopy (forward) if tpl1.Spec.Multiorch != nil { tpl1.Spec.Multiorch.Replicas = ptr.To(int32(999)) - if tpl2.Spec.Multiorch.Replicas != nil && *tpl2.Spec.Multiorch.Replicas == 999 { - t.Error("DeepCopy failed - modifications leaked from tpl1 to tpl2") - } + c.False( + tpl2.Spec.Multiorch.Replicas != nil && *tpl2.Spec.Multiorch.Replicas == 999, + "DeepCopy failed - modifications leaked from tpl1 to tpl2", + ) } // Verify DeepCopy (reverse): mutating a cached result must not @@ -468,16 +460,16 @@ func TestResolver_Caching(t *testing.T) { if tpl2.Spec.Multiorch != nil { tpl2.Spec.Multiorch.Replicas = ptr.To(int32(777)) tpl3, err := r.ResolveShardTemplate(t.Context(), "default") - if err != nil { - t.Fatalf("Third ResolveShardTemplate failed: %v", err) - } - if tpl3.Spec.Multiorch.Replicas != nil && *tpl3.Spec.Multiorch.Replicas == 777 { - t.Error("Cache corruption - mutating a cached result polluted subsequent resolves") - } + c.Require().NoError(err, "Third ResolveShardTemplate failed") + c.False( + tpl3.Spec.Multiorch.Replicas != nil && *tpl3.Spec.Multiorch.Replicas == 777, + "Cache corruption - mutating a cached result polluted subsequent resolves", + ) } }) t.Run("FallbackNotCached", func(t *testing.T) { + c := assert.NewCollecting(t) // Verify that fallback empty templates are NOT cached var getCalls int emptyClient := fake.NewClientBuilder().WithScheme(scheme).Build() @@ -495,25 +487,17 @@ func TestResolver_Caching(t *testing.T) { // Call with empty name (triggers fallback) _, err := r.ResolveShardTemplate(t.Context(), "") - if err != nil { - t.Fatalf("Fallback resolve failed: %v", err) - } + c.Require().NoError(err, "Fallback resolve failed") // Should have attempted Get once and got NotFound - if getCalls != 1 { - t.Errorf("Expected 1 Get call for fallback, got %d", getCalls) - } + c.Eq(1, getCalls, "Expected 1 Get call for fallback, got") // Second call - fallback should NOT be cached, so another Get attempt _, err = r.ResolveShardTemplate(t.Context(), "") - if err != nil { - t.Fatalf("Second fallback resolve failed: %v", err) - } + c.Require().NoError(err, "Second fallback resolve failed") // Since fallback is not cached, we expect another Get call - if getCalls != 2 { - t.Errorf("Expected 2 Get calls (fallback not cached), got %d", getCalls) - } + c.Eq(2, getCalls, "Expected 2 Get calls (fallback not cached), got") }) } @@ -532,50 +516,37 @@ func TestSharedHelpers(t *testing.T) { {"Claims Set", corev1.ResourceRequirements{Claims: []corev1.ResourceClaim{}}, false}, } for _, tc := range tests { - if got := isResourcesZero(tc.res); got != tc.want { - t.Errorf("%s: got %v, want %v", tc.name, got, tc.want) - } + got := isResourcesZero(tc.res) + assert.NewCollecting(t).Eq(tc.want, got, "%s: got %v, want", tc.name, got) } }) t.Run("defaultEtcdSpec", func(t *testing.T) { + c := assert.NewCollecting(t) spec := &multigresv1alpha1.EtcdSpec{} defaultEtcdSpec(spec, "/test/global") - if spec.Image != DefaultEtcdImage { - t.Errorf("Image: got %q, want %q", spec.Image, DefaultEtcdImage) - } - if spec.Storage.Size != DefaultEtcdStorageSize { - t.Errorf("Storage: got %q, want %q", spec.Storage.Size, DefaultEtcdStorageSize) - } - if *spec.Replicas != DefaultEtcdReplicas { - t.Errorf("Replicas: got %d, want %d", *spec.Replicas, DefaultEtcdReplicas) - } - if isResourcesZero(spec.Resources) { - t.Error("Resources should be defaulted") - } + c.Eq(DefaultEtcdImage, spec.Image, "Image: got") + c.Eq(DefaultEtcdStorageSize, spec.Storage.Size, "Storage: got") + c.Eq(DefaultEtcdReplicas, *spec.Replicas, "Replicas: got") + c.False(isResourcesZero(spec.Resources), "Resources should be defaulted") // Test Preservation spec2 := &multigresv1alpha1.EtcdSpec{Image: "custom"} defaultEtcdSpec(spec2, "/test/global") - if spec2.Image != "custom" { - t.Error("Should preserve existing image") - } + c.Eq("custom", spec2.Image, "Should preserve existing image") }) t.Run("defaultStatelessSpec", func(t *testing.T) { + c := assert.NewCollecting(t) spec := &multigresv1alpha1.StatelessSpec{} res := corev1.ResourceRequirements{ Requests: corev1.ResourceList{corev1.ResourceCPU: parseQty("1")}, } defaultStatelessSpec(spec, res, 5) - if *spec.Replicas != 5 { - t.Errorf("Replicas: got %d, want 5", *spec.Replicas) - } - if cmp.Diff(spec.Resources, res) != "" { - t.Error("Resources not copied correctly") - } + c.Eq(5, *spec.Replicas, "Replicas: got") + c.EqDiff(spec.Resources, res, "Resources not copied correctly") // Test DeepCopy independence res.Requests[corev1.ResourceCPU] = parseQty("999") @@ -584,9 +555,7 @@ func TestSharedHelpers(t *testing.T) { // Wait, Requests[...] returns a Value. // We need to capture it to check it. val := spec.Resources.Requests[corev1.ResourceCPU] - if val.String() == "999" { - t.Error("Shared pointer detected in defaultStatelessSpec") - } + c.NotEq("999", val.String(), "Shared pointer detected in defaultStatelessSpec") // Test Preservation spec2 := &multigresv1alpha1.StatelessSpec{ @@ -596,15 +565,12 @@ func TestSharedHelpers(t *testing.T) { }, } defaultStatelessSpec(spec2, res, 5) - if *spec2.Replicas != 10 { - t.Error("Should preserve existing Replicas") - } - if spec2.Resources.Requests.Cpu().String() != "5m" { - t.Error("Should preserve existing Resources") - } + c.Eq(10, *spec2.Replicas, "Should preserve existing Replicas") + c.Eq("5m", spec2.Resources.Requests.Cpu().String(), "Should preserve existing Resources") }) t.Run("mergeStatelessSpec", func(t *testing.T) { + c := assert.NewCollecting(t) base := &multigresv1alpha1.StatelessSpec{ PodAnnotations: map[string]string{"a": "1"}, PodLabels: map[string]string{"l1": "v1"}, @@ -622,24 +588,15 @@ func TestSharedHelpers(t *testing.T) { } mergeStatelessSpec(base, override) - if len(base.PodAnnotations) != 2 { - t.Errorf("Map merge failed (Annotations), got %v", base.PodAnnotations) - } - if len(base.PodLabels) != 2 { - t.Errorf("Map merge failed (Labels), got %v", base.PodLabels) - } - if *base.Replicas != 3 { - t.Error("Replicas not merged") - } - if base.Resources.Requests == nil { - t.Error("Resources not merged") - } - if base.Affinity == nil { - t.Error("Affinity not merged") - } + c.Len(base.PodAnnotations, 2, "Map merge failed (Annotations), got") + c.Len(base.PodLabels, 2, "Map merge failed (Labels), got") + c.Eq(3, *base.Replicas, "Replicas not merged") + c.NotNil(base.Resources.Requests, "Resources not merged") + c.NotNil(base.Affinity, "Affinity not merged") }) t.Run("mergePodPlacementSpec", func(t *testing.T) { + c := assert.NewCollecting(t) override := &multigresv1alpha1.PodPlacementSpec{ Tolerations: []corev1.Toleration{ { @@ -654,20 +611,19 @@ func TestSharedHelpers(t *testing.T) { var base *multigresv1alpha1.PodPlacementSpec mergePodPlacementSpec(&base, override) - if base == nil { - t.Fatal("Placement not initialized") - } - if len(base.Tolerations) != 1 { - t.Fatalf("Tolerations not merged, got %v", base.Tolerations) - } + c.Require().NotNil(base, "Placement not initialized") + c.Require().Len(base.Tolerations, 1, "Tolerations not merged, got") override.Tolerations[0].Value = "changed" - if base.Tolerations[0].Value != "customer-pg" { - t.Error("Tolerations should be deep-copied during merge") - } + c.Eq( + "customer-pg", + base.Tolerations[0].Value, + "Tolerations should be deep-copied during merge", + ) }) t.Run("mergePodPlacementSpec clears inherited tolerations", func(t *testing.T) { + c := assert.NewAborting(t) base := &multigresv1alpha1.PodPlacementSpec{ Tolerations: []corev1.Toleration{ { @@ -682,12 +638,8 @@ func TestSharedHelpers(t *testing.T) { override := &multigresv1alpha1.PodPlacementSpec{} mergePodPlacementSpec(&base, override) - if base == nil { - t.Fatal("Placement should remain initialized") - } - if len(base.Tolerations) != 0 { - t.Fatalf("Expected inherited tolerations to be cleared, got %v", base.Tolerations) - } + c.NotNil(base, "Placement should remain initialized") + c.Empty(base.Tolerations, "Expected inherited tolerations to be cleared, got") }) } diff --git a/pkg/resolver/shard_test.go b/pkg/resolver/shard_test.go index 7d71d7b1..f16aae09 100644 --- a/pkg/resolver/shard_test.go +++ b/pkg/resolver/shard_test.go @@ -16,6 +16,8 @@ import ( "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestResolver_ResolveShard(t *testing.T) { @@ -317,6 +319,7 @@ func TestResolver_ResolveShard(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(tc.objects...).Build() r := NewResolver(c, ns) @@ -329,35 +332,25 @@ func TestResolver_ResolveShard(t *testing.T) { }, ) if tc.wantErr { - if err == nil { - t.Error("Expected error") - } + ck.Error(err, "Expected error") return } - if err != nil { - t.Fatalf("Unexpected error: %v", err) - } + ck.Require().NoError(err, "Unexpected error") orch, pools, pvcPolicy := &resolved.Multiorch, resolved.Pools, resolved.PVCDeletionPolicy - if diff := cmp.Diff( + ck.EqDiffOpts( tc.wantOrch, orch, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Orch Diff (-want +got):\n%s", diff) - } - if diff := cmp.Diff( + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Orch Diff", + ) + ck.EqDiffOpts( tc.wantPools, pools, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Pools Diff (-want +got):\n%s", diff) - } - if diff := cmp.Diff(tc.wantPVCPolicy, pvcPolicy); diff != "" { - t.Errorf("PVC Policy Diff (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Pools Diff", + ) + ck.EqDiff(tc.wantPVCPolicy, pvcPolicy, "PVC Policy Diff") }) } } @@ -412,6 +405,7 @@ func TestResolver_ResolveShardTemplate(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder(). WithScheme(scheme). WithObjects(tc.existingObjects...). @@ -420,9 +414,7 @@ func TestResolver_ResolveShardTemplate(t *testing.T) { res, err := r.ResolveShardTemplate(t.Context(), tc.reqName) if tc.wantErr { - if err == nil { - t.Fatal("Expected error, got nil") - } + ck.Require().Error(err, "Expected error, got nil") if tc.errContains != "" && !strings.Contains(err.Error(), tc.errContains) { t.Errorf( "Error message mismatch: got %q, want substring %q", @@ -436,25 +428,20 @@ func TestResolver_ResolveShardTemplate(t *testing.T) { } if !tc.wantFound { - if res == nil { - t.Fatal( - "Expected non-nil result structure even for not-found implicit fallback", - ) - } - if res.GetName() != "" { - t.Errorf("Expected empty result, got object with name %q", res.GetName()) - } + ck.Require(). + NotNil(res, "Expected non-nil result structure even for not-found implicit fallback") + ck.Eq("", res.GetName(), "Expected empty result, got object with name") return } - if got, want := res.GetName(), tc.wantResName; got != want { - t.Errorf("Result name mismatch: got %q, want %q", got, want) - } + got, want := res.GetName(), tc.wantResName + ck.Eq(want, got, "Result name mismatch: got") }) } } func TestMergePoolSpec_RuntimeIdentity(t *testing.T) { + c := assert.NewCollecting(t) base := multigresv1alpha1.PoolSpec{ Postgres: multigresv1alpha1.ContainerConfig{ RunAsUser: ptr.To(int64(1000)), @@ -481,12 +468,20 @@ func TestMergePoolSpec_RuntimeIdentity(t *testing.T) { wantUser, wantGroup int64, ) { t.Helper() - if config.RunAsUser == nil || *config.RunAsUser != wantUser { - t.Errorf("%s runAsUser = %v, want %d", name, config.RunAsUser, wantUser) - } - if config.RunAsGroup == nil || *config.RunAsGroup != wantGroup { - t.Errorf("%s runAsGroup = %v, want %d", name, config.RunAsGroup, wantGroup) - } + c.False( + config.RunAsUser == nil || *config.RunAsUser != wantUser, + "%s runAsUser = %v, want %d", + name, + config.RunAsUser, + wantUser, + ) + c.False( + config.RunAsGroup == nil || *config.RunAsGroup != wantGroup, + "%s runAsGroup = %v, want %d", + name, + config.RunAsGroup, + wantGroup, + ) } assertPoolIdentity("postgres", got.Postgres, 1000, 2001) assertPoolIdentity("multipooler", got.Multipooler, 1000, 2002) @@ -502,6 +497,7 @@ func TestMergeShardConfig_RuntimeIdentityPartialOverride(t *testing.T) { t.Run("override multipooler UID can match template postgres UID", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) resolved := mergeShardConfig( &multigresv1alpha1.ShardTemplate{ @@ -530,16 +526,21 @@ func TestMergeShardConfig_RuntimeIdentityPartialOverride(t *testing.T) { ) got := resolved.Pools["rw"] - if got.Postgres.RunAsUser == nil || *got.Postgres.RunAsUser != 1000 { - t.Fatalf("postgres runAsUser = %v, want 1000", got.Postgres.RunAsUser) - } - if got.Multipooler.RunAsUser == nil || *got.Multipooler.RunAsUser != 1000 { - t.Fatalf("multipooler runAsUser = %v, want 1000", got.Multipooler.RunAsUser) - } + c.False( + got.Postgres.RunAsUser == nil || *got.Postgres.RunAsUser != 1000, + "postgres runAsUser = %v, want 1000", + got.Postgres.RunAsUser, + ) + c.False( + got.Multipooler.RunAsUser == nil || *got.Multipooler.RunAsUser != 1000, + "multipooler runAsUser = %v, want 1000", + got.Multipooler.RunAsUser, + ) }) t.Run("mismatched override remains visible for resolved validation", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) resolved := mergeShardConfig( &multigresv1alpha1.ShardTemplate{ @@ -568,12 +569,16 @@ func TestMergeShardConfig_RuntimeIdentityPartialOverride(t *testing.T) { ) got := resolved.Pools["rw"] - if got.Postgres.RunAsUser == nil || *got.Postgres.RunAsUser != 1000 { - t.Fatalf("postgres runAsUser = %v, want 1000", got.Postgres.RunAsUser) - } - if got.Multipooler.RunAsUser == nil || *got.Multipooler.RunAsUser != 2000 { - t.Fatalf("multipooler runAsUser = %v, want 2000", got.Multipooler.RunAsUser) - } + c.False( + got.Postgres.RunAsUser == nil || *got.Postgres.RunAsUser != 1000, + "postgres runAsUser = %v, want 1000", + got.Postgres.RunAsUser, + ) + c.False( + got.Multipooler.RunAsUser == nil || *got.Multipooler.RunAsUser != 2000, + "multipooler runAsUser = %v, want 2000", + got.Multipooler.RunAsUser, + ) }) } @@ -988,6 +993,7 @@ func TestMergeShardConfig(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) resolved := mergeShardConfig( tc.tpl, tc.overrides, @@ -997,21 +1003,18 @@ func TestMergeShardConfig(t *testing.T) { ) orch, pools := resolved.Multiorch, resolved.Pools - if diff := cmp.Diff( + c.EqDiffOpts( tc.wantOrch, orch, - cmpopts.IgnoreUnexported(resource.Quantity{}), - ); diff != "" { - t.Errorf("Orch mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff( + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{})}, + "Orch mismatch", + ) + c.EqDiffOpts( tc.wantPools, pools, - cmpopts.IgnoreUnexported(resource.Quantity{}), - cmpopts.EquateEmpty(), - ); diff != "" { - t.Errorf("Pools mismatch (-want +got):\n%s", diff) - } + []cmp.Option{cmpopts.IgnoreUnexported(resource.Quantity{}), cmpopts.EquateEmpty()}, + "Pools mismatch", + ) }) } } @@ -1029,9 +1032,7 @@ func TestMergeShardConfig_InitdbArgs(t *testing.T) { }, nil, nil, nil, nil, ).InitdbArgs - if initdbArgs != "--locale-provider=icu" { - t.Errorf("initdbArgs = %q, want %q", initdbArgs, "--locale-provider=icu") - } + assert.NewCollecting(t).Eq("--locale-provider=icu", initdbArgs, "initdbArgs") }) t.Run("overrides override template", func(t *testing.T) { @@ -1047,9 +1048,7 @@ func TestMergeShardConfig_InitdbArgs(t *testing.T) { }, nil, nil, nil, ).InitdbArgs - if initdbArgs != "--data-checksums" { - t.Errorf("initdbArgs = %q, want %q", initdbArgs, "--data-checksums") - } + assert.NewCollecting(t).Eq("--data-checksums", initdbArgs, "initdbArgs") }) t.Run("inline overrides template", func(t *testing.T) { @@ -1066,9 +1065,7 @@ func TestMergeShardConfig_InitdbArgs(t *testing.T) { }, nil, nil, ).InitdbArgs - if initdbArgs != "--data-checksums" { - t.Errorf("initdbArgs = %q, want %q", initdbArgs, "--data-checksums") - } + assert.NewCollecting(t).Eq("--data-checksums", initdbArgs, "initdbArgs") }) t.Run("inline overrides both template and overrides", func(t *testing.T) { @@ -1087,9 +1084,7 @@ func TestMergeShardConfig_InitdbArgs(t *testing.T) { }, nil, nil, ).InitdbArgs - if initdbArgs != "--wal-segsize=64" { - t.Errorf("initdbArgs = %q, want %q", initdbArgs, "--wal-segsize=64") - } + assert.NewCollecting(t).Eq("--wal-segsize=64", initdbArgs, "initdbArgs") }) t.Run("no InitdbArgs anywhere", func(t *testing.T) { @@ -1098,9 +1093,7 @@ func TestMergeShardConfig_InitdbArgs(t *testing.T) { &multigresv1alpha1.ShardTemplate{}, nil, nil, nil, nil, ).InitdbArgs - if initdbArgs != "" { - t.Errorf("initdbArgs = %q, want empty", initdbArgs) - } + assert.NewCollecting(t).Eq("", initdbArgs, "initdbArgs") }) t.Run("empty override does not clear template value", func(t *testing.T) { @@ -1116,10 +1109,7 @@ func TestMergeShardConfig_InitdbArgs(t *testing.T) { }, nil, nil, nil, ).InitdbArgs - if initdbArgs != "--locale-provider=icu" { - t.Errorf("initdbArgs = %q, want %q (empty override should not clear template)", - initdbArgs, "--locale-provider=icu") - } + assert.NewCollecting(t).Eq("--locale-provider=icu", initdbArgs, "initdbArgs") }) } @@ -1137,10 +1127,8 @@ func TestResolver_ClientErrors_Shard(t *testing.T) { r := NewResolver(mc, "default") _, err := r.ResolveShardTemplate(t.Context(), "any") - if err == nil || - err.Error() != "failed to get ShardTemplate: simulated database connection error" { - t.Errorf("Error mismatch: got %v, want simulated error", err) - } + assert.NewCollecting(t).False(err == nil || + err.Error() != "failed to get ShardTemplate: simulated database connection error", "Error mismatch: got %v, want simulated error", err) } func TestResolveShard_PVCDeletionPolicy(t *testing.T) { @@ -1148,6 +1136,7 @@ func TestResolveShard_PVCDeletionPolicy(t *testing.T) { _ = multigresv1alpha1.AddToScheme(scheme) t.Run("From Template", func(t *testing.T) { + c := assert.NewCollecting(t) r := &Resolver{ Client: fake.NewClientBuilder(). WithScheme(scheme). @@ -1167,16 +1156,17 @@ func TestResolveShard_PVCDeletionPolicy(t *testing.T) { resolved, err := r.ResolveShard(t.Context(), &multigresv1alpha1.ShardConfig{ ShardTemplate: "tpl-pvc", }, ResolveShardOptions{}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") policy := resolved.PVCDeletionPolicy - if policy == nil || policy.WhenDeleted != multigresv1alpha1.DeletePVCRetentionPolicy { - t.Errorf("Expected Template PVCDeletionPolicy=Delete, got %v", policy) - } + c.False( + policy == nil || policy.WhenDeleted != multigresv1alpha1.DeletePVCRetentionPolicy, + "Expected Template PVCDeletionPolicy=Delete, got %v", + policy, + ) }) t.Run("Pool Level Override", func(t *testing.T) { + c := assert.NewCollecting(t) r := &Resolver{ Client: fake.NewClientBuilder().WithScheme(scheme).Build(), Namespace: "default", @@ -1195,17 +1185,13 @@ func TestResolveShard_PVCDeletionPolicy(t *testing.T) { }, }, }, ResolveShardOptions{}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") pools := resolved.Pools if p, ok := pools["custom-pool"]; !ok { t.Fatal("Expected custom-pool to exist") } else { - if p.PVCDeletionPolicy == nil || - p.PVCDeletionPolicy.WhenDeleted != multigresv1alpha1.RetainPVCRetentionPolicy { - t.Errorf("Expected Pool PVCDeletionPolicy=Retain, got %v", p.PVCDeletionPolicy) - } + c.False(p.PVCDeletionPolicy == nil || + p.PVCDeletionPolicy.WhenDeleted != multigresv1alpha1.RetainPVCRetentionPolicy, "Expected Pool PVCDeletionPolicy=Retain, got %v", p.PVCDeletionPolicy) } }) } @@ -1220,9 +1206,7 @@ func TestDefaultBackupConfig(t *testing.T) { Filesystem: &multigresv1alpha1.FilesystemBackupConfig{}, } defaultBackupConfig(cfg) - if cfg.Filesystem.Path != DefaultBackupPath { - t.Errorf("Path = %q, want %q", cfg.Filesystem.Path, DefaultBackupPath) - } + assert.NewCollecting(t).Eq(DefaultBackupPath, cfg.Filesystem.Path, "Path") }) t.Run("sets default storage size", func(t *testing.T) { @@ -1232,17 +1216,13 @@ func TestDefaultBackupConfig(t *testing.T) { Filesystem: &multigresv1alpha1.FilesystemBackupConfig{}, } defaultBackupConfig(cfg) - if cfg.Filesystem.Storage.Size != DefaultBackupStorageSize { - t.Errorf( - "Storage.Size = %q, want %q", - cfg.Filesystem.Storage.Size, - DefaultBackupStorageSize, - ) - } + assert.NewCollecting(t). + Eq(DefaultBackupStorageSize, cfg.Filesystem.Storage.Size, "Storage.Size") }) t.Run("does not override existing values", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cfg := &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ @@ -1251,26 +1231,19 @@ func TestDefaultBackupConfig(t *testing.T) { }, } defaultBackupConfig(cfg) - if cfg.Filesystem.Path != "/custom" { - t.Errorf("Path = %q, want /custom", cfg.Filesystem.Path) - } - if cfg.Filesystem.Storage.Size != "50Gi" { - t.Errorf("Storage.Size = %q, want 50Gi", cfg.Filesystem.Storage.Size) - } + c.Eq("/custom", cfg.Filesystem.Path, "Path") + c.Eq("50Gi", cfg.Filesystem.Storage.Size, "Storage.Size") }) t.Run("creates filesystem struct if nil", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cfg := &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, } defaultBackupConfig(cfg) - if cfg.Filesystem == nil { - t.Fatal("Filesystem = nil, want non-nil") - } - if cfg.Filesystem.Path != DefaultBackupPath { - t.Errorf("Path = %q, want %q", cfg.Filesystem.Path, DefaultBackupPath) - } + c.Require().NotNil(cfg.Filesystem, "Filesystem = nil, want non-nil") + c.Eq(DefaultBackupPath, cfg.Filesystem.Path, "Path") }) t.Run("does not touch s3 config", func(t *testing.T) { @@ -1281,21 +1254,16 @@ func TestDefaultBackupConfig(t *testing.T) { } defaultBackupConfig(cfg) // Should not create Filesystem struct for S3 type - if cfg.Filesystem != nil { - t.Errorf("Filesystem should be nil for S3 type, got %+v", cfg.Filesystem) - } + assert.NewCollecting(t).Nil(cfg.Filesystem, "Filesystem should be nil for S3 type, got") }) t.Run("sets default type when empty", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cfg := &multigresv1alpha1.BackupConfig{} defaultBackupConfig(cfg) - if cfg.Type != multigresv1alpha1.BackupTypeFilesystem { - t.Errorf("Type = %q, want %q", cfg.Type, multigresv1alpha1.BackupTypeFilesystem) - } - if cfg.Filesystem == nil { - t.Fatal("Filesystem = nil, want non-nil") - } + c.Eq(multigresv1alpha1.BackupTypeFilesystem, cfg.Type, "Type") + c.Require().NotNil(cfg.Filesystem, "Filesystem = nil, want non-nil") }) } @@ -1306,6 +1274,7 @@ func TestResolveShard_InheritedBackup(t *testing.T) { t.Run("inherited backup propagates to resolved config", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme).Build() r := NewResolver(c, "default") @@ -1328,23 +1297,16 @@ func TestResolveShard_InheritedBackup(t *testing.T) { }, ResolveShardOptions{InheritedBackup: inherited}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") backupCfg := resolved.Backup - if backupCfg == nil { - t.Fatal("backup config should not be nil") - } - if backupCfg.Type != multigresv1alpha1.BackupTypeFilesystem { - t.Errorf("Type = %q, want filesystem", backupCfg.Type) - } - if backupCfg.Filesystem.Path != "/inherited-path" { - t.Errorf("Path = %q, want /inherited-path", backupCfg.Filesystem.Path) - } + ck.Require().NotNil(backupCfg, "backup config should not be nil") + ck.Eq(multigresv1alpha1.BackupTypeFilesystem, backupCfg.Type, "Type") + ck.Eq("/inherited-path", backupCfg.Filesystem.Path, "Path") }) t.Run("shard backup overrides inherited", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme).Build() r := NewResolver(c, "default") @@ -1372,17 +1334,14 @@ func TestResolveShard_InheritedBackup(t *testing.T) { }, ResolveShardOptions{InheritedBackup: inherited}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") backupCfg := resolved.Backup - if backupCfg.Filesystem.Path != "/shard-override" { - t.Errorf("Path = %q, want /shard-override", backupCfg.Filesystem.Path) - } + ck.Eq("/shard-override", backupCfg.Filesystem.Path, "Path") }) t.Run("nil inherited gets filesystem default", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme).Build() r := NewResolver(c, "default") @@ -1397,19 +1356,11 @@ func TestResolveShard_InheritedBackup(t *testing.T) { }, ResolveShardOptions{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") backupCfg := resolved.Backup - if backupCfg == nil { - t.Fatal("backup config should not be nil (should get defaults)") - } - if backupCfg.Type != multigresv1alpha1.BackupTypeFilesystem { - t.Errorf("Type = %q, want filesystem", backupCfg.Type) - } - if backupCfg.Filesystem.Path != DefaultBackupPath { - t.Errorf("Path = %q, want %q", backupCfg.Filesystem.Path, DefaultBackupPath) - } + ck.Require().NotNil(backupCfg, "backup config should not be nil (should get defaults)") + ck.Eq(multigresv1alpha1.BackupTypeFilesystem, backupCfg.Type, "Type") + ck.Eq(DefaultBackupPath, backupCfg.Filesystem.Path, "Path") }) } @@ -1433,9 +1384,8 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { }, nil, nil, nil, nil, ).PostgresConfigRef - if ref == nil || ref.Name != "template-config" || ref.Key != "postgresql.conf" { - t.Errorf("postgresConfigRef = %v, want %v", ref, templateRef) - } + assert.NewCollecting(t). + False(ref == nil || ref.Name != "template-config" || ref.Key != "postgresql.conf", "postgresConfigRef = %v, want %v", ref, templateRef) }) t.Run("overrides replace template ref", func(t *testing.T) { @@ -1451,9 +1401,8 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { }, nil, nil, nil, ).PostgresConfigRef - if ref == nil || ref.Name != "override-config" || ref.Key != "custom.conf" { - t.Errorf("postgresConfigRef = %v, want %v", ref, overrideRef) - } + assert.NewCollecting(t). + False(ref == nil || ref.Name != "override-config" || ref.Key != "custom.conf", "postgresConfigRef = %v, want %v", ref, overrideRef) }) t.Run("inline replaces template and overrides", func(t *testing.T) { @@ -1472,9 +1421,8 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { }, nil, nil, ).PostgresConfigRef - if ref == nil || ref.Name != "inline-config" || ref.Key != "inline.conf" { - t.Errorf("postgresConfigRef = %v, want %v", ref, inlineRef) - } + assert.NewCollecting(t). + False(ref == nil || ref.Name != "inline-config" || ref.Key != "inline.conf", "postgresConfigRef = %v, want %v", ref, inlineRef) }) t.Run("nil everywhere returns nil", func(t *testing.T) { @@ -1483,9 +1431,7 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { &multigresv1alpha1.ShardTemplate{}, nil, nil, nil, nil, ).PostgresConfigRef - if ref != nil { - t.Errorf("postgresConfigRef = %v, want nil", ref) - } + assert.NewCollecting(t).Nil(ref, "postgresConfigRef") }) t.Run("only overrides set ref", func(t *testing.T) { @@ -1497,9 +1443,8 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { }, nil, nil, nil, ).PostgresConfigRef - if ref == nil || ref.Name != "override-config" { - t.Errorf("postgresConfigRef = %v, want %v", ref, overrideRef) - } + assert.NewCollecting(t). + False(ref == nil || ref.Name != "override-config", "postgresConfigRef = %v, want %v", ref, overrideRef) }) t.Run("only inline sets ref", func(t *testing.T) { @@ -1511,9 +1456,8 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { }, nil, nil, ).PostgresConfigRef - if ref == nil || ref.Name != "inline-config" { - t.Errorf("postgresConfigRef = %v, want %v", ref, inlineRef) - } + assert.NewCollecting(t). + False(ref == nil || ref.Name != "inline-config", "postgresConfigRef = %v, want %v", ref, inlineRef) }) t.Run("nil overrides do not clear template ref", func(t *testing.T) { @@ -1527,13 +1471,8 @@ func TestMergeShardConfig_PostgresConfigRef(t *testing.T) { &multigresv1alpha1.ShardOverrides{}, nil, nil, nil, ).PostgresConfigRef - if ref == nil || ref.Name != "template-config" { - t.Errorf( - "postgresConfigRef = %v, want %v (nil override should not clear template)", - ref, - templateRef, - ) - } + assert.NewCollecting(t). + False(ref == nil || ref.Name != "template-config", "postgresConfigRef = %v, want %v (nil override should not clear template)", ref, templateRef) }) } @@ -1546,13 +1485,12 @@ func TestMergeShardConfig_PostgresConfig(t *testing.T) { &multigresv1alpha1.ShardTemplate{}, nil, nil, nil, nil, ).PostgresConfig - if cfg != nil { - t.Errorf("postgresConfig = %v, want nil", cfg) - } + assert.NewCollecting(t).Nil(cfg, "postgresConfig") }) t.Run("layers merge per key with inline winning", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cfg := mergeShardConfig( &multigresv1alpha1.ShardTemplate{ Spec: multigresv1alpha1.ShardTemplateSpec{ @@ -1580,13 +1518,9 @@ func TestMergeShardConfig_PostgresConfig(t *testing.T) { "shared_buffers": "2GB", // only in template "work_mem": "8MB", // only in override } - if len(cfg) != len(want) { - t.Fatalf("postgresConfig = %v, want %v", cfg, want) - } + c.Require().Len(cfg, len(want), "postgresConfig = %v, want %v", cfg, want) for k, v := range want { - if cfg[k] != v { - t.Errorf("postgresConfig[%q] = %q, want %q", k, cfg[k], v) - } + c.Eq(v, cfg[k], "postgresConfig[%q] = %q, want", k, cfg[k]) } }) @@ -1603,8 +1537,7 @@ func TestMergeShardConfig_PostgresConfig(t *testing.T) { }, nil, nil, nil, ) - if tplMap["max_connections"] != "100" { - t.Errorf("template map mutated: %v", tplMap) - } + assert.NewCollecting(t). + Eq("100", tplMap["max_connections"], "template map mutated: %v", tplMap) }) } diff --git a/pkg/resolver/validation_test.go b/pkg/resolver/validation_test.go index 9e5ae9fa..2af58be4 100644 --- a/pkg/resolver/validation_test.go +++ b/pkg/resolver/validation_test.go @@ -13,6 +13,8 @@ import ( "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) // TestResolver_ValidateClusterIntegrity verifies checking template existence. @@ -139,17 +141,19 @@ func TestResolver_ValidateClusterIntegrity(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) r := NewResolver(fakeClient, "default") err := r.ValidateClusterIntegrity(t.Context(), tc.cluster) if tc.wantErr == "" { - if err != nil { - t.Errorf("Expected nil error, got %v", err) - } + c.NoError(err, "Expected nil error, got") } else { - if err == nil || !strings.Contains(err.Error(), tc.wantErr) { - t.Errorf("Expected error containing '%s', got %v", tc.wantErr, err) - } + c.False( + err == nil || !strings.Contains(err.Error(), tc.wantErr), + "Expected error containing '%s', got %v", + tc.wantErr, + err, + ) } }) } @@ -1782,6 +1786,7 @@ func TestResolver_ValidateClusterLogic(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + ck := assert.NewCollecting(t) var c client.Client = fakeClient if tc.customClient != nil { c = tc.customClient @@ -1794,20 +1799,19 @@ func TestResolver_ValidateClusterLogic(t *testing.T) { // Check Error if tc.wantErr == "" { - if err != nil { - t.Errorf("Expected nil error, got %v", err) - } + ck.NoError(err, "Expected nil error, got") } else { - if err == nil || !strings.Contains(err.Error(), tc.wantErr) { - t.Errorf("Expected error containing '%s', got %v", tc.wantErr, err) - } + ck.False( + err == nil || !strings.Contains(err.Error(), tc.wantErr), + "Expected error containing '%s', got %v", + tc.wantErr, + err, + ) } // Check Warnings if len(tc.wantWarnings) > 0 { - if len(warnings) == 0 { - t.Errorf("Expected warnings containing %v, got none", tc.wantWarnings) - } + ck.NotEmpty(warnings, "Expected warnings containing %v, got none", tc.wantWarnings) for _, want := range tc.wantWarnings { found := false for _, got := range warnings { @@ -1816,14 +1820,14 @@ func TestResolver_ValidateClusterLogic(t *testing.T) { break } } - if !found { - t.Errorf("Expected warning containing '%s', got %v", want, warnings) - } + ck.True(found, "Expected warning containing '%s', got %v", want, warnings) } } - if tc.wantNoWarnings && len(warnings) > 0 { - t.Errorf("Expected no warnings, got %v", warnings) - } + ck.False( + tc.wantNoWarnings && len(warnings) > 0, + "Expected no warnings, got %v", + warnings, + ) }) } } diff --git a/pkg/resource-handler/controller/cell/cell_controller_internal_test.go b/pkg/resource-handler/controller/cell/cell_controller_internal_test.go index 3162bfa6..dda557f6 100644 --- a/pkg/resource-handler/controller/cell/cell_controller_internal_test.go +++ b/pkg/resource-handler/controller/cell/cell_controller_internal_test.go @@ -21,6 +21,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) // TestReconcileMultigatewayDeployment_InvalidScheme tests the error path when BuildMultigatewayDeployment fails. @@ -51,9 +53,8 @@ func TestReconcileMultigatewayDeployment_InvalidScheme(t *testing.T) { } err := reconciler.reconcileMultigatewayDeployment(context.Background(), cell) - if err == nil { - t.Error("reconcileMultigatewayDeployment() should error with invalid scheme") - } + assert.NewCollecting(t). + Error(err, "reconcileMultigatewayDeployment() should error with invalid scheme") } // TestReconcileMultigatewayService_InvalidScheme tests the error path when BuildMultigatewayService fails. @@ -81,9 +82,8 @@ func TestReconcileMultigatewayService_InvalidScheme(t *testing.T) { } err := reconciler.reconcileMultigatewayService(context.Background(), cell) - if err == nil { - t.Error("reconcileMultigatewayService() should error with invalid scheme") - } + assert.NewCollecting(t). + Error(err, "reconcileMultigatewayService() should error with invalid scheme") } // TestUpdateStatus_MultigatewayDeploymentNotFound tests the NotFound path in updateStatus. @@ -116,12 +116,8 @@ func TestUpdateStatus_MultigatewayDeploymentNotFound(t *testing.T) { // Call updateStatus when Multigateway Deployment doesn't exist yet err := reconciler.updateStatus(context.Background(), cell) - if err != nil { - t.Errorf( - "updateStatus() should not error when Multigateway Deployment not found, got: %v", - err, - ) - } + assert.NewCollecting(t). + NoError(err, "updateStatus() should not error when Multigateway Deployment not found, got") } // TestReconcileMultigatewayDeployment_PatchError tests error path on Patch Multigateway Deployment. @@ -163,9 +159,8 @@ func TestReconcileMultigatewayDeployment_PatchError(t *testing.T) { } err := reconciler.reconcileMultigatewayDeployment(context.Background(), cell) - if err == nil { - t.Error("reconcileMultigatewayDeployment() should error on Patch failure") - } + assert.NewCollecting(t). + Error(err, "reconcileMultigatewayDeployment() should error on Patch failure") } // TestReconcileMultigatewayService_PatchError tests error path on Patch Multigateway Service. @@ -207,9 +202,8 @@ func TestReconcileMultigatewayService_PatchError(t *testing.T) { } err := reconciler.reconcileMultigatewayService(context.Background(), cell) - if err == nil { - t.Error("reconcileMultigatewayService() should error on Patch failure") - } + assert.NewCollecting(t). + Error(err, "reconcileMultigatewayService() should error on Patch failure") } // TestUpdateStatus_GetError tests error path on Get Multigateway Deployment (not NotFound). @@ -250,13 +244,12 @@ func TestUpdateStatus_GetError(t *testing.T) { } err := reconciler.updateStatus(context.Background(), cell) - if err == nil { - t.Error("updateStatus() should error on Get failure") - } + assert.NewCollecting(t).Error(err, "updateStatus() should error on Get failure") } // TestSetConditions_ZeroReplicas tests setConditions when deployments have zero replicas. func TestSetConditions_ZeroReplicas(t *testing.T) { + c := assert.NewCollecting(t) reconciler := &CellReconciler{ Recorder: record.NewFakeRecorder(100), } @@ -282,31 +275,18 @@ func TestSetConditions_ZeroReplicas(t *testing.T) { reconciler.setConditions(cell, mgDeploy) conditions := cell.Status.Conditions - if len(conditions) != 2 { - t.Fatalf("setConditions() should set 2 conditions, got %d", len(conditions)) - } + c.Require(). + Len(conditions, 2, "setConditions() should set 2 conditions, got %d", len(conditions)) availCond := conditions[0] - if availCond.Type != "Available" { - t.Errorf("Condition type = %s, want Available", availCond.Type) - } - if availCond.Status != metav1.ConditionFalse { - t.Errorf("Condition status = %s, want False (zero replicas)", availCond.Status) - } - if availCond.Reason != "MultigatewayUnavailable" { - t.Errorf("Condition reason = %s, want MultigatewayUnavailable", availCond.Reason) - } + c.Eq("Available", availCond.Type, "Condition type") + c.Eq(metav1.ConditionFalse, availCond.Status, "Condition status") + c.Eq("MultigatewayUnavailable", availCond.Reason, "Condition reason") readyCond := conditions[1] - if readyCond.Type != "Ready" { - t.Errorf("Condition type = %s, want Ready", readyCond.Type) - } - if readyCond.Status != metav1.ConditionFalse { - t.Errorf("Condition status = %s, want False", readyCond.Status) - } - if readyCond.Reason != "MultigatewayNotReady" { - t.Errorf("Condition reason = %s, want MultigatewayNotReady", readyCond.Reason) - } + c.Eq("Ready", readyCond.Type, "Condition type") + c.Eq(metav1.ConditionFalse, readyCond.Status, "Condition status") + c.Eq("MultigatewayNotReady", readyCond.Reason, "Condition reason") } // TestSetupWithManager tests the manager setup function. @@ -324,9 +304,7 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Metrics: metricsserver.Options{BindAddress: "0"}, }) - if err != nil { - t.Fatalf("Failed to create manager: %v", err) - } + assert.NewAborting(t).NoError(err, "Failed to create manager") return mgr } @@ -337,9 +315,7 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Recorder: record.NewFakeRecorder(100), } - if err := r.SetupWithManager(mgr); err != nil { - t.Errorf("SetupWithManager() error = %v", err) - } + assert.NewCollecting(t).NoError(r.SetupWithManager(mgr), "SetupWithManager() error =") }) t.Run("with options", func(t *testing.T) { @@ -349,16 +325,15 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Recorder: record.NewFakeRecorder(100), } - if err := r.SetupWithManager(mgr, controller.Options{ + assert.NewCollecting(t).NoError(r.SetupWithManager(mgr, controller.Options{ MaxConcurrentReconciles: 1, SkipNameValidation: ptr.To(true), - }); err != nil { - t.Errorf("SetupWithManager() with opts error = %v", err) - } + }), "SetupWithManager() with opts error =") }) } func TestUpdateStatus_DegradedOnCrashLoop(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -419,10 +394,7 @@ func TestUpdateStatus_DegradedOnCrashLoop(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := reconciler.updateStatus(context.Background(), cell); err != nil { - t.Fatalf("updateStatus() unexpected error: %v", err) - } - if cell.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("expected PhaseDegraded, got %q", cell.Status.Phase) - } + c.Require(). + NoError(reconciler.updateStatus(context.Background(), cell), "updateStatus() unexpected error") + c.Eq(multigresv1alpha1.PhaseDegraded, cell.Status.Phase, "expected PhaseDegraded, got") } diff --git a/pkg/resource-handler/controller/cell/cell_controller_test.go b/pkg/resource-handler/controller/cell/cell_controller_test.go index eefef580..2fd15bd1 100644 --- a/pkg/resource-handler/controller/cell/cell_controller_test.go +++ b/pkg/resource-handler/controller/cell/cell_controller_test.go @@ -19,6 +19,8 @@ import ( "github.com/multigres/multigres-operator/pkg/resource-handler/controller/cell" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) // buildHashedName helper to generate the expected hashed name for tests @@ -37,20 +39,13 @@ type conditionAssertion struct { // assertConditions verifies conditions match expectations func assertConditions(t testing.TB, got []metav1.Condition, want ...conditionAssertion) { t.Helper() - if len(got) != len(want) { - t.Fatalf("condition count = %d, want %d", len(got), len(want)) - } + c := assert.NewCollecting(t) + c.Require().Len(got, len(want), "condition count = %d, want", len(got)) for i, w := range want { g := got[i] - if g.Type != w.Type { - t.Errorf("condition[%d].Type = %q, want %q", i, g.Type, w.Type) - } - if g.Status != w.Status { - t.Errorf("condition[%d].Status = %q, want %q", i, g.Status, w.Status) - } - if g.Reason != w.Reason { - t.Errorf("condition[%d].Reason = %q, want %q", i, g.Reason, w.Reason) - } + c.Eq(w.Type, g.Type, "condition[%d].Type = %q, want", i, g.Type) + c.Eq(w.Status, g.Status, "condition[%d].Status = %q, want", i, g.Status) + c.Eq(w.Reason, g.Reason, "condition[%d].Reason = %q, want", i, g.Reason) } } @@ -90,35 +85,26 @@ func TestCellReconciler_Reconcile(t *testing.T) { existingObjects: []client.Object{}, wantRequeue: true, assertFunc: func(t *testing.T, c client.Client, cell *multigresv1alpha1.Cell) { + ck := assert.NewCollecting(t) hashedName := buildHashedName( cell.Labels["multigres.com/cluster"], string(cell.Spec.Name), ) // Verify Multigateway Deployment was created mgDeploy := &appsv1.Deployment{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedName, Namespace: "default"}, - mgDeploy); err != nil { - t.Errorf("Multigateway Deployment should exist: %v", err) - } + mgDeploy), "Multigateway Deployment should exist") // Verify Multigateway Service was created mgSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedName, Namespace: "default"}, - mgSvc); err != nil { - t.Errorf("Multigateway Service should exist: %v", err) - } + mgSvc), "Multigateway Service should exist") // Verify defaults const wantReplicas int32 = 1 - if *mgDeploy.Spec.Replicas != wantReplicas { - t.Errorf( - "Multigateway Deployment replicas = %d, want %d", - *mgDeploy.Spec.Replicas, - wantReplicas, - ) - } + ck.Eq(wantReplicas, *mgDeploy.Spec.Replicas, "Multigateway Deployment replicas") }, }, "update existing resources": { @@ -161,6 +147,7 @@ func TestCellReconciler_Reconcile(t *testing.T) { }, }, assertFunc: func(t *testing.T, c client.Client, cell *multigresv1alpha1.Cell) { + ck := assert.NewCollecting(t) hashedName := buildHashedName( cell.Labels["multigres.com/cluster"], string(cell.Spec.Name), @@ -170,26 +157,17 @@ func TestCellReconciler_Reconcile(t *testing.T) { Name: hashedName, Namespace: "default", }, mgDeploy) - if err != nil { - t.Fatalf("Failed to get Multigateway Deployment: %v", err) - } + ck.Require().NoError(err, "Failed to get Multigateway Deployment") - if *mgDeploy.Spec.Replicas != 5 { - t.Errorf( - "Multigateway Deployment replicas = %d, want 5", - *mgDeploy.Spec.Replicas, - ) - } + ck.Eq(5, *mgDeploy.Spec.Replicas, "Multigateway Deployment replicas") - if len(mgDeploy.Spec.Template.Spec.Containers) == 0 { - t.Fatal("Multigateway Deployment has no containers") - } - if mgDeploy.Spec.Template.Spec.Containers[0].Image != "custom/multigateway:v1.0.0" { - t.Errorf( - "Multigateway image = %s, want custom/multigateway:v1.0.0", - mgDeploy.Spec.Template.Spec.Containers[0].Image, - ) - } + ck.Require(). + NotEmpty(mgDeploy.Spec.Template.Spec.Containers, "Multigateway Deployment has no containers") + ck.Eq( + "custom/multigateway:v1.0.0", + mgDeploy.Spec.Template.Spec.Containers[0].Image, + "Multigateway image", + ) }, }, @@ -224,11 +202,9 @@ func TestCellReconciler_Reconcile(t *testing.T) { string(cell.Spec.Name), ) mgDeploy := &appsv1.Deployment{} - if err := c.Get(t.Context(), + assert.NewCollecting(t).Error(c.Get(t.Context(), types.NamespacedName{Name: hashedName, Namespace: "default"}, - mgDeploy); err == nil { - t.Errorf("Multigateway Deployment should NOT exist") - } + mgDeploy), "Multigateway Deployment should NOT exist") }, }, @@ -269,12 +245,11 @@ func TestCellReconciler_Reconcile(t *testing.T) { }, }, assertFunc: func(t *testing.T, c client.Client, cell *multigresv1alpha1.Cell) { + ck := assert.NewCollecting(t) updatedCell := &multigresv1alpha1.Cell{} - if err := c.Get(t.Context(), + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-cell-ready", Namespace: "default"}, - updatedCell); err != nil { - t.Fatalf("Failed to get Cell: %v", err) - } + updatedCell), "Failed to get Cell") assertConditions( t, @@ -294,9 +269,8 @@ func TestCellReconciler_Reconcile(t *testing.T) { if got, want := updatedCell.Status.GatewayReplicas, int32(2); got != want { t.Errorf("GatewayReplicas = %d, want %d", got, want) } - if got, want := updatedCell.Status.GatewayReadyReplicas, int32(2); got != want { - t.Errorf("GatewayReadyReplicas = %d, want %d", got, want) - } + got, want := updatedCell.Status.GatewayReadyReplicas, int32(2) + ck.Eq(want, got, "GatewayReadyReplicas") }, }, "not ready status - partial replicas": { @@ -338,11 +312,9 @@ func TestCellReconciler_Reconcile(t *testing.T) { }, assertFunc: func(t *testing.T, c client.Client, cell *multigresv1alpha1.Cell) { updatedCell := &multigresv1alpha1.Cell{} - if err := c.Get(t.Context(), + assert.NewAborting(t).NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-cell-partial", Namespace: "default"}, - updatedCell); err != nil { - t.Fatalf("Failed to get Cell: %v", err) - } + updatedCell), "Failed to get Cell") // With 2/3 replicas ready: Available=True (service is up), Ready=False (not converged) assertConditions( @@ -381,25 +353,24 @@ func TestCellReconciler_Reconcile(t *testing.T) { }, existingObjects: []client.Object{}, assertFunc: func(t *testing.T, c client.Client, cell *multigresv1alpha1.Cell) { + ck := assert.NewCollecting(t) updatedCell := &multigresv1alpha1.Cell{} - if err := c.Get(t.Context(), + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-cell-pending-deletion", Namespace: "default"}, - updatedCell); err != nil { - t.Fatalf("Failed to get Cell: %v", err) - } + updatedCell), "Failed to get Cell") found := false for _, cond := range updatedCell.Status.Conditions { if cond.Type == multigresv1alpha1.ConditionReadyForDeletion { found = true - if cond.Status != metav1.ConditionTrue { - t.Errorf("expected ReadyForDeletion=True, got %s", cond.Status) - } + ck.Eq( + metav1.ConditionTrue, + cond.Status, + "expected ReadyForDeletion=True, got", + ) } } - if !found { - t.Errorf("expected ConditionReadyForDeletion to be present") - } + ck.True(found, "expected ConditionReadyForDeletion to be present") }, }, ////---------------------------------------- @@ -575,6 +546,7 @@ func TestCellReconciler_Reconcile(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) // Calculate hashed name clusterName := tc.cell.Labels["multigres.com/cluster"] @@ -628,9 +600,7 @@ func TestCellReconciler_Reconcile(t *testing.T) { } if !cellInExisting { err := fakeClient.Create(t.Context(), tc.cell) - if err != nil { - t.Fatalf("Failed to create Cell: %v", err) - } + c.Require().NoError(err, "Failed to create Cell") } // Reconcile @@ -650,13 +620,12 @@ func TestCellReconciler_Reconcile(t *testing.T) { return } - if (result.RequeueAfter != 0) != tc.wantRequeue { - t.Errorf( - "Reconcile() requeue = %v, want requeue = %v", - result.RequeueAfter, - tc.wantRequeue, - ) - } + c.Eq( + tc.wantRequeue, + (result.RequeueAfter != 0), + "Reconcile() requeue = %v, want requeue =", + result.RequeueAfter, + ) // Run custom assertions if provided if tc.assertFunc != nil { @@ -667,6 +636,7 @@ func TestCellReconciler_Reconcile(t *testing.T) { } func TestCellReconciler_ReconcileNotFound(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -691,10 +661,6 @@ func TestCellReconciler_ReconcileNotFound(t *testing.T) { } result, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Errorf("Reconcile() should not error on NotFound, got: %v", err) - } - if result.RequeueAfter > 0 { - t.Errorf("Reconcile() should not requeue on NotFound") - } + c.NoError(err, "Reconcile() should not error on NotFound, got") + c.LessOrEqual(0, result.RequeueAfter, "Reconcile() should not requeue on NotFound") } diff --git a/pkg/resource-handler/controller/cell/integration_test.go b/pkg/resource-handler/controller/cell/integration_test.go index 164c8f9c..32865843 100644 --- a/pkg/resource-handler/controller/cell/integration_test.go +++ b/pkg/resource-handler/controller/cell/integration_test.go @@ -25,6 +25,8 @@ import ( "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestSetupWithManager(t *testing.T) { @@ -41,15 +43,13 @@ func TestSetupWithManager(t *testing.T) { ), ) - if err := (&cellcontroller.CellReconciler{ + assert.NewAborting(t).NoError((&cellcontroller.CellReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("cell-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") } func TestCellReconciliation(t *testing.T) { @@ -88,9 +88,11 @@ func TestCellReconciliation(t *testing.T) { Images: multigresv1alpha1.CellImages{ Multigateway: "ghcr.io/multigres/multigres:main", }, - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{ - Replicas: ptr.To(int32(2)), - }}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(2)), + }, + }, GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ Address: "global-topo:2379", RootPath: "/multigres/global", @@ -104,19 +106,39 @@ func TestCellReconciliation(t *testing.T) { wantResources: []client.Object{ &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "test-cell-multigateway", "multigateway", "zone1", "usw1-az1"), + Name: "test-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "test-cell-multigateway", + "multigateway", + "zone1", + "usw1-az1", + ), OwnerReferences: cellOwnerRefs(t, "test-cell"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(2)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(cellLabels(t, "test-cell-multigateway", "multigateway", "zone1", "usw1-az1")), + MatchLabels: metadata.GetSelectorLabels( + cellLabels( + t, + "test-cell-multigateway", + "multigateway", + "zone1", + "usw1-az1", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: cellLabels(t, "test-cell-multigateway", "multigateway", "zone1", "usw1-az1"), + Labels: cellLabels( + t, + "test-cell-multigateway", + "multigateway", + "zone1", + "usw1-az1", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -182,9 +204,15 @@ func TestCellReconciliation(t *testing.T) { }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "test-cell-multigateway", "multigateway", "zone1", "usw1-az1"), + Name: "test-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "test-cell-multigateway", + "multigateway", + "zone1", + "usw1-az1", + ), OwnerReferences: cellOwnerRefs(t, "test-cell"), }, Spec: corev1.ServiceSpec{ @@ -194,7 +222,15 @@ func TestCellReconciliation(t *testing.T) { tcpServicePort(t, "grpc", 15170), tcpServicePort(t, "postgres", 5432), }, - Selector: metadata.GetSelectorLabels(cellLabels(t, "test-cell-multigateway", "multigateway", "zone1", "usw1-az1")), + Selector: metadata.GetSelectorLabels( + cellLabels( + t, + "test-cell-multigateway", + "multigateway", + "zone1", + "usw1-az1", + ), + ), }, }, }, @@ -220,9 +256,11 @@ func TestCellReconciliation(t *testing.T) { Images: multigresv1alpha1.CellImages{ Multigateway: "ghcr.io/multigres/multigres:main", }, - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{ - Replicas: ptr.To(int32(3)), - }}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(3)), + }, + }, GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ Address: "global-topo:2379", RootPath: "/multigres/global", @@ -236,19 +274,39 @@ func TestCellReconciliation(t *testing.T) { wantResources: []client.Object{ &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "custom-replicas-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "custom-replicas-cell-multigateway", "multigateway", "zone2", "usw1-az2"), + Name: "custom-replicas-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "custom-replicas-cell-multigateway", + "multigateway", + "zone2", + "usw1-az2", + ), OwnerReferences: cellOwnerRefs(t, "custom-replicas-cell"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(3)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(cellLabels(t, "custom-replicas-cell-multigateway", "multigateway", "zone2", "usw1-az2")), + MatchLabels: metadata.GetSelectorLabels( + cellLabels( + t, + "custom-replicas-cell-multigateway", + "multigateway", + "zone2", + "usw1-az2", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: cellLabels(t, "custom-replicas-cell-multigateway", "multigateway", "zone2", "usw1-az2"), + Labels: cellLabels( + t, + "custom-replicas-cell-multigateway", + "multigateway", + "zone2", + "usw1-az2", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -314,9 +372,15 @@ func TestCellReconciliation(t *testing.T) { }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "custom-replicas-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "custom-replicas-cell-multigateway", "multigateway", "zone2", "usw1-az2"), + Name: "custom-replicas-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "custom-replicas-cell-multigateway", + "multigateway", + "zone2", + "usw1-az2", + ), OwnerReferences: cellOwnerRefs(t, "custom-replicas-cell"), }, Spec: corev1.ServiceSpec{ @@ -326,7 +390,15 @@ func TestCellReconciliation(t *testing.T) { tcpServicePort(t, "grpc", 15170), tcpServicePort(t, "postgres", 5432), }, - Selector: metadata.GetSelectorLabels(cellLabels(t, "custom-replicas-cell-multigateway", "multigateway", "zone2", "usw1-az2")), + Selector: metadata.GetSelectorLabels( + cellLabels( + t, + "custom-replicas-cell-multigateway", + "multigateway", + "zone2", + "usw1-az2", + ), + ), }, }, }, @@ -352,9 +424,11 @@ func TestCellReconciliation(t *testing.T) { Images: multigresv1alpha1.CellImages{ Multigateway: "custom/multigateway:v1.0.0", }, - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{ - Replicas: ptr.To(int32(2)), - }}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(2)), + }, + }, GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ Address: "global-topo:2379", RootPath: "/multigres/global", @@ -368,19 +442,39 @@ func TestCellReconciliation(t *testing.T) { wantResources: []client.Object{ &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "custom-images-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "custom-images-cell-multigateway", "multigateway", "zone3", "usw1-az3"), + Name: "custom-images-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "custom-images-cell-multigateway", + "multigateway", + "zone3", + "usw1-az3", + ), OwnerReferences: cellOwnerRefs(t, "custom-images-cell"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(2)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(cellLabels(t, "custom-images-cell-multigateway", "multigateway", "zone3", "usw1-az3")), + MatchLabels: metadata.GetSelectorLabels( + cellLabels( + t, + "custom-images-cell-multigateway", + "multigateway", + "zone3", + "usw1-az3", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: cellLabels(t, "custom-images-cell-multigateway", "multigateway", "zone3", "usw1-az3"), + Labels: cellLabels( + t, + "custom-images-cell-multigateway", + "multigateway", + "zone3", + "usw1-az3", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -446,9 +540,15 @@ func TestCellReconciliation(t *testing.T) { }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "custom-images-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "custom-images-cell-multigateway", "multigateway", "zone3", "usw1-az3"), + Name: "custom-images-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "custom-images-cell-multigateway", + "multigateway", + "zone3", + "usw1-az3", + ), OwnerReferences: cellOwnerRefs(t, "custom-images-cell"), }, Spec: corev1.ServiceSpec{ @@ -458,7 +558,15 @@ func TestCellReconciliation(t *testing.T) { tcpServicePort(t, "grpc", 15170), tcpServicePort(t, "postgres", 5432), }, - Selector: metadata.GetSelectorLabels(cellLabels(t, "custom-images-cell-multigateway", "multigateway", "zone3", "usw1-az3")), + Selector: metadata.GetSelectorLabels( + cellLabels( + t, + "custom-images-cell-multigateway", + "multigateway", + "zone3", + "usw1-az3", + ), + ), }, }, }, @@ -484,18 +592,20 @@ func TestCellReconciliation(t *testing.T) { Images: multigresv1alpha1.CellImages{ Multigateway: "ghcr.io/multigres/multigres:main", }, - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{ - Replicas: ptr.To(int32(2)), - Affinity: &corev1.Affinity{ - NodeAffinity: &corev1.NodeAffinity{ - RequiredDuringSchedulingIgnoredDuringExecution: &corev1.NodeSelector{ - NodeSelectorTerms: []corev1.NodeSelectorTerm{ - { - MatchExpressions: []corev1.NodeSelectorRequirement{ - { - Key: "node-type", - Operator: corev1.NodeSelectorOpIn, - Values: []string{"gateway"}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(2)), + Affinity: &corev1.Affinity{ + NodeAffinity: &corev1.NodeAffinity{ + RequiredDuringSchedulingIgnoredDuringExecution: &corev1.NodeSelector{ + NodeSelectorTerms: []corev1.NodeSelectorTerm{ + { + MatchExpressions: []corev1.NodeSelectorRequirement{ + { + Key: "node-type", + Operator: corev1.NodeSelectorOpIn, + Values: []string{"gateway"}, + }, }, }, }, @@ -503,7 +613,7 @@ func TestCellReconciliation(t *testing.T) { }, }, }, - }}, + }, GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ Address: "global-topo:2379", RootPath: "/multigres/global", @@ -517,19 +627,39 @@ func TestCellReconciliation(t *testing.T) { wantResources: []client.Object{ &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "affinity-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "affinity-cell-multigateway", "multigateway", "zone4", "usw1-az4"), + Name: "affinity-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "affinity-cell-multigateway", + "multigateway", + "zone4", + "usw1-az4", + ), OwnerReferences: cellOwnerRefs(t, "affinity-cell"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(2)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(cellLabels(t, "affinity-cell-multigateway", "multigateway", "zone4", "usw1-az4")), + MatchLabels: metadata.GetSelectorLabels( + cellLabels( + t, + "affinity-cell-multigateway", + "multigateway", + "zone4", + "usw1-az4", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: cellLabels(t, "affinity-cell-multigateway", "multigateway", "zone4", "usw1-az4"), + Labels: cellLabels( + t, + "affinity-cell-multigateway", + "multigateway", + "zone4", + "usw1-az4", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -612,9 +742,15 @@ func TestCellReconciliation(t *testing.T) { }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "affinity-cell-multigateway", - Namespace: "default", - Labels: cellLabels(t, "affinity-cell-multigateway", "multigateway", "zone4", "usw1-az4"), + Name: "affinity-cell-multigateway", + Namespace: "default", + Labels: cellLabels( + t, + "affinity-cell-multigateway", + "multigateway", + "zone4", + "usw1-az4", + ), OwnerReferences: cellOwnerRefs(t, "affinity-cell"), }, Spec: corev1.ServiceSpec{ @@ -624,7 +760,15 @@ func TestCellReconciliation(t *testing.T) { tcpServicePort(t, "grpc", 15170), tcpServicePort(t, "postgres", 5432), }, - Selector: metadata.GetSelectorLabels(cellLabels(t, "affinity-cell-multigateway", "multigateway", "zone4", "usw1-az4")), + Selector: metadata.GetSelectorLabels( + cellLabels( + t, + "affinity-cell-multigateway", + "multigateway", + "zone4", + "usw1-az4", + ), + ), }, }, }, @@ -634,6 +778,7 @@ func TestCellReconciliation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) ctx := t.Context() mgr := testutil.SetUpEnvtestManager(t, scheme, testutil.WithCRDPaths( @@ -659,16 +804,12 @@ func TestCellReconciliation(t *testing.T) { Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("cell-controller"), } - if err := cellReconciler.SetupWithManager(mgr, controller.Options{ - // Needed for the parallel test runs + // Needed for the parallel test runs + c.Require().NoError(cellReconciler.SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") - if err := client.Create(ctx, tc.cell); err != nil { - t.Fatalf("Failed to create the initial item, %v", err) - } + c.Require().NoError(client.Create(ctx, tc.cell), "Failed to create the initial item") markManagedLocalTopoServerHealthy(t, ctx, client, tc.cell) // Patch wantResources with hashed names @@ -715,9 +856,7 @@ func TestCellReconciliation(t *testing.T) { } } - if err := watcher.WaitForMatch(tc.wantResources...); err != nil { - t.Errorf("Resources mismatch:\n%v", err) - } + c.NoError(watcher.WaitForMatch(tc.wantResources...), "Resources mismatch:\n") }) } } @@ -731,6 +870,7 @@ func markManagedLocalTopoServerHealthy( cell *multigresv1alpha1.Cell, ) { t.Helper() + c := assert.NewAborting(t) if cell.Spec.TopoServer == nil || cell.Spec.TopoServer.Etcd == nil { return } @@ -740,7 +880,7 @@ func markManagedLocalTopoServerHealthy( Namespace: cell.Namespace, Name: cellcontroller.BuildLocalTopoServerName(cell), } - if err := wait.PollUntilContextTimeout(ctx, 100*time.Millisecond, 10*time.Second, true, + c.NoError(wait.PollUntilContextTimeout(ctx, 100*time.Millisecond, 10*time.Second, true, func(ctx context.Context) (bool, error) { if err := k8sClient.Get(ctx, key, toposerver); err != nil { if apierrors.IsNotFound(err) { @@ -749,15 +889,16 @@ func markManagedLocalTopoServerHealthy( return false, err } return true, nil - }); err != nil { - t.Fatalf("Timed out waiting for managed local TopoServer %s/%s: %v", key.Namespace, key.Name, err) - } + }), "Timed out waiting for managed local TopoServer %s/%s", key.Namespace, key.Name) toposerver.Status.Phase = multigresv1alpha1.PhaseHealthy toposerver.Status.ObservedGeneration = toposerver.Generation - if err := k8sClient.Status().Update(ctx, toposerver); err != nil { - t.Fatalf("Failed to mark managed local TopoServer %s/%s healthy: %v", key.Namespace, key.Name, err) - } + c.NoError( + k8sClient.Status().Update(ctx, toposerver), + "Failed to mark managed local TopoServer %s/%s healthy", + key.Namespace, + key.Name, + ) } // cellLabels returns standard labels for cell resources in tests @@ -795,5 +936,10 @@ func tcpPort(t testing.TB, name string, port int32) corev1.ContainerPort { // tcpServicePort creates a TCP service port with named target func tcpServicePort(t testing.TB, name string, port int32) corev1.ServicePort { t.Helper() - return corev1.ServicePort{Name: name, Port: port, TargetPort: intstr.FromString(name), Protocol: corev1.ProtocolTCP} + return corev1.ServicePort{ + Name: name, + Port: port, + TargetPort: intstr.FromString(name), + Protocol: corev1.ProtocolTCP, + } } diff --git a/pkg/resource-handler/controller/cell/local_toposerver_test.go b/pkg/resource-handler/controller/cell/local_toposerver_test.go index cc56a34c..9a98f1fc 100644 --- a/pkg/resource-handler/controller/cell/local_toposerver_test.go +++ b/pkg/resource-handler/controller/cell/local_toposerver_test.go @@ -19,15 +19,16 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestBuildLocalTopoServer(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("AddToScheme() error = %v", err) - } + c.Require().NoError(multigresv1alpha1.AddToScheme(scheme), "AddToScheme() error =") cell := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ @@ -51,24 +52,12 @@ func TestBuildLocalTopoServer(t *testing.T) { } got, err := BuildLocalTopoServer(cell, scheme) - if err != nil { - t.Fatalf("BuildLocalTopoServer() error = %v", err) - } - if got == nil { - t.Fatal("BuildLocalTopoServer() = nil, want TopoServer") - } - if got.Name != BuildLocalTopoServerName(cell) { - t.Errorf("name = %q, want %q", got.Name, BuildLocalTopoServerName(cell)) - } - if got.Namespace != "default" { - t.Errorf("namespace = %q, want default", got.Namespace) - } - if got.Labels[metadata.LabelMultigresCluster] != "cluster" { - t.Errorf("cluster label = %q, want cluster", got.Labels[metadata.LabelMultigresCluster]) - } - if got.Labels[metadata.LabelMultigresCell] != "zone-a" { - t.Errorf("cell label = %q, want zone-a", got.Labels[metadata.LabelMultigresCell]) - } + c.Require().NoError(err, "BuildLocalTopoServer() error =") + c.Require().NotNil(got, "BuildLocalTopoServer() = nil, want TopoServer") + c.Eq(BuildLocalTopoServerName(cell), got.Name, "name") + c.Eq("default", got.Namespace, "namespace") + c.Eq("cluster", got.Labels[metadata.LabelMultigresCluster], "cluster label") + c.Eq("zone-a", got.Labels[metadata.LabelMultigresCell], "cell label") if got.Spec.Etcd == nil || got.Spec.Etcd.RootPath != "/multigres/zone-a" { t.Fatalf("etcd spec = %#v, want root path /multigres/zone-a", got.Spec.Etcd) } @@ -83,6 +72,7 @@ func TestBuildLocalTopoServer(t *testing.T) { func TestBuildLocalTopoServerExternalReturnsNil(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) cell := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{Name: "cluster-zone-a", Namespace: "default"}, @@ -98,16 +88,13 @@ func TestBuildLocalTopoServerExternalReturnsNil(t *testing.T) { } got, err := BuildLocalTopoServer(cell, runtime.NewScheme()) - if err != nil { - t.Fatalf("BuildLocalTopoServer() error = %v", err) - } - if got != nil { - t.Fatalf("BuildLocalTopoServer() = %#v, want nil", got) - } + c.NoError(err, "BuildLocalTopoServer() error =") + c.Nil(got, "BuildLocalTopoServer()") } func TestBuildLocalTopoServerNameIsSafeForTopoServerChildren(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) cell := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ @@ -116,17 +103,25 @@ func TestBuildLocalTopoServerNameIsSafeForTopoServerChildren(t *testing.T) { } got := BuildLocalTopoServerName(cell) - if len(got) > 52 { - t.Fatalf("managed TopoServer name length = %d, want <= 52: %q", len(got), got) - } - if len(got+"-headless") > 63 { - t.Fatalf("managed TopoServer headless service name length = %d, want <= 63: %q", - len(got+"-headless"), got+"-headless") - } + c.LessOrEqual( + 52, + len(got), + "managed TopoServer name length = %d, want <= 52: %q", + len(got), + got, + ) + c.LessOrEqual( + 63, + len(got+"-headless"), + "managed TopoServer headless service name length = %d, want <= 63: %q", + len(got+"-headless"), + got+"-headless", + ) } func TestCellReconcilerWaitsForManagedLocalTopoServer(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -145,33 +140,26 @@ func TestCellReconcilerWaitsForManagedLocalTopoServer(t *testing.T) { result, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, }) - if err != nil { - t.Fatalf("Reconcile() error = %v", err) - } - if result.RequeueAfter != localTopoServerRecheckDelay { - t.Fatalf("RequeueAfter = %v, want %v", result.RequeueAfter, localTopoServerRecheckDelay) - } + c.NoError(err, "Reconcile() error =") + c.Eq(localTopoServerRecheckDelay, result.RequeueAfter, "RequeueAfter") toposerver := &multigresv1alpha1.TopoServer{} - if err := fakeClient.Get(t.Context(), client.ObjectKey{ + c.NoError(fakeClient.Get(t.Context(), client.ObjectKey{ Name: BuildLocalTopoServerName(cell), Namespace: cell.Namespace, - }, toposerver); err != nil { - t.Fatalf("managed TopoServer should exist: %v", err) - } + }, toposerver), "managed TopoServer should exist") deployment := &appsv1.Deployment{} err = fakeClient.Get(t.Context(), client.ObjectKey{ Name: "test-cluster-zone-a-multigateway", Namespace: cell.Namespace, }, deployment) - if !errors.IsNotFound(err) { - t.Fatalf("Multigateway Deployment get error = %v, want NotFound", err) - } + c.True(errors.IsNotFound(err), "Multigateway Deployment get error = %v, want NotFound", err) } func TestCellReconcilerSetsWaitingStatusForManagedLocalTopoServer(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -194,9 +182,7 @@ func TestCellReconcilerSetsWaitingStatusForManagedLocalTopoServer(t *testing.T) Reason: "MultigatewayReady", ObservedGeneration: cell.Generation, }) - if err := fakeClient.Status().Update(t.Context(), cell); err != nil { - t.Fatalf("failed to seed Cell status: %v", err) - } + c.NoError(fakeClient.Status().Update(t.Context(), cell), "failed to seed Cell status") reconciler := &CellReconciler{ Client: fakeClient, @@ -204,36 +190,32 @@ func TestCellReconcilerSetsWaitingStatusForManagedLocalTopoServer(t *testing.T) Recorder: record.NewFakeRecorder(10), } - if _, err := reconciler.Reconcile(t.Context(), ctrl.Request{ + _, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, - }); err != nil { - t.Fatalf("Reconcile() error = %v", err) - } + }) + assert.NewAborting(t).NoError(err, "Reconcile() error =") updatedCell := &multigresv1alpha1.Cell{} - if err := fakeClient.Get(t.Context(), client.ObjectKey{ + c.NoError(fakeClient.Get(t.Context(), client.ObjectKey{ Name: cell.Name, Namespace: cell.Namespace, - }, updatedCell); err != nil { - t.Fatalf("Cell get error = %v", err) - } - if updatedCell.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Fatalf("Phase = %s, want Progressing", updatedCell.Status.Phase) - } + }, updatedCell), "Cell get error =") + c.Eq(multigresv1alpha1.PhaseProgressing, updatedCell.Status.Phase, "Phase") for _, conditionType := range []string{"Available", "Ready"} { condition := meta.FindStatusCondition(updatedCell.Status.Conditions, conditionType) - if condition == nil { - t.Fatalf("%s condition not found", conditionType) - } + c.NotNil(condition, "%s condition not found", conditionType) if condition.Status != metav1.ConditionFalse || condition.Reason != "LocalTopoServerNotReady" { t.Fatalf("%s = %s/%s, want False/LocalTopoServerNotReady", conditionType, condition.Status, condition.Reason) } - if condition.ObservedGeneration != updatedCell.Generation { - t.Fatalf("%s observedGeneration = %d, want %d", - conditionType, condition.ObservedGeneration, updatedCell.Generation) - } + c.Eq( + updatedCell.Generation, + condition.ObservedGeneration, + "%s observedGeneration = %d, want", + conditionType, + condition.ObservedGeneration, + ) } } @@ -272,9 +254,8 @@ func TestCellReconcilerDeletesStaleManagedLocalTopoServer(t *testing.T) { Name: toposerver.Name, Namespace: toposerver.Namespace, }, got) - if !errors.IsNotFound(err) { - t.Fatalf("stale local TopoServer get error = %v, want NotFound", err) - } + assert.NewAborting(t). + True(errors.IsNotFound(err), "stale local TopoServer get error = %v, want NotFound", err) } func TestCellReconcilerIgnoresUnownedLocalTopoServerWhenNoManagedTopoDesired(t *testing.T) { @@ -302,23 +283,21 @@ func TestCellReconcilerIgnoresUnownedLocalTopoServerWhenNoManagedTopoDesired(t * Recorder: record.NewFakeRecorder(10), } - if _, err := reconciler.Reconcile(t.Context(), ctrl.Request{ + _, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, - }); err != nil { - t.Fatalf("Reconcile() error = %v", err) - } + }) + assert.NewAborting(t).NoError(err, "Reconcile() error =") got := &multigresv1alpha1.TopoServer{} - if err := fakeClient.Get(t.Context(), client.ObjectKey{ + assert.NewAborting(t).NoError(fakeClient.Get(t.Context(), client.ObjectKey{ Name: toposerver.Name, Namespace: toposerver.Namespace, - }, got); err != nil { - t.Fatalf("unowned local TopoServer should be left alone: %v", err) - } + }, got), "unowned local TopoServer should be left alone") } func TestCellReconcilerRefusesManagedLocalTopoServerNameConflict(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -339,16 +318,18 @@ func TestCellReconcilerRefusesManagedLocalTopoServerNameConflict(t *testing.T) { _, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, }) - if err == nil { - t.Fatal("Reconcile() error = nil, want name conflict") - } - if !strings.Contains(err.Error(), "not controlled by Cell") { - t.Fatalf("Reconcile() error = %v, want not controlled by Cell", err) - } + c.Error(err, "Reconcile() error = nil, want name conflict") + c.StrContains( + err.Error(), + "not controlled by Cell", + "Reconcile() error = %v, want not controlled by Cell", + err, + ) } func TestCellReconcilerPendingDeletionWaitsForManagedLocalTopoServer(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -371,27 +352,19 @@ func TestCellReconcilerPendingDeletionWaitsForManagedLocalTopoServer(t *testing. result, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, }) - if err != nil { - t.Fatalf("Reconcile() error = %v", err) - } - if result.RequeueAfter != localTopoServerRecheckDelay { - t.Fatalf("RequeueAfter = %v, want %v", result.RequeueAfter, localTopoServerRecheckDelay) - } + c.NoError(err, "Reconcile() error =") + c.Eq(localTopoServerRecheckDelay, result.RequeueAfter, "RequeueAfter") updatedCell := &multigresv1alpha1.Cell{} - if err := fakeClient.Get(t.Context(), client.ObjectKey{ + c.NoError(fakeClient.Get(t.Context(), client.ObjectKey{ Name: cell.Name, Namespace: cell.Namespace, - }, updatedCell); err != nil { - t.Fatalf("Cell get error = %v", err) - } + }, updatedCell), "Cell get error =") condition := meta.FindStatusCondition( updatedCell.Status.Conditions, multigresv1alpha1.ConditionReadyForDeletion, ) - if condition == nil { - t.Fatal("ReadyForDeletion condition not found") - } + c.NotNil(condition, "ReadyForDeletion condition not found") if condition.Status != metav1.ConditionFalse || condition.Reason != "LocalTopoServerDeleting" { t.Fatalf("ReadyForDeletion = %s/%s, want False/LocalTopoServerDeleting", condition.Status, condition.Reason) @@ -402,6 +375,7 @@ func TestCellReconcilerPendingDeletionDeletesObservedLocalTopoServerWhenSpecNoLo t *testing.T, ) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -430,25 +404,20 @@ func TestCellReconcilerPendingDeletionDeletesObservedLocalTopoServerWhenSpecNoLo result, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, }) - if err != nil { - t.Fatalf("Reconcile() error = %v", err) - } - if result.RequeueAfter != localTopoServerRecheckDelay { - t.Fatalf("RequeueAfter = %v, want %v", result.RequeueAfter, localTopoServerRecheckDelay) - } + c.NoError(err, "Reconcile() error =") + c.Eq(localTopoServerRecheckDelay, result.RequeueAfter, "RequeueAfter") got := &multigresv1alpha1.TopoServer{} err = fakeClient.Get(t.Context(), client.ObjectKey{ Name: toposerver.Name, Namespace: toposerver.Namespace, }, got) - if !errors.IsNotFound(err) { - t.Fatalf("local TopoServer get error = %v, want NotFound", err) - } + c.True(errors.IsNotFound(err), "local TopoServer get error = %v, want NotFound", err) } func TestCellReconcilerLocalTopoServerReadyRequiresObservedGeneration(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -468,16 +437,13 @@ func TestCellReconcilerLocalTopoServerReadyRequiresObservedGeneration(t *testing } ready, err := reconciler.localTopoServerReady(t.Context(), cell) - if err != nil { - t.Fatalf("localTopoServerReady() error = %v", err) - } - if ready { - t.Fatal("localTopoServerReady() = true, want false for stale observedGeneration") - } + c.NoError(err, "localTopoServerReady() error =") + c.False(ready, "localTopoServerReady() = true, want false for stale observedGeneration") } func TestCellReconcilerPendingDeletionReadyAfterManagedLocalTopoServerDeleted(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := cellTestScheme(t) cell := managedLocalTopoCell("test-cell") @@ -499,27 +465,19 @@ func TestCellReconcilerPendingDeletionReadyAfterManagedLocalTopoServerDeleted(t result, err := reconciler.Reconcile(t.Context(), ctrl.Request{ NamespacedName: types.NamespacedName{Name: cell.Name, Namespace: cell.Namespace}, }) - if err != nil { - t.Fatalf("Reconcile() error = %v", err) - } - if result.RequeueAfter != 0 { - t.Fatalf("RequeueAfter = %v, want 0", result.RequeueAfter) - } + c.NoError(err, "Reconcile() error =") + c.Eq(0, result.RequeueAfter, "RequeueAfter") updatedCell := &multigresv1alpha1.Cell{} - if err := fakeClient.Get(t.Context(), client.ObjectKey{ + c.NoError(fakeClient.Get(t.Context(), client.ObjectKey{ Name: cell.Name, Namespace: cell.Namespace, - }, updatedCell); err != nil { - t.Fatalf("Cell get error = %v", err) - } + }, updatedCell), "Cell get error =") condition := meta.FindStatusCondition( updatedCell.Status.Conditions, multigresv1alpha1.ConditionReadyForDeletion, ) - if condition == nil { - t.Fatal("ReadyForDeletion condition not found") - } + c.NotNil(condition, "ReadyForDeletion condition not found") if condition.Status != metav1.ConditionTrue || condition.Reason != "LocalTopoServerDeleted" { t.Fatalf("ReadyForDeletion = %s/%s, want True/LocalTopoServerDeleted", condition.Status, condition.Reason) @@ -528,16 +486,11 @@ func TestCellReconcilerPendingDeletionReadyAfterManagedLocalTopoServerDeleted(t func cellTestScheme(t testing.TB) *runtime.Scheme { t.Helper() + c := assert.NewAborting(t) scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("AddToScheme(multigres) error = %v", err) - } - if err := appsv1.AddToScheme(scheme); err != nil { - t.Fatalf("AddToScheme(apps) error = %v", err) - } - if err := corev1.AddToScheme(scheme); err != nil { - t.Fatalf("AddToScheme(core) error = %v", err) - } + c.NoError(multigresv1alpha1.AddToScheme(scheme), "AddToScheme(multigres) error =") + c.NoError(appsv1.AddToScheme(scheme), "AddToScheme(apps) error =") + c.NoError(corev1.AddToScheme(scheme), "AddToScheme(core) error =") return scheme } diff --git a/pkg/resource-handler/controller/cell/multigateway_test.go b/pkg/resource-handler/controller/cell/multigateway_test.go index 85a1336e..0d29c2a2 100644 --- a/pkg/resource-handler/controller/cell/multigateway_test.go +++ b/pkg/resource-handler/controller/cell/multigateway_test.go @@ -6,7 +6,6 @@ import ( "testing" "time" - "github.com/google/go-cmp/cmp" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/resource" @@ -18,6 +17,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestBuildMultigatewayDeployment(t *testing.T) { @@ -1475,9 +1476,7 @@ func TestBuildMultigatewayDeployment(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildMultigatewayDeployment() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildMultigatewayDeployment() mismatch") }) } } @@ -1503,9 +1502,8 @@ func TestBuildCellNodeSelector(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { cell := &multigresv1alpha1.Cell{Spec: tc.spec} - if diff := cmp.Diff(tc.want, buildCellNodeSelector(cell)); diff != "" { - t.Errorf("buildCellNodeSelector() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t). + EqDiff(tc.want, buildCellNodeSelector(cell), "buildCellNodeSelector() mismatch") }) } } @@ -1531,6 +1529,7 @@ func TestBuildMultigatewayDeployment_ProjectRefAnnotation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewAborting(t) cell := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ Name: "test-cell", @@ -1557,9 +1556,7 @@ func TestBuildMultigatewayDeployment_ProjectRefAnnotation(t *testing.T) { } deploy, err := BuildMultigatewayDeployment(cell, scheme) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") if got := deploy.Spec.Template.Annotations[metadata.AnnotationProjectRef]; got != tc.want { t.Fatalf("annotation %q = %q, want %q", metadata.AnnotationProjectRef, got, tc.want) @@ -1571,15 +1568,15 @@ func TestBuildMultigatewayDeployment_ProjectRefAnnotation(t *testing.T) { metadata.LabelAppManagedBy: metadata.ManagedByMultigres, } for key, want := range assertedLabels { - if got := deploy.Spec.Template.Labels[key]; got != want { - t.Fatalf("label %q = %q, want %q", key, got, want) - } + got := deploy.Spec.Template.Labels[key] + c.Eq(want, got, "label %q = %q, want", key, got) } }) } } func TestBuildMultigatewayDeployment_OmitsPrometheusScrapeAnnotations(t *testing.T) { + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1618,9 +1615,7 @@ func TestBuildMultigatewayDeployment_OmitsPrometheusScrapeAnnotations(t *testing } deploy, err := BuildMultigatewayDeployment(cell, scheme) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") if _, ok := deploy.Spec.Template.Annotations[metadata.AnnotationPrometheusScrape]; ok { t.Fatalf("annotation %q should be omitted", metadata.AnnotationPrometheusScrape) @@ -1631,9 +1626,7 @@ func TestBuildMultigatewayDeployment_OmitsPrometheusScrapeAnnotations(t *testing if _, ok := deploy.Spec.Template.Annotations[metadata.AnnotationPrometheusPath]; ok { t.Fatalf("annotation %q should be omitted", metadata.AnnotationPrometheusPath) } - if got := deploy.Spec.Template.Annotations["custom-annotation"]; got != "keep-me" { - t.Fatalf("custom annotation = %q, want %q", got, "keep-me") - } + c.Eq("keep-me", deploy.Spec.Template.Annotations["custom-annotation"], "custom annotation") } func TestBuildMultigatewayService(t *testing.T) { @@ -1959,14 +1952,13 @@ func TestBuildMultigatewayService(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildMultigatewayService() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildMultigatewayService() mismatch") }) } } func TestBuildMultigatewayDeployment_Observability(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1991,12 +1983,8 @@ func TestBuildMultigatewayDeployment_Observability(t *testing.T) { }, } deploy, err := BuildMultigatewayDeployment(cellObj, scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment failed: %v", err) - } - if len(deploy.Spec.Template.Spec.Volumes) == 0 { - t.Errorf("expected volumes to contain otel config") - } + c.Require().NoError(err, "BuildMultigatewayDeployment failed") + c.NotEmpty(deploy.Spec.Template.Spec.Volumes, "expected volumes to contain otel config") env := deploy.Spec.Template.Spec.Containers[0].Env assertMultigatewayEnvVar( t, @@ -2014,9 +2002,8 @@ func assertMultigatewayEnvVar(t *testing.T, envVars []corev1.EnvVar, name, want t.Helper() for _, envVar := range envVars { if envVar.Name == name { - if envVar.Value != want { - t.Fatalf("env var %q = %q, want %q", name, envVar.Value, want) - } + assert.NewAborting(t). + Eq(want, envVar.Value, "env var %q = %q, want", name, envVar.Value) return } } @@ -2027,9 +2014,7 @@ func assertMultigatewayResourceAttribute(t *testing.T, envVars []corev1.EnvVar, t.Helper() for _, envVar := range envVars { if envVar.Name == "OTEL_RESOURCE_ATTRIBUTES" { - if !strings.Contains(envVar.Value, want) { - t.Fatalf("OTEL_RESOURCE_ATTRIBUTES = %q, want it to contain %q", envVar.Value, want) - } + assert.NewAborting(t).StrContains(envVar.Value, want, "OTEL_RESOURCE_ATTRIBUTES") return } } @@ -2041,6 +2026,7 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { _ = multigresv1alpha1.AddToScheme(scheme) t.Run("internalTLS and certCommonName independently enable both TLS modes", func(t *testing.T) { + c := assert.NewCollecting(t) cellObj := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ Name: "test-tls", @@ -2063,9 +2049,7 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { }, } deploy, err := BuildMultigatewayDeployment(cellObj, scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment failed: %v", err) - } + c.Require().NoError(err, "BuildMultigatewayDeployment failed") // Verify TLS volume exists var foundVol bool @@ -2076,12 +2060,11 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { if v.Secret == nil { t.Error("TLS volume should use a Secret source") } else { - if v.Secret.SecretName != multigresv1alpha1.CertSecretName { - t.Errorf( - "TLS volume secretName = %q, want %q", - v.Secret.SecretName, multigresv1alpha1.CertSecretName, - ) - } + c.Eq( + multigresv1alpha1.CertSecretName, + v.Secret.SecretName, + "TLS volume secretName", + ) if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf( "TLS volume defaultMode = %v, want 0444", @@ -2096,19 +2079,12 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { t.Error("internal TLS volume should use a Secret source") } else { wantSecretName := "multigateway.test-cluster.default.multigres.internal" - if v.Secret.SecretName != wantSecretName { - t.Errorf( - "internal TLS volume secretName = %q, want %q", - v.Secret.SecretName, wantSecretName, - ) - } - if strings.Contains(v.Secret.SecretName, cellObj.Spec.CertCommonName) { - t.Errorf( - "internal TLS volume secretName %q must not contain public CertCommonName %q", - v.Secret.SecretName, - cellObj.Spec.CertCommonName, - ) - } + c.Eq(wantSecretName, v.Secret.SecretName, "internal TLS volume secretName") + c.NotStrContains( + v.Secret.SecretName, + cellObj.Spec.CertCommonName, + "internal TLS volume secretName", + ) if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf( "internal TLS volume defaultMode = %v, want 0444", @@ -2118,12 +2094,12 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { } } } - if !foundVol { - t.Errorf("expected TLS volume %q in pod spec", tlsVolumeName) - } - if !foundInternalVol { - t.Errorf("expected internal TLS volume %q in pod spec", internalTLSVolumeName) - } + c.True(foundVol, "expected TLS volume %q in pod spec", tlsVolumeName) + c.True( + foundInternalVol, + "expected internal TLS volume %q in pod spec", + internalTLSVolumeName, + ) // Verify TLS volumeMount exists container := deploy.Spec.Template.Spec.Containers[0] @@ -2132,32 +2108,21 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { for _, m := range container.VolumeMounts { if m.Name == tlsVolumeName { foundMount = true - if m.MountPath != tlsMountPath { - t.Errorf("TLS mount path = %q, want %q", m.MountPath, tlsMountPath) - } - if !m.ReadOnly { - t.Error("TLS mount should be readOnly") - } + c.Eq(tlsMountPath, m.MountPath, "TLS mount path") + c.True(m.ReadOnly, "TLS mount should be readOnly") } if m.Name == internalTLSVolumeName { foundInternalMount = true - if m.MountPath != internalTLSMountPath { - t.Errorf( - "internal TLS mount path = %q, want %q", - m.MountPath, internalTLSMountPath, - ) - } - if !m.ReadOnly { - t.Error("internal TLS mount should be readOnly") - } + c.Eq(internalTLSMountPath, m.MountPath, "internal TLS mount path") + c.True(m.ReadOnly, "internal TLS mount should be readOnly") } } - if !foundMount { - t.Errorf("expected TLS volumeMount %q in container", tlsVolumeName) - } - if !foundInternalMount { - t.Errorf("expected internal TLS volumeMount %q in container", internalTLSVolumeName) - } + c.True(foundMount, "expected TLS volumeMount %q in container", tlsVolumeName) + c.True( + foundInternalMount, + "expected internal TLS volumeMount %q in container", + internalTLSVolumeName, + ) // Verify TLS args are appended args := container.Args @@ -2176,24 +2141,15 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { "--pg-tls-key-file", tlsKeyFile, } // The TLS args should be appended as a group. - if len(args) < len(wantArgs) { - t.Fatalf("expected at least %d args, got %d", len(wantArgs), len(args)) - } + c.Require().GreaterOrEqual(len(wantArgs), len(args), "expected at least") tailArgs := args[len(args)-len(wantArgs):] - if diff := cmp.Diff(wantArgs, tailArgs); diff != "" { - t.Errorf("TLS args mismatch (-want +got):\n%s", diff) - } + c.EqDiff(wantArgs, tailArgs, "TLS args mismatch") serverName := tailArgs[len(wantArgs)-6] - if strings.Contains(serverName, cellObj.Spec.CertCommonName) { - t.Errorf( - "multipooler gRPC server name %q must not contain public CertCommonName %q", - serverName, - cellObj.Spec.CertCommonName, - ) - } + c.NotStrContains(serverName, cellObj.Spec.CertCommonName, "multipooler gRPC server name") }) t.Run("internalTLS enables internal TLS without certCommonName", func(t *testing.T) { + c := assert.NewCollecting(t) cellObj := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ Name: "test-no-public-tls", @@ -2215,19 +2171,12 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { }, } deploy, err := BuildMultigatewayDeployment(cellObj, scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment failed: %v", err) - } + c.Require().NoError(err, "BuildMultigatewayDeployment failed") var foundInternalVol bool for _, v := range deploy.Spec.Template.Spec.Volumes { - if v.Name == tlsVolumeName || - (v.Secret != nil && v.Secret.SecretName == multigresv1alpha1.CertSecretName) { - t.Errorf( - "public generated-certs TLS volume %q should not be present when CertCommonName is empty", - v.Name, - ) - } + c.False(v.Name == tlsVolumeName || + (v.Secret != nil && v.Secret.SecretName == multigresv1alpha1.CertSecretName), "public generated-certs TLS volume %q should not be present when CertCommonName is empty", v.Name) if v.Name == internalTLSVolumeName { foundInternalVol = true if v.Secret == nil { @@ -2235,12 +2184,7 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { continue } wantSecretName := "multigateway.test-cluster.default.multigres.internal" - if v.Secret.SecretName != wantSecretName { - t.Errorf( - "internal TLS volume secretName = %q, want %q", - v.Secret.SecretName, wantSecretName, - ) - } + c.Eq(wantSecretName, v.Secret.SecretName, "internal TLS volume secretName") if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf( "internal TLS volume defaultMode = %v, want 0444", @@ -2249,43 +2193,34 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { } } } - if !foundInternalVol { - t.Errorf("expected internal TLS volume %q in pod spec", internalTLSVolumeName) - } + c.True( + foundInternalVol, + "expected internal TLS volume %q in pod spec", + internalTLSVolumeName, + ) container := deploy.Spec.Template.Spec.Containers[0] var foundInternalMount bool for _, m := range container.VolumeMounts { - if m.Name == tlsVolumeName { - t.Errorf( - "public TLS volumeMount %q should not be present when CertCommonName is empty", - tlsVolumeName, - ) - } + c.NotEq(tlsVolumeName, m.Name, "public TLS volumeMount") if m.Name == internalTLSVolumeName { foundInternalMount = true - if m.MountPath != internalTLSMountPath { - t.Errorf( - "internal TLS mount path = %q, want %q", - m.MountPath, internalTLSMountPath, - ) - } - if !m.ReadOnly { - t.Error("internal TLS mount should be readOnly") - } + c.Eq(internalTLSMountPath, m.MountPath, "internal TLS mount path") + c.True(m.ReadOnly, "internal TLS mount should be readOnly") } } - if !foundInternalMount { - t.Errorf("expected internal TLS volumeMount %q in container", internalTLSVolumeName) - } + c.True( + foundInternalMount, + "expected internal TLS volumeMount %q in container", + internalTLSVolumeName, + ) for _, arg := range container.Args { - if arg == "--pg-tls-cert-file" || arg == "--pg-tls-key-file" { - t.Errorf( - "public PostgreSQL TLS arg %q should not be present when CertCommonName is empty", - arg, - ) - } + c.False( + arg == "--pg-tls-cert-file" || arg == "--pg-tls-key-file", + "public PostgreSQL TLS arg %q should not be present when CertCommonName is empty", + arg, + ) } wantInternalArgs := []string{ @@ -2300,20 +2235,13 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { "multipooler.test-cluster.default.multigres.internal", "--multipooler-grpc-require-tls", } - if len(container.Args) < len(wantInternalArgs) { - t.Fatalf( - "expected at least %d args, got %d", - len(wantInternalArgs), - len(container.Args), - ) - } + c.Require().GreaterOrEqual(len(wantInternalArgs), len(container.Args), "expected at least") tailArgs := container.Args[len(container.Args)-len(wantInternalArgs):] - if diff := cmp.Diff(wantInternalArgs, tailArgs); diff != "" { - t.Errorf("internal TLS args mismatch (-want +got):\n%s", diff) - } + c.EqDiff(wantInternalArgs, tailArgs, "internal TLS args mismatch") }) t.Run("certCommonName enables public TLS without internalTLS", func(t *testing.T) { + c := assert.NewCollecting(t) cellObj := &multigresv1alpha1.Cell{ ObjectMeta: metav1.ObjectMeta{ Name: "test-public-tls-only", @@ -2333,9 +2261,7 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { }, } deploy, err := BuildMultigatewayDeployment(cellObj, scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment failed: %v", err) - } + c.Require().NoError(err, "BuildMultigatewayDeployment failed") internalSecretName := "multigateway.test-cluster.default.multigres.internal" var foundPublicVolume bool @@ -2345,13 +2271,11 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { if volume.Secret == nil { t.Error("public TLS volume should use a Secret source") } else { - if volume.Secret.SecretName != multigresv1alpha1.CertSecretName { - t.Errorf( - "public TLS volume secretName = %q, want %q", - volume.Secret.SecretName, - multigresv1alpha1.CertSecretName, - ) - } + c.Eq( + multigresv1alpha1.CertSecretName, + volume.Secret.SecretName, + "public TLS volume secretName", + ) if volume.Secret.DefaultMode == nil || *volume.Secret.DefaultMode != 0o444 { t.Errorf( "public TLS volume defaultMode = %v, want 0444", @@ -2360,52 +2284,34 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { } } } - if volume.Name == internalTLSVolumeName || - (volume.Secret != nil && volume.Secret.SecretName == internalSecretName) { - t.Errorf( - "internal TLS secret volume %q should not be present when InternalTLS is nil", - volume.Name, - ) - } - } - if !foundPublicVolume { - t.Errorf("expected public generated-certs TLS volume %q in pod spec", tlsVolumeName) + c.False(volume.Name == internalTLSVolumeName || + (volume.Secret != nil && volume.Secret.SecretName == internalSecretName), "internal TLS secret volume %q should not be present when InternalTLS is nil", volume.Name) } + c.True( + foundPublicVolume, + "expected public generated-certs TLS volume %q in pod spec", + tlsVolumeName, + ) container := deploy.Spec.Template.Spec.Containers[0] var foundPublicMount bool for _, mount := range container.VolumeMounts { if mount.Name == tlsVolumeName { foundPublicMount = true - if mount.MountPath != tlsMountPath { - t.Errorf("public TLS mount path = %q, want %q", mount.MountPath, tlsMountPath) - } - if !mount.ReadOnly { - t.Error("public TLS mount should be readOnly") - } - } - if mount.Name == internalTLSVolumeName { - t.Errorf( - "internal TLS volumeMount %q should not be present when InternalTLS is nil", - internalTLSVolumeName, - ) + c.Eq(tlsMountPath, mount.MountPath, "public TLS mount path") + c.True(mount.ReadOnly, "public TLS mount should be readOnly") } + c.NotEq(internalTLSVolumeName, mount.Name, "internal TLS volumeMount") } - if !foundPublicMount { - t.Errorf("expected public TLS volumeMount %q in container", tlsVolumeName) - } + c.True(foundPublicMount, "expected public TLS volumeMount %q in container", tlsVolumeName) wantPublicArgs := []string{ "--pg-tls-cert-file", tlsCertFile, "--pg-tls-key-file", tlsKeyFile, } - if len(container.Args) < len(wantPublicArgs) { - t.Fatalf("expected at least %d args, got %d", len(wantPublicArgs), len(container.Args)) - } + c.Require().GreaterOrEqual(len(wantPublicArgs), len(container.Args), "expected at least") tailArgs := container.Args[len(container.Args)-len(wantPublicArgs):] - if diff := cmp.Diff(wantPublicArgs, tailArgs); diff != "" { - t.Errorf("public PostgreSQL TLS args mismatch (-want +got):\n%s", diff) - } + c.EqDiff(wantPublicArgs, tailArgs, "public PostgreSQL TLS args mismatch") internalTLSArgNames := map[string]struct{}{ "--grpc-cert": {}, @@ -2419,9 +2325,8 @@ func TestBuildMultigatewayDeployment_TLS(t *testing.T) { "--multipooler-grpc-require-tls": {}, } for _, arg := range container.Args { - if _, found := internalTLSArgNames[arg]; found { - t.Errorf("internal TLS arg %q should not be present when InternalTLS is nil", arg) - } + _, found := internalTLSArgNames[arg] + c.False(found, "internal TLS arg %q should not be present when InternalTLS is nil", arg) } }) } @@ -2449,29 +2354,29 @@ func TestBuildMultigatewayDeployment_Buffer(t *testing.T) { gatewayArgs := func(t *testing.T, cellObj *multigresv1alpha1.Cell) []string { t.Helper() deploy, err := BuildMultigatewayDeployment(cellObj, scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment failed: %v", err) - } + assert.NewAborting(t).NoError(err, "BuildMultigatewayDeployment failed") return deploy.Spec.Template.Spec.Containers[0].Args } t.Run("default resolved cell emits --buffer-enabled only", func(t *testing.T) { + c := assert.NewCollecting(t) // A Cell resolved from a default MultigresCluster carries // buffer.enabled: true and nothing else. args := gatewayArgs(t, buildCell(&multigresv1alpha1.GatewayBufferConfig{ Enabled: ptr.To(true), })) - if !slices.Contains(args, "--buffer-enabled=true") { - t.Errorf("expected --buffer-enabled=true in args, got %v", args) - } + c.Contains(args, "--buffer-enabled=true", "expected --buffer-enabled=true in args, got") for _, arg := range args { - if strings.HasPrefix(arg, "--buffer-") && arg != "--buffer-enabled=true" { - t.Errorf("unexpected buffer flag %q for default config", arg) - } + c.False( + strings.HasPrefix(arg, "--buffer-") && arg != "--buffer-enabled=true", + "unexpected buffer flag %q for default config", + arg, + ) } }) t.Run("explicit values land verbatim", func(t *testing.T) { + c := assert.NewCollecting(t) args := gatewayArgs(t, buildCell(&multigresv1alpha1.GatewayBufferConfig{ Enabled: ptr.To(true), Window: &metav1.Duration{Duration: 20 * time.Second}, @@ -2489,20 +2394,16 @@ func TestBuildMultigatewayDeployment_Buffer(t *testing.T) { "--buffer-drain-concurrency", "4", } idx := slices.Index(args, "--buffer-enabled=true") - if idx < 0 || len(args) < idx+len(want) { - t.Fatalf("buffer args missing or truncated, got %v", args) - } - if diff := cmp.Diff(want, args[idx:idx+len(want)]); diff != "" { - t.Errorf("buffer args mismatch (-want +got):\n%s", diff) - } + c.Require(). + False(idx < 0 || len(args) < idx+len(want), "buffer args missing or truncated, got %v", args) + c.EqDiff(want, args[idx:idx+len(want)], "buffer args mismatch") }) t.Run("nil buffer emits no buffer flags", func(t *testing.T) { // Legacy Cell CRs written before this field existed. for _, arg := range gatewayArgs(t, buildCell(nil)) { - if strings.HasPrefix(arg, "--buffer-") { - t.Errorf("unexpected buffer flag %q for nil buffer config", arg) - } + assert.NewCollecting(t). + False(strings.HasPrefix(arg, "--buffer-"), "unexpected buffer flag %q for nil buffer config", arg) } }) @@ -2511,9 +2412,8 @@ func TestBuildMultigatewayDeployment_Buffer(t *testing.T) { Enabled: ptr.To(false), Window: &metav1.Duration{Duration: 20 * time.Second}, })) - if !slices.Contains(args, "--buffer-enabled=false") { - t.Errorf("expected --buffer-enabled=false when disabled, got %v", args) - } + assert.NewCollecting(t). + Contains(args, "--buffer-enabled=false", "expected --buffer-enabled=false when disabled, got") }) } @@ -2545,43 +2445,41 @@ func TestBuildMultigatewayDeployment_TopoClientTLS(t *testing.T) { } t.Run("presents the client certificate when the reference carries it", func(t *testing.T) { + ck := assert.NewCollecting(t) cell := baseCell() secret := multigresv1alpha1.TopoClientCertSecretName("test-cluster") cell.Spec.GlobalTopoServer.CASecret = secret cell.Spec.GlobalTopoServer.ClientCertSecret = secret got, err := BuildMultigatewayDeployment(cell, scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment() error = %v", err) - } + ck.Require().NoError(err, "BuildMultigatewayDeployment() error =") c := got.Spec.Template.Spec.Containers[0] - if !slices.Contains(c.Args, "--topo-etcd-tls-cert") { - t.Errorf("missing topo client TLS flags: %v", c.Args) - } + ck.Contains(c.Args, "--topo-etcd-tls-cert", "missing topo client TLS flags") var mounted bool for _, m := range c.VolumeMounts { if m.Name == multigresv1alpha1.TopoClientTLSVolumeName { mounted = true } } - if !mounted { - t.Error("multigateway does not mount the topo client certificate") - } + ck.True(mounted, "multigateway does not mount the topo client certificate") }) t.Run("renders unchanged when the reference carries no credential", func(t *testing.T) { + ck := assert.NewCollecting(t) got, err := BuildMultigatewayDeployment(baseCell(), scheme) - if err != nil { - t.Fatalf("BuildMultigatewayDeployment() error = %v", err) - } + ck.Require().NoError(err, "BuildMultigatewayDeployment() error =") c := got.Spec.Template.Spec.Containers[0] - if slices.Contains(c.Args, "--topo-etcd-tls-cert") { - t.Error("topo TLS flag present with no credential on the reference") - } + ck.NotContains( + c.Args, + "--topo-etcd-tls-cert", + "topo TLS flag present with no credential on the reference", + ) for _, v := range got.Spec.Template.Spec.Volumes { - if v.Name == multigresv1alpha1.TopoClientTLSVolumeName { - t.Error("topo client volume present with no credential on the reference") - } + ck.NotEq( + multigresv1alpha1.TopoClientTLSVolumeName, + v.Name, + "topo client volume present with no credential on the reference", + ) } }) } diff --git a/pkg/resource-handler/controller/shard/configmap_test.go b/pkg/resource-handler/controller/shard/configmap_test.go index 1c28eeca..e306ae55 100644 --- a/pkg/resource-handler/controller/shard/configmap_test.go +++ b/pkg/resource-handler/controller/shard/configmap_test.go @@ -4,12 +4,13 @@ import ( "strings" "testing" - "github.com/google/go-cmp/cmp" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildPgHbaConfigMap(t *testing.T) { @@ -58,38 +59,33 @@ func TestBuildPgHbaConfigMap(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) scheme := tc.scheme if scheme == nil { scheme = defaultScheme } cm, err := BuildPgHbaConfigMap(tc.shard, scheme) - if (err != nil) != tc.wantErr { - t.Fatalf("BuildPgHbaConfigMap() error = %v, wantErr %v", err, tc.wantErr) - } + c.Require().ErrorWhen(tc.wantErr, err, "BuildPgHbaConfigMap() error") if tc.wantErr { return } - if cm.Name != PgHbaConfigMapName(tc.shard.Name) { - t.Errorf("ConfigMap name = %v, want %v", cm.Name, PgHbaConfigMapName(tc.shard.Name)) - } - if cm.Namespace != tc.shard.Namespace { - t.Errorf("ConfigMap namespace = %v, want %v", cm.Namespace, tc.shard.Namespace) - } + c.Eq(PgHbaConfigMapName(tc.shard.Name), cm.Name, "ConfigMap name") + c.Eq(tc.shard.Namespace, cm.Namespace, "ConfigMap namespace") // Verify owner reference - if len(cm.OwnerReferences) != 1 { - t.Fatalf("Expected 1 owner reference, got %d", len(cm.OwnerReferences)) - } + c.Require(). + Len(cm.OwnerReferences, 1, "Expected 1 owner reference, got %d", len(cm.OwnerReferences)) ownerRef := cm.OwnerReferences[0] if ownerRef.Name != tc.shard.Name || ownerRef.Kind != "Shard" { t.Errorf("Owner reference = %+v, want Shard/%s", ownerRef, tc.shard.Name) } - if !ptr.Deref(ownerRef.Controller, false) { - t.Error("Expected owner reference to be controller") - } + c.True( + ptr.Deref(ownerRef.Controller, false), + "Expected owner reference to be controller", + ) // Verify labels expectedLabels := map[string]string{ @@ -100,20 +96,18 @@ func TestBuildPgHbaConfigMap(t *testing.T) { "app.kubernetes.io/managed-by": "multigres-operator", "multigres.com/cluster": tc.shard.Labels["multigres.com/cluster"], } - if diff := cmp.Diff(expectedLabels, cm.Labels); diff != "" { - t.Errorf("Labels mismatch (-want +got):\n%s", diff) - } + c.EqDiff(expectedLabels, cm.Labels, "Labels mismatch") // Verify template content exists template, ok := cm.Data["pg_hba_template.conf"] - if !ok { - t.Fatal("ConfigMap missing pg_hba_template.conf key") - } + c.Require().True(ok, "ConfigMap missing pg_hba_template.conf key") // Verify the template matches what's embedded (source of truth) - if template != DefaultPgHbaTemplate { - t.Error("Template content doesn't match DefaultPgHbaTemplate") - } + c.Eq( + DefaultPgHbaTemplate, + template, + "Template content doesn't match DefaultPgHbaTemplate", + ) }) } } @@ -132,34 +126,28 @@ func TestBuildPostgresConfigMap(t *testing.T) { } t.Run("stores rendered content under the config key with an owner ref", func(t *testing.T) { + c := assert.NewCollecting(t) rendered := "# rendered\nmax_connections = 200\n" cm, err := BuildPostgresConfigMap(shard, rendered, scheme) - if err != nil { - t.Fatalf("BuildPostgresConfigMap() error = %v", err) - } - if cm.Name != PostgresConfigMapName(shard.Name) { - t.Errorf("name = %q, want %q", cm.Name, PostgresConfigMapName(shard.Name)) - } - if cm.Namespace != shard.Namespace { - t.Errorf("namespace = %q, want %q", cm.Namespace, shard.Namespace) - } - if got := cm.Data[PostgresConfigMapKey]; got != rendered { - t.Errorf("Data[%q] = %q, want %q", PostgresConfigMapKey, got, rendered) - } + c.Require().NoError(err, "BuildPostgresConfigMap() error =") + c.Eq(PostgresConfigMapName(shard.Name), cm.Name, "name") + c.Eq(shard.Namespace, cm.Namespace, "namespace") + got := cm.Data[PostgresConfigMapKey] + c.Eq(rendered, got, "Data[%q] = %q, want", PostgresConfigMapKey, got) if len(cm.OwnerReferences) != 1 || cm.OwnerReferences[0].Name != shard.Name || cm.OwnerReferences[0].Kind != "Shard" { t.Errorf("owner reference = %+v, want Shard/%s", cm.OwnerReferences, shard.Name) } - if !ptr.Deref(cm.OwnerReferences[0].Controller, false) { - t.Error("expected owner reference to be controller") - } + c.True( + ptr.Deref(cm.OwnerReferences[0].Controller, false), + "expected owner reference to be controller", + ) }) t.Run("returns error on invalid scheme", func(t *testing.T) { - if _, err := BuildPostgresConfigMap(shard, "x", runtime.NewScheme()); err == nil { - t.Error("expected error with empty scheme") - } + _, err := BuildPostgresConfigMap(shard, "x", runtime.NewScheme()) + assert.NewCollecting(t).Error(err, "expected error with empty scheme") }) } @@ -177,47 +165,38 @@ func TestBuildPostgresExporterQueriesConfigMap(t *testing.T) { } t.Run("stores embedded queries under the queries key with an owner ref", func(t *testing.T) { + c := assert.NewCollecting(t) cm, err := BuildPostgresExporterQueriesConfigMap(shard, scheme) - if err != nil { - t.Fatalf("BuildPostgresExporterQueriesConfigMap() error = %v", err) - } - if cm.Name != PostgresExporterQueriesConfigMapName(shard.Name) { - t.Errorf( - "name = %q, want %q", - cm.Name, - PostgresExporterQueriesConfigMapName(shard.Name), - ) - } - if cm.Namespace != shard.Namespace { - t.Errorf("namespace = %q, want %q", cm.Namespace, shard.Namespace) - } - if got := cm.Data[PostgresExporterQueriesConfigMapKey]; got != DefaultPostgresExporterQueries { - t.Errorf( - "Data[%q] doesn't match DefaultPostgresExporterQueries", - PostgresExporterQueriesConfigMapKey, - ) - } + c.Require().NoError(err, "BuildPostgresExporterQueriesConfigMap() error =") + c.Eq(PostgresExporterQueriesConfigMapName(shard.Name), cm.Name, "name") + c.Eq(shard.Namespace, cm.Namespace, "namespace") + c.Eq( + DefaultPostgresExporterQueries, + cm.Data[PostgresExporterQueriesConfigMapKey], + "Data[%q] doesn't match DefaultPostgresExporterQueries", + PostgresExporterQueriesConfigMapKey, + ) if len(cm.OwnerReferences) != 1 || cm.OwnerReferences[0].Name != shard.Name || cm.OwnerReferences[0].Kind != "Shard" { t.Errorf("owner reference = %+v, want Shard/%s", cm.OwnerReferences, shard.Name) } - if !ptr.Deref(cm.OwnerReferences[0].Controller, false) { - t.Error("expected owner reference to be controller") - } + c.True( + ptr.Deref(cm.OwnerReferences[0].Controller, false), + "expected owner reference to be controller", + ) }) t.Run("returns error on invalid scheme", func(t *testing.T) { - if _, err := BuildPostgresExporterQueriesConfigMap(shard, runtime.NewScheme()); err == nil { - t.Error("expected error with empty scheme") - } + _, err := BuildPostgresExporterQueriesConfigMap(shard, runtime.NewScheme()) + assert.NewCollecting(t).Error(err, "expected error with empty scheme") }) } func TestDefaultPostgresExporterQueriesEmbedded(t *testing.T) { - if DefaultPostgresExporterQueries == "" { - t.Fatal("DefaultPostgresExporterQueries is empty - go:embed may have failed") - } + c := assert.NewCollecting(t) + c.Require(). + NotEq("", DefaultPostgresExporterQueries, "DefaultPostgresExporterQueries is empty - go:embed may have failed") // Each top-level key becomes the exporter's metric-name prefix. for _, queryName := range []string{ @@ -228,17 +207,19 @@ func TestDefaultPostgresExporterQueriesEmbedded(t *testing.T) { "connection_stats:", "max_connections:", } { - if !strings.Contains(DefaultPostgresExporterQueries, "\n"+queryName) { - t.Errorf("DefaultPostgresExporterQueries missing query block %q", queryName) - } + c.StrContains( + DefaultPostgresExporterQueries, + "\n"+queryName, + "DefaultPostgresExporterQueries missing query block %q", + queryName, + ) } } func TestDefaultPgHbaTemplateEmbedded(t *testing.T) { + c := assert.NewCollecting(t) // Verify the embedded template is not empty - if DefaultPgHbaTemplate == "" { - t.Error("DefaultPgHbaTemplate is empty - go:embed may have failed") - } + c.NotEq("", DefaultPgHbaTemplate, "DefaultPgHbaTemplate is empty - go:embed may have failed") // Verify critical configuration lines exist // We check for the presence of rules, ignoring multiple spaces @@ -276,12 +257,11 @@ func TestDefaultPgHbaTemplateEmbedded(t *testing.T) { break } } - if !found { - t.Errorf( - "DefaultPgHbaTemplate missing %s (expected line containing all of: %v)", - check.desc, - check.mustContain, - ) - } + c.True( + found, + "DefaultPgHbaTemplate missing %s (expected line containing all of: %v)", + check.desc, + check.mustContain, + ) } } diff --git a/pkg/resource-handler/controller/shard/containers_test.go b/pkg/resource-handler/controller/shard/containers_test.go index 10f006cd..2e9276d9 100644 --- a/pkg/resource-handler/controller/shard/containers_test.go +++ b/pkg/resource-handler/controller/shard/containers_test.go @@ -1,12 +1,10 @@ package shard import ( + "reflect" "regexp" - "strings" "testing" - "github.com/google/go-cmp/cmp" - "github.com/stretchr/testify/assert" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/resource" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -14,6 +12,8 @@ import ( "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildMultipoolerContainer(t *testing.T) { @@ -375,9 +375,7 @@ func TestBuildMultipoolerContainer(t *testing.T) { tc.serviceID, ) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("buildMultipoolerContainer() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "buildMultipoolerContainer() mismatch") }) } } @@ -429,9 +427,7 @@ func TestBuildPostgresExporterContainer(t *testing.T) { }, } - if diff := cmp.Diff(want, got); diff != "" { - t.Fatalf("buildPostgresExporterContainer() mismatch (-want +got):\n%s", diff) - } + assert.NewAborting(t).EqDiff(want, got, "buildPostgresExporterContainer() mismatch") } func TestPoolContainers_CustomPostgresSuperuser(t *testing.T) { @@ -519,6 +515,7 @@ func TestPoolContainers_PostgresPasswordFile(t *testing.T) { } func TestPoolContainers_PostgresPasswordSecretRef(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard"}, Spec: multigresv1alpha1.ShardSpec{ @@ -537,15 +534,8 @@ func TestPoolContainers_PostgresPasswordSecretRef(t *testing.T) { if v.Name != PostgresPasswordVolumeName { continue } - if v.Secret == nil { - t.Fatal("postgres password volume should use Secret source") - } - if v.Secret.SecretName != "multigres-admin-password" { - t.Errorf( - "postgres password SecretName = %q, want multigres-admin-password", - v.Secret.SecretName, - ) - } + c.Require().NotNil(v.Secret, "postgres password volume should use Secret source") + c.Eq("multigres-admin-password", v.Secret.SecretName, "postgres password SecretName") if len(v.Secret.Items) != 1 || v.Secret.Items[0].Key != PostgresPasswordSecretKey || v.Secret.Items[0].Path != PostgresPasswordSecretKey { @@ -574,6 +564,7 @@ func TestPoolContainers_PostgresInitSecretsRef(t *testing.T) { } t.Run("ref set projects custom key to default filename", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard"}, Spec: multigresv1alpha1.ShardSpec{ @@ -589,18 +580,9 @@ func TestPoolContainers_PostgresInitSecretsRef(t *testing.T) { volumes := buildPoolVolumes(shard, "zone1") vol := findVolume(volumes, PostgresInitSecretsVolumeName) - if vol == nil { - t.Fatal("expected postgres init-secrets Secret volume in pool volumes") - } - if vol.Secret == nil { - t.Fatal("postgres init-secrets volume should use Secret source") - } - if vol.Secret.SecretName != "multigres-init-secrets" { - t.Errorf( - "postgres init-secrets SecretName = %q, want multigres-init-secrets", - vol.Secret.SecretName, - ) - } + ck.Require().NotNil(vol, "expected postgres init-secrets Secret volume in pool volumes") + ck.Require().NotNil(vol.Secret, "postgres init-secrets volume should use Secret source") + ck.Eq("multigres-init-secrets", vol.Secret.SecretName, "postgres init-secrets SecretName") if len(vol.Secret.Items) != 1 || vol.Secret.Items[0].Key != "custom-key.json" || vol.Secret.Items[0].Path != PostgresInitSecretsFileName { @@ -639,9 +621,8 @@ func TestPoolContainers_PostgresInitSecretsRef(t *testing.T) { volumes := buildPoolVolumes(shard, "zone1") vol := findVolume(volumes, PostgresInitSecretsVolumeName) - if vol == nil { - t.Fatal("expected postgres init-secrets Secret volume in pool volumes") - } + assert.NewAborting(t). + NotNil(vol, "expected postgres init-secrets Secret volume in pool volumes") if len(vol.Secret.Items) != 1 || vol.Secret.Items[0].Key != PostgresInitSecretsFileName { t.Errorf( "postgres init-secrets Secret items = %+v, want default key %q", @@ -662,9 +643,8 @@ func TestPoolContainers_PostgresInitSecretsRef(t *testing.T) { } volumes := buildPoolVolumes(shard, "zone1") - if findVolume(volumes, PostgresInitSecretsVolumeName) != nil { - t.Error("expected no postgres init-secrets Secret volume when ref is nil") - } + assert.NewCollecting(t). + Nil(findVolume(volumes, PostgresInitSecretsVolumeName), "expected no postgres init-secrets Secret volume when ref is nil") c := buildPgctldSidecar(shard, pool) assertNotContainsEnvVar(t, c.Env, "POSTGRES_INIT_SECRETS_FILE") @@ -775,9 +755,7 @@ func TestBuildMultiorchContainer(t *testing.T) { t.Run(name, func(t *testing.T) { got := buildMultiorchContainer(tc.shard, tc.cellName) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("buildMultiorchContainer() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "buildMultiorchContainer() mismatch") }) } } @@ -816,52 +794,54 @@ func otelShard() *multigresv1alpha1.Shard { func TestBuildPgctldSidecar(t *testing.T) { t.Run("default image", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{Spec: multigresv1alpha1.ShardSpec{}} c := buildPgctldSidecar(shard, multigresv1alpha1.PoolSpec{}) - if c.Image != multigresv1alpha1.DefaultPostgresImage { - t.Errorf("Image = %q, want %q", c.Image, multigresv1alpha1.DefaultPostgresImage) - } - assert.Equal(t, DefaultPostgresUID, *c.SecurityContext.RunAsUser) - assert.Equal(t, DefaultPostgresGID, *c.SecurityContext.RunAsGroup) - if c.Command[0] != "/usr/local/bin/pgctld" { - t.Errorf("Command = %v, want /usr/local/bin/pgctld", c.Command) - } + ck.Eq(multigresv1alpha1.DefaultPostgresImage, c.Image, "Image") + ck.EqDeep(DefaultPostgresUID, *c.SecurityContext.RunAsUser) + ck.EqDeep(DefaultPostgresGID, *c.SecurityContext.RunAsGroup) + ck.Eq( + "/usr/local/bin/pgctld", + c.Command[0], + "Command = %v, want /usr/local/bin/pgctld", + c.Command, + ) assertContainsFlag(t, c.Args, "--http-port=15400") - if c.StartupProbe == nil || c.StartupProbe.HTTPGet.Path != "/live" { - t.Errorf("expected StartupProbe to hit /live, got %v", c.StartupProbe) - } - if c.LivenessProbe == nil || c.LivenessProbe.HTTPGet.Path != "/live" { - t.Errorf("expected LivenessProbe to hit /live, got %v", c.LivenessProbe) - } + ck.False( + c.StartupProbe == nil || c.StartupProbe.HTTPGet.Path != "/live", + "expected StartupProbe to hit /live, got %v", + c.StartupProbe, + ) + ck.False( + c.LivenessProbe == nil || c.LivenessProbe.HTTPGet.Path != "/live", + "expected LivenessProbe to hit /live, got %v", + c.LivenessProbe, + ) wantReadinessCommand := []string{ "pg_isready", "-h", "/var/lib/pooler/pg_sockets", "-p", "5432", } - if c.ReadinessProbe == nil || + ck.False(c.ReadinessProbe == nil || c.ReadinessProbe.Exec == nil || - !assert.ObjectsAreEqual(c.ReadinessProbe.Exec.Command, wantReadinessCommand) { - t.Errorf( - "expected ReadinessProbe command %v, got %v", + !reflect.DeepEqual( + c.ReadinessProbe.Exec.Command, wantReadinessCommand, - c.ReadinessProbe, - ) - } + ), "expected ReadinessProbe command %v, got %v", wantReadinessCommand, c.ReadinessProbe) }) t.Run("custom image", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ Spec: multigresv1alpha1.ShardSpec{ Images: multigresv1alpha1.ShardImages{Postgres: "custom/pgctld:v1"}, }, } c := buildPgctldSidecar(shard, multigresv1alpha1.PoolSpec{}) - if c.Image != "custom/pgctld:v1" { - t.Errorf("Image = %q, want %q", c.Image, "custom/pgctld:v1") - } + ck.Eq("custom/pgctld:v1", c.Image, "Image") // The numeric identity must not depend on the image reference: pgctld // declares USER postgres by name, so without it RunAsNonRoot makes the // kubelet reject the container regardless of which tag is in use. - assert.Equal(t, ptr.To(DefaultPostgresUID), c.SecurityContext.RunAsUser) - assert.Equal(t, ptr.To(DefaultPostgresGID), c.SecurityContext.RunAsGroup) + ck.EqDeep(ptr.To(DefaultPostgresUID), c.SecurityContext.RunAsUser) + ck.EqDeep(ptr.To(DefaultPostgresGID), c.SecurityContext.RunAsGroup) }) t.Run("with observability", func(t *testing.T) { @@ -968,10 +948,8 @@ func TestBuildPgctldSidecar(t *testing.T) { assertContainsEnvVar(t, c.Env, "POSTGRES_INITDB_ARGS") for _, e := range c.Env { if e.Name == "POSTGRES_INITDB_ARGS" { - if e.Value != "--locale-provider=icu --icu-locale=en_US.UTF-8" { - t.Errorf("POSTGRES_INITDB_ARGS = %q, want %q", - e.Value, "--locale-provider=icu --icu-locale=en_US.UTF-8") - } + assert.NewCollecting(t). + Eq("--locale-provider=icu --icu-locale=en_US.UTF-8", e.Value, "POSTGRES_INITDB_ARGS") return } } @@ -999,13 +977,8 @@ func TestBuildPgctldSidecar(t *testing.T) { c := buildPgctldSidecar(shard, multigresv1alpha1.PoolSpec{}) assertContainsEnvVar(t, c.Env, "POSTGRES_INITDB_EXTRA_CONF") for _, e := range c.Env { - if e.Name == "POSTGRES_INITDB_EXTRA_CONF" && e.Value != PostgresConfigFilePath { - t.Errorf( - "POSTGRES_INITDB_EXTRA_CONF = %q, want %q", - e.Value, - PostgresConfigFilePath, - ) - } + assert.NewCollecting(t). + False(e.Name == "POSTGRES_INITDB_EXTRA_CONF" && e.Value != PostgresConfigFilePath, "POSTGRES_INITDB_EXTRA_CONF = %q, want %q", e.Value, PostgresConfigFilePath) } }) } @@ -1019,20 +992,13 @@ func TestBuildPgctldSidecar(t *testing.T) { "without ref": {Spec: multigresv1alpha1.ShardSpec{}}, } { t.Run(name, func(t *testing.T) { + ck := assert.NewCollecting(t) c := buildPgctldSidecar(shard, multigresv1alpha1.PoolSpec{}) assertContainsVolumeMount(t, c.VolumeMounts, PostgresConfigVolumeName) for _, m := range c.VolumeMounts { if m.Name == PostgresConfigVolumeName { - if m.MountPath != PostgresConfigMountPath { - t.Errorf( - "postgres config mount path = %q, want %q", - m.MountPath, - PostgresConfigMountPath, - ) - } - if !m.ReadOnly { - t.Error("postgres config volume mount should be read-only") - } + ck.Eq(PostgresConfigMountPath, m.MountPath, "postgres config mount path") + ck.True(m.ReadOnly, "postgres config volume mount should be read-only") } } }) @@ -1043,18 +1009,14 @@ func TestBuildPgctldSidecar(t *testing.T) { func TestS3EnvVars(t *testing.T) { t.Run("nil backup returns nil", func(t *testing.T) { got := s3EnvVars(nil) - if got != nil { - t.Errorf("s3EnvVars(nil) = %v, want nil", got) - } + assert.NewCollecting(t).Nil(got, "s3EnvVars(nil)") }) t.Run("filesystem backup returns nil", func(t *testing.T) { got := s3EnvVars(&multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, }) - if got != nil { - t.Errorf("s3EnvVars(filesystem) = %v, want nil", got) - } + assert.NewCollecting(t).Nil(got, "s3EnvVars(filesystem)") }) t.Run("s3 with region only", func(t *testing.T) { @@ -1065,12 +1027,12 @@ func TestS3EnvVars(t *testing.T) { Region: "eu-west-1", }, }) - if len(got) != 1 || got[0].Name != "AWS_REGION" || got[0].Value != "eu-west-1" { - t.Errorf("s3EnvVars(region-only) = %v, want [{AWS_REGION eu-west-1}]", got) - } + assert.NewCollecting(t). + False(len(got) != 1 || got[0].Name != "AWS_REGION" || got[0].Value != "eu-west-1", "s3EnvVars(region-only) = %v, want [{AWS_REGION eu-west-1}]", got) }) t.Run("s3 with credentials secret", func(t *testing.T) { + c := assert.NewCollecting(t) got := s3EnvVars(&multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeS3, S3: &multigresv1alpha1.S3BackupConfig{ @@ -1079,9 +1041,7 @@ func TestS3EnvVars(t *testing.T) { CredentialsSecret: "my-secret", }, }) - if len(got) != 3 { - t.Fatalf("s3EnvVars(full) returned %d vars, want 3", len(got)) - } + c.Require().Len(got, 3, "s3EnvVars(full) returned %d vars, want 3", len(got)) assertContainsEnvVar(t, got, "AWS_REGION") assertContainsEnvVar(t, got, "AWS_ACCESS_KEY_ID") assertContainsEnvVar(t, got, "AWS_SECRET_ACCESS_KEY") @@ -1089,13 +1049,9 @@ func TestS3EnvVars(t *testing.T) { // Verify it references the correct secret for _, e := range got { if e.Name == "AWS_ACCESS_KEY_ID" { - if e.ValueFrom == nil || e.ValueFrom.SecretKeyRef == nil { - t.Fatal("AWS_ACCESS_KEY_ID missing SecretKeyRef") - } - if e.ValueFrom.SecretKeyRef.Name != "my-secret" { - t.Errorf("AWS_ACCESS_KEY_ID secret = %q, want %q", - e.ValueFrom.SecretKeyRef.Name, "my-secret") - } + c.Require(). + False(e.ValueFrom == nil || e.ValueFrom.SecretKeyRef == nil, "AWS_ACCESS_KEY_ID missing SecretKeyRef") + c.Eq("my-secret", e.ValueFrom.SecretKeyRef.Name, "AWS_ACCESS_KEY_ID secret") } } }) @@ -1112,12 +1068,8 @@ func TestS3EnvVars(t *testing.T) { }, }) // Should only have AWS_REGION, no credential env vars - if len(got) != 1 { - t.Fatalf( - "s3EnvVars(serviceAccountName-only) returned %d vars, want 1 (AWS_REGION only)", - len(got), - ) - } + assert.NewAborting(t). + Len(got, 1, "s3EnvVars(serviceAccountName-only) returned %d vars, want 1 (AWS_REGION only)", len(got)) assertContainsEnvVar(t, got, "AWS_REGION") assertNotContainsEnvVar(t, got, "AWS_ACCESS_KEY_ID") assertNotContainsEnvVar(t, got, "AWS_SECRET_ACCESS_KEY") @@ -1133,15 +1085,14 @@ func TestS3EnvVars(t *testing.T) { }, }) // Should have 2 vars: KEY_ID and SECRET_KEY, but no REGION - if len(got) != 2 { - t.Fatalf("s3EnvVars(no-region) returned %d vars, want 2", len(got)) - } + assert.NewAborting(t).Len(got, 2, "s3EnvVars(no-region) returned %d vars, want 2", len(got)) assertNotContainsEnvVar(t, got, "AWS_REGION") assertContainsEnvVar(t, got, "AWS_ACCESS_KEY_ID") }) } func TestBuildSharedBackupVolume_S3(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Labels: map[string]string{"multigres.com/cluster": "test-cluster"}, @@ -1161,15 +1112,9 @@ func TestBuildSharedBackupVolume_S3(t *testing.T) { } vol := buildSharedBackupVolume(shard) - if vol.Name != BackupVolumeName { - t.Errorf("volume name = %q, want %q", vol.Name, BackupVolumeName) - } - if vol.EmptyDir == nil { - t.Error("S3 backup volume should use EmptyDir, got PVC or other source") - } - if vol.PersistentVolumeClaim != nil { - t.Error("S3 backup volume should NOT use PersistentVolumeClaim") - } + c.Eq(BackupVolumeName, vol.Name, "volume name") + c.NotNil(vol.EmptyDir, "S3 backup volume should use EmptyDir, got PVC or other source") + c.Nil(vol.PersistentVolumeClaim, "S3 backup volume should NOT use PersistentVolumeClaim") } func assertContainsFlag(t *testing.T, args []string, want string) { @@ -1184,16 +1129,14 @@ func assertContainsFlag(t *testing.T, args []string, want string) { func assertFlagValue(t *testing.T, args []string, flag, want string) { t.Helper() + c := assert.NewCollecting(t) for i, arg := range args { if arg != flag { continue } - if i+1 >= len(args) { - t.Fatalf("%s has no value in args %v", flag, args) - } - if got := args[i+1]; got != want { - t.Errorf("%s = %q, want %q", flag, got, want) - } + c.Require().Less(len(args), i+1, "%s has no value in args %v", flag, args) + got := args[i+1] + c.Eq(want, got, "%s = %q, want", flag, got) return } t.Errorf("args %v does not contain flag %q", args, flag) @@ -1367,9 +1310,7 @@ func assertEnvVarValue(t *testing.T, envVars []corev1.EnvVar, name, want string) t.Helper() for _, e := range envVars { if e.Name == name { - if e.Value != want { - t.Errorf("env var %q = %q, want %q", name, e.Value, want) - } + assert.NewCollecting(t).Eq(want, e.Value, "env var %q = %q, want", name, e.Value) return } } @@ -1380,9 +1321,7 @@ func assertOTELResourceAttribute(t *testing.T, envVars []corev1.EnvVar, want str t.Helper() for _, e := range envVars { if e.Name == "OTEL_RESOURCE_ATTRIBUTES" { - if !strings.Contains(e.Value, want) { - t.Errorf("OTEL_RESOURCE_ATTRIBUTES = %q, want it to contain %q", e.Value, want) - } + assert.NewCollecting(t).StrContains(e.Value, want, "OTEL_RESOURCE_ATTRIBUTES") return } } @@ -1421,16 +1360,13 @@ func assertContainsVolumeMount(t *testing.T, mounts []corev1.VolumeMount, name s func assertReadOnlyVolumeMount(t *testing.T, mounts []corev1.VolumeMount, name, mountPath string) { t.Helper() + c := assert.NewCollecting(t) for _, m := range mounts { if m.Name != name { continue } - if m.MountPath != mountPath { - t.Errorf("mount %q path = %q, want %q", name, m.MountPath, mountPath) - } - if !m.ReadOnly { - t.Errorf("mount %q should be read-only", name) - } + c.Eq(mountPath, m.MountPath, "mount %q path = %q, want", name, m.MountPath) + c.True(m.ReadOnly, "mount %q should be read-only", name) return } t.Errorf("volume mounts %v does not contain mount %q", mounts, name) @@ -1450,12 +1386,11 @@ func TestBuildPgBackRestCertVolume(t *testing.T) { t.Run("nil when no backup", func(t *testing.T) { shard := &multigresv1alpha1.Shard{Spec: multigresv1alpha1.ShardSpec{}} vol := buildPgBackRestCertVolume(shard) - if vol != nil { - t.Fatalf("expected nil volume when no backup, got %+v", vol) - } + assert.NewAborting(t).Nil(vol, "expected nil volume when no backup, got") }) t.Run("auto-generated projected volume", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard"}, Spec: multigresv1alpha1.ShardSpec{ @@ -1465,44 +1400,25 @@ func TestBuildPgBackRestCertVolume(t *testing.T) { }, } vol := buildPgBackRestCertVolume(shard) - if vol == nil { - t.Fatal("expected non-nil volume for auto-generated certs") - } - if vol.Name != PgBackRestCertVolumeName { - t.Errorf("volume name = %q, want %q", vol.Name, PgBackRestCertVolumeName) - } - if vol.Projected == nil { - t.Fatal("expected projected volume source for auto-generated certs") - } - if len(vol.Projected.Sources) != 2 { - t.Fatalf("expected 2 projection sources, got %d", len(vol.Projected.Sources)) - } + c.Require().NotNil(vol, "expected non-nil volume for auto-generated certs") + c.Eq(PgBackRestCertVolumeName, vol.Name, "volume name") + c.Require(). + NotNil(vol.Projected, "expected projected volume source for auto-generated certs") + c.Require(). + Len(vol.Projected.Sources, 2, "expected 2 projection sources, got %d", len(vol.Projected.Sources)) // Verify CA source caSource := vol.Projected.Sources[0] - if caSource.Secret.Name != "test-shard-pgbackrest-ca" { - t.Errorf( - "CA secret name = %q, want %q", - caSource.Secret.Name, - "test-shard-pgbackrest-ca", - ) - } + c.Eq("test-shard-pgbackrest-ca", caSource.Secret.Name, "CA secret name") if len(caSource.Secret.Items) != 1 || caSource.Secret.Items[0].Key != "ca.crt" { t.Errorf("CA items = %+v, want [{Key:ca.crt Path:ca.crt}]", caSource.Secret.Items) } // Verify TLS source (key renaming) tlsSource := vol.Projected.Sources[1] - if tlsSource.Secret.Name != "test-shard-pgbackrest-tls" { - t.Errorf( - "TLS secret name = %q, want %q", - tlsSource.Secret.Name, - "test-shard-pgbackrest-tls", - ) - } - if len(tlsSource.Secret.Items) != 2 { - t.Fatalf("expected 2 TLS items, got %d", len(tlsSource.Secret.Items)) - } + c.Eq("test-shard-pgbackrest-tls", tlsSource.Secret.Name, "TLS secret name") + c.Require(). + Len(tlsSource.Secret.Items, 2, "expected 2 TLS items, got %d", len(tlsSource.Secret.Items)) if tlsSource.Secret.Items[0].Key != "tls.crt" || tlsSource.Secret.Items[0].Path != "pgbackrest.crt" { t.Errorf( @@ -1520,6 +1436,7 @@ func TestBuildPgBackRestCertVolume(t *testing.T) { }) t.Run("user-provided Secret volume", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ Spec: multigresv1alpha1.ShardSpec{ Backup: &multigresv1alpha1.BackupConfig{ @@ -1534,33 +1451,19 @@ func TestBuildPgBackRestCertVolume(t *testing.T) { }, } vol := buildPgBackRestCertVolume(shard) - if vol == nil { - t.Fatal("expected non-nil volume for user-provided certs") - } - if vol.Projected == nil { - t.Fatal( - "expected projected volume for user-provided certs (key renaming for cert-manager compat)", - ) - } - if len(vol.Projected.Sources) != 1 { - t.Fatalf( - "expected 1 projection source for user-provided, got %d", - len(vol.Projected.Sources), - ) - } + c.Require().NotNil(vol, "expected non-nil volume for user-provided certs") + c.Require(). + NotNil(vol.Projected, "expected projected volume for user-provided certs (key renaming for cert-manager compat)") + c.Require(). + Len(vol.Projected.Sources, 1, "expected 1 projection source for user-provided, got %d", len(vol.Projected.Sources)) src := vol.Projected.Sources[0] - if src.Secret.Name != "my-custom-certs" { - t.Errorf("secret name = %q, want %q", src.Secret.Name, "my-custom-certs") - } - if len(src.Secret.Items) != 3 { - t.Fatalf( - "expected 3 items (ca.crt, tls.crt→pgbackrest.crt, tls.key→pgbackrest.key), got %d", - len(src.Secret.Items), - ) - } + c.Eq("my-custom-certs", src.Secret.Name, "secret name") + c.Require(). + Len(src.Secret.Items, 3, "expected 3 items (ca.crt, tls.crt→pgbackrest.crt, tls.key→pgbackrest.key), got %d", len(src.Secret.Items)) }) t.Run("auto-generated when PgBackRestTLS is nil", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "shard1"}, Spec: multigresv1alpha1.ShardSpec{ @@ -1571,15 +1474,13 @@ func TestBuildPgBackRestCertVolume(t *testing.T) { }, } vol := buildPgBackRestCertVolume(shard) - if vol == nil { - t.Fatal("expected non-nil volume when PgBackRestTLS is nil (auto-generated)") - } - if vol.Projected == nil { - t.Error("expected projected volume for auto-generated fallback") - } + c.Require(). + NotNil(vol, "expected non-nil volume when PgBackRestTLS is nil (auto-generated)") + c.NotNil(vol.Projected, "expected projected volume for auto-generated fallback") }) t.Run("auto-generated when SecretName is empty", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "shard1"}, Spec: multigresv1alpha1.ShardSpec{ @@ -1592,12 +1493,8 @@ func TestBuildPgBackRestCertVolume(t *testing.T) { }, } vol := buildPgBackRestCertVolume(shard) - if vol == nil { - t.Fatal("expected non-nil volume when SecretName is empty (auto-generated)") - } - if vol.Projected == nil { - t.Error("expected projected volume for auto-generated fallback") - } + c.Require().NotNil(vol, "expected non-nil volume when SecretName is empty (auto-generated)") + c.NotNil(vol.Projected, "expected projected volume for auto-generated fallback") }) } @@ -1605,9 +1502,7 @@ func TestBuildPgBackRestCipherKeyVolume(t *testing.T) { t.Run("nil when no backup", func(t *testing.T) { shard := &multigresv1alpha1.Shard{Spec: multigresv1alpha1.ShardSpec{}} vol := buildPgBackRestCipherKeyVolume(shard) - if vol != nil { - t.Fatalf("expected nil volume when no backup, got %+v", vol) - } + assert.NewAborting(t).Nil(vol, "expected nil volume when no backup, got") }) t.Run("nil when backup configured but encryption disabled", func(t *testing.T) { @@ -1619,12 +1514,11 @@ func TestBuildPgBackRestCipherKeyVolume(t *testing.T) { }, } vol := buildPgBackRestCipherKeyVolume(shard) - if vol != nil { - t.Fatalf("expected nil volume when encryption disabled, got %+v", vol) - } + assert.NewAborting(t).Nil(vol, "expected nil volume when encryption disabled, got") }) t.Run("user-provided secret volume", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard"}, Spec: multigresv1alpha1.ShardSpec{ @@ -1637,18 +1531,10 @@ func TestBuildPgBackRestCipherKeyVolume(t *testing.T) { }, } vol := buildPgBackRestCipherKeyVolume(shard) - if vol == nil { - t.Fatal("expected non-nil volume for user-provided cipher key") - } - if vol.Name != PgBackRestCipherKeyVolumeName { - t.Errorf("volume name = %q, want %q", vol.Name, PgBackRestCipherKeyVolumeName) - } - if vol.Secret == nil { - t.Fatal("expected Secret volume source") - } - if vol.Secret.SecretName != "my-cipher-secret" { - t.Errorf("secret name = %q, want %q", vol.Secret.SecretName, "my-cipher-secret") - } + c.Require().NotNil(vol, "expected non-nil volume for user-provided cipher key") + c.Eq(PgBackRestCipherKeyVolumeName, vol.Name, "volume name") + c.Require().NotNil(vol.Secret, "expected Secret volume source") + c.Eq("my-cipher-secret", vol.Secret.SecretName, "secret name") if vol.Secret.DefaultMode == nil || *vol.Secret.DefaultMode != 0o444 { t.Errorf("defaultMode = %v, want 0444", vol.Secret.DefaultMode) } @@ -1671,9 +1557,8 @@ func TestPgctldContainer_PgBackRestCertArgs(t *testing.T) { // Verify volume mount is read-only for _, m := range c.VolumeMounts { - if m.Name == PgBackRestCertVolumeName && !m.ReadOnly { - t.Error("pgbackrest cert volume mount should be read-only") - } + assert.NewCollecting(t). + False(m.Name == PgBackRestCertVolumeName && !m.ReadOnly, "pgbackrest cert volume mount should be read-only") } }) @@ -1777,9 +1662,8 @@ func TestMultipoolerSidecar_PgBackRestCertArgs(t *testing.T) { ) assertContainsVolumeMount(t, c.VolumeMounts, PgBackRestCipherKeyVolumeName) for _, m := range c.VolumeMounts { - if m.Name == PgBackRestCipherKeyVolumeName && !m.ReadOnly { - t.Error("pgbackrest cipher key volume mount should be read-only") - } + assert.NewCollecting(t). + False(m.Name == PgBackRestCipherKeyVolumeName && !m.ReadOnly, "pgbackrest cipher key volume mount should be read-only") } }) @@ -1803,6 +1687,7 @@ func TestMultipoolerSidecar_PgBackRestCertArgs(t *testing.T) { func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { t.Run("postgres password secret volume present", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -1820,16 +1705,8 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { if v.Name != PostgresPasswordVolumeName { continue } - if v.Secret == nil { - t.Fatal("postgres password volume should use Secret source") - } - if v.Secret.SecretName != "multigres-admin-password" { - t.Errorf( - "postgres password SecretName = %q, want %q", - v.Secret.SecretName, - "multigres-admin-password", - ) - } + c.Require().NotNil(v.Secret, "postgres password volume should use Secret source") + c.Eq("multigres-admin-password", v.Secret.SecretName, "postgres password SecretName") if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf("postgres password defaultMode = %v, want 0444", v.Secret.DefaultMode) } @@ -1866,9 +1743,8 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { break } } - if !found { - t.Error("expected pgbackrest-certs volume when backup configured") - } + assert.NewCollecting(t). + True(found, "expected pgbackrest-certs volume when backup configured") }) t.Run("no cert volume when no backup", func(t *testing.T) { @@ -1880,9 +1756,8 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { } volumes := buildPoolVolumes(shard, "zone1") for _, v := range volumes { - if v.Name == PgBackRestCertVolumeName { - t.Error("cert volume should not be present when no backup configured") - } + assert.NewCollecting(t). + NotEq(PgBackRestCertVolumeName, v.Name, "cert volume should not be present when no backup configured") } }) @@ -1909,9 +1784,8 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { break } } - if !found { - t.Error("expected pgbackrest-cipher volume when encryption configured") - } + assert.NewCollecting(t). + True(found, "expected pgbackrest-cipher volume when encryption configured") }) t.Run("no cipher key volume when encryption disabled", func(t *testing.T) { @@ -1928,13 +1802,13 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { } volumes := buildPoolVolumes(shard, "zone1") for _, v := range volumes { - if v.Name == PgBackRestCipherKeyVolumeName { - t.Error("cipher key volume should not be present when encryption disabled") - } + assert.NewCollecting(t). + NotEq(PgBackRestCipherKeyVolumeName, v.Name, "cipher key volume should not be present when encryption disabled") } }) t.Run("always projects the operator-owned postgres config ConfigMap", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -1949,13 +1823,13 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { for _, v := range volumes { if v.Name == PostgresConfigVolumeName { found = true - if v.ConfigMap == nil { - t.Fatal("postgres config volume should use ConfigMap source") - } - if v.ConfigMap.Name != PostgresConfigMapName("test-shard") { - t.Errorf("postgres config ConfigMap name = %q, want %q", - v.ConfigMap.Name, PostgresConfigMapName("test-shard")) - } + c.Require(). + NotNil(v.ConfigMap, "postgres config volume should use ConfigMap source") + c.Eq( + PostgresConfigMapName("test-shard"), + v.ConfigMap.Name, + "postgres config ConfigMap name", + ) if len(v.ConfigMap.Items) != 1 || v.ConfigMap.Items[0].Key != PostgresConfigMapKey || v.ConfigMap.Items[0].Path != "postgresql.conf" { @@ -1968,12 +1842,11 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { break } } - if !found { - t.Error("expected postgres-config volume in pool volumes") - } + c.True(found, "expected postgres-config volume in pool volumes") }) t.Run("internal multipooler tls volume is present when enabled", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -1990,17 +1863,9 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { if v.Name != ShardTLSVolumeName { continue } - if v.Secret == nil { - t.Fatal("shard TLS volume should use Secret source") - } + c.Require().NotNil(v.Secret, "shard TLS volume should use Secret source") wantSecretName := "multipooler.test.default.multigres.internal" - if v.Secret.SecretName != wantSecretName { - t.Errorf( - "shard TLS secret = %q, want internal secret %q", - v.Secret.SecretName, - wantSecretName, - ) - } + c.Eq(wantSecretName, v.Secret.SecretName, "shard TLS secret") if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf( "shard TLS secret defaultMode = %v, want 0444", @@ -2027,12 +1892,8 @@ func TestBuildPoolVolumes_CertVolumePresence(t *testing.T) { } for _, volume := range buildPoolVolumes(shard, "zone1") { - if volume.Name == ShardTLSVolumeName { - t.Errorf( - "internal TLS volume should be absent for config %+v", - internalTLS, - ) - } + assert.NewCollecting(t). + NotEq(ShardTLSVolumeName, volume.Name, "internal TLS volume should be absent for config %+v", internalTLS) } } }) @@ -2043,34 +1904,28 @@ func TestBuildPoolServiceID(t *testing.T) { podName := "minimal-postgres-default-0-inf-pool-default-zone-a-a3a0d77b-1" id1 := BuildPoolServiceID(podName) id2 := BuildPoolServiceID(podName) - if id1 != id2 { - t.Errorf("non-deterministic: %q != %q", id1, id2) - } + assert.NewCollecting(t).Eq(id2, id1, "non-deterministic") }) t.Run("format", func(t *testing.T) { id := BuildPoolServiceID("some-pod-name") pattern := regexp.MustCompile(`^p-[0-9a-f]{8}$`) - if !pattern.MatchString(id) { - t.Errorf("BuildPoolServiceID(%q) = %q, want format p-[0-9a-f]{8}", "some-pod-name", id) - } + assert.NewCollecting(t). + True(pattern.MatchString(id), "BuildPoolServiceID(%q) = %q, want format p-[0-9a-f]{8}", "some-pod-name", id) }) t.Run("different inputs produce different outputs", func(t *testing.T) { id1 := BuildPoolServiceID("pod-a") id2 := BuildPoolServiceID("pod-b") - if id1 == id2 { - t.Errorf("collision: BuildPoolServiceID(%q) == BuildPoolServiceID(%q) == %q", - "pod-a", "pod-b", id1) - } + assert.NewCollecting(t). + NotEq(id2, id1, "collision: BuildPoolServiceID(%q) == BuildPoolServiceID(%q) ==", "pod-a", "pod-b") }) t.Run("length is always 10", func(t *testing.T) { for _, name := range []string{"a", "short", "a-very-long-pod-name-that-goes-on-and-on"} { id := BuildPoolServiceID(name) - if len(id) != 10 { - t.Errorf("BuildPoolServiceID(%q) = %q (len %d), want len 10", name, id, len(id)) - } + assert.NewCollecting(t). + Len(id, 10, "BuildPoolServiceID(%q) = %q (len %d), want len 10", name, id, len(id)) } }) } diff --git a/pkg/resource-handler/controller/shard/disruption_test.go b/pkg/resource-handler/controller/shard/disruption_test.go index cea08c4b..cd7ad29b 100644 --- a/pkg/resource-handler/controller/shard/disruption_test.go +++ b/pkg/resource-handler/controller/shard/disruption_test.go @@ -22,6 +22,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/poolerclient" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) type disruptionTopo struct { @@ -108,40 +110,31 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { {name: "missing committed member blocks cleanup", stillInCohort: true}, } { t.Run(tc.name, func(t *testing.T) { + c := assert.NewAborting(t) r, shard, groups, rpc, responses := fourToTwoFixture(t) name := BuildPoolPodName(shard, "main", "b", 1) key := client.ObjectKey{Namespace: shard.Namespace, Name: name} pod := &corev1.Pod{} - if err := r.Get(t.Context(), key, pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Get(t.Context(), key, pod)) pod.Status.Phase = corev1.PodPending pod.Status.Conditions = []corev1.PodCondition{ {Type: corev1.PodReady, Status: corev1.ConditionFalse}, } - if err := r.Status().Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Status().Update(t.Context(), pod)) groups["b"][name] = pod.DeepCopy() // Keep the caller's pod stale: scheduling evidence must come from // the uncached API reader used by the disruption preflight. pod.Spec.NodeName = tc.nodeName - if err := r.Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Update(t.Context(), pod)) if tc.scheduled { pod.Status.Conditions = append(pod.Status.Conditions, corev1.PodCondition{ Type: corev1.PodScheduled, Status: corev1.ConditionTrue, }) - if err := r.Status().Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Status().Update(t.Context(), pod)) } r.APIReader = r.Client store, err := r.CreateTopoStore(shard) - if err != nil { - t.Fatal(err) - } + c.NoError(err) topo := store.(*disruptionTopo) topo.poolers = slices.DeleteFunc( topo.poolers, @@ -168,14 +161,13 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { groups["b"][name], &shardRolloutTracker{}, ) - if err != nil || allowed != tc.wantAction { - t.Fatalf( - "preflight with stale target: allowed=%v err=%v; want %v", - allowed, - err, - tc.wantAction, - ) - } + c.False( + err != nil || allowed != tc.wantAction, + "preflight with stale target: allowed=%v err=%v; want %v", + allowed, + err, + tc.wantAction, + ) tracker := &shardRolloutTracker{} action, _, err := r.handleScaleDown( t.Context(), @@ -188,9 +180,7 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { false, tracker, ) - if err != nil { - t.Fatal(err) - } + c.NoError(err) if action != tc.wantAction || tracker.waitingForRecovery == tc.wantAction { t.Fatalf( "action=%v waitingForRecovery=%v; want action=%v", @@ -202,9 +192,7 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { if !tc.wantAction { return } - if !tracker.HasStarted() { - t.Fatal("cleanup must reserve this pass's disruption") - } + c.True(tracker.HasStarted(), "cleanup must reserve this pass's disruption") // Complete the ordinary drain/PVC cleanup path, then ensure the // other cell's excess primary is no longer blocked by this pod. for range 3 { @@ -212,9 +200,7 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { t.Fatal(err) } } - if err := r.Get(t.Context(), key, pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Get(t.Context(), key, pod)) groups["b"][name] = pod if _, _, err := r.handleScaleDown( t.Context(), @@ -243,9 +229,12 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { false, &shardRolloutTracker{}, ) - if err != nil || !action { - t.Fatalf("next scale-down did not resume: action=%v err=%v", action, err) - } + c.False( + err != nil || !action, + "next scale-down did not resume: action=%v err=%v", + action, + err, + ) }) } } @@ -253,28 +242,21 @@ func TestScaleDownUnregisteredExtra(t *testing.T) { func TestScaleDownWaitsForOtherCellsExtraPodDrain(t *testing.T) { for _, state := range []string{metadata.DrainStateRequested, metadata.DrainStateDraining, metadata.DrainStateAcknowledged, metadata.DrainStateReadyForDeletion, "terminating"} { t.Run(state, func(t *testing.T) { + c := assert.NewAborting(t) r, shard, groups, _, _ := fourToTwoFixture(t) pod := &corev1.Pod{} key := client.ObjectKey{ Namespace: shard.Namespace, Name: BuildPoolPodName(shard, "main", "a", 1), } - if err := r.Get(t.Context(), key, pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Get(t.Context(), key, pod)) if state == "terminating" { pod.Finalizers = []string{"test/hold"} - if err := r.Update(t.Context(), pod); err != nil { - t.Fatal(err) - } - if err := r.Delete(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Update(t.Context(), pod)) + c.NoError(r.Delete(t.Context(), pod)) } else { pod.Annotations = map[string]string{metadata.AnnotationDrainState: state} - if err := r.Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Update(t.Context(), pod)) } // No shared in-memory tracker survives this new reconciliation. action, _, err := r.handleScaleDown( @@ -288,14 +270,13 @@ func TestScaleDownWaitsForOtherCellsExtraPodDrain(t *testing.T) { false, &shardRolloutTracker{}, ) - if err != nil || action { - t.Fatalf("overlapping drain: action=%v err=%v", action, err) - } + c.False(err != nil || action, "overlapping drain: action=%v err=%v", action, err) }) } } func TestScaleDownReplicaFirstAndWaitsForCohortRecovery(t *testing.T) { + c := assert.NewAborting(t) r, shard, groups, rpc, responses := fourToTwoFixture(t) run := func(cell string) (bool, *shardRolloutTracker) { t.Helper() @@ -311,9 +292,7 @@ func TestScaleDownReplicaFirstAndWaitsForCohortRecovery(t *testing.T) { false, tracker, ) - if err != nil { - t.Fatal(err) - } + c.NoError(err) return action, tracker } if action, _ := run("a"); action { @@ -326,9 +305,7 @@ func TestScaleDownReplicaFirstAndWaitsForCohortRecovery(t *testing.T) { t.Fatal("second reconcile started an overlapping drain") } removed := groups["b"][BuildPoolPodName(shard, "main", "b", 1)] - if err := r.Delete(t.Context(), removed); err != nil { - t.Fatal(err) - } + c.NoError(r.Delete(t.Context(), removed)) delete(groups["b"], removed.Name) if action, tracker := run("a"); action || !tracker.waitingForRecovery { t.Fatal("must requeue while deleted member remains in committed cohort") @@ -343,9 +320,8 @@ func TestScaleDownReplicaFirstAndWaitsForCohortRecovery(t *testing.T) { ) rpc.SetStatusResponse(topoclient.ComponentIDString(response.ConsensusStatus.Id), response) } - if action, _ := run("a"); !action { - t.Fatal("scale-down did not resume after cohort recovery") - } + action, _ := run("a") + c.True(action, "scale-down did not resume after cohort recovery") } func TestDisruptionReadsUncachedState(t *testing.T) { @@ -364,9 +340,8 @@ func TestDisruptionReadsUncachedState(t *testing.T) { } r.APIReader = fake.NewClientBuilder().WithScheme(r.Scheme).WithObjects(objects...).Build() healthy, err := r.isShardHealthy(t.Context(), shard) - if err != nil || healthy { - t.Fatalf("uncached drain ignored: healthy=%v err=%v", healthy, err) - } + assert.NewAborting(t). + False(err != nil || healthy, "uncached drain ignored: healthy=%v err=%v", healthy, err) } func TestDisruptionWithoutObservationsFailsClosed(t *testing.T) { @@ -390,22 +365,19 @@ func TestDisruptionWithoutObservationsFailsClosed(t *testing.T) { } func TestScaleDownCleanupReservesShardDisruption(t *testing.T) { + c := assert.NewAborting(t) r, shard, groups, _, _ := fourToTwoFixture(t) name := BuildPoolPodName(shard, "main", "b", 1) pod := &corev1.Pod{} - if err := r.Get( + c.NoError(r.Get( t.Context(), client.ObjectKey{Namespace: shard.Namespace, Name: name}, pod, - ); err != nil { - t.Fatal(err) - } + )) pod.Annotations = map[string]string{ metadata.AnnotationDrainState: metadata.DrainStateReadyForDeletion, } - if err := r.Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Update(t.Context(), pod)) groups["b"][name] = pod tracker := &shardRolloutTracker{} action, _, err := r.handleScaleDown( @@ -419,9 +391,12 @@ func TestScaleDownCleanupReservesShardDisruption(t *testing.T) { false, tracker, ) - if err != nil || !action || !tracker.HasStarted() { - t.Fatalf("cleanup did not reserve disruption: action=%v err=%v", action, err) - } + c.False( + err != nil || !action || !tracker.HasStarted(), + "cleanup did not reserve disruption: action=%v err=%v", + action, + err, + ) action, _, err = r.handleScaleDown( t.Context(), shard, @@ -433,9 +408,12 @@ func TestScaleDownCleanupReservesShardDisruption(t *testing.T) { false, tracker, ) - if err != nil || action { - t.Fatalf("cleanup allowed another drain in same pass: action=%v err=%v", action, err) - } + c.False( + err != nil || action, + "cleanup allowed another drain in same pass: action=%v err=%v", + action, + err, + ) } func (s *disruptionTopo) Close() error { return nil } @@ -464,10 +442,9 @@ func observeHealthyDisruption( poolName, cellName string, ) (*rpcclient.FakeClient, []*md.StatusResponse) { t.Helper() + c := assert.NewAborting(t) pods := &corev1.PodList{} - if err := r.List(t.Context(), pods, client.InNamespace(shard.Namespace)); err != nil { - t.Fatal(err) - } + c.NoError(r.List(t.Context(), pods, client.InNamespace(shard.Namespace))) slices.SortFunc(pods.Items, func(a, b corev1.Pod) int { if a.Name < b.Name { return -1 @@ -506,9 +483,7 @@ func observeHealthyDisruption( }, }, } - if err := r.Create(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Create(t.Context(), pod)) pods.Items = append(pods.Items, *pod) } if missing > 0 { @@ -534,16 +509,12 @@ func observeHealthyDisruption( if pod.Labels[metadata.LabelMultigresCell] == "" { pod.Labels[metadata.LabelMultigresCell] = cellName } - if err := r.Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Update(t.Context(), pod)) if len(pod.Status.Conditions) == 0 { pod.Status.Conditions = []corev1.PodCondition{ {Type: corev1.PodReady, Status: corev1.ConditionTrue}, } - if err := r.Client.Status().Update(t.Context(), pod); err != nil { - t.Fatal(err) - } + c.NoError(r.Client.Status().Update(t.Context(), pod)) } ids[i] = &cm.ID{Name: pod.Name, Cell: pod.Labels[metadata.LabelMultigresCell]} if shard.Status.PodRoles[pod.Name] == "PRIMARY" { diff --git a/pkg/resource-handler/controller/shard/integration_test.go b/pkg/resource-handler/controller/shard/integration_test.go index 730dd795..8b9cb34d 100644 --- a/pkg/resource-handler/controller/shard/integration_test.go +++ b/pkg/resource-handler/controller/shard/integration_test.go @@ -16,7 +16,6 @@ import ( "github.com/multigres/multigres/go/common/topoclient/memorytopo" cm "github.com/multigres/multigres/go/pb/clustermetadata" md "github.com/multigres/multigres/go/pb/multipoolermanagerdata" - "github.com/stretchr/testify/require" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" policyv1 "k8s.io/api/policy/v1" @@ -36,6 +35,8 @@ import ( "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestSetupWithManager(t *testing.T) { @@ -53,15 +54,13 @@ func TestSetupWithManager(t *testing.T) { ), ) - if err := (&shardcontroller.ShardReconciler{ + assert.NewAborting(t).NoError((&shardcontroller.ShardReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("shard-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") } func setTestPostgresPasswordSecretRef(shard *multigresv1alpha1.Shard) { @@ -90,9 +89,9 @@ func createTestPostgresPasswordSecret( "password": []byte("postgres"), }, } - if err := c.Create(ctx, secret); client.IgnoreAlreadyExists(err) != nil { - t.Fatalf("Failed to create postgres password Secret: %v", err) - } + err := c.Create(ctx, secret) + assert.NewAborting(t). + NoError(client.IgnoreAlreadyExists(err), "Failed to create postgres password Secret: %v", err) } func createTestPostgresInitSecretsSecret( @@ -108,12 +107,14 @@ func createTestPostgresInitSecretsSecret( Namespace: namespace, }, Data: map[string][]byte{ - key: []byte(`{"roles":{"app":"app-password"},"database_settings":{"testdb":{"work_mem":"64MB"}}}`), + key: []byte( + `{"roles":{"app":"app-password"},"database_settings":{"testdb":{"work_mem":"64MB"}}}`, + ), }, } - if err := c.Create(ctx, secret); client.IgnoreAlreadyExists(err) != nil { - t.Fatalf("Failed to create postgres init-secrets Secret: %v", err) - } + err := c.Create(ctx, secret) + assert.NewAborting(t). + NoError(client.IgnoreAlreadyExists(err), "Failed to create postgres init-secrets Secret: %v", err) } func TestShardReconciliation(t *testing.T) { @@ -172,8 +173,11 @@ func TestShardReconciliation(t *testing.T) { }, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: "/backups", Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: "/backups", + Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, + }, }, }, }, @@ -181,19 +185,36 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Deployment for zone-a &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-shard-multiorch-zone-a", - Namespace: "default", - Labels: shardLabels(t, "test-shard-multiorch-zone-a", "multiorch", "zone-a"), + Name: "test-shard-multiorch-zone-a", + Namespace: "default", + Labels: shardLabels( + t, + "test-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "test-shard"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(1)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(shardLabels(t, "test-shard-multiorch-zone-a", "multiorch", "zone-a")), + MatchLabels: metadata.GetSelectorLabels( + shardLabels( + t, + "test-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: shardLabels(t, "test-shard-multiorch-zone-a", "multiorch", "zone-a"), + Labels: shardLabels( + t, + "test-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -254,9 +275,14 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Service for zone-a &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-shard-multiorch-zone-a", - Namespace: "default", - Labels: shardLabels(t, "test-shard-multiorch-zone-a", "multiorch", "zone-a"), + Name: "test-shard-multiorch-zone-a", + Namespace: "default", + Labels: shardLabels( + t, + "test-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "test-shard"), }, Spec: corev1.ServiceSpec{ @@ -265,25 +291,44 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "http", 15300), tcpServicePort(t, "grpc", 15370), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "test-shard-multiorch-zone-a", "multiorch", "zone-a")), + Selector: metadata.GetSelectorLabels( + shardLabels(t, "test-shard-multiorch-zone-a", "multiorch", "zone-a"), + ), }, }, // Multiorch Deployment for zone-b &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-shard-multiorch-zone-b", - Namespace: "default", - Labels: shardLabels(t, "test-shard-multiorch-zone-b", "multiorch", "zone-b"), + Name: "test-shard-multiorch-zone-b", + Namespace: "default", + Labels: shardLabels( + t, + "test-shard-multiorch-zone-b", + "multiorch", + "zone-b", + ), OwnerReferences: shardOwnerRefs(t, "test-shard"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(1)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(shardLabels(t, "test-shard-multiorch-zone-b", "multiorch", "zone-b")), + MatchLabels: metadata.GetSelectorLabels( + shardLabels( + t, + "test-shard-multiorch-zone-b", + "multiorch", + "zone-b", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: shardLabels(t, "test-shard-multiorch-zone-b", "multiorch", "zone-b"), + Labels: shardLabels( + t, + "test-shard-multiorch-zone-b", + "multiorch", + "zone-b", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -344,9 +389,14 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Service for zone-b &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-shard-multiorch-zone-b", - Namespace: "default", - Labels: shardLabels(t, "test-shard-multiorch-zone-b", "multiorch", "zone-b"), + Name: "test-shard-multiorch-zone-b", + Namespace: "default", + Labels: shardLabels( + t, + "test-shard-multiorch-zone-b", + "multiorch", + "zone-b", + ), OwnerReferences: shardOwnerRefs(t, "test-shard"), }, Spec: corev1.ServiceSpec{ @@ -355,14 +405,21 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "http", 15300), tcpServicePort(t, "grpc", 15370), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "test-shard-multiorch-zone-b", "multiorch", "zone-b")), + Selector: metadata.GetSelectorLabels( + shardLabels(t, "test-shard-multiorch-zone-b", "multiorch", "zone-b"), + ), }, }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "test-shard-pool-primary-zone-a-headless", - Namespace: "default", - Labels: shardLabels(t, "test-shard-pool-primary-zone-a", "shard-pool", "zone-a"), + Name: "test-shard-pool-primary-zone-a-headless", + Namespace: "default", + Labels: shardLabels( + t, + "test-shard-pool-primary-zone-a", + "shard-pool", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "test-shard"), }, Spec: corev1.ServiceSpec{ @@ -374,7 +431,14 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "postgres", 5432), tcpServicePort(t, "metrics", 9187), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "test-shard-pool-primary-zone-a", "shard-pool", "zone-a")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "test-shard-pool-primary-zone-a", + "shard-pool", + "zone-a", + ), + ), PublishNotReadyAddresses: true, }, }, @@ -427,27 +491,47 @@ func TestShardReconciliation(t *testing.T) { }, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: "/backups", Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: "/backups", + Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, + }, }, }, }, wantResources: []client.Object{ &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "init-secrets-shard-multiorch-zone-a", - Namespace: "default", - Labels: shardLabels(t, "init-secrets-shard-multiorch-zone-a", "multiorch", "zone-a"), + Name: "init-secrets-shard-multiorch-zone-a", + Namespace: "default", + Labels: shardLabels( + t, + "init-secrets-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "init-secrets-shard"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(1)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(shardLabels(t, "init-secrets-shard-multiorch-zone-a", "multiorch", "zone-a")), + MatchLabels: metadata.GetSelectorLabels( + shardLabels( + t, + "init-secrets-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: shardLabels(t, "init-secrets-shard-multiorch-zone-a", "multiorch", "zone-a"), + Labels: shardLabels( + t, + "init-secrets-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -507,9 +591,14 @@ func TestShardReconciliation(t *testing.T) { }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "init-secrets-shard-multiorch-zone-a", - Namespace: "default", - Labels: shardLabels(t, "init-secrets-shard-multiorch-zone-a", "multiorch", "zone-a"), + Name: "init-secrets-shard-multiorch-zone-a", + Namespace: "default", + Labels: shardLabels( + t, + "init-secrets-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "init-secrets-shard"), }, Spec: corev1.ServiceSpec{ @@ -518,14 +607,26 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "http", 15300), tcpServicePort(t, "grpc", 15370), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "init-secrets-shard-multiorch-zone-a", "multiorch", "zone-a")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "init-secrets-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), + ), }, }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "init-secrets-shard-pool-primary-zone-a-headless", - Namespace: "default", - Labels: shardLabels(t, "init-secrets-shard-pool-primary-zone-a", "shard-pool", "zone-a"), + Name: "init-secrets-shard-pool-primary-zone-a-headless", + Namespace: "default", + Labels: shardLabels( + t, + "init-secrets-shard-pool-primary-zone-a", + "shard-pool", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "init-secrets-shard"), }, Spec: corev1.ServiceSpec{ @@ -537,7 +638,14 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "postgres", 5432), tcpServicePort(t, "metrics", 9187), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "init-secrets-shard-pool-primary-zone-a", "shard-pool", "zone-a")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "init-secrets-shard-pool-primary-zone-a", + "shard-pool", + "zone-a", + ), + ), PublishNotReadyAddresses: true, }, }, @@ -590,8 +698,11 @@ func TestShardReconciliation(t *testing.T) { }, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: "/backups", Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: "/backups", + Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, + }, }, }, }, @@ -599,19 +710,36 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Deployment for zone-a &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "delete-policy-shard-multiorch-zone-a", - Namespace: "default", - Labels: shardLabels(t, "delete-policy-shard-multiorch-zone-a", "multiorch", "zone-a"), + Name: "delete-policy-shard-multiorch-zone-a", + Namespace: "default", + Labels: shardLabels( + t, + "delete-policy-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "delete-policy-shard"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(1)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(shardLabels(t, "delete-policy-shard-multiorch-zone-a", "multiorch", "zone-a")), + MatchLabels: metadata.GetSelectorLabels( + shardLabels( + t, + "delete-policy-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: shardLabels(t, "delete-policy-shard-multiorch-zone-a", "multiorch", "zone-a"), + Labels: shardLabels( + t, + "delete-policy-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -672,9 +800,14 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Service for zone-a &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "delete-policy-shard-multiorch-zone-a", - Namespace: "default", - Labels: shardLabels(t, "delete-policy-shard-multiorch-zone-a", "multiorch", "zone-a"), + Name: "delete-policy-shard-multiorch-zone-a", + Namespace: "default", + Labels: shardLabels( + t, + "delete-policy-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "delete-policy-shard"), }, Spec: corev1.ServiceSpec{ @@ -683,14 +816,26 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "http", 15300), tcpServicePort(t, "grpc", 15370), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "delete-policy-shard-multiorch-zone-a", "multiorch", "zone-a")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "delete-policy-shard-multiorch-zone-a", + "multiorch", + "zone-a", + ), + ), }, }, &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "delete-policy-shard-pool-primary-zone-a-headless", - Namespace: "default", - Labels: shardLabels(t, "delete-policy-shard-pool-primary-zone-a", "shard-pool", "zone-a"), + Name: "delete-policy-shard-pool-primary-zone-a-headless", + Namespace: "default", + Labels: shardLabels( + t, + "delete-policy-shard-pool-primary-zone-a", + "shard-pool", + "zone-a", + ), OwnerReferences: shardOwnerRefs(t, "delete-policy-shard"), }, Spec: corev1.ServiceSpec{ @@ -702,7 +847,14 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "postgres", 5432), tcpServicePort(t, "metrics", 9187), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "delete-policy-shard-pool-primary-zone-a", "shard-pool", "zone-a")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "delete-policy-shard-pool-primary-zone-a", + "shard-pool", + "zone-a", + ), + ), PublishNotReadyAddresses: true, }, }, @@ -751,8 +903,11 @@ func TestShardReconciliation(t *testing.T) { }, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: "/backups", Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: "/backups", + Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, + }, }, }, }, @@ -760,19 +915,36 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Deployment for zone1 &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "multi-cell-shard-multiorch-zone1", - Namespace: "default", - Labels: shardLabels(t, "multi-cell-shard-multiorch-zone1", "multiorch", "zone1"), + Name: "multi-cell-shard-multiorch-zone1", + Namespace: "default", + Labels: shardLabels( + t, + "multi-cell-shard-multiorch-zone1", + "multiorch", + "zone1", + ), OwnerReferences: shardOwnerRefs(t, "multi-cell-shard"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(1)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(shardLabels(t, "multi-cell-shard-multiorch-zone1", "multiorch", "zone1")), + MatchLabels: metadata.GetSelectorLabels( + shardLabels( + t, + "multi-cell-shard-multiorch-zone1", + "multiorch", + "zone1", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: shardLabels(t, "multi-cell-shard-multiorch-zone1", "multiorch", "zone1"), + Labels: shardLabels( + t, + "multi-cell-shard-multiorch-zone1", + "multiorch", + "zone1", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -833,9 +1005,14 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Service for zone1 &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "multi-cell-shard-multiorch-zone1", - Namespace: "default", - Labels: shardLabels(t, "multi-cell-shard-multiorch-zone1", "multiorch", "zone1"), + Name: "multi-cell-shard-multiorch-zone1", + Namespace: "default", + Labels: shardLabels( + t, + "multi-cell-shard-multiorch-zone1", + "multiorch", + "zone1", + ), OwnerReferences: shardOwnerRefs(t, "multi-cell-shard"), }, Spec: corev1.ServiceSpec{ @@ -844,25 +1021,49 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "http", 15300), tcpServicePort(t, "grpc", 15370), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "multi-cell-shard-multiorch-zone1", "multiorch", "zone1")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "multi-cell-shard-multiorch-zone1", + "multiorch", + "zone1", + ), + ), }, }, // Multiorch Deployment for zone2 &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ - Name: "multi-cell-shard-multiorch-zone2", - Namespace: "default", - Labels: shardLabels(t, "multi-cell-shard-multiorch-zone2", "multiorch", "zone2"), + Name: "multi-cell-shard-multiorch-zone2", + Namespace: "default", + Labels: shardLabels( + t, + "multi-cell-shard-multiorch-zone2", + "multiorch", + "zone2", + ), OwnerReferences: shardOwnerRefs(t, "multi-cell-shard"), }, Spec: appsv1.DeploymentSpec{ Replicas: ptr.To(int32(1)), Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(shardLabels(t, "multi-cell-shard-multiorch-zone2", "multiorch", "zone2")), + MatchLabels: metadata.GetSelectorLabels( + shardLabels( + t, + "multi-cell-shard-multiorch-zone2", + "multiorch", + "zone2", + ), + ), }, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{ - Labels: shardLabels(t, "multi-cell-shard-multiorch-zone2", "multiorch", "zone2"), + Labels: shardLabels( + t, + "multi-cell-shard-multiorch-zone2", + "multiorch", + "zone2", + ), Annotations: map[string]string{ "multigres.com/project-ref": "test-cluster", }, @@ -923,9 +1124,14 @@ func TestShardReconciliation(t *testing.T) { // Multiorch Service for zone2 &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "multi-cell-shard-multiorch-zone2", - Namespace: "default", - Labels: shardLabels(t, "multi-cell-shard-multiorch-zone2", "multiorch", "zone2"), + Name: "multi-cell-shard-multiorch-zone2", + Namespace: "default", + Labels: shardLabels( + t, + "multi-cell-shard-multiorch-zone2", + "multiorch", + "zone2", + ), OwnerReferences: shardOwnerRefs(t, "multi-cell-shard"), }, Spec: corev1.ServiceSpec{ @@ -934,15 +1140,27 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "http", 15300), tcpServicePort(t, "grpc", 15370), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "multi-cell-shard-multiorch-zone2", "multiorch", "zone2")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "multi-cell-shard-multiorch-zone2", + "multiorch", + "zone2", + ), + ), }, }, // Headless Service for zone1 &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "multi-cell-shard-pool-primary-zone1-headless", - Namespace: "default", - Labels: shardLabels(t, "multi-cell-shard-pool-primary-zone1", "shard-pool", "zone1"), + Name: "multi-cell-shard-pool-primary-zone1-headless", + Namespace: "default", + Labels: shardLabels( + t, + "multi-cell-shard-pool-primary-zone1", + "shard-pool", + "zone1", + ), OwnerReferences: shardOwnerRefs(t, "multi-cell-shard"), }, Spec: corev1.ServiceSpec{ @@ -954,16 +1172,28 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "postgres", 5432), tcpServicePort(t, "metrics", 9187), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "multi-cell-shard-pool-primary-zone1", "shard-pool", "zone1")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "multi-cell-shard-pool-primary-zone1", + "shard-pool", + "zone1", + ), + ), PublishNotReadyAddresses: true, }, }, // Headless Service for zone2 &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ - Name: "multi-cell-shard-pool-primary-zone2-headless", - Namespace: "default", - Labels: shardLabels(t, "multi-cell-shard-pool-primary-zone2", "shard-pool", "zone2"), + Name: "multi-cell-shard-pool-primary-zone2-headless", + Namespace: "default", + Labels: shardLabels( + t, + "multi-cell-shard-pool-primary-zone2", + "shard-pool", + "zone2", + ), OwnerReferences: shardOwnerRefs(t, "multi-cell-shard"), }, Spec: corev1.ServiceSpec{ @@ -975,7 +1205,14 @@ func TestShardReconciliation(t *testing.T) { tcpServicePort(t, "postgres", 5432), tcpServicePort(t, "metrics", 9187), }, - Selector: metadata.GetSelectorLabels(shardLabels(t, "multi-cell-shard-pool-primary-zone2", "shard-pool", "zone2")), + Selector: metadata.GetSelectorLabels( + shardLabels( + t, + "multi-cell-shard-pool-primary-zone2", + "shard-pool", + "zone2", + ), + ), PublishNotReadyAddresses: true, }, }, @@ -986,6 +1223,7 @@ func TestShardReconciliation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) ctx := t.Context() mgr := testutil.SetUpEnvtestManager(t, scheme, testutil.WithCRDPaths( @@ -1023,7 +1261,11 @@ func TestShardReconciliation(t *testing.T) { return case <-ticker.C: podList := &corev1.PodList{} - if err := k8sClient.List(ctx, podList, client.InNamespace(tc.shard.Namespace)); err != nil { + if err := k8sClient.List( + ctx, + podList, + client.InNamespace(tc.shard.Namespace), + ); err != nil { continue } for i := range podList.Items { @@ -1053,12 +1295,10 @@ func TestShardReconciliation(t *testing.T) { Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("shard-controller"), } - if err := shardReconciler.SetupWithManager(mgr, controller.Options{ - // Needed for the parallel test runs + // Needed for the parallel test runs + ck.Require().NoError(shardReconciler.SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") setTestPostgresPasswordSecretRef(tc.shard) createTestPostgresPasswordSecret(t, ctx, k8sClient, tc.shard.Namespace) @@ -1067,11 +1307,17 @@ func TestShardReconciliation(t *testing.T) { if key == "" { key = shardcontroller.PostgresInitSecretsFileName } - createTestPostgresInitSecretsSecret(t, ctx, k8sClient, tc.shard.Namespace, ref.Name, key) - } - if err := k8sClient.Create(ctx, tc.shard); err != nil { - t.Fatalf("Failed to create the initial item, %v", err) + createTestPostgresInitSecretsSecret( + t, + ctx, + k8sClient, + tc.shard.Namespace, + ref.Name, + key, + ) } + ck.Require(). + NoError(k8sClient.Create(ctx, tc.shard), "Failed to create the initial item") // Patch wantResources with hashed names for _, obj := range tc.wantResources { @@ -1146,9 +1392,7 @@ func TestShardReconciliation(t *testing.T) { // ones the controller creates. These test shards have no // PostgresConfigRef, so the ref content is empty. _, hashes, err := shardcontroller.RenderPostgresConfig(tc.shard, "") - if err != nil { - t.Fatalf("Failed to render postgres config hash: %v", err) - } + ck.Require().NoError(err, "Failed to render postgres config hash") if tc.shard.Annotations == nil { tc.shard.Annotations = map[string]string{} } @@ -1164,34 +1408,47 @@ func TestShardReconciliation(t *testing.T) { replicas = *poolSpec.ReplicasPerCell } for i := 0; i < int(replicas); i++ { - pod, err := shardcontroller.BuildPoolPod(tc.shard, string(poolName), string(cellName), poolSpec, i, mgr.GetScheme()) - if err != nil { - t.Fatalf("Failed to build pod: %v", err) - } + pod, err := shardcontroller.BuildPoolPod( + tc.shard, + string(poolName), + string(cellName), + poolSpec, + i, + mgr.GetScheme(), + ) + ck.Require().NoError(err, "Failed to build pod") filteredResources = append(filteredResources, pod) - pvc, err := shardcontroller.BuildPoolDataPVC(tc.shard, string(poolName), string(cellName), poolSpec, i, shardcontroller.ShouldDeletePVCOnShardRemoval(tc.shard, poolSpec), mgr.GetScheme()) - if err != nil { - t.Fatalf("Failed to build pvc: %v", err) - } + pvc, err := shardcontroller.BuildPoolDataPVC( + tc.shard, + string(poolName), + string(cellName), + poolSpec, + i, + shardcontroller.ShouldDeletePVCOnShardRemoval(tc.shard, poolSpec), + mgr.GetScheme(), + ) + ck.Require().NoError(err, "Failed to build pvc") filteredResources = append(filteredResources, pvc) } // Shared backup PVC is per-shard, not per-pod or per-cell. - if tc.shard.Spec.Backup != nil && tc.shard.Spec.Backup.Type == multigresv1alpha1.BackupTypeFilesystem && !backupPVCAdded { - backupPVC, err := shardcontroller.BuildSharedBackupPVC(tc.shard, shardcontroller.ShouldDeleteShardLevelPVCOnRemoval(tc.shard), mgr.GetScheme()) - if err != nil { - t.Fatalf("Failed to build backup pvc: %v", err) - } + if tc.shard.Spec.Backup != nil && + tc.shard.Spec.Backup.Type == multigresv1alpha1.BackupTypeFilesystem && + !backupPVCAdded { + backupPVC, err := shardcontroller.BuildSharedBackupPVC( + tc.shard, + shardcontroller.ShouldDeleteShardLevelPVCOnRemoval(tc.shard), + mgr.GetScheme(), + ) + ck.Require().NoError(err, "Failed to build backup pvc") filteredResources = append(filteredResources, backupPVC) backupPVCAdded = true } } } - if err := watcher.WaitForMatch(filteredResources...); err != nil { - t.Errorf("Resources mismatch:\n%v", err) - } + ck.NoError(watcher.WaitForMatch(filteredResources...), "Resources mismatch:\n") }) } } @@ -1249,7 +1506,12 @@ func tcpPort(t testing.TB, name string, port int32) corev1.ContainerPort { // tcpServicePort creates a TCP service port with named target func tcpServicePort(t testing.TB, name string, port int32) corev1.ServicePort { t.Helper() - return corev1.ServicePort{Name: name, Port: port, TargetPort: intstr.FromString(name), Protocol: corev1.ProtocolTCP} + return corev1.ServicePort{ + Name: name, + Port: port, + TargetPort: intstr.FromString(name), + Protocol: corev1.ProtocolTCP, + } } // multipoolerPorts returns the standard multipooler container ports @@ -1264,6 +1526,7 @@ func multipoolerPorts(t testing.TB) []corev1.ContainerPort { func TestReconcileDeletions(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1278,15 +1541,13 @@ func TestReconcileDeletions(t *testing.T) { ) // Setup controller with manager - if err := (&shardcontroller.ShardReconciler{ + c.NoError((&shardcontroller.ShardReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("shard-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") ctx := t.Context() k8sClient := mgr.GetClient() @@ -1333,17 +1594,18 @@ func TestReconcileDeletions(t *testing.T) { }, }, Backup: &multigresv1alpha1.BackupConfig{ - Type: multigresv1alpha1.BackupTypeFilesystem, - Filesystem: &multigresv1alpha1.FilesystemBackupConfig{Path: "/backups", Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}}, + Type: multigresv1alpha1.BackupTypeFilesystem, + Filesystem: &multigresv1alpha1.FilesystemBackupConfig{ + Path: "/backups", + Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, + }, }, }, } setTestPostgresPasswordSecretRef(shard) createTestPostgresPasswordSecret(t, ctx, k8sClient, shard.Namespace) - if err := k8sClient.Create(ctx, shard); err != nil { - t.Fatalf("Failed to create Shard: %v", err) - } + c.NoError(k8sClient.Create(ctx, shard), "Failed to create Shard") cm := &corev1.ConfigMap{ ObjectMeta: metav1.ObjectMeta{ @@ -1356,24 +1618,27 @@ func TestReconcileDeletions(t *testing.T) { // We use polling to avoid strict content matching on Data pollFound := false for i := 0; i < 20; i++ { - err := k8sClient.Get(ctx, types.NamespacedName{Name: shardcontroller.PgHbaConfigMapName("test-shard-deletion-reconcile"), Namespace: "default"}, cm) + err := k8sClient.Get( + ctx, + types.NamespacedName{ + Name: shardcontroller.PgHbaConfigMapName("test-shard-deletion-reconcile"), + Namespace: "default", + }, + cm, + ) if err == nil { pollFound = true break } time.Sleep(500 * time.Millisecond) } - if !pollFound { - t.Fatalf("ConfigMap not initially created") - } + c.True(pollFound, "ConfigMap not initially created") // 2. Delete ConfigMap // We need to fetch it first to get UID/ResourceVersion for proper deletion if needed, // strictly speaking not needed for k8s deletion by name if we construct it, // but better to be safe with client usage. - if err := k8sClient.Delete(ctx, cm); err != nil { - t.Fatalf("Failed to delete ConfigMap: %v", err) - } + c.NoError(k8sClient.Delete(ctx, cm), "Failed to delete ConfigMap") // 3. Wait for ConfigMap to be recreated // Since the controller watches ConfigMaps, the deletion event should trigger Reconcile. @@ -1391,7 +1656,14 @@ func TestReconcileDeletions(t *testing.T) { default: } - err := k8sClient.Get(ctx, types.NamespacedName{Name: shardcontroller.PgHbaConfigMapName("test-shard-deletion-reconcile"), Namespace: "default"}, cm) + err := k8sClient.Get( + ctx, + types.NamespacedName{ + Name: shardcontroller.PgHbaConfigMapName("test-shard-deletion-reconcile"), + Namespace: "default", + }, + cm, + ) if err == nil { found = true break @@ -1399,13 +1671,12 @@ func TestReconcileDeletions(t *testing.T) { time.Sleep(interval) } - if !found { - t.Fatalf("ConfigMap was not recreated") - } + c.True(found, "ConfigMap was not recreated") } func TestShardReconciliation_DanglingPostgresInitSecretsRef(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1419,15 +1690,13 @@ func TestShardReconciliation_DanglingPostgresInitSecretsRef(t *testing.T) { ), ) - if err := (&shardcontroller.ShardReconciler{ + c.Require().NoError((&shardcontroller.ShardReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("shard-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") ctx := t.Context() k8sClient := mgr.GetClient() @@ -1481,11 +1750,9 @@ func TestShardReconciliation_DanglingPostgresInitSecretsRef(t *testing.T) { setTestPostgresPasswordSecretRef(shard) createTestPostgresPasswordSecret(t, ctx, k8sClient, shard.Namespace) - if err := k8sClient.Create(ctx, shard); err != nil { - t.Fatalf("Failed to create Shard: %v", err) - } + c.Require().NoError(k8sClient.Create(ctx, shard), "Failed to create Shard") - require.Eventually(t, func() bool { + c.Require().EventuallyTrue(10*time.Second, 100*time.Millisecond, func() bool { events := &corev1.EventList{} if err := k8sClient.List(ctx, events, client.InNamespace(shard.Namespace)); err != nil { return false @@ -1498,25 +1765,21 @@ func TestShardReconciliation_DanglingPostgresInitSecretsRef(t *testing.T) { } } return false - }, 10*time.Second, 100*time.Millisecond, - "expected a ConfigError event for the dangling init-secrets reference") + }, "expected a ConfigError event for the dangling init-secrets reference") podList := &corev1.PodList{} - if err := k8sClient.List(ctx, podList, + c.Require().NoError(k8sClient.List(ctx, podList, client.InNamespace(shard.Namespace), client.MatchingLabels{ "multigres.com/cluster": "test-cluster", "multigres.com/shard": "0", }, - ); err != nil { - t.Fatalf("Failed to list pods: %v", err) - } - if len(podList.Items) != 0 { - t.Errorf( - "expected no pool pods for shard with dangling PostgresInitSecretsRef, got %d", - len(podList.Items), - ) - } + ), "Failed to list pods") + c.Empty( + podList.Items, + "expected no pool pods for shard with dangling PostgresInitSecretsRef, got %d", + len(podList.Items), + ) } // TestReloadVsRestartRollout drives a real reconcile against envtest and proves @@ -1554,6 +1817,7 @@ func testReloadVsRestartRollout( blockedReason string, ) { t.Helper() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -1627,22 +1891,18 @@ func testReloadVsRestartRollout( blocked: make(chan string, 1), started: make(chan string, 1), } - if err := (&shardcontroller.ShardReconciler{ + c.NoError((&shardcontroller.ShardReconciler{ Client: mgr.GetClient(), APIReader: mgr.GetAPIReader(), Scheme: mgr.GetScheme(), Recorder: recorder, PoolerClients: poolerclient.Static(rpc), CreateTopoStore: topoFactory, - }).SetupWithManager(mgr, controller.Options{SkipNameValidation: ptr.To(true)}); err != nil { - t.Fatalf("Failed to create controller: %v", err) - } + }).SetupWithManager(mgr, controller.Options{SkipNameValidation: ptr.To(true)}), "Failed to create controller") setTestPostgresPasswordSecretRef(shard) createTestPostgresPasswordSecret(t, ctx, k8sClient, shard.Namespace) - if err := k8sClient.Create(ctx, shard); err != nil { - t.Fatalf("Failed to create Shard: %v", err) - } + c.NoError(k8sClient.Create(ctx, shard), "Failed to create Shard") // Keep pool pods Ready in the background so the controller does not block. go markPoolPodsReady(ctx, k8sClient, clusterName) @@ -1654,7 +1914,7 @@ func testReloadVsRestartRollout( // Wait for every desired member before testing either configuration change. var original corev1.PodList - require.Eventually(t, func() bool { + c.EventuallyTrue(30*time.Second, 200*time.Millisecond, func() bool { if err := mgr.GetClient().List( ctx, &original, @@ -1670,7 +1930,7 @@ func testReloadVsRestartRollout( } } return true - }, 30*time.Second, 200*time.Millisecond, "wait for the complete ready cohort") + }, "wait for the complete ready cohort") // Multipoolers register only after their pods exist. Registering before // the controller cache sees them lets dead-pooler cleanup mark them shut down. registerPoolers() @@ -1681,29 +1941,23 @@ func testReloadVsRestartRollout( waitForConfigMapContains(t, ctx, k8sClient, shardName, "work_mem = '8MB'") // Reaching the reload RPC proves the controller processed the pool rollout // decision, not merely the earlier ConfigMap write. - require.Eventually(t, func() bool { + c.EventuallyTrue(30*time.Second, 100*time.Millisecond, func() bool { for _, call := range rpc.GetCallLog() { if strings.HasPrefix(call, "ReloadConfig") { return true } } return false - }, 30*time.Second, 100*time.Millisecond, "reload-only change did not reach ReloadConfig") + }, "reload-only change did not reach ReloadConfig") assertUnchanged := func() { t.Helper() for _, orig := range original.Items { var got corev1.Pod - require.NoError(t, mgr.GetAPIReader().Get(ctx, client.ObjectKeyFromObject(&orig), &got)) - require.Equal(t, orig.UID, got.UID, "pod %s was recreated", orig.Name) - require.Empty( - t, - got.Annotations[metadata.AnnotationDrainState], - "pod %s was drained", - orig.Name, - ) - require.Equal( - t, + c.NoError(mgr.GetAPIReader().Get(ctx, client.ObjectKeyFromObject(&orig), &got)) + c.EqDeep(orig.UID, got.UID, "pod %s was recreated", orig.Name) + c.Empty(got.Annotations[metadata.AnnotationDrainState], "pod %s was drained", orig.Name) + c.EqDeep( orig.Annotations[metadata.AnnotationSpecHash], got.Annotations[metadata.AnnotationSpecHash], ) @@ -1716,39 +1970,36 @@ func testReloadVsRestartRollout( if blockedReason != "" { // Observe an actual failed preflight, not merely the absence of a drain // before the controller processes the config update. - require.Eventually(t, func() bool { + c.EventuallyTrue(40*time.Second, 200*time.Millisecond, func() bool { select { case message := <-recorder.blocked: return strings.Contains(message, blockedReason) default: return false } - }, 40*time.Second, 200*time.Millisecond, "did not observe a blocked disruption: %s", blockedReason) + }, "did not observe a blocked disruption: %s", blockedReason) assertUnchanged() return } // A drain can advance through all annotations between polling intervals in // envtest. Observe its initiation directly instead of racing pod deletion. - require.Eventually(t, func() bool { + c.EventuallyTrue(40*time.Second, 100*time.Millisecond, func() bool { select { case message := <-recorder.started: return strings.Contains(message, "Initiated drain for drifted replica pod") default: return false } - }, 40*time.Second, 100*time.Millisecond, "restart change did not initiate a replica drain") + }, "restart change did not initiate a replica drain") var pods corev1.PodList - require.NoError( - t, - k8sClient.List(ctx, &pods, client.InNamespace(shard.Namespace), poolSelector), - ) + c.NoError(k8sClient.List(ctx, &pods, client.InNamespace(shard.Namespace), poolSelector)) draining := 0 for _, pod := range pods.Items { if pod.Annotations[metadata.AnnotationDrainState] != "" { draining++ } } - require.LessOrEqual(t, draining, 1, "restart initiated overlapping drains") + c.LessOrEqual(1, draining, "restart initiated overlapping drains") } // Observe preflight events without depending on the API event broadcaster's @@ -1789,8 +2040,9 @@ func rolloutDataPlane( missingStatus bool, ) (*rpcclient.FakeClient, func(*multigresv1alpha1.Shard) (topoclient.Store, error), func()) { t.Helper() + c := assert.NewAborting(t) store, factory := memorytopo.NewServerAndFactory(t.Context(), "zone1") - t.Cleanup(func() { require.NoError(t, store.Close()) }) + t.Cleanup(func() { c.NoError(store.Close()) }) ids := make([]*cm.ID, replicas) for i := range ids { ids[i] = &cm.ID{ @@ -1858,7 +2110,7 @@ func rolloutDataPlane( ), nil }, func() { for _, pooler := range poolers { - require.NoError(t, store.RegisterMultipooler(t.Context(), pooler, false)) + c.NoError(store.RegisterMultipooler(t.Context(), pooler, false)) } } } @@ -1922,7 +2174,7 @@ func setInlineConfig( shardName, key, val string, ) { t.Helper() - if err := retry.RetryOnConflict(retry.DefaultRetry, func() error { + assert.NewAborting(t).NoError(retry.RetryOnConflict(retry.DefaultRetry, func() error { s := &multigresv1alpha1.Shard{} if err := c.Get( ctx, @@ -1937,9 +2189,7 @@ func setInlineConfig( } s.Spec.PostgresConfig[key] = val return c.Patch(ctx, s, client.MergeFrom(base)) - }); err != nil { - t.Fatalf("update shard inline config %s=%s: %v", key, val, err) - } + }), "update shard inline config %s=%s", key, val) } func waitForConfigMapContains( diff --git a/pkg/resource-handler/controller/shard/maintenance_surge_test.go b/pkg/resource-handler/controller/shard/maintenance_surge_test.go index f04fd51e..f84a5a9b 100644 --- a/pkg/resource-handler/controller/shard/maintenance_surge_test.go +++ b/pkg/resource-handler/controller/shard/maintenance_surge_test.go @@ -15,19 +15,20 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestMaintenanceSurgeLifecycleForRollingUpdate(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := maintenanceSurgeTestScheme(t) shard := maintenanceSurgeTestShard() poolName := "primary" cellName := "zone-a" pool := shard.Spec.Pools[multigresv1alpha1.PoolName(poolName)] target, err := BuildPoolPod(shard, poolName, cellName, pool, 0, scheme) - if err != nil { - t.Fatalf("build target pod: %v", err) - } + ck.NoError(err, "build target pod") target.Annotations[metadata.AnnotationSpecHash] = "stale" setReady(target, true) @@ -53,84 +54,76 @@ func TestMaintenanceSurgeLifecycleForRollingUpdate(t *testing.T) { 1, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("create maintenance surge: %v", err) - } - if !acted || active != 0 { - t.Fatalf("create result = active %d, acted %v; want 0, true", active, acted) - } + ck.NoError(err, "create maintenance surge") + ck.False( + !acted || active != 0, + "create result = active %d, acted %v; want 0, true", + active, + acted, + ) surgeName := BuildPoolPodName(shard, poolName, cellName, 1) surge := &corev1.Pod{} - if err := c.Get( + ck.NoError(c.Get( t.Context(), types.NamespacedName{Name: surgeName, Namespace: shard.Namespace}, surge, - ); err != nil { - t.Fatalf("get maintenance surge: %v", err) - } + ), "get maintenance surge") if !isMaintenanceSurge(surge) { t.Fatalf("pod %s is missing the maintenance surge annotation", surge.Name) } setReady(surge, true) - if err := c.Status().Update(t.Context(), surge); err != nil { - t.Fatalf("mark maintenance surge ready: %v", err) - } + ck.NoError(c.Status().Update(t.Context(), surge), "mark maintenance surge ready") localPods, localPVCs := getLocalPoolObjects(t, c, shard, poolName, cellName) active, acted, err = r.reconcileCellMaintenanceSurge( t.Context(), shard, poolName, cellName, pool, localPods, localPVCs, 1, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("retain maintenance surge: %v", err) - } - if acted || active != 1 { - t.Fatalf("retain result = active %d, acted %v; want 1, false", active, acted) - } + ck.NoError(err, "retain maintenance surge") + ck.False( + acted || active != 1, + "retain result = active %d, acted %v; want 1, false", + active, + acted, + ) target = localPods[target.Name] desiredTarget, err := BuildPoolPod(shard, poolName, cellName, pool, 0, scheme) - if err != nil { - t.Fatalf("build desired target: %v", err) - } + ck.NoError(err, "build desired target") base := target.DeepCopy() desiredHash := desiredTarget.Annotations[metadata.AnnotationSpecHash] target.Annotations[metadata.AnnotationSpecHash] = desiredHash - if err := c.Patch(t.Context(), target, client.MergeFrom(base)); err != nil { - t.Fatalf("mark target current: %v", err) - } + ck.NoError(c.Patch(t.Context(), target, client.MergeFrom(base)), "mark target current") localPods, localPVCs = getLocalPoolObjects(t, c, shard, poolName, cellName) active, acted, err = r.reconcileCellMaintenanceSurge( t.Context(), shard, poolName, cellName, pool, localPods, localPVCs, 1, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("release maintenance surge: %v", err) - } - if acted || active != 0 { - t.Fatalf("release result = active %d, acted %v; want 0, false", active, acted) - } + ck.NoError(err, "release maintenance surge") + ck.False( + acted || active != 0, + "release result = active %d, acted %v; want 0, false", + active, + acted, + ) } func TestExplicitMaintenanceRequestWaitsForSurge(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := maintenanceSurgeTestScheme(t) shard := maintenanceSurgeTestShard() poolName := "primary" cellName := "zone-a" pool := shard.Spec.Pools[multigresv1alpha1.PoolName(poolName)] target, err := BuildPoolPod(shard, poolName, cellName, pool, 0, scheme) - if err != nil { - t.Fatalf("build target pod: %v", err) - } + ck.NoError(err, "build target pod") target.Annotations[metadata.AnnotationMaintenanceRequested] = maintenanceAnnotationTrue setReady(target, true) peer, err := BuildPoolPod(shard, poolName, "zone-b", pool, 0, scheme) - if err != nil { - t.Fatalf("build peer pod: %v", err) - } + ck.NoError(err, "build peer pod") setReady(peer, true) c := fake.NewClientBuilder(). @@ -155,51 +148,55 @@ func TestExplicitMaintenanceRequestWaitsForSurge(t *testing.T) { 1, &shardRolloutTracker{}, ) - if err != nil || !acted { - t.Fatalf("create explicit maintenance surge: acted %v, err %v", acted, err) - } + ck.False( + err != nil || !acted, + "create explicit maintenance surge: acted %v, err %v", + acted, + err, + ) updatedTarget := &corev1.Pod{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(target), updatedTarget); err != nil { - t.Fatalf("get target before surge readiness: %v", err) - } - if updatedTarget.Annotations[metadata.AnnotationMaintenanceReady] != "" { - t.Fatal("maintenance request became ready before the surge was ready") - } + ck.NoError( + c.Get(t.Context(), client.ObjectKeyFromObject(target), updatedTarget), + "get target before surge readiness", + ) + ck.Eq( + "", + updatedTarget.Annotations[metadata.AnnotationMaintenanceReady], + "maintenance request became ready before the surge was ready", + ) surge := &corev1.Pod{} - if err := c.Get( + ck.NoError(c.Get( t.Context(), types.NamespacedName{ Name: BuildPoolPodName(shard, poolName, cellName, 1), Namespace: shard.Namespace, }, surge, - ); err != nil { - t.Fatalf("get explicit maintenance surge: %v", err) - } + ), "get explicit maintenance surge") setReady(surge, true) - if err := c.Status().Update(t.Context(), surge); err != nil { - t.Fatalf("mark explicit maintenance surge ready: %v", err) - } + ck.NoError(c.Status().Update(t.Context(), surge), "mark explicit maintenance surge ready") localPods, localPVCs := getLocalPoolObjects(t, c, shard, poolName, cellName) _, acted, err = r.reconcileCellMaintenanceSurge( t.Context(), shard, poolName, cellName, pool, localPods, localPVCs, 1, &shardRolloutTracker{}, ) - if err != nil || !acted { - t.Fatalf("publish maintenance readiness: acted %v, err %v", acted, err) - } - if err := c.Get(t.Context(), client.ObjectKeyFromObject(target), updatedTarget); err != nil { - t.Fatalf("get maintenance-ready target: %v", err) - } - if updatedTarget.Annotations[metadata.AnnotationMaintenanceReady] != maintenanceAnnotationTrue { - t.Fatal("maintenance readiness was not published after the surge became ready") - } + ck.False(err != nil || !acted, "publish maintenance readiness: acted %v, err %v", acted, err) + ck.NoError( + c.Get(t.Context(), client.ObjectKeyFromObject(target), updatedTarget), + "get maintenance-ready target", + ) + ck.Eq( + maintenanceAnnotationTrue, + updatedTarget.Annotations[metadata.AnnotationMaintenanceReady], + "maintenance readiness was not published after the surge became ready", + ) } func TestScaleUpPromotesMaintenanceSurgesToDesiredCapacity(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := maintenanceSurgeTestScheme(t) shard := maintenanceSurgeTestShard() poolName := "primary" @@ -212,9 +209,7 @@ func TestScaleUpPromotesMaintenanceSurgesToDesiredCapacity(t *testing.T) { for _, cellName := range []string{"zone-a", "zone-b"} { for index := 0; index < 2; index++ { pod, err := BuildPoolPod(shard, poolName, cellName, pool, index, scheme) - if err != nil { - t.Fatalf("build pooler %s/%d: %v", cellName, index, err) - } + ck.NoError(err, "build pooler %s/%d", cellName, index) if index == 1 { pod.Annotations[metadata.AnnotationMaintenanceSurge] = maintenanceAnnotationTrue } @@ -236,19 +231,11 @@ func TestScaleUpPromotesMaintenanceSurgesToDesiredCapacity(t *testing.T) { // The PDB must treat deterministic indices 0 and 1 as the four desired // replicas immediately, even before stale surge annotations are cleaned up. - if err := r.reconcileShardPDB(t.Context(), shard); err != nil { - t.Fatalf("reconcile shard PDB: %v", err) - } + ck.NoError(r.reconcileShardPDB(t.Context(), shard), "reconcile shard PDB") pdb, err := BuildShardPodDisruptionBudget(shard, scheme) - if err != nil { - t.Fatalf("build shard PDB: %v", err) - } - if err := c.Get(t.Context(), client.ObjectKeyFromObject(pdb), pdb); err != nil { - t.Fatalf("get shard PDB: %v", err) - } - if got := pdb.Spec.MinAvailable.IntValue(); got != 3 { - t.Fatalf("minAvailable after scale-up = %d, want 3", got) - } + ck.NoError(err, "build shard PDB") + ck.NoError(c.Get(t.Context(), client.ObjectKeyFromObject(pdb), pdb), "get shard PDB") + ck.Eq(3, pdb.Spec.MinAvailable.IntValue(), "minAvailable after scale-up") localPods, localPVCs := getLocalPoolObjects(t, c, shard, poolName, "zone-a") active, acted, err := r.reconcileCellMaintenanceSurge( @@ -262,23 +249,23 @@ func TestScaleUpPromotesMaintenanceSurgesToDesiredCapacity(t *testing.T) { 2, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("promote surge after scale-up: %v", err) - } - if !acted || active != 0 { - t.Fatalf("promotion result = active %d, acted %v; want 0, true", active, acted) - } + ck.NoError(err, "promote surge after scale-up") + ck.False( + !acted || active != 0, + "promotion result = active %d, acted %v; want 0, true", + active, + acted, + ) promoted := &corev1.Pod{} key := types.NamespacedName{ Name: BuildPoolPodName(shard, poolName, "zone-a", 1), Namespace: shard.Namespace, } - if err := c.Get(t.Context(), key, promoted); err != nil { - t.Fatalf("get promoted pooler: %v", err) - } - if isMaintenanceSurge(promoted) { - t.Fatal("desired pooler retained the maintenance surge annotation") - } + ck.NoError(c.Get(t.Context(), key, promoted), "get promoted pooler") + ck.False( + isMaintenanceSurge(promoted), + "desired pooler retained the maintenance surge annotation", + ) } func maintenanceSurgeTestScheme(t *testing.T) *runtime.Scheme { @@ -289,9 +276,7 @@ func maintenanceSurgeTestScheme(t *testing.T) *runtime.Scheme { "policy": policyv1.AddToScheme, "multigres": multigresv1alpha1.AddToScheme, } { - if err := add(scheme); err != nil { - t.Fatalf("add %s scheme: %v", name, err) - } + assert.NewAborting(t).NoError(add(scheme), "add %s scheme", name) } return scheme } @@ -331,16 +316,19 @@ func getLocalPoolObjects( cellName string, ) (map[string]*corev1.Pod, map[string]*corev1.PersistentVolumeClaim) { t.Helper() + ck := assert.NewAborting(t) labels := buildPoolLabelsWithCell(shard, poolName, cellName) selector := client.MatchingLabels(metadata.GetSelectorLabels(labels)) pods := &corev1.PodList{} - if err := c.List(t.Context(), pods, client.InNamespace(shard.Namespace), selector); err != nil { - t.Fatalf("list local pods: %v", err) - } + ck.NoError( + c.List(t.Context(), pods, client.InNamespace(shard.Namespace), selector), + "list local pods", + ) pvcs := &corev1.PersistentVolumeClaimList{} - if err := c.List(t.Context(), pvcs, client.InNamespace(shard.Namespace), selector); err != nil { - t.Fatalf("list local PVCs: %v", err) - } + ck.NoError( + c.List(t.Context(), pvcs, client.InNamespace(shard.Namespace), selector), + "list local PVCs", + ) podsByName := make(map[string]*corev1.Pod, len(pods.Items)) for i := range pods.Items { podsByName[pods.Items[i].Name] = &pods.Items[i] diff --git a/pkg/resource-handler/controller/shard/multiorch_test.go b/pkg/resource-handler/controller/shard/multiorch_test.go index b1772502..1af120e6 100644 --- a/pkg/resource-handler/controller/shard/multiorch_test.go +++ b/pkg/resource-handler/controller/shard/multiorch_test.go @@ -3,7 +3,6 @@ package shard import ( "testing" - "github.com/google/go-cmp/cmp" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -14,6 +13,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" nameutil "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) func TestBuildMultiorchDeployment(t *testing.T) { @@ -626,9 +627,7 @@ func TestBuildMultiorchDeployment(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildMultiorchDeployment() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildMultiorchDeployment() mismatch") }) } } @@ -654,6 +653,7 @@ func TestBuildMultiorchDeployment_ProjectRefAnnotation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -678,9 +678,7 @@ func TestBuildMultiorchDeployment_ProjectRefAnnotation(t *testing.T) { } deploy, err := BuildMultiorchDeployment(shard, "zone-a", scheme) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") if got := deploy.Spec.Template.Annotations[metadata.AnnotationProjectRef]; got != tc.want { t.Fatalf("annotation %q = %q, want %q", metadata.AnnotationProjectRef, got, tc.want) @@ -692,15 +690,15 @@ func TestBuildMultiorchDeployment_ProjectRefAnnotation(t *testing.T) { metadata.LabelAppManagedBy: metadata.ManagedByMultigres, } for key, want := range assertedLabels { - if got := deploy.Spec.Template.Labels[key]; got != want { - t.Fatalf("label %q = %q, want %q", key, got, want) - } + got := deploy.Spec.Template.Labels[key] + c.Eq(want, got, "label %q = %q, want", key, got) } }) } } func TestBuildMultiorchDeployment_OmitsPrometheusScrapeAnnotations(t *testing.T) { + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -735,9 +733,7 @@ func TestBuildMultiorchDeployment_OmitsPrometheusScrapeAnnotations(t *testing.T) } deploy, err := BuildMultiorchDeployment(shard, "zone-a", scheme) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") if _, ok := deploy.Spec.Template.Annotations[metadata.AnnotationPrometheusScrape]; ok { t.Fatalf("annotation %q should be omitted", metadata.AnnotationPrometheusScrape) @@ -748,12 +744,11 @@ func TestBuildMultiorchDeployment_OmitsPrometheusScrapeAnnotations(t *testing.T) if _, ok := deploy.Spec.Template.Annotations[metadata.AnnotationPrometheusPath]; ok { t.Fatalf("annotation %q should be omitted", metadata.AnnotationPrometheusPath) } - if got := deploy.Spec.Template.Annotations["custom-annotation"]; got != "keep-me" { - t.Fatalf("custom annotation = %q, want %q", got, "keep-me") - } + c.Eq("keep-me", deploy.Spec.Template.Annotations["custom-annotation"], "custom annotation") } func TestBuildMultiorchDeployment_ShardTLSVolume(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -777,9 +772,7 @@ func TestBuildMultiorchDeployment_ShardTLSVolume(t *testing.T) { } got, err := BuildMultiorchDeployment(shard, "zone-a", scheme) - if err != nil { - t.Fatalf("BuildMultiorchDeployment() error = %v", err) - } + c.Require().NoError(err, "BuildMultiorchDeployment() error =") found := false for _, v := range got.Spec.Template.Spec.Volumes { @@ -787,17 +780,9 @@ func TestBuildMultiorchDeployment_ShardTLSVolume(t *testing.T) { continue } found = true - if v.Secret == nil { - t.Fatal("shard TLS volume should use Secret source") - } + c.Require().NotNil(v.Secret, "shard TLS volume should use Secret source") wantSecretName := "multiorch.test-cluster.default.multigres.internal" //nolint:gosec // test constant - if v.Secret.SecretName != wantSecretName { - t.Errorf( - "shard TLS secret = %q, want internal secret %q", - v.Secret.SecretName, - wantSecretName, - ) - } + c.Eq(wantSecretName, v.Secret.SecretName, "shard TLS secret") if v.Secret.DefaultMode == nil || *v.Secret.DefaultMode != 0o444 { t.Errorf( "shard TLS secret defaultMode = %v, want 0444", @@ -805,9 +790,7 @@ func TestBuildMultiorchDeployment_ShardTLSVolume(t *testing.T) { ) } } - if !found { - t.Error("expected internal multiorch TLS volume when internal TLS is enabled") - } + c.True(found, "expected internal multiorch TLS volume when internal TLS is enabled") for _, tc := range []struct { name string @@ -820,20 +803,19 @@ func TestBuildMultiorchDeployment_ShardTLSVolume(t *testing.T) { }, } { t.Run("absent when "+tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) disabledShard := shard.DeepCopy() disabledShard.Spec.InternalTLS = tc.internalTLS disabled, err := BuildMultiorchDeployment(disabledShard, "zone-a", scheme) - if err != nil { - t.Fatalf("BuildMultiorchDeployment() error = %v", err) - } + c.Require().NoError(err, "BuildMultiorchDeployment() error =") for _, volume := range disabled.Spec.Template.Spec.Volumes { - if volume.Name == ShardTLSVolumeName { - t.Errorf( - "internal TLS volume should be absent for config %+v", - tc.internalTLS, - ) - } + c.NotEq( + ShardTLSVolumeName, + volume.Name, + "internal TLS volume should be absent for config %+v", + tc.internalTLS, + ) } }) } @@ -1028,9 +1010,7 @@ func TestBuildMultiorchService(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildMultiorchService() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildMultiorchService() mismatch") }) } } diff --git a/pkg/resource-handler/controller/shard/pool_pod_test.go b/pkg/resource-handler/controller/shard/pool_pod_test.go index 48e13468..81a5028d 100644 --- a/pkg/resource-handler/controller/shard/pool_pod_test.go +++ b/pkg/resource-handler/controller/shard/pool_pod_test.go @@ -7,8 +7,6 @@ import ( "strings" "testing" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" @@ -16,6 +14,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func newTestShard() *multigresv1alpha1.Shard { @@ -54,28 +54,20 @@ func testScheme() *runtime.Scheme { } func TestBuildPoolPod_BasicStructure(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() pool := newTestPoolSpec() pod, err := BuildPoolPod(shard, "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if pod.Namespace != "default" { - t.Errorf("namespace = %q, want %q", pod.Namespace, "default") - } + c.Eq("default", pod.Namespace, "namespace") // Verify owner reference - if len(pod.OwnerReferences) != 1 { - t.Fatalf("expected 1 owner reference, got %d", len(pod.OwnerReferences)) - } - if pod.OwnerReferences[0].Name != "test-shard" { - t.Errorf("owner name = %q, want %q", pod.OwnerReferences[0].Name, "test-shard") - } - if pod.OwnerReferences[0].Kind != "Shard" { - t.Errorf("owner kind = %q, want %q", pod.OwnerReferences[0].Kind, "Shard") - } + c.Require(). + Len(pod.OwnerReferences, 1, "expected 1 owner reference, got %d", len(pod.OwnerReferences)) + c.Eq("test-shard", pod.OwnerReferences[0].Name, "owner name") + c.Eq("Shard", pod.OwnerReferences[0].Kind, "owner kind") // Verify labels expectedLabels := map[string]string{ @@ -90,14 +82,12 @@ func TestBuildPoolPod_BasicStructure(t *testing.T) { "multigres.com/tablegroup": "default", } for k, want := range expectedLabels { - if got := pod.Labels[k]; got != want { - t.Errorf("label %q = %q, want %q", k, got, want) - } + got := pod.Labels[k] + c.Eq(want, got, "label %q = %q, want", k, got) } - if got := pod.Annotations[metadata.AnnotationProjectRef]; got != "test-cluster" { - t.Errorf("annotation %q = %q, want %q", metadata.AnnotationProjectRef, got, "test-cluster") - } + got := pod.Annotations[metadata.AnnotationProjectRef] + c.Eq("test-cluster", got, "annotation %q = %q, want", metadata.AnnotationProjectRef, got) if len(pod.Spec.ReadinessGates) != 1 || pod.Spec.ReadinessGates[0].ConditionType != PoolerDataReadyCondition { t.Errorf( @@ -126,26 +116,23 @@ func TestBuildPoolPod_ProjectRefAnnotation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewAborting(t) shard := newTestShard() shard.Annotations = tc.annotations pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") - if got := pod.Annotations[metadata.AnnotationProjectRef]; got != tc.want { - t.Fatalf("annotation %q = %q, want %q", metadata.AnnotationProjectRef, got, tc.want) - } + got := pod.Annotations[metadata.AnnotationProjectRef] + c.Eq(tc.want, got, "annotation %q = %q, want", metadata.AnnotationProjectRef, got) }) } } func TestBuildPoolPod_PrometheusScrapeAnnotations(t *testing.T) { + c := assert.NewAborting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") wantAnnotations := map[string]string{ metadata.AnnotationPrometheusScrape: "true", @@ -153,55 +140,30 @@ func TestBuildPoolPod_PrometheusScrapeAnnotations(t *testing.T) { metadata.AnnotationPrometheusPath: "/metrics", } for key, want := range wantAnnotations { - if got := pod.Annotations[key]; got != want { - t.Fatalf("annotation %q = %q, want %q", key, got, want) - } + got := pod.Annotations[key] + c.Eq(want, got, "annotation %q = %q, want", key, got) } } func TestBuildPoolPod_Containers(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if len(pod.Spec.InitContainers) != 1 { - t.Fatalf( - "expected 1 init container (pgctld sidecar), got %d", - len(pod.Spec.InitContainers), - ) - } - if pod.Spec.InitContainers[0].Name != "postgres" { - t.Errorf( - "init container name = %q, want %q", - pod.Spec.InitContainers[0].Name, - "postgres", - ) - } + c.Require(). + Len(pod.Spec.InitContainers, 1, "expected 1 init container (pgctld sidecar), got %d", len(pod.Spec.InitContainers)) + c.Eq("postgres", pod.Spec.InitContainers[0].Name, "init container name") - if len(pod.Spec.Containers) != 2 { - t.Fatalf( - "expected 2 containers (multipooler + postgres-exporter), got %d", - len(pod.Spec.Containers), - ) - } - if pod.Spec.Containers[0].Name != "multipooler" { - t.Errorf("container name = %q, want %q", pod.Spec.Containers[0].Name, "multipooler") - } - if pod.Spec.Containers[1].Name != "postgres-exporter" { - t.Errorf( - "container name = %q, want %q", - pod.Spec.Containers[1].Name, - "postgres-exporter", - ) - } + c.Require(). + Len(pod.Spec.Containers, 2, "expected 2 containers (multipooler + postgres-exporter), got %d", len(pod.Spec.Containers)) + c.Eq("multipooler", pod.Spec.Containers[0].Name, "container name") + c.Eq("postgres-exporter", pod.Spec.Containers[1].Name, "container name") } func TestBuildPoolPod_Volumes(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") volumeNames := make(map[string]bool) for _, v := range pod.Spec.Volumes { @@ -216,26 +178,21 @@ func TestBuildPoolPod_Volumes(t *testing.T) { PostgresPasswordVolumeName, } for _, name := range required { - if !volumeNames[name] { - t.Errorf("missing required volume %q", name) - } + c.False(!volumeNames[name], "missing required volume %q", name) } // Verify data volume references PVC for _, v := range pod.Spec.Volumes { if v.Name == DataVolumeName { - if v.PersistentVolumeClaim == nil { - t.Fatal("data volume should reference a PVC") - } + c.Require().NotNil(v.PersistentVolumeClaim, "data volume should reference a PVC") pvcName := v.PersistentVolumeClaim.ClaimName - if pvcName == "" { - t.Error("data volume PVC claim name is empty") - } + c.NotEq("", pvcName, "data volume PVC claim name is empty") } } } func TestBuildPoolPod_UsesShardWideBackupPVC(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.Backup = &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, @@ -245,13 +202,9 @@ func TestBuildPoolPod_UsesShardWideBackupPVC(t *testing.T) { pool := newTestPoolSpec() pool.Cells = []multigresv1alpha1.CellName{"zone-a", "zone-b"} podA, err := BuildPoolPod(shard, "main", "zone-a", pool, 0, testScheme()) - if err != nil { - t.Fatalf("build zone-a pooler pod: %v", err) - } + c.Require().NoError(err, "build zone-a pooler pod") podB, err := BuildPoolPod(shard, "main", "zone-b", pool, 0, testScheme()) - if err != nil { - t.Fatalf("build zone-b pooler pod: %v", err) - } + c.Require().NoError(err, "build zone-b pooler pod") backupClaim := func(pod *corev1.Pod) string { for _, volume := range pod.Spec.Volumes { @@ -263,34 +216,24 @@ func TestBuildPoolPod_UsesShardWideBackupPVC(t *testing.T) { } want := BuildSharedBackupPVCName(shard) - if got := backupClaim(podA); got != want { - t.Errorf("zone-a backup claim = %q, want %q", got, want) - } - if got := backupClaim(podB); got != want { - t.Errorf("zone-b backup claim = %q, want %q", got, want) - } + c.Eq(want, backupClaim(podA), "zone-a backup claim") + c.Eq(want, backupClaim(podB), "zone-b backup claim") } func TestBuildPoolPod_PostgresPasswordFile(t *testing.T) { + ck := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") passwordVolume := findVolume(pod.Spec.Volumes, PostgresPasswordVolumeName) - if passwordVolume == nil { - t.Fatalf("missing postgres password volume %q", PostgresPasswordVolumeName) - } - if passwordVolume.Secret == nil { - t.Fatal("postgres password volume should use Secret source") - } - if passwordVolume.Secret.SecretName != "multigres-admin-password" { - t.Errorf( - "postgres password SecretName = %q, want %q", - passwordVolume.Secret.SecretName, - "multigres-admin-password", - ) - } + ck.Require(). + NotNil(passwordVolume, "missing postgres password volume %q", PostgresPasswordVolumeName) + ck.Require().NotNil(passwordVolume.Secret, "postgres password volume should use Secret source") + ck.Eq( + "multigres-admin-password", + passwordVolume.Secret.SecretName, + "postgres password SecretName", + ) if passwordVolume.Secret.DefaultMode == nil || *passwordVolume.Secret.DefaultMode != 0o444 { t.Errorf("postgres password defaultMode = %v, want 0444", passwordVolume.Secret.DefaultMode) } @@ -327,6 +270,7 @@ func TestBuildPoolPod_PostgresPasswordFile(t *testing.T) { } func TestBuildPoolPod_PostgresPasswordSecretRef(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.PostgresPasswordSecretRef = multigresv1alpha1.PostgresPasswordSecretRef{ Name: "multigres-admin-password", @@ -334,23 +278,17 @@ func TestBuildPoolPod_PostgresPasswordSecretRef(t *testing.T) { } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") passwordVolume := findVolume(pod.Spec.Volumes, PostgresPasswordVolumeName) - if passwordVolume == nil { - t.Fatalf("missing postgres password volume %q", PostgresPasswordVolumeName) - } - if passwordVolume.Secret == nil { - t.Fatal("postgres password volume should use Secret source") - } - if passwordVolume.Secret.SecretName != "multigres-admin-password" { - t.Errorf( - "postgres password SecretName = %q, want multigres-admin-password", - passwordVolume.Secret.SecretName, - ) - } + c.Require(). + NotNil(passwordVolume, "missing postgres password volume %q", PostgresPasswordVolumeName) + c.Require().NotNil(passwordVolume.Secret, "postgres password volume should use Secret source") + c.Eq( + "multigres-admin-password", + passwordVolume.Secret.SecretName, + "postgres password SecretName", + ) if len(passwordVolume.Secret.Items) != 1 || passwordVolume.Secret.Items[0].Key != "current" || passwordVolume.Secret.Items[0].Path != PostgresPasswordSecretKey { @@ -372,6 +310,7 @@ func TestBuildPoolPod_PostgresPasswordSecretRef(t *testing.T) { } func TestBuildPoolPod_PostgresInitSecretsRef(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.PostgresInitSecretsRef = &multigresv1alpha1.PostgresInitSecretsRef{ Name: "multigres-init-secrets", @@ -379,23 +318,18 @@ func TestBuildPoolPod_PostgresInitSecretsRef(t *testing.T) { } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") initSecretsVolume := findVolume(pod.Spec.Volumes, PostgresInitSecretsVolumeName) - if initSecretsVolume == nil { - t.Fatalf("missing postgres init-secrets volume %q", PostgresInitSecretsVolumeName) - } - if initSecretsVolume.Secret == nil { - t.Fatal("postgres init-secrets volume should use Secret source") - } - if initSecretsVolume.Secret.SecretName != "multigres-init-secrets" { - t.Errorf( - "postgres init-secrets SecretName = %q, want multigres-init-secrets", - initSecretsVolume.Secret.SecretName, - ) - } + c.Require(). + NotNil(initSecretsVolume, "missing postgres init-secrets volume %q", PostgresInitSecretsVolumeName) + c.Require(). + NotNil(initSecretsVolume.Secret, "postgres init-secrets volume should use Secret source") + c.Eq( + "multigres-init-secrets", + initSecretsVolume.Secret.SecretName, + "postgres init-secrets SecretName", + ) if len(initSecretsVolume.Secret.Items) != 1 || initSecretsVolume.Secret.Items[0].Key != "custom-key.json" || initSecretsVolume.Secret.Items[0].Path != PostgresInitSecretsFileName { @@ -425,21 +359,20 @@ func TestBuildPoolPod_PostgresInitSecretsRef(t *testing.T) { } func TestBuildPoolPod_PostgresInitSecretsRef_Absent(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if v := findVolume(pod.Spec.Volumes, PostgresInitSecretsVolumeName); v != nil { - t.Errorf("expected no postgres init-secrets volume, got %+v", v) - } + c.Nil( + findVolume(pod.Spec.Volumes, PostgresInitSecretsVolumeName), + "expected no postgres init-secrets volume, got", + ) } func TestComputeSpecHash_ChangesOnPostgresInitSecretsRef(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") wantHash := ComputeSpecHash(pod) shardWithRef := newTestShard() @@ -447,27 +380,21 @@ func TestComputeSpecHash_ChangesOnPostgresInitSecretsRef(t *testing.T) { Name: "multigres-init-secrets", } podWithRef, err := BuildPoolPod(shardWithRef, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if got := ComputeSpecHash(podWithRef); got == wantHash { - t.Error("spec hash should differ when postgres init-secrets ref is set vs unset") - } + c.NotEq( + wantHash, + ComputeSpecHash(podWithRef), + "spec hash should differ when postgres init-secrets ref is set vs unset", + ) } func TestBuildPoolPod_SecurityContext(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if pod.Spec.SecurityContext != nil { - t.Errorf( - "pod security context = %v, want nil when fsGroup is not configured", - pod.Spec.SecurityContext, - ) - } + c.Nil(pod.Spec.SecurityContext, "pod security context") if pod.Spec.TerminationGracePeriodSeconds == nil || *pod.Spec.TerminationGracePeriodSeconds != 30 { @@ -487,9 +414,7 @@ func TestBuildPoolPod_FSGroupDoesNotOverrideContainerRuntimeIdentity(t *testing. pool.FSGroup = ptr.To(int64(2000)) pod, err := BuildPoolPod(shard, "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") // Overriding the image must not drop the numeric identity. pgctld // declares USER postgres by name, so leaving RunAsUser unset pairs @@ -522,9 +447,7 @@ func TestBuildPoolPod_FSGroupDoesNotOverrideContainerRuntimeIdentity(t *testing. pool.FSGroup = ptr.To(int64(2000)) pod, err := BuildPoolPod(newTestShard(), "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") assertPoolPodFSGroup(t, pod, 2000) assertContainerIdentity( @@ -559,9 +482,7 @@ func TestBuildPoolPod_FSGroupDoesNotOverrideContainerRuntimeIdentity(t *testing. pool.Multipooler.RunAsGroup = ptr.To(int64(3001)) pod, err := BuildPoolPod(shard, "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") assertPoolPodFSGroup(t, pod, 2000) assertContainerIdentity( @@ -593,9 +514,7 @@ func TestBuildPoolPod_FSGroupDoesNotOverrideContainerRuntimeIdentity(t *testing. pool.Postgres.RunAsGroup = ptr.To(int64(1001)) pod, err := BuildPoolPod(shard, "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") assertPoolPodFSGroup(t, pod, 2000) assertContainerIdentity( @@ -618,7 +537,7 @@ func TestBuildPoolPod_FSGroupDoesNotOverrideContainerRuntimeIdentity(t *testing. pool.Multipooler.RunAsUser = ptr.To(int64(1000)) _, err := BuildPoolPod(newTestShard(), "main", "z1", pool, 0, testScheme()) - require.ErrorContains(t, err, "requires matching postgres runAsUser") + assert.NewAborting(t).ErrorContains(err, "requires matching postgres runAsUser") }) t.Run("rejects mismatched explicit shared data UIDs", func(t *testing.T) { @@ -628,44 +547,44 @@ func TestBuildPoolPod_FSGroupDoesNotOverrideContainerRuntimeIdentity(t *testing. pool.Multipooler.RunAsUser = ptr.To(int64(1000)) _, err := BuildPoolPod(newTestShard(), "main", "z1", pool, 0, testScheme()) - require.ErrorContains(t, err, "must match because both access PGDATA") + assert.NewAborting(t).ErrorContains(err, "must match because both access PGDATA") }) } func TestBuildContainerSecurityContext(t *testing.T) { t.Run("image identity", func(t *testing.T) { + c := assert.NewCollecting(t) sc := buildContainerSecurityContext(nil, nil) - assert.True(t, *sc.RunAsNonRoot) - assert.Nil(t, sc.RunAsUser) - assert.Nil(t, sc.RunAsGroup) + c.True(*sc.RunAsNonRoot) + c.Nil(sc.RunAsUser) + c.Nil(sc.RunAsGroup) }) t.Run("explicit identity", func(t *testing.T) { + c := assert.NewCollecting(t) sc := buildContainerSecurityContext(ptr.To(int64(1000)), ptr.To(int64(1001))) - assert.True(t, *sc.RunAsNonRoot) - assert.Equal(t, int64(1000), *sc.RunAsUser) - assert.Equal(t, int64(1001), *sc.RunAsGroup) + c.True(*sc.RunAsNonRoot) + c.EqDeep(int64(1000), *sc.RunAsUser) + c.EqDeep(int64(1001), *sc.RunAsGroup) }) t.Run("non-root user with root group", func(t *testing.T) { + c := assert.NewCollecting(t) sc := buildContainerSecurityContext(ptr.To(int64(1000)), ptr.To(int64(0))) - assert.True(t, *sc.RunAsNonRoot) - assert.Equal(t, int64(1000), *sc.RunAsUser) - assert.Equal(t, int64(0), *sc.RunAsGroup) + c.True(*sc.RunAsNonRoot) + c.EqDeep(int64(1000), *sc.RunAsUser) + c.EqDeep(int64(0), *sc.RunAsGroup) }) } func assertPoolPodFSGroup(t *testing.T, pod *corev1.Pod, want int64) { t.Helper() - if pod.Spec.SecurityContext == nil { - t.Fatal("pod security context is nil") - } - if pod.Spec.SecurityContext.FSGroup == nil { - t.Fatal("pod fsGroup is nil") - } - assert.Equal(t, want, *pod.Spec.SecurityContext.FSGroup) - assert.Nil(t, pod.Spec.SecurityContext.RunAsUser) - assert.Nil(t, pod.Spec.SecurityContext.RunAsGroup) + c := assert.NewCollecting(t) + c.Require().NotNil(pod.Spec.SecurityContext, "pod security context is nil") + c.Require().NotNil(pod.Spec.SecurityContext.FSGroup, "pod fsGroup is nil") + c.EqDeep(want, *pod.Spec.SecurityContext.FSGroup) + c.Nil(pod.Spec.SecurityContext.RunAsUser) + c.Nil(pod.Spec.SecurityContext.RunAsGroup) } func assertContainerIdentity( @@ -675,61 +594,49 @@ func assertContainerIdentity( wantGroup *int64, ) { t.Helper() - if container.SecurityContext == nil { - t.Fatalf("container %q security context is nil", container.Name) - } - assert.True(t, *container.SecurityContext.RunAsNonRoot) - assert.Equal(t, wantUser, container.SecurityContext.RunAsUser, container.Name) - assert.Equal(t, wantGroup, container.SecurityContext.RunAsGroup, container.Name) + c := assert.NewCollecting(t) + c.Require(). + NotNil(container.SecurityContext, "container %q security context is nil", container.Name) + c.True(*container.SecurityContext.RunAsNonRoot) + c.EqDeep(wantUser, container.SecurityContext.RunAsUser, container.Name) + c.EqDeep(wantGroup, container.SecurityContext.RunAsGroup, container.Name) } func TestBuildPoolPod_SpecHash(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") hash, ok := pod.Annotations[metadata.AnnotationSpecHash] - if !ok { - t.Fatal("spec-hash annotation missing") - } - if hash == "" { - t.Error("spec-hash annotation is empty") - } - if len(hash) != 8 { - t.Errorf("spec-hash length = %d, want 8 (FNV-1a 32-bit hex)", len(hash)) - } + c.Require().True(ok, "spec-hash annotation missing") + c.NotEq("", hash, "spec-hash annotation is empty") + c.Len(hash, 8, "spec-hash length = %d, want 8 (FNV-1a 32-bit hex)", len(hash)) } func TestComputeSpecHash_ChangesOnPostgresPasswordFileSpec(t *testing.T) { pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") wantHash := ComputeSpecHash(pod) t.Run("volume", func(t *testing.T) { oldPod := pod.DeepCopy() oldPod.Spec.Volumes = removeVolume(oldPod.Spec.Volumes, PostgresPasswordVolumeName) - if got := ComputeSpecHash(oldPod); got == wantHash { - t.Error("spec hash should differ when postgres password volume is removed") - } + assert.NewCollecting(t). + NotEq(wantHash, ComputeSpecHash(oldPod), "spec hash should differ when postgres password volume is removed") }) t.Run("pgctld env", func(t *testing.T) { oldPod := pod.DeepCopy() useLegacyPasswordEnv(&oldPod.Spec.InitContainers[0]) - if got := ComputeSpecHash(oldPod); got == wantHash { - t.Error("spec hash should differ when pgctld password env changes") - } + assert.NewCollecting(t). + NotEq(wantHash, ComputeSpecHash(oldPod), "spec hash should differ when pgctld password env changes") }) t.Run("multipooler env", func(t *testing.T) { oldPod := pod.DeepCopy() useLegacyPasswordEnv(&oldPod.Spec.Containers[0]) - if got := ComputeSpecHash(oldPod); got == wantHash { - t.Error("spec hash should differ when multipooler password env changes") - } + assert.NewCollecting(t). + NotEq(wantHash, ComputeSpecHash(oldPod), "spec hash should differ when multipooler password env changes") }) t.Run("volume mount", func(t *testing.T) { @@ -738,9 +645,8 @@ func TestComputeSpecHash_ChangesOnPostgresPasswordFileSpec(t *testing.T) { oldPod.Spec.Containers[0].VolumeMounts, PostgresPasswordVolumeName, ) - if got := ComputeSpecHash(oldPod); got == wantHash { - t.Error("spec hash should differ when postgres password volume mount is removed") - } + assert.NewCollecting(t). + NotEq(wantHash, ComputeSpecHash(oldPod), "spec hash should differ when postgres password volume mount is removed") }) t.Run("volume mount read-only", func(t *testing.T) { @@ -750,24 +656,21 @@ func TestComputeSpecHash_ChangesOnPostgresPasswordFileSpec(t *testing.T) { changedPod.Spec.Containers[0].VolumeMounts[i].ReadOnly = false } } - if got := ComputeSpecHash(changedPod); got == wantHash { - t.Error("spec hash should differ when postgres password volume mount readOnly changes") - } + assert.NewCollecting(t). + NotEq(wantHash, ComputeSpecHash(changedPod), "spec hash should differ when postgres password volume mount readOnly changes") }) } func TestBuildPoolPod_NoFinalizers(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if len(pod.Finalizers) != 0 { - t.Errorf("finalizers = %v, want none", pod.Finalizers) - } + c.Empty(pod.Finalizers, "finalizers") } func TestBuildPoolPod_Affinity(t *testing.T) { + c := assert.NewAborting(t) pool := newTestPoolSpec() pool.Affinity = &corev1.Affinity{ NodeAffinity: &corev1.NodeAffinity{ @@ -786,16 +689,16 @@ func TestBuildPoolPod_Affinity(t *testing.T) { } pod, err := BuildPoolPod(newTestShard(), "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") - if pod.Spec.Affinity == nil || pod.Spec.Affinity.NodeAffinity == nil { - t.Fatal("affinity not set on pod") - } + c.False( + pod.Spec.Affinity == nil || pod.Spec.Affinity.NodeAffinity == nil, + "affinity not set on pod", + ) } func TestBuildPoolPod_Tolerations(t *testing.T) { + c := assert.NewCollecting(t) pool := newTestPoolSpec() pool.Tolerations = []corev1.Toleration{ { @@ -807,19 +710,12 @@ func TestBuildPoolPod_Tolerations(t *testing.T) { } pod, err := BuildPoolPod(newTestShard(), "main", "z1", pool, 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if len(pod.Spec.Tolerations) != 1 { - t.Fatalf("expected 1 toleration, got %d", len(pod.Spec.Tolerations)) - } - if pod.Spec.Tolerations[0].Key != "dedicated" { - t.Errorf("toleration key = %q, want %q", pod.Spec.Tolerations[0].Key, "dedicated") - } - if pod.Spec.Tolerations[0].Value != "database" { - t.Errorf("toleration value = %q, want %q", pod.Spec.Tolerations[0].Value, "database") - } + c.Require(). + Len(pod.Spec.Tolerations, 1, "expected 1 toleration, got %d", len(pod.Spec.Tolerations)) + c.Eq("dedicated", pod.Spec.Tolerations[0].Key, "toleration key") + c.Eq("database", pod.Spec.Tolerations[0].Value, "toleration value") } func TestComputeSpecHash_ChangesOnTolerations(t *testing.T) { @@ -840,9 +736,7 @@ func TestComputeSpecHash_ChangesOnTolerations(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when tolerations change") - } + assert.NewCollecting(t).NotEq(hash2, hash1, "spec hash should differ when tolerations change") } func TestComputeSpecHash_ChangesOnFSGroup(t *testing.T) { @@ -856,49 +750,41 @@ func TestComputeSpecHash_ChangesOnFSGroup(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when fsGroup changes") - } + assert.NewCollecting(t).NotEq(hash2, hash1, "spec hash should differ when fsGroup changes") } func TestComputeSpecHash_ChangesOnRuntimeIdentity(t *testing.T) { + c := assert.NewCollecting(t) pool1 := newTestPoolSpec() pool1.Postgres.RunAsUser = ptr.To(int64(1000)) pool1.Multipooler.RunAsUser = ptr.To(int64(1000)) pod1, err := BuildPoolPod(newTestShard(), "main", "z1", pool1, 0, testScheme()) - require.NoError(t, err) + c.Require().NoError(err) pool2 := newTestPoolSpec() pool2.Postgres.RunAsUser = ptr.To(int64(2000)) pool2.Multipooler.RunAsUser = ptr.To(int64(2000)) pod2, err := BuildPoolPod(newTestShard(), "main", "z1", pool2, 0, testScheme()) - require.NoError(t, err) + c.Require().NoError(err) - assert.NotEqual( - t, + c.NotEqDeep( pod1.Annotations[metadata.AnnotationSpecHash], pod2.Annotations[metadata.AnnotationSpecHash], ) } func TestBuildPoolPod_NodeSelector(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.CellTopologyLabels = map[multigresv1alpha1.CellName]map[string]string{ "z1": {"topology.kubernetes.io/zone": "us-east-1a"}, } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if pod.Spec.NodeSelector == nil { - t.Fatal("node selector is nil") - } - if pod.Spec.NodeSelector["topology.kubernetes.io/zone"] != "us-east-1a" { - t.Errorf("node selector zone = %q, want %q", - pod.Spec.NodeSelector["topology.kubernetes.io/zone"], "us-east-1a") - } + c.Require().NotNil(pod.Spec.NodeSelector, "node selector is nil") + c.Eq("us-east-1a", pod.Spec.NodeSelector["topology.kubernetes.io/zone"], "node selector zone") } func TestBuildPoolPod_Hostname(t *testing.T) { @@ -917,6 +803,7 @@ func TestBuildPoolPod_Hostname(t *testing.T) { } for name, tt := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Labels["multigres.com/cluster"] = tt.clusterName shard.Namespace = tt.namespace @@ -926,7 +813,7 @@ func TestBuildPoolPod_Hostname(t *testing.T) { multigresv1alpha1.PoolName(tt.poolName): pool, } pod, err := BuildPoolPod(shard, tt.poolName, tt.cellName, pool, tt.index, testScheme()) - require.NoError(t, err) + c.Require().NoError(err) svc, err := BuildPoolHeadlessService( shard, tt.poolName, @@ -934,28 +821,27 @@ func TestBuildPoolPod_Hostname(t *testing.T) { pool, testScheme(), ) - require.NoError(t, err) + c.Require().NoError(err) - assert.Equal(t, pod.Name, pod.Spec.Hostname) - assert.Equal(t, svc.Name, pod.Spec.Subdomain) - assert.True(t, svc.Spec.PublishNotReadyAddresses) + c.EqDeep(pod.Name, pod.Spec.Hostname) + c.EqDeep(svc.Name, pod.Spec.Subdomain) + c.True(svc.Spec.PublishNotReadyAddresses) for key, value := range svc.Spec.Selector { - assert.Equal(t, value, pod.Labels[key], "headless service must select the pod") + c.EqDeep(value, pod.Labels[key], "headless service must select the pod") } hostname := fmt.Sprintf("%s.%s.%s.svc.cluster.local", pod.Name, svc.Name, tt.namespace) var hostnameArgs []string for _, container := range pod.Spec.Containers { for _, arg := range container.Args { if strings.HasPrefix(arg, "--hostname=") { - assert.Equal(t, "multipooler", container.Name) + c.EqDeep("multipooler", container.Name) hostnameArgs = append(hostnameArgs, arg) } } } - assert.Equal(t, []string{"--hostname=" + hostname}, hostnameArgs) + c.EqDeep([]string{"--hostname=" + hostname}, hostnameArgs) cert := &x509.Certificate{DNSNames: pgBackRestPoolDNSNames(shard)} - assert.NoError( - t, + c.NoError( cert.VerifyHostname(hostname), "advertised address must match backup TLS SANs", ) @@ -968,18 +854,15 @@ func TestBuildPoolPod_Hostname(t *testing.T) { func(arg string) bool { return strings.HasPrefix(arg, "--hostname=") }, ) } - assert.Equal(t, ComputeSpecHash(pod), pod.Annotations[metadata.AnnotationSpecHash]) - assert.NotEqual( - t, - ComputeSpecHash(legacyPod), - pod.Annotations[metadata.AnnotationSpecHash], - ) + c.EqDeep(ComputeSpecHash(pod), pod.Annotations[metadata.AnnotationSpecHash]) + c.NotEqDeep(ComputeSpecHash(legacyPod), pod.Annotations[metadata.AnnotationSpecHash]) }) } } func TestBuildPoolPod_ServiceAccountName(t *testing.T) { t.Run("set when S3 serviceAccountName configured", func(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.Backup = &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeS3, @@ -990,29 +873,19 @@ func TestBuildPoolPod_ServiceAccountName(t *testing.T) { }, } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pod.Spec.ServiceAccountName != "multigres-backup" { - t.Errorf( - "ServiceAccountName = %q, want %q", - pod.Spec.ServiceAccountName, - "multigres-backup", - ) - } + c.Require().NoError(err, "unexpected error") + c.Eq("multigres-backup", pod.Spec.ServiceAccountName, "ServiceAccountName") }) t.Run("empty when no backup config", func(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pod.Spec.ServiceAccountName != "" { - t.Errorf("ServiceAccountName = %q, want empty", pod.Spec.ServiceAccountName) - } + c.Require().NoError(err, "unexpected error") + c.Eq("", pod.Spec.ServiceAccountName, "ServiceAccountName") }) t.Run("empty when S3 has no serviceAccountName", func(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.Backup = &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeS3, @@ -1022,26 +895,19 @@ func TestBuildPoolPod_ServiceAccountName(t *testing.T) { }, } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pod.Spec.ServiceAccountName != "" { - t.Errorf("ServiceAccountName = %q, want empty", pod.Spec.ServiceAccountName) - } + c.Require().NoError(err, "unexpected error") + c.Eq("", pod.Spec.ServiceAccountName, "ServiceAccountName") }) t.Run("empty when backup is filesystem type", func(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.Backup = &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pod.Spec.ServiceAccountName != "" { - t.Errorf("ServiceAccountName = %q, want empty", pod.Spec.ServiceAccountName) - } + c.Require().NoError(err, "unexpected error") + c.Eq("", pod.Spec.ServiceAccountName, "ServiceAccountName") }) } @@ -1063,9 +929,8 @@ func TestComputeSpecHash_ChangesOnServiceAccountName(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when ServiceAccountName is added") - } + assert.NewCollecting(t). + NotEq(hash2, hash1, "spec hash should differ when ServiceAccountName is added") } func TestComputeSpecHash_Deterministic(t *testing.T) { @@ -1075,16 +940,15 @@ func TestComputeSpecHash_Deterministic(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 != hash2 { - t.Errorf("spec hash not deterministic: %q != %q", hash1, hash2) - } + assert.NewCollecting(t).Eq(hash2, hash1, "spec hash not deterministic") } func TestPodNeedsUpdateWhenReadinessProtectionsChange(t *testing.T) { + c := assert.NewAborting(t) shard := newTestShard() pool := newTestPoolSpec() desired, err := BuildPoolPod(shard, "main", "z1", pool, 0, testScheme()) - require.NoError(t, err) + c.NoError(err) legacy := desired.DeepCopy() legacy.Spec.ReadinessGates = nil @@ -1096,9 +960,10 @@ func TestPodNeedsUpdateWhenReadinessProtectionsChange(t *testing.T) { } legacy.Annotations[metadata.AnnotationSpecHash] = ComputeSpecHash(legacy) - if !podNeedsUpdate(legacy, shard, "main", "z1", pool, 0, testScheme()) { - t.Fatal("pod without the current readiness gates and probes must be replaced") - } + c.True( + podNeedsUpdate(legacy, shard, "main", "z1", pool, 0, testScheme()), + "pod without the current readiness gates and probes must be replaced", + ) } func TestComputeSpecHashIncludesReadinessGatesAndProbes(t *testing.T) { @@ -1110,7 +975,7 @@ func TestComputeSpecHashIncludesReadinessGatesAndProbes(t *testing.T) { 0, testScheme(), ) - require.NoError(t, err) + assert.NewAborting(t).NoError(err) wantHash := ComputeSpecHash(desired) tests := map[string]func(*corev1.Pod){ @@ -1131,9 +996,8 @@ func TestComputeSpecHashIncludesReadinessGatesAndProbes(t *testing.T) { t.Run(name, func(t *testing.T) { changed := desired.DeepCopy() mutate(changed) - if got := ComputeSpecHash(changed); got == wantHash { - t.Fatalf("spec hash did not change when %s changed", name) - } + assert.NewAborting(t). + NotEq(wantHash, ComputeSpecHash(changed), "spec hash did not change when %s changed", name) }) } } @@ -1164,9 +1028,7 @@ func TestComputeSpecHash_ChangesOnDrift(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when affinity changes") - } + assert.NewCollecting(t).NotEq(hash2, hash1, "spec hash should differ when affinity changes") } func TestComputeSpecHash_ChangesOnValueFromDrift(t *testing.T) { @@ -1195,9 +1057,8 @@ func TestComputeSpecHash_ChangesOnValueFromDrift(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when ValueFrom secret name changes") - } + assert.NewCollecting(t). + NotEq(hash2, hash1, "spec hash should differ when ValueFrom secret name changes") } func TestComputeSpecHash_ChangesOnEnvFromDrift(t *testing.T) { @@ -1218,12 +1079,12 @@ func TestComputeSpecHash_ChangesOnEnvFromDrift(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when EnvFrom config map name changes") - } + assert.NewCollecting(t). + NotEq(hash2, hash1, "spec hash should differ when EnvFrom config map name changes") } func TestBuildPoolPodName_Truncation(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Labels["multigres.com/cluster"] = "very-long-cluster-name-for-testing" shard.Spec.DatabaseName = "long-database-name" @@ -1232,27 +1093,18 @@ func TestBuildPoolPodName_Truncation(t *testing.T) { name := BuildPoolPodName(shard, "main-pool", "us-east-1a", 99) - if len(name) > 63 { - t.Errorf("pod name %q exceeds 63 chars (len=%d)", name, len(name)) - } - if !strings.HasSuffix(name, "-99") { - t.Errorf("pod name %q should end with -99", name) - } + c.LessOrEqual(63, len(name), "pod name %q exceeds 63 chars (len=%d)", name, len(name)) + c.True(strings.HasSuffix(name, "-99"), "pod name %q should end with -99", name) } func TestBuildPoolPodName_ShortName(t *testing.T) { + c := assert.NewCollecting(t) name := BuildPoolPodName(newTestShard(), "main", "z1", 0) - if len(name) > 63 { - t.Errorf("pod name %q exceeds 63 chars (len=%d)", name, len(name)) - } - if !strings.HasSuffix(name, "-0") { - t.Errorf("pod name %q should end with -0", name) - } + c.LessOrEqual(63, len(name), "pod name %q exceeds 63 chars (len=%d)", name, len(name)) + c.True(strings.HasSuffix(name, "-0"), "pod name %q should end with -0", name) // Pod name should contain meaningful parts - if !strings.Contains(name, "pool") { - t.Errorf("pod name %q should contain 'pool'", name) - } + c.StrContains(name, "pool", "pod name") } func TestComputeSpecHash_ChangesOnPostgresConfigHash(t *testing.T) { @@ -1267,9 +1119,8 @@ func TestComputeSpecHash_ChangesOnPostgresConfigHash(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when postgres config hash annotation is added") - } + assert.NewCollecting(t). + NotEq(hash2, hash1, "spec hash should differ when postgres config hash annotation is added") } func TestComputeSpecHash_ChangesOnDifferentPostgresConfigHash(t *testing.T) { @@ -1287,37 +1138,31 @@ func TestComputeSpecHash_ChangesOnDifferentPostgresConfigHash(t *testing.T) { hash1 := ComputeSpecHash(pod1) hash2 := ComputeSpecHash(pod2) - if hash1 == hash2 { - t.Error("spec hash should differ when postgres config hash value changes") - } + assert.NewCollecting(t). + NotEq(hash2, hash1, "spec hash should differ when postgres config hash value changes") } func TestBuildPoolPod_PropagatesPostgresConfigHash(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Annotations = map[string]string{ metadata.AnnotationPostgresConfigHash: "deadbeef", } pod, err := BuildPoolPod(shard, "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") got := pod.Annotations[metadata.AnnotationPostgresConfigHash] - if got != "deadbeef" { - t.Errorf("postgres config hash annotation = %q, want %q", got, "deadbeef") - } + c.Eq("deadbeef", got, "postgres config hash annotation") } func TestBuildPoolPod_OmitsPostgresConfigHashWhenAbsent(t *testing.T) { + c := assert.NewCollecting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if _, ok := pod.Annotations[metadata.AnnotationPostgresConfigHash]; ok { - t.Error("postgres config hash annotation should not be present when shard has none") - } + _, ok := pod.Annotations[metadata.AnnotationPostgresConfigHash] + c.False(ok, "postgres config hash annotation should not be present when shard has none") } func findVolume(volumes []corev1.Volume, name string) *corev1.Volume { diff --git a/pkg/resource-handler/controller/shard/pool_pvc_test.go b/pkg/resource-handler/controller/shard/pool_pvc_test.go index 0fbe50f1..db9a4c85 100644 --- a/pkg/resource-handler/controller/shard/pool_pvc_test.go +++ b/pkg/resource-handler/controller/shard/pool_pvc_test.go @@ -9,27 +9,22 @@ import ( "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildPoolDataPVC_BasicStructure(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() pool := newTestPoolSpec() pvc, err := BuildPoolDataPVC(shard, "main", "z1", pool, 0, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if pvc.Namespace != "default" { - t.Errorf("namespace = %q, want %q", pvc.Namespace, "default") - } + c.Eq("default", pvc.Namespace, "namespace") - if len(pvc.OwnerReferences) != 0 { - t.Fatalf( - "expected 0 owner references with deleteOnShardRemoval=false, got %d", - len(pvc.OwnerReferences), - ) - } + c.Require(). + Empty(pvc.OwnerReferences, "expected 0 owner references with deleteOnShardRemoval=false, got %d", len(pvc.OwnerReferences)) expectedLabels := map[string]string{ "app.kubernetes.io/component": PoolComponentName, @@ -39,28 +34,24 @@ func TestBuildPoolDataPVC_BasicStructure(t *testing.T) { "multigres.com/shard": "0-inf", } for k, want := range expectedLabels { - if got := pvc.Labels[k]; got != want { - t.Errorf("label %q = %q, want %q", k, got, want) - } + got := pvc.Labels[k] + c.Eq(want, got, "label %q = %q, want", k, got) } } func TestBuildPoolDataPVC_StorageDefaults(t *testing.T) { + c := assert.NewCollecting(t) pool := multigresv1alpha1.PoolSpec{ Storage: multigresv1alpha1.StorageSpec{}, } pvc, err := BuildPoolDataPVC(newTestShard(), "main", "z1", pool, 0, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") // Default size got := pvc.Spec.Resources.Requests[corev1.ResourceStorage] want := resource.MustParse(DefaultDataVolumeSize) - if got.Cmp(want) != 0 { - t.Errorf("storage size = %s, want %s", got.String(), want.String()) - } + c.Eq(0, got.Cmp(want), "storage size = %s, want %s", got.String(), want.String()) // Default access mode if len(pvc.Spec.AccessModes) != 1 || pvc.Spec.AccessModes[0] != corev1.ReadWriteOnce { @@ -74,6 +65,7 @@ func TestBuildPoolDataPVC_StorageDefaults(t *testing.T) { } func TestBuildPoolDataPVC_CustomStorage(t *testing.T) { + c := assert.NewCollecting(t) pool := multigresv1alpha1.PoolSpec{ Storage: multigresv1alpha1.StorageSpec{ Class: "fast-ssd", @@ -83,16 +75,12 @@ func TestBuildPoolDataPVC_CustomStorage(t *testing.T) { } pvc, err := BuildPoolDataPVC(newTestShard(), "main", "z1", pool, 0, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") // Custom size got := pvc.Spec.Resources.Requests[corev1.ResourceStorage] want := resource.MustParse("50Gi") - if got.Cmp(want) != 0 { - t.Errorf("storage size = %s, want %s", got.String(), want.String()) - } + c.Eq(0, got.Cmp(want), "storage size = %s, want %s", got.String(), want.String()) // Custom access mode if len(pvc.Spec.AccessModes) != 1 || pvc.Spec.AccessModes[0] != corev1.ReadWriteMany { @@ -120,9 +108,7 @@ func TestBuildPoolDataPVC_NameConsistency(t *testing.T) { } } - if dataPVCRef != pvc.Name { - t.Errorf("pod references PVC %q but BuildPoolDataPVC created %q", dataPVCRef, pvc.Name) - } + assert.NewCollecting(t).Eq(pvc.Name, dataPVCRef, "pod references PVC") } func TestBuildPoolDataPVCName_MatchesPodReference(t *testing.T) { @@ -165,59 +151,45 @@ func TestBuildPoolDataPVCName_MatchesPodReference(t *testing.T) { } } - if pvcName != podPVCRef { - t.Errorf("BuildPoolDataPVCName() = %q, pod references %q", pvcName, podPVCRef) - } + assert.NewCollecting(t).Eq(podPVCRef, pvcName, "BuildPoolDataPVCName()") }) } } func TestBuildPoolDataPVC_OwnerRefWithDeletePolicy(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() pool := newTestPoolSpec() pvc, err := BuildPoolDataPVC(shard, "main", "z1", pool, 0, true, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if len(pvc.OwnerReferences) != 1 { - t.Fatalf( - "expected 1 owner reference with deleteOnShardRemoval=true, got %d", - len(pvc.OwnerReferences), - ) - } + c.Require(). + Len(pvc.OwnerReferences, 1, "expected 1 owner reference with deleteOnShardRemoval=true, got %d", len(pvc.OwnerReferences)) ref := pvc.OwnerReferences[0] - if ref.Name != shard.Name { - t.Errorf("ownerRef name = %q, want %q", ref.Name, shard.Name) - } - if ref.UID != shard.UID { - t.Errorf("ownerRef UID = %q, want %q", ref.UID, shard.UID) - } - if ref.Controller == nil || !*ref.Controller { - t.Error("ownerRef Controller should be true") - } + c.Eq(shard.Name, ref.Name, "ownerRef name") + c.Eq(shard.UID, ref.UID, "ownerRef UID") + c.False(ref.Controller == nil || !*ref.Controller, "ownerRef Controller should be true") } func TestBuildPoolDataPVC_NoOwnerRefWithRetainPolicy(t *testing.T) { + c := assert.NewAborting(t) shard := newTestShard() pool := newTestPoolSpec() pvc, err := BuildPoolDataPVC(shard, "main", "z1", pool, 0, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") - if len(pvc.OwnerReferences) != 0 { - t.Fatalf( - "expected 0 owner references with deleteOnShardRemoval=false, got %d", - len(pvc.OwnerReferences), - ) - } + c.Empty( + pvc.OwnerReferences, + "expected 0 owner references with deleteOnShardRemoval=false, got %d", + len(pvc.OwnerReferences), + ) } func TestBuildSharedBackupPVC_OwnerRefWithDeletePolicy(t *testing.T) { + c := assert.NewCollecting(t) shard := newTestShard() shard.Spec.Backup = &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, @@ -227,24 +199,17 @@ func TestBuildSharedBackupPVC_OwnerRefWithDeletePolicy(t *testing.T) { } pvc, err := BuildSharedBackupPVC(shard, true, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if len(pvc.OwnerReferences) != 1 { - t.Fatalf( - "expected 1 owner reference with deleteOnShardRemoval=true, got %d", - len(pvc.OwnerReferences), - ) - } + c.Require(). + Len(pvc.OwnerReferences, 1, "expected 1 owner reference with deleteOnShardRemoval=true, got %d", len(pvc.OwnerReferences)) ref := pvc.OwnerReferences[0] - if ref.Name != shard.Name { - t.Errorf("ownerRef name = %q, want %q", ref.Name, shard.Name) - } + c.Eq(shard.Name, ref.Name, "ownerRef name") } func TestBuildSharedBackupPVC_NoOwnerRefWithRetainPolicy(t *testing.T) { + c := assert.NewAborting(t) shard := newTestShard() shard.Spec.Backup = &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeFilesystem, @@ -254,19 +219,17 @@ func TestBuildSharedBackupPVC_NoOwnerRefWithRetainPolicy(t *testing.T) { } pvc, err := BuildSharedBackupPVC(shard, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") - if len(pvc.OwnerReferences) != 0 { - t.Fatalf( - "expected 0 owner references with deleteOnShardRemoval=false, got %d", - len(pvc.OwnerReferences), - ) - } + c.Empty( + pvc.OwnerReferences, + "expected 0 owner references with deleteOnShardRemoval=false, got %d", + len(pvc.OwnerReferences), + ) } func TestBuildShardPodDisruptionBudget(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -283,46 +246,28 @@ func TestBuildShardPodDisruptionBudget(t *testing.T) { } pdb, err := BuildShardPodDisruptionBudget(shard, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") - if pdb.Namespace != "default" { - t.Errorf("namespace = %q, want %q", pdb.Namespace, "default") - } + c.Eq("default", pdb.Namespace, "namespace") // Three desired members preserve two available members. if pdb.Spec.MinAvailable == nil || pdb.Spec.MinAvailable.IntValue() != 2 { t.Errorf("minAvailable = %v, want 2", pdb.Spec.MinAvailable) } - if pdb.Spec.MaxUnavailable != nil { - t.Errorf("maxUnavailable = %v, want nil", pdb.Spec.MaxUnavailable) - } + c.Nil(pdb.Spec.MaxUnavailable, "maxUnavailable") // Selector should match all pool pods in the shard, not one pool or cell. - if pdb.Spec.Selector == nil { - t.Fatal("selector is nil") - } + c.Require().NotNil(pdb.Spec.Selector, "selector is nil") sel := pdb.Spec.Selector.MatchLabels if _, ok := sel["multigres.com/cell"]; ok { t.Errorf("selector must not be scoped to a cell: %#v", sel) } - if _, ok := sel["multigres.com/pool"]; ok { - t.Errorf("selector must not be scoped to a pool: %#v", sel) - } - if sel["multigres.com/shard"] != "0-inf" { - t.Errorf("selector shard = %q, want %q", sel["multigres.com/shard"], "0-inf") - } - if sel["app.kubernetes.io/component"] != PoolComponentName { - t.Errorf( - "selector component = %q, want %q", - sel["app.kubernetes.io/component"], - PoolComponentName, - ) - } + _, ok := sel["multigres.com/pool"] + c.False(ok, "selector must not be scoped to a pool: %#v", sel) + c.Eq("0-inf", sel["multigres.com/shard"], "selector shard") + c.Eq(PoolComponentName, sel["app.kubernetes.io/component"], "selector component") // Owner reference - if len(pdb.OwnerReferences) != 1 { - t.Fatalf("expected 1 owner reference, got %d", len(pdb.OwnerReferences)) - } + c.Require(). + Len(pdb.OwnerReferences, 1, "expected 1 owner reference, got %d", len(pdb.OwnerReferences)) } diff --git a/pkg/resource-handler/controller/shard/pool_service_test.go b/pkg/resource-handler/controller/shard/pool_service_test.go index 89e65b01..21e16f6e 100644 --- a/pkg/resource-handler/controller/shard/pool_service_test.go +++ b/pkg/resource-handler/controller/shard/pool_service_test.go @@ -3,7 +3,6 @@ package shard import ( "testing" - "github.com/google/go-cmp/cmp" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" @@ -11,6 +10,8 @@ import ( "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildPoolHeadlessService(t *testing.T) { @@ -338,9 +339,7 @@ func TestBuildPoolHeadlessService(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildPoolHeadlessService() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildPoolHeadlessService() mismatch") }) } } diff --git a/pkg/resource-handler/controller/shard/ports_test.go b/pkg/resource-handler/controller/shard/ports_test.go index bcd29407..3cb855c6 100644 --- a/pkg/resource-handler/controller/shard/ports_test.go +++ b/pkg/resource-handler/controller/shard/ports_test.go @@ -5,6 +5,8 @@ import ( corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/util/intstr" + + "github.com/multigres/testkit/assert" ) func TestBuildMultipoolerContainerPorts(t *testing.T) { @@ -36,6 +38,7 @@ func TestBuildMultipoolerContainerPorts(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := buildMultipoolerContainerPorts() if len(got) != len(tt.want) { @@ -48,25 +51,21 @@ func TestBuildMultipoolerContainerPorts(t *testing.T) { } for i, port := range got { - if port.Name != tt.want[i].Name { - t.Errorf("port[%d].Name = %s, want %s", i, port.Name, tt.want[i].Name) - } - if port.ContainerPort != tt.want[i].ContainerPort { - t.Errorf( - "port[%d].ContainerPort = %d, want %d", - i, - port.ContainerPort, - tt.want[i].ContainerPort, - ) - } - if port.Protocol != tt.want[i].Protocol { - t.Errorf( - "port[%d].Protocol = %s, want %s", - i, - port.Protocol, - tt.want[i].Protocol, - ) - } + c.Eq(tt.want[i].Name, port.Name, "port[%d].Name = %s, want", i, port.Name) + c.Eq( + tt.want[i].ContainerPort, + port.ContainerPort, + "port[%d].ContainerPort = %d, want", + i, + port.ContainerPort, + ) + c.Eq( + tt.want[i].Protocol, + port.Protocol, + "port[%d].Protocol = %s, want", + i, + port.Protocol, + ) } }) } @@ -110,6 +109,7 @@ func TestBuildPoolHeadlessServicePorts(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := buildPoolHeadlessServicePorts() if len(got) != len(tt.want) { @@ -122,28 +122,22 @@ func TestBuildPoolHeadlessServicePorts(t *testing.T) { } for i, port := range got { - if port.Name != tt.want[i].Name { - t.Errorf("port[%d].Name = %s, want %s", i, port.Name, tt.want[i].Name) - } - if port.Port != tt.want[i].Port { - t.Errorf("port[%d].Port = %d, want %d", i, port.Port, tt.want[i].Port) - } - if port.TargetPort != tt.want[i].TargetPort { - t.Errorf( - "port[%d].TargetPort = %v, want %v", - i, - port.TargetPort, - tt.want[i].TargetPort, - ) - } - if port.Protocol != tt.want[i].Protocol { - t.Errorf( - "port[%d].Protocol = %s, want %s", - i, - port.Protocol, - tt.want[i].Protocol, - ) - } + c.Eq(tt.want[i].Name, port.Name, "port[%d].Name = %s, want", i, port.Name) + c.Eq(tt.want[i].Port, port.Port, "port[%d].Port = %d, want", i, port.Port) + c.Eq( + tt.want[i].TargetPort, + port.TargetPort, + "port[%d].TargetPort = %v, want", + i, + port.TargetPort, + ) + c.Eq( + tt.want[i].Protocol, + port.Protocol, + "port[%d].Protocol = %s, want", + i, + port.Protocol, + ) } }) } @@ -173,6 +167,7 @@ func TestBuildMultiorchContainerPorts(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := buildMultiorchContainerPorts() if len(got) != len(tt.want) { @@ -185,25 +180,21 @@ func TestBuildMultiorchContainerPorts(t *testing.T) { } for i, port := range got { - if port.Name != tt.want[i].Name { - t.Errorf("port[%d].Name = %s, want %s", i, port.Name, tt.want[i].Name) - } - if port.ContainerPort != tt.want[i].ContainerPort { - t.Errorf( - "port[%d].ContainerPort = %d, want %d", - i, - port.ContainerPort, - tt.want[i].ContainerPort, - ) - } - if port.Protocol != tt.want[i].Protocol { - t.Errorf( - "port[%d].Protocol = %s, want %s", - i, - port.Protocol, - tt.want[i].Protocol, - ) - } + c.Eq(tt.want[i].Name, port.Name, "port[%d].Name = %s, want", i, port.Name) + c.Eq( + tt.want[i].ContainerPort, + port.ContainerPort, + "port[%d].ContainerPort = %d, want", + i, + port.ContainerPort, + ) + c.Eq( + tt.want[i].Protocol, + port.Protocol, + "port[%d].Protocol = %s, want", + i, + port.Protocol, + ) } }) } @@ -235,6 +226,7 @@ func TestBuildMultiorchServicePorts(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := buildMultiorchServicePorts() if len(got) != len(tt.want) { @@ -247,28 +239,22 @@ func TestBuildMultiorchServicePorts(t *testing.T) { } for i, port := range got { - if port.Name != tt.want[i].Name { - t.Errorf("port[%d].Name = %s, want %s", i, port.Name, tt.want[i].Name) - } - if port.Port != tt.want[i].Port { - t.Errorf("port[%d].Port = %d, want %d", i, port.Port, tt.want[i].Port) - } - if port.TargetPort != tt.want[i].TargetPort { - t.Errorf( - "port[%d].TargetPort = %v, want %v", - i, - port.TargetPort, - tt.want[i].TargetPort, - ) - } - if port.Protocol != tt.want[i].Protocol { - t.Errorf( - "port[%d].Protocol = %s, want %s", - i, - port.Protocol, - tt.want[i].Protocol, - ) - } + c.Eq(tt.want[i].Name, port.Name, "port[%d].Name = %s, want", i, port.Name) + c.Eq(tt.want[i].Port, port.Port, "port[%d].Port = %d, want", i, port.Port) + c.Eq( + tt.want[i].TargetPort, + port.TargetPort, + "port[%d].TargetPort = %v, want", + i, + port.TargetPort, + ) + c.Eq( + tt.want[i].Protocol, + port.Protocol, + "port[%d].Protocol = %s, want", + i, + port.Protocol, + ) } }) } @@ -293,6 +279,7 @@ func TestBuildPostgresExporterContainerPorts(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) got := buildPostgresExporterContainerPorts() if len(got) != len(tt.want) { @@ -305,25 +292,21 @@ func TestBuildPostgresExporterContainerPorts(t *testing.T) { } for i, port := range got { - if port.Name != tt.want[i].Name { - t.Errorf("port[%d].Name = %s, want %s", i, port.Name, tt.want[i].Name) - } - if port.ContainerPort != tt.want[i].ContainerPort { - t.Errorf( - "port[%d].ContainerPort = %d, want %d", - i, - port.ContainerPort, - tt.want[i].ContainerPort, - ) - } - if port.Protocol != tt.want[i].Protocol { - t.Errorf( - "port[%d].Protocol = %s, want %s", - i, - port.Protocol, - tt.want[i].Protocol, - ) - } + c.Eq(tt.want[i].Name, port.Name, "port[%d].Name = %s, want", i, port.Name) + c.Eq( + tt.want[i].ContainerPort, + port.ContainerPort, + "port[%d].ContainerPort = %d, want", + i, + port.ContainerPort, + ) + c.Eq( + tt.want[i].Protocol, + port.Protocol, + "port[%d].Protocol = %s, want", + i, + port.Protocol, + ) } }) } diff --git a/pkg/resource-handler/controller/shard/postgres_config_test.go b/pkg/resource-handler/controller/shard/postgres_config_test.go index d67abae6..f7d540da 100644 --- a/pkg/resource-handler/controller/shard/postgres_config_test.go +++ b/pkg/resource-handler/controller/shard/postgres_config_test.go @@ -16,36 +16,33 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/postgresconfig" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestSetPostgresConfigStatus(t *testing.T) { r := &ShardReconciler{} t.Run("settled clears InProgress and stamps time", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} r.setPostgresConfigStatus(shard, false, nil) st := shard.Status.PostgresConfig - if st == nil || st.InProgress { - t.Fatalf("status = %+v, want InProgress false", st) - } - if st.LastAppliedAt == nil { - t.Error("LastAppliedAt should be set when settled") - } - if st.Error != "" { - t.Errorf("Error = %q, want empty", st.Error) - } + c.Require().False(st == nil || st.InProgress, "status = %+v, want InProgress false", st) + c.NotNil(st.LastAppliedAt, "LastAppliedAt should be set when settled") + c.Eq("", st.Error, "Error") }) t.Run("in-progress sets InProgress and does not stamp", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} r.setPostgresConfigStatus(shard, true, nil) st := shard.Status.PostgresConfig - if !st.InProgress { - t.Error("InProgress should be true during a rollout") - } - if st.LastAppliedAt != nil { - t.Error("LastAppliedAt should not be stamped while a rollout is in progress") - } + c.True(st.InProgress, "InProgress should be true during a rollout") + c.Nil( + st.LastAppliedAt, + "LastAppliedAt should not be stamped while a rollout is in progress", + ) }) // The key fix: a rollout driven by a PostgresConfigRef edit (or a new @@ -60,9 +57,8 @@ func TestSetPostgresConfigStatus(t *testing.T) { }, } r.setPostgresConfigStatus(shard, true, nil) - if !shard.Status.PostgresConfig.InProgress { - t.Error("InProgress should be true for a content-driven rollout at a steady generation") - } + assert.NewCollecting(t). + True(shard.Status.PostgresConfig.InProgress, "InProgress should be true for a content-driven rollout at a steady generation") }) t.Run("settling after a rollout re-stamps LastAppliedAt", func(t *testing.T) { @@ -77,9 +73,8 @@ func TestSetPostgresConfigStatus(t *testing.T) { } r.setPostgresConfigStatus(shard, false, nil) st := shard.Status.PostgresConfig - if st.InProgress { - t.Error("InProgress should clear once the config settles") - } + assert.NewCollecting(t). + False(st.InProgress, "InProgress should clear once the config settles") if st.LastAppliedAt == nil || !st.LastAppliedAt.After(past.Time) { t.Errorf("LastAppliedAt should advance on settle, got %v", st.LastAppliedAt) } @@ -99,15 +94,12 @@ func TestSetPostgresConfigStatus(t *testing.T) { }) t.Run("config error is reported and is not in progress", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{} r.setPostgresConfigStatus(shard, false, errTest) st := shard.Status.PostgresConfig - if st.InProgress { - t.Error("InProgress should be false when a config error is reported") - } - if !strings.Contains(st.Error, "boom") { - t.Errorf("Error = %q, want it to contain the failure", st.Error) - } + c.False(st.InProgress, "InProgress should be false when a config error is reported") + c.StrContains(st.Error, "boom", "Error") }) } @@ -117,15 +109,12 @@ func TestRenderEffectiveConfig(t *testing.T) { _ = corev1.AddToScheme(scheme) t.Run("no ref returns a stable hash", func(t *testing.T) { + c := assert.NewCollecting(t) r := &ShardReconciler{Client: fake.NewClientBuilder().WithScheme(scheme).Build()} shard := &multigresv1alpha1.Shard{ObjectMeta: metav1.ObjectMeta{Name: "s1"}} rc := r.renderEffectiveConfig(context.Background(), shard) - if rc.err != nil { - t.Fatalf("unexpected error: %v", rc.err) - } - if len(rc.restartHash) != 64 { - t.Errorf("hash length = %d, want 64", len(rc.restartHash)) - } + c.Require().NoError(rc.err, "unexpected error") + c.Len(rc.restartHash, 64, "hash length = %d, want 64", len(rc.restartHash)) }) t.Run("missing ref ConfigMap surfaces an error", func(t *testing.T) { @@ -136,9 +125,8 @@ func TestRenderEffectiveConfig(t *testing.T) { PostgresConfigRef: &multigresv1alpha1.PostgresConfigRef{Name: "missing", Key: "k"}, }, } - if r.renderEffectiveConfig(context.Background(), shard).err == nil { - t.Error("expected error for missing ConfigMap") - } + assert.NewCollecting(t). + Error(r.renderEffectiveConfig(context.Background(), shard).err, "expected error for missing ConfigMap") }) } @@ -148,6 +136,7 @@ func TestRenderEffectiveConfig(t *testing.T) { // reload-hash — the wiring that makes a removal-only reload-safe change verifiable // (see postgresconfig.TestReloadMarkerDetectsRemoval for the removal semantics). func TestRenderEffectiveConfig_ReloadMarker(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -159,9 +148,7 @@ func TestRenderEffectiveConfig_ReloadMarker(t *testing.T) { Spec: multigresv1alpha1.ShardSpec{PostgresConfig: cfg}, } rc := r.renderEffectiveConfig(context.Background(), shard) - if rc.err != nil { - t.Fatalf("renderEffectiveConfig error: %v", rc.err) - } + c.Require().NoError(rc.err, "renderEffectiveConfig error") return rc } @@ -171,34 +158,32 @@ func TestRenderEffectiveConfig_ReloadMarker(t *testing.T) { // The marker lands in the delivered file and in the RPC expected settings, // with its value equal to the reload-hash. - if !strings.Contains(rc.content, postgresconfig.ReloadMarkerGUC+" = '"+rc.reloadHash+"'") { - t.Errorf("delivered config missing marker line for %q:\n%s", - postgresconfig.ReloadMarkerGUC, rc.content) - } - if got := rc.reloadSettings[postgresconfig.ReloadMarkerGUC]; got != rc.reloadHash { - t.Errorf("reloadSettings[%q] = %q, want reload-hash %q", - postgresconfig.ReloadMarkerGUC, got, rc.reloadHash) - } + c.StrContains( + rc.content, + postgresconfig.ReloadMarkerGUC+" = '"+rc.reloadHash+"'", + "delivered config missing marker line for %q:\n", + postgresconfig.ReloadMarkerGUC, + ) + got := rc.reloadSettings[postgresconfig.ReloadMarkerGUC] + c.Eq( + rc.reloadHash, + got, + "reloadSettings[%q] = %q, want reload-hash", + postgresconfig.ReloadMarkerGUC, + got, + ) // Removing the reload-safe param moves the reload-hash (hence the marker) but // leaves the restart-hash untouched — still a reload, not a pod recreation. rcRemoved := render(nil) - if rc.restartHash != rcRemoved.restartHash { - t.Errorf( - "restart-hash moved on a reload-only removal: %s -> %s", - rc.restartHash, - rcRemoved.restartHash, - ) - } - if rc.reloadHash == rcRemoved.reloadHash { - t.Fatalf( - "reload-hash did not move when the reload-safe param was removed (still %s)", - rc.reloadHash, - ) - } - if rc.reloadSettings[postgresconfig.ReloadMarkerGUC] == rcRemoved.reloadSettings[postgresconfig.ReloadMarkerGUC] { - t.Error("marker did not move on a reload-safe removal; a stale mount would pass the gate") - } + c.Eq(rcRemoved.restartHash, rc.restartHash, "restart-hash moved on a reload-only removal") + c.Require(). + NotEq(rcRemoved.reloadHash, rc.reloadHash, "reload-hash did not move when the reload-safe param was removed (still") + c.NotEq( + rcRemoved.reloadSettings[postgresconfig.ReloadMarkerGUC], + rc.reloadSettings[postgresconfig.ReloadMarkerGUC], + "marker did not move on a reload-safe removal; a stale mount would pass the gate", + ) } var errTest = errTestType("boom") @@ -220,9 +205,8 @@ func TestShardClusterName(t *testing.T) { ShardName: "0", }, } - if got := shardClusterName(shard); got != "mycluster/mydb/mytg/0" { - t.Errorf("shardClusterName = %q, want mycluster/mydb/mytg/0", got) - } + assert.NewCollecting(t). + Eq("mycluster/mydb/mytg/0", shardClusterName(shard), "shardClusterName") }) t.Run("drops empty components", func(t *testing.T) { @@ -232,19 +216,16 @@ func TestShardClusterName(t *testing.T) { }, Spec: multigresv1alpha1.ShardSpec{ShardName: "0"}, } - if got := shardClusterName(shard); got != "c/0" { - t.Errorf("shardClusterName = %q, want c/0", got) - } + assert.NewCollecting(t).Eq("c/0", shardClusterName(shard), "shardClusterName") }) t.Run("falls back to object name without cluster label", func(t *testing.T) { shard := &multigresv1alpha1.Shard{ObjectMeta: metav1.ObjectMeta{Name: "fallback-shard"}} - if got := shardClusterName(shard); got != "fallback-shard" { - t.Errorf("shardClusterName = %q, want fallback-shard", got) - } + assert.NewCollecting(t).Eq("fallback-shard", shardClusterName(shard), "shardClusterName") }) t.Run("rendered config carries the shard cluster_name", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "obj-name", @@ -257,12 +238,12 @@ func TestShardClusterName(t *testing.T) { }, } rendered, _, err := renderPostgresConfig(shard, "") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !strings.Contains(rendered, "cluster_name = 'mycluster/mydb/mytg/0'") { - t.Errorf("rendered config missing shard cluster_name:\n%s", rendered) - } + c.Require().NoError(err, "unexpected error") + c.StrContains( + rendered, + "cluster_name = 'mycluster/mydb/mytg/0'", + "rendered config missing shard cluster_name:\n", + ) }) } @@ -270,6 +251,7 @@ func TestReduceShardResources(t *testing.T) { t.Run( "max mem/cpu, min disk across pools with limit-then-request fallback", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ Spec: multigresv1alpha1.ShardSpec{ Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ @@ -301,23 +283,16 @@ func TestReduceShardResources(t *testing.T) { }, } mem, cpu, disk := reduceShardResources(shard) - if mem != 2*(1<<30) { - t.Errorf("mem = %d, want 2Gi (max, pool b request)", mem) - } - if cpu != 4000 { - t.Errorf("cpu = %d millicores, want 4000 (max, pool b request)", cpu) - } - if disk != 5*(1<<30) { - t.Errorf("disk = %d, want 5Gi (min across pools)", disk) - } + c.Eq(2*(1<<30), mem, "mem") + c.Eq(4000, cpu, "cpu") + c.Eq(5*(1<<30), disk, "disk") }, ) t.Run("no pools returns zeros", func(t *testing.T) { mem, cpu, disk := reduceShardResources(&multigresv1alpha1.Shard{}) - if mem != 0 || cpu != 0 || disk != 0 { - t.Errorf("got (%d, %d, %d), want all zero", mem, cpu, disk) - } + assert.NewCollecting(t). + False(mem != 0 || cpu != 0 || disk != 0, "got (%d, %d, %d), want all zero", mem, cpu, disk) }) t.Run("pools without resources return zeros", func(t *testing.T) { @@ -325,15 +300,15 @@ func TestReduceShardResources(t *testing.T) { Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{"a": {}}, }} mem, cpu, disk := reduceShardResources(shard) - if mem != 0 || cpu != 0 || disk != 0 { - t.Errorf("got (%d, %d, %d), want all zero", mem, cpu, disk) - } + assert.NewCollecting(t). + False(mem != 0 || cpu != 0 || disk != 0, "got (%d, %d, %d), want all zero", mem, cpu, disk) }) } // A shard with no PostgresConfigRef still gets an operator-owned ConfigMap // rendering the baseline — the operator always owns the file. func TestReconcilePostgresConfig_RendersBaselineWithoutRef(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -347,33 +322,30 @@ func TestReconcilePostgresConfig_RendersBaselineWithoutRef(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme} cfg := r.renderEffectiveConfig(context.Background(), shard) - if err := r.reconcilePostgresConfig(context.Background(), shard, cfg); err != nil { - t.Fatalf("reconcilePostgresConfig() error = %v", err) - } + ck.Require(). + NoError(r.reconcilePostgresConfig(context.Background(), shard, cfg), "reconcilePostgresConfig() error =") // The operator ConfigMap must exist with the rendered baseline. got := &corev1.ConfigMap{} - if err := c.Get(context.Background(), client.ObjectKey{ + ck.Require().NoError(c.Get(context.Background(), client.ObjectKey{ Namespace: "default", Name: PostgresConfigMapName("s1"), - }, got); err != nil { - t.Fatalf("operator ConfigMap not created: %v", err) - } + }, got), "operator ConfigMap not created") rendered := got.Data[PostgresConfigMapKey] - if !strings.Contains(rendered, "shared_buffers = 64MB") { - t.Errorf("rendered baseline missing default shared_buffers:\n%s", rendered) - } + ck.StrContains( + rendered, + "shared_buffers = 64MB", + "rendered baseline missing default shared_buffers:\n", + ) // The content hash annotation must be stamped. - if len(shard.Annotations[metadata.AnnotationPostgresConfigHash]) != 64 { - t.Errorf("hash annotation = %q, want a 64-char SHA-256 hex", - shard.Annotations[metadata.AnnotationPostgresConfigHash]) - } + ck.Len(shard.Annotations[metadata.AnnotationPostgresConfigHash], 64, "hash annotation") } // A shard with a PostgresConfigRef renders the baseline plus the ref content // into the operator ConfigMap. func TestReconcilePostgresConfig_MergesRefContent(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -396,42 +368,34 @@ func TestReconcilePostgresConfig_MergesRefContent(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme} cfg := r.renderEffectiveConfig(context.Background(), shard) - if err := r.reconcilePostgresConfig(context.Background(), shard, cfg); err != nil { - t.Fatalf("reconcilePostgresConfig() error = %v", err) - } + ck.Require(). + NoError(r.reconcilePostgresConfig(context.Background(), shard, cfg), "reconcilePostgresConfig() error =") got := &corev1.ConfigMap{} - if err := c.Get(context.Background(), client.ObjectKey{ + ck.Require().NoError(c.Get(context.Background(), client.ObjectKey{ Namespace: "default", Name: PostgresConfigMapName("s1"), - }, got); err != nil { - t.Fatalf("operator ConfigMap not created: %v", err) - } + }, got), "operator ConfigMap not created") rendered := got.Data[PostgresConfigMapKey] // Baseline present, and the ref rendered BEFORE it so the operator's // resource-derived baseline wins last-write-wins: the baseline's // shared_buffers = 64MB must override the ref's 8GB. The deprecated ref must // not override the operator's sizing math — only inline spec.postgresConfig. - if !strings.Contains(rendered, "shared_buffers = 64MB") { - t.Errorf("rendered config missing baseline:\n%s", rendered) - } - if !strings.Contains(rendered, "shared_buffers = '8GB'") { - t.Errorf("rendered config missing ref content:\n%s", rendered) - } - if strings.Index( - rendered, - "shared_buffers = '8GB'", - ) > strings.Index( + ck.StrContains(rendered, "shared_buffers = 64MB", "rendered config missing baseline:\n") + ck.StrContains(rendered, "shared_buffers = '8GB'", "rendered config missing ref content:\n") + ck.LessOrEqual(strings.Index( rendered, "shared_buffers = 64MB", - ) { - t.Errorf("ref content should precede the baseline so the baseline wins:\n%s", rendered) - } + ), strings.Index( + rendered, + "shared_buffers = '8GB'", + ), "ref content should precede the baseline so the baseline wins:\n%s", rendered) } // The inline spec.postgresConfig map is rendered last, so it overrides both the // baseline and any PostgresConfigRef content. func TestReconcilePostgresConfig_InlineMapWins(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -455,23 +419,21 @@ func TestReconcilePostgresConfig_InlineMapWins(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme} cfg := r.renderEffectiveConfig(context.Background(), shard) - if err := r.reconcilePostgresConfig(context.Background(), shard, cfg); err != nil { - t.Fatalf("reconcilePostgresConfig() error = %v", err) - } + ck.Require(). + NoError(r.reconcilePostgresConfig(context.Background(), shard, cfg), "reconcilePostgresConfig() error =") got := &corev1.ConfigMap{} - if err := c.Get(context.Background(), client.ObjectKey{ + ck.Require().NoError(c.Get(context.Background(), client.ObjectKey{ Namespace: "default", Name: PostgresConfigMapName("s1"), - }, got); err != nil { - t.Fatalf("operator ConfigMap not created: %v", err) - } + }, got), "operator ConfigMap not created") rendered := got.Data[PostgresConfigMapKey] - if !strings.Contains(rendered, "work_mem = '64MB'") { - t.Errorf("rendered config missing inline override:\n%s", rendered) - } + ck.StrContains(rendered, "work_mem = '64MB'", "rendered config missing inline override:\n") // The inline map must appear after the ref content so it wins last-write-wins. - if strings.Index(rendered, "work_mem = '64MB'") < strings.Index(rendered, "work_mem = '1MB'") { - t.Errorf("inline map should follow the ref content:\n%s", rendered) - } + ck.GreaterOrEqual( + strings.Index(rendered, "work_mem = '1MB'"), + strings.Index(rendered, "work_mem = '64MB'"), + "inline map should follow the ref content:\n%s", + rendered, + ) } diff --git a/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go b/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go index cd83c8e0..0ced1c72 100644 --- a/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_data_plane_posture_internal_test.go @@ -27,6 +27,8 @@ import ( "github.com/multigres/multigres-operator/pkg/data-handler/poolerclient" "github.com/multigres/multigres-operator/pkg/data-handler/posture" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) type countingPoolerResolver struct { @@ -58,16 +60,11 @@ func (r *countingPoolerResolver) ClientFor( func postureTestScheme(t *testing.T) *runtime.Scheme { t.Helper() + c := assert.NewAborting(t) scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("add shard scheme: %v", err) - } - if err := appsv1.AddToScheme(scheme); err != nil { - t.Fatalf("add apps scheme: %v", err) - } - if err := corev1.AddToScheme(scheme); err != nil { - t.Fatalf("add core scheme: %v", err) - } + c.NoError(multigresv1alpha1.AddToScheme(scheme), "add shard scheme") + c.NoError(appsv1.AddToScheme(scheme), "add apps scheme") + c.NoError(corev1.AddToScheme(scheme), "add core scheme") return scheme } @@ -139,24 +136,24 @@ func postureTestStore(t *testing.T) (topoclient.Store, topoclient.ComponentID) { factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), ) id := &clustermetadata.ID{Cell: "cell1", Name: "pooler-0"} - if err := store.RegisterMultipooler(context.Background(), &clustermetadata.Multipooler{ - Id: id, - Hostname: "pooler-0", - ShardKey: &clustermetadata.ShardKey{ - Database: "database", - TableGroup: "table-group", - Shard: "0", - }, - RoutingState: &clustermetadata.RoutingState{ - Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, - }, - }, false); err != nil { - t.Fatalf("register pooler: %v", err) - } + assert.NewAborting(t). + NoError(store.RegisterMultipooler(context.Background(), &clustermetadata.Multipooler{ + Id: id, + Hostname: "pooler-0", + ShardKey: &clustermetadata.ShardKey{ + Database: "database", + TableGroup: "table-group", + Shard: "0", + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + }, + }, false), "register pooler") return store, topoclient.ComponentIDString(id) } func TestUpdateStatusPublishesPostureFailureAndPhaseTogether(t *testing.T) { + ck := assert.NewCollecting(t) shard := postureTestShard() shard.Status.PodPostures = map[string]string{"pooler-0": "PRIMARY"} shard.Status.Conditions = []metav1.Condition{{ @@ -168,20 +165,19 @@ func TestUpdateStatusPublishesPostureFailureAndPhaseTogether(t *testing.T) { }} r, c := postureTestReconciler(t, shard, nil) - if err := r.updateStatus(t.Context(), shard, renderedConfig{}); err != nil { - t.Fatalf("updateStatus() error = %v", err) - } + ck.Require(). + NoError(r.updateStatus(t.Context(), shard, renderedConfig{}), "updateStatus() error =") got := &multigresv1alpha1.Shard{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(shard), got); err != nil { - t.Fatalf("get updated shard: %v", err) - } - if got.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("phase = %q, want %q", got.Status.Phase, multigresv1alpha1.PhaseDegraded) - } - if got.Status.PodPostures["pooler-0"] != "PRIMARY" { - t.Errorf("podPostures = %v, want pooler-0 PRIMARY", got.Status.PodPostures) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(shard), got), "get updated shard") + ck.Eq(multigresv1alpha1.PhaseDegraded, got.Status.Phase, "phase") + ck.Eq( + "PRIMARY", + got.Status.PodPostures["pooler-0"], + "podPostures = %v, want pooler-0 PRIMARY", + got.Status.PodPostures, + ) postureFailurePersisted := false for _, condition := range got.Status.Conditions { if condition.Type == posture.ConditionConsistent && @@ -196,6 +192,7 @@ func TestUpdateStatusPublishesPostureFailureAndPhaseTogether(t *testing.T) { } func TestUpdateStatusKeepsIncompletePostureOutOfHealthy(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() shard.Status.Conditions = []metav1.Condition{{ Type: posture.ConditionConsistent, @@ -206,15 +203,13 @@ func TestUpdateStatusKeepsIncompletePostureOutOfHealthy(t *testing.T) { }} r, _ := postureTestReconciler(t, shard, nil) - if err := r.updateStatus(t.Context(), shard, renderedConfig{}); err != nil { - t.Fatalf("updateStatus() error = %v", err) - } - if shard.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf("phase = %q, want %q", shard.Status.Phase, multigresv1alpha1.PhaseProgressing) - } + c.Require(). + NoError(r.updateStatus(t.Context(), shard, renderedConfig{}), "updateStatus() error =") + c.Eq(multigresv1alpha1.PhaseProgressing, shard.Status.Phase, "phase") } func TestReconcilePostureDebouncesFirstInconsistency(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() store, poolerID := postureTestStore(t) defer func() { _ = store.Close() }() @@ -228,36 +223,34 @@ func TestReconcilePostureDebouncesFirstInconsistency(t *testing.T) { r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) retryAfter, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("first reconcilePosture() error = %v", err) - } - if retryAfter != postureDebounceRequeueDelay { - t.Error("first inconsistent posture observation did not request a requeue") - } - if got := withDataPlaneRequeue( + c.Require().NoError(err, "first reconcilePosture() error =") + c.Eq( + postureDebounceRequeueDelay, + retryAfter, + "first inconsistent posture observation did not request a requeue", + ) + c.Eq(postureDebounceRequeueDelay, withDataPlaneRequeue( ctrl.Result{}, retryAfter, false, - ).RequeueAfter; got != postureDebounceRequeueDelay { - t.Errorf("first requeue delay = %v, want %v", got, postureDebounceRequeueDelay) - } + ).RequeueAfter, "first requeue delay") if conditionIsFalse(shard.Status.Conditions, posture.ConditionConsistent) { t.Errorf("conditions = %#v, want no failure on first observation", shard.Status.Conditions) } retryAfter, err = r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("second reconcilePosture() error = %v", err) - } + c.Require().NoError(err, "second reconcilePosture() error =") // The mismatch is now accepted into status (PostureConsistent=False), but // this mock pooler never reports IsInitialized/PostgresReady, so it has // also never reached posture readiness. An accepted-but-not-ready shard // must keep requesting a requeue: nothing but this controller's own // backoff will ever look again, since neither a status recovery in // topology nor a role fix changes a Kubernetes object. - if retryAfter <= 0 { - t.Error("second inconsistent posture observation (still not ready) requested no requeue") - } + c.Greater( + 0, + retryAfter, + "second inconsistent posture observation (still not ready) requested no requeue", + ) if !conditionIsFalse(shard.Status.Conditions, posture.ConditionConsistent) { t.Errorf( "conditions = %#v, want posture failure on second observation", @@ -267,6 +260,7 @@ func TestReconcilePostureDebouncesFirstInconsistency(t *testing.T) { } func TestReconcilePostureDebouncesFirstIncompleteObservation(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() shard.Status.Conditions = []metav1.Condition{{ Type: posture.ConditionConsistent, @@ -283,19 +277,17 @@ func TestReconcilePostureDebouncesFirstIncompleteObservation(t *testing.T) { r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) retryAfter, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("first reconcilePosture() error = %v", err) - } - if retryAfter != postureDebounceRequeueDelay { - t.Error("first incomplete posture observation did not request a requeue") - } - if got := withDataPlaneRequeue( + c.Require().NoError(err, "first reconcilePosture() error =") + c.Eq( + postureDebounceRequeueDelay, + retryAfter, + "first incomplete posture observation did not request a requeue", + ) + c.Eq(postureDebounceRequeueDelay, withDataPlaneRequeue( ctrl.Result{}, retryAfter, false, - ).RequeueAfter; got != postureDebounceRequeueDelay { - t.Errorf("first requeue delay = %v, want %v", got, postureDebounceRequeueDelay) - } + ).RequeueAfter, "first requeue delay") if !conditionIsTrue(shard.Status.Conditions, posture.ConditionConsistent) { t.Errorf( "conditions = %#v, want prior consistent condition preserved on first blip", @@ -304,30 +296,26 @@ func TestReconcilePostureDebouncesFirstIncompleteObservation(t *testing.T) { } retryAfter, err = r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("second reconcilePosture() error = %v", err) - } + c.Require().NoError(err, "second reconcilePosture() error =") // Accepted into status (Unknown/ObservationIncomplete), but a pooler whose // Status RPC errors has also never reached posture readiness, so this must // still requeue rather than strand the pod until the 10h resync. - if retryAfter <= 0 { - t.Error("second incomplete posture observation (still not ready) requested no requeue") - } + c.Greater( + 0, + retryAfter, + "second incomplete posture observation (still not ready) requested no requeue", + ) for _, condition := range shard.Status.Conditions { if condition.Type != posture.ConditionConsistent { continue } - if condition.Status != metav1.ConditionUnknown || - condition.Reason != "ObservationIncomplete" { - t.Errorf( - "condition = %#v, want Unknown/ObservationIncomplete on second blip", - condition, - ) - } + c.False(condition.Status != metav1.ConditionUnknown || + condition.Reason != "ObservationIncomplete", "condition = %#v, want Unknown/ObservationIncomplete on second blip", condition) } } func TestReconcileDataPlaneRequeuesFirstPostureStrike(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() store, poolerID := postureTestStore(t) pod := postureTestPod() @@ -356,21 +344,16 @@ func TestReconcileDataPlaneRequeuesFirstPostureStrike(t *testing.T) { restartHash: "restart", reloadHash: "reload", }) - if err != nil { - t.Fatalf("reconcileDataPlane() error = %v", err) - } - if result.RequeueAfter != postureDebounceRequeueDelay { - t.Errorf("requeue delay = %v, want %v", result.RequeueAfter, postureDebounceRequeueDelay) - } + c.Require().NoError(err, "reconcileDataPlane() error =") + c.Eq(postureDebounceRequeueDelay, result.RequeueAfter, "requeue delay") if !callLogHas(rpc.GetCallLog(), "ReloadConfig") { t.Errorf("reload phase did not run, call log = %v", rpc.GetCallLog()) } - if resolver.calls != 1 { - t.Errorf("ClientFor calls = %d, want exactly 1 per reconcile", resolver.calls) - } + c.Eq(1, resolver.calls, "ClientFor calls") } func TestReconcileDataPlaneContinuesWithoutPoolerClient(t *testing.T) { + ck := assert.NewCollecting(t) shard := postureTestShard() lastBackup := metav1.NewTime(time.Now().Add(-time.Hour).Truncate(time.Second)) shard.Status.LastBackupTime = &lastBackup @@ -406,57 +389,36 @@ func TestReconcileDataPlaneContinuesWithoutPoolerClient(t *testing.T) { } result, err := r.reconcileDataPlane(t.Context(), shard, renderedConfig{}) - if err != nil { - t.Fatalf("reconcileDataPlane() error = %v", err) - } - if result.RequeueAfter != poolerClientRetryDelay { - t.Errorf("requeue delay = %v, want %v", result.RequeueAfter, poolerClientRetryDelay) - } - if resolver.calls != 1 { - t.Errorf("ClientFor calls = %d, want 1", resolver.calls) - } + ck.Require().NoError(err, "reconcileDataPlane() error =") + ck.Eq(poolerClientRetryDelay, result.RequeueAfter, "requeue delay") + ck.Eq(1, resolver.calls, "ClientFor calls") got := &multigresv1alpha1.Shard{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(shard), got); err != nil { - t.Fatalf("get updated shard: %v", err) - } - if got.Status.PodRoles["pooler-0"] != "REPLICA" { - t.Errorf("podRoles = %v, want topology-derived pooler-0 REPLICA", got.Status.PodRoles) - } - if got.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf("phase = %q, want %q", got.Status.Phase, multigresv1alpha1.PhaseProgressing) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(shard), got), "get updated shard") + ck.Eq( + "REPLICA", + got.Status.PodRoles["pooler-0"], + "podRoles = %v, want topology-derived pooler-0 REPLICA", + got.Status.PodRoles, + ) + ck.Eq(multigresv1alpha1.PhaseProgressing, got.Status.Phase, "phase") postureCondition := findPostureCondition( got.Status.Conditions, posture.ConditionConsistent, ) - if postureCondition == nil { - t.Fatalf( - "conditions = %#v, want %s condition", - got.Status.Conditions, - posture.ConditionConsistent, - ) - } - if postureCondition.Status != metav1.ConditionUnknown || + ck.Require(). + NotNil(postureCondition, "conditions = %#v, want %s condition", got.Status.Conditions, posture.ConditionConsistent) + ck.False(postureCondition.Status != metav1.ConditionUnknown || postureCondition.Reason != "PoolerClientUnavailable" || - !strings.Contains(postureCondition.Message, "certificate not issued") { - t.Errorf( - "condition = %#v, want Unknown/PoolerClientUnavailable with resolver error", - postureCondition, - ) - } - if calls := rpc.GetCallLog(); len(calls) != 0 { - t.Errorf("RPC phases ran with resolver error, call log = %v", calls) - } + !strings.Contains( + postureCondition.Message, + "certificate not issued", + ), "condition = %#v, want Unknown/PoolerClientUnavailable with resolver error", postureCondition) + ck.Empty(rpc.GetCallLog(), "RPC phases ran with resolver error, call log =") backupCondition := findPostureCondition(got.Status.Conditions, backuphealth.ConditionHealthy) - if backupCondition == nil || backupCondition.Status != metav1.ConditionUnknown || - backupCondition.Reason != reasonBackupCheckUnavailable { - t.Errorf( - "backup condition = %#v, want Unknown/%s", - backupCondition, - reasonBackupCheckUnavailable, - ) - } + ck.False(backupCondition == nil || backupCondition.Status != metav1.ConditionUnknown || + backupCondition.Reason != reasonBackupCheckUnavailable, "backup condition = %#v, want Unknown/%s", backupCondition, reasonBackupCheckUnavailable) if got.Status.LastBackupTime == nil || !got.Status.LastBackupTime.Equal(&lastBackup) || got.Status.LastBackupType != "full" { t.Errorf( @@ -467,9 +429,10 @@ func TestReconcileDataPlaneContinuesWithoutPoolerClient(t *testing.T) { } recorder := r.Recorder.(*record.FakeRecorder) - if !containsEventWithReason(drainEvents(recorder), "PoolerClientUnavailable") { - t.Error("expected PoolerClientUnavailable event") - } + ck.True( + containsEventWithReason(drainEvents(recorder), "PoolerClientUnavailable"), + "expected PoolerClientUnavailable event", + ) secondStore, _ := postureTestStore(t) r.CreateTopoStore = func(*multigresv1alpha1.Shard) (topoclient.Store, error) { @@ -478,12 +441,12 @@ func TestReconcileDataPlaneContinuesWithoutPoolerClient(t *testing.T) { if _, err := r.reconcileDataPlane(t.Context(), got, renderedConfig{}); err != nil { t.Fatalf("second reconcileDataPlane() error = %v", err) } - if events := drainEvents(recorder); containsEventWithReason(events, "PoolerClientUnavailable") { - t.Errorf( - "repeated resolver failure re-emitted PoolerClientUnavailable, events = %v", - events, - ) - } + events := drainEvents(recorder) + ck.False( + containsEventWithReason(events, "PoolerClientUnavailable"), + "repeated resolver failure re-emitted PoolerClientUnavailable, events = %v", + events, + ) } func TestSetBackupUnknownPreservesConfirmedFailure(t *testing.T) { @@ -496,13 +459,12 @@ func TestSetBackupUnknownPreservesConfirmedFailure(t *testing.T) { }} setBackupUnknownUnlessFalse(shard, "RPC unavailable") condition := findPostureCondition(shard.Status.Conditions, backuphealth.ConditionHealthy) - if condition == nil || condition.Status != metav1.ConditionFalse || - condition.Reason != "BackupStale" { - t.Errorf("backup condition = %#v, want preserved False/BackupStale", condition) - } + assert.NewCollecting(t).False(condition == nil || condition.Status != metav1.ConditionFalse || + condition.Reason != "BackupStale", "backup condition = %#v, want preserved False/BackupStale", condition) } func TestReconcileDataPlaneResolverFailurePreservesConfirmedPostureFailure(t *testing.T) { + ck := assert.NewCollecting(t) shard := postureTestShard() shard.Status.Conditions = []metav1.Condition{{ Type: posture.ConditionConsistent, @@ -525,36 +487,24 @@ func TestReconcileDataPlaneResolverFailurePreservesConfirmedPostureFailure(t *te // Simulate one unsettled observation immediately before the transport // outage. Resolver failure must break that sequence. - if strikes := r.recordPostureObservation(shard, true); strikes != 1 { - t.Fatalf("initial strikes = %d, want 1", strikes) - } + ck.Require().Eq(1, r.recordPostureObservation(shard, true), "initial strikes") result, err := r.reconcileDataPlane(t.Context(), shard, renderedConfig{}) - if err != nil { - t.Fatalf("reconcileDataPlane() error = %v", err) - } - if result.RequeueAfter != poolerClientRetryDelay { - t.Errorf("requeue delay = %v, want %v", result.RequeueAfter, poolerClientRetryDelay) - } + ck.Require().NoError(err, "reconcileDataPlane() error =") + ck.Eq(poolerClientRetryDelay, result.RequeueAfter, "requeue delay") got := &multigresv1alpha1.Shard{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(shard), got); err != nil { - t.Fatalf("get updated shard: %v", err) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(shard), got), "get updated shard") condition := findPostureCondition(got.Status.Conditions, posture.ConditionConsistent) - if condition == nil || condition.Status != metav1.ConditionFalse || - condition.Reason != "MultiplePrimaries" { - t.Errorf("condition = %#v, want preserved False/MultiplePrimaries", condition) - } - if got.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("phase = %q, want %q", got.Status.Phase, multigresv1alpha1.PhaseDegraded) - } - if strikes := r.recordPostureObservation(shard, true); strikes != 1 { - t.Errorf("first post-outage strikes = %d, want 1", strikes) - } + ck.False(condition == nil || condition.Status != metav1.ConditionFalse || + condition.Reason != "MultiplePrimaries", "condition = %#v, want preserved False/MultiplePrimaries", condition) + ck.Eq(multigresv1alpha1.PhaseDegraded, got.Status.Phase, "phase") + ck.Eq(1, r.recordPostureObservation(shard, true), "first post-outage strikes") } func TestReconcileDataPlaneMarksBackupUnknownOnRPCFailure(t *testing.T) { + ck := assert.NewCollecting(t) shard := postureTestShard() lastBackup := metav1.NewTime(time.Now().Add(-time.Hour).Truncate(time.Second)) shard.Status.LastBackupTime = &lastBackup @@ -596,18 +546,14 @@ func TestReconcileDataPlaneMarksBackupUnknownOnRPCFailure(t *testing.T) { return store, nil } - if _, err := r.reconcileDataPlane(t.Context(), shard, renderedConfig{}); err != nil { - t.Fatalf("reconcileDataPlane() error = %v", err) - } + _, err := r.reconcileDataPlane(t.Context(), shard, renderedConfig{}) + assert.NewAborting(t).NoError(err, "reconcileDataPlane() error =") got := &multigresv1alpha1.Shard{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(shard), got); err != nil { - t.Fatalf("get updated shard: %v", err) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(shard), got), "get updated shard") condition := findPostureCondition(got.Status.Conditions, backuphealth.ConditionHealthy) - if condition == nil || condition.Status != metav1.ConditionUnknown || - condition.Reason != reasonBackupCheckUnavailable { - t.Errorf("backup condition = %#v, want Unknown/%s", condition, reasonBackupCheckUnavailable) - } + ck.False(condition == nil || condition.Status != metav1.ConditionUnknown || + condition.Reason != reasonBackupCheckUnavailable, "backup condition = %#v, want Unknown/%s", condition, reasonBackupCheckUnavailable) if got.Status.LastBackupTime == nil || !got.Status.LastBackupTime.Equal(&lastBackup) || got.Status.LastBackupType != "full" { t.Errorf( @@ -633,6 +579,7 @@ func TestReconcilePostureClearsStaleUnknownReasonDuringDebounce(t *testing.T) { reasonAwaitingPoolerRegistration, } { t.Run(previousReason, func(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() shard.Status.Conditions = []metav1.Condition{{ Type: posture.ConditionConsistent, @@ -644,22 +591,18 @@ func TestReconcilePostureClearsStaleUnknownReasonDuringDebounce(t *testing.T) { r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) retryAfter, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("reconcilePosture() error = %v", err) - } - if retryAfter != postureDebounceRequeueDelay { - t.Fatal("first recovered unsettled observation did not request debounce requeue") - } + c.Require().NoError(err, "reconcilePosture() error =") + c.Require(). + Eq(postureDebounceRequeueDelay, retryAfter, "first recovered unsettled observation did not request debounce requeue") condition := findPostureCondition(shard.Status.Conditions, posture.ConditionConsistent) - if condition == nil || condition.Status != metav1.ConditionUnknown || - condition.Reason != reasonObservationPending { - t.Errorf("condition = %#v, want Unknown/ObservationPending", condition) - } + c.False(condition == nil || condition.Status != metav1.ConditionUnknown || + condition.Reason != reasonObservationPending, "condition = %#v, want Unknown/ObservationPending", condition) }) } } func TestReconcilePostureRequeuesWhileTopologyHasNoPoolers(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() shard.Status.Conditions = []metav1.Condition{{ Type: posture.ConditionConsistent, @@ -680,17 +623,12 @@ func TestReconcilePostureRequeuesWhileTopologyHasNoPoolers(t *testing.T) { rpc := rpcclient.NewFakeClient() r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) retryAfter, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("reconcilePosture() error = %v", err) - } - if retryAfter != poolerRegistrationRetryDelay { - t.Fatal("empty topology did not request a bootstrap requeue") - } + c.Require().NoError(err, "reconcilePosture() error =") + c.Require(). + Eq(poolerRegistrationRetryDelay, retryAfter, "empty topology did not request a bootstrap requeue") condition := findPostureCondition(shard.Status.Conditions, posture.ConditionConsistent) - if condition == nil || condition.Status != metav1.ConditionUnknown || - condition.Reason != "AwaitingPoolerRegistration" { - t.Errorf("condition = %#v, want Unknown/AwaitingPoolerRegistration", condition) - } + c.False(condition == nil || condition.Status != metav1.ConditionUnknown || + condition.Reason != "AwaitingPoolerRegistration", "condition = %#v, want Unknown/AwaitingPoolerRegistration", condition) } func TestWithDataPlaneRequeueUsesEarliestDelay(t *testing.T) { @@ -727,9 +665,7 @@ func TestWithDataPlaneRequeueUsesEarliestDelay(t *testing.T) { tt.postureRetryAfter, tt.poolerClientUnavailable, ) - if result.RequeueAfter != tt.want { - t.Errorf("requeue delay = %v, want %v", result.RequeueAfter, tt.want) - } + assert.NewCollecting(t).Eq(tt.want, result.RequeueAfter, "requeue delay") }) } } diff --git a/pkg/resource-handler/controller/shard/reconcile_deletion_internal_test.go b/pkg/resource-handler/controller/shard/reconcile_deletion_internal_test.go index 30c66999..e20fd3ac 100644 --- a/pkg/resource-handler/controller/shard/reconcile_deletion_internal_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_deletion_internal_test.go @@ -18,6 +18,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestHandleDeletion_ErrorPaths(t *testing.T) { @@ -68,9 +70,7 @@ func TestHandleDeletion_ErrorPaths(t *testing.T) { } _, err := r.handleDeletion(context.Background(), shard) - if err == nil { - t.Error("expected error on Pod list failure") - } + assert.NewCollecting(t).Error(err, "expected error on Pod list failure") }) t.Run("error listing deployments", func(t *testing.T) { @@ -92,12 +92,11 @@ func TestHandleDeletion_ErrorPaths(t *testing.T) { } _, err := r.handleDeletion(context.Background(), shard) - if err == nil { - t.Error("expected error on Deployment list failure") - } + assert.NewCollecting(t).Error(err, "expected error on Deployment list failure") }) t.Run("deletion with existing deployments deletes them", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() deploy := &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ @@ -120,9 +119,7 @@ func TestHandleDeletion_ErrorPaths(t *testing.T) { } _, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } + ck.Require().NoError(err, "handleDeletion returned error") got := &appsv1.Deployment{} err = c.Get( @@ -130,8 +127,10 @@ func TestHandleDeletion_ErrorPaths(t *testing.T) { types.NamespacedName{Name: "mo-deploy", Namespace: "default"}, got, ) - if !errors.IsNotFound(err) { - t.Errorf("deployment should have been deleted, but Get returned: %v", err) - } + ck.True( + errors.IsNotFound(err), + "deployment should have been deleted, but Get returned: %v", + err, + ) }) } diff --git a/pkg/resource-handler/controller/shard/reconcile_deletion_test.go b/pkg/resource-handler/controller/shard/reconcile_deletion_test.go index d56eb92f..8596a953 100644 --- a/pkg/resource-handler/controller/shard/reconcile_deletion_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_deletion_test.go @@ -17,6 +17,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestHandleDeletion(t *testing.T) { @@ -73,6 +75,7 @@ func TestHandleDeletion(t *testing.T) { t.Run("scheduled pod gets deleted during shard deletion", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) pod := makePod("pod-0", true, "") shard := baseShard.DeepCopy() @@ -91,12 +94,12 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != podTerminationRequeueDelay { - t.Errorf("Expected requeue while pod still terminating, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq( + podTerminationRequeueDelay, + result.RequeueAfter, + "Expected requeue while pod still terminating, got", + ) // Pod should be deleted updatedPod := &corev1.Pod{} @@ -105,13 +108,12 @@ func TestHandleDeletion(t *testing.T) { types.NamespacedName{Name: "pod-0", Namespace: "default"}, updatedPod, ) - if err == nil { - t.Error("Expected pod to be deleted") - } + ck.Error(err, "Expected pod to be deleted") }) t.Run("unscheduled pod gets deleted during shard deletion", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) pod := makePod("pod-unsched", false, "") shard := baseShard.DeepCopy() @@ -130,15 +132,12 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != podTerminationRequeueDelay { - t.Errorf( - "Expected requeue while unscheduled pod terminating, got %v", - result.RequeueAfter, - ) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq( + podTerminationRequeueDelay, + result.RequeueAfter, + "Expected requeue while unscheduled pod terminating, got", + ) // Pod should be deleted updatedPod := &corev1.Pod{} @@ -147,13 +146,12 @@ func TestHandleDeletion(t *testing.T) { types.NamespacedName{Name: "pod-unsched", Namespace: "default"}, updatedPod, ) - if err == nil { - t.Error("Expected pod to be deleted") - } + ck.Error(err, "Expected pod to be deleted") }) t.Run("ready-for-deletion pod gets deleted", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) pod := makePod("pod-rfd", true, metadata.DrainStateReadyForDeletion) shard := baseShard.DeepCopy() @@ -172,12 +170,12 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != podTerminationRequeueDelay { - t.Errorf("Expected requeue while drained pod terminating, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq( + podTerminationRequeueDelay, + result.RequeueAfter, + "Expected requeue while drained pod terminating, got", + ) // Pod should be deleted updatedPod := &corev1.Pod{} @@ -186,13 +184,12 @@ func TestHandleDeletion(t *testing.T) { types.NamespacedName{Name: "pod-rfd", Namespace: "default"}, updatedPod, ) - if err == nil { - t.Error("Expected pod to be deleted") - } + ck.Error(err, "Expected pod to be deleted") }) t.Run("mid-drain pod gets deleted directly", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) pod := makePod("pod-draining", true, metadata.DrainStateDraining) pod.Annotations[metadata.AnnotationDrainRequestedAt] = time.Now().UTC().Format(time.RFC3339) @@ -212,12 +209,12 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != podTerminationRequeueDelay { - t.Errorf("Expected requeue while pod still terminating, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq( + podTerminationRequeueDelay, + result.RequeueAfter, + "Expected requeue while pod still terminating, got", + ) // Pod should be deleted directly, no drain wait updatedPod := &corev1.Pod{} @@ -226,13 +223,12 @@ func TestHandleDeletion(t *testing.T) { types.NamespacedName{Name: "pod-draining", Namespace: "default"}, updatedPod, ) - if err == nil { - t.Error("Expected pod to be deleted") - } + ck.Error(err, "Expected pod to be deleted") }) t.Run("expired drain pod gets deleted directly", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) pod := makePod("pod-timeout", true, metadata.DrainStateDraining) pod.Annotations[metadata.AnnotationDrainRequestedAt] = time.Now(). @@ -255,12 +251,12 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != podTerminationRequeueDelay { - t.Errorf("Expected requeue while pod still terminating, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq( + podTerminationRequeueDelay, + result.RequeueAfter, + "Expected requeue while pod still terminating, got", + ) // Pod should be deleted updatedPod := &corev1.Pod{} @@ -269,13 +265,12 @@ func TestHandleDeletion(t *testing.T) { types.NamespacedName{Name: "pod-timeout", Namespace: "default"}, updatedPod, ) - if err == nil { - t.Error("Expected pod to be deleted") - } + ck.Error(err, "Expected pod to be deleted") }) t.Run("cleanup completes when no pods exist", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() @@ -293,16 +288,13 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != 0 { - t.Error("Expected no requeue when no pods exist") - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq(0, result.RequeueAfter, "Expected no requeue when no pods exist") }) t.Run("PVC orphaned only after pods are gone", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() shard.Spec.PVCDeletionPolicy = &multigresv1alpha1.PVCDeletionPolicy{ @@ -332,42 +324,46 @@ func TestHandleDeletion(t *testing.T) { // First pass: pod still present -> requeue, PVC must NOT be orphaned yet. result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != podTerminationRequeueDelay { - t.Errorf("Expected requeue while pod alive, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq( + podTerminationRequeueDelay, + result.RequeueAfter, + "Expected requeue while pod alive, got", + ) got := &corev1.PersistentVolumeClaim{} - if err := c.Get(context.Background(), - types.NamespacedName{Name: "data-pvc-0", Namespace: "default"}, got); err != nil { - t.Fatalf("failed to get PVC: %v", err) - } + ck.Require().NoError(c.Get( + context.Background(), + types.NamespacedName{ + Name: "data-pvc-0", + Namespace: "default", + }, + got, + ), "failed to get PVC") if _, ok := got.Labels[metadata.LabelOrphan]; ok { t.Error("PVC must not be orphaned while a pod still references it") } // Second pass: pod deleted by the first pass -> PVC orphaned now. result, err = r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion (2nd pass) returned error: %v", err) - } - if result.RequeueAfter != 0 { - t.Errorf("Expected no requeue once pods are gone, got %v", result.RequeueAfter) - } - if err := c.Get(context.Background(), - types.NamespacedName{Name: "data-pvc-0", Namespace: "default"}, got); err != nil { - t.Fatalf("failed to get PVC after cleanup: %v", err) - } - if _, ok := got.Labels[metadata.LabelOrphan]; !ok { - t.Error("PVC should be orphaned once all pods are gone") - } + ck.Require().NoError(err, "handleDeletion (2nd pass) returned error") + ck.Eq(0, result.RequeueAfter, "Expected no requeue once pods are gone, got") + ck.Require().NoError(c.Get( + context.Background(), + types.NamespacedName{ + Name: "data-pvc-0", + Namespace: "default", + }, + got, + ), "failed to get PVC after cleanup") + _, ok := got.Labels[metadata.LabelOrphan] + ck.True(ok, "PVC should be orphaned once all pods are gone") }) t.Run( "MultigresCluster being deleted hard-deletes PVC instead of orphaning", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() shard.Spec.PVCDeletionPolicy = &multigresv1alpha1.PVCDeletionPolicy{ @@ -403,26 +399,22 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != 0 { - t.Errorf("Expected no requeue once pods are gone, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq(0, result.RequeueAfter, "Expected no requeue once pods are gone, got") got := &corev1.PersistentVolumeClaim{} err = c.Get(context.Background(), types.NamespacedName{Name: "data-pvc-churn", Namespace: "default"}, got) - if !apierrors.IsNotFound(err) { - t.Errorf( - "PVC should be hard-deleted while cluster is being deleted, got err=%v", - err, - ) - } + ck.True( + apierrors.IsNotFound(err), + "PVC should be hard-deleted while cluster is being deleted, got err=%v", + err, + ) }, ) t.Run("pod stuck terminating past timeout does not block PVC cleanup", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() shard.Spec.PVCDeletionPolicy = &multigresv1alpha1.PVCDeletionPolicy{ @@ -454,20 +446,19 @@ func TestHandleDeletion(t *testing.T) { } result, err := r.handleDeletion(context.Background(), shard) - if err != nil { - t.Fatalf("handleDeletion returned error: %v", err) - } - if result.RequeueAfter != 0 { - t.Errorf("Expected no requeue for pod stuck past timeout, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "handleDeletion returned error") + ck.Eq(0, result.RequeueAfter, "Expected no requeue for pod stuck past timeout, got") got := &corev1.PersistentVolumeClaim{} - if err := c.Get(context.Background(), - types.NamespacedName{Name: "data-pvc-stuck", Namespace: "default"}, got); err != nil { - t.Fatalf("failed to get PVC: %v", err) - } - if _, ok := got.Labels[metadata.LabelOrphan]; !ok { - t.Error("PVC should be orphaned even when a pod is stuck terminating past the timeout") - } + ck.Require().NoError(c.Get( + context.Background(), + types.NamespacedName{ + Name: "data-pvc-stuck", + Namespace: "default", + }, + got, + ), "failed to get PVC") + _, ok := got.Labels[metadata.LabelOrphan] + ck.True(ok, "PVC should be orphaned even when a pod is stuck terminating past the timeout") }) } @@ -505,6 +496,7 @@ func TestHandlePendingDeletion(t *testing.T) { t.Run("No pods — sets ReadyForDeletion immediately", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() c := fake.NewClientBuilder(). @@ -521,19 +513,13 @@ func TestHandlePendingDeletion(t *testing.T) { } result, err := r.handlePendingDeletion(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result.RequeueAfter != 0 { - t.Error("Expected no requeue when no pods exist") - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(0, result.RequeueAfter, "Expected no requeue when no pods exist") updated := &multigresv1alpha1.Shard{} - if err := c.Get(t.Context(), types.NamespacedName{ + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{ Name: shard.Name, Namespace: shard.Namespace, - }, updated); err != nil { - t.Fatalf("failed to get shard: %v", err) - } + }, updated), "failed to get shard") found := false for _, cond := range updated.Status.Conditions { @@ -542,13 +528,12 @@ func TestHandlePendingDeletion(t *testing.T) { found = true } } - if !found { - t.Error("Expected ReadyForDeletion condition to be True") - } + ck.True(found, "Expected ReadyForDeletion condition to be True") }) t.Run("Pods without drain — initiates drain and requeues", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() pod := &corev1.Pod{ @@ -576,29 +561,24 @@ func TestHandlePendingDeletion(t *testing.T) { } result, err := r.handlePendingDeletion(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result.RequeueAfter == 0 { - t.Error("Expected requeue when pods exist") - } + ck.Require().NoError(err, "unexpected error") + ck.NotEq(0, result.RequeueAfter, "Expected requeue when pods exist") // Verify drain was initiated. updatedPod := &corev1.Pod{} - if err := c.Get(t.Context(), types.NamespacedName{ + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{ Name: pod.Name, Namespace: pod.Namespace, - }, updatedPod); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updatedPod.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf("Expected drain state %q, got %q", - metadata.DrainStateRequested, - updatedPod.Annotations[metadata.AnnotationDrainState]) - } + }, updatedPod), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updatedPod.Annotations[metadata.AnnotationDrainState], + "Expected drain state", + ) }) t.Run("All pods ready-for-deletion — sets ReadyForDeletion", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() pod := &corev1.Pod{ @@ -629,27 +609,22 @@ func TestHandlePendingDeletion(t *testing.T) { } result, err := r.handlePendingDeletion(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // The pod was deleted, so we need to requeue to check again. - if result.RequeueAfter == 0 { - t.Error("Expected requeue after deleting drained pods") - } + ck.NotEq(0, result.RequeueAfter, "Expected requeue after deleting drained pods") // Verify the pod was deleted. updatedPod := &corev1.Pod{} err = c.Get(t.Context(), types.NamespacedName{ Name: pod.Name, Namespace: pod.Namespace, }, updatedPod) - if err == nil { - t.Error("Expected pod to be deleted") - } + ck.Error(err, "Expected pod to be deleted") }) t.Run("Mix of draining and undrained pods — requeues", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() drainingPod := &corev1.Pod{ @@ -691,24 +666,18 @@ func TestHandlePendingDeletion(t *testing.T) { } result, err := r.handlePendingDeletion(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result.RequeueAfter == 0 { - t.Error("Expected requeue when pods are still draining") - } + ck.Require().NoError(err, "unexpected error") + ck.NotEq(0, result.RequeueAfter, "Expected requeue when pods are still draining") // Verify the undrained pod now has drain state set. updatedPod := &corev1.Pod{} - if err := c.Get(t.Context(), types.NamespacedName{ + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{ Name: undrainedPod.Name, Namespace: undrainedPod.Namespace, - }, updatedPod); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updatedPod.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf("Expected drain state %q, got %q", - metadata.DrainStateRequested, - updatedPod.Annotations[metadata.AnnotationDrainState]) - } + }, updatedPod), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updatedPod.Annotations[metadata.AnnotationDrainState], + "Expected drain state", + ) }) } diff --git a/pkg/resource-handler/controller/shard/reconcile_quarantine_internal_test.go b/pkg/resource-handler/controller/shard/reconcile_quarantine_internal_test.go index 2237b368..ca6d63f3 100644 --- a/pkg/resource-handler/controller/shard/reconcile_quarantine_internal_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_quarantine_internal_test.go @@ -21,6 +21,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/backuphealth" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) // markBackupHealthy sets the shard's backup-health condition to True so @@ -132,9 +134,8 @@ func qrStoreWithQuarantinedReason( Reason: reason, } } - if err := store.RegisterMultipooler(context.Background(), mp, false); err != nil { - t.Fatalf("register pooler %s: %v", name, err) - } + assert.NewAborting(t). + NoError(store.RegisterMultipooler(context.Background(), mp, false), "register pooler %s", name) } return store } @@ -169,6 +170,7 @@ func TestReconcileQuarantineRemediation(t *testing.T) { old := time.Now().Add(-1 * time.Hour) t.Run("wipes a quarantined replica when pool healthy", func(t *testing.T) { + c := assert.NewCollecting(t) shard := qrShard() markBackupHealthy(shard) badPod := qrPod(shard, 1, old, false) // quarantined replica, old, not ready @@ -178,21 +180,20 @@ func TestReconcileQuarantineRemediation(t *testing.T) { store := qrStoreWithQuarantined(t, shard, 1) acted, err := r.reconcileQuarantineRemediation(context.Background(), store, shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !acted { - t.Fatal("expected remediation to act on the quarantined pod") - } - if exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)) { - t.Error("expected quarantined pod to be deleted") - } - if exists(t, r, &corev1.PersistentVolumeClaim{}, client.ObjectKeyFromObject(badPVC)) { - t.Error("expected quarantined pod's data PVC to be deleted (wiped)") - } - if !exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(goodPod)) { - t.Error("healthy sibling pod should be untouched") - } + c.Require().NoError(err, "unexpected error") + c.Require().True(acted, "expected remediation to act on the quarantined pod") + c.False( + exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)), + "expected quarantined pod to be deleted", + ) + c.False( + exists(t, r, &corev1.PersistentVolumeClaim{}, client.ObjectKeyFromObject(badPVC)), + "expected quarantined pod's data PVC to be deleted (wiped)", + ) + c.True( + exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(goodPod)), + "healthy sibling pod should be untouched", + ) // The remediation event should carry the topology quarantine reason. rec := r.Recorder.(*record.FakeRecorder) @@ -202,12 +203,15 @@ func TestReconcileQuarantineRemediation(t *testing.T) { foundReason = true } } - if !foundReason { - t.Errorf("expected a remediation event containing the quarantine reason %q", qrReason) - } + c.True( + foundReason, + "expected a remediation event containing the quarantine reason %q", + qrReason, + ) }) t.Run("wipes with an empty reason -> event falls back to 'unspecified'", func(t *testing.T) { + c := assert.NewCollecting(t) shard := qrShard() markBackupHealthy(shard) badPod := qrPod(shard, 1, old, false) @@ -217,15 +221,12 @@ func TestReconcileQuarantineRemediation(t *testing.T) { store := qrStoreWithQuarantinedReason(t, shard, "", 1) // no reason recorded acted, err := r.reconcileQuarantineRemediation(context.Background(), store, shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !acted { - t.Fatal("expected remediation to act even without a recorded reason") - } - if exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)) { - t.Error("expected quarantined pod to be deleted") - } + c.Require().NoError(err, "unexpected error") + c.Require().True(acted, "expected remediation to act even without a recorded reason") + c.False( + exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)), + "expected quarantined pod to be deleted", + ) rec := r.Recorder.(*record.FakeRecorder) foundUnspecified := false @@ -234,12 +235,11 @@ func TestReconcileQuarantineRemediation(t *testing.T) { foundUnspecified = true } } - if !foundUnspecified { - t.Error("expected the event to fall back to 'unspecified' on empty reason") - } + c.True(foundUnspecified, "expected the event to fall back to 'unspecified' on empty reason") }) t.Run("defers when the pod is too young (stale-record guard)", func(t *testing.T) { + c := assert.NewCollecting(t) shard := qrShard() markBackupHealthy(shard) badPod := qrPod(shard, 1, time.Now(), false) // just created @@ -249,21 +249,20 @@ func TestReconcileQuarantineRemediation(t *testing.T) { store := qrStoreWithQuarantined(t, shard, 1) acted, err := r.reconcileQuarantineRemediation(context.Background(), store, shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if acted { - t.Error("expected no action for a too-young pod") - } - if !exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)) { - t.Error("young quarantined pod must not be deleted yet") - } - if !exists(t, r, &corev1.PersistentVolumeClaim{}, client.ObjectKeyFromObject(badPVC)) { - t.Error("young quarantined pod's PVC must not be deleted yet") - } + c.Require().NoError(err, "unexpected error") + c.False(acted, "expected no action for a too-young pod") + c.True( + exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)), + "young quarantined pod must not be deleted yet", + ) + c.True( + exists(t, r, &corev1.PersistentVolumeClaim{}, client.ObjectKeyFromObject(badPVC)), + "young quarantined pod's PVC must not be deleted yet", + ) }) t.Run("defers when another pod in the pool is unhealthy", func(t *testing.T) { + c := assert.NewCollecting(t) shard := qrShard() markBackupHealthy(shard) badPod := qrPod(shard, 1, old, false) // quarantined @@ -273,18 +272,16 @@ func TestReconcileQuarantineRemediation(t *testing.T) { store := qrStoreWithQuarantined(t, shard, 1) // only idx 1 quarantined acted, err := r.reconcileQuarantineRemediation(context.Background(), store, shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if acted { - t.Error("expected no action while another pool pod is unhealthy") - } - if !exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)) { - t.Error("quarantined pod must not be wiped during a broader outage") - } + c.Require().NoError(err, "unexpected error") + c.False(acted, "expected no action while another pool pod is unhealthy") + c.True( + exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)), + "quarantined pod must not be wiped during a broader outage", + ) }) t.Run("defers when no healthy backup exists (never wipe the last copy)", func(t *testing.T) { + c := assert.NewCollecting(t) shard := qrShard() // no backup-health condition => not healthy badPod := qrPod(shard, 1, old, false) goodPod := qrPod(shard, 0, old, true) @@ -293,21 +290,20 @@ func TestReconcileQuarantineRemediation(t *testing.T) { store := qrStoreWithQuarantined(t, shard, 1) acted, err := r.reconcileQuarantineRemediation(context.Background(), store, shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if acted { - t.Error("expected no action when there is no healthy backup to restore from") - } - if !exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)) { - t.Error("must not delete the pod without a healthy backup") - } - if !exists(t, r, &corev1.PersistentVolumeClaim{}, client.ObjectKeyFromObject(badPVC)) { - t.Error("must not wipe the data PVC without a healthy backup") - } + c.Require().NoError(err, "unexpected error") + c.False(acted, "expected no action when there is no healthy backup to restore from") + c.True( + exists(t, r, &corev1.Pod{}, client.ObjectKeyFromObject(badPod)), + "must not delete the pod without a healthy backup", + ) + c.True( + exists(t, r, &corev1.PersistentVolumeClaim{}, client.ObjectKeyFromObject(badPVC)), + "must not wipe the data PVC without a healthy backup", + ) }) t.Run("no quarantined poolers -> no action", func(t *testing.T) { + c := assert.NewCollecting(t) shard := qrShard() p0 := qrPod(shard, 0, old, true) p1 := qrPod(shard, 1, old, true) @@ -315,11 +311,7 @@ func TestReconcileQuarantineRemediation(t *testing.T) { store := qrStoreWithQuarantined(t, shard) // none quarantined acted, err := r.reconcileQuarantineRemediation(context.Background(), store, shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if acted { - t.Error("expected no action when nothing is quarantined") - } + c.Require().NoError(err, "unexpected error") + c.False(acted, "expected no action when nothing is quarantined") }) } diff --git a/pkg/resource-handler/controller/shard/reconcile_readiness_test.go b/pkg/resource-handler/controller/shard/reconcile_readiness_test.go index c827da41..3518d6b7 100644 --- a/pkg/resource-handler/controller/shard/reconcile_readiness_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_readiness_test.go @@ -12,17 +12,16 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/posture" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestReconcilePoolerReadiness(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := runtime.NewScheme() - if err := corev1.AddToScheme(scheme); err != nil { - t.Fatalf("add Pod scheme: %v", err) - } - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("add Shard scheme: %v", err) - } + ck.NoError(corev1.AddToScheme(scheme), "add Pod scheme") + ck.NoError(multigresv1alpha1.AddToScheme(scheme), "add Shard scheme") shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -51,35 +50,23 @@ func TestReconcilePoolerReadiness(t *testing.T) { Build() r := &ShardReconciler{Client: c, Scheme: scheme} - if err := r.reconcilePoolerReadiness(t.Context(), shard, map[string]posture.Readiness{ + ck.NoError(r.reconcilePoolerReadiness(t.Context(), shard, map[string]posture.Readiness{ pod.Name: {Ready: true, Reason: "DataPlaneReady", Message: "ready"}, - }); err != nil { - t.Fatalf("reconcile readiness: %v", err) - } + }), "reconcile readiness") updated := &corev1.Pod{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(pod), updated); err != nil { - t.Fatalf("get updated pod: %v", err) - } + ck.NoError(c.Get(t.Context(), client.ObjectKeyFromObject(pod), updated), "get updated pod") condition := readinessCondition(updated.Status.Conditions) - if condition == nil || + ck.False(condition == nil || condition.Status != corev1.ConditionTrue || - condition.Reason != "DataPlaneReady" { - t.Fatalf("readiness condition = %#v, want true DataPlaneReady", condition) - } + condition.Reason != "DataPlaneReady", "readiness condition = %#v, want true DataPlaneReady", condition) - if err := r.reconcilePoolerReadiness(t.Context(), shard, nil); err != nil { - t.Fatalf("reconcile missing observation: %v", err) - } - if err := c.Get(t.Context(), client.ObjectKeyFromObject(pod), updated); err != nil { - t.Fatalf("get unready pod: %v", err) - } + ck.NoError(r.reconcilePoolerReadiness(t.Context(), shard, nil), "reconcile missing observation") + ck.NoError(c.Get(t.Context(), client.ObjectKeyFromObject(pod), updated), "get unready pod") condition = readinessCondition(updated.Status.Conditions) - if condition == nil || + ck.False(condition == nil || condition.Status != corev1.ConditionFalse || - condition.Reason != "ObservationUnavailable" { - t.Fatalf("readiness condition = %#v, want false ObservationUnavailable", condition) - } + condition.Reason != "ObservationUnavailable", "readiness condition = %#v, want false ObservationUnavailable", condition) } func readinessCondition(conditions []corev1.PodCondition) *corev1.PodCondition { diff --git a/pkg/resource-handler/controller/shard/reconcile_shared_infra_test.go b/pkg/resource-handler/controller/shard/reconcile_shared_infra_test.go index 69b7e5e7..bb04d5c9 100644 --- a/pkg/resource-handler/controller/shard/reconcile_shared_infra_test.go +++ b/pkg/resource-handler/controller/shard/reconcile_shared_infra_test.go @@ -8,6 +8,8 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestPgBackRestPoolDNSNames(t *testing.T) { @@ -27,9 +29,7 @@ func TestPgBackRestPoolDNSNames(t *testing.T) { t.Run("no pools yields no DNS names", func(t *testing.T) { shard := baseShard() - if got := pgBackRestPoolDNSNames(shard); len(got) != 0 { - t.Errorf("expected no DNS names, got %v", got) - } + assert.NewCollecting(t).Empty(pgBackRestPoolDNSNames(shard), "expected no DNS names, got") }) t.Run("single pool, single cell yields a scoped wildcard for that svc", func(t *testing.T) { @@ -79,9 +79,8 @@ func TestPgBackRestPoolDNSNames(t *testing.T) { "empty": {}, } - if got := pgBackRestPoolDNSNames(shard); len(got) != 0 { - t.Errorf("expected no DNS names for a pool with no cells, got %v", got) - } + assert.NewCollecting(t). + Empty(pgBackRestPoolDNSNames(shard), "expected no DNS names for a pool with no cells, got") }) } @@ -89,17 +88,14 @@ func TestPgBackRestPoolDNSNames(t *testing.T) { // elements, ignoring order. func assertSameStringSet(t *testing.T, got, want []string) { t.Helper() + c := assert.NewAborting(t) gotSorted := append([]string(nil), got...) wantSorted := append([]string(nil), want...) sort.Strings(gotSorted) sort.Strings(wantSorted) - if len(gotSorted) != len(wantSorted) { - t.Fatalf("got %v, want %v", got, want) - } + c.Len(gotSorted, len(wantSorted), "got %v, want %v", got, want) for i := range gotSorted { - if gotSorted[i] != wantSorted[i] { - t.Fatalf("got %v, want %v", got, want) - } + c.Eq(wantSorted[i], gotSorted[i], "got %v, want %v", got, want) } } diff --git a/pkg/resource-handler/controller/shard/registration_requeue_test.go b/pkg/resource-handler/controller/shard/registration_requeue_test.go index b9a5ff99..fd284815 100644 --- a/pkg/resource-handler/controller/shard/registration_requeue_test.go +++ b/pkg/resource-handler/controller/shard/registration_requeue_test.go @@ -22,6 +22,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/data-handler/poolerclient" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) // fakeClock lets a test drive recordNotConverged's elapsed-time math without @@ -71,10 +73,8 @@ func TestReadinessBackoffDelayClampsElapsedTime(t *testing.T) { // than on one draw. for range 200 { got := readinessBackoffDelay(tc.elapsed) - if got < tc.min || got > tc.max { - t.Fatalf("elapsed=%v: delay %v outside [%v, %v]", - tc.elapsed, got, tc.min, tc.max) - } + assert.NewAborting(t). + False(got < tc.min || got > tc.max, "elapsed=%v: delay %v outside [%v, %v]", tc.elapsed, got, tc.min, tc.max) } } } @@ -90,10 +90,9 @@ func TestReadinessBackoffDelayIsNeverZero(t *testing.T) { 30 * time.Second, time.Minute, time.Hour, } { for range 50 { - if got := readinessBackoffDelay(elapsed); got <= 0 { - t.Fatalf("elapsed=%v produced a non-positive delay %v, "+ - "which controller-runtime reads as no requeue", elapsed, got) - } + got := readinessBackoffDelay(elapsed) + assert.NewAborting(t).Greater(0, got, "elapsed=%v produced a non-positive delay %v, "+ + "which controller-runtime reads as no requeue", elapsed, got) } } } @@ -108,9 +107,7 @@ func TestReadinessBackoffDelayJitters(t *testing.T) { for range 200 { seen[readinessBackoffDelay(30*time.Second)] = true } - if len(seen) < 10 { - t.Fatalf("only %d distinct delays across 200 draws; jitter is not applied", len(seen)) - } + assert.NewAborting(t).GreaterOrEqual(10, len(seen), "only") } // shardNamed is the minimum a strike counter reads. @@ -134,9 +131,7 @@ func TestPostureStrikesLeaveNoEntryOnceSettled(t *testing.T) { r.recordPostureObservation(s, false) } - if got := len(r.postureStrikes); got != 0 { - t.Fatalf("a thousand shards seen and settled left %d entries, want 0", got) - } + assert.NewAborting(t).Eq(0, len(r.postureStrikes), "a thousand shards seen and settled left") } // TestPostureStrikesDoNotSurviveRecreation pins the intended semantic: a @@ -146,6 +141,7 @@ func TestPostureStrikesLeaveNoEntryOnceSettled(t *testing.T) { // do not pick a requeue delay, that is notConvergedSince's job. func TestPostureStrikesDoNotSurviveRecreation(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) r := &ShardReconciler{} s := shardNamed("ns", "shard-0") @@ -153,37 +149,32 @@ func TestPostureStrikesDoNotSurviveRecreation(t *testing.T) { for range 5 { r.recordPostureObservation(s, true) } - if got := r.recordPostureObservation(s, false); got != 0 { - t.Fatalf("a settled observation reported %d strikes, want 0", got) - } + c.Eq(0, r.recordPostureObservation(s, false), "a settled observation reported") // The replacement is a different object at the same key, which is what // the tablegroup controller creates after a Shard is deleted. - if got := r.recordPostureObservation(shardNamed("ns", "shard-0"), true); got != 1 { - t.Fatalf("a recreated shard opened at %d strikes, want 1", got) - } + c.Eq( + 1, + r.recordPostureObservation(shardNamed("ns", "shard-0"), true), + "a recreated shard opened at", + ) } // TestPostureStrikesCountConsecutiveUnsettled pins what the counter is for, // so settling on delete cannot be "fixed" into never counting at all. func TestPostureStrikesCountConsecutiveUnsettled(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) r := &ShardReconciler{} s := shardNamed("ns", "shard-0") for want := 1; want <= 3; want++ { - if got := r.recordPostureObservation(s, true); got != want { - t.Fatalf("consecutive unsettled observation %d reported %d strikes", want, got) - } + c.Eq(want, r.recordPostureObservation(s, true), "consecutive unsettled observation") } // Shards are counted independently, which is the only reason the map has // keys at all. - if got := r.recordPostureObservation(shardNamed("ns", "other"), true); got != 1 { - t.Fatalf("a second shard opened at %d strikes, want 1", got) - } - if got := r.recordPostureObservation(s, true); got != 4 { - t.Fatalf("the first shard reported %d strikes after a second shard, want 4", got) - } + c.Eq(1, r.recordPostureObservation(shardNamed("ns", "other"), true), "a second shard opened at") + c.Eq(4, r.recordPostureObservation(s, true), "the first shard reported") } // TestNotConvergedSinceLeavesNoEntryOnceSettled mirrors @@ -200,9 +191,7 @@ func TestNotConvergedSinceLeavesNoEntryOnceSettled(t *testing.T) { r.recordNotConverged(s, false) } - if got := len(r.notConvergedSince); got != 0 { - t.Fatalf("a thousand shards seen and settled left %d entries, want 0", got) - } + assert.NewAborting(t).Eq(0, len(r.notConvergedSince), "a thousand shards seen and settled left") } // TestNotConvergedSinceTracksElapsedTime pins the elapsed-time semantics: the @@ -211,38 +200,36 @@ func TestNotConvergedSinceLeavesNoEntryOnceSettled(t *testing.T) { // it so a later not-converged spell starts over rather than resuming. func TestNotConvergedSinceTracksElapsedTime(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) clk := &fakeClock{t: time.Unix(1_700_000_000, 0)} r := &ShardReconciler{Clock: clk.now} s := shardNamed("ns", "shard-0") - if got := r.recordNotConverged(s, true); got != 0 { - t.Fatalf("first not-converged observation reported elapsed %v, want 0", got) - } + c.Eq(0, r.recordNotConverged(s, true), "first not-converged observation reported elapsed") clk.advance(37 * time.Second) - if got := r.recordNotConverged(s, true); got != 37*time.Second { - t.Fatalf("second not-converged observation reported elapsed %v, want 37s", got) - } + c.Eq( + 37*time.Second, + r.recordNotConverged(s, true), + "second not-converged observation reported elapsed", + ) // A burst of same-instant calls (a wave of unrelated pod events) must not // itself advance the elapsed time. - if got := r.recordNotConverged(s, true); got != 37*time.Second { - t.Fatalf( - "third not-converged observation (no time passed) reported elapsed %v, want 37s", got, - ) - } - - if got := r.recordNotConverged(s, false); got != 0 { - t.Fatalf("a settled observation reported elapsed %v, want 0", got) - } - if got := r.recordNotConverged(s, true); got != 0 { - t.Fatalf("a fresh not-converged spell reported elapsed %v, want 0 (not resumed)", got) - } + c.Eq( + 37*time.Second, + r.recordNotConverged(s, true), + "third not-converged observation (no time passed) reported elapsed", + ) + + c.Eq(0, r.recordNotConverged(s, false), "a settled observation reported elapsed") + c.Eq(0, r.recordNotConverged(s, true), "a fresh not-converged spell reported elapsed") } // TestForgetStrikesDropsBothCounters pins that a Shard's strike entries do // not survive forgetStrikes, in either counter. func TestForgetStrikesDropsBothCounters(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) r := &ShardReconciler{} s := shardNamed("ns", "shard-0") @@ -251,12 +238,8 @@ func TestForgetStrikesDropsBothCounters(t *testing.T) { r.forgetStrikes(s.Namespace, s.Name) - if got := len(r.postureStrikes); got != 0 { - t.Fatalf("posture strikes: %d entries survived forgetStrikes, want 0", got) - } - if got := len(r.notConvergedSince); got != 0 { - t.Fatalf("not-converged-since: %d entries survived forgetStrikes, want 0", got) - } + c.Eq(0, len(r.postureStrikes), "posture strikes") + c.Eq(0, len(r.notConvergedSince), "not-converged-since") } // TestHandleDeletionForgetsStrikes drives the deletion cleanup through the @@ -264,6 +247,7 @@ func TestForgetStrikesDropsBothCounters(t *testing.T) { // directly, so a regression that stops handleDeletion from reaching it is // caught here rather than only in the helper's own unit test. func TestHandleDeletionForgetsStrikes(t *testing.T) { + ck := assert.NewCollecting(t) shard := postureTestShard() shard.Finalizers = []string{shardFinalizer} now := metav1.Now() @@ -285,16 +269,14 @@ func TestHandleDeletionForgetsStrikes(t *testing.T) { r.postureStrikes = map[string]int{key: 2} r.notConvergedSince = map[string]time.Time{key: time.Now()} - if _, err := r.handleDeletion(t.Context(), shard); err != nil { - t.Fatalf("handleDeletion() error = %v", err) - } + _, err := r.handleDeletion(t.Context(), shard) + ck.Require().NoError(err, "handleDeletion() error =") if _, ok := r.postureStrikes[key]; ok { t.Errorf("posture strikes entry for %s survived handleDeletion", key) } - if _, ok := r.notConvergedSince[key]; ok { - t.Errorf("not-converged-since entry for %s survived handleDeletion", key) - } + _, ok := r.notConvergedSince[key] + ck.False(ok, "not-converged-since entry for %s survived handleDeletion", key) } // TestReconcileForgetsStrikesOnNotFound drives the not-found cleanup through @@ -302,6 +284,7 @@ func TestHandleDeletionForgetsStrikes(t *testing.T) { // once handleDeletion above has already run and removed the finalizer) must // still have its strike entries dropped, as a backstop. func TestReconcileForgetsStrikesOnNotFound(t *testing.T) { + ck := assert.NewCollecting(t) scheme := postureTestScheme(t) c := fake.NewClientBuilder().WithScheme(scheme).Build() r := &ShardReconciler{ @@ -317,16 +300,14 @@ func TestReconcileForgetsStrikesOnNotFound(t *testing.T) { req := ctrl.Request{ NamespacedName: types.NamespacedName{Namespace: "default", Name: "gone-shard"}, } - if _, err := r.Reconcile(t.Context(), req); err != nil { - t.Fatalf("Reconcile() error = %v", err) - } + _, err := r.Reconcile(t.Context(), req) + ck.Require().NoError(err, "Reconcile() error =") if _, ok := r.postureStrikes[key]; ok { t.Errorf("posture strikes entry for %s survived Reconcile on a missing Shard", key) } - if _, ok := r.notConvergedSince[key]; ok { - t.Errorf("not-converged-since entry for %s survived Reconcile on a missing Shard", key) - } + _, ok := r.notConvergedSince[key] + ck.False(ok, "not-converged-since entry for %s survived Reconcile on a missing Shard", key) } // gateTestReconciler is postureTestReconciler plus a Pod status subresource, @@ -370,20 +351,19 @@ func registeredReplica( ) topoclient.ComponentID { t.Helper() id := &clustermetadata.ID{Cell: cell, Name: name} - if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ - Id: id, - Hostname: name, - ShardKey: &clustermetadata.ShardKey{ - Database: string(shard.Spec.DatabaseName), - TableGroup: string(shard.Spec.TableGroupName), - Shard: string(shard.Spec.ShardName), - }, - RoutingState: &clustermetadata.RoutingState{ - Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, - }, - }, false); err != nil { - t.Fatalf("register pooler %s: %v", name, err) - } + assert.NewAborting(t). + NoError(store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + Id: id, + Hostname: name, + ShardKey: &clustermetadata.ShardKey{ + Database: string(shard.Spec.DatabaseName), + TableGroup: string(shard.Spec.TableGroupName), + Shard: string(shard.Spec.ShardName), + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + }, + }, false), "register pooler %s", name) componentID := topoclient.ComponentIDString(id) rpc.SetStatusResponse(componentID, readyStatusResponse(id)) @@ -438,20 +418,19 @@ func notYetSettledReplica( ) *clustermetadata.ID { t.Helper() id := &clustermetadata.ID{Cell: cell, Name: name} - if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ - Id: id, - Hostname: name, - ShardKey: &clustermetadata.ShardKey{ - Database: string(shard.Spec.DatabaseName), - TableGroup: string(shard.Spec.TableGroupName), - Shard: string(shard.Spec.ShardName), - }, - RoutingState: &clustermetadata.RoutingState{ - Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, - }, - }, false); err != nil { - t.Fatalf("register pooler %s: %v", name, err) - } + assert.NewAborting(t). + NoError(store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + Id: id, + Hostname: name, + ShardKey: &clustermetadata.ShardKey{ + Database: string(shard.Spec.DatabaseName), + TableGroup: string(shard.Spec.TableGroupName), + Shard: string(shard.Spec.ShardName), + }, + RoutingState: &clustermetadata.RoutingState{ + Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, + }, + }, false), "register pooler %s", name) componentID := topoclient.ComponentIDString(id) rpc.SetStatusResponse(componentID, &multipoolermanagerdatapb.StatusResponse{ @@ -478,15 +457,14 @@ func notYetSettledReplica( // selector that forgets to scope to pool pods seeds it AwaitingRegistration // forever and this shard never returns 0. func TestReconcilePostureConvergedShardReturnsZero(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() shard.Labels[metadata.LabelMultigresDatabase] = "database" shard.Labels[metadata.LabelMultigresTableGroup] = "table-group" shard.Labels[metadata.LabelMultigresShard] = "0" dep, err := BuildMultiorchDeployment(shard, "cell1", postureTestScheme(t)) - if err != nil { - t.Fatalf("BuildMultiorchDeployment() error = %v", err) - } + c.Require().NoError(err, "BuildMultiorchDeployment() error =") orch := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{ Name: "multiorch-abc", Namespace: shard.Namespace, Labels: dep.Spec.Template.Labels, }} @@ -501,19 +479,14 @@ func TestReconcilePostureConvergedShardReturnsZero(t *testing.T) { r, _ := postureTestReconciler(t, shard, rpc, pool, orch) delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("reconcilePosture() error = %v", err) - } - if delay != 0 { - t.Errorf("delay = %v, want 0 for a converged shard", delay) - } + c.Require().NoError(err, "reconcilePosture() error =") + c.Eq(0, delay, "delay") key := fmt.Sprintf("%s/%s", shard.Namespace, shard.Name) if _, ok := r.postureStrikes[key]; ok { t.Errorf("posture strikes entry left for a converged shard") } - if _, ok := r.notConvergedSince[key]; ok { - t.Errorf("not-converged-since entry left for a converged shard") - } + _, ok := r.notConvergedSince[key] + c.False(ok, "not-converged-since entry left for a converged shard") if conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { t.Errorf("conditions = %#v, want no failure for a converged shard", shard.Status.Conditions) } @@ -526,6 +499,7 @@ func TestReconcilePostureConvergedShardReturnsZero(t *testing.T) { // failing pod has also never reached posture readiness, so this must keep // requesting a requeue rather than stranding it until the 10h resync. func TestReconcilePostureAcceptedIncompleteObservationStillRequeues(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) @@ -541,23 +515,19 @@ func TestReconcilePostureAcceptedIncompleteObservationStillRequeues(t *testing.T r, _ := postureTestReconciler(t, shard, rpc, p0, p1) first, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("first reconcilePosture() error = %v", err) - } - if first != postureDebounceRequeueDelay { - t.Errorf("first delay = %v, want the %v debounce", first, postureDebounceRequeueDelay) - } + c.Require().NoError(err, "first reconcilePosture() error =") + c.Eq(postureDebounceRequeueDelay, first, "first delay") for pass := 2; pass <= 4; pass++ { delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("pass %d reconcilePosture() error = %v", pass, err) - } - if delay <= 0 { - t.Errorf( - "pass %d: delay = %v, want non-zero (RPC failure still unresolved)", pass, delay, - ) - } + c.Require().NoError(err, "pass %d reconcilePosture() error =", pass) + c.Greater( + 0, + delay, + "pass %d: delay = %v, want non-zero (RPC failure still unresolved)", + pass, + delay, + ) } } @@ -568,13 +538,14 @@ func TestReconcilePostureAcceptedIncompleteObservationStillRequeues(t *testing.T // threshold it is accepted as PostureConsistent=False, and this must keep // requesting a requeue. func TestReconcilePostureAcceptedMismatchAndNotReadyStillRequeues(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) defer func() { _ = store.Close() }() rpc := rpcclient.NewFakeClient() id := &clustermetadata.ID{Cell: "cell1", Name: "pooler-0"} - if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + c.Require().NoError(store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ Id: id, Hostname: "pooler-0", ShardKey: &clustermetadata.ShardKey{ @@ -585,9 +556,7 @@ func TestReconcilePostureAcceptedMismatchAndNotReadyStillRequeues(t *testing.T) RoutingState: &clustermetadata.RoutingState{ Role: clustermetadata.RoutingRole_ROUTING_ROLE_REPLICA, }, - }, false); err != nil { - t.Fatalf("register pooler: %v", err) - } + }, false), "register pooler") componentID := topoclient.ComponentIDString(id) rpc.SetStatusResponse(componentID, &multipoolermanagerdatapb.StatusResponse{ Status: &multipoolermanagerdatapb.Status{ @@ -597,21 +566,19 @@ func TestReconcilePostureAcceptedMismatchAndNotReadyStillRequeues(t *testing.T) r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) first, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("first reconcilePosture() error = %v", err) - } - if first != postureDebounceRequeueDelay { - t.Errorf("first delay = %v, want the %v debounce", first, postureDebounceRequeueDelay) - } + c.Require().NoError(err, "first reconcilePosture() error =") + c.Eq(postureDebounceRequeueDelay, first, "first delay") for pass := 2; pass <= 4; pass++ { delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("pass %d reconcilePosture() error = %v", pass, err) - } - if delay <= 0 { - t.Errorf("pass %d: delay = %v, want non-zero (mismatch still unresolved)", pass, delay) - } + c.Require().NoError(err, "pass %d reconcilePosture() error =", pass) + c.Greater( + 0, + delay, + "pass %d: delay = %v, want non-zero (mismatch still unresolved)", + pass, + delay, + ) if !conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { t.Errorf("pass %d: conditions = %#v, want PostureConsistent=False once accepted", pass, shard.Status.Conditions) @@ -629,6 +596,7 @@ func TestReconcilePostureAcceptedMismatchAndNotReadyStillRequeues(t *testing.T) // who is primary must not be left Degraded until the 10h resync once that // disagreement is accepted into status. func TestReconcilePostureAcceptedMismatchWithReadyPodsStillRequeues(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) @@ -639,25 +607,19 @@ func TestReconcilePostureAcceptedMismatchWithReadyPodsStillRequeues(t *testing.T r, _ := postureTestReconciler(t, shard, rpc, postureTestPod()) first, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("first reconcilePosture() error = %v", err) - } - if first != postureDebounceRequeueDelay { - t.Errorf("first delay = %v, want the %v debounce", first, postureDebounceRequeueDelay) - } + c.Require().NoError(err, "first reconcilePosture() error =") + c.Eq(postureDebounceRequeueDelay, first, "first delay") for pass := 2; pass <= 4; pass++ { delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("pass %d reconcilePosture() error = %v", pass, err) - } - if delay <= 0 { - t.Errorf( - "pass %d: delay = %v, want non-zero (mismatch still unresolved, though the pod is ready)", - pass, - delay, - ) - } + c.Require().NoError(err, "pass %d reconcilePosture() error =", pass) + c.Greater( + 0, + delay, + "pass %d: delay = %v, want non-zero (mismatch still unresolved, though the pod is ready)", + pass, + delay, + ) if !conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { t.Errorf("pass %d: conditions = %#v, want PostureConsistent=False once accepted", pass, shard.Status.Conditions) @@ -675,6 +637,7 @@ func TestReconcilePostureAcceptedMismatchWithReadyPodsStillRequeues(t *testing.T // regardless of how many reconcile passes it took to get there, so a mutation // that turns the backoff back into a per-pass count cannot pass by chance. func TestReconcilePostureRequeuesWhileAPodAwaitsItsPooler(t *testing.T) { + c := assert.NewCollecting(t) shard := postureTestShard() _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) @@ -701,13 +664,16 @@ func TestReconcilePostureRequeuesWhileAPodAwaitsItsPooler(t *testing.T) { } { clk.advance(tc.advance) delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("reconcilePosture() error = %v", err) - } + c.Require().NoError(err, "reconcilePosture() error =") min, max := wantDelayRange(tc.elapsed) - if delay < min || delay > max { - t.Errorf("at elapsed=%v: delay = %v, want in [%v, %v]", tc.elapsed, delay, min, max) - } + c.False( + delay < min || delay > max, + "at elapsed=%v: delay = %v, want in [%v, %v]", + tc.elapsed, + delay, + min, + max, + ) } } @@ -724,6 +690,7 @@ func TestReconcilePostureRequeuesWhileAPodAwaitsItsPooler(t *testing.T) { // fixture to a committed primary and checks the requeue stops, the gate goes // True, and the not-converged-since entry is gone. func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { + ck := assert.NewCollecting(t) shard := postureTestShard() _, factory := memorytopo.NewServerAndFactory(t.Context(), "cell1") store := topoclient.NewWithFactory(factory, "", []string{""}, topoclient.NewDefaultTopoConfig()) @@ -751,20 +718,22 @@ func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { } { clk.advance(tc.advance) delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("reconcilePosture() error = %v", err) - } + ck.Require().NoError(err, "reconcilePosture() error =") min, max := wantDelayRange(tc.elapsed) - if delay < min || delay > max { - t.Errorf("at elapsed=%v: delay = %v, want in [%v, %v]", tc.elapsed, delay, min, max) - } + ck.False( + delay < min || delay > max, + "at elapsed=%v: delay = %v, want in [%v, %v]", + tc.elapsed, + delay, + min, + max, + ) } for _, pod := range []*corev1.Pod{pod0, pod1} { got := &corev1.Pod{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(pod), got); err != nil { - t.Fatalf("get pod %s: %v", pod.Name, err) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(pod), got), "get pod %s", pod.Name) condition := readinessCondition(got.Status.Conditions) if condition == nil || condition.Status != corev1.ConditionFalse { t.Errorf("pod %s readiness condition = %#v, want False while waiting for a primary", @@ -776,7 +745,7 @@ func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { // pooler-0 as leader, so both are cohort-eligible members of the same // durability rule and pooler-0's postgres reports PRIMARY, matching the // topology role a leader-designate needs. - if err := store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ + ck.Require().NoError(store.RegisterMultipooler(t.Context(), &clustermetadata.Multipooler{ Id: id0, Hostname: "pooler-0", ShardKey: &clustermetadata.ShardKey{ @@ -787,9 +756,7 @@ func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { RoutingState: &clustermetadata.RoutingState{ Role: clustermetadata.RoutingRole_ROUTING_ROLE_PRIMARY, }, - }, true); err != nil { - t.Fatalf("promote pooler-0 in topology: %v", err) - } + }, true), "promote pooler-0 in topology") rule := &clustermetadata.ShardRule{ RuleNumber: &clustermetadata.RuleNumber{CoordinatorTerm: 1}, LeaderId: id0, @@ -827,16 +794,11 @@ func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { clk.advance(time.Second) delay, err := r.reconcilePosture(t.Context(), store, shard, rpc) - if err != nil { - t.Fatalf("reconcilePosture() after election error = %v", err) - } - if delay != 0 { - t.Errorf("delay after a primary is elected = %v, want 0", delay) - } + ck.Require().NoError(err, "reconcilePosture() after election error =") + ck.Eq(0, delay, "delay after a primary is elected") key := fmt.Sprintf("%s/%s", shard.Namespace, shard.Name) - if _, ok := r.notConvergedSince[key]; ok { - t.Errorf("not-converged-since entry left after a primary is elected") - } + _, ok := r.notConvergedSince[key] + ck.False(ok, "not-converged-since entry left after a primary is elected") if conditionIsFalse(shard.Status.Conditions, "PostureConsistent") { t.Errorf( "conditions = %#v, want no failure once a primary is elected", @@ -846,9 +808,8 @@ func TestReconcilePostureBacksOffThenClearsOnceAPrimaryIsElected(t *testing.T) { for _, pod := range []*corev1.Pod{pod0, pod1} { got := &corev1.Pod{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(pod), got); err != nil { - t.Fatalf("get pod %s: %v", pod.Name, err) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(pod), got), "get pod %s", pod.Name) condition := readinessCondition(got.Status.Conditions) if condition == nil || condition.Status != corev1.ConditionTrue { t.Errorf("pod %s readiness condition = %#v, want True once a primary is elected", diff --git a/pkg/resource-handler/controller/shard/reload_internal_test.go b/pkg/resource-handler/controller/shard/reload_internal_test.go index 1edbea82..fc3cd019 100644 --- a/pkg/resource-handler/controller/shard/reload_internal_test.go +++ b/pkg/resource-handler/controller/shard/reload_internal_test.go @@ -20,6 +20,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) const ( @@ -31,13 +33,10 @@ const ( func reloadTestScheme(t *testing.T) *runtime.Scheme { t.Helper() + c := assert.NewAborting(t) s := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(s); err != nil { - t.Fatalf("add multigres scheme: %v", err) - } - if err := corev1.AddToScheme(s); err != nil { - t.Fatalf("add corev1 scheme: %v", err) - } + c.NoError(multigresv1alpha1.AddToScheme(s), "add multigres scheme") + c.NoError(corev1.AddToScheme(s), "add corev1 scheme") return s } @@ -107,13 +106,12 @@ func reloadTestStore(t *testing.T) (topoclient.Store, topoclient.ComponentID) { factory, "", []string{""}, topoclient.NewDefaultTopoConfig(), ) id := &clustermetadata.ID{Cell: reloadTestCell, Name: reloadTestPod} - if err := store.RegisterMultipooler(context.Background(), &clustermetadata.Multipooler{ - Id: id, - Hostname: reloadTestPod, - ShardKey: &clustermetadata.ShardKey{Database: "db", TableGroup: "tg", Shard: "0"}, - }, false); err != nil { - t.Fatalf("register pooler: %v", err) - } + assert.NewAborting(t). + NoError(store.RegisterMultipooler(context.Background(), &clustermetadata.Multipooler{ + Id: id, + Hostname: reloadTestPod, + ShardKey: &clustermetadata.ShardKey{Database: "db", TableGroup: "tg", Shard: "0"}, + }, false), "register pooler") return store, topoclient.ComponentIDString(id) } @@ -142,19 +140,18 @@ func callLogHas(log []string, method string) bool { func podReloadHash(t *testing.T, r *ShardReconciler) string { t.Helper() got := &corev1.Pod{} - if err := r.Get( + assert.NewAborting(t).NoError(r.Get( context.Background(), client.ObjectKey{Namespace: "ns", Name: reloadTestPod}, got, - ); err != nil { - t.Fatalf("get pod: %v", err) - } + ), "get pod") return got.Annotations[metadata.AnnotationPostgresReloadHash] } // TestReconcileReloadStateStampsWhenVerified: the RPC confirms the reload took // effect (config_load_time set), so the pod is stamped current. func TestReconcileReloadStateStampsWhenVerified(t *testing.T) { + c := assert.NewCollecting(t) scheme := reloadTestScheme(t) shard := reloadTestShard() store, poolerID := reloadTestStore(t) @@ -175,28 +172,19 @@ func TestReconcileReloadStateStampsWhenVerified(t *testing.T) { reloadTestRendered(), rpc, ) - if err != nil { - t.Fatalf("reconcileReloadState: %v", err) - } - if wait != 0 { - t.Errorf("wait = %v, want 0 (reload completed)", wait) - } + c.Require().NoError(err, "reconcileReloadState") + c.Eq(0, wait, "wait") if !callLogHas(rpc.GetCallLog(), "ReloadConfig") { t.Errorf("ReloadConfig was not called; call log = %v", rpc.GetCallLog()) } - if h := podReloadHash(t, r); h != reloadTestDesired { - t.Errorf( - "pod reload-hash = %q, want %q (stamped after verified reload)", - h, - reloadTestDesired, - ) - } + c.Eq(reloadTestDesired, podReloadHash(t, r), "pod reload-hash") } // TestReconcileReloadStateNotSyncedRetries: the RPC returns no config_load_time // (mounted file not yet caught up, or postgres down) — the pod must NOT be // stamped and the step must requeue. func TestReconcileReloadStateNotSyncedRetries(t *testing.T) { + c := assert.NewCollecting(t) scheme := reloadTestScheme(t) shard := reloadTestShard() store, poolerID := reloadTestStore(t) @@ -218,20 +206,19 @@ func TestReconcileReloadStateNotSyncedRetries(t *testing.T) { reloadTestRendered(), rpc, ) - if err != nil { - t.Fatalf("reconcileReloadState: %v", err) - } - if wait != reloadRetryDelay { - t.Errorf("wait = %v, want %v (retry until file syncs)", wait, reloadRetryDelay) - } - if h := podReloadHash(t, r); h == reloadTestDesired { - t.Errorf("pod reload-hash was stamped despite an unsynced file") - } + c.Require().NoError(err, "reconcileReloadState") + c.Eq(reloadRetryDelay, wait, "wait") + c.NotEq( + reloadTestDesired, + podReloadHash(t, r), + "pod reload-hash was stamped despite an unsynced file", + ) } // TestReconcileReloadStateNeedsRestart: a reload-classified setting actually // needs a restart — surfaced (requeue), not stamped. func TestReconcileReloadStateNeedsRestart(t *testing.T) { + c := assert.NewCollecting(t) scheme := reloadTestScheme(t) shard := reloadTestShard() store, poolerID := reloadTestStore(t) @@ -255,15 +242,13 @@ func TestReconcileReloadStateNeedsRestart(t *testing.T) { reloadTestRendered(), rpc, ) - if err != nil { - t.Fatalf("reconcileReloadState: %v", err) - } - if wait != reloadRetryDelay { - t.Errorf("wait = %v, want %v", wait, reloadRetryDelay) - } - if h := podReloadHash(t, r); h == reloadTestDesired { - t.Errorf("pod reload-hash was stamped despite needs_restart") - } + c.Require().NoError(err, "reconcileReloadState") + c.Eq(reloadRetryDelay, wait, "wait") + c.NotEq( + reloadTestDesired, + podReloadHash(t, r), + "pod reload-hash was stamped despite needs_restart", + ) } // TestReconcileReloadStateSkips: pods that are already current, draining, or @@ -299,15 +284,14 @@ func TestReconcileReloadStateSkips(t *testing.T) { pod := reloadTestPodObj(tc.reloadHash, tc.restart, tc.extraAnn) r := newReloadReconciler(scheme, rpc, shard, pod) - if _, err := r.reconcileReloadState( + _, err := r.reconcileReloadState( context.Background(), store, shard, reloadTestRendered(), rpc, - ); err != nil { - t.Fatalf("reconcileReloadState: %v", err) - } + ) + assert.NewAborting(t).NoError(err, "reconcileReloadState") if callLogHas(rpc.GetCallLog(), "ReloadConfig") { t.Errorf( "ReloadConfig should not be called for %q; call log = %v", diff --git a/pkg/resource-handler/controller/shard/secret_test.go b/pkg/resource-handler/controller/shard/secret_test.go index 2f4201b3..4acfeda6 100644 --- a/pkg/resource-handler/controller/shard/secret_test.go +++ b/pkg/resource-handler/controller/shard/secret_test.go @@ -11,6 +11,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) const testPostgresAuthRefName = "multigres-admin-ref" @@ -78,9 +80,8 @@ func TestReconcilePostgresPasswordSecret_ValidatesExternalRef(t *testing.T) { Scheme: scheme, } - if err := reconciler.reconcilePostgresPasswordSecret(context.Background(), shard); err != nil { - t.Fatalf("reconcilePostgresPasswordSecret() error = %v", err) - } + assert.NewAborting(t). + NoError(reconciler.reconcilePostgresPasswordSecret(context.Background(), shard), "reconcilePostgresPasswordSecret() error =") } func TestReconcilePostgresInitSecretsSecret(t *testing.T) { @@ -254,6 +255,7 @@ func TestReconcilePostgresInitSecretsSecret(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) shard := newShard() tc.configureShard(shard) @@ -268,23 +270,26 @@ func TestReconcilePostgresInitSecretsSecret(t *testing.T) { } err := reconciler.reconcilePostgresInitSecretsSecret(context.Background(), shard) - if tc.wantErr && err == nil { - t.Fatal("expected error, got nil") - } - if !tc.wantErr && err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().False(tc.wantErr && err == nil, "expected error, got nil") + c.Require().False(!tc.wantErr && err != nil, "unexpected error: %v", err) if err != nil { msg := err.Error() - if tc.wantErrText != "" && !strings.Contains(msg, tc.wantErrText) { - t.Errorf("error = %q, want it to contain %q", msg, tc.wantErrText) - } - if strings.Contains(msg, payloadRoleName) { - t.Errorf("error message must not contain role names, got: %q", msg) - } - if strings.Contains(msg, payloadPassword) { - t.Errorf("error message must not contain payload values, got: %q", msg) - } + c.False( + tc.wantErrText != "" && !strings.Contains(msg, tc.wantErrText), + "error = %q, want it to contain %q", + msg, + tc.wantErrText, + ) + c.NotStrContains( + msg, + payloadRoleName, + "error message must not contain role names, got", + ) + c.NotStrContains( + msg, + payloadPassword, + "error message must not contain payload values, got", + ) } }) } diff --git a/pkg/resource-handler/controller/shard/shard_controller_internal_test.go b/pkg/resource-handler/controller/shard/shard_controller_internal_test.go index 344f3ce7..c774a331 100644 --- a/pkg/resource-handler/controller/shard/shard_controller_internal_test.go +++ b/pkg/resource-handler/controller/shard/shard_controller_internal_test.go @@ -33,6 +33,8 @@ import ( "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/pkg/util/name" + + "github.com/multigres/testkit/assert" ) // Helper functions moved from shard_controller_test_util_test.go @@ -158,9 +160,8 @@ func TestSetConditions(t *testing.T) { // Use go-cmp for exact match, ignoring LastTransitionTime opts := cmpopts.IgnoreFields(metav1.Condition{}, "LastTransitionTime") - if diff := cmp.Diff(tc.want, got, opts); diff != "" { - t.Errorf("setConditions() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t). + EqDiffOpts(tc.want, got, []cmp.Option{opts}, "setConditions() mismatch") }) } } @@ -168,6 +169,7 @@ func TestSetConditions(t *testing.T) { // TestBuildMultiorchContainer_WithImage tests buildMultiorchContainer with custom image. // This tests the image override path that was missing coverage. func TestBuildMultiorchContainer_WithImage(t *testing.T) { + c := assert.NewCollecting(t) customImage := "custom/multiorch:v1.2.3" shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -188,12 +190,8 @@ func TestBuildMultiorchContainer_WithImage(t *testing.T) { container := buildMultiorchContainer(shard, "zone1") - if container.Image != customImage { - t.Errorf("buildMultiorchContainer() image = %s, want %s", container.Image, customImage) - } - if container.Name != "multiorch" { - t.Errorf("buildMultiorchContainer() name = %s, want multiorch", container.Name) - } + c.Eq(customImage, container.Image, "buildMultiorchContainer() image") + c.Eq("multiorch", container.Name, "buildMultiorchContainer() name") } // TestReconcile_InvalidScheme tests the error path when Build* functions fail due to invalid scheme. @@ -317,9 +315,8 @@ func TestReconcile_InvalidScheme(t *testing.T) { } err := tc.reconcileFunc(reconciler, context.Background(), shard) - if err == nil { - t.Errorf("reconcile function should error with invalid scheme") - } + assert.NewCollecting(t). + Error(err, "reconcile function should error with invalid scheme") }) } } @@ -361,9 +358,8 @@ func TestUpdateStatus_PoolPodsNotFound(t *testing.T) { // Call updateStatus when pool Pods don't exist yet err := reconciler.updateStatus(context.Background(), shard, renderedConfig{}) - if err != nil { - t.Errorf("updateStatus() should not error when pool Pods not found, got: %v", err) - } + assert.NewCollecting(t). + NoError(err, "updateStatus() should not error when pool Pods not found, got") } // TestReconcile_PatchError tests error path on Patch operations. @@ -524,9 +520,7 @@ func TestReconcile_PatchError(t *testing.T) { } err := tc.reconcileFunc(reconciler, context.Background(), shard) - if err == nil { - t.Errorf("reconcile function should error on Patch failure") - } + assert.NewCollecting(t).Error(err, "reconcile function should error on Patch failure") }) } } @@ -534,6 +528,7 @@ func TestReconcile_PatchError(t *testing.T) { // TestReconcile_PostgresSecretError verifies the error path in Reconcile when // reconcilePostgresPasswordSecret fails (lines 81-92 of shard_controller.go). func TestReconcile_PostgresSecretError(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -579,15 +574,14 @@ func TestReconcile_PostgresSecretError(t *testing.T) { } _, err := reconciler.Reconcile(t.Context(), req) - if err == nil { - t.Fatal("Reconcile should return an error when reconcilePostgresPasswordSecret fails") - } - if !strings.Contains( + c.Require(). + Error(err, "Reconcile should return an error when reconcilePostgresPasswordSecret fails") + c.StrContains( err.Error(), `failed to get postgres password Secret "missing-postgres-password"`, - ) { - t.Errorf("unexpected error: %v", err) - } + "unexpected error: %v", + err, + ) } // TestUpdateStatus_Multiorch tests updateStatus with different Multiorch deployment scenarios. @@ -674,6 +668,7 @@ func TestUpdateStatus_Multiorch(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -715,31 +710,19 @@ func TestUpdateStatus_Multiorch(t *testing.T) { } err := reconciler.updateStatus(context.Background(), shard, renderedConfig{}) - if tc.expectError && err == nil { - t.Error("updateStatus() should error but didn't") - } - if !tc.expectError && err != nil { - t.Errorf("updateStatus() unexpected error: %v", err) - } + c.False(tc.expectError && err == nil, "updateStatus() should error but didn't") + c.False(!tc.expectError && err != nil, "updateStatus() unexpected error: %v", err) // For non-error cases, verify OrchReady status if !tc.expectError { updatedShard := &multigresv1alpha1.Shard{} - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(shard), updatedShard, - ); err != nil { - t.Fatalf("Failed to get shard: %v", err) - } + ), "Failed to get shard") - if updatedShard.Status.OrchReady != tc.expectOrchReady { - t.Errorf( - "OrchReady = %v, want %v", - updatedShard.Status.OrchReady, - tc.expectOrchReady, - ) - } + c.Eq(tc.expectOrchReady, updatedShard.Status.OrchReady, "OrchReady") } }) } @@ -788,9 +771,7 @@ func TestUpdateStatus_GetError(t *testing.T) { } err := reconciler.updateStatus(context.Background(), shard, renderedConfig{}) - if err == nil { - t.Error("updateStatus() should error on Get failure") - } + assert.NewCollecting(t).Error(err, "updateStatus() should error on Get failure") } // statusPatchCapture wraps a client.Client to snapshot the state of the @@ -825,6 +806,7 @@ func (w *capturingStatusWriter) Patch( // TestUpdateStatus_FieldOwner verifies that the SSA status patch uses // "multigres-resource-handler" as the field owner, not "multigres-operator". func TestUpdateStatus_FieldOwner(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -861,23 +843,17 @@ func TestUpdateStatus_FieldOwner(t *testing.T) { } err := reconciler.updateStatus(context.Background(), shard, renderedConfig{}) - if err != nil { - t.Fatalf("updateStatus() unexpected error: %v", err) - } + c.Require().NoError(err, "updateStatus() unexpected error") // Verify the field owner is "multigres-resource-handler" foundFieldOwner := false for _, opt := range capture.capturedOpts { if fo, ok := opt.(client.FieldOwner); ok { - if string(fo) != "multigres-resource-handler" { - t.Errorf("field owner = %q, want %q", string(fo), "multigres-resource-handler") - } + c.Eq("multigres-resource-handler", string(fo), "field owner") foundFieldOwner = true } } - if !foundFieldOwner { - t.Error("no FieldOwner option found in Status().Patch() call") - } + c.True(foundFieldOwner, "no FieldOwner option found in Status().Patch() call") } // TestHandleScaleDown_ConcurrentDrainPrevention verifies that handleScaleDown @@ -1087,6 +1063,7 @@ func TestHandleScaleDown_ConcurrentDrainPrevention(t *testing.T) { for testName, tc := range tests { t.Run(testName, func(t *testing.T) { + c := assert.NewCollecting(t) shard := baseShard.DeepCopy() objects := make([]client.Object, 0, len(tc.pods)+1) @@ -1127,16 +1104,10 @@ func TestHandleScaleDown_ConcurrentDrainPrevention(t *testing.T) { tc.replicas, tc.actionTaken, ) - if err != nil { - t.Fatalf("handleScaleDown() unexpected error: %v", err) - } + c.Require().NoError(err, "handleScaleDown() unexpected error") - if gotAction != tc.wantAction { - t.Errorf("actionTaken = %v, want %v", gotAction, tc.wantAction) - } - if gotInProgress != tc.wantInProgress { - t.Errorf("inProgress = %v, want %v", gotInProgress, tc.wantInProgress) - } + c.Eq(tc.wantAction, gotAction, "actionTaken") + c.Eq(tc.wantInProgress, gotInProgress, "inProgress") for _, p := range tc.pods { updated := &corev1.Pod{} @@ -1189,9 +1160,7 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Metrics: metricsserver.Options{BindAddress: "0"}, }) - if err != nil { - t.Fatalf("Failed to create manager: %v", err) - } + assert.NewAborting(t).NoError(err, "Failed to create manager") return mgr } @@ -1203,9 +1172,7 @@ func TestSetupWithManager(t *testing.T) { Recorder: record.NewFakeRecorder(100), APIReader: mgr.GetClient(), } - if err := r.SetupWithManager(mgr); err != nil { - t.Errorf("SetupWithManager() error = %v", err) - } + assert.NewCollecting(t).NoError(r.SetupWithManager(mgr), "SetupWithManager() error =") }) t.Run("with options", func(t *testing.T) { @@ -1216,12 +1183,10 @@ func TestSetupWithManager(t *testing.T) { Recorder: record.NewFakeRecorder(100), APIReader: mgr.GetClient(), } - if err := r.SetupWithManager(mgr, controller.Options{ + assert.NewCollecting(t).NoError(r.SetupWithManager(mgr, controller.Options{ MaxConcurrentReconciles: 1, SkipNameValidation: ptr.To(true), - }); err != nil { - t.Errorf("SetupWithManager() with opts error = %v", err) - } + }), "SetupWithManager() with opts error =") }) } @@ -1270,6 +1235,7 @@ func TestClusterIsChurning(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) var objs []client.Object if tc.cluster != nil { @@ -1279,12 +1245,8 @@ func TestClusterIsChurning(t *testing.T) { r := &ShardReconciler{Client: fakeClient} got, err := r.clusterIsChurning(context.Background(), "default", tc.clusterName) - if err != nil { - t.Fatalf("clusterIsChurning() returned unexpected error: %v", err) - } - if got != tc.want { - t.Errorf("clusterIsChurning() = %v, want %v", got, tc.want) - } + c.Require().NoError(err, "clusterIsChurning() returned unexpected error") + c.Eq(tc.want, got, "clusterIsChurning()") }) } } @@ -1397,6 +1359,7 @@ func TestCleanupDrainedPod_PVCDeletion(t *testing.T) { for tn, tc := range tests { t.Run(tn, func(t *testing.T) { + c := assert.NewCollecting(t) shard := baseShard.DeepCopy() pod := makePod(tc.podName) @@ -1427,9 +1390,7 @@ func TestCleanupDrainedPod_PVCDeletion(t *testing.T) { err := reconciler.cleanupDrainedPod( context.Background(), shard, pod, poolName, poolSpec, replicas, ) - if err != nil { - t.Fatalf("cleanupDrainedPod() returned unexpected error: %v", err) - } + c.Require().NoError(err, "cleanupDrainedPod() returned unexpected error") pvcAfter := &corev1.PersistentVolumeClaim{} getErr := fakeClient.Get( @@ -1438,31 +1399,26 @@ func TestCleanupDrainedPod_PVCDeletion(t *testing.T) { pvcAfter, ) pvcExists := getErr == nil - if pvcExists != tc.wantPVC { - t.Fatalf( - "PVC %s exists = %v, want %v (err=%v)", - tc.pvcName, pvcExists, tc.wantPVC, getErr, - ) - } + c.Require(). + Eq(tc.wantPVC, pvcExists, "PVC %s exists = %v, want %v (err=%v)", tc.pvcName, pvcExists, tc.wantPVC, getErr) if !pvcExists { return } _, hasOrphanLabel := pvcAfter.Labels[metadata.LabelOrphan] - if hasOrphanLabel != tc.wantOrphan { - t.Errorf( - "PVC %s orphan-since present = %v, want %v", - tc.pvcName, hasOrphanLabel, tc.wantOrphan, - ) - } + c.Eq( + tc.wantOrphan, + hasOrphanLabel, + "PVC %s orphan-since present = %v, want", + tc.pvcName, + hasOrphanLabel, + ) podAfter := &corev1.Pod{} - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKey{Namespace: "default", Name: tc.podName}, podAfter, - ); err != nil { - t.Fatalf("failed to get pod after cleanup: %v", err) - } + ), "failed to get pod after cleanup") }) } } @@ -1486,6 +1442,7 @@ func TestHandleExternalDeletion(t *testing.T) { } t.Run("unscheduled pod is ignored", func(t *testing.T) { + ck := assert.NewAborting(t) pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ Name: "pod-unsched", @@ -1501,21 +1458,18 @@ func TestHandleExternalDeletion(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard, pod).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.handleExternalDeletion(context.Background(), shard, pod); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.NoError(r.handleExternalDeletion(context.Background(), shard, pod), "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } + ), "failed to get pod") }) t.Run("scheduled pod without drain annotation gets drain initiated", func(t *testing.T) { + ck := assert.NewCollecting(t) pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ Name: "pod-sched", @@ -1533,34 +1487,31 @@ func TestHandleExternalDeletion(t *testing.T) { rec := record.NewFakeRecorder(10) r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: rec} - if err := r.handleExternalDeletion(context.Background(), shard, pod); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.handleExternalDeletion(context.Background(), shard, pod), "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf("drain state = %q, want %q", - updated.Annotations[metadata.AnnotationDrainState], metadata.DrainStateRequested) - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updated.Annotations[metadata.AnnotationDrainState], + "drain state", + ) select { case event := <-rec.Events: - if !strings.Contains(event, "ExternalDeletion") { - t.Errorf("expected ExternalDeletion event, got %q", event) - } + ck.StrContains(event, "ExternalDeletion", "expected ExternalDeletion event, got") default: t.Error("expected ExternalDeletion event") } }) t.Run("scheduled pod with existing drain annotation is left alone", func(t *testing.T) { + ck := assert.NewCollecting(t) pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ Name: "pod-draining", @@ -1579,25 +1530,24 @@ func TestHandleExternalDeletion(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard, pod).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.handleExternalDeletion(context.Background(), shard, pod); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.handleExternalDeletion(context.Background(), shard, pod), "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateDraining { - t.Errorf("drain state changed to %q, should remain %q", - updated.Annotations[metadata.AnnotationDrainState], metadata.DrainStateDraining) - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateDraining, + updated.Annotations[metadata.AnnotationDrainState], + "drain state changed to", + ) }) t.Run("unscheduled pod without scheduled condition is ignored", func(t *testing.T) { + ck := assert.NewAborting(t) pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ Name: "pod-no-conditions", @@ -1609,18 +1559,14 @@ func TestHandleExternalDeletion(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard, pod).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.handleExternalDeletion(context.Background(), shard, pod); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.NoError(r.handleExternalDeletion(context.Background(), shard, pod), "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } + ), "failed to get pod") }) t.Run("error initiating drain for scheduled pod", func(t *testing.T) { @@ -1649,9 +1595,7 @@ func TestHandleExternalDeletion(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.handleExternalDeletion(context.Background(), shard, pod) - if err == nil { - t.Error("expected error when initiateDrain fails") - } + assert.NewCollecting(t).Error(err, "expected error when initiateDrain fails") }) } @@ -1674,9 +1618,8 @@ func TestReconcilePgBackRestCerts(t *testing.T) { APIReader: c, } - if err := r.reconcilePgBackRestCerts(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcilePgBackRestCerts(context.Background(), shard), "unexpected error") }) t.Run("user-provided secret with valid keys succeeds", func(t *testing.T) { @@ -1711,12 +1654,12 @@ func TestReconcilePgBackRestCerts(t *testing.T) { APIReader: c, } - if err := r.reconcilePgBackRestCerts(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcilePgBackRestCerts(context.Background(), shard), "unexpected error") }) t.Run("user-provided secret not found returns error", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -1740,15 +1683,12 @@ func TestReconcilePgBackRestCerts(t *testing.T) { } err := r.reconcilePgBackRestCerts(context.Background(), shard) - if err == nil { - t.Error("expected error for missing secret") - } - if !strings.Contains(err.Error(), "not found") { - t.Errorf("expected 'not found' error, got: %v", err) - } + ck.Error(err, "expected error for missing secret") + ck.StrContains(err.Error(), "not found", "expected 'not found' error, got: %v", err) }) t.Run("user-provided secret missing required key returns error", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -1781,12 +1721,13 @@ func TestReconcilePgBackRestCerts(t *testing.T) { } err := r.reconcilePgBackRestCerts(context.Background(), shard) - if err == nil { - t.Error("expected error for missing key") - } - if !strings.Contains(err.Error(), "tls.key") { - t.Errorf("expected error about missing 'tls.key', got: %v", err) - } + ck.Error(err, "expected error for missing key") + ck.StrContains( + err.Error(), + "tls.key", + "expected error about missing 'tls.key', got: %v", + err, + ) }) } @@ -1796,6 +1737,7 @@ func TestReconcileBackupCipherSecret(t *testing.T) { _ = corev1.AddToScheme(scheme) t.Run("nil backup config returns nil and creates no secret", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{}, @@ -1809,20 +1751,16 @@ func TestReconcileBackupCipherSecret(t *testing.T) { APIReader: c, } - if err := r.reconcileBackupCipherSecret(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.reconcileBackupCipherSecret(context.Background(), shard), "unexpected error") var secrets corev1.SecretList - if err := c.List(context.Background(), &secrets); err != nil { - t.Fatalf("failed to list secrets: %v", err) - } - if len(secrets.Items) != 0 { - t.Errorf("expected no secrets created, got %d", len(secrets.Items)) - } + ck.Require().NoError(c.List(context.Background(), &secrets), "failed to list secrets") + ck.Empty(secrets.Items, "expected no secrets created, got %d", len(secrets.Items)) }) t.Run("backup without encryption returns nil and creates no secret", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -1840,20 +1778,16 @@ func TestReconcileBackupCipherSecret(t *testing.T) { APIReader: c, } - if err := r.reconcileBackupCipherSecret(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.reconcileBackupCipherSecret(context.Background(), shard), "unexpected error") var secrets corev1.SecretList - if err := c.List(context.Background(), &secrets); err != nil { - t.Fatalf("failed to list secrets: %v", err) - } - if len(secrets.Items) != 0 { - t.Errorf("expected no secrets created, got %d", len(secrets.Items)) - } + ck.Require().NoError(c.List(context.Background(), &secrets), "failed to list secrets") + ck.Empty(secrets.Items, "expected no secrets created, got %d", len(secrets.Items)) }) t.Run("user-provided secret valid, no operator secret created", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -1884,20 +1818,21 @@ func TestReconcileBackupCipherSecret(t *testing.T) { APIReader: c, } - if err := r.reconcileBackupCipherSecret(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.reconcileBackupCipherSecret(context.Background(), shard), "unexpected error") var secrets corev1.SecretList - if err := c.List(context.Background(), &secrets); err != nil { - t.Fatalf("failed to list secrets: %v", err) - } - if len(secrets.Items) != 1 { - t.Errorf("expected only the user-provided secret to exist, got %d", len(secrets.Items)) - } + ck.Require().NoError(c.List(context.Background(), &secrets), "failed to list secrets") + ck.Len( + secrets.Items, + 1, + "expected only the user-provided secret to exist, got %d", + len(secrets.Items), + ) }) t.Run("user-provided secret not found returns error", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -1922,15 +1857,12 @@ func TestReconcileBackupCipherSecret(t *testing.T) { } err := r.reconcileBackupCipherSecret(context.Background(), shard) - if err == nil { - t.Error("expected error for missing secret") - } - if !strings.Contains(err.Error(), "not found") { - t.Errorf("expected 'not found' error, got: %v", err) - } + ck.Error(err, "expected error for missing secret") + ck.StrContains(err.Error(), "not found", "expected 'not found' error, got: %v", err) }) t.Run("user-provided secret missing required key returns error", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -1962,12 +1894,14 @@ func TestReconcileBackupCipherSecret(t *testing.T) { } err := r.reconcileBackupCipherSecret(context.Background(), shard) - if err == nil { - t.Error("expected error for missing key") - } - if !strings.Contains(err.Error(), PgBackRestCipherKeyDataKey) { - t.Errorf("expected error about missing %q, got: %v", PgBackRestCipherKeyDataKey, err) - } + ck.Error(err, "expected error for missing key") + ck.StrContains( + err.Error(), + PgBackRestCipherKeyDataKey, + "expected error about missing %q, got: %v", + PgBackRestCipherKeyDataKey, + err, + ) }) t.Run("empty secret name returns error", func(t *testing.T) { @@ -1993,9 +1927,7 @@ func TestReconcileBackupCipherSecret(t *testing.T) { } err := r.reconcileBackupCipherSecret(context.Background(), shard) - if err == nil { - t.Error("expected error for empty secret name") - } + assert.NewCollecting(t).Error(err, "expected error for empty secret name") }) } @@ -2024,6 +1956,7 @@ func TestCreateMissingResources(t *testing.T) { cellName := "zone1" t.Run("terminal pod (Failed) is deleted for recreation", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, 0) @@ -2055,12 +1988,8 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 1, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken for terminal pod deletion") - } + ck.Require().NoError(err, "unexpected error") + ck.True(actionTaken, "expected actionTaken for terminal pod deletion") // Pod should be deleted err = c.Get( @@ -2068,12 +1997,11 @@ func TestCreateMissingResources(t *testing.T) { types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, ) - if !errors.IsNotFound(err) { - t.Errorf("terminal pod should be deleted, but Get returned: %v", err) - } + ck.True(errors.IsNotFound(err), "terminal pod should be deleted, but Get returned: %v", err) }) t.Run("reused orphan PVC has orphan-since label cleared", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, 0) @@ -2106,18 +2034,14 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 1, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") got := &corev1.PersistentVolumeClaim{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), types.NamespacedName{Name: pvcName, Namespace: "default"}, got, - ); err != nil { - t.Fatalf("PVC must still exist: %v", err) - } + ), "PVC must still exist") if _, ok := got.Labels[metadata.LabelOrphan]; ok { t.Errorf( "orphan-since label should be cleared when PVC is reused, still present: %v", @@ -2125,16 +2049,15 @@ func TestCreateMissingResources(t *testing.T) { ) } // pod (re)created on the reused PVC - if err := c.Get( + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("pod should be created on reused PVC: %v", err) - } + ), "pod should be created on reused PVC") }) t.Run("terminal pod (Succeeded) is deleted for recreation", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, 0) @@ -2163,17 +2086,14 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 1, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken for terminal pod deletion") - } + ck.Require().NoError(err, "unexpected error") + ck.True(actionTaken, "expected actionTaken for terminal pod deletion") }) t.Run( "externally deleted pod (DeletionTimestamp + finalizer) calls handleExternalDeletion", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, 0) @@ -2213,33 +2133,26 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 1, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken for externally deleted pod") - } + ck.Require().NoError(err, "unexpected error") + ck.True(actionTaken, "expected actionTaken for externally deleted pod") // Should have initiated drain updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf( - "drain state = %q, want %q", - updated.Annotations[metadata.AnnotationDrainState], - metadata.DrainStateRequested, - ) - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updated.Annotations[metadata.AnnotationDrainState], + "drain state", + ) }, ) t.Run("not-ready pod does not block creation of other replicas", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName0 := BuildPoolPodName(shard, poolName, cellName, 0) podName1 := BuildPoolPodName(shard, poolName, cellName, 1) @@ -2284,31 +2197,24 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 2, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken for pod creation") - } + ck.Require().NoError(err, "unexpected error") + ck.True(actionTaken, "expected actionTaken for pod creation") // Pod 1 SHOULD be created even though pod 0 is not ready - if err := c.Get( + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: podName1, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("pod-1 should have been created despite pod-0 being not ready: %v", err) - } - if err := c.Get( + ), "pod-1 should have been created despite pod-0 being not ready") + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: pvcName1, Namespace: "default"}, &corev1.PersistentVolumeClaim{}, - ); err != nil { - t.Errorf("pvc-1 should have been created despite pod-0 being not ready: %v", err) - } + ), "pvc-1 should have been created despite pod-0 being not ready") }) t.Run("all missing pods and PVCs created in one pass", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard).Build() @@ -2324,34 +2230,27 @@ func TestCreateMissingResources(t *testing.T) { map[string]*corev1.PersistentVolumeClaim{}, 3, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken for pod creation") - } + ck.Require().NoError(err, "unexpected error") + ck.True(actionTaken, "expected actionTaken for pod creation") for i := 0; i < 3; i++ { podName := BuildPoolPodName(shard, poolName, cellName, i) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, i) - if err := c.Get( + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("pod-%d should exist: %v", i, err) - } - if err := c.Get( + ), "pod-%d should exist", i) + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: pvcName, Namespace: "default"}, &corev1.PersistentVolumeClaim{}, - ); err != nil { - t.Errorf("pvc-%d should exist: %v", i, err) - } + ), "pvc-%d should exist", i) } }) t.Run("actionTaken blocks terminal pod deletion", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName0 := BuildPoolPodName(shard, poolName, cellName, 0) podName1 := BuildPoolPodName(shard, poolName, cellName, 1) @@ -2392,9 +2291,7 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 2, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // Only one of the two terminal pods should have been deleted (actionTaken blocks second) deleted := 0 @@ -2409,12 +2306,11 @@ func TestCreateMissingResources(t *testing.T) { deleted++ } } - if deleted != 1 { - t.Errorf("expected exactly 1 terminal pod deleted (sequential gating), got %d", deleted) - } + ck.Eq(1, deleted, "expected exactly 1 terminal pod deleted (sequential gating), got") }) t.Run("missing pod is created", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, 0) @@ -2435,28 +2331,20 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 1, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken for pod creation") - } + ck.Require().NoError(err, "unexpected error") + ck.True(actionTaken, "expected actionTaken for pod creation") // Pod and PVC should exist - if err := c.Get( + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("pod should exist: %v", err) - } - if err := c.Get( + ), "pod should exist") + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: pvcName, Namespace: "default"}, &corev1.PersistentVolumeClaim{}, - ); err != nil { - t.Errorf("PVC should exist: %v", err) - } + ), "PVC should exist") }) t.Run("error creating pod", func(t *testing.T) { @@ -2477,9 +2365,7 @@ func TestCreateMissingResources(t *testing.T) { context.Background(), shard, poolName, cellName, poolSpec, map[string]*corev1.Pod{}, map[string]*corev1.PersistentVolumeClaim{}, 1, ) - if err == nil { - t.Error("expected error on pod create failure") - } + assert.NewCollecting(t).Error(err, "expected error on pod create failure") }) t.Run("error creating PVC", func(t *testing.T) { @@ -2500,9 +2386,7 @@ func TestCreateMissingResources(t *testing.T) { context.Background(), shard, poolName, cellName, poolSpec, map[string]*corev1.Pod{}, map[string]*corev1.PersistentVolumeClaim{}, 1, ) - if err == nil { - t.Error("expected error on PVC create failure") - } + assert.NewCollecting(t).Error(err, "expected error on PVC create failure") }) t.Run("error deleting terminal pod", func(t *testing.T) { @@ -2539,9 +2423,7 @@ func TestCreateMissingResources(t *testing.T) { existingPVCs, 1, ) - if err == nil { - t.Error("expected error on terminal pod delete failure") - } + assert.NewCollecting(t).Error(err, "expected error on terminal pod delete failure") }) } @@ -2565,32 +2447,20 @@ func TestResolvePodIndex(t *testing.T) { for tn, tc := range tests { t.Run(tn, func(t *testing.T) { got, ok := resolvePodIndex(tc.podName) - if got != tc.want || ok != tc.wantOK { - t.Errorf( - "resolvePodIndex(%q) = (%d, %v), want (%d, %v)", - tc.podName, - got, - ok, - tc.want, - tc.wantOK, - ) - } + assert.NewCollecting(t). + False(got != tc.want || ok != tc.wantOK, "resolvePodIndex(%q) = (%d, %v), want (%d, %v)", tc.podName, got, ok, tc.want, tc.wantOK) }) } } func TestIsPodReady(t *testing.T) { t.Run("nil pod", func(t *testing.T) { - if isPodReady(nil) { - t.Error("expected false for nil pod") - } + assert.NewCollecting(t).False(isPodReady(nil), "expected false for nil pod") }) t.Run("no conditions", func(t *testing.T) { pod := &corev1.Pod{} - if isPodReady(pod) { - t.Error("expected false for pod with no conditions") - } + assert.NewCollecting(t).False(isPodReady(pod), "expected false for pod with no conditions") }) t.Run("ready condition false", func(t *testing.T) { @@ -2601,9 +2471,7 @@ func TestIsPodReady(t *testing.T) { }, }, } - if isPodReady(pod) { - t.Error("expected false for pod with PodReady=False") - } + assert.NewCollecting(t).False(isPodReady(pod), "expected false for pod with PodReady=False") }) t.Run("ready condition true", func(t *testing.T) { @@ -2614,9 +2482,7 @@ func TestIsPodReady(t *testing.T) { }, }, } - if !isPodReady(pod) { - t.Error("expected true for pod with PodReady=True") - } + assert.NewCollecting(t).True(isPodReady(pod), "expected true for pod with PodReady=True") }) t.Run("only non-ready conditions", func(t *testing.T) { @@ -2628,9 +2494,8 @@ func TestIsPodReady(t *testing.T) { }, }, } - if isPodReady(pod) { - t.Error("expected false for pod with no PodReady condition") - } + assert.NewCollecting(t). + False(isPodReady(pod), "expected false for pod with no PodReady condition") }) } @@ -2643,15 +2508,13 @@ func TestIsPoolHealthy(t *testing.T) { } t.Run("empty pool is unhealthy when replicas are expected", func(t *testing.T) { - if isPoolHealthy(map[string]*corev1.Pod{}, 1, shard) { - t.Error("expected empty pool with 1 expected replica to be unhealthy") - } + assert.NewCollecting(t). + False(isPoolHealthy(map[string]*corev1.Pod{}, 1, shard), "expected empty pool with 1 expected replica to be unhealthy") }) t.Run("empty pool is healthy when no replicas are expected", func(t *testing.T) { - if !isPoolHealthy(map[string]*corev1.Pod{}, 0, shard) { - t.Error("expected empty pool with 0 expected replicas to be healthy") - } + assert.NewCollecting(t). + True(isPoolHealthy(map[string]*corev1.Pod{}, 0, shard), "expected empty pool with 0 expected replicas to be healthy") }) t.Run("draining pod makes pool unhealthy", func(t *testing.T) { @@ -2670,9 +2533,8 @@ func TestIsPoolHealthy(t *testing.T) { }, }, } - if isPoolHealthy(pods, 1, shard) { - t.Error("draining pod should make the pool unhealthy") - } + assert.NewCollecting(t). + False(isPoolHealthy(pods, 1, shard), "draining pod should make the pool unhealthy") }) t.Run("pod being deleted makes pool unhealthy", func(t *testing.T) { @@ -2691,9 +2553,8 @@ func TestIsPoolHealthy(t *testing.T) { }, }, } - if isPoolHealthy(pods, 1, shard) { - t.Error("terminating pod should make the pool unhealthy") - } + assert.NewCollecting(t). + False(isPoolHealthy(pods, 1, shard), "terminating pod should make the pool unhealthy") }) t.Run("extra pod (index >= replicas) is excluded from health check", func(t *testing.T) { @@ -2715,9 +2576,8 @@ func TestIsPoolHealthy(t *testing.T) { }, }, } - if !isPoolHealthy(pods, 1, shard) { - t.Error("extra pod at index 2 with replicas=1 should be excluded from health check") - } + assert.NewCollecting(t). + True(isPoolHealthy(pods, 1, shard), "extra pod at index 2 with replicas=1 should be excluded from health check") }) t.Run("QUARANTINED pod that is not ready does not block health check", func(t *testing.T) { @@ -2733,9 +2593,8 @@ func TestIsPoolHealthy(t *testing.T) { }, }, } - if !isPoolHealthy(pods, 1, quarantinedShard) { - t.Error("QUARANTINED pod should be excluded from health check") - } + assert.NewCollecting(t). + True(isPoolHealthy(pods, 1, quarantinedShard), "QUARANTINED pod should be excluded from health check") }) } @@ -2753,18 +2612,16 @@ func TestPodNeedsUpdate(t *testing.T) { Annotations: map[string]string{metadata.AnnotationSpecHash: "old"}, }, } - if podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s) { - t.Error("pod with deletion timestamp should not need update") - } + assert.NewCollecting(t). + False(podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s), "pod with deletion timestamp should not need update") }) t.Run("pod missing spec-hash annotation needs update", func(t *testing.T) { pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{}, } - if !podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s) { - t.Error("pod without spec-hash annotation should need update") - } + assert.NewCollecting(t). + True(podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s), "pod without spec-hash annotation should need update") }) t.Run("pod with matching spec-hash does not need update", func(t *testing.T) { @@ -2775,9 +2632,8 @@ func TestPodNeedsUpdate(t *testing.T) { Annotations: map[string]string{metadata.AnnotationSpecHash: hash}, }, } - if podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s) { - t.Error("pod with matching spec-hash should not need update") - } + assert.NewCollecting(t). + False(podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s), "pod with matching spec-hash should not need update") }) t.Run("pod with old spec-hash needs update", func(t *testing.T) { @@ -2786,9 +2642,8 @@ func TestPodNeedsUpdate(t *testing.T) { Annotations: map[string]string{metadata.AnnotationSpecHash: "old-hash"}, }, } - if !podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s) { - t.Error("pod with old spec-hash should need update") - } + assert.NewCollecting(t). + True(podNeedsUpdate(pod, shard, "main", "z1", pool, 0, s), "pod with old spec-hash should need update") }) t.Run("build error assumes no update needed", func(t *testing.T) { @@ -2798,9 +2653,8 @@ func TestPodNeedsUpdate(t *testing.T) { Annotations: map[string]string{metadata.AnnotationSpecHash: "some-hash"}, }, } - if podNeedsUpdate(pod, shard, "main", "z1", pool, 0, emptyScheme) { - t.Error("build failure should assume no update needed") - } + assert.NewCollecting(t). + False(podNeedsUpdate(pod, shard, "main", "z1", pool, 0, emptyScheme), "build failure should assume no update needed") }) } @@ -2810,6 +2664,7 @@ func TestInitiateDrain(t *testing.T) { _ = corev1.AddToScheme(scheme) t.Run("sets drain annotations on pod with nil annotations", func(t *testing.T) { + ck := assert.NewCollecting(t) pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ Name: "pod-nil-ann", @@ -2820,25 +2675,24 @@ func TestInitiateDrain(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(pod).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.initiateDrain(context.Background(), pod); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(r.initiateDrain(context.Background(), pod), "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf("drain state = %q, want %q", - updated.Annotations[metadata.AnnotationDrainState], metadata.DrainStateRequested) - } - if updated.Annotations[metadata.AnnotationDrainRequestedAt] == "" { - t.Error("drain requested-at timestamp should be set") - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updated.Annotations[metadata.AnnotationDrainState], + "drain state", + ) + ck.NotEq( + "", + updated.Annotations[metadata.AnnotationDrainRequestedAt], + "drain requested-at timestamp should be set", + ) }) t.Run("error on patch failure", func(t *testing.T) { @@ -2858,9 +2712,7 @@ func TestInitiateDrain(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.initiateDrain(context.Background(), pod) - if err == nil { - t.Error("expected error on patch failure") - } + assert.NewCollecting(t).Error(err, "expected error on patch failure") }) } @@ -2890,6 +2742,7 @@ func TestHandleRollingUpdates(t *testing.T) { } t.Run("no drifted pods sets RollingUpdate condition to false", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} @@ -2899,25 +2752,20 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{}, 0, false, false, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") found := false for _, cond := range shard.Status.Conditions { if cond.Type == "RollingUpdate" { found = true - if cond.Status != metav1.ConditionFalse { - t.Errorf("RollingUpdate condition status = %s, want False", cond.Status) - } + ck.Eq(metav1.ConditionFalse, cond.Status, "RollingUpdate condition status") } } - if !found { - t.Error("RollingUpdate condition not set") - } + ck.True(found, "RollingUpdate condition not set") }) t.Run("actionTaken skips drain initiation", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pod := &corev1.Pod{ @@ -2938,25 +2786,24 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, true, false, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // Pod should NOT have drain annotation (actionTaken blocks) updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != "" { - t.Error("drain annotation should not be set when actionTaken is true") - } + ), "failed to get pod") + ck.Eq( + "", + updated.Annotations[metadata.AnnotationDrainState], + "drain annotation should not be set when actionTaken is true", + ) }) t.Run("isAnyPodDraining skips drain initiation", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) pod := &corev1.Pod{ @@ -2977,24 +2824,23 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, false, true, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != "" { - t.Error("drain annotation should not be set when isAnyPodDraining is true") - } + ), "failed to get pod") + ck.Eq( + "", + updated.Annotations[metadata.AnnotationDrainState], + "drain annotation should not be set when isAnyPodDraining is true", + ) }) t.Run("primary-only drift initiates drain for primary", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) shard.Status.PodRoles = map[string]string{podName: "PRIMARY"} @@ -3018,25 +2864,23 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, false, false, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf("primary pod should have drain requested, got %q", - updated.Annotations[metadata.AnnotationDrainState]) - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updated.Annotations[metadata.AnnotationDrainState], + "primary pod should have drain requested, got", + ) }) t.Run("primary with existing drain annotation is skipped", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 0) shard.Status.PodRoles = map[string]string{podName: "PRIMARY"} @@ -3060,26 +2904,24 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, false, false, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod), updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } + ), "failed to get pod") // Should still be "draining", not changed to "requested" - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateDraining { - t.Errorf("primary pod drain state should remain %q, got %q", - metadata.DrainStateDraining, updated.Annotations[metadata.AnnotationDrainState]) - } + ck.Eq( + metadata.DrainStateDraining, + updated.Annotations[metadata.AnnotationDrainState], + "primary pod drain state should remain", + ) }) t.Run("replica is drained before primary", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName0 := BuildPoolPodName(shard, poolName, cellName, 0) podName1 := BuildPoolPodName(shard, poolName, cellName, 1) @@ -3110,37 +2952,33 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName0: pod0, podName1: pod1}, 2, false, false, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // Replica should have drain requested updated0 := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod0), updated0, - ); err != nil { - t.Fatalf("failed to get pod0: %v", err) - } - if updated0.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf("replica pod should have drain requested, got %q", - updated0.Annotations[metadata.AnnotationDrainState]) - } + ), "failed to get pod0") + ck.Eq( + metadata.DrainStateRequested, + updated0.Annotations[metadata.AnnotationDrainState], + "replica pod should have drain requested, got", + ) // Primary should NOT have drain annotation (replica first) updated1 := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(pod1), updated1, - ); err != nil { - t.Fatalf("failed to get pod1: %v", err) - } - if updated1.Annotations[metadata.AnnotationDrainState] != "" { - t.Errorf("primary pod should NOT have drain annotation yet, got %q", - updated1.Annotations[metadata.AnnotationDrainState]) - } + ), "failed to get pod1") + ck.Eq( + "", + updated1.Annotations[metadata.AnnotationDrainState], + "primary pod should NOT have drain annotation yet, got", + ) }) t.Run("error initiating drain for replica", func(t *testing.T) { @@ -3169,9 +3007,7 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, false, false, &shardRolloutTracker{}, ) - if err == nil { - t.Error("expected error when drain initiation fails") - } + assert.NewCollecting(t).Error(err, "expected error when drain initiation fails") }) t.Run("error initiating drain for primary", func(t *testing.T) { @@ -3200,9 +3036,7 @@ func TestHandleRollingUpdates(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, false, false, &shardRolloutTracker{}, ) - if err == nil { - t.Error("expected error when primary drain initiation fails") - } + assert.NewCollecting(t).Error(err, "expected error when primary drain initiation fails") }) } @@ -3210,6 +3044,7 @@ func TestHandleRollingUpdates(t *testing.T) { // rollout tracker already marked started blocks a new drain even when // isShardHealthy would otherwise say the shard is healthy. func TestHandleRollingUpdates_RolloutTrackerBlocksSameShardPass(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -3262,19 +3097,16 @@ func TestHandleRollingUpdates_RolloutTrackerBlocksSameShardPass(t *testing.T) { map[string]*corev1.Pod{podName: pod}, 1, false, false, rollout, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") updated := &corev1.Pod{} - if err := c.Get(context.Background(), client.ObjectKeyFromObject(pod), updated); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != "" { - t.Error( - "drain annotation should not be set when the rollout tracker already started this pass", - ) - } + ck.Require(). + NoError(c.Get(context.Background(), client.ObjectKeyFromObject(pod), updated), "failed to get pod") + ck.Eq( + "", + updated.Annotations[metadata.AnnotationDrainState], + "drain annotation should not be set when the rollout tracker already started this pass", + ) } // TestIsShardHealthy_MissingPodCountsAsUnhealthy verifies that a pool/cell @@ -3337,9 +3169,7 @@ func TestIsShardHealthy_MissingPodCountsAsUnhealthy(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} healthy, err := r.isShardHealthy(context.Background(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t).NoError(err, "unexpected error") if healthy { t.Error( "isShardHealthy should be false when pool-1/zone1 has no pod at all, " + @@ -3369,12 +3199,12 @@ func TestReconcileSharedBackupPVC(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.reconcileSharedBackupPVC(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(r.reconcileSharedBackupPVC(context.Background(), shard), "unexpected error") }) t.Run("nil backup creates PVC with defaults", func(t *testing.T) { + ck := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -3390,19 +3220,15 @@ func TestReconcileSharedBackupPVC(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.reconcileSharedBackupPVC(context.Background(), shard); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.NoError(r.reconcileSharedBackupPVC(context.Background(), shard), "unexpected error") pvcName := BuildSharedBackupPVCName(shard) pvc := &corev1.PersistentVolumeClaim{} - if err := c.Get( + ck.NoError(c.Get( context.Background(), types.NamespacedName{Name: pvcName, Namespace: "default"}, pvc, - ); err != nil { - t.Fatalf("PVC should exist: %v", err) - } + ), "PVC should exist") }) t.Run("error on patch failure", func(t *testing.T) { @@ -3430,14 +3256,13 @@ func TestReconcileSharedBackupPVC(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.reconcileSharedBackupPVC(context.Background(), shard) - if err == nil { - t.Error("expected error on PVC patch failure") - } + assert.NewCollecting(t).Error(err, "expected error on PVC patch failure") }) } func TestBuildSharedBackupPVC_Variants(t *testing.T) { t.Run("filesystem backup with custom storage class", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -3461,12 +3286,8 @@ func TestBuildSharedBackupPVC_Variants(t *testing.T) { } pvc, err := BuildSharedBackupPVC(shard, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pvc == nil { - t.Fatal("expected non-nil PVC") - } + c.NoError(err, "unexpected error") + c.NotNil(pvc, "expected non-nil PVC") if pvc.Spec.StorageClassName == nil || *pvc.Spec.StorageClassName != "premium-ssd" { t.Errorf("storage class = %v, want premium-ssd", pvc.Spec.StorageClassName) } @@ -3476,6 +3297,7 @@ func TestBuildSharedBackupPVC_Variants(t *testing.T) { }) t.Run("nil filesystem config uses defaults", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", Namespace: "default", @@ -3493,15 +3315,9 @@ func TestBuildSharedBackupPVC_Variants(t *testing.T) { } pvc, err := BuildSharedBackupPVC(shard, false, testScheme()) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pvc == nil { - t.Fatal("expected non-nil PVC") - } - if pvc.Spec.StorageClassName != nil { - t.Errorf("storage class should be nil by default, got %v", pvc.Spec.StorageClassName) - } + c.Require().NoError(err, "unexpected error") + c.Require().NotNil(pvc, "expected non-nil PVC") + c.Nil(pvc.Spec.StorageClassName, "storage class should be nil by default, got") }) } @@ -3556,9 +3372,7 @@ func TestCleanupDrainedPod_ErrorPaths(t *testing.T) { } err := r.cleanupDrainedPod(context.Background(), shard, pod, poolName, poolSpec, 3) - if err == nil { - t.Error("expected error on PVC Get failure") - } + assert.NewCollecting(t).Error(err, "expected error on PVC Get failure") }) t.Run("error orphaning PVC", func(t *testing.T) { @@ -3596,12 +3410,11 @@ func TestCleanupDrainedPod_ErrorPaths(t *testing.T) { } err := r.cleanupDrainedPod(context.Background(), shard, pod, poolName, poolSpec, 3) - if err == nil { - t.Error("expected error on PVC orphan-patch failure") - } + assert.NewCollecting(t).Error(err, "expected error on PVC orphan-patch failure") }) t.Run("nil PVC deletion policy orphans PVC", func(t *testing.T) { + ck := assert.NewCollecting(t) shard := baseShard.DeepCopy() podName := BuildPoolPodName(shard, poolName, cellName, 5) pvcName := BuildPoolDataPVCName(shard, poolName, cellName, 5) @@ -3630,21 +3443,16 @@ func TestCleanupDrainedPod_ErrorPaths(t *testing.T) { multigresv1alpha1.PoolSpec{}, 3, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") got := &corev1.PersistentVolumeClaim{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), types.NamespacedName{Name: pvcName, Namespace: "default"}, got, - ); err != nil { - t.Fatalf("PVC must still exist, got err: %v", err) - } - if _, ok := got.Labels[metadata.LabelOrphan]; !ok { - t.Errorf("PVC %s missing orphan-since label", pvcName) - } + ), "PVC must still exist, got err") + _, ok := got.Labels[metadata.LabelOrphan] + ck.True(ok, "PVC %s missing orphan-since label", pvcName) }) } @@ -3692,9 +3500,7 @@ func TestReconcilePoolPods_ErrorPaths(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err == nil { - t.Error("expected error on pod list failure") - } + assert.NewCollecting(t).Error(err, "expected error on pod list failure") }) t.Run("error listing PVCs", func(t *testing.T) { @@ -3718,9 +3524,7 @@ func TestReconcilePoolPods_ErrorPaths(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err == nil { - t.Error("expected error on PVC list failure") - } + assert.NewCollecting(t).Error(err, "expected error on PVC list failure") }) } @@ -3789,9 +3593,8 @@ func TestHandleScaleDown_ErrorPaths(t *testing.T) { context.Background(), shard, poolName, multigresv1alpha1.PoolSpec{}, existingPods, 1, 1, false, ) - if err == nil { - t.Error("expected error on drain initiation failure for extra pod") - } + assert.NewCollecting(t). + Error(err, "expected error on drain initiation failure for extra pod") }) t.Run("error deleting ready-for-deletion pod after cleanup", func(t *testing.T) { @@ -3825,9 +3628,7 @@ func TestHandleScaleDown_ErrorPaths(t *testing.T) { context.Background(), shard, poolName, multigresv1alpha1.PoolSpec{}, existingPods, 1, 1, false, ) - if err == nil { - t.Error("expected error on pod delete failure after cleanup") - } + assert.NewCollecting(t).Error(err, "expected error on pod delete failure after cleanup") }) t.Run("error handling external deletion of extra pod", func(t *testing.T) { @@ -3863,9 +3664,7 @@ func TestHandleScaleDown_ErrorPaths(t *testing.T) { context.Background(), shard, poolName, multigresv1alpha1.PoolSpec{}, existingPods, 1, 1, false, ) - if err == nil { - t.Error("expected error on external deletion handling failure") - } + assert.NewCollecting(t).Error(err, "expected error on external deletion handling failure") }) } @@ -3875,9 +3674,7 @@ func TestSelectPodToDrain_NilPod(t *testing.T) { t.Run("empty list returns nil", func(t *testing.T) { result := r.selectPodToDrain(context.Background(), []*corev1.Pod{}, shard) - if result != nil { - t.Errorf("expected nil for empty list, got %v", result) - } + assert.NewCollecting(t).Nil(result, "expected nil for empty list, got") }) t.Run("nil entries are skipped", func(t *testing.T) { @@ -3893,13 +3690,13 @@ func TestSelectPodToDrain_NilPod(t *testing.T) { }, } result := r.selectPodToDrain(context.Background(), pods, shard) - if result == nil || result.Name != "pod-0" { - t.Errorf("expected pod-0, got %v", result) - } + assert.NewCollecting(t). + False(result == nil || result.Name != "pod-0", "expected pod-0, got %v", result) }) } func TestReconcile_BackupCertsError(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -3946,12 +3743,8 @@ func TestReconcile_BackupCertsError(t *testing.T) { req := ctrl.Request{NamespacedName: client.ObjectKeyFromObject(shard)} _, err := reconciler.Reconcile(t.Context(), req) - if err == nil { - t.Error("expected error when pgBackRest TLS secret not found") - } - if !strings.Contains(err.Error(), "not found") { - t.Errorf("expected 'not found' error, got: %v", err) - } + ck.Error(err, "expected error when pgBackRest TLS secret not found") + ck.StrContains(err.Error(), "not found", "expected 'not found' error, got: %v", err) } func TestResolvePodRole_FQDNPrefix(t *testing.T) { @@ -3963,12 +3756,11 @@ func TestResolvePodRole_FQDNPrefix(t *testing.T) { }, } role := resolvePodRole(shard, "my-pod") - if role != "PRIMARY" { - t.Errorf("expected PRIMARY via FQDN prefix match, got %q", role) - } + assert.NewCollecting(t).Eq("PRIMARY", role, "expected PRIMARY via FQDN prefix match, got") } func TestBuildSharedBackupPVC_InvalidStorageSize(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -3992,12 +3784,13 @@ func TestBuildSharedBackupPVC_InvalidStorageSize(t *testing.T) { scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) _, err := BuildSharedBackupPVC(shard, false, scheme) - if err == nil { - t.Fatal("expected error for invalid storage size") - } - if !strings.Contains(err.Error(), "invalid storage size") { - t.Errorf("expected 'invalid storage size' error, got: %v", err) - } + c.Require().Error(err, "expected error for invalid storage size") + c.StrContains( + err.Error(), + "invalid storage size", + "expected 'invalid storage size' error, got: %v", + err, + ) } func TestReconcilePoolPods_ErrorPropagation(t *testing.T) { @@ -4044,9 +3837,7 @@ func TestReconcilePoolPods_ErrorPropagation(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err == nil { - t.Fatal("expected error from createMissingResources") - } + assert.NewAborting(t).Error(err, "expected error from createMissingResources") }) t.Run("handleScaleDown error propagates", func(t *testing.T) { @@ -4111,9 +3902,7 @@ func TestReconcilePoolPods_ErrorPropagation(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err == nil { - t.Fatal("expected error from handleScaleDown") - } + assert.NewAborting(t).Error(err, "expected error from handleScaleDown") }) t.Run("handleRollingUpdates error propagates", func(t *testing.T) { @@ -4160,13 +3949,12 @@ func TestReconcilePoolPods_ErrorPropagation(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err == nil { - t.Fatal("expected error from handleRollingUpdates") - } + assert.NewAborting(t).Error(err, "expected error from handleRollingUpdates") }) } func TestCreateMissingResources_PodBuildError(t *testing.T) { + c := assert.NewCollecting(t) emptyScheme := runtime.NewScheme() shard := &multigresv1alpha1.Shard{ @@ -4218,12 +4006,13 @@ func TestCreateMissingResources_PodBuildError(t *testing.T) { t.Context(), shardCopy, poolName, cellName, poolSpec, map[string]*corev1.Pod{}, existingPVCs, 1, ) - if err == nil { - t.Fatal("expected error from BuildPoolPod with empty scheme") - } - if !strings.Contains(err.Error(), "failed to build pod") { - t.Errorf("expected 'failed to build pod' error, got: %v", err) - } + c.Require().Error(err, "expected error from BuildPoolPod with empty scheme") + c.StrContains( + err.Error(), + "failed to build pod", + "expected 'failed to build pod' error, got: %v", + err, + ) } func TestCreateMissingResources_ExternalDeletionError(t *testing.T) { @@ -4289,12 +4078,12 @@ func TestCreateMissingResources_ExternalDeletionError(t *testing.T) { t.Context(), shard, poolName, cellName, poolSpec, existingPods, existingPVCs, 1, ) - if err == nil { - t.Fatal("expected error from handleExternalDeletion in createMissingResources") - } + assert.NewAborting(t). + Error(err, "expected error from handleExternalDeletion in createMissingResources") } func TestHandleScaleDown_ExternalDeletionSetsActionTaken(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4340,15 +4129,12 @@ func TestHandleScaleDown_ExternalDeletionSetsActionTaken(t *testing.T) { multigresv1alpha1.PoolSpec{ReplicasPerCell: ptr.To(int32(1))}, existingPods, 1, 1, false, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !actionTaken { - t.Error("expected actionTaken=true after external deletion of extra pod") - } + c.Require().NoError(err, "unexpected error") + c.True(actionTaken, "expected actionTaken=true after external deletion of extra pod") } func TestHandleRollingUpdates_SkipsUpToDatePods(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4420,24 +4206,19 @@ func TestHandleRollingUpdates_SkipsUpToDatePods(t *testing.T) { existingPods, 1, false, false, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require().NoError(err, "unexpected error") var updatedPod1 corev1.Pod - if err := base.Get( + c.Require().NoError(base.Get( t.Context(), types.NamespacedName{Name: pod1Name, Namespace: "default"}, &updatedPod1, - ); err != nil { - t.Fatalf("failed to get pod1: %v", err) - } - if updatedPod1.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf( - "expected drifted pod1 to have drain requested, got: %q", - updatedPod1.Annotations[metadata.AnnotationDrainState], - ) - } + ), "failed to get pod1") + c.Eq( + metadata.DrainStateRequested, + updatedPod1.Annotations[metadata.AnnotationDrainState], + "expected drifted pod1 to have drain requested, got", + ) } func TestSelectPodToDrain_AllNilEntries(t *testing.T) { @@ -4445,9 +4226,7 @@ func TestSelectPodToDrain_AllNilEntries(t *testing.T) { shard := &multigresv1alpha1.Shard{} pods := []*corev1.Pod{nil, nil, nil} result := r.selectPodToDrain(t.Context(), pods, shard) - if result != nil { - t.Errorf("expected nil when all entries are nil, got %v", result) - } + assert.NewCollecting(t).Nil(result, "expected nil when all entries are nil, got") } func TestReconcileSharedBackupPVC_BuildErrorAndNilReturn(t *testing.T) { @@ -4456,6 +4235,7 @@ func TestReconcileSharedBackupPVC_BuildErrorAndNilReturn(t *testing.T) { _ = corev1.AddToScheme(scheme) t.Run("build error propagates", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard", @@ -4479,16 +4259,18 @@ func TestReconcileSharedBackupPVC_BuildErrorAndNilReturn(t *testing.T) { base := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard.DeepCopy()).Build() r := &ShardReconciler{Client: base, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.reconcileSharedBackupPVC(t.Context(), shard) - if err == nil { - t.Fatal("expected error from build failure") - } - if !strings.Contains(err.Error(), "failed to build shared backup PVC") { - t.Errorf("expected build PVC error, got: %v", err) - } + c.Require().Error(err, "expected error from build failure") + c.StrContains( + err.Error(), + "failed to build shared backup PVC", + "expected build PVC error, got: %v", + err, + ) }) } func TestReconcile_Deletion(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -4532,22 +4314,17 @@ func TestReconcile_Deletion(t *testing.T) { req := ctrl.Request{NamespacedName: client.ObjectKeyFromObject(shard)} result, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if result.RequeueAfter != 0 { - t.Errorf("expected no requeue, got %v", result.RequeueAfter) - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(0, result.RequeueAfter, "expected no requeue, got") var updated multigresv1alpha1.Shard if err := c.Get(t.Context(), client.ObjectKeyFromObject(shard), &updated); err != nil { - if !errors.IsNotFound(err) { - t.Fatalf("unexpected error fetching shard: %v", err) - } + ck.Require().True(errors.IsNotFound(err), "unexpected error fetching shard: %v", err) } } func TestReconcileShardPDB_Error(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4575,15 +4352,12 @@ func TestReconcileShardPDB_Error(t *testing.T) { }) r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.reconcileShardPDB(t.Context(), shard) - if err == nil { - t.Fatal("expected error from PDB reconciliation") - } - if !strings.Contains(err.Error(), "failed to apply shard PDB") { - t.Errorf("expected PDB error, got: %v", err) - } + ck.Require().Error(err, "expected error from PDB reconciliation") + ck.StrContains(err.Error(), "failed to apply shard PDB", "expected PDB error, got: %v", err) } func TestUpdateStatus_ProgressingPhase(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4639,18 +4413,13 @@ func TestUpdateStatus_ProgressingPhase(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: recorder} err := r.updateStatus(t.Context(), shard, renderedConfig{}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if shard.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf("expected PhaseProgressing, got %q", shard.Status.Phase) - } - if shard.Status.Message == "" { - t.Error("expected non-empty status message for Progressing phase") - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(multigresv1alpha1.PhaseProgressing, shard.Status.Phase, "expected PhaseProgressing, got") + ck.NotEq("", shard.Status.Message, "expected non-empty status message for Progressing phase") } func TestUpdatePoolsStatus_PoolEmptyEvent(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4681,27 +4450,20 @@ func TestUpdatePoolsStatus_PoolEmptyEvent(t *testing.T) { cellsSet := make(map[multigresv1alpha1.CellName]bool) pools, err := r.updatePoolsStatus(t.Context(), shard, cellsSet, "", "") totalPods, readyPods := pools.totalPods, pools.readyPods - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if readyPods != 0 { - t.Errorf("expected 0 ready pods, got %d", readyPods) - } - if totalPods != 1 { - t.Errorf("expected 1 total pod (desired replicas), got %d", totalPods) - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(0, readyPods, "expected 0 ready pods, got") + ck.Eq(1, totalPods, "expected 1 total pod (desired replicas), got") select { case event := <-recorder.Events: - if !strings.Contains(event, "PoolEmpty") { - t.Errorf("expected PoolEmpty event, got: %s", event) - } + ck.StrContains(event, "PoolEmpty", "expected PoolEmpty event, got") default: t.Error("expected PoolEmpty event to be recorded") } } func TestReconcilePool_PoolPodsError(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4736,15 +4498,17 @@ func TestReconcilePool_PoolPodsError(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.reconcilePool(t.Context(), shard, "primary", poolSpec, &shardRolloutTracker{}) - if err == nil { - t.Fatal("expected error from reconcilePoolPods within reconcilePool") - } - if !strings.Contains(err.Error(), "failed to reconcile pool pods") { - t.Errorf("expected pool pods error, got: %v", err) - } + ck.Require().Error(err, "expected error from reconcilePoolPods within reconcilePool") + ck.StrContains( + err.Error(), + "failed to reconcile pool pods", + "expected pool pods error, got: %v", + err, + ) } func TestUpdateStatus_HealthyPhase(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4823,18 +4587,13 @@ func TestUpdateStatus_HealthyPhase(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: recorder} err := r.updateStatus(t.Context(), shard, renderedConfig{}) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if shard.Status.Phase != multigresv1alpha1.PhaseHealthy { - t.Errorf("expected PhaseHealthy, got %q", shard.Status.Phase) - } - if shard.Status.Message != "Ready" { - t.Errorf("expected 'Ready' message, got %q", shard.Status.Message) - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(multigresv1alpha1.PhaseHealthy, shard.Status.Phase, "expected PhaseHealthy, got") + ck.Eq("Ready", shard.Status.Message, "expected 'Ready' message, got") } func TestUpdatePoolsStatus_TerminatingPodExcluded(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -4885,17 +4644,11 @@ func TestUpdatePoolsStatus_TerminatingPodExcluded(t *testing.T) { cellsSet := make(map[multigresv1alpha1.CellName]bool) pools, err := r.updatePoolsStatus(t.Context(), shard, cellsSet, "", "") totalPods, readyPods := pools.totalPods, pools.readyPods - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // Pod is terminating, so it should be excluded. // totalPods = desired replicas (1), readyPods = 0 (terminating pod excluded) - if readyPods != 0 { - t.Errorf("expected 0 ready pods (terminating pod excluded), got %d", readyPods) - } - if totalPods != 1 { - t.Errorf("expected 1 total pod (desired replicas), got %d", totalPods) - } + ck.Eq(0, readyPods, "expected 0 ready pods (terminating pod excluded), got") + ck.Eq(1, totalPods, "expected 1 total pod (desired replicas), got") } func TestIsDrainStale(t *testing.T) { @@ -4929,9 +4682,7 @@ func TestIsDrainStale(t *testing.T) { // Build a pod with matching spec-hash for index 4 matchingPod := func(index int, drainState string) *corev1.Pod { desired, err := BuildPoolPod(shard, "main", "z1", shard.Spec.Pools["main"], index, scheme) - if err != nil { - t.Fatalf("BuildPoolPod failed: %v", err) - } + assert.NewAborting(t).NoError(err, "BuildPoolPod failed") hash := ComputeSpecHash(desired) return &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ @@ -4949,18 +4700,14 @@ func TestIsDrainStale(t *testing.T) { t.Run("CancelsStaleScaleDownDrain", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateRequested) - if !r.isDrainStale(shard, pod, metadata.DrainStateRequested) { - t.Error("expected drain to be stale (pod within replicas, spec matches)") - } + assert.NewCollecting(t). + True(r.isDrainStale(shard, pod, metadata.DrainStateRequested), "expected drain to be stale (pod within replicas, spec matches)") }) t.Run("DoesNotCancelDrainingState", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateDraining) - if r.isDrainStale(shard, pod, metadata.DrainStateDraining) { - t.Error( - "expected drain NOT to be stale in Draining state (standby removal already sent)", - ) - } + assert.NewCollecting(t). + False(r.isDrainStale(shard, pod, metadata.DrainStateDraining), "expected drain NOT to be stale in Draining state (standby removal already sent)") }) t.Run("DoesNotCancelExtraPodDrain", func(t *testing.T) { @@ -4971,55 +4718,50 @@ func TestIsDrainStale(t *testing.T) { ReplicasPerCell: ptr.To(int32(4)), } pod := matchingPod(4, metadata.DrainStateRequested) - if r.isDrainStale(smallShard, pod, metadata.DrainStateRequested) { - t.Error("expected drain NOT to be stale (pod is extra)") - } + assert.NewCollecting(t). + False(r.isDrainStale(smallShard, pod, metadata.DrainStateRequested), "expected drain NOT to be stale (pod is extra)") }) t.Run("MissingLabels", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateRequested) // Clear labels to hit `if poolName == "" || cellName == ""` pod.Labels = nil - if r.isDrainStale(shard, pod, metadata.DrainStateRequested) { - t.Error("expected drain NOT to be stale (missing labels)") - } + assert.NewCollecting(t). + False(r.isDrainStale(shard, pod, metadata.DrainStateRequested), "expected drain NOT to be stale (missing labels)") }) t.Run("MissingPoolInSpec", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateRequested) // Change pool to one that doesn't exist in spec pod.Labels[metadata.LabelMultigresPool] = "nonexistent" - if r.isDrainStale(shard, pod, metadata.DrainStateRequested) { - t.Error("expected drain NOT to be stale (pool not in spec)") - } + assert.NewCollecting(t). + False(r.isDrainStale(shard, pod, metadata.DrainStateRequested), "expected drain NOT to be stale (pool not in spec)") }) t.Run("DoesNotCancelAcknowledgedDrain", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateAcknowledged) - if r.isDrainStale(shard, pod, metadata.DrainStateAcknowledged) { - t.Error("expected drain NOT to be stale (past point of no return)") - } + assert.NewCollecting(t). + False(r.isDrainStale(shard, pod, metadata.DrainStateAcknowledged), "expected drain NOT to be stale (past point of no return)") }) t.Run("DoesNotCancelDrainOnDeletingPod", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateRequested) now := metav1.Now() pod.DeletionTimestamp = &now - if r.isDrainStale(shard, pod, metadata.DrainStateRequested) { - t.Error("expected drain NOT to be stale (pod is being deleted)") - } + assert.NewCollecting(t). + False(r.isDrainStale(shard, pod, metadata.DrainStateRequested), "expected drain NOT to be stale (pod is being deleted)") }) t.Run("DoesNotCancelWhenSpecDrifted", func(t *testing.T) { pod := matchingPod(4, metadata.DrainStateRequested) pod.Annotations[metadata.AnnotationSpecHash] = "wrong-hash" - if r.isDrainStale(shard, pod, metadata.DrainStateRequested) { - t.Error("expected drain NOT to be stale (spec-hash mismatch)") - } + assert.NewCollecting(t). + False(r.isDrainStale(shard, pod, metadata.DrainStateRequested), "expected drain NOT to be stale (spec-hash mismatch)") }) } func TestUpdatePoolsStatus_DrainAnnotationExcludedFromReady(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -5073,15 +4815,9 @@ func TestUpdatePoolsStatus_DrainAnnotationExcludedFromReady(t *testing.T) { cellsSet := make(map[multigresv1alpha1.CellName]bool) pools, err := r.updatePoolsStatus(t.Context(), shard, cellsSet, "", "") totalPods, readyPods := pools.totalPods, pools.readyPods - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if readyPods != 0 { - t.Errorf("expected 0 ready pods (draining pod excluded), got %d", readyPods) - } - if totalPods != 1 { - t.Errorf("expected 1 total pod (desired replicas), got %d", totalPods) - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(0, readyPods, "expected 0 ready pods (draining pod excluded), got") + ck.Eq(1, totalPods, "expected 1 total pod (desired replicas), got") } func TestUpdatePoolsStatus_DegradedOnCrashLoop(t *testing.T) { @@ -5121,6 +4857,7 @@ func TestUpdatePoolsStatus_DegradedOnCrashLoop(t *testing.T) { podName := BuildPoolPodName(shard, "primary", "zone1", 0) t.Run("CrashLoopBackOff", func(t *testing.T) { + ck := assert.NewCollecting(t) pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ Name: podName, @@ -5159,15 +4896,17 @@ func TestUpdatePoolsStatus_DegradedOnCrashLoop(t *testing.T) { Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.updateStatus(t.Context(), shard, renderedConfig{}); err != nil { - t.Fatalf("unexpected error: %v", err) - } - if shard.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("expected PhaseDegraded for CrashLoopBackOff pod, got %q", shard.Status.Phase) - } + ck.Require(). + NoError(r.updateStatus(t.Context(), shard, renderedConfig{}), "unexpected error") + ck.Eq( + multigresv1alpha1.PhaseDegraded, + shard.Status.Phase, + "expected PhaseDegraded for CrashLoopBackOff pod, got", + ) }) t.Run("OOMKilled", func(t *testing.T) { + ck := assert.NewCollecting(t) s := shard.DeepCopy() pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ @@ -5207,15 +4946,16 @@ func TestUpdatePoolsStatus_DegradedOnCrashLoop(t *testing.T) { Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.updateStatus(t.Context(), s, renderedConfig{}); err != nil { - t.Fatalf("unexpected error: %v", err) - } - if s.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("expected PhaseDegraded for OOMKilled pod, got %q", s.Status.Phase) - } + ck.Require().NoError(r.updateStatus(t.Context(), s, renderedConfig{}), "unexpected error") + ck.Eq( + multigresv1alpha1.PhaseDegraded, + s.Status.Phase, + "expected PhaseDegraded for OOMKilled pod, got", + ) }) t.Run("RunningPodNotDegraded", func(t *testing.T) { + ck := assert.NewCollecting(t) s := shard.DeepCopy() pod := &corev1.Pod{ ObjectMeta: metav1.ObjectMeta{ @@ -5258,15 +4998,12 @@ func TestUpdatePoolsStatus_DegradedOnCrashLoop(t *testing.T) { Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.updateStatus(t.Context(), s, renderedConfig{}); err != nil { - t.Fatalf("unexpected error: %v", err) - } - if s.Status.Phase == multigresv1alpha1.PhaseDegraded { - t.Errorf( - "expected Progressing (not Degraded) for running pod with prior restarts, got %q", - s.Status.Phase, - ) - } + ck.Require().NoError(r.updateStatus(t.Context(), s, renderedConfig{}), "unexpected error") + ck.NotEq( + multigresv1alpha1.PhaseDegraded, + s.Status.Phase, + "expected Progressing (not Degraded) for running pod with prior restarts, got", + ) }) } @@ -5334,6 +5071,7 @@ func TestUpdatePoolsStatus_ConfigApplyFailing(t *testing.T) { } t.Run("crash-looping pod on desired config is not settled", func(t *testing.T) { + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme). WithObjects(shard, crashLoopingPod(desiredHash)).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} @@ -5341,21 +5079,19 @@ func TestUpdatePoolsStatus_ConfigApplyFailing(t *testing.T) { pools, err := r.updatePoolsStatus( t.Context(), shard, make(map[multigresv1alpha1.CellName]bool), desiredHash, "", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !pools.poolDegraded { - t.Error("expected poolDegraded=true for a crash-looping pod") - } + ck.Require().NoError(err, "unexpected error") + ck.True(pools.poolDegraded, "expected poolDegraded=true for a crash-looping pod") // The pod already carries the desired hash, so there is no content drift; // configInProgress can only be true here via the apply-failing path (a // desired-config pod crash-looping), which is exactly what must not settle. - if !pools.configInProgress { - t.Error("expected configInProgress=true: desired config is on a crash-looping pod") - } + ck.True( + pools.configInProgress, + "expected configInProgress=true: desired config is on a crash-looping pod", + ) }) t.Run("crash-looping pod on stale config is unsettled via drift", func(t *testing.T) { + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme). WithObjects(shard, crashLoopingPod("stale-hash")).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} @@ -5363,20 +5099,18 @@ func TestUpdatePoolsStatus_ConfigApplyFailing(t *testing.T) { pools, err := r.updatePoolsStatus( t.Context(), shard, make(map[multigresv1alpha1.CellName]bool), desiredHash, "", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !pools.poolDegraded { - t.Error("expected poolDegraded=true for a crash-looping pod") - } + ck.Require().NoError(err, "unexpected error") + ck.True(pools.poolDegraded, "expected poolDegraded=true for a crash-looping pod") // The pod carries a stale hash, so config is unsettled via content drift. - if !pools.configInProgress { - t.Error("expected configInProgress=true: the pod carries a stale hash") - } + ck.True( + pools.configInProgress, + "expected configInProgress=true: the pod carries a stale hash", + ) }) } func TestUpdateStatus_DegradedOnMultiorchCrashLoop(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -5461,15 +5195,17 @@ func TestUpdateStatus_DegradedOnMultiorchCrashLoop(t *testing.T) { Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.updateStatus(t.Context(), shard, renderedConfig{}); err != nil { - t.Fatalf("unexpected error: %v", err) - } - if shard.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("expected PhaseDegraded for crash-looping Multiorch, got %q", shard.Status.Phase) - } - if shard.Status.Message != "One or more Multiorch pods are crash-looping" { - t.Errorf("expected Multiorch-specific degraded message, got %q", shard.Status.Message) - } + ck.Require().NoError(r.updateStatus(t.Context(), shard, renderedConfig{}), "unexpected error") + ck.Eq( + multigresv1alpha1.PhaseDegraded, + shard.Status.Phase, + "expected PhaseDegraded for crash-looping Multiorch, got", + ) + ck.Eq( + "One or more Multiorch pods are crash-looping", + shard.Status.Message, + "expected Multiorch-specific degraded message, got", + ) } func TestExpandPVCIfNeeded(t *testing.T) { @@ -5488,6 +5224,7 @@ func TestExpandPVCIfNeeded(t *testing.T) { } t.Run("no-op when sizes are equal", func(t *testing.T) { + ck := assert.NewCollecting(t) pvc := &corev1.PersistentVolumeClaim{ ObjectMeta: metav1.ObjectMeta{Name: "data-pvc-0", Namespace: "default"}, Spec: corev1.PersistentVolumeClaimSpec{ @@ -5504,19 +5241,17 @@ func TestExpandPVCIfNeeded(t *testing.T) { poolSpec := multigresv1alpha1.PoolSpec{ Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, } - if err := r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec), "unexpected error") got := &corev1.PersistentVolumeClaim{} _ = c.Get(t.Context(), types.NamespacedName{Name: "data-pvc-0", Namespace: "default"}, got) current := got.Spec.Resources.Requests[corev1.ResourceStorage] - if current.Cmp(resource.MustParse("10Gi")) != 0 { - t.Errorf("expected 10Gi, got %s", current.String()) - } + ck.Eq(0, current.Cmp(resource.MustParse("10Gi")), "expected 10Gi, got %s", current.String()) }) t.Run("patches PVC when desired is larger", func(t *testing.T) { + ck := assert.NewCollecting(t) pvc := &corev1.PersistentVolumeClaim{ ObjectMeta: metav1.ObjectMeta{Name: "data-pvc-1", Namespace: "default"}, Spec: corev1.PersistentVolumeClaimSpec{ @@ -5534,19 +5269,22 @@ func TestExpandPVCIfNeeded(t *testing.T) { poolSpec := multigresv1alpha1.PoolSpec{ Storage: multigresv1alpha1.StorageSpec{Size: "20Gi"}, } - if err := r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec), "unexpected error") got := &corev1.PersistentVolumeClaim{} _ = c.Get(t.Context(), types.NamespacedName{Name: "data-pvc-1", Namespace: "default"}, got) current := got.Spec.Resources.Requests[corev1.ResourceStorage] - if current.Cmp(resource.MustParse("20Gi")) != 0 { - t.Errorf("expected 20Gi after expansion, got %s", current.String()) - } + ck.Eq( + 0, + current.Cmp(resource.MustParse("20Gi")), + "expected 20Gi after expansion, got %s", + current.String(), + ) }) t.Run("no-op when desired is smaller", func(t *testing.T) { + ck := assert.NewCollecting(t) pvc := &corev1.PersistentVolumeClaim{ ObjectMeta: metav1.ObjectMeta{Name: "data-pvc-2", Namespace: "default"}, Spec: corev1.PersistentVolumeClaimSpec{ @@ -5563,19 +5301,22 @@ func TestExpandPVCIfNeeded(t *testing.T) { poolSpec := multigresv1alpha1.PoolSpec{ Storage: multigresv1alpha1.StorageSpec{Size: "10Gi"}, } - if err := r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec), "unexpected error") got := &corev1.PersistentVolumeClaim{} _ = c.Get(t.Context(), types.NamespacedName{Name: "data-pvc-2", Namespace: "default"}, got) current := got.Spec.Resources.Requests[corev1.ResourceStorage] - if current.Cmp(resource.MustParse("20Gi")) != 0 { - t.Errorf("expected 20Gi unchanged, got %s", current.String()) - } + ck.Eq( + 0, + current.Cmp(resource.MustParse("20Gi")), + "expected 20Gi unchanged, got %s", + current.String(), + ) }) t.Run("handles nil requests map", func(t *testing.T) { + ck := assert.NewCollecting(t) pvc := &corev1.PersistentVolumeClaim{ ObjectMeta: metav1.ObjectMeta{Name: "data-pvc-3", Namespace: "default"}, } @@ -5586,16 +5327,13 @@ func TestExpandPVCIfNeeded(t *testing.T) { poolSpec := multigresv1alpha1.PoolSpec{ Storage: multigresv1alpha1.StorageSpec{Size: "5Gi"}, } - if err := r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require(). + NoError(r.expandPVCIfNeeded(t.Context(), shard, pvc, poolSpec), "unexpected error") got := &corev1.PersistentVolumeClaim{} _ = c.Get(t.Context(), types.NamespacedName{Name: "data-pvc-3", Namespace: "default"}, got) current := got.Spec.Resources.Requests[corev1.ResourceStorage] - if current.Cmp(resource.MustParse("5Gi")) != 0 { - t.Errorf("expected 5Gi, got %s", current.String()) - } + ck.Eq(0, current.Cmp(resource.MustParse("5Gi")), "expected 5Gi, got %s", current.String()) }) } @@ -5615,25 +5353,22 @@ func TestPVCNeedsFilesystemResize(t *testing.T) { }, }, } - if !pvcNeedsFilesystemResize(pvcs, "data-pvc-0") { - t.Error("expected true for FileSystemResizePending condition") - } + assert.NewCollecting(t). + True(pvcNeedsFilesystemResize(pvcs, "data-pvc-0"), "expected true for FileSystemResizePending condition") }) t.Run("returns false when no condition", func(t *testing.T) { pvcs := map[string]*corev1.PersistentVolumeClaim{ "data-pvc-0": {}, } - if pvcNeedsFilesystemResize(pvcs, "data-pvc-0") { - t.Error("expected false when no conditions") - } + assert.NewCollecting(t). + False(pvcNeedsFilesystemResize(pvcs, "data-pvc-0"), "expected false when no conditions") }) t.Run("returns false for unknown PVC", func(t *testing.T) { pvcs := map[string]*corev1.PersistentVolumeClaim{} - if pvcNeedsFilesystemResize(pvcs, "missing") { - t.Error("expected false for unknown PVC name") - } + assert.NewCollecting(t). + False(pvcNeedsFilesystemResize(pvcs, "missing"), "expected false for unknown PVC name") }) } @@ -5675,9 +5410,8 @@ func TestResolvePodRole(t *testing.T) { PodRoles: tc.podRoles, }, } - if got := resolvePodRole(shard, tc.podName); got != tc.want { - t.Errorf("resolvePodRole() = %q, want %q", got, tc.want) - } + assert.NewCollecting(t). + Eq(tc.want, resolvePodRole(shard, tc.podName), "resolvePodRole()") }) } } @@ -5726,14 +5460,14 @@ func TestIsPoolerPruningEnabled(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() - if got := isPoolerPruningEnabled(tc.shard); got != tc.want { - t.Errorf("isPoolerPruningEnabled() = %v, want %v", got, tc.want) - } + assert.NewCollecting(t). + Eq(tc.want, isPoolerPruningEnabled(tc.shard), "isPoolerPruningEnabled()") }) } } func TestEnqueueFromPostgresConfigMap(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -5804,12 +5538,8 @@ func TestEnqueueFromPostgresConfigMap(t *testing.T) { requests := reconciler.enqueueFromPostgresConfigMap(context.Background(), cm) - if len(requests) != 1 { - t.Fatalf("expected 1 request, got %d", len(requests)) - } - if requests[0].Name != "shard-with-ref" { - t.Errorf("enqueued shard = %q, want %q", requests[0].Name, "shard-with-ref") - } + c.Require().Len(requests, 1, "expected 1 request, got %d", len(requests)) + c.Eq("shard-with-ref", requests[0].Name, "enqueued shard") } func TestRenderEffectiveConfig_RefHashing(t *testing.T) { @@ -5818,6 +5548,7 @@ func TestRenderEffectiveConfig_RefHashing(t *testing.T) { _ = corev1.AddToScheme(scheme) t.Run("produces a deterministic hash over the rendered config", func(t *testing.T) { + ck := assert.NewCollecting(t) cm := &corev1.ConfigMap{ ObjectMeta: metav1.ObjectMeta{Name: "pg-config", Namespace: "default"}, Data: map[string]string{"custom.conf": "shared_buffers = '8GB'"}, @@ -5837,25 +5568,18 @@ func TestRenderEffectiveConfig_RefHashing(t *testing.T) { rc := r.renderEffectiveConfig(context.Background(), shard) hash, err := rc.restartHash, rc.err - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(hash) != 64 { - t.Errorf("hash length = %d, want 64 (SHA-256 hex)", len(hash)) - } + ck.Require().NoError(err, "unexpected error") + ck.Len(hash, 64, "hash length = %d, want 64 (SHA-256 hex)", len(hash)) // Same content should produce the same hash. rc2 := r.renderEffectiveConfig(context.Background(), shard) hash2, err := rc2.restartHash, rc2.err - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if hash2 != hash { - t.Errorf("hash not deterministic: %q != %q", hash, hash2) - } + ck.Require().NoError(err, "unexpected error") + ck.Eq(hash, hash2, "hash not deterministic") }) t.Run("different content produces different hash", func(t *testing.T) { + ck := assert.NewCollecting(t) cm1 := &corev1.ConfigMap{ ObjectMeta: metav1.ObjectMeta{Name: "pg-v1", Namespace: "default"}, Data: map[string]string{"pg.conf": "shared_buffers = '4GB'"}, @@ -5889,20 +5613,15 @@ func TestRenderEffectiveConfig_RefHashing(t *testing.T) { rc1 := r.renderEffectiveConfig(context.Background(), shard1) h1, err := rc1.restartHash, rc1.err - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") rc2 := r.renderEffectiveConfig(context.Background(), shard2) h2, err := rc2.restartHash, rc2.err - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if h1 == h2 { - t.Error("different ConfigMap content should produce different hashes") - } + ck.Require().NoError(err, "unexpected error") + ck.NotEq(h2, h1, "different ConfigMap content should produce different hashes") }) t.Run("missing key returns error", func(t *testing.T) { + ck := assert.NewCollecting(t) cm := &corev1.ConfigMap{ ObjectMeta: metav1.ObjectMeta{Name: "pg-config", Namespace: "default"}, Data: map[string]string{"other.conf": "value"}, @@ -5921,12 +5640,8 @@ func TestRenderEffectiveConfig_RefHashing(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme} err := r.renderEffectiveConfig(context.Background(), shard).err - if err == nil { - t.Fatal("expected error for missing key") - } - if !strings.Contains(err.Error(), "missing-key") { - t.Errorf("error should mention missing key, got: %v", err) - } + ck.Require().Error(err, "expected error for missing key") + ck.StrContains(err.Error(), "missing-key", "error should mention missing key, got: %v", err) }) t.Run("missing ConfigMap returns error", func(t *testing.T) { @@ -5943,9 +5658,8 @@ func TestRenderEffectiveConfig_RefHashing(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).Build() r := &ShardReconciler{Client: c, Scheme: scheme} - if r.renderEffectiveConfig(context.Background(), shard).err == nil { - t.Error("expected error for missing ConfigMap") - } + assert.NewCollecting(t). + Error(r.renderEffectiveConfig(context.Background(), shard).err, "expected error for missing ConfigMap") }) } @@ -5990,9 +5704,8 @@ func TestReconcilePoolPods_AdditionalErrorPaths(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err == nil || !strings.Contains(err.Error(), "failed to create PVC") { - t.Fatalf("expected PVC creation error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to create PVC"), "expected PVC creation error, got %v", err) }) t.Run("markPodPVCOrphan network error", func(t *testing.T) { @@ -6034,9 +5747,8 @@ func TestReconcilePoolPods_AdditionalErrorPaths(t *testing.T) { r := &ShardReconciler{Client: fails, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} err := r.cleanupDrainedPod(context.Background(), shard, pod, "main", poolSpec, 1) - if err == nil || !strings.Contains(err.Error(), "failed to mark PVC") { - t.Fatalf("expected PVC orphan-patch error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to mark PVC"), "expected PVC orphan-patch error, got %v", err) }) } @@ -6097,6 +5809,7 @@ func TestUpdatePoolsStatus_ReloadPending(t *testing.T) { } t.Run("stale reload-hash is in progress even when restart-hash matches", func(t *testing.T) { + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme). WithObjects(shard, readyPod(desiredRestart, "reload-stale")).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} @@ -6108,15 +5821,15 @@ func TestUpdatePoolsStatus_ReloadPending(t *testing.T) { desiredRestart, desiredReload, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !pools.configInProgress { - t.Error("expected configInProgress=true: reload-hash is stale (reload pending)") - } + ck.Require().NoError(err, "unexpected error") + ck.True( + pools.configInProgress, + "expected configInProgress=true: reload-hash is stale (reload pending)", + ) }) t.Run("both hashes current is settled", func(t *testing.T) { + ck := assert.NewCollecting(t) c := fake.NewClientBuilder().WithScheme(scheme). WithObjects(shard, readyPod(desiredRestart, desiredReload)).Build() r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} @@ -6128,11 +5841,10 @@ func TestUpdatePoolsStatus_ReloadPending(t *testing.T) { desiredRestart, desiredReload, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if pools.configInProgress { - t.Error("expected configInProgress=false: both hashes match desired") - } + ck.Require().NoError(err, "unexpected error") + ck.False( + pools.configInProgress, + "expected configInProgress=false: both hashes match desired", + ) }) } diff --git a/pkg/resource-handler/controller/shard/shard_controller_test.go b/pkg/resource-handler/controller/shard/shard_controller_test.go index ea822b4a..e5a8e31f 100644 --- a/pkg/resource-handler/controller/shard/shard_controller_test.go +++ b/pkg/resource-handler/controller/shard/shard_controller_test.go @@ -21,6 +21,8 @@ import ( "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) type reconcileTestCase struct { @@ -71,48 +73,39 @@ func TestShardReconciler_Reconcile(t *testing.T) { }, existingObjects: []client.Object{}, assertFunc: func(t *testing.T, c client.Client, shard *multigresv1alpha1.Shard) { + ck := assert.NewCollecting(t) // Verify Multiorch Deployment was created (with cell suffix) moDeploy := &appsv1.Deployment{} hashedMoName := buildHashedMultiorchName(shard, "zone1") - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedMoName, Namespace: "default"}, - moDeploy); err != nil { - t.Errorf("Multiorch Deployment should exist: %v", err) - } + moDeploy), "Multiorch Deployment should exist") // Verify Multiorch Service was created (with cell suffix) moSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedMoName, Namespace: "default"}, - moSvc); err != nil { - t.Errorf("Multiorch Service should exist: %v", err) - } + moSvc), "Multiorch Service should exist") // Verify Pool Pod and PVC were created podName := BuildPoolPodName(shard, "primary", "zone1", 0) pod := &corev1.Pod{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: podName, Namespace: "default"}, - pod); err != nil { - t.Errorf("Pool Pod should exist: %v", err) - } + pod), "Pool Pod should exist") pvcName := BuildPoolDataPVCName(shard, "primary", "zone1", 0) pvc := &corev1.PersistentVolumeClaim{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: pvcName, Namespace: "default"}, - pvc); err != nil { - t.Errorf("Pool PVC should exist: %v", err) - } + pvc), "Pool PVC should exist") // Verify Pool headless Service was created (with cell suffix) hashedHeadless := buildHashedPoolHeadlessServiceName(shard, "primary", "zone1") poolSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedHeadless, Namespace: "default"}, - poolSvc); err != nil { - t.Errorf("Pool headless Service should exist: %v", err) - } + poolSvc), "Pool headless Service should exist") }, }, "create resources for Shard with multiple pools": { @@ -149,38 +142,33 @@ func TestShardReconciler_Reconcile(t *testing.T) { }, existingObjects: []client.Object{}, assertFunc: func(t *testing.T, c client.Client, shard *multigresv1alpha1.Shard) { + ck := assert.NewCollecting(t) // Verify replica pool pods for i := 0; i < 2; i++ { podName := BuildPoolPodName(shard, "replica", "zone1", i) - if err := c.Get( + ck.NoError(c.Get( t.Context(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("Replica pool Pod %d should exist: %v", i, err) - } + ), "Replica pool Pod %d should exist", i) } // Verify read-pool pods for i := 0; i < 3; i++ { podName := BuildPoolPodName(shard, "read-pool", "zone1", i) - if err := c.Get( + ck.NoError(c.Get( t.Context(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("read-pool Pod %d should exist: %v", i, err) - } + ), "read-pool Pod %d should exist", i) } // Verify both headless services hashReplicaHeadless := buildHashedPoolHeadlessServiceName(shard, "replica", "zone1") replicaSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashReplicaHeadless, Namespace: "default"}, - replicaSvc); err != nil { - t.Errorf("Replica pool headless Service should exist: %v", err) - } + replicaSvc), "Replica pool headless Service should exist") hashReadPoolHeadless := buildHashedPoolHeadlessServiceName( shard, @@ -188,11 +176,9 @@ func TestShardReconciler_Reconcile(t *testing.T) { "zone1", ) readPoolSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashReadPoolHeadless, Namespace: "default"}, - readPoolSvc); err != nil { - t.Errorf("read-pool headless Service should exist: %v", err) - } + readPoolSvc), "read-pool headless Service should exist") }, }, "Multiorch infers cells from pools when not specified": { @@ -221,22 +207,19 @@ func TestShardReconciler_Reconcile(t *testing.T) { }, existingObjects: []client.Object{}, assertFunc: func(t *testing.T, c client.Client, shard *multigresv1alpha1.Shard) { + ck := assert.NewCollecting(t) // Multiorch should be deployed to both zone1 and zone2 hashedMo1 := buildHashedMultiorchName(shard, "zone1") mo1 := &appsv1.Deployment{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedMo1, Namespace: "default"}, - mo1); err != nil { - t.Errorf("Multiorch Deployment for zone1 should exist: %v", err) - } + mo1), "Multiorch Deployment for zone1 should exist") hashedMo2 := buildHashedMultiorchName(shard, "zone2") mo2 := &appsv1.Deployment{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedMo2, Namespace: "default"}, - mo2); err != nil { - t.Errorf("Multiorch Deployment for zone2 should exist: %v", err) - } + mo2), "Multiorch Deployment for zone2 should exist") }, }, "create one shard-wide backup PVC for all active cells including pool-only cells": { @@ -278,27 +261,22 @@ func TestShardReconciler_Reconcile(t *testing.T) { }, existingObjects: []client.Object{}, assertFunc: func(t *testing.T, c client.Client, shard *multigresv1alpha1.Shard) { + ck := assert.NewCollecting(t) // Verify Multiorch Deployment was ONLY created for zone1 hashedMo1 := buildHashedMultiorchName(shard, "zone1") - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedMo1, Namespace: "default"}, - &appsv1.Deployment{}); err != nil { - t.Errorf("Multiorch Deployment for zone1 should exist: %v", err) - } + &appsv1.Deployment{}), "Multiorch Deployment for zone1 should exist") hashedMo2 := buildHashedMultiorchName(shard, "zone2") - if err := c.Get(t.Context(), + ck.Error(c.Get(t.Context(), types.NamespacedName{Name: hashedMo2, Namespace: "default"}, - &appsv1.Deployment{}); err == nil { - t.Errorf("Multiorch Deployment for zone2 should NOT exist") - } + &appsv1.Deployment{}), "Multiorch Deployment for zone2 should NOT exist") // All poolers share one backup PVC, even across cells. hashedPvc := buildHashedBackupPVCName(shard) - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: hashedPvc, Namespace: "default"}, - &corev1.PersistentVolumeClaim{}); err != nil { - t.Errorf("Shard-wide backup PVC should exist: %v", err) - } + &corev1.PersistentVolumeClaim{}), "Shard-wide backup PVC should exist") }, }, "error when Multiorch and pools have no cells specified": { @@ -381,46 +359,39 @@ func TestShardReconciler_Reconcile(t *testing.T) { }, existingObjects: []client.Object{}, assertFunc: func(t *testing.T, c client.Client, shard *multigresv1alpha1.Shard) { + ck := assert.NewCollecting(t) // Verify Pods for zone1 for i := 0; i < 2; i++ { podName := BuildPoolPodName(shard, "primary", "zone1", i) - if err := c.Get( + ck.NoError(c.Get( t.Context(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("Zone1 Pod %d should exist: %v", i, err) - } + ), "Zone1 Pod %d should exist", i) } // Verify Pods for zone2 for i := 0; i < 2; i++ { podName := BuildPoolPodName(shard, "primary", "zone2", i) - if err := c.Get( + ck.NoError(c.Get( t.Context(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("Zone2 Pod %d should exist: %v", i, err) - } + ), "Zone2 Pod %d should exist", i) } // Verify headless Services for both cells hashSvc1 := buildHashedPoolHeadlessServiceName(shard, "primary", "zone1") svc1 := &corev1.Service{} - if err := c.Get(t.Context(), + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{Name: hashSvc1, Namespace: "default"}, - svc1); err != nil { - t.Fatalf("Headless Service for zone1 should exist: %v", err) - } + svc1), "Headless Service for zone1 should exist") hashSvc2 := buildHashedPoolHeadlessServiceName(shard, "primary", "zone2") svc2 := &corev1.Service{} - if err := c.Get(t.Context(), + ck.Require().NoError(c.Get(t.Context(), types.NamespacedName{Name: hashSvc2, Namespace: "default"}, - svc2); err != nil { - t.Fatalf("Headless Service for zone2 should exist: %v", err) - } + svc2), "Headless Service for zone2 should exist") }, }, "update existing resources": { @@ -476,13 +447,11 @@ func TestShardReconciler_Reconcile(t *testing.T) { // Verify scale up to 5 pods for i := 0; i < 5; i++ { podName := BuildPoolPodName(shard, "primary", "zone1", i) - if err := c.Get( + assert.NewCollecting(t).NoError(c.Get( t.Context(), types.NamespacedName{Name: podName, Namespace: "default"}, &corev1.Pod{}, - ); err != nil { - t.Errorf("Zone1 Pod %d should exist: %v", i, err) - } + ), "Zone1 Pod %d should exist", i) } }, }, @@ -544,11 +513,9 @@ func TestShardReconciler_Reconcile(t *testing.T) { // Verify Multiorch Deployment was NOT created moDeploy := &appsv1.Deployment{} hashedMoName := buildHashedMultiorchName(shard, "zone1") - if err := c.Get(t.Context(), + assert.NewCollecting(t).Error(c.Get(t.Context(), types.NamespacedName{Name: hashedMoName, Namespace: "default"}, - moDeploy); err == nil { - t.Errorf("Multiorch Deployment should NOT exist") - } + moDeploy), "Multiorch Deployment should NOT exist") }, }, @@ -1105,9 +1072,7 @@ func TestShardReconciler_Reconcile(t *testing.T) { } if !shardInExisting { err := fakeClient.Create(t.Context(), tc.shard) - if err != nil { - t.Fatalf("Failed to create Shard: %v", err) - } + assert.NewAborting(t).NoError(err, "Failed to create Shard") } // Check headers @@ -1184,6 +1149,7 @@ func TestShardReconciler_Reconcile(t *testing.T) { } func TestShardReconciler_ReconcileNotFound(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -1210,12 +1176,8 @@ func TestShardReconciler_ReconcileNotFound(t *testing.T) { } result, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Errorf("Reconcile() should not error on NotFound, got: %v", err) - } - if result.RequeueAfter > 0 { - t.Errorf("Reconcile() should not requeue on NotFound") - } + c.NoError(err, "Reconcile() should not error on NotFound, got") + c.LessOrEqual(0, result.RequeueAfter, "Reconcile() should not requeue on NotFound") } func TestShardReconciler_UpdateStatus(t *testing.T) { @@ -1226,6 +1188,7 @@ func TestShardReconciler_UpdateStatus(t *testing.T) { _ = multigresv1alpha1.AddToScheme(scheme) t.Run("all_replicas_ready_status", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard-ready", @@ -1295,47 +1258,34 @@ func TestShardReconciler_UpdateStatus(t *testing.T) { Recorder: record.NewFakeRecorder(100), } - if err := r.updateStatus(context.Background(), shard, renderedConfig{}); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + c.Require(). + NoError(r.updateStatus(context.Background(), shard, renderedConfig{}), "updateStatus failed") updatedShard := &multigresv1alpha1.Shard{} - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(shard), updatedShard, - ); err != nil { - t.Fatalf("Failed to get Shard: %v", err) - } + ), "Failed to get Shard") foundTrue := false for _, cond := range updatedShard.Status.Conditions { if cond.Type == "Available" { - if cond.Status != metav1.ConditionFalse { - t.Errorf("Condition status = %s, want %s", cond.Status, metav1.ConditionFalse) - } - if cond.Reason != "NotAllPodsReady" { - t.Errorf("Condition reason = %s, want %s", cond.Reason, "NotAllPodsReady") - } + c.Eq(metav1.ConditionFalse, cond.Status, "Condition status") + c.Eq("NotAllPodsReady", cond.Reason, "Condition reason") foundTrue = true } } - if !foundTrue { - t.Errorf("Condition %s not found", "Available") - } - if updatedShard.Status.PoolsReady { - t.Error("PoolsReady should be false when 1/3 pools are ready") - } - if updatedShard.Status.Phase != multigresv1alpha1.PhaseProgressing { - t.Errorf( - "Expected Phase to be %s, got %s", - multigresv1alpha1.PhaseProgressing, - updatedShard.Status.Phase, - ) - } + c.True(foundTrue, "Condition %s not found", "Available") + c.False( + updatedShard.Status.PoolsReady, + "PoolsReady should be false when 1/3 pools are ready", + ) + c.Eq(multigresv1alpha1.PhaseProgressing, updatedShard.Status.Phase, "Expected Phase to be") }) t.Run("status_with_multiple_pools", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard-multi", @@ -1411,25 +1361,21 @@ func TestShardReconciler_UpdateStatus(t *testing.T) { Recorder: record.NewFakeRecorder(100), } - if err := r.updateStatus(context.Background(), shard, renderedConfig{}); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + c.Require(). + NoError(r.updateStatus(context.Background(), shard, renderedConfig{}), "updateStatus failed") updatedShard := &multigresv1alpha1.Shard{} - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( context.Background(), client.ObjectKeyFromObject(shard), updatedShard, - ); err != nil { - t.Fatalf("Failed to get Shard: %v", err) - } + ), "Failed to get Shard") - if !updatedShard.Status.PoolsReady { - t.Error("PoolsReady should be true when all pools are ready") - } + c.True(updatedShard.Status.PoolsReady, "PoolsReady should be true when all pools are ready") }) t.Run("multi_cell_pool_aggregates_across_cells", func(t *testing.T) { + c := assert.NewCollecting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ Name: "test-shard-multicell", @@ -1522,28 +1468,25 @@ func TestShardReconciler_UpdateStatus(t *testing.T) { pools, err := r.updatePoolsStatus( context.Background(), shard, cellsSet, "", "", ) - if err != nil { - t.Fatalf("updatePoolsStatus failed: %v", err) - } + c.Require().NoError(err, "updatePoolsStatus failed") totalPods, readyPods := pools.totalPods, pools.readyPods // Verify aggregate: desired for primary is 3 pods per cell * 2 cells = 6 pods - if totalPods != 6 { - t.Errorf("totalPods = %d, want 6", totalPods) - } - if readyPods != 5 { - t.Errorf("readyPods = %d, want 5", readyPods) - } + c.Eq(6, totalPods, "totalPods") + c.Eq(5, readyPods, "readyPods") // Verify both cells are tracked - if !cellsSet["zone1"] || !cellsSet["zone2"] { - t.Errorf("cellsSet = %v, want both zone1 and zone2", cellsSet) - } + c.False( + !cellsSet["zone1"] || !cellsSet["zone2"], + "cellsSet = %v, want both zone1 and zone2", + cellsSet, + ) }) } func TestScaleDownPodSelection(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) r := &ShardReconciler{} shard := &multigresv1alpha1.Shard{ @@ -1585,28 +1528,35 @@ func TestScaleDownPodSelection(t *testing.T) { // 1. All ready, 1 is primary. Should pick 2 (highest index and not primary) selected := r.selectPodToDrain(context.Background(), pods, shard) - if selected == nil || selected.Name != "pod-2" { - t.Errorf("Expected pod-2 to be selected, got %v", selected) - } + c.False( + selected == nil || selected.Name != "pod-2", + "Expected pod-2 to be selected, got %v", + selected, + ) // 2. Pod 0 is NOT ready. Should pick 0. pods[0].Status.Conditions[0].Status = corev1.ConditionFalse selected = r.selectPodToDrain(context.Background(), pods, shard) - if selected == nil || selected.Name != "pod-0" { - t.Errorf("Expected pod-0 (not ready) to be selected, got %v", selected) - } + c.False( + selected == nil || selected.Name != "pod-0", + "Expected pod-0 (not ready) to be selected, got %v", + selected, + ) // 3. Delete pod Roles, shouldn't panic, falls back to highest index shard.Status.PodRoles = nil pods[0].Status.Conditions[0].Status = corev1.ConditionTrue selected = r.selectPodToDrain(context.Background(), pods, shard) - if selected == nil || selected.Name != "pod-2" { - t.Errorf("Expected pod-2 (highest index) to be selected, got %v", selected) - } + c.False( + selected == nil || selected.Name != "pod-2", + "Expected pod-2 (highest index) to be selected, got %v", + selected, + ) } func TestScaleDown_ExternallyDeletedExtraPod(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -1679,13 +1629,9 @@ func TestScaleDown_ExternallyDeletedExtraPod(t *testing.T) { pod.Finalizers = []string{"kubernetes.io/test"} } - if err := c.Create(context.Background(), pod); err != nil { - t.Fatalf("failed to create pod: %v", err) - } + ck.Require().NoError(c.Create(context.Background(), pod), "failed to create pod") if i == 2 { - if err := c.Delete(t.Context(), pod); err != nil { - t.Fatal(err) - } + ck.Require().NoError(c.Delete(t.Context(), pod)) } } @@ -1698,9 +1644,7 @@ func TestScaleDown_ExternallyDeletedExtraPod(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // 2. Extra pod with DeletionTimestamp should have drain requested by handleExternalDeletion var extraPod corev1.Pod @@ -1712,16 +1656,13 @@ func TestScaleDown_ExternallyDeletedExtraPod(t *testing.T) { }, &extraPod, ) - if err != nil { - t.Fatalf("failed to get extra pod: %v", err) - } + ck.Require().NoError(err, "failed to get extra pod") - if extraPod.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf( - "Expected extra internally deleted pod to be marked for drain, got %v", - extraPod.Annotations[metadata.AnnotationDrainState], - ) - } + ck.Eq( + metadata.DrainStateRequested, + extraPod.Annotations[metadata.AnnotationDrainState], + "Expected extra internally deleted pod to be marked for drain, got", + ) } func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { @@ -1773,6 +1714,7 @@ func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { t.Run("blocks drain when pool has non-ready pod", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := shardObj.DeepCopy() shard.Status.PodRoles = map[string]string{ podName0: "PRIMARY", @@ -1807,36 +1749,27 @@ func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { 2, // effectiveReplicas: no maintenance surge in this test false, ) - if err != nil { - t.Fatalf("handleScaleDown returned error: %v", err) - } + ck.Require().NoError(err, "handleScaleDown returned error") - if actionTaken { - t.Error("Expected no action taken (health gate should block)") - } + ck.False(actionTaken, "Expected no action taken (health gate should block)") // Pod-2 should NOT have drain annotation updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), types.NamespacedName{Name: podName2, Namespace: "default"}, updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != "" { - t.Errorf( - "Expected no drain annotation (health gate should block), got %q", - updated.Annotations[metadata.AnnotationDrainState], - ) - } + ), "failed to get pod") + ck.Eq( + "", + updated.Annotations[metadata.AnnotationDrainState], + "Expected no drain annotation (health gate should block), got", + ) // Verify ScaleDownBlocked event was emitted select { case event := <-rec.Events: - if !strings.Contains(event, "ScaleDownBlocked") { - t.Errorf("Expected ScaleDownBlocked event, got %q", event) - } + ck.StrContains(event, "ScaleDownBlocked", "Expected ScaleDownBlocked event, got") default: t.Error("Expected ScaleDownBlocked event to be emitted") } @@ -1844,6 +1777,7 @@ func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { t.Run("allows drain when all pods are healthy", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := shardObj.DeepCopy() shard.Status.PodRoles = map[string]string{ podName0: "PRIMARY", @@ -1878,34 +1812,27 @@ func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { 2, // effectiveReplicas: no maintenance surge in this test false, ) - if err != nil { - t.Fatalf("handleScaleDown returned error: %v", err) - } + ck.Require().NoError(err, "handleScaleDown returned error") - if !actionTaken { - t.Error("Expected action taken (drain should proceed)") - } + ck.True(actionTaken, "Expected action taken (drain should proceed)") // Pod-2 SHOULD have drain annotation updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), types.NamespacedName{Name: podName2, Namespace: "default"}, updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf( - "Expected drain annotation %q, got %q", - metadata.DrainStateRequested, - updated.Annotations[metadata.AnnotationDrainState], - ) - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updated.Annotations[metadata.AnnotationDrainState], + "Expected drain annotation", + ) }) t.Run("shard tracker blocks a second pool scale-down", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := shardObj.DeepCopy() shard.Status.PodRoles = map[string]string{ podName0: "PRIMARY", @@ -1934,26 +1861,17 @@ func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { t.Context(), shard, poolName, multigresv1alpha1.PoolSpec{}, existingPods, 2, 2, false, tracker, ) - if err != nil { - t.Fatalf("handleScaleDown returned error: %v", err) - } - if actionTaken { - t.Fatal("expected the shard-wide tracker to block a second drain") - } + ck.Require().NoError(err, "handleScaleDown returned error") + ck.Require().False(actionTaken, "expected the shard-wide tracker to block a second drain") updated := &corev1.Pod{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(pods[2]), updated); err != nil { - t.Fatalf("get extra pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != "" { - t.Errorf( - "drain state = %q, want empty", - updated.Annotations[metadata.AnnotationDrainState], - ) - } + ck.Require(). + NoError(c.Get(t.Context(), client.ObjectKeyFromObject(pods[2]), updated), "get extra pod") + ck.Eq("", updated.Annotations[metadata.AnnotationDrainState], "drain state") }) t.Run("unhealthy extra pod does not block its own removal", func(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) shard := shardObj.DeepCopy() shard.Status.PodRoles = map[string]string{ podName0: "PRIMARY", @@ -1991,35 +1909,31 @@ func TestScaleDown_HealthGateBlocksDrain(t *testing.T) { 2, // effectiveReplicas: no maintenance surge in this test false, ) - if err != nil { - t.Fatalf("handleScaleDown returned error: %v", err) - } + ck.Require().NoError(err, "handleScaleDown returned error") - if !actionTaken { - t.Error("Expected action taken (unhealthy extra pod should not block its own removal)") - } + ck.True( + actionTaken, + "Expected action taken (unhealthy extra pod should not block its own removal)", + ) // Pod-2 SHOULD have drain annotation despite being unhealthy updated := &corev1.Pod{} - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), types.NamespacedName{Name: podName2, Namespace: "default"}, updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != metadata.DrainStateRequested { - t.Errorf( - "Expected drain annotation %q, got %q", - metadata.DrainStateRequested, - updated.Annotations[metadata.AnnotationDrainState], - ) - } + ), "failed to get pod") + ck.Eq( + metadata.DrainStateRequested, + updated.Annotations[metadata.AnnotationDrainState], + "Expected drain annotation", + ) }) } func TestRollingUpdateOrder(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -2083,12 +1997,9 @@ func TestRollingUpdateOrder(t *testing.T) { }, }, } - if err := c.Create(context.Background(), pod); err != nil { - t.Fatalf("failed to create pod: %v", err) - } - if err := c.Status().Update(context.Background(), pod); err != nil { - t.Fatalf("failed to set pod status: %v", err) - } + ck.Require().NoError(c.Create(context.Background(), pod), "failed to create pod") + ck.Require(). + NoError(c.Status().Update(context.Background(), pod), "failed to set pod status") } observeHealthyDisruption(t, r, shardObj, "primary", "zone1") @@ -2102,35 +2013,27 @@ func TestRollingUpdateOrder(t *testing.T) { poolSpec, &shardRolloutTracker{}, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.Require().NoError(err, "unexpected error") // Check that EXACTLY ONE REPLICA has drain-requested annotation. Pod 1 is PRIMARY. // So either Pod 0 or Pod 2 should be marked for drain, not Pod 1. drainCount := 0 for i := 0; i < 3; i++ { var pod corev1.Pod - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), types.NamespacedName{ Name: BuildPoolPodName(shardObj, "primary", "zone1", i), Namespace: "default", }, &pod, - ); err != nil { - t.Fatalf("pod %d should still exist: %v", i, err) - } + ), "pod %d should still exist", i) if pod.Annotations[metadata.AnnotationDrainState] == metadata.DrainStateRequested { drainCount++ - if i == 1 { - t.Errorf("Primary pod was marked for drain before replicas!") - } + ck.NotEq(1, i, "Primary pod was marked for drain before replicas!") } } - if drainCount != 1 { - t.Errorf("Expected exactly 1 pod to have drain-requested annotation, got %d", drainCount) - } + ck.Eq(1, drainCount, "Expected exactly 1 pod to have drain-requested annotation, got") } // TestRollingUpdateWaitsForSiblingCell verifies that a pool does not start @@ -2141,6 +2044,7 @@ func TestRollingUpdateOrder(t *testing.T) { // once. func TestRollingUpdateWaitsForSiblingCell(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = corev1.AddToScheme(scheme) @@ -2219,36 +2123,30 @@ func TestRollingUpdateWaitsForSiblingCell(t *testing.T) { r := &ShardReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} for _, pod := range []*corev1.Pod{zone1Pod, zone2Pod} { - if err := c.Create(context.Background(), pod); err != nil { - t.Fatalf("failed to create pod %s: %v", pod.Name, err) - } - if err := c.Status().Update(context.Background(), pod); err != nil { - t.Fatalf("failed to set status for pod %s: %v", pod.Name, err) - } + ck.Require(). + NoError(c.Create(context.Background(), pod), "failed to create pod %s", pod.Name) + ck.Require(). + NoError(c.Status().Update(context.Background(), pod), "failed to set status for pod %s", pod.Name) } - if err := r.reconcilePoolPods( + ck.Require().NoError(r.reconcilePoolPods( context.Background(), shardObj, "pool-2", "zone2", poolSpec, &shardRolloutTracker{}, - ); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ), "unexpected error") var updated corev1.Pod - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKeyFromObject(zone2Pod), &updated, - ); err != nil { - t.Fatalf("failed to get pod: %v", err) - } - if updated.Annotations[metadata.AnnotationDrainState] != "" { - t.Error( - "pool-2 should not drain its drifted pod while pool-1's pod in zone1 is not Ready", - ) - } + ), "failed to get pod") + ck.Eq( + "", + updated.Annotations[metadata.AnnotationDrainState], + "pool-2 should not drain its drifted pod while pool-1's pod in zone1 is not Ready", + ) } diff --git a/pkg/resource-handler/controller/shard/shard_pdb_test.go b/pkg/resource-handler/controller/shard/shard_pdb_test.go index 76cd79ea..2934df6b 100644 --- a/pkg/resource-handler/controller/shard/shard_pdb_test.go +++ b/pkg/resource-handler/controller/shard/shard_pdb_test.go @@ -17,6 +17,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestShardMinAvailable(t *testing.T) { @@ -36,19 +38,16 @@ func TestShardMinAvailable(t *testing.T) { shard := &multigresv1alpha1.Shard{Spec: multigresv1alpha1.ShardSpec{ Replicas: ptr.To(tc.replicas), }} - if got := shardMinAvailable(shard); got != tc.want { - t.Fatalf("shardMinAvailable() = %d, want %d", got, tc.want) - } + assert.NewAborting(t).Eq(tc.want, shardMinAvailable(shard), "shardMinAvailable()") }) } } func TestBuildShardPodDisruptionBudgetsDoNotOverlap(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("add Shard scheme: %v", err) - } + c.Require().NoError(multigresv1alpha1.AddToScheme(scheme), "add Shard scheme") shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -75,35 +74,22 @@ func TestBuildShardPodDisruptionBudgetsDoNotOverlap(t *testing.T) { } pdbs, err := BuildShardPodDisruptionBudgets(shard, scheme) - if err != nil { - t.Fatalf("build PDBs: %v", err) - } - if len(pdbs) != 1 { - t.Fatalf("PDB count = %d, want one shard-wide budget", len(pdbs)) - } - if got := pdbs[0].Spec.MinAvailable.IntValue(); got != 3 { - t.Errorf("shard minAvailable = %d, want 3", got) - } + c.Require().NoError(err, "build PDBs") + c.Require().Len(pdbs, 1, "PDB count = %d, want one shard-wide budget", len(pdbs)) + c.Eq(3, pdbs[0].Spec.MinAvailable.IntValue(), "shard minAvailable") selector := pdbs[0].Spec.Selector.MatchLabels - if _, scopedToCell := selector[metadata.LabelMultigresCell]; scopedToCell { - t.Errorf("shard PDB must not select a cell: %#v", selector) - } - if _, scopedToPool := selector[metadata.LabelMultigresPool]; scopedToPool { - t.Errorf("shard PDB must not select a pool: %#v", selector) - } + _, scopedToCell := selector[metadata.LabelMultigresCell] + c.False(scopedToCell, "shard PDB must not select a cell: %#v", selector) + _, scopedToPool := selector[metadata.LabelMultigresPool] + c.False(scopedToPool, "shard PDB must not select a pool: %#v", selector) } func TestReconcileShardPDBReplacesLegacyPoolCellPDBs(t *testing.T) { + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("add Shard scheme: %v", err) - } - if err := policyv1.AddToScheme(scheme); err != nil { - t.Fatalf("add policy scheme: %v", err) - } - if err := corev1.AddToScheme(scheme); err != nil { - t.Fatalf("add Pod scheme: %v", err) - } + ck.Require().NoError(multigresv1alpha1.AddToScheme(scheme), "add Shard scheme") + ck.Require().NoError(policyv1.AddToScheme(scheme), "add policy scheme") + ck.Require().NoError(corev1.AddToScheme(scheme), "add Pod scheme") shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -122,9 +108,7 @@ func TestReconcileShardPDBReplacesLegacyPoolCellPDBs(t *testing.T) { } desired, err := BuildShardPodDisruptionBudget(shard, scheme) - if err != nil { - t.Fatalf("build desired PDB: %v", err) - } + ck.Require().NoError(err, "build desired PDB") legacy := &policyv1.PodDisruptionBudget{ ObjectMeta: metav1.ObjectMeta{ @@ -135,9 +119,8 @@ func TestReconcileShardPDBReplacesLegacyPoolCellPDBs(t *testing.T) { } legacy.Labels[metadata.LabelMultigresPool] = "primary" legacy.Labels[metadata.LabelMultigresCell] = "zone1" - if err := ctrl.SetControllerReference(shard, legacy, scheme); err != nil { - t.Fatalf("set legacy owner reference: %v", err) - } + ck.Require(). + NoError(ctrl.SetControllerReference(shard, legacy, scheme), "set legacy owner reference") unmanaged := legacy.DeepCopy() unmanaged.Name = "unmanaged-pdb" @@ -149,17 +132,13 @@ func TestReconcileShardPDBReplacesLegacyPoolCellPDBs(t *testing.T) { Build() r := &ShardReconciler{Client: c, Scheme: scheme} - if err := r.reconcileShardPDB(t.Context(), shard); err != nil { - t.Fatalf("reconcile shard PDB: %v", err) - } + ck.Require().NoError(r.reconcileShardPDB(t.Context(), shard), "reconcile shard PDB") - if err := c.Get( + ck.NoError(c.Get( t.Context(), client.ObjectKeyFromObject(desired), &policyv1.PodDisruptionBudget{}, - ); err != nil { - t.Errorf("shard-wide PDB should exist: %v", err) - } + ), "shard-wide PDB should exist") if err := c.Get( t.Context(), client.ObjectKeyFromObject(legacy), @@ -167,17 +146,16 @@ func TestReconcileShardPDBReplacesLegacyPoolCellPDBs(t *testing.T) { ); !apierrors.IsNotFound(err) { t.Errorf("legacy PDB should be deleted, got: %v", err) } - if err := c.Get( + ck.NoError(c.Get( t.Context(), client.ObjectKeyFromObject(unmanaged), &policyv1.PodDisruptionBudget{}, - ); err != nil { - t.Errorf("unmanaged PDB should be preserved: %v", err) - } + ), "unmanaged PDB should be preserved") } func TestReconcileShardPDBCountsMaintenanceSurge(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := maintenanceSurgeTestScheme(t) shard := maintenanceSurgeTestShard() shard.Spec.Replicas = ptr.To(int32(3)) @@ -192,18 +170,10 @@ func TestReconcileShardPDBCountsMaintenanceSurge(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(shard, surge).Build() r := &ShardReconciler{Client: c, Scheme: scheme} - if err := r.reconcileShardPDB(t.Context(), shard); err != nil { - t.Fatalf("reconcile shard PDB: %v", err) - } + ck.NoError(r.reconcileShardPDB(t.Context(), shard), "reconcile shard PDB") desired, err := BuildShardPodDisruptionBudget(shard, scheme) - if err != nil { - t.Fatalf("build shard PDB: %v", err) - } + ck.NoError(err, "build shard PDB") actual := &policyv1.PodDisruptionBudget{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(desired), actual); err != nil { - t.Fatalf("get shard PDB: %v", err) - } - if got := actual.Spec.MinAvailable.IntValue(); got != 3 { - t.Fatalf("minAvailable with one surge = %d, want 3", got) - } + ck.NoError(c.Get(t.Context(), client.ObjectKeyFromObject(desired), actual), "get shard PDB") + ck.Eq(3, actual.Spec.MinAvailable.IntValue(), "minAvailable with one surge") } diff --git a/pkg/resource-handler/controller/shard/storage_class_guard_test.go b/pkg/resource-handler/controller/shard/storage_class_guard_test.go index 70137e36..3c5389ee 100644 --- a/pkg/resource-handler/controller/shard/storage_class_guard_test.go +++ b/pkg/resource-handler/controller/shard/storage_class_guard_test.go @@ -6,7 +6,6 @@ import ( "errors" "maps" "slices" - "strings" "testing" "time" @@ -28,6 +27,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestValidateStorageClassDependencies(t *testing.T) { @@ -54,6 +55,7 @@ func TestValidateStorageClassDependencies(t *testing.T) { } t.Run("nothing explicit reports one not-specified verdict", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -65,18 +67,21 @@ func TestValidateStorageClassDependencies(t *testing.T) { r := newReconciler(shard) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if check.status != metav1.ConditionTrue || check.reason != storageClassNotSpecifiedReason { - t.Fatalf("unexpected verdict: %+v", check) - } - if check.backupDependency != nil || check.poolDependency != nil { - t.Fatalf("expected no dependency errors, got %+v", check) - } + c.NoError(err, "unexpected error") + c.False( + check.status != metav1.ConditionTrue || check.reason != storageClassNotSpecifiedReason, + "unexpected verdict: %+v", + check, + ) + c.False( + check.backupDependency != nil || check.poolDependency != nil, + "expected no dependency errors, got %+v", + check, + ) }) t.Run("explicit classes all present report found", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -95,12 +100,12 @@ func TestValidateStorageClassDependencies(t *testing.T) { ) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if check.status != metav1.ConditionTrue || check.reason != storageClassFoundReason { - t.Fatalf("unexpected verdict: %+v", check) - } + c.NoError(err, "unexpected error") + c.False( + check.status != metav1.ConditionTrue || check.reason != storageClassFoundReason, + "unexpected verdict: %+v", + check, + ) }) // One writer means one verdict, and in the mixed cases the merged verdict @@ -109,6 +114,7 @@ func TestValidateStorageClassDependencies(t *testing.T) { // was explicit. The reason is published on the condition, so both directions // of "some of it is explicit" are pinned here. t.Run("explicit backup class with no pool class reports found", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -124,15 +130,16 @@ func TestValidateStorageClassDependencies(t *testing.T) { ) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if check.status != metav1.ConditionTrue || check.reason != storageClassFoundReason { - t.Fatalf("unexpected verdict: %+v", check) - } + c.NoError(err, "unexpected error") + c.False( + check.status != metav1.ConditionTrue || check.reason != storageClassFoundReason, + "unexpected verdict: %+v", + check, + ) }) t.Run("explicit pool class with no backup class reports found", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -149,15 +156,16 @@ func TestValidateStorageClassDependencies(t *testing.T) { ) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if check.status != metav1.ConditionTrue || check.reason != storageClassFoundReason { - t.Fatalf("unexpected verdict: %+v", check) - } + c.NoError(err, "unexpected error") + c.False( + check.status != metav1.ConditionTrue || check.reason != storageClassFoundReason, + "unexpected verdict: %+v", + check, + ) }) t.Run("missing backup class reports the backup dependency", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{Backup: filesystemBackup("missing-sc")}, @@ -165,24 +173,23 @@ func TestValidateStorageClassDependencies(t *testing.T) { r := newReconciler(shard) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if check.status != metav1.ConditionFalse || check.reason != storageClassNotFoundReason { - t.Fatalf("unexpected verdict: %+v", check) - } - if !isMissingStorageClassDependency(check.backupDependency) { - t.Fatalf("expected backup dependency error, got %v", check.backupDependency) - } - if check.poolDependency != nil { - t.Fatalf("expected no pool dependency error, got %v", check.poolDependency) - } - if !strings.Contains(check.message, `"missing-sc"`) { - t.Fatalf("message must name the class, got %q", check.message) - } + c.NoError(err, "unexpected error") + c.False( + check.status != metav1.ConditionFalse || check.reason != storageClassNotFoundReason, + "unexpected verdict: %+v", + check, + ) + c.True( + isMissingStorageClassDependency(check.backupDependency), + "expected backup dependency error, got %v", + check.backupDependency, + ) + c.NoError(check.poolDependency, "expected no pool dependency error, got") + c.StrContains(check.message, `"missing-sc"`, "message must name the class, got") }) t.Run("missing pool class reports the pool dependency", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -196,24 +203,23 @@ func TestValidateStorageClassDependencies(t *testing.T) { r := newReconciler(shard) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if check.status != metav1.ConditionFalse || check.reason != storageClassNotFoundReason { - t.Fatalf("unexpected verdict: %+v", check) - } - if !isMissingStorageClassDependency(check.poolDependency) { - t.Fatalf("expected pool dependency error, got %v", check.poolDependency) - } - if check.backupDependency != nil { - t.Fatalf("expected no backup dependency error, got %v", check.backupDependency) - } - if !strings.Contains(check.message, "primary") { - t.Fatalf("message must name the pool, got %q", check.message) - } + c.NoError(err, "unexpected error") + c.False( + check.status != metav1.ConditionFalse || check.reason != storageClassNotFoundReason, + "unexpected verdict: %+v", + check, + ) + c.True( + isMissingStorageClassDependency(check.poolDependency), + "expected pool dependency error, got %v", + check.poolDependency, + ) + c.NoError(check.backupDependency, "expected no backup dependency error, got") + c.StrContains(check.message, "primary", "message must name the pool, got") }) t.Run("missing backup class wins over a missing pool class", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -231,18 +237,15 @@ func TestValidateStorageClassDependencies(t *testing.T) { r := newReconciler(shard) check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if !isMissingStorageClassDependency(check.backupDependency) || - check.poolDependency != nil { - t.Fatalf("expected only the backup dependency, got %+v", check) - } + c.NoError(err, "unexpected error") + c.False(!isMissingStorageClassDependency(check.backupDependency) || + check.poolDependency != nil, "expected only the backup dependency, got %+v", check) }) // Spec.Pools is a map, so an unordered scan would report whichever missing // pool Go's iteration reached first and flap the condition message. t.Run("the reported pool is stable across calls", func(t *testing.T) { + c := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "test-shard", Namespace: "default"}, Spec: multigresv1alpha1.ShardSpec{ @@ -264,19 +267,13 @@ func TestValidateStorageClassDependencies(t *testing.T) { want := "" for i := range 30 { check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.NoError(err, "unexpected error") if i == 0 { want = check.message } - if check.message != want { - t.Fatalf("message flapped: %q then %q", want, check.message) - } - } - if !strings.Contains(want, "aaa") { - t.Fatalf("expected the first pool by name, got %q", want) + c.Eq(want, check.message, "message flapped") } + c.StrContains(want, "aaa", "expected the first pool by name, got") }) } @@ -298,6 +295,7 @@ func TestValidateStorageClassDependencies(t *testing.T) { // test/suite. func TestSetStorageClassCondition_AppliesOnlyConditions(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -321,9 +319,7 @@ func TestSetStorageClassCondition_AppliesOnlyConditions(t *testing.T) { fakeClient := testutil.NewFakeClientWithFailures(baseClient, &testutil.FailureConfig{ OnStatusPatch: func(obj client.Object) error { raw, err := json.Marshal(obj) - if err != nil { - t.Fatalf("marshal patch payload: %v", err) - } + c.NoError(err, "marshal patch payload") captured = raw return nil }, @@ -336,32 +332,24 @@ func TestSetStorageClassCondition_AppliesOnlyConditions(t *testing.T) { reason: storageClassNotSpecifiedReason, message: "No explicit backup filesystem or pool StorageClass configured; using cluster default", } - if err := r.setStorageClassCondition(t.Context(), shard, check); err != nil { - t.Fatalf("setStorageClassCondition: %v", err) - } - if captured == nil { - t.Fatal("guard did not apply a status patch") - } + c.NoError(r.setStorageClassCondition(t.Context(), shard, check), "setStorageClassCondition") + c.NotNil(captured, "guard did not apply a status patch") var payload struct { Status map[string]json.RawMessage `json:"status"` } - if err := json.Unmarshal(captured, &payload); err != nil { - t.Fatalf("unmarshal patch payload: %v", err) - } + c.NoError(json.Unmarshal(captured, &payload), "unmarshal patch payload") keys := slices.Sorted(maps.Keys(payload.Status)) - if !slices.Equal(keys, []string{"conditions"}) { - t.Fatalf("guard payload must own status.conditions only, got %v", keys) - } + c.EqDiff([]string{"conditions"}, keys, "guard payload must own status.conditions only, got") var conditions []metav1.Condition - if err := json.Unmarshal(payload.Status["conditions"], &conditions); err != nil { - t.Fatalf("unmarshal conditions: %v", err) - } - if len(conditions) != 1 || conditions[0].Type != conditionStorageClassValid { - t.Fatalf("guard payload must carry exactly the %s condition, got %+v", - conditionStorageClassValid, conditions) - } + c.NoError(json.Unmarshal(payload.Status["conditions"], &conditions), "unmarshal conditions") + c.False( + len(conditions) != 1 || conditions[0].Type != conditionStorageClassValid, + "guard payload must carry exactly the %s condition, got %+v", + conditionStorageClassValid, + conditions, + ) if conditions[0].Status != check.status || conditions[0].Reason != check.reason || conditions[0].Message != check.message { t.Fatalf("condition does not match the verdict: %+v", conditions[0]) @@ -380,6 +368,7 @@ func TestSetStorageClassCondition_AppliesOnlyConditions(t *testing.T) { // of this assertion is TestShardStatusQuiesces under test/suite. func TestStorageClassCondition_IsStableAcrossReconciles(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -414,44 +403,33 @@ func TestStorageClassCondition_IsStableAcrossReconciles(t *testing.T) { reconcileStorageClasses := func() []byte { check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("validate: %v", err) - } - if err := r.setStorageClassCondition(t.Context(), shard, check); err != nil { - t.Fatalf("set condition: %v", err) - } + c.NoError(err, "validate") + c.NoError(r.setStorageClassCondition(t.Context(), shard, check), "set condition") var got multigresv1alpha1.Shard - if err := baseClient.Get(t.Context(), client.ObjectKeyFromObject(shard), &got); err != nil { - t.Fatalf("read shard: %v", err) - } + c.NoError( + baseClient.Get(t.Context(), client.ObjectKeyFromObject(shard), &got), + "read shard", + ) cond := findCondition(got.Status.Conditions, conditionStorageClassValid) - if cond == nil { - t.Fatalf("no %s condition", conditionStorageClassValid) - } + c.NotNil(cond, "no %s condition", conditionStorageClassValid) raw, err := json.Marshal(cond) - if err != nil { - t.Fatalf("marshal condition: %v", err) - } + c.NoError(err, "marshal condition") return raw } first := reconcileStorageClasses() - if patches != 1 { - t.Fatalf("first cycle must apply the condition once, applied %d times", patches) - } + c.Eq(1, patches, "first cycle must apply the condition once, applied") second := reconcileStorageClasses() - if string(first) != string(second) { - t.Fatalf( - "StorageClassValid condition moved between reconciles:\n %s\n %s", - first, - second, - ) - } - if patches != 1 { - t.Fatalf("second cycle rewrote a settled condition: %d applies total", patches) - } + c.Eq( + string(second), + string(first), + "StorageClassValid condition moved between reconciles:\n %s\n %s", + first, + second, + ) + c.Eq(1, patches, "second cycle rewrote a settled condition") } // guardStatusApply is one status apply captured under the guard's field @@ -494,6 +472,7 @@ type guardStatusApply struct { // TestShardStatusQuiesces under test/suite. func TestReconcile_StorageClassConditionSettlesAcrossReconciles(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -541,27 +520,16 @@ func TestReconcile_StorageClassConditionSettlesAcrossReconciles(t *testing.T) { popts := (&client.SubResourcePatchOptions{}).ApplyOptions(opts) if popts.FieldManager == "multigres-resource-handler-guard" { raw, err := json.Marshal(obj) - if err != nil { - t.Fatalf("marshal guard payload: %v", err) - } + ck.NoError(err, "marshal guard payload") var payload struct { Status map[string]json.RawMessage `json:"status"` } - if err := json.Unmarshal(raw, &payload); err != nil { - t.Fatalf("unmarshal guard payload: %v", err) - } + ck.NoError(json.Unmarshal(raw, &payload), "unmarshal guard payload") var conditions []metav1.Condition if raw, ok := payload.Status["conditions"]; ok { - if err := json.Unmarshal(raw, &conditions); err != nil { - t.Fatalf("unmarshal guard conditions: %v", err) - } - } - if len(conditions) != 1 { - t.Fatalf( - "guard payload must carry exactly one condition, got %+v", - conditions, - ) + ck.NoError(json.Unmarshal(raw, &conditions), "unmarshal guard conditions") } + ck.Len(conditions, 1, "guard payload must carry exactly one condition, got") guardApplies = append(guardApplies, guardStatusApply{ statusKeys: slices.Sorted(maps.Keys(payload.Status)), condition: conditions[0], @@ -583,30 +551,34 @@ func TestReconcile_StorageClassConditionSettlesAcrossReconciles(t *testing.T) { assertOnlyOwnsConditions := func(t *testing.T, apply guardStatusApply) { t.Helper() - if !slices.Equal(apply.statusKeys, []string{"conditions"}) { - t.Fatalf("guard payload must own status.conditions only, got %v", apply.statusKeys) - } - if apply.condition.Type != conditionStorageClassValid { - t.Fatalf("guard payload must carry the %s condition, got %+v", - conditionStorageClassValid, apply.condition) - } + c := assert.NewAborting(t) + c.EqDiff( + []string{"conditions"}, + apply.statusKeys, + "guard payload must own status.conditions only, got", + ) + c.Eq( + conditionStorageClassValid, + apply.condition.Type, + "guard payload must carry the %s condition, got %+v", + conditionStorageClassValid, + apply.condition, + ) } if _, err := r.Reconcile(t.Context(), req); err != nil { t.Fatalf("first reconcile: %v", err) } - if len(guardApplies) != 1 { - t.Fatalf( - "first reconcile must make exactly one guard apply (a second writer is back), got %d", - len(guardApplies), - ) - } + ck.Len( + guardApplies, + 1, + "first reconcile must make exactly one guard apply (a second writer is back), got %d", + len(guardApplies), + ) first := guardApplies[0] assertOnlyOwnsConditions(t, first) - if first.condition.Status != metav1.ConditionTrue || - first.condition.Reason != storageClassNotSpecifiedReason { - t.Fatalf("unexpected first verdict: %+v", first.condition) - } + ck.False(first.condition.Status != metav1.ConditionTrue || + first.condition.Reason != storageClassNotSpecifiedReason, "unexpected first verdict: %+v", first.condition) // The fake client's SSA status apply replaces the whole object rather than // scoping to the applied fields, which drops Spec. Restore it exactly as the @@ -614,39 +586,32 @@ func TestReconcile_StorageClassConditionSettlesAcrossReconciles(t *testing.T) { // second reconcile sees the same spec as the first rather than erroring on a // shard with no pools. var stored multigresv1alpha1.Shard - if err := fakeClient.Get(t.Context(), req.NamespacedName, &stored); err != nil { - t.Fatalf("read shard before restoring spec: %v", err) - } + ck.NoError( + fakeClient.Get(t.Context(), req.NamespacedName, &stored), + "read shard before restoring spec", + ) stored.Spec = shard.Spec - if err := fakeClient.Update(t.Context(), &stored); err != nil { - t.Fatalf("restore shard spec: %v", err) - } - - if _, err := r.Reconcile(t.Context(), req); err != nil { - t.Fatalf("second reconcile: %v", err) - } - if len(guardApplies) != 2 { - t.Fatalf( - "second reconcile must make exactly one guard apply of its own (a second writer "+ - "is back), got %d total", - len(guardApplies), - ) - } + ck.NoError(fakeClient.Update(t.Context(), &stored), "restore shard spec") + + _, err := r.Reconcile(t.Context(), req) + ck.NoError(err, "second reconcile") + ck.Len( + guardApplies, + 2, + "second reconcile must make exactly one guard apply of its own (a second writer "+ + "is back), got %d total", + len(guardApplies), + ) second := guardApplies[1] assertOnlyOwnsConditions(t, second) - if second.condition.Status != first.condition.Status || + ck.False(second.condition.Status != first.condition.Status || second.condition.Reason != first.condition.Reason || - second.condition.Message != first.condition.Message { - t.Fatalf( - "StorageClassValid verdict flapped between reconciles:\n %+v\n %+v", - first.condition, - second.condition, - ) - } + second.condition.Message != first.condition.Message, "StorageClassValid verdict flapped between reconciles:\n %+v\n %+v", first.condition, second.condition) } func TestReconcile_MissingStorageClassReturnsDependencyRequeueEvenWhenPVCExists(t *testing.T) { + ck := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -712,24 +677,20 @@ func TestReconcile_MissingStorageClassReturnsDependencyRequeueEvenWhenPVCExists( result, err := r.Reconcile(t.Context(), ctrl.Request{ NamespacedName: client.ObjectKeyFromObject(shard), }) - if err != nil { - t.Fatalf("expected non-error dependency requeue, got error: %v", err) - } - if result.RequeueAfter != storageClassDependencyRequeue { - t.Fatalf("requeueAfter = %v, want %v", result.RequeueAfter, storageClassDependencyRequeue) - } + ck.NoError(err, "expected non-error dependency requeue, got error") + ck.Eq(storageClassDependencyRequeue, result.RequeueAfter, "requeueAfter") } func TestIsMissingStorageClassDependencyWrapped(t *testing.T) { + c := assert.NewAborting(t) err := errors.New("other") - if isMissingStorageClassDependency(err) { - t.Fatal("expected false for non-dependency error") - } + c.False(isMissingStorageClassDependency(err), "expected false for non-dependency error") wrapped := errors.Join(errors.New("outer"), &missingStorageClassDependencyError{className: "x"}) - if !isMissingStorageClassDependency(wrapped) { - t.Fatal("expected true for wrapped missing dependency error") - } + c.True( + isMissingStorageClassDependency(wrapped), + "expected true for wrapped missing dependency error", + ) } func findCondition(conditions []metav1.Condition, conditionType string) *metav1.Condition { @@ -753,6 +714,7 @@ func TestShardReconciler_FieldOwnershipIsolation(t *testing.T) { t.Run("updateStatus patch contains only Available condition", func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) shard := &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{ @@ -821,27 +783,16 @@ func TestShardReconciler_FieldOwnershipIsolation(t *testing.T) { Recorder: record.NewFakeRecorder(100), } - if err := r.updateStatus(t.Context(), shard, renderedConfig{}); err != nil { - t.Fatalf("updateStatus: %v", err) - } + ck.NoError(r.updateStatus(t.Context(), shard, renderedConfig{}), "updateStatus") patchShard, ok := capturedPatchObj.(*multigresv1alpha1.Shard) - if !ok { - t.Fatalf("expected *Shard patch, got %T", capturedPatchObj) - } + ck.True(ok, "expected *Shard patch, got %T", capturedPatchObj) for _, c := range patchShard.Status.Conditions { - if c.Type == conditionStorageClassValid { - t.Fatalf( - "updateStatus patch must not contain %s condition", - conditionStorageClassValid, - ) - } + ck.NotEq(conditionStorageClassValid, c.Type, "updateStatus patch must not contain") } availCond := findCondition(patchShard.Status.Conditions, "Available") - if availCond == nil { - t.Fatal("updateStatus patch must contain Available condition") - } + ck.NotNil(availCond, "updateStatus patch must contain Available condition") }) } @@ -904,9 +855,8 @@ func childExists(t *testing.T, c client.Client, obj client.Object, namespace, na if err == nil { return true } - if !apierrors.IsNotFound(err) { - t.Fatalf("unexpected error reading %T %s: %v", obj, name, err) - } + assert.NewAborting(t). + True(apierrors.IsNotFound(err), "unexpected error reading %T %s: %v", obj, name, err) return false } @@ -918,6 +868,7 @@ func childExists(t *testing.T, c client.Client, obj client.Object, namespace, na // a missing pool class must stop after it and before anything that consumes // pool storage. func TestReconcile_MissingBackupStorageClassStopsBeforeTheSharedBackupPVC(t *testing.T) { + ck := assert.NewCollecting(t) scheme := storageClassGateScheme() shard := storageClassGateShard("missing-backup-sc", "") @@ -938,32 +889,32 @@ func TestReconcile_MissingBackupStorageClassStopsBeforeTheSharedBackupPVC(t *tes result, err := r.Reconcile(t.Context(), ctrl.Request{ NamespacedName: client.ObjectKeyFromObject(shard), }) - if err != nil { - t.Fatalf("expected non-error dependency requeue, got error: %v", err) - } - if result.RequeueAfter != storageClassDependencyRequeue { - t.Fatalf("requeueAfter = %v, want %v", result.RequeueAfter, storageClassDependencyRequeue) - } + ck.Require().NoError(err, "expected non-error dependency requeue, got error") + ck.Require().Eq(storageClassDependencyRequeue, result.RequeueAfter, "requeueAfter") ns := shard.Namespace - if !childExists(t, c, &corev1.ConfigMap{}, ns, PgHbaConfigMapName(shard.Name)) { - t.Error("pg_hba ConfigMap is missing: the backup gate moved above the shared ConfigMaps") - } - if childExists(t, c, &appsv1.Deployment{}, ns, buildHashedMultiorchName(shard, "zone1")) { - t.Error("Multiorch Deployment was created: the backup gate moved below the Multiorch block") - } + ck.True( + childExists(t, c, &corev1.ConfigMap{}, ns, PgHbaConfigMapName(shard.Name)), + "pg_hba ConfigMap is missing: the backup gate moved above the shared ConfigMaps", + ) + ck.False( + childExists(t, c, &appsv1.Deployment{}, ns, buildHashedMultiorchName(shard, "zone1")), + "Multiorch Deployment was created: the backup gate moved below the Multiorch block", + ) if childExists(t, c, &corev1.PersistentVolumeClaim{}, ns, BuildSharedBackupPVCName(shard)) { t.Error( "shared backup PVC was created against a StorageClass that does not exist: " + "the backup gate moved below the backup PVC block", ) } - if childExists(t, c, &corev1.ConfigMap{}, ns, PostgresConfigMapName(shard.Name)) { - t.Error("postgres config ConfigMap was created: the reconcile ran past both gates") - } + ck.False( + childExists(t, c, &corev1.ConfigMap{}, ns, PostgresConfigMapName(shard.Name)), + "postgres config ConfigMap was created: the reconcile ran past both gates", + ) } func TestReconcile_MissingPoolStorageClassStopsAfterTheSharedBackupPVC(t *testing.T) { + ck := assert.NewCollecting(t) scheme := storageClassGateScheme() shard := storageClassGateShard("backup-sc", "missing-pool-sc") @@ -988,26 +939,24 @@ func TestReconcile_MissingPoolStorageClassStopsAfterTheSharedBackupPVC(t *testin result, err := r.Reconcile(t.Context(), ctrl.Request{ NamespacedName: client.ObjectKeyFromObject(shard), }) - if err != nil { - t.Fatalf("expected non-error dependency requeue, got error: %v", err) - } - if result.RequeueAfter != storageClassDependencyRequeue { - t.Fatalf("requeueAfter = %v, want %v", result.RequeueAfter, storageClassDependencyRequeue) - } + ck.Require().NoError(err, "expected non-error dependency requeue, got error") + ck.Require().Eq(storageClassDependencyRequeue, result.RequeueAfter, "requeueAfter") ns := shard.Namespace - if !childExists(t, c, &appsv1.Deployment{}, ns, buildHashedMultiorchName(shard, "zone1")) { - t.Error("Multiorch Deployment is missing: the pool gate moved above the Multiorch block") - } + ck.True( + childExists(t, c, &appsv1.Deployment{}, ns, buildHashedMultiorchName(shard, "zone1")), + "Multiorch Deployment is missing: the pool gate moved above the Multiorch block", + ) if !childExists(t, c, &corev1.PersistentVolumeClaim{}, ns, BuildSharedBackupPVCName(shard)) { t.Error( "shared backup PVC is missing: a missing pool class must not stop the reconcile " + "before the backup PVC, whose own StorageClass is present", ) } - if childExists(t, c, &corev1.ConfigMap{}, ns, PostgresConfigMapName(shard.Name)) { - t.Error("postgres config ConfigMap was created: the pool gate moved below it") - } + ck.False( + childExists(t, c, &corev1.ConfigMap{}, ns, PostgresConfigMapName(shard.Name)), + "postgres config ConfigMap was created: the pool gate moved below it", + ) if childExists(t, c, &corev1.Pod{}, ns, BuildPoolPodName(shard, "primary", "zone1", 0)) { t.Error( "pool pod was created against a StorageClass that does not exist: " + @@ -1027,15 +976,12 @@ func persistStorageClassCondition( build func(generation int64) metav1.Condition, ) { t.Helper() + ck := assert.NewAborting(t) var shard multigresv1alpha1.Shard - if err := c.Get(t.Context(), key, &shard); err != nil { - t.Fatalf("read shard: %v", err) - } + ck.NoError(c.Get(t.Context(), key, &shard), "read shard") shard.Status.Conditions = []metav1.Condition{build(shard.Generation)} - if err := c.Status().Update(t.Context(), &shard); err != nil { - t.Fatalf("seed condition: %v", err) - } + ck.NoError(c.Status().Update(t.Context(), &shard), "seed condition") } // TestSetStorageClassCondition_RepublishesWhenOneComparedFieldDiffers takes the @@ -1139,6 +1085,7 @@ func TestSetStorageClassCondition_RepublishesWhenOneComparedFieldDiffers(t *test for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1172,29 +1119,24 @@ func TestSetStorageClassCondition_RepublishesWhenOneComparedFieldDiffers(t *test Recorder: record.NewFakeRecorder(10), } - if err := r.setStorageClassCondition(t.Context(), shard, settled); err != nil { - t.Fatalf("setStorageClassCondition: %v", err) - } + c.NoError( + r.setStorageClassCondition(t.Context(), shard, settled), + "setStorageClassCondition", + ) want := 0 if tc.wantPatch { want = 1 } - if patches != want { - t.Fatalf("applied %d patches, want %d", patches, want) - } + c.Eq(want, patches, "applied") if !tc.wantPatch { return } var got multigresv1alpha1.Shard - if err := baseClient.Get(t.Context(), key, &got); err != nil { - t.Fatalf("read shard: %v", err) - } + c.NoError(baseClient.Get(t.Context(), key, &got), "read shard") cond := findCondition(got.Status.Conditions, conditionStorageClassValid) - if cond == nil { - t.Fatalf("no %s condition", conditionStorageClassValid) - } + c.NotNil(cond, "no %s condition", conditionStorageClassValid) if cond.Status != settled.status || cond.Reason != settled.reason || cond.Message != settled.message || cond.ObservedGeneration != got.Generation { t.Fatalf("published condition does not match the verdict: %+v", *cond) @@ -1210,6 +1152,7 @@ func TestSetStorageClassCondition_RepublishesWhenOneComparedFieldDiffers(t *test // message and lastTransitionTime all have to move with it. func TestStorageClassCondition_UpdatesWhenTheVerdictChanges(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1246,31 +1189,20 @@ func TestStorageClassCondition_UpdatesWhenTheVerdictChanges(t *testing.T) { t.Helper() check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("validate: %v", err) - } - if err := r.setStorageClassCondition(t.Context(), shard, check); err != nil { - t.Fatalf("set condition: %v", err) - } + c.Require().NoError(err, "validate") + c.Require().NoError(r.setStorageClassCondition(t.Context(), shard, check), "set condition") var got multigresv1alpha1.Shard - if err := baseClient.Get(t.Context(), key, &got); err != nil { - t.Fatalf("read shard: %v", err) - } + c.Require().NoError(baseClient.Get(t.Context(), key, &got), "read shard") cond := findCondition(got.Status.Conditions, conditionStorageClassValid) - if cond == nil { - t.Fatalf("no %s condition", conditionStorageClassValid) - } + c.Require().NotNil(cond, "no %s condition", conditionStorageClassValid) return *cond } before := cycle() - if before.Status != metav1.ConditionFalse || before.Reason != storageClassNotFoundReason { - t.Fatalf("first cycle must report the missing class: %+v", before) - } - if patches != 1 { - t.Fatalf("first cycle applied %d patches, want 1", patches) - } + c.Require(). + False(before.Status != metav1.ConditionFalse || before.Reason != storageClassNotFoundReason, "first cycle must report the missing class: %+v", before) + c.Require().Eq(1, patches, "first cycle applied") // Backdated because metav1.Time serialises at second precision and both // cycles run inside the same second, which would make a rewritten @@ -1287,35 +1219,22 @@ func TestStorageClassCondition_UpdatesWhenTheVerdictChanges(t *testing.T) { } }) - if err := baseClient.Create(t.Context(), &storagev1.StorageClass{ + c.Require().NoError(baseClient.Create(t.Context(), &storagev1.StorageClass{ ObjectMeta: metav1.ObjectMeta{Name: "appears-later"}, - }); err != nil { - t.Fatalf("create StorageClass: %v", err) - } + }), "create StorageClass") after := cycle() - if patches != 2 { - t.Fatalf("the changed verdict was skipped: %d patches total", patches) - } - if after.Status != metav1.ConditionTrue { - t.Errorf("status = %s, want %s", after.Status, metav1.ConditionTrue) - } - if after.Reason != storageClassFoundReason { - t.Errorf("reason = %s, want %s", after.Reason, storageClassFoundReason) - } - if after.Message == before.Message { - t.Errorf("message did not move off the missing-class text: %q", after.Message) - } - if want := "All explicitly configured StorageClasses are present"; after.Message != want { - t.Errorf("message = %q, want %q", after.Message, want) - } - if !after.LastTransitionTime.After(backdated.Time) { - t.Errorf( - "lastTransitionTime = %s, want it moved past %s: the condition transitioned", - after.LastTransitionTime, - backdated, - ) - } + c.Require().Eq(2, patches, "the changed verdict was skipped") + c.Eq(metav1.ConditionTrue, after.Status, "status") + c.Eq(storageClassFoundReason, after.Reason, "reason") + c.NotEq(before.Message, after.Message, "message did not move off the missing-class text") + c.Eq("All explicitly configured StorageClasses are present", after.Message, "message") + c.True( + after.LastTransitionTime.After(backdated.Time), + "lastTransitionTime = %s, want it moved past %s: the condition transitioned", + after.LastTransitionTime, + backdated, + ) } // TestStorageClassCondition_PreservesLastTransitionTimeWithoutATransition @@ -1325,6 +1244,7 @@ func TestStorageClassCondition_UpdatesWhenTheVerdictChanges(t *testing.T) { // lastTransitionTime stays put, matching meta.SetStatusCondition. func TestStorageClassCondition_PreservesLastTransitionTimeWithoutATransition(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -1369,30 +1289,16 @@ func TestStorageClassCondition_PreservesLastTransitionTimeWithoutATransition(t * r := &ShardReconciler{Client: fakeClient, Scheme: scheme, Recorder: record.NewFakeRecorder(50)} check, err := r.validateStorageClassDependencies(t.Context(), shard) - if err != nil { - t.Fatalf("validate: %v", err) - } - if err := r.setStorageClassCondition(t.Context(), shard, check); err != nil { - t.Fatalf("set condition: %v", err) - } - if patches != 1 { - t.Fatalf("the stale message was not republished: %d patches", patches) - } + c.NoError(err, "validate") + c.NoError(r.setStorageClassCondition(t.Context(), shard, check), "set condition") + c.Eq(1, patches, "the stale message was not republished") var got multigresv1alpha1.Shard - if err := baseClient.Get(t.Context(), key, &got); err != nil { - t.Fatalf("read shard: %v", err) - } + c.NoError(baseClient.Get(t.Context(), key, &got), "read shard") cond := findCondition(got.Status.Conditions, conditionStorageClassValid) - if cond == nil { - t.Fatalf("no %s condition", conditionStorageClassValid) - } - if cond.Message == staleMessage { - t.Fatalf("message was not rewritten: %q", cond.Message) - } - if cond.Status != metav1.ConditionTrue { - t.Fatalf("status = %s, want %s", cond.Status, metav1.ConditionTrue) - } + c.NotNil(cond, "no %s condition", conditionStorageClassValid) + c.NotEq(staleMessage, cond.Message, "message was not rewritten") + c.Eq(metav1.ConditionTrue, cond.Status, "status") if !cond.LastTransitionTime.Time.Equal(backdated.Time) { t.Errorf( "lastTransitionTime = %s, want it preserved at %s: the status did not transition", diff --git a/pkg/resource-handler/controller/shard/topo_client_tls_test.go b/pkg/resource-handler/controller/shard/topo_client_tls_test.go index 4e3f6d40..eae910ad 100644 --- a/pkg/resource-handler/controller/shard/topo_client_tls_test.go +++ b/pkg/resource-handler/controller/shard/topo_client_tls_test.go @@ -6,6 +6,8 @@ import ( corev1 "k8s.io/api/core/v1" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) // topoTLSShard returns a shard whose global topology reference carries the @@ -61,18 +63,16 @@ func hasArg(args []string, flag string) bool { // server, never a container that also mounts postgres state, and the pod never // shares a process namespace. func TestPoolPod_TopoClientTLSMountBoundary(t *testing.T) { + ck := assert.NewCollecting(t) pod, err := BuildPoolPod(topoTLSShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("BuildPoolPod() error = %v", err) - } + ck.Require().NoError(err, "BuildPoolPod() error =") multipooler := containerByName(pod.Spec.Containers, "multipooler") - if multipooler == nil { - t.Fatal("multipooler container missing") - } - if !mountsVolume(multipooler, multigresv1alpha1.TopoClientTLSVolumeName) { - t.Error("topo client certificate is not mounted into multipooler") - } + ck.Require().NotNil(multipooler, "multipooler container missing") + ck.True( + mountsVolume(multipooler, multigresv1alpha1.TopoClientTLSVolumeName), + "topo client certificate is not mounted into multipooler", + ) // No other container in the pod may mount the credential. pgctld runs as an // init (native sidecar) container named postgres; the exporter sits beside @@ -93,25 +93,24 @@ func TestPoolPod_TopoClientTLSMountBoundary(t *testing.T) { // also mount: those share the data PVC and the socket directory, so anything // projected onto them is readable by the postgres superuser. postgres := containerByName(pod.Spec.InitContainers, "postgres") - if postgres == nil { - t.Fatal("postgres (pgctld) container missing") - } + ck.Require().NotNil(postgres, "postgres (pgctld) container missing") sharedWithPostgres := map[string]struct{}{} for _, m := range postgres.VolumeMounts { sharedWithPostgres[m.Name] = struct{}{} } - if _, shared := sharedWithPostgres[multigresv1alpha1.TopoClientTLSVolumeName]; shared { - t.Error("topo client certificate shares a volume with the postgres container") - } + _, shared := sharedWithPostgres[multigresv1alpha1.TopoClientTLSVolumeName] + ck.False(shared, "topo client certificate shares a volume with the postgres container") for _, m := range multipooler.VolumeMounts { if m.Name != multigresv1alpha1.TopoClientTLSVolumeName { continue } // The mount path must be its own, not nested under the shared data or // socket directories. - if m.MountPath == DataMountPath || m.MountPath == SocketDirMountPath { - t.Errorf("topo client certificate mounted on a shared path %q", m.MountPath) - } + ck.False( + m.MountPath == DataMountPath || m.MountPath == SocketDirMountPath, + "topo client certificate mounted on a shared path %q", + m.MountPath, + ) } // Shared PID plus ptrace would expose one container's mounted files through @@ -124,15 +123,16 @@ func TestPoolPod_TopoClientTLSMountBoundary(t *testing.T) { // With no topo client credential on the reference (topology TLS off), the pool // pod renders exactly as before: no topo client volume, mount or flags anywhere. func TestPoolPod_TopoClientTLSOffRendersUnchanged(t *testing.T) { + ck := assert.NewAborting(t) pod, err := BuildPoolPod(newTestShard(), "main", "z1", newTestPoolSpec(), 0, testScheme()) - if err != nil { - t.Fatalf("BuildPoolPod() error = %v", err) - } + ck.NoError(err, "BuildPoolPod() error =") for _, v := range pod.Spec.Volumes { - if v.Name == multigresv1alpha1.TopoClientTLSVolumeName { - t.Fatal("topo client volume present with topology TLS off") - } + ck.NotEq( + multigresv1alpha1.TopoClientTLSVolumeName, + v.Name, + "topo client volume present with topology TLS off", + ) } all := append([]corev1.Container{}, pod.Spec.InitContainers...) all = append(all, pod.Spec.Containers...) @@ -150,6 +150,7 @@ func TestPoolPod_TopoClientTLSOffRendersUnchanged(t *testing.T) { // multipooler and multiorch both open topology connections, so both present the // client certificate through the three etcd TLS flags when it is configured. func TestTopoSpeakingContainers_PresentClientCert(t *testing.T) { + c := assert.NewCollecting(t) shard := topoTLSShard() pool := newTestPoolSpec() @@ -160,16 +161,19 @@ func TestTopoSpeakingContainers_PresentClientCert(t *testing.T) { "multipooler": mp.Args, "multiorch": orch.Args, } { - if !hasArg(args, "--topo-etcd-tls-cert") || + c.False(!hasArg(args, "--topo-etcd-tls-cert") || !hasArg(args, "--topo-etcd-tls-key") || - !hasArg(args, "--topo-etcd-tls-ca") { - t.Errorf("%s is missing topo client TLS flags: %v", name, args) - } - } - if !mountsVolume(&mp, multigresv1alpha1.TopoClientTLSVolumeName) { - t.Error("multipooler does not mount the topo client certificate") - } - if !mountsVolume(&orch, multigresv1alpha1.TopoClientTLSVolumeName) { - t.Error("multiorch does not mount the topo client certificate") - } + !hasArg( + args, + "--topo-etcd-tls-ca", + ), "%s is missing topo client TLS flags: %v", name, args) + } + c.True( + mountsVolume(&mp, multigresv1alpha1.TopoClientTLSVolumeName), + "multipooler does not mount the topo client certificate", + ) + c.True( + mountsVolume(&orch, multigresv1alpha1.TopoClientTLSVolumeName), + "multiorch does not mount the topo client certificate", + ) } diff --git a/pkg/resource-handler/controller/storage/pvc_test.go b/pkg/resource-handler/controller/storage/pvc_test.go index f6918cff..4ec6cabe 100644 --- a/pkg/resource-handler/controller/storage/pvc_test.go +++ b/pkg/resource-handler/controller/storage/pvc_test.go @@ -3,10 +3,11 @@ package storage import ( "testing" - "github.com/google/go-cmp/cmp" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/resource" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + + "github.com/multigres/testkit/assert" ) func TestBuildPVCTemplate(t *testing.T) { @@ -152,9 +153,7 @@ func TestBuildPVCTemplate(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := BuildPVCTemplate(tc.name, tc.storageClassName, tc.storageSize, tc.accessModes) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildPVCTemplate() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildPVCTemplate() mismatch") }) } } diff --git a/pkg/resource-handler/controller/toposerver/certificate_test.go b/pkg/resource-handler/controller/toposerver/certificate_test.go index 9a50bc28..5d461504 100644 --- a/pkg/resource-handler/controller/toposerver/certificate_test.go +++ b/pkg/resource-handler/controller/toposerver/certificate_test.go @@ -19,6 +19,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/certs" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func certScheme() *runtime.Scheme { @@ -52,6 +54,7 @@ func certTestTopoServer(tls *multigresv1alpha1.TopoTLSConfig) *multigresv1alpha1 } func TestBuildServingCertificate(t *testing.T) { + c := assert.NewCollecting(t) scheme := certScheme() toposerver := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{ Enabled: ptr.To(true), @@ -59,29 +62,18 @@ func TestBuildServingCertificate(t *testing.T) { }) got, err := BuildServingCertificate(toposerver, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } - if got == nil { - t.Fatal("BuildServingCertificate() = nil, want a Certificate") - } + c.Require().NoError(err, "BuildServingCertificate() error =") + c.Require().NotNil(got, "BuildServingCertificate() = nil, want a Certificate") wantName := "test-cluster-global-topo-topo-server-tls" - if got.GetName() != wantName { - t.Errorf("name = %q, want %q", got.GetName(), wantName) - } - if got.GetNamespace() != "supabase" { - t.Errorf("namespace = %q, want supabase", got.GetNamespace()) - } + c.Eq(wantName, got.GetName(), "name") + c.Eq("supabase", got.GetNamespace(), "namespace") ownerRefs := got.GetOwnerReferences() - if len(ownerRefs) != 1 || ownerRefs[0].Kind != "TopoServer" { - t.Fatalf("ownerReferences = %+v, want one TopoServer ref", ownerRefs) - } + c.Require(). + False(len(ownerRefs) != 1 || ownerRefs[0].Kind != "TopoServer", "ownerReferences = %+v, want one TopoServer ref", ownerRefs) spec, ok := got.Object["spec"].(map[string]any) - if !ok { - t.Fatal("spec is not a map") - } + c.Require().True(ok, "spec is not a map") // Both Services the controller creates have to verify: the client Service // (BuildClientService) and the headless peer Service (BuildHeadlessService). @@ -96,20 +88,18 @@ func TestBuildServingCertificate(t *testing.T) { "test-cluster-global-topo-headless.supabase.svc.cluster.local", "*.test-cluster-global-topo-headless.supabase.svc.cluster.local", } - if diff := cmp.Diff(wantDNSNames, spec["dnsNames"]); diff != "" { - t.Errorf("dnsNames mismatch (-want +got):\n%s", diff) - } + c.Eq("", cmp.Diff(wantDNSNames, spec["dnsNames"]), "dnsNames mismatch (-want +got):\n") wantSubject := fmt.Sprintf( certs.LiteralSubjectTemplate, "test-cluster-global-topo.supabase.svc.cluster.local", ) - if diff := cmp.Diff(wantSubject, spec["literalSubject"]); diff != "" { - t.Errorf("literalSubject mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(wantName, spec["secretName"]); diff != "" { - t.Errorf("secretName mismatch (-want +got):\n%s", diff) - } + c.Eq( + "", + cmp.Diff(wantSubject, spec["literalSubject"]), + "literalSubject mismatch (-want +got):\n", + ) + c.Eq("", cmp.Diff(wantName, spec["secretName"]), "secretName mismatch (-want +got):\n") // The topology server is shared infrastructure, so it takes the issuer from // the topology TLS config rather than any single cluster's issuer. @@ -118,32 +108,23 @@ func TestBuildServingCertificate(t *testing.T) { "kind": "ClusterIssuer", "group": "cert-manager.io", } - if diff := cmp.Diff(wantIssuerRef, spec["issuerRef"]); diff != "" { - t.Errorf("issuerRef mismatch (-want +got):\n%s", diff) - } + c.Eq("", cmp.Diff(wantIssuerRef, spec["issuerRef"]), "issuerRef mismatch (-want +got):\n") } func TestBuildServingCertificateSANsCoverBothServices(t *testing.T) { + c := assert.NewAborting(t) scheme := certScheme() toposerver := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) clientSvc, err := BuildClientService(toposerver, scheme) - if err != nil { - t.Fatalf("BuildClientService() error = %v", err) - } + c.NoError(err, "BuildClientService() error =") headlessSvc, err := BuildHeadlessService(toposerver, scheme) - if err != nil { - t.Fatalf("BuildHeadlessService() error = %v", err) - } + c.NoError(err, "BuildHeadlessService() error =") cert, err := BuildServingCertificate(toposerver, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } + c.NoError(err, "BuildServingCertificate() error =") sans, _, err := unstructured.NestedSlice(cert.Object, "spec", "dnsNames") - if err != nil { - t.Fatalf("NestedSlice(dnsNames) error = %v", err) - } + c.NoError(err, "NestedSlice(dnsNames) error =") covered := make(map[string]struct{}, len(sans)) for _, s := range sans { covered[s.(string)] = struct{}{} @@ -158,17 +139,14 @@ func TestBuildServingCertificateSANsCoverBothServices(t *testing.T) { } func TestBuildServingCertificateDefaultIssuer(t *testing.T) { + c := assert.NewCollecting(t) scheme := certScheme() toposerver := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) cert, err := BuildServingCertificate(toposerver, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } + c.Require().NoError(err, "BuildServingCertificate() error =") issuer, _, _ := unstructured.NestedString(cert.Object, "spec", "issuerRef", "name") - if issuer != certs.DefaultIssuerName { - t.Errorf("issuerRef.name = %q, want %q", issuer, certs.DefaultIssuerName) - } + c.Eq(certs.DefaultIssuerName, issuer, "issuerRef.name") } func TestBuildServingCertificateDisabled(t *testing.T) { @@ -179,13 +157,10 @@ func TestBuildServingCertificateDisabled(t *testing.T) { "disabled": {Enabled: ptr.To(false)}, } { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) got, err := BuildServingCertificate(certTestTopoServer(tls), scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } - if got != nil { - t.Errorf("BuildServingCertificate() = %v, want nil", got) - } + c.Require().NoError(err, "BuildServingCertificate() error =") + c.Nil(got, "BuildServingCertificate()") }) } } @@ -194,6 +169,7 @@ func TestReconcileCertificate(t *testing.T) { certName := "test-cluster-global-topo-topo-server-tls" t.Run("applies the certificate when enabled", func(t *testing.T) { + ck := assert.NewAborting(t) scheme := certScheme() toposerver := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(toposerver).Build() @@ -203,28 +179,26 @@ func TestReconcileCertificate(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), toposerver); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.NoError( + r.reconcileCertificate(context.Background(), toposerver), + "reconcileCertificate() error =", + ) got := &unstructured.Unstructured{} got.SetGroupVersionKind(certs.GVK) - if err := c.Get( + ck.NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: certName}, got, - ); err != nil { - t.Fatalf("expected serving Certificate, got error %v", err) - } + ), "expected serving Certificate, got error") }) t.Run("reports a foreign certificate collision when enabled", func(t *testing.T) { + ck := assert.NewCollecting(t) scheme := certScheme() toposerver := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) foreign, err := BuildServingCertificate(toposerver, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } + ck.Require().NoError(err, "BuildServingCertificate() error =") foreign.SetOwnerReferences([]metav1.OwnerReference{{ APIVersion: "example.com/v1", Kind: "Other", @@ -243,31 +217,24 @@ func TestReconcileCertificate(t *testing.T) { key := client.ObjectKey{Namespace: "supabase", Name: certName} before := &unstructured.Unstructured{} before.SetGroupVersionKind(certs.GVK) - if err := c.Get(context.Background(), key, before); err != nil { - t.Fatalf("Get() error = %v", err) - } + ck.Require().NoError(c.Get(context.Background(), key, before), "Get() error =") - if err := r.reconcileCertificate(context.Background(), toposerver); err == nil { - t.Fatal("reconcileCertificate() error = nil, want collision error") - } + ck.Require(). + Error(r.reconcileCertificate(context.Background(), toposerver), "reconcileCertificate() error = nil, want collision error") got := &unstructured.Unstructured{} got.SetGroupVersionKind(certs.GVK) - if err := c.Get(context.Background(), key, got); err != nil { - t.Fatalf("foreign Certificate was modified or deleted: %v", err) - } - if diff := cmp.Diff(before.Object, got.Object); diff != "" { - t.Errorf("foreign Certificate changed (-want +got):\n%s", diff) - } + ck.Require(). + NoError(c.Get(context.Background(), key, got), "foreign Certificate was modified or deleted") + ck.EqDiff(before.Object, got.Object, "foreign Certificate changed") }) t.Run("prunes the certificate and its secret when disabled", func(t *testing.T) { + ck := assert.NewCollecting(t) scheme := certScheme() enabled := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) existing, err := BuildServingCertificate(enabled, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } + ck.Require().NoError(err, "BuildServingCertificate() error =") secret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{Name: certName, Namespace: "supabase"}, } @@ -283,29 +250,25 @@ func TestReconcileCertificate(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), toposerver); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.Require(). + NoError(r.reconcileCertificate(context.Background(), toposerver), "reconcileCertificate() error =") got := &unstructured.Unstructured{} got.SetGroupVersionKind(certs.GVK) - if err := c.Get( + ck.Error(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: certName}, got, - ); err == nil { - t.Error("expected serving Certificate to be deleted") - } - if err := c.Get( + ), "expected serving Certificate to be deleted") + ck.Error(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: certName}, &corev1.Secret{}, - ); err == nil { - t.Error("expected generated Secret to be deleted") - } + ), "expected generated Secret to be deleted") }) t.Run("issues nothing when TLS is unset", func(t *testing.T) { + ck := assert.NewCollecting(t) scheme := certScheme() toposerver := certTestTopoServer(nil) c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(toposerver).Build() @@ -315,29 +278,23 @@ func TestReconcileCertificate(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), toposerver); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.Require(). + NoError(r.reconcileCertificate(context.Background(), toposerver), "reconcileCertificate() error =") list := &unstructured.UnstructuredList{} list.SetGroupVersionKind(certs.GVK) - if err := c.List(context.Background(), list); err != nil { - t.Fatalf("List() error = %v", err) - } - if len(list.Items) != 0 { - t.Errorf("got %d Certificates, want 0", len(list.Items)) - } + ck.Require().NoError(c.List(context.Background(), list), "List() error =") + ck.Empty(list.Items, "got %d Certificates, want 0", len(list.Items)) }) // Removing the TLS block is as common as setting enabled to false, and it // must not leave private key material behind either. t.Run("cleans up when the TLS block is removed", func(t *testing.T) { + ck := assert.NewCollecting(t) scheme := certScheme() enabled := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) existing, err := BuildServingCertificate(enabled, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } + ck.Require().NoError(err, "BuildServingCertificate() error =") secret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{Name: certName, Namespace: "supabase"}, } @@ -353,29 +310,29 @@ func TestReconcileCertificate(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), toposerver); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.Require(). + NoError(r.reconcileCertificate(context.Background(), toposerver), "reconcileCertificate() error =") key := client.ObjectKey{Namespace: "supabase", Name: certName} got := &unstructured.Unstructured{} got.SetGroupVersionKind(certs.GVK) - if err := c.Get(context.Background(), key, got); err == nil { - t.Error("expected orphaned Certificate to be deleted") - } - if err := c.Get(context.Background(), key, &corev1.Secret{}); err == nil { - t.Error("expected orphaned Secret to be deleted") - } + ck.Error( + c.Get(context.Background(), key, got), + "expected orphaned Certificate to be deleted", + ) + ck.Error( + c.Get(context.Background(), key, &corev1.Secret{}), + "expected orphaned Secret to be deleted", + ) }) // A same-named Certificate owned by something else is not ours to delete. t.Run("leaves an unowned certificate alone", func(t *testing.T) { + ck := assert.NewCollecting(t) scheme := certScheme() enabled := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) foreign, err := BuildServingCertificate(enabled, scheme) - if err != nil { - t.Fatalf("BuildServingCertificate() error = %v", err) - } + ck.Require().NoError(err, "BuildServingCertificate() error =") foreign.SetOwnerReferences(nil) toposerver := certTestTopoServer(nil) @@ -389,19 +346,16 @@ func TestReconcileCertificate(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.reconcileCertificate(context.Background(), toposerver); err != nil { - t.Fatalf("reconcileCertificate() error = %v", err) - } + ck.Require(). + NoError(r.reconcileCertificate(context.Background(), toposerver), "reconcileCertificate() error =") got := &unstructured.Unstructured{} got.SetGroupVersionKind(certs.GVK) - if err := c.Get( + ck.NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: certName}, got, - ); err != nil { - t.Errorf("unowned Certificate was deleted: %v", err) - } + ), "unowned Certificate was deleted") }) } @@ -409,40 +363,34 @@ func TestReconcileCertificate(t *testing.T) { // enforcement existed: plaintext listeners and no serving certificate mount. // This is the invariant that keeps the change safe to merge with the gate off. func TestTopoTLSOffRendersPlaintext(t *testing.T) { + c := assert.NewCollecting(t) scheme := certScheme() sts, err := BuildStatefulSet(certTestTopoServer(nil), scheme) - if err != nil { - t.Fatalf("BuildStatefulSet() error = %v", err) - } + c.Require().NoError(err, "BuildStatefulSet() error =") for _, vol := range sts.Spec.Template.Spec.Volumes { - if vol.Name == TopoServerTLSVolumeName { - t.Fatalf("serving certificate volume present with topology TLS off") - } + c.Require(). + NotEq(TopoServerTLSVolumeName, vol.Name, "serving certificate volume present with topology TLS off") } env := etcdEnvMap(t, sts) - if got := env["ETCD_LISTEN_CLIENT_URLS"]; got != "http://[::]:2379" { - t.Errorf("ETCD_LISTEN_CLIENT_URLS = %q, want plaintext", got) - } - if _, ok := env["ETCD_CLIENT_CERT_AUTH"]; ok { - t.Errorf("ETCD_CLIENT_CERT_AUTH set with topology TLS off") - } + c.Eq("http://[::]:2379", env["ETCD_LISTEN_CLIENT_URLS"], "ETCD_LISTEN_CLIENT_URLS") + _, ok := env["ETCD_CLIENT_CERT_AUTH"] + c.False(ok, "ETCD_CLIENT_CERT_AUTH set with topology TLS off") } // With topology TLS on, etcd serves its client and peer listeners over TLS, // requires a client certificate on each, and mounts the issued serving // certificate. The metrics listener stays plaintext so the probes keep working. func TestTopoTLSOnRequiresClientCerts(t *testing.T) { + c := assert.NewCollecting(t) scheme := certScheme() sts, err := BuildStatefulSet(certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{ Enabled: ptr.To(true), IssuerName: "multigres-infra-issuer", }), scheme) - if err != nil { - t.Fatalf("BuildStatefulSet() error = %v", err) - } + c.Require().NoError(err, "BuildStatefulSet() error =") var servingVol *corev1.Volume for i := range sts.Spec.Template.Spec.Volumes { @@ -450,9 +398,7 @@ func TestTopoTLSOnRequiresClientCerts(t *testing.T) { servingVol = &sts.Spec.Template.Spec.Volumes[i] } } - if servingVol == nil { - t.Fatal("serving certificate volume missing with topology TLS on") - } + c.Require().NotNil(servingVol, "serving certificate volume missing with topology TLS on") if servingVol.Secret == nil || servingVol.Secret.SecretName != multigresv1alpha1.TopoServerCertSecretName( "test-cluster-global-topo", @@ -464,18 +410,15 @@ func TestTopoTLSOnRequiresClientCerts(t *testing.T) { for _, m := range sts.Spec.Template.Spec.Containers[0].VolumeMounts { if m.Name == TopoServerTLSVolumeName { mounted = true - if !m.ReadOnly || m.MountPath != TopoServerTLSMountPath { - t.Errorf( - "serving certificate mount = %+v, want read-only at %s", - m, - TopoServerTLSMountPath, - ) - } + c.False( + !m.ReadOnly || m.MountPath != TopoServerTLSMountPath, + "serving certificate mount = %+v, want read-only at %s", + m, + TopoServerTLSMountPath, + ) } } - if !mounted { - t.Error("serving certificate is not mounted into the etcd container") - } + c.True(mounted, "serving certificate is not mounted into the etcd container") env := etcdEnvMap(t, sts) wantEnv := map[string]string{ @@ -492,9 +435,8 @@ func TestTopoTLSOnRequiresClientCerts(t *testing.T) { "ETCD_PEER_TRUSTED_CA_FILE": TopoServerTLSCAFile, } for k, want := range wantEnv { - if got := env[k]; got != want { - t.Errorf("%s = %q, want %q", k, got, want) - } + got := env[k] + c.Eq(want, got, "%s = %q, want", k, got) } } diff --git a/pkg/resource-handler/controller/toposerver/container_env_test.go b/pkg/resource-handler/controller/toposerver/container_env_test.go index f9351a66..8748d92f 100644 --- a/pkg/resource-handler/controller/toposerver/container_env_test.go +++ b/pkg/resource-handler/controller/toposerver/container_env_test.go @@ -3,13 +3,15 @@ package toposerver import ( "testing" - "github.com/google/go-cmp/cmp" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" corev1 "k8s.io/api/core/v1" "k8s.io/utils/ptr" + + "github.com/multigres/testkit/assert" ) func TestMaintenanceEnvironment(t *testing.T) { + c := assert.NewCollecting(t) ts := certTestTopoServer(nil) ts.Spec.Etcd.Maintenance = &multigresv1alpha1.EtcdMaintenanceConfig{ AutoCompactionMode: "revision", @@ -17,14 +19,10 @@ func TestMaintenanceEnvironment(t *testing.T) { QuotaBackendBytes: ptr.To(int64(512 << 20)), } sts, err := BuildStatefulSet(ts, certScheme()) - if err != nil { - t.Fatal(err) - } + c.Require().NoError(err) env := etcdEnvMap(t, sts) for key, want := range map[string]string{"ETCD_AUTO_COMPACTION_MODE": "revision", "ETCD_AUTO_COMPACTION_RETENTION": "20000", "ETCD_QUOTA_BACKEND_BYTES": "536870912"} { - if env[key] != want { - t.Errorf("%s=%q, want %q", key, env[key], want) - } + c.Eq(want, env[key], "%s=%q, want", key, env[key]) } for _, config := range []*multigresv1alpha1.EtcdMaintenanceConfig{ {AutoCompactionRetention: "0h"}, @@ -60,9 +58,7 @@ func TestBuildPodIdentityEnv(t *testing.T) { }, } - if diff := cmp.Diff(want, got); diff != "" { - t.Errorf("buildPodIdentityEnv() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(want, got, "buildPodIdentityEnv() mismatch") } func TestBuildEtcdConfigEnv(t *testing.T) { @@ -143,9 +139,7 @@ func TestBuildEtcdConfigEnv(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := buildEtcdConfigEnv(tc.toposerverName, tc.serviceName, tc.namespace, false) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("buildEtcdConfigEnv() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "buildEtcdConfigEnv() mismatch") }) } } @@ -211,9 +205,7 @@ func TestBuildEtcdClusterPeerList(t *testing.T) { tc.replicas, "http", ) - if got != tc.want { - t.Errorf("buildEtcdClusterPeerList() = %v, want %v", got, tc.want) - } + assert.NewCollecting(t).Eq(tc.want, got, "buildEtcdClusterPeerList()") }) } } @@ -369,9 +361,7 @@ func TestBuildContainerEnv(t *testing.T) { tc.serviceName, false, ) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildContainerEnv() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildContainerEnv() mismatch") }) } } diff --git a/pkg/resource-handler/controller/toposerver/integration_test.go b/pkg/resource-handler/controller/toposerver/integration_test.go index 96fa332c..5cc42eca 100644 --- a/pkg/resource-handler/controller/toposerver/integration_test.go +++ b/pkg/resource-handler/controller/toposerver/integration_test.go @@ -22,6 +22,8 @@ import ( toposervercontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/toposerver" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestSetupWithManager(t *testing.T) { @@ -39,15 +41,13 @@ func TestSetupWithManager(t *testing.T) { ), ) - if err := (&toposervercontroller.TopoServerReconciler{ + assert.NewAborting(t).NoError((&toposervercontroller.TopoServerReconciler{ Client: mgr.GetClient(), Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("toposerver-controller"), }).SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") } func TestTopoServerReconciliation(t *testing.T) { @@ -93,7 +93,9 @@ func TestTopoServerReconciliation(t *testing.T) { Replicas: ptr.To(int32(3)), ServiceName: "test-toposerver-headless", Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(toposerverLabels(t, "test-cluster")), + MatchLabels: metadata.GetSelectorLabels( + toposerverLabels(t, "test-cluster"), + ), }, PersistentVolumeClaimRetentionPolicy: &appsv1.StatefulSetPersistentVolumeClaimRetentionPolicy{ WhenDeleted: appsv1.RetainPersistentVolumeClaimRetentionPolicyType, @@ -137,14 +139,35 @@ func TestTopoServerReconciliation(t *testing.T) { }, {Name: "ETCD_NAME", Value: "$(POD_NAME)"}, {Name: "ETCD_DATA_DIR", Value: "/var/lib/etcd"}, - {Name: "ETCD_LISTEN_CLIENT_URLS", Value: "http://[::]:2379"}, - {Name: "ETCD_LISTEN_PEER_URLS", Value: "http://[::]:2380"}, - {Name: "ETCD_LISTEN_METRICS_URLS", Value: "http://[::]:2381"}, - {Name: "ETCD_ADVERTISE_CLIENT_URLS", Value: "http://$(POD_NAME).test-toposerver-headless.$(POD_NAMESPACE).svc.cluster.local:2379"}, - {Name: "ETCD_INITIAL_ADVERTISE_PEER_URLS", Value: "http://$(POD_NAME).test-toposerver-headless.$(POD_NAMESPACE).svc.cluster.local:2380"}, + { + Name: "ETCD_LISTEN_CLIENT_URLS", + Value: "http://[::]:2379", + }, + { + Name: "ETCD_LISTEN_PEER_URLS", + Value: "http://[::]:2380", + }, + { + Name: "ETCD_LISTEN_METRICS_URLS", + Value: "http://[::]:2381", + }, + { + Name: "ETCD_ADVERTISE_CLIENT_URLS", + Value: "http://$(POD_NAME).test-toposerver-headless.$(POD_NAMESPACE).svc.cluster.local:2379", + }, + { + Name: "ETCD_INITIAL_ADVERTISE_PEER_URLS", + Value: "http://$(POD_NAME).test-toposerver-headless.$(POD_NAMESPACE).svc.cluster.local:2380", + }, {Name: "ETCD_INITIAL_CLUSTER_STATE", Value: "new"}, - {Name: "ETCD_INITIAL_CLUSTER_TOKEN", Value: "test-toposerver"}, - {Name: "ETCD_INITIAL_CLUSTER", Value: "test-toposerver-0=http://test-toposerver-0.test-toposerver-headless.default.svc.cluster.local:2380,test-toposerver-1=http://test-toposerver-1.test-toposerver-headless.default.svc.cluster.local:2380,test-toposerver-2=http://test-toposerver-2.test-toposerver-headless.default.svc.cluster.local:2380"}, + { + Name: "ETCD_INITIAL_CLUSTER_TOKEN", + Value: "test-toposerver", + }, + { + Name: "ETCD_INITIAL_CLUSTER", + Value: "test-toposerver-0=http://test-toposerver-0.test-toposerver-headless.default.svc.cluster.local:2380,test-toposerver-1=http://test-toposerver-1.test-toposerver-headless.default.svc.cluster.local:2380,test-toposerver-2=http://test-toposerver-2.test-toposerver-headless.default.svc.cluster.local:2380", + }, {Name: "ETCD_AUTO_COMPACTION_MODE", Value: "periodic"}, {Name: "ETCD_AUTO_COMPACTION_RETENTION", Value: "1h"}, {Name: "ETCD_QUOTA_BACKEND_BYTES", Value: "2147483648"}, @@ -155,8 +178,10 @@ func TestTopoServerReconciliation(t *testing.T) { StartupProbe: &corev1.Probe{ ProbeHandler: corev1.ProbeHandler{ HTTPGet: &corev1.HTTPGetAction{ - Path: "/readyz", - Port: intstr.FromInt32(toposervercontroller.MetricsPort), + Path: "/readyz", + Port: intstr.FromInt32( + toposervercontroller.MetricsPort, + ), Scheme: corev1.URISchemeHTTP, }, }, @@ -168,8 +193,10 @@ func TestTopoServerReconciliation(t *testing.T) { LivenessProbe: &corev1.Probe{ ProbeHandler: corev1.ProbeHandler{ HTTPGet: &corev1.HTTPGetAction{ - Path: "/livez", - Port: intstr.FromInt32(toposervercontroller.MetricsPort), + Path: "/livez", + Port: intstr.FromInt32( + toposervercontroller.MetricsPort, + ), Scheme: corev1.URISchemeHTTP, }, }, @@ -181,8 +208,10 @@ func TestTopoServerReconciliation(t *testing.T) { ReadinessProbe: &corev1.Probe{ ProbeHandler: corev1.ProbeHandler{ HTTPGet: &corev1.HTTPGetAction{ - Path: "/readyz", - Port: intstr.FromInt32(toposervercontroller.MetricsPort), + Path: "/readyz", + Port: intstr.FromInt32( + toposervercontroller.MetricsPort, + ), Scheme: corev1.URISchemeHTTP, }, }, @@ -199,7 +228,9 @@ func TestTopoServerReconciliation(t *testing.T) { { ObjectMeta: metav1.ObjectMeta{Name: "data"}, Spec: corev1.PersistentVolumeClaimSpec{ - AccessModes: []corev1.PersistentVolumeAccessMode{corev1.ReadWriteOnce}, + AccessModes: []corev1.PersistentVolumeAccessMode{ + corev1.ReadWriteOnce, + }, Resources: corev1.VolumeResourceRequirements{ Requests: corev1.ResourceList{ corev1.ResourceStorage: resource.MustParse("10Gi"), @@ -245,7 +276,9 @@ func TestTopoServerReconciliation(t *testing.T) { tcpServicePort(t, "client", 2379), tcpServicePort(t, "peer", 2380), }, - Selector: metadata.GetSelectorLabels(toposerverLabels(t, "test-cluster")), + Selector: metadata.GetSelectorLabels( + toposerverLabels(t, "test-cluster"), + ), PublishNotReadyAddresses: true, }, }, @@ -278,7 +311,9 @@ func TestTopoServerReconciliation(t *testing.T) { Replicas: ptr.To(int32(3)), ServiceName: "delete-policy-topo-headless", Selector: &metav1.LabelSelector{ - MatchLabels: metadata.GetSelectorLabels(toposerverLabels(t, "test-cluster")), + MatchLabels: metadata.GetSelectorLabels( + toposerverLabels(t, "test-cluster"), + ), }, PersistentVolumeClaimRetentionPolicy: &appsv1.StatefulSetPersistentVolumeClaimRetentionPolicy{ WhenDeleted: appsv1.DeletePersistentVolumeClaimRetentionPolicyType, @@ -322,14 +357,35 @@ func TestTopoServerReconciliation(t *testing.T) { }, {Name: "ETCD_NAME", Value: "$(POD_NAME)"}, {Name: "ETCD_DATA_DIR", Value: "/var/lib/etcd"}, - {Name: "ETCD_LISTEN_CLIENT_URLS", Value: "http://[::]:2379"}, - {Name: "ETCD_LISTEN_PEER_URLS", Value: "http://[::]:2380"}, - {Name: "ETCD_LISTEN_METRICS_URLS", Value: "http://[::]:2381"}, - {Name: "ETCD_ADVERTISE_CLIENT_URLS", Value: "http://$(POD_NAME).delete-policy-topo-headless.$(POD_NAMESPACE).svc.cluster.local:2379"}, - {Name: "ETCD_INITIAL_ADVERTISE_PEER_URLS", Value: "http://$(POD_NAME).delete-policy-topo-headless.$(POD_NAMESPACE).svc.cluster.local:2380"}, + { + Name: "ETCD_LISTEN_CLIENT_URLS", + Value: "http://[::]:2379", + }, + { + Name: "ETCD_LISTEN_PEER_URLS", + Value: "http://[::]:2380", + }, + { + Name: "ETCD_LISTEN_METRICS_URLS", + Value: "http://[::]:2381", + }, + { + Name: "ETCD_ADVERTISE_CLIENT_URLS", + Value: "http://$(POD_NAME).delete-policy-topo-headless.$(POD_NAMESPACE).svc.cluster.local:2379", + }, + { + Name: "ETCD_INITIAL_ADVERTISE_PEER_URLS", + Value: "http://$(POD_NAME).delete-policy-topo-headless.$(POD_NAMESPACE).svc.cluster.local:2380", + }, {Name: "ETCD_INITIAL_CLUSTER_STATE", Value: "new"}, - {Name: "ETCD_INITIAL_CLUSTER_TOKEN", Value: "delete-policy-topo"}, - {Name: "ETCD_INITIAL_CLUSTER", Value: "delete-policy-topo-0=http://delete-policy-topo-0.delete-policy-topo-headless.default.svc.cluster.local:2380,delete-policy-topo-1=http://delete-policy-topo-1.delete-policy-topo-headless.default.svc.cluster.local:2380,delete-policy-topo-2=http://delete-policy-topo-2.delete-policy-topo-headless.default.svc.cluster.local:2380"}, + { + Name: "ETCD_INITIAL_CLUSTER_TOKEN", + Value: "delete-policy-topo", + }, + { + Name: "ETCD_INITIAL_CLUSTER", + Value: "delete-policy-topo-0=http://delete-policy-topo-0.delete-policy-topo-headless.default.svc.cluster.local:2380,delete-policy-topo-1=http://delete-policy-topo-1.delete-policy-topo-headless.default.svc.cluster.local:2380,delete-policy-topo-2=http://delete-policy-topo-2.delete-policy-topo-headless.default.svc.cluster.local:2380", + }, {Name: "ETCD_AUTO_COMPACTION_MODE", Value: "periodic"}, {Name: "ETCD_AUTO_COMPACTION_RETENTION", Value: "1h"}, {Name: "ETCD_QUOTA_BACKEND_BYTES", Value: "2147483648"}, @@ -340,8 +396,10 @@ func TestTopoServerReconciliation(t *testing.T) { StartupProbe: &corev1.Probe{ ProbeHandler: corev1.ProbeHandler{ HTTPGet: &corev1.HTTPGetAction{ - Path: "/readyz", - Port: intstr.FromInt32(toposervercontroller.MetricsPort), + Path: "/readyz", + Port: intstr.FromInt32( + toposervercontroller.MetricsPort, + ), Scheme: corev1.URISchemeHTTP, }, }, @@ -353,8 +411,10 @@ func TestTopoServerReconciliation(t *testing.T) { LivenessProbe: &corev1.Probe{ ProbeHandler: corev1.ProbeHandler{ HTTPGet: &corev1.HTTPGetAction{ - Path: "/livez", - Port: intstr.FromInt32(toposervercontroller.MetricsPort), + Path: "/livez", + Port: intstr.FromInt32( + toposervercontroller.MetricsPort, + ), Scheme: corev1.URISchemeHTTP, }, }, @@ -366,8 +426,10 @@ func TestTopoServerReconciliation(t *testing.T) { ReadinessProbe: &corev1.Probe{ ProbeHandler: corev1.ProbeHandler{ HTTPGet: &corev1.HTTPGetAction{ - Path: "/readyz", - Port: intstr.FromInt32(toposervercontroller.MetricsPort), + Path: "/readyz", + Port: intstr.FromInt32( + toposervercontroller.MetricsPort, + ), Scheme: corev1.URISchemeHTTP, }, }, @@ -384,9 +446,13 @@ func TestTopoServerReconciliation(t *testing.T) { { ObjectMeta: metav1.ObjectMeta{Name: "data"}, Spec: corev1.PersistentVolumeClaimSpec{ - AccessModes: []corev1.PersistentVolumeAccessMode{corev1.ReadWriteOnce}, + AccessModes: []corev1.PersistentVolumeAccessMode{ + corev1.ReadWriteOnce, + }, Resources: corev1.VolumeResourceRequirements{ - Requests: corev1.ResourceList{corev1.ResourceStorage: resource.MustParse("10Gi")}, + Requests: corev1.ResourceList{ + corev1.ResourceStorage: resource.MustParse("10Gi"), + }, }, VolumeMode: ptr.To(corev1.PersistentVolumeFilesystem), }, @@ -428,7 +494,9 @@ func TestTopoServerReconciliation(t *testing.T) { tcpServicePort(t, "client", 2379), tcpServicePort(t, "peer", 2380), }, - Selector: metadata.GetSelectorLabels(toposerverLabels(t, "test-cluster")), + Selector: metadata.GetSelectorLabels( + toposerverLabels(t, "test-cluster"), + ), PublishNotReadyAddresses: true, }, }, @@ -439,6 +507,7 @@ func TestTopoServerReconciliation(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) ctx := t.Context() mgr := testutil.SetUpEnvtestManager(t, scheme, testutil.WithCRDPaths( @@ -463,23 +532,17 @@ func TestTopoServerReconciliation(t *testing.T) { Scheme: mgr.GetScheme(), Recorder: mgr.GetEventRecorderFor("toposerver-controller"), } - if err := toposerverReconciler.SetupWithManager(mgr, controller.Options{ - // Needed for the parallel test runs + // Needed for the parallel test runs + c.Require().NoError(toposerverReconciler.SetupWithManager(mgr, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("Failed to create controller, %v", err) - } + }), "Failed to create controller") - if err := client.Create(ctx, tc.toposerver); err != nil { - t.Fatalf("Failed to create the initial item, %v", err) - } + c.Require(). + NoError(client.Create(ctx, tc.toposerver), "Failed to create the initial item") - if err := watcher.WaitForMatch(tc.wantResources...); err != nil { - t.Errorf("Resources mismatch:\n%v", err) - } + c.NoError(watcher.WaitForMatch(tc.wantResources...), "Resources mismatch:\n") }) } - } // Test helpers @@ -518,5 +581,10 @@ func tcpPort(t testing.TB, name string, port int32) corev1.ContainerPort { // tcpServicePort creates a TCP service port with named target func tcpServicePort(t testing.TB, name string, port int32) corev1.ServicePort { t.Helper() - return corev1.ServicePort{Name: name, Port: port, TargetPort: intstr.FromString(name), Protocol: corev1.ProtocolTCP} + return corev1.ServicePort{ + Name: name, + Port: port, + TargetPort: intstr.FromString(name), + Protocol: corev1.ProtocolTCP, + } } diff --git a/pkg/resource-handler/controller/toposerver/maintenance_client_test.go b/pkg/resource-handler/controller/toposerver/maintenance_client_test.go index 595b57f8..3c58d512 100644 --- a/pkg/resource-handler/controller/toposerver/maintenance_client_test.go +++ b/pkg/resource-handler/controller/toposerver/maintenance_client_test.go @@ -15,14 +15,15 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestMaintenanceTLSUsesUncachedCredentialAndFailsClosed(t *testing.T) { + c := assert.NewAborting(t) ts := certTestTopoServer(&multigresv1alpha1.TopoTLSConfig{Enabled: ptr.To(true)}) pub, key, err := ed25519.GenerateKey(rand.Reader) - if err != nil { - t.Fatal(err) - } + c.NoError(err) template := &x509.Certificate{ SerialNumber: big.NewInt(1), NotBefore: time.Now().Add(-time.Hour), @@ -36,13 +37,9 @@ func TestMaintenanceTLSUsesUncachedCredentialAndFailsClosed(t *testing.T) { }, } der, err := x509.CreateCertificate(rand.Reader, template, template, pub, key) - if err != nil { - t.Fatal(err) - } + c.NoError(err) keyDER, err := x509.MarshalPKCS8PrivateKey(key) - if err != nil { - t.Fatal(err) - } + c.NoError(err) certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) secret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{ @@ -65,6 +62,7 @@ func TestMaintenanceTLSUsesUncachedCredentialAndFailsClosed(t *testing.T) { {"invalid CA", func(s *corev1.Secret) { s.Data["ca.crt"] = []byte("invalid") }, true}, } { t.Run(test.name, func(t *testing.T) { + c := assert.NewAborting(t) s := secret.DeepCopy() test.change(s) r := &TopoServerReconciler{ @@ -72,13 +70,9 @@ func TestMaintenanceTLSUsesUncachedCredentialAndFailsClosed(t *testing.T) { APIReader: fake.NewClientBuilder().WithScheme(certScheme()).WithObjects(s).Build(), } cfg, err := r.maintenanceTLSConfig(t.Context(), ts) - if (err != nil) != test.wantErr { - t.Fatalf("error=%v", err) - } - if !test.wantErr && - (cfg.MinVersion < tls.VersionTLS12 || cfg.InsecureSkipVerify || cfg.RootCAs == nil || len(cfg.Certificates) != 1) { - t.Fatal("invalid TLS configuration") - } + c.ErrorWhen(test.wantErr, err, "error=") + c.False(!test.wantErr && + (cfg.MinVersion < tls.VersionTLS12 || cfg.InsecureSkipVerify || cfg.RootCAs == nil || len(cfg.Certificates) != 1), "invalid TLS configuration") }) } } diff --git a/pkg/resource-handler/controller/toposerver/maintenance_etcd_test.go b/pkg/resource-handler/controller/toposerver/maintenance_etcd_test.go index 3316d8e7..8d91674c 100644 --- a/pkg/resource-handler/controller/toposerver/maintenance_etcd_test.go +++ b/pkg/resource-handler/controller/toposerver/maintenance_etcd_test.go @@ -21,30 +21,27 @@ import ( "go.etcd.io/etcd/api/v3/v3rpc/rpctypes" clientv3 "go.etcd.io/etcd/client/v3" "k8s.io/client-go/tools/record" + + "github.com/multigres/testkit/assert" ) // These tests launch disposable local etcd processes, never a configured // Kubernetes cluster. Prefer the same binary as envtest when available. func startMaintenanceEtcd(t *testing.T) (*memberClients, []string) { t.Helper() + ck := assert.NewAborting(t) binary := filepath.Join(os.Getenv("KUBEBUILDER_ASSETS"), "etcd") if _, err := os.Stat(binary); err != nil { var lookupErr error binary, lookupErr = exec.LookPath("etcd") - if lookupErr != nil { - t.Fatal("integration test requires etcd or KUBEBUILDER_ASSETS") - } + ck.NoError(lookupErr, "integration test requires etcd or KUBEBUILDER_ASSETS") } allocateURL := func() string { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } + ck.NoError(err) url := "http://" + listener.Addr().String() - if err := listener.Close(); err != nil { - t.Fatal(err) - } + ck.NoError(listener.Close()) return url } endpoints, peers, cluster := make([]string, 3), make([]string, 3), make([]string, 3) @@ -56,9 +53,7 @@ func startMaintenanceEtcd(t *testing.T) (*memberClients, []string) { for i := range endpoints { dir := t.TempDir() logFile, err := os.Create(filepath.Join(dir, "etcd.log")) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) cmd := exec.Command( binary, "--name", @@ -104,9 +99,7 @@ func startMaintenanceEtcd(t *testing.T) (*memberClients, []string) { c := &memberClients{clients: map[string]*clientv3.Client{}} for _, ep := range endpoints { cl, err := clientv3.New(clientv3.Config{Endpoints: []string{ep}, DialTimeout: time.Second}) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) c.clients[ep] = cl if c.first == nil { c.first = cl @@ -126,6 +119,7 @@ func startMaintenanceEtcd(t *testing.T) (*memberClients, []string) { } func TestLiveEtcdCompactionAndMaintenance(t *testing.T) { + ck := assert.NewAborting(t) c, endpoints := startMaintenanceEtcd(t) ctx, cancel := context.WithTimeout(t.Context(), 90*time.Second) defer cancel() @@ -133,9 +127,7 @@ func TestLiveEtcdCompactionAndMaintenance(t *testing.T) { var firstSize int64 for cycle := range 5 { first, err := c.first.Put(ctx, "/maintenance-test/data", value) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) for range 63 { if _, err := c.first.Put(ctx, "/maintenance-test/data", value); err != nil { t.Fatal(err) @@ -151,15 +143,11 @@ func TestLiveEtcdCompactionAndMaintenance(t *testing.T) { if errors.Is(err, rpctypes.ErrCompacted) { break } - if err != nil || time.Now().After(deadline) { - t.Fatalf("history was not compacted: %v", err) - } + ck.False(err != nil || time.Now().After(deadline), "history was not compacted: %v", err) time.Sleep(200 * time.Millisecond) } statuses, err := healthyMembers(ctx, c, endpoints) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) s := statuses[endpoints[0]] if cycle == 0 { firstSize = s.DbSize @@ -187,14 +175,12 @@ func TestLiveEtcdCompactionAndMaintenance(t *testing.T) { endpoints, topoclient.NewDefaultTopoConfig(), ) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) defer func() { _ = store.Close() }() owner := &multigresv1alpha1.MultigresCluster{} register := func() { t.Helper() - if err := topo.RegisterDatabaseFromSpec( + ck.NoError(topo.RegisterDatabaseFromSpec( ctx, store, record.NewFakeRecorder(10), @@ -203,15 +189,11 @@ func TestLiveEtcdCompactionAndMaintenance(t *testing.T) { []string{"cell1"}, nil, "", - ); err != nil { - t.Fatal(err) - } + )) } register() before, err := c.first.Get(ctx, "/operator-test", clientv3.WithPrefix()) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) for range 5 { register() if _, err := healthyMembers(ctx, c, endpoints); err != nil { @@ -219,53 +201,38 @@ func TestLiveEtcdCompactionAndMaintenance(t *testing.T) { } } after, err := c.first.Get(ctx, "/operator-test", clientv3.WithPrefix()) - if err != nil { - t.Fatal(err) - } - if before.Header.Revision != after.Header.Revision { - t.Fatalf( - "no-op reconciles advanced etcd revision: %d -> %d", - before.Header.Revision, - after.Header.Revision, - ) - } + ck.NoError(err) + ck.Eq(after.Header.Revision, before.Header.Revision, "no-op reconciles advanced etcd revision") // Reclaim each member independently, transferring leadership first where // needed and requiring healthy, linearizable reads between every operation. for _, ep := range endpoints { statuses, err := healthyMembers(ctx, c, endpoints) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) s := statuses[ep] if s.Header.MemberId == s.Leader { for _, other := range endpoints { if other != ep { - if err := c.MoveLeader(ctx, ep, statuses[other].Header.MemberId); err != nil { - t.Fatal(err) - } + ck.NoError(c.MoveLeader(ctx, ep, statuses[other].Header.MemberId)) break } } } - if err := c.Defragment(ctx, ep); err != nil { - t.Fatal(err) - } + ck.NoError(c.Defragment(ctx, ep)) statuses, err = healthyMembers(ctx, c, endpoints) - if err != nil { - t.Fatal(err) - } - if statuses[ep].DbSize >= s.DbSize { - t.Fatalf( - "defragmentation did not shrink %s: %d -> %d", - ep, - s.DbSize, - statuses[ep].DbSize, - ) - } + ck.NoError(err) + ck.Less( + s.DbSize, + statuses[ep].DbSize, + "defragmentation did not shrink %s: %d ->", + ep, + s.DbSize, + ) } got, err := c.first.Get(ctx, "/maintenance-test/data") - if err != nil || len(got.Kvs) != 1 || string(got.Kvs[0].Value) != value { - t.Fatalf("live data lost after maintenance: %v", err) - } + ck.False( + err != nil || len(got.Kvs) != 1 || string(got.Kvs[0].Value) != value, + "live data lost after maintenance: %v", + err, + ) } diff --git a/pkg/resource-handler/controller/toposerver/maintenance_reconcile_test.go b/pkg/resource-handler/controller/toposerver/maintenance_reconcile_test.go index 424b4fc5..8e007916 100644 --- a/pkg/resource-handler/controller/toposerver/maintenance_reconcile_test.go +++ b/pkg/resource-handler/controller/toposerver/maintenance_reconcile_test.go @@ -7,7 +7,6 @@ import ( "testing" "time" - "github.com/google/go-cmp/cmp" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" policyv1 "k8s.io/api/policy/v1" @@ -18,6 +17,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/certs" + + "github.com/multigres/testkit/assert" ) func TestReconcileRepairsMaintenanceDependencies(t *testing.T) { @@ -34,43 +35,30 @@ func TestReconcileRepairsMaintenanceDependencies(t *testing.T) { {name: "member recovered", age: 3 * time.Minute}, } { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) r, ts, etcd := maintenanceFixture(t) - if err := policyv1.AddToScheme(r.Scheme); err != nil { - t.Fatal(err) - } + c.Require().NoError(policyv1.AddToScheme(r.Scheme)) key := client.ObjectKeyFromObject(ts) before := &appsv1.StatefulSet{} - if err := r.Get(t.Context(), key, before); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Get(t.Context(), key, before)) if tc.tls { ts.Spec.TLS = &multigresv1alpha1.TopoTLSConfig{ Enabled: ptr.To(true), IssuerName: "topology-issuer", } - if err := r.Update(t.Context(), ts); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Update(t.Context(), ts)) desired, err := BuildStatefulSet(ts, r.Scheme) - if err != nil { - t.Fatal(err) - } + c.Require().NoError(err) before.Spec = desired.Spec - if err := r.Update(t.Context(), before); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Update(t.Context(), before)) } ts.Status.EtcdMaintenance = &multigresv1alpha1.EtcdMaintenanceStatus{ LastAttemptTime: metav1.NewTime(time.Now().Add(-tc.age)), Endpoint: maintenanceEndpoints(ts)[0], InProgress: true, } - if err := r.Status().Update(t.Context(), ts); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Status().Update(t.Context(), ts)) ts.Spec.Etcd.Image = "etcd:pending-rollout" - if err := r.Update(t.Context(), ts); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Update(t.Context(), ts)) healthErr := errors.New("member is still unavailable") if tc.unhealthy { @@ -92,15 +80,16 @@ func TestReconcileRepairsMaintenanceDependencies(t *testing.T) { result, err := r.Reconcile(t.Context(), ctrl.Request{NamespacedName: key}) if tc.unhealthy { - if !errors.Is(err, healthErr) { - t.Errorf("expected member health error, got %v", err) - } + c.ErrorIs(err, healthErr, "expected member health error, got") } else if err != nil { t.Errorf("Reconcile() error = %v", err) } - if tc.wantActive && result.RequeueAfter != statusRecheckDelay { - t.Errorf("requeue = %v, want %v", result.RequeueAfter, statusRecheckDelay) - } + c.False( + tc.wantActive && result.RequeueAfter != statusRecheckDelay, + "requeue = %v, want %v", + result.RequeueAfter, + statusRecheckDelay, + ) for _, name := range []string{ts.Name + "-headless", ts.Name} { svc := &corev1.Service{} if err := r.Get( @@ -111,13 +100,13 @@ func TestReconcileRepairsMaintenanceDependencies(t *testing.T) { t.Errorf("maintenance blocked Service repair: %v", err) continue } - if !metav1.IsControlledBy(svc, ts) { - t.Errorf("Service %s is not owned by the TopoServer", name) - } - if name == ts.Name+"-headless" && - (svc.Spec.ClusterIP != corev1.ClusterIPNone || !svc.Spec.PublishNotReadyAddresses) { - t.Error("headless Service does not publish member DNS during recovery") - } + c.True( + metav1.IsControlledBy(svc, ts), + "Service %s is not owned by the TopoServer", + name, + ) + c.False(name == ts.Name+"-headless" && + (svc.Spec.ClusterIP != corev1.ClusterIPNone || !svc.Spec.PublishNotReadyAddresses), "headless Service does not publish member DNS during recovery") } if tc.tls { cert, err := certs.Get( @@ -126,19 +115,16 @@ func TestReconcileRepairsMaintenanceDependencies(t *testing.T) { ts.Namespace, multigresv1alpha1.TopoServerCertName(ts.Name), ) - if err != nil || cert == nil { - t.Errorf( - "maintenance blocked Certificate repair: certificate=%v error=%v", - cert, - err, - ) - } + c.False( + err != nil || cert == nil, + "maintenance blocked Certificate repair: certificate=%v error=%v", + cert, + err, + ) } fresh := &multigresv1alpha1.TopoServer{} - if err := r.Get(t.Context(), key, fresh); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Get(t.Context(), key, fresh)) if fresh.Status.EtcdMaintenance == nil || fresh.Status.EtcdMaintenance.InProgress != tc.wantActive { t.Errorf( @@ -148,22 +134,19 @@ func TestReconcileRepairsMaintenanceDependencies(t *testing.T) { ) } after := &appsv1.StatefulSet{} - if err := r.Get(t.Context(), key, after); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Get(t.Context(), key, after)) if tc.wantActive { - if diff := cmp.Diff(before.Spec, after.Spec); diff != "" { - t.Errorf("StatefulSet changed during maintenance (-before +after):\n%s", diff) - } + c.EqDiff(before.Spec, after.Spec, "StatefulSet changed during maintenance") } else if after.Spec.Template.Spec.Containers[0].Image != string(ts.Spec.Etcd.Image) { t.Error("StatefulSet update did not resume after maintenance recovery") } - if (connectionAttempts > 0) != (tc.age >= maintenanceTimeout) { - t.Errorf("maintenance connection attempts = %d", connectionAttempts) - } - if len(etcd.defragged) != 0 { - t.Error("recovery started another defragmentation") - } + c.Eq( + (tc.age >= maintenanceTimeout), + (connectionAttempts > 0), + "maintenance connection attempts = %d", + connectionAttempts, + ) + c.Empty(etcd.defragged, "recovery started another defragmentation") }) } } diff --git a/pkg/resource-handler/controller/toposerver/maintenance_status_integration_test.go b/pkg/resource-handler/controller/toposerver/maintenance_status_integration_test.go index d959add0..cb317a34 100644 --- a/pkg/resource-handler/controller/toposerver/maintenance_status_integration_test.go +++ b/pkg/resource-handler/controller/toposerver/maintenance_status_integration_test.go @@ -13,9 +13,12 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/client-go/tools/record" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) func TestMaintenanceReservationSurvivesStatusApply(t *testing.T) { + ck := assert.NewAborting(t) scheme := certScheme() _ = appsv1.AddToScheme(scheme) cfg := testutil.SetUpEnvtest( @@ -23,22 +26,14 @@ func TestMaintenanceReservationSurvivesStatusApply(t *testing.T) { testutil.WithCRDPaths(filepath.Join("../../../..", "config", "crd", "bases")), ) c, err := client.New(cfg, client.Options{Scheme: scheme}) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) ts := certTestTopoServer(nil) ts.Namespace = "default" ts.UID = "" - if err := c.Create(t.Context(), ts); err != nil { - t.Fatal(err) - } + ck.NoError(c.Create(t.Context(), ts)) sts, err := BuildStatefulSet(ts, scheme) - if err != nil { - t.Fatal(err) - } - if err := c.Create(t.Context(), sts); err != nil { - t.Fatal(err) - } + ck.NoError(err) + ck.NoError(c.Create(t.Context(), sts)) r := &TopoServerReconciler{ Client: c, APIReader: c, @@ -51,22 +46,17 @@ func TestMaintenanceReservationSurvivesStatusApply(t *testing.T) { Endpoint: maintenanceEndpoints(ts)[0], InProgress: true, } - if err := r.saveMaintenance(t.Context(), ts, state); err != nil { - t.Fatal(err) - } + ck.NoError(r.saveMaintenance(t.Context(), ts, state)) if err := r.saveMaintenance(t.Context(), stale, state); !apierrors.IsConflict(err) { t.Fatalf("stale reservation should conflict, got %v", err) } // The ordinary status writer intentionally omits maintenance fields; SSA // must preserve the independently owned reservation, including after restart. - if err := r.updateStatus(t.Context(), stale); err != nil { - t.Fatal(err) - } + ck.NoError(r.updateStatus(t.Context(), stale)) fresh := &multigresv1alpha1.TopoServer{} - if err := c.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh); err != nil { - t.Fatal(err) - } - if fresh.Status.EtcdMaintenance == nil || !fresh.Status.EtcdMaintenance.InProgress { - t.Fatal("status apply removed active maintenance reservation") - } + ck.NoError(c.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh)) + ck.False( + fresh.Status.EtcdMaintenance == nil || !fresh.Status.EtcdMaintenance.InProgress, + "status apply removed active maintenance reservation", + ) } diff --git a/pkg/resource-handler/controller/toposerver/maintenance_test.go b/pkg/resource-handler/controller/toposerver/maintenance_test.go index 83123b0e..aef58fd1 100644 --- a/pkg/resource-handler/controller/toposerver/maintenance_test.go +++ b/pkg/resource-handler/controller/toposerver/maintenance_test.go @@ -22,6 +22,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/interceptor" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) type fakeEtcdMaintenance struct { @@ -74,6 +76,7 @@ func maintenanceFixture( t *testing.T, ) (*TopoServerReconciler, *multigresv1alpha1.TopoServer, *fakeEtcdMaintenance) { t.Helper() + ck := assert.NewAborting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -84,9 +87,7 @@ func maintenanceFixture( DefragmentationEnabled: ptr.To(true), } sts, err := BuildStatefulSet(ts, scheme) - if err != nil { - t.Fatal(err) - } + ck.NoError(err) sts.UID, sts.Generation = "sts-uid", 1 sts.Status = appsv1.StatefulSetStatus{ ObservedGeneration: 1, @@ -133,9 +134,7 @@ func maintenanceFixture( WithStatusSubresource(&multigresv1alpha1.TopoServer{}, &appsv1.StatefulSet{}, &corev1.Pod{}). WithObjects(objects...). Build() - if err := c.Get(t.Context(), client.ObjectKeyFromObject(ts), ts); err != nil { - t.Fatal(err) - } + ck.NoError(c.Get(t.Context(), client.ObjectKeyFromObject(ts), ts)) r := &TopoServerReconciler{ Client: c, APIReader: c, @@ -147,28 +146,22 @@ func maintenanceFixture( } func TestEtcdMaintenanceSerializesMembersAndRestarts(t *testing.T) { + c := assert.NewAborting(t) r, ts, f := maintenanceFixture(t) - if err := r.reconcileMaintenance(t.Context(), ts); err != nil { - t.Fatal(err) - } + c.NoError(r.reconcileMaintenance(t.Context(), ts)) if len(f.defragged) != 1 || len(f.moved) != 1 { t.Fatalf("defrags=%v leader transfers=%v", f.defragged, f.moved) } fresh := &multigresv1alpha1.TopoServer{} - if err := r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh); err != nil { - t.Fatal(err) - } - if fresh.Status.EtcdMaintenance == nil || fresh.Status.EtcdMaintenance.InProgress { - t.Fatal("completed reservation not persisted") - } + c.NoError(r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh)) + c.False( + fresh.Status.EtcdMaintenance == nil || fresh.Status.EtcdMaintenance.InProgress, + "completed reservation not persisted", + ) // A new controller instance, or a stale reconcile, must obey the persisted interval. r2 := *r - if err := r2.reconcileMaintenance(t.Context(), fresh); err != nil { - t.Fatal(err) - } - if len(f.defragged) != 1 { - t.Fatal("maintenance repeated inside the interval") - } + c.NoError(r2.reconcileMaintenance(t.Context(), fresh)) + c.Len(f.defragged, 1, "maintenance repeated inside the interval") } func TestEtcdMaintenanceHealthGates(t *testing.T) { @@ -219,9 +212,7 @@ func TestEtcdMaintenanceHealthGates(t *testing.T) { sts := &appsv1.StatefulSet{} _ = r.Get(t.Context(), client.ObjectKeyFromObject(ts), sts) sts.Status.UpdateRevision = "rev2" - if err := r.Status().Update(t.Context(), sts); err != nil { - t.Fatal(err) - } + assert.NewAborting(t).NoError(r.Status().Update(t.Context(), sts)) }, false}, {"reservation conflict", func(r *TopoServerReconciler, _ *multigresv1alpha1.TopoServer, _ *fakeEtcdMaintenance) { r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{SubResourcePatch: func(context.Context, client.Client, string, client.Object, client.Patch, ...client.SubResourcePatchOption) error { @@ -230,15 +221,12 @@ func TestEtcdMaintenanceHealthGates(t *testing.T) { }, true}, } { t.Run(test.name, func(t *testing.T) { + c := assert.NewAborting(t) r, ts, f := maintenanceFixture(t) test.mutate(r, ts, f) err := r.reconcileMaintenance(t.Context(), ts) - if (err != nil) != test.wantErr { - t.Fatalf("error=%v", err) - } - if len(f.defragged) != 0 { - t.Fatalf("unsafe defragmentation: %v", f.defragged) - } + c.ErrorWhen(test.wantErr, err, "error=") + c.Empty(f.defragged, "unsafe defragmentation") }) } } @@ -279,15 +267,18 @@ func TestEtcdMaintenanceRejectsInvalidMemberResponses(t *testing.T) { }, } { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) r, ts, f := maintenanceFixture(t) tc.mutate(f) err := r.reconcileMaintenance(t.Context(), ts) - if err == nil || !strings.Contains(err.Error(), tc.wantError) { - t.Fatalf("expected %q, got %v", tc.wantError, err) - } - if tc.wantCause != nil && !errors.Is(err, tc.wantCause) { - t.Errorf("error %v does not wrap %v", err, tc.wantCause) - } + c.Require(). + False(err == nil || !strings.Contains(err.Error(), tc.wantError), "expected %q, got %v", tc.wantError, err) + c.False( + tc.wantCause != nil && !errors.Is(err, tc.wantCause), + "error %v does not wrap %v", + err, + tc.wantCause, + ) if len(f.moved) != 0 || len(f.defragged) != 0 { t.Errorf( "maintenance changed unhealthy members: transfers=%v defrags=%v", @@ -296,70 +287,51 @@ func TestEtcdMaintenanceRejectsInvalidMemberResponses(t *testing.T) { ) } fresh := &multigresv1alpha1.TopoServer{} - if err := r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh); err != nil { - t.Fatal(err) - } - if fresh.Status.EtcdMaintenance != nil { - t.Error("failed health checks created a maintenance reservation") - } + c.Require().NoError(r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh)) + c.Nil( + fresh.Status.EtcdMaintenance, + "failed health checks created a maintenance reservation", + ) }) } } func TestEtcdMaintenanceIncompleteLeadershipTransfer(t *testing.T) { + c := assert.NewCollecting(t) r, ts, f := maintenanceFixture(t) f.skipLeaderTransfer = true err := r.reconcileMaintenance(t.Context(), ts) - if err == nil || err.Error() != "etcd leadership transfer has not completed" { - t.Fatalf("expected incomplete leadership transfer, got %v", err) - } + c.Require(). + False(err == nil || err.Error() != "etcd leadership transfer has not completed", "expected incomplete leadership transfer, got %v", err) if len(f.moved) != 1 || f.moved[0] != 2 { t.Errorf("leadership transfer attempts = %v, want [2]", f.moved) } - if len(f.defragged) != 0 { - t.Errorf( - "defragmented a member whose leadership transfer did not complete: %v", - f.defragged, - ) - } + c.Empty(f.defragged, "defragmented a member whose leadership transfer did not complete") fresh := &multigresv1alpha1.TopoServer{} - if err := r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh); err != nil { - t.Fatal(err) - } + c.Require().NoError(r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh)) state := fresh.Status.EtcdMaintenance - if state == nil || !state.InProgress || state.Endpoint != maintenanceEndpoints(ts)[0] { - t.Fatalf("incomplete transfer did not retain the target reservation: %+v", state) - } + c.Require(). + False(state == nil || !state.InProgress || state.Endpoint != maintenanceEndpoints(ts)[0], "incomplete transfer did not retain the target reservation: %+v", state) r2 := *r - if err := r2.reconcileMaintenance(t.Context(), fresh); err != nil { - t.Fatal(err) - } - if len(f.moved) != 1 || len(f.defragged) != 0 { - t.Error("maintenance restarted while the transfer reservation remained active") - } + c.Require().NoError(r2.reconcileMaintenance(t.Context(), fresh)) + c.False( + len(f.moved) != 1 || len(f.defragged) != 0, + "maintenance restarted while the transfer reservation remained active", + ) } func TestEtcdMaintenanceInterruptedOperation(t *testing.T) { + c := assert.NewAborting(t) r, ts, f := maintenanceFixture(t) f.defragErr = context.DeadlineExceeded - if err := r.reconcileMaintenance(t.Context(), ts); err == nil { - t.Fatal("expected timeout") - } + c.Error(r.reconcileMaintenance(t.Context(), ts), "expected timeout") fresh := &multigresv1alpha1.TopoServer{} - if err := r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh); err != nil { - t.Fatal(err) - } - if !fresh.Status.EtcdMaintenance.InProgress { - t.Fatal("uncertain operation released its reservation") - } + c.NoError(r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh)) + c.True(fresh.Status.EtcdMaintenance.InProgress, "uncertain operation released its reservation") waiting, err := r.resumeMaintenance(t.Context(), fresh) - if err != nil || !waiting { - t.Fatalf("resume immediately: waiting=%v err=%v", waiting, err) - } + c.False(err != nil || !waiting, "resume immediately: waiting=%v err=%v", waiting, err) fresh.Status.EtcdMaintenance.LastAttemptTime = metav1.NewTime(time.Now().Add(-3 * time.Minute)) - if err := r.Status().Update(t.Context(), fresh); err != nil { - t.Fatal(err) - } + c.NoError(r.Status().Update(t.Context(), fresh)) f.healthErr = errors.New("previous member still unavailable") if waiting, err = r.resumeMaintenance(t.Context(), fresh); !waiting || err == nil { t.Fatalf("unhealthy resume: waiting=%v err=%v", waiting, err) @@ -368,25 +340,16 @@ func TestEtcdMaintenanceInterruptedOperation(t *testing.T) { if waiting, err = r.resumeMaintenance(t.Context(), fresh); waiting || err != nil { t.Fatalf("healthy resume: waiting=%v err=%v", waiting, err) } - if err := r.reconcileMaintenance(t.Context(), fresh); err != nil { - t.Fatal(err) - } - if len(f.defragged) != 1 { - t.Fatal("interrupted operation started another defrag") - } + c.NoError(r.reconcileMaintenance(t.Context(), fresh)) + c.Len(f.defragged, 1, "interrupted operation started another defrag") } func TestEtcdMaintenancePostHealthFailureKeepsReservation(t *testing.T) { + c := assert.NewAborting(t) r, ts, f := maintenanceFixture(t) f.afterDefrag = func() { f.healthErr = errors.New("member unhealthy after defrag") } - if err := r.reconcileMaintenance(t.Context(), ts); err == nil { - t.Fatal("expected post-defrag health failure") - } + c.Error(r.reconcileMaintenance(t.Context(), ts), "expected post-defrag health failure") fresh := &multigresv1alpha1.TopoServer{} - if err := r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh); err != nil { - t.Fatal(err) - } - if !fresh.Status.EtcdMaintenance.InProgress { - t.Fatal("failed post-check released reservation") - } + c.NoError(r.Get(t.Context(), client.ObjectKeyFromObject(ts), fresh)) + c.True(fresh.Status.EtcdMaintenance.InProgress, "failed post-check released reservation") } diff --git a/pkg/resource-handler/controller/toposerver/pdb_test.go b/pkg/resource-handler/controller/toposerver/pdb_test.go index 8db580df..770aeb86 100644 --- a/pkg/resource-handler/controller/toposerver/pdb_test.go +++ b/pkg/resource-handler/controller/toposerver/pdb_test.go @@ -3,7 +3,6 @@ package toposerver import ( "testing" - "github.com/google/go-cmp/cmp" policyv1 "k8s.io/api/policy/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" @@ -12,9 +11,12 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestBuildPodDisruptionBudget(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -28,9 +30,7 @@ func TestBuildPodDisruptionBudget(t *testing.T) { } got, err := BuildPodDisruptionBudget(toposerver, scheme) - if err != nil { - t.Fatalf("BuildPodDisruptionBudget() error = %v", err) - } + c.Require().NoError(err, "BuildPodDisruptionBudget() error =") labels := metadata.BuildStandardLabels("test-cluster", ComponentName) metadata.AddClusterLabel(labels, "test-cluster") @@ -58,9 +58,7 @@ func TestBuildPodDisruptionBudget(t *testing.T) { }, } - if diff := cmp.Diff(want, got); diff != "" { - t.Errorf("BuildPodDisruptionBudget() mismatch (-want +got):\n%s", diff) - } + c.EqDiff(want, got, "BuildPodDisruptionBudget() mismatch") } func TestBuildPodDisruptionBudgetInvalidScheme(t *testing.T) { @@ -68,7 +66,6 @@ func TestBuildPodDisruptionBudgetInvalidScheme(t *testing.T) { &multigresv1alpha1.TopoServer{ObjectMeta: metav1.ObjectMeta{Name: "test"}}, runtime.NewScheme(), ) - if err == nil { - t.Fatal("BuildPodDisruptionBudget() should fail with an unregistered scheme") - } + assert.NewAborting(t). + Error(err, "BuildPodDisruptionBudget() should fail with an unregistered scheme") } diff --git a/pkg/resource-handler/controller/toposerver/ports_test.go b/pkg/resource-handler/controller/toposerver/ports_test.go index 5b186335..934e5aa5 100644 --- a/pkg/resource-handler/controller/toposerver/ports_test.go +++ b/pkg/resource-handler/controller/toposerver/ports_test.go @@ -3,12 +3,13 @@ package toposerver import ( "testing" - "github.com/google/go-cmp/cmp" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/util/intstr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildContainerPorts(t *testing.T) { @@ -47,9 +48,7 @@ func TestBuildContainerPorts(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := buildContainerPorts(tc.toposerver) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("buildContainerPorts() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "buildContainerPorts() mismatch") }) } } @@ -87,9 +86,7 @@ func TestBuildHeadlessServicePorts(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := buildHeadlessServicePorts(tc.toposerver) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("buildHeadlessServicePorts() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "buildHeadlessServicePorts() mismatch") }) } } @@ -121,9 +118,7 @@ func TestBuildClientServicePorts(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := buildClientServicePorts(tc.toposerver) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("buildClientServicePorts() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "buildClientServicePorts() mismatch") }) } } diff --git a/pkg/resource-handler/controller/toposerver/service_test.go b/pkg/resource-handler/controller/toposerver/service_test.go index 76db5eaa..6a267609 100644 --- a/pkg/resource-handler/controller/toposerver/service_test.go +++ b/pkg/resource-handler/controller/toposerver/service_test.go @@ -3,7 +3,6 @@ package toposerver import ( "testing" - "github.com/google/go-cmp/cmp" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" @@ -11,6 +10,8 @@ import ( "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildHeadlessService(t *testing.T) { @@ -100,9 +101,7 @@ func TestBuildHeadlessService(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildHeadlessService() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildHeadlessService() mismatch") }) } } @@ -187,9 +186,7 @@ func TestBuildClientService(t *testing.T) { return } - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildClientService() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildClientService() mismatch") }) } } diff --git a/pkg/resource-handler/controller/toposerver/statefulset_test.go b/pkg/resource-handler/controller/toposerver/statefulset_test.go index 2b084b7b..b3e42471 100644 --- a/pkg/resource-handler/controller/toposerver/statefulset_test.go +++ b/pkg/resource-handler/controller/toposerver/statefulset_test.go @@ -3,7 +3,6 @@ package toposerver import ( "testing" - "github.com/google/go-cmp/cmp" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/resource" @@ -13,6 +12,8 @@ import ( "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildStatefulSet(t *testing.T) { @@ -647,14 +648,13 @@ func TestBuildStatefulSet(t *testing.T) { corev1.EnvVar{Name: "ETCD_QUOTA_BACKEND_BYTES", Value: "2147483648"}, ) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildStatefulSet() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildStatefulSet() mismatch") }) } } func TestBuildStatefulSetPlacementControls(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) @@ -692,34 +692,23 @@ func TestBuildStatefulSetPlacementControls(t *testing.T) { } got, err := BuildStatefulSet(toposerver, scheme) - if err != nil { - t.Fatalf("BuildStatefulSet() error = %v", err) - } + c.Require().NoError(err, "BuildStatefulSet() error =") podSpec := got.Spec.Template.Spec - if diff := cmp.Diff(placement.NodeSelector, podSpec.NodeSelector); diff != "" { - t.Errorf("NodeSelector mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(placement.Affinity, podSpec.Affinity); diff != "" { - t.Errorf("Affinity mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(placement.Tolerations, podSpec.Tolerations); diff != "" { - t.Errorf("Tolerations mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff( + c.EqDiff(placement.NodeSelector, podSpec.NodeSelector, "NodeSelector mismatch") + c.EqDiff(placement.Affinity, podSpec.Affinity, "Affinity mismatch") + c.EqDiff(placement.Tolerations, podSpec.Tolerations, "Tolerations mismatch") + c.EqDiff( placement.TopologySpreadConstraints, podSpec.TopologySpreadConstraints, - ); diff != "" { - t.Errorf("TopologySpreadConstraints mismatch (-want +got):\n%s", diff) - } + "TopologySpreadConstraints mismatch", + ) // The built object must not alias the source custom resource. podSpec.NodeSelector["node-pool"] = "other" podSpec.Affinity.PodAntiAffinity.RequiredDuringSchedulingIgnoredDuringExecution[0].TopologyKey = "other" podSpec.TopologySpreadConstraints[0].TopologyKey = "other" - if placement.NodeSelector["node-pool"] != "topology" || + c.Require().False(placement.NodeSelector["node-pool"] != "topology" || placement.Affinity.PodAntiAffinity.RequiredDuringSchedulingIgnoredDuringExecution[0].TopologyKey != corev1.LabelHostname || - placement.TopologySpreadConstraints[0].TopologyKey != corev1.LabelTopologyZone { - t.Fatal("BuildStatefulSet() aliased placement fields from the TopoServer") - } + placement.TopologySpreadConstraints[0].TopologyKey != corev1.LabelTopologyZone, "BuildStatefulSet() aliased placement fields from the TopoServer") } diff --git a/pkg/resource-handler/controller/toposerver/storage_class_guard_test.go b/pkg/resource-handler/controller/toposerver/storage_class_guard_test.go index d14629c8..e2390606 100644 --- a/pkg/resource-handler/controller/toposerver/storage_class_guard_test.go +++ b/pkg/resource-handler/controller/toposerver/storage_class_guard_test.go @@ -14,6 +14,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) func TestValidateEtcdStorageClassDependency(t *testing.T) { @@ -25,6 +27,7 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { t.Run("empty storage class sets True/NotSpecified condition", func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{Name: "test-ts", Namespace: "default"}, Spec: multigresv1alpha1.TopoServerSpec{}, @@ -36,23 +39,21 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { Build() r := &TopoServerReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.validateEtcdStorageClassDependency(t.Context(), ts); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.NoError(r.validateEtcdStorageClassDependency(t.Context(), ts), "unexpected error") var updated multigresv1alpha1.TopoServer - if err := c.Get(t.Context(), client.ObjectKeyFromObject(ts), &updated); err != nil { - t.Fatalf("failed to read toposerver: %v", err) - } + ck.NoError( + c.Get(t.Context(), client.ObjectKeyFromObject(ts), &updated), + "failed to read toposerver", + ) cond := findCondition(updated.Status.Conditions, conditionStorageClassValid) - if cond == nil || cond.Status != metav1.ConditionTrue || - cond.Reason != storageClassNotSpecifiedReason { - t.Fatalf("unexpected condition: %#v", cond) - } + ck.False(cond == nil || cond.Status != metav1.ConditionTrue || + cond.Reason != storageClassNotSpecifiedReason, "unexpected condition: %#v", cond) }) t.Run("ready storage class sets True/Ready condition", func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{Name: "test-ts", Namespace: "default"}, Spec: multigresv1alpha1.TopoServerSpec{ @@ -72,25 +73,23 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { Build() r := &TopoServerReconciler{Client: c, Scheme: scheme, Recorder: record.NewFakeRecorder(10)} - if err := r.validateEtcdStorageClassDependency(t.Context(), ts); err != nil { - t.Fatalf("unexpected error: %v", err) - } + ck.NoError(r.validateEtcdStorageClassDependency(t.Context(), ts), "unexpected error") var updated multigresv1alpha1.TopoServer - if err := c.Get(t.Context(), client.ObjectKeyFromObject(ts), &updated); err != nil { - t.Fatalf("failed to read toposerver: %v", err) - } + ck.NoError( + c.Get(t.Context(), client.ObjectKeyFromObject(ts), &updated), + "failed to read toposerver", + ) cond := findCondition(updated.Status.Conditions, conditionStorageClassValid) - if cond == nil || cond.Status != metav1.ConditionTrue || - cond.Reason != storageClassReadyReason { - t.Fatalf("unexpected condition: %#v", cond) - } + ck.False(cond == nil || cond.Status != metav1.ConditionTrue || + cond.Reason != storageClassReadyReason, "unexpected condition: %#v", cond) }) t.Run( "immediate binding mode sets False condition and returns dependency error", func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{Name: "test-ts", Namespace: "default"}, Spec: multigresv1alpha1.TopoServerSpec{ @@ -115,23 +114,21 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { } err := r.validateEtcdStorageClassDependency(t.Context(), ts) - if err == nil || !isStorageClassDependencyError(err) { - t.Fatalf("expected StorageClass dependency error, got: %v", err) - } + ck.False( + err == nil || !isStorageClassDependencyError(err), + "expected StorageClass dependency error, got: %v", + err, + ) var updated multigresv1alpha1.TopoServer - if getErr := c.Get( + ck.NoError(c.Get( t.Context(), client.ObjectKeyFromObject(ts), &updated, - ); getErr != nil { - t.Fatalf("failed to read toposerver: %v", getErr) - } + ), "failed to read toposerver") cond := findCondition(updated.Status.Conditions, conditionStorageClassValid) - if cond == nil || cond.Status != metav1.ConditionFalse || - cond.Reason != storageClassBindingModeReason { - t.Fatalf("unexpected condition: %#v", cond) - } + ck.False(cond == nil || cond.Status != metav1.ConditionFalse || + cond.Reason != storageClassBindingModeReason, "unexpected condition: %#v", cond) }, ) @@ -139,6 +136,7 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { "missing storage class sets False/NotFound condition and returns dependency error", func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{Name: "test-ts", Namespace: "default"}, Spec: multigresv1alpha1.TopoServerSpec{ @@ -159,28 +157,27 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { } err := r.validateEtcdStorageClassDependency(t.Context(), ts) - if err == nil || !isMissingStorageClassDependency(err) { - t.Fatalf("expected missing dependency error, got: %v", err) - } + ck.False( + err == nil || !isMissingStorageClassDependency(err), + "expected missing dependency error, got: %v", + err, + ) var updated multigresv1alpha1.TopoServer - if getErr := c.Get( + ck.NoError(c.Get( t.Context(), client.ObjectKeyFromObject(ts), &updated, - ); getErr != nil { - t.Fatalf("failed to read toposerver: %v", getErr) - } + ), "failed to read toposerver") cond := findCondition(updated.Status.Conditions, conditionStorageClassValid) - if cond == nil || cond.Status != metav1.ConditionFalse || - cond.Reason != storageClassNotFoundReason { - t.Fatalf("unexpected condition: %#v", cond) - } + ck.False(cond == nil || cond.Status != metav1.ConditionFalse || + cond.Reason != storageClassNotFoundReason, "unexpected condition: %#v", cond) }, ) t.Run("API error propagates without setting condition", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{Name: "test-ts", Namespace: "default"}, Spec: multigresv1alpha1.TopoServerSpec{ @@ -204,38 +201,37 @@ func TestValidateEtcdStorageClassDependency(t *testing.T) { } err := r.validateEtcdStorageClassDependency(t.Context(), ts) - if err == nil { - t.Fatal("expected error, got nil") - } - if isMissingStorageClassDependency(err) { - t.Fatal("expected non-dependency error, got dependency error") - } + c.Error(err, "expected error, got nil") + c.False( + isMissingStorageClassDependency(err), + "expected non-dependency error, got dependency error", + ) var updated multigresv1alpha1.TopoServer - if getErr := baseClient.Get( + c.NoError(baseClient.Get( t.Context(), client.ObjectKeyFromObject(ts), &updated, - ); getErr != nil { - t.Fatalf("failed to read toposerver: %v", getErr) - } + ), "failed to read toposerver") cond := findCondition(updated.Status.Conditions, conditionStorageClassValid) - if cond != nil && cond.Status == metav1.ConditionFalse { - t.Fatalf("condition should not be False on API error, got: %#v", cond) - } + c.False( + cond != nil && cond.Status == metav1.ConditionFalse, + "condition should not be False on API error, got: %#v", + cond, + ) }) } func TestIsMissingStorageClassDependencyWrapped(t *testing.T) { + c := assert.NewAborting(t) err := errors.New("other") - if isMissingStorageClassDependency(err) { - t.Fatal("expected false for non-dependency error") - } + c.False(isMissingStorageClassDependency(err), "expected false for non-dependency error") wrapped := errors.Join(errors.New("outer"), &missingStorageClassDependencyError{className: "x"}) - if !isMissingStorageClassDependency(wrapped) { - t.Fatal("expected true for wrapped missing dependency error") - } + c.True( + isMissingStorageClassDependency(wrapped), + "expected true for wrapped missing dependency error", + ) } // findCondition returns the condition with the given type, or nil if not found. diff --git a/pkg/resource-handler/controller/toposerver/toposerver_controller_internal_test.go b/pkg/resource-handler/controller/toposerver/toposerver_controller_internal_test.go index c5d0a41c..1e553f1d 100644 --- a/pkg/resource-handler/controller/toposerver/toposerver_controller_internal_test.go +++ b/pkg/resource-handler/controller/toposerver/toposerver_controller_internal_test.go @@ -20,6 +20,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) // TestReconcileStatefulSet_InvalidScheme tests the error path when BuildStatefulSet fails. @@ -48,9 +50,7 @@ func TestReconcileStatefulSet_InvalidScheme(t *testing.T) { } err := reconciler.reconcileStatefulSet(context.Background(), toposerver) - if err == nil { - t.Error("reconcileStatefulSet() should error with invalid scheme") - } + assert.NewCollecting(t).Error(err, "reconcileStatefulSet() should error with invalid scheme") } // TestReconcileHeadlessService_InvalidScheme tests the error path when BuildHeadlessService fails. @@ -76,9 +76,8 @@ func TestReconcileHeadlessService_InvalidScheme(t *testing.T) { } err := reconciler.reconcileHeadlessService(context.Background(), toposerver) - if err == nil { - t.Error("reconcileHeadlessService() should error with invalid scheme") - } + assert.NewCollecting(t). + Error(err, "reconcileHeadlessService() should error with invalid scheme") } // TestReconcileClientService_InvalidScheme tests the error path when BuildClientService fails. @@ -104,9 +103,7 @@ func TestReconcileClientService_InvalidScheme(t *testing.T) { } err := reconciler.reconcileClientService(context.Background(), toposerver) - if err == nil { - t.Error("reconcileClientService() should error with invalid scheme") - } + assert.NewCollecting(t).Error(err, "reconcileClientService() should error with invalid scheme") } // TestUpdateStatus_StatefulSetNotFound tests the NotFound path in updateStatus. @@ -137,9 +134,8 @@ func TestUpdateStatus_StatefulSetNotFound(t *testing.T) { // Call updateStatus when StatefulSet doesn't exist yet err := reconciler.updateStatus(context.Background(), toposerver) - if err != nil { - t.Errorf("updateStatus() should not error when StatefulSet not found, got: %v", err) - } + assert.NewCollecting(t). + NoError(err, "updateStatus() should not error when StatefulSet not found, got") } // TestReconcileClientService_PatchError tests error path on Patch client Service. @@ -179,9 +175,7 @@ func TestReconcileClientService_PatchError(t *testing.T) { } err := reconciler.reconcileClientService(context.Background(), toposerver) - if err == nil { - t.Error("reconcileClientService() should error on Patch failure") - } + assert.NewCollecting(t).Error(err, "reconcileClientService() should error on Patch failure") } // TestUpdateStatus_GetError tests error path on Get StatefulSet (not NotFound). @@ -215,9 +209,7 @@ func TestUpdateStatus_GetError(t *testing.T) { } err := reconciler.updateStatus(context.Background(), toposerver) - if err == nil { - t.Error("updateStatus() should error on Get failure") - } + assert.NewCollecting(t).Error(err, "updateStatus() should error on Get failure") } // TestSetupWithManager tests the manager setup function. @@ -235,25 +227,23 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Metrics: metricsserver.Options{BindAddress: "0"}, }) - if err != nil { - t.Fatalf("Failed to create manager: %v", err) - } + assert.NewAborting(t).NoError(err, "Failed to create manager") return mgr } t.Run("default options", func(t *testing.T) { + c := assert.NewCollecting(t) mgr := createMgr() r := &TopoServerReconciler{ Client: mgr.GetClient(), Scheme: scheme, Recorder: record.NewFakeRecorder(100), } - if err := r.SetupWithManager(mgr); err != nil { - t.Errorf("SetupWithManager() error = %v", err) - } - if r.APIReader != mgr.GetAPIReader() { - t.Error("maintenance must use the manager's uncached API reader") - } + c.NoError(r.SetupWithManager(mgr), "SetupWithManager() error =") + c.False( + r.APIReader != mgr.GetAPIReader(), + "maintenance must use the manager's uncached API reader", + ) }) t.Run("with options", func(t *testing.T) { @@ -263,12 +253,10 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Recorder: record.NewFakeRecorder(100), } - if err := r.SetupWithManager(mgr, controller.Options{ + assert.NewCollecting(t).NoError(r.SetupWithManager(mgr, controller.Options{ MaxConcurrentReconciles: 1, SkipNameValidation: ptr.To(true), - }); err != nil { - t.Errorf("SetupWithManager() with opts error = %v", err) - } + }), "SetupWithManager() with opts error =") }) for name, reader := range map[string]client.Reader{ @@ -276,6 +264,7 @@ func TestSetupWithManager(t *testing.T) { "shared setup preserves injected reader": fake.NewClientBuilder().WithScheme(scheme).Build(), } { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) mgr := createMgr() r := &TopoServerReconciler{ Client: mgr.GetClient(), @@ -283,23 +272,20 @@ func TestSetupWithManager(t *testing.T) { Scheme: scheme, Recorder: record.NewFakeRecorder(100), } - if err := r.SetupWithManagerReconciler(mgr, r, controller.Options{ + c.Require().NoError(r.SetupWithManagerReconciler(mgr, r, controller.Options{ SkipNameValidation: ptr.To(true), - }); err != nil { - t.Fatalf("SetupWithManagerReconciler() error = %v", err) - } + }), "SetupWithManagerReconciler() error =") want := reader if want == nil { want = mgr.GetAPIReader() } - if r.APIReader != want { - t.Error("shared setup selected the wrong maintenance reader") - } + c.False(r.APIReader != want, "shared setup selected the wrong maintenance reader") }) } } func TestUpdateStatus_DegradedOnCrashLoop(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -356,10 +342,7 @@ func TestUpdateStatus_DegradedOnCrashLoop(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := reconciler.updateStatus(context.Background(), toposerver); err != nil { - t.Fatalf("updateStatus() unexpected error: %v", err) - } - if toposerver.Status.Phase != multigresv1alpha1.PhaseDegraded { - t.Errorf("expected PhaseDegraded, got %q", toposerver.Status.Phase) - } + c.Require(). + NoError(reconciler.updateStatus(context.Background(), toposerver), "updateStatus() unexpected error") + c.Eq(multigresv1alpha1.PhaseDegraded, toposerver.Status.Phase, "expected PhaseDegraded, got") } diff --git a/pkg/resource-handler/controller/toposerver/toposerver_controller_test.go b/pkg/resource-handler/controller/toposerver/toposerver_controller_test.go index 47db4ece..84ab3ca2 100644 --- a/pkg/resource-handler/controller/toposerver/toposerver_controller_test.go +++ b/pkg/resource-handler/controller/toposerver/toposerver_controller_test.go @@ -20,6 +20,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) func TestTopoServerReconciler_Reconcile(t *testing.T) { @@ -51,43 +53,30 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { existingObjects: []client.Object{}, wantRequeue: true, assertFunc: func(t *testing.T, c client.Client, toposerver *multigresv1alpha1.TopoServer) { + ck := assert.NewCollecting(t) // Verify all workload, disruption, and service resources were created. sts := &appsv1.StatefulSet{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-toposerver", Namespace: "default"}, - sts); err != nil { - t.Errorf("StatefulSet should exist: %v", err) - } + sts), "StatefulSet should exist") pdb := &policyv1.PodDisruptionBudget{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-toposerver", Namespace: "default"}, - pdb); err != nil { - t.Errorf("PodDisruptionBudget should exist: %v", err) - } + pdb), "PodDisruptionBudget should exist") headlessSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-toposerver-headless", Namespace: "default"}, - headlessSvc); err != nil { - t.Errorf("Headless Service should exist: %v", err) - } + headlessSvc), "Headless Service should exist") clientSvc := &corev1.Service{} - if err := c.Get(t.Context(), + ck.NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-toposerver", Namespace: "default"}, - clientSvc); err != nil { - t.Errorf("Client Service should exist: %v", err) - } + clientSvc), "Client Service should exist") // Verify defaults - if *sts.Spec.Replicas != int32(3) { - t.Errorf( - "StatefulSet replicas = %d, want %d", - *sts.Spec.Replicas, - int32(3), - ) - } + ck.Eq(int32(3), *sts.Spec.Replicas, "StatefulSet replicas") }, }, "update existing resources": { @@ -134,25 +123,21 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { }, }, assertFunc: func(t *testing.T, c client.Client, toposerver *multigresv1alpha1.TopoServer) { + ck := assert.NewCollecting(t) sts := &appsv1.StatefulSet{} err := c.Get(t.Context(), types.NamespacedName{ Name: "existing-toposerver", Namespace: "default", }, sts) - if err != nil { - t.Fatalf("Failed to get StatefulSet: %v", err) - } + ck.Require().NoError(err, "Failed to get StatefulSet") - if *sts.Spec.Replicas != 5 { - t.Errorf("StatefulSet replicas = %d, want 5", *sts.Spec.Replicas) - } + ck.Eq(5, *sts.Spec.Replicas, "StatefulSet replicas") - if sts.Spec.Template.Spec.Containers[0].Image != "quay.io/coreos/etcd:v3.5.15" { - t.Errorf( - "StatefulSet image = %s, want quay.io/coreos/etcd:v3.5.15", - sts.Spec.Template.Spec.Containers[0].Image, - ) - } + ck.Eq( + "quay.io/coreos/etcd:v3.5.15", + sts.Spec.Template.Spec.Containers[0].Image, + "StatefulSet image", + ) }, }, @@ -180,9 +165,8 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { err := c.Get(t.Context(), types.NamespacedName{Name: "test-missing-sc", Namespace: "default"}, sts) - if err == nil { - t.Error("StatefulSet should not exist when StorageClass is missing") - } + assert.NewCollecting(t). + Error(err, "StatefulSet should not exist when StorageClass is missing") }, }, "existing StorageClass proceeds normally and creates StatefulSet": { @@ -207,11 +191,9 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { wantRequeue: true, assertFunc: func(t *testing.T, c client.Client, toposerver *multigresv1alpha1.TopoServer) { sts := &appsv1.StatefulSet{} - if err := c.Get(t.Context(), + assert.NewCollecting(t).NoError(c.Get(t.Context(), types.NamespacedName{Name: "test-existing-sc", Namespace: "default"}, - sts); err != nil { - t.Errorf("StatefulSet should exist when StorageClass is present: %v", err) - } + sts), "StatefulSet should exist when StorageClass is present") }, }, "immediate StorageClass requeues without creating StatefulSet": { @@ -237,9 +219,8 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { assertFunc: func(t *testing.T, c client.Client, toposerver *multigresv1alpha1.TopoServer) { sts := &appsv1.StatefulSet{} err := c.Get(t.Context(), client.ObjectKeyFromObject(toposerver), sts) - if err == nil { - t.Error("StatefulSet should not exist with immediate volume binding") - } + assert.NewCollecting(t). + Error(err, "StatefulSet should not exist with immediate volume binding") }, }, @@ -367,6 +348,7 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) // Fake clients register unstructured certificate types lazily. Each // parallel subtest must own its scheme to avoid concurrent mutation. @@ -408,9 +390,7 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { } if !toposerverInExisting { err := fakeClient.Create(t.Context(), tc.toposerver) - if err != nil { - t.Fatalf("Failed to create TopoServer: %v", err) - } + c.Require().NoError(err, "Failed to create TopoServer") } // Reconcile @@ -430,13 +410,12 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { return } - if (result.RequeueAfter != 0) != tc.wantRequeue { - t.Errorf( - "Reconcile() requeue = %v, want requeue = %v", - result.RequeueAfter, - tc.wantRequeue, - ) - } + c.Eq( + tc.wantRequeue, + (result.RequeueAfter != 0), + "Reconcile() requeue = %v, want requeue =", + result.RequeueAfter, + ) // Run custom assertions if provided if tc.assertFunc != nil { @@ -447,6 +426,7 @@ func TestTopoServerReconciler_Reconcile(t *testing.T) { } func TestTopoServerReconciler_ReconcileNotFound(t *testing.T) { + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = multigresv1alpha1.AddToScheme(scheme) _ = appsv1.AddToScheme(scheme) @@ -473,12 +453,8 @@ func TestTopoServerReconciler_ReconcileNotFound(t *testing.T) { } result, err := reconciler.Reconcile(t.Context(), req) - if err != nil { - t.Errorf("Reconcile() should not error on NotFound, got: %v", err) - } - if result.RequeueAfter > 0 { - t.Errorf("Reconcile() should not requeue on NotFound") - } + c.NoError(err, "Reconcile() should not error on NotFound, got") + c.LessOrEqual(0, result.RequeueAfter, "Reconcile() should not requeue on NotFound") } func TestTopoServerReconciler_UpdateStatus(t *testing.T) { @@ -490,6 +466,7 @@ func TestTopoServerReconciler_UpdateStatus(t *testing.T) { _ = storagev1.AddToScheme(scheme) t.Run("all_replicas_ready_status", func(t *testing.T) { + c := assert.NewCollecting(t) toposerver := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ Name: "test-toposerver-ready", @@ -530,29 +507,21 @@ func TestTopoServerReconciler_UpdateStatus(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.updateStatus(t.Context(), toposerver); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + c.Require().NoError(r.updateStatus(t.Context(), toposerver), "updateStatus failed") updatedTopoServer := &multigresv1alpha1.TopoServer{} - if err := fakeClient.Get( + c.Require().NoError(fakeClient.Get( t.Context(), client.ObjectKeyFromObject(toposerver), updatedTopoServer, - ); err != nil { - t.Fatalf("Failed to get TopoServer: %v", err) - } + ), "Failed to get TopoServer") if len(updatedTopoServer.Status.Conditions) == 0 { t.Error("Status.Conditions should not be empty") } else { readyCondition := updatedTopoServer.Status.Conditions[0] - if readyCondition.Type != "Ready" { - t.Errorf("Condition type = %s, want Ready", readyCondition.Type) - } - if readyCondition.Status != metav1.ConditionTrue { - t.Errorf("Condition status = %s, want True", readyCondition.Status) - } + c.Eq("Ready", readyCondition.Type, "Condition type") + c.Eq(metav1.ConditionTrue, readyCondition.Status, "Condition status") } }) @@ -582,9 +551,8 @@ func TestTopoServerReconciler_UpdateStatus(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.updateStatus(t.Context(), toposerver); err != nil { - t.Fatalf("updateStatus failed: %v", err) - } + assert.NewAborting(t). + NoError(r.updateStatus(t.Context(), toposerver), "updateStatus failed") toposerverUpdated := &multigresv1alpha1.TopoServer{} _ = fakeClient.Get(t.Context(), client.ObjectKeyFromObject(toposerver), toposerverUpdated) if len(toposerverUpdated.Status.Conditions) > 0 && @@ -609,6 +577,7 @@ func TestTopoServerReconciler_FieldOwnershipIsolation(t *testing.T) { t.Run("updateStatus patch contains exactly one condition (Ready)", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ @@ -656,29 +625,26 @@ func TestTopoServerReconciler_FieldOwnershipIsolation(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.updateStatus(t.Context(), ts); err != nil { - t.Fatalf("updateStatus: %v", err) - } + c.NoError(r.updateStatus(t.Context(), ts), "updateStatus") patchTS, ok := capturedPatchObj.(*multigresv1alpha1.TopoServer) - if !ok { - t.Fatalf("expected *TopoServer patch, got %T", capturedPatchObj) - } + c.True(ok, "expected *TopoServer patch, got %T", capturedPatchObj) // Exactly one condition: Ready. - if len(patchTS.Status.Conditions) != 1 { - t.Fatalf("updateStatus patch must contain exactly 1 condition, got %d: %v", - len(patchTS.Status.Conditions), patchTS.Status.Conditions) - } - if patchTS.Status.Conditions[0].Type != "Ready" { - t.Fatalf("expected Ready condition, got %s", patchTS.Status.Conditions[0].Type) - } + c.Len( + patchTS.Status.Conditions, + 1, + "updateStatus patch must contain exactly 1 condition, got %d", + len(patchTS.Status.Conditions), + ) + c.Eq("Ready", patchTS.Status.Conditions[0].Type, "expected Ready condition, got") }) t.Run( "guard patch contains exactly one condition (StorageClassValid) and no other status fields", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) ts := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ @@ -718,50 +684,34 @@ func TestTopoServerReconciler_FieldOwnershipIsolation(t *testing.T) { Recorder: record.NewFakeRecorder(10), } - if err := r.validateEtcdStorageClassDependency(t.Context(), ts); err != nil { - t.Fatalf("guard: %v", err) - } + c.NoError(r.validateEtcdStorageClassDependency(t.Context(), ts), "guard") patchTS, ok := capturedPatchObj.(*multigresv1alpha1.TopoServer) - if !ok { - t.Fatalf("expected *TopoServer patch, got %T", capturedPatchObj) - } + c.True(ok, "expected *TopoServer patch, got %T", capturedPatchObj) // Exactly one condition: StorageClassValid. - if len(patchTS.Status.Conditions) != 1 { - t.Fatalf("guard patch must contain exactly 1 condition, got %d: %v", - len(patchTS.Status.Conditions), patchTS.Status.Conditions) - } + c.Len( + patchTS.Status.Conditions, + 1, + "guard patch must contain exactly 1 condition, got %d", + len(patchTS.Status.Conditions), + ) scCond := &patchTS.Status.Conditions[0] - if scCond.Type != conditionStorageClassValid { - t.Fatalf("expected %s condition, got %s", conditionStorageClassValid, scCond.Type) - } + c.Eq(conditionStorageClassValid, scCond.Type, "expected") if scCond.Status != metav1.ConditionTrue || scCond.Reason != storageClassReadyReason { t.Fatalf("unexpected condition: status=%s reason=%s", scCond.Status, scCond.Reason) } // No other status fields should be set in the guard patch. - if patchTS.Status.Phase != "" { - t.Fatalf("guard patch must not set Phase, got %q", patchTS.Status.Phase) - } - if patchTS.Status.Message != "" { - t.Fatalf("guard patch must not set Message, got %q", patchTS.Status.Message) - } - if patchTS.Status.ClientService != "" { - t.Fatalf( - "guard patch must not set ClientService, got %q", - patchTS.Status.ClientService, - ) - } - if patchTS.Status.PeerService != "" { - t.Fatalf("guard patch must not set PeerService, got %q", patchTS.Status.PeerService) - } - if patchTS.Status.ObservedGeneration != 0 { - t.Fatalf( - "guard patch must not set ObservedGeneration, got %d", - patchTS.Status.ObservedGeneration, - ) - } + c.Eq("", patchTS.Status.Phase, "guard patch must not set Phase, got") + c.Eq("", patchTS.Status.Message, "guard patch must not set Message, got") + c.Eq("", patchTS.Status.ClientService, "guard patch must not set ClientService, got") + c.Eq("", patchTS.Status.PeerService, "guard patch must not set PeerService, got") + c.Eq( + 0, + patchTS.Status.ObservedGeneration, + "guard patch must not set ObservedGeneration, got", + ) }, ) } @@ -775,6 +725,7 @@ func TestTopoServerReconciler_StandaloneMocks(t *testing.T) { _ = storagev1.AddToScheme(scheme) t.Run("ignore_deleted", func(t *testing.T) { + c := assert.NewAborting(t) toposerver := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ Name: "test-deleted", @@ -784,9 +735,7 @@ func TestTopoServerReconciler_StandaloneMocks(t *testing.T) { fakeClient := fake.NewClientBuilder().WithScheme(scheme).WithObjects(toposerver).Build() // Now delete it so DeletionTimestamp is set - if err := fakeClient.Delete(t.Context(), toposerver); err != nil { - t.Fatalf("failed to delete: %v", err) - } + c.NoError(fakeClient.Delete(t.Context(), toposerver), "failed to delete") r := &TopoServerReconciler{ Client: fakeClient, @@ -799,15 +748,12 @@ func TestTopoServerReconciler_StandaloneMocks(t *testing.T) { NamespacedName: types.NamespacedName{Name: "test-deleted", Namespace: "default"}, }, ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if res.RequeueAfter != 0 { - t.Fatalf("expected no requeue") - } + c.NoError(err, "unexpected error") + c.Eq(0, res.RequeueAfter, "expected no requeue") }) t.Run("healthy_phase", func(t *testing.T) { + c := assert.NewAborting(t) toposerver := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ Name: "test-healthy", @@ -820,9 +766,7 @@ func TestTopoServerReconciler_StandaloneMocks(t *testing.T) { // Build the exact STS and Pods needed so reconcile reaches healthy phase. sts, err := BuildStatefulSet(toposerver, scheme) - if err != nil { - t.Fatalf("failed to build sts: %v", err) - } + c.NoError(err, "failed to build sts") sts.Generation = 1 sts.Status = appsv1.StatefulSetStatus{ Replicas: *sts.Spec.Replicas, @@ -867,17 +811,19 @@ func TestTopoServerReconciler_StandaloneMocks(t *testing.T) { NamespacedName: types.NamespacedName{Name: "test-healthy", Namespace: "default"}, }, ) - if reconcileErr != nil { - t.Fatalf("unexpected error: %v", reconcileErr) - } + c.NoError(reconcileErr, "unexpected error") updatedTopo := &multigresv1alpha1.TopoServer{} _ = fakeClient.Get(t.Context(), client.ObjectKeyFromObject(toposerver), updatedTopo) - if res.RequeueAfter != 0 { - t.Fatalf("expected no requeue, but got requeueAfter %v, phase is %s, msg is %s", - res.RequeueAfter, updatedTopo.Status.Phase, updatedTopo.Status.Message) - } + c.Eq( + 0, + res.RequeueAfter, + "expected no requeue, but got requeueAfter %v, phase is %s, msg is %s", + res.RequeueAfter, + updatedTopo.Status.Phase, + updatedTopo.Status.Message, + ) }) t.Run("list_error_in_updateStatus", func(t *testing.T) { @@ -914,9 +860,10 @@ func TestTopoServerReconciler_StandaloneMocks(t *testing.T) { NamespacedName: types.NamespacedName{Name: "test-list-err", Namespace: "default"}, }, ) - if reconcileErr == nil || - !strings.Contains(reconcileErr.Error(), "failed to list toposerver pods") { - t.Fatalf("expected list pods error, got %v", reconcileErr) - } + assert.NewAborting(t).False(reconcileErr == nil || + !strings.Contains( + reconcileErr.Error(), + "failed to list toposerver pods", + ), "expected list pods error, got %v", reconcileErr) }) } diff --git a/pkg/testutil/compare_internal_test.go b/pkg/testutil/compare_internal_test.go index 9df18447..4a236dd6 100644 --- a/pkg/testutil/compare_internal_test.go +++ b/pkg/testutil/compare_internal_test.go @@ -5,6 +5,8 @@ import ( "github.com/google/go-cmp/cmp" corev1 "k8s.io/api/core/v1" + + "github.com/multigres/testkit/assert" ) // TestFilterByFieldName tests the filter function directly. @@ -32,14 +34,8 @@ func TestFilterByFieldName(t *testing.T) { t.Parallel() filter := filterByFieldName(tc.fieldName) got := filter(tc.path) - if got != tc.want { - t.Errorf( - "filterByFieldName(%s) with empty path = %v, want %v", - tc.fieldName, - got, - tc.want, - ) - } + assert.NewCollecting(t). + Eq(tc.want, got, "filterByFieldName(%s) with empty path = %v, want", tc.fieldName, got) }) } } @@ -64,7 +60,6 @@ func TestFilterByFieldName_Integration(t *testing.T) { // Should match when ignoring Status diff := cmp.Diff(svc1, svc2, IgnoreStatus()) - if diff != "" { - t.Errorf("Services should match when ignoring Status, but found diff:\n%s", diff) - } + assert.NewCollecting(t). + Eq("", diff, "Services should match when ignoring Status, but found diff:\n") } diff --git a/pkg/testutil/compare_test.go b/pkg/testutil/compare_test.go index 3d046b1b..f007149a 100644 --- a/pkg/testutil/compare_test.go +++ b/pkg/testutil/compare_test.go @@ -10,6 +10,8 @@ import ( "k8s.io/apimachinery/pkg/util/intstr" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) func TestComparisonOptions(t *testing.T) { @@ -164,9 +166,8 @@ func TestComparisonOptions(t *testing.T) { t.Run(name, func(t *testing.T) { t.Parallel() diff := cmp.Diff(tc.obj1, tc.obj2, tc.options...) - if diff != "" { - t.Errorf("%s should make objects match, but found diff:\n%s", name, diff) - } + assert.NewCollecting(t). + Eq("", diff, "%s should make objects match, but found diff:\n", name) }) } } @@ -221,9 +222,8 @@ func TestIgnoreProbeDefaults(t *testing.T) { testutil.IgnoreMetaRuntimeFields(), testutil.IgnorePodSpecDefaults(), ) - if diff != "" { - t.Errorf("IgnoreProbeDefaults should ignore probe defaults, but found diff:\n%s", diff) - } + assert.NewCollecting(t). + Eq("", diff, "IgnoreProbeDefaults should ignore probe defaults, but found diff:\n") } func TestIgnorePVCRuntimeFields(t *testing.T) { @@ -259,10 +259,6 @@ func TestIgnorePVCRuntimeFields(t *testing.T) { testutil.IgnorePVCRuntimeFields(), testutil.IgnoreMetaRuntimeFields(), ) - if diff != "" { - t.Errorf( - "IgnorePVCRuntimeFields should ignore finalizers and VolumeMode, but found diff:\n%s", - diff, - ) - } + assert.NewCollecting(t). + Eq("", diff, "IgnorePVCRuntimeFields should ignore finalizers and VolumeMode, but found diff:\n") } diff --git a/pkg/testutil/e2e_test.go b/pkg/testutil/e2e_test.go index 71cd46cc..fe5d8176 100644 --- a/pkg/testutil/e2e_test.go +++ b/pkg/testutil/e2e_test.go @@ -15,6 +15,8 @@ import ( "sigs.k8s.io/e2e-framework/pkg/env" "sigs.k8s.io/e2e-framework/pkg/envconf" + + "github.com/multigres/testkit/assert" ) // --------------------------------------------------------------------------- @@ -22,25 +24,14 @@ import ( // --------------------------------------------------------------------------- func TestDefaultE2EOpts(t *testing.T) { + c := assert.NewCollecting(t) o := defaultE2EOpts() - if !o.parallel { - t.Error("parallel should be true by default") - } - if o.clusterWaitTime != 5*time.Minute { - t.Errorf("clusterWaitTime = %v, want 5m", o.clusterWaitTime) - } - if o.kindConfigPath != "" { - t.Errorf("kindConfigPath = %q, want empty", o.kindConfigPath) - } - if len(o.images) != 0 { - t.Errorf("images = %v, want empty", o.images) - } - if len(o.setupFuncs) != 0 { - t.Errorf("setupFuncs should be empty") - } - if len(o.finishFuncs) != 0 { - t.Errorf("finishFuncs should be empty") - } + c.True(o.parallel, "parallel should be true by default") + c.Eq(5*time.Minute, o.clusterWaitTime, "clusterWaitTime") + c.Eq("", o.kindConfigPath, "kindConfigPath") + c.Empty(o.images, "images") + c.Empty(o.setupFuncs, "setupFuncs should be empty") + c.Empty(o.finishFuncs, "finishFuncs should be empty") } // --------------------------------------------------------------------------- @@ -50,33 +41,24 @@ func TestDefaultE2EOpts(t *testing.T) { func TestWithSequential(t *testing.T) { o := defaultE2EOpts() WithSequential()(o) - if o.parallel { - t.Error("parallel should be false after WithSequential") - } + assert.NewCollecting(t).False(o.parallel, "parallel should be false after WithSequential") } func TestWithE2EKindConfig(t *testing.T) { o := defaultE2EOpts() WithE2EKindConfig("/path/to/config.yaml")(o) - if o.kindConfigPath != "/path/to/config.yaml" { - t.Errorf("kindConfigPath = %q, want %q", o.kindConfigPath, "/path/to/config.yaml") - } + assert.NewCollecting(t).Eq("/path/to/config.yaml", o.kindConfigPath, "kindConfigPath") } func TestWithImage(t *testing.T) { + c := assert.NewCollecting(t) o := defaultE2EOpts() WithImage("img1:latest")(o) WithImage("img2:v2")(o) - if len(o.images) != 2 { - t.Fatalf("len(images) = %d, want 2", len(o.images)) - } - if o.images[0] != "img1:latest" { - t.Errorf("images[0] = %q, want %q", o.images[0], "img1:latest") - } - if o.images[1] != "img2:v2" { - t.Errorf("images[1] = %q, want %q", o.images[1], "img2:v2") - } + c.Require().Len(o.images, 2, "len(images) = %d, want 2", len(o.images)) + c.Eq("img1:latest", o.images[0], "images[0]") + c.Eq("img2:v2", o.images[1], "images[1]") } func TestWithSetup(t *testing.T) { @@ -84,9 +66,7 @@ func TestWithSetup(t *testing.T) { WithSetup(func(ctx context.Context, cfg *envconf.Config) (context.Context, error) { return ctx, nil })(o) - if len(o.setupFuncs) != 1 { - t.Errorf("len(setupFuncs) = %d, want 1", len(o.setupFuncs)) - } + assert.NewCollecting(t).Len(o.setupFuncs, 1, "len(setupFuncs) = %d, want 1", len(o.setupFuncs)) } func TestWithFinish(t *testing.T) { @@ -94,17 +74,14 @@ func TestWithFinish(t *testing.T) { WithFinish(func(ctx context.Context, cfg *envconf.Config) (context.Context, error) { return ctx, nil })(o) - if len(o.finishFuncs) != 1 { - t.Errorf("len(finishFuncs) = %d, want 1", len(o.finishFuncs)) - } + assert.NewCollecting(t). + Len(o.finishFuncs, 1, "len(finishFuncs) = %d, want 1", len(o.finishFuncs)) } func TestWithClusterWait(t *testing.T) { o := defaultE2EOpts() WithClusterWait(10 * time.Minute)(o) - if o.clusterWaitTime != 10*time.Minute { - t.Errorf("clusterWaitTime = %v, want 10m", o.clusterWaitTime) - } + assert.NewCollecting(t).Eq(10*time.Minute, o.clusterWaitTime, "clusterWaitTime") } // --------------------------------------------------------------------------- @@ -122,7 +99,10 @@ func TestSanitizeE2EName(t *testing.T) { {"Test Foo Bar", "test-foo-bar"}, {"UPPER", "upper"}, // Truncated to 50 chars - {"a-very-long-name-that-exceeds-fifty-characters-and-should-be-truncated", "a-very-long-name-that-exceeds-fifty-characters-and"}, + { + "a-very-long-name-that-exceeds-fifty-characters-and-should-be-truncated", + "a-very-long-name-that-exceeds-fifty-characters-and", + }, // Trailing dashes stripped {"-leading-and-trailing-", "leading-and-trailing"}, } @@ -130,9 +110,8 @@ func TestSanitizeE2EName(t *testing.T) { for _, tt := range tests { t.Run(tt.input, func(t *testing.T) { got := sanitizeE2EName(tt.input) - if got != tt.want { - t.Errorf("sanitizeE2EName(%q) = %q, want %q", tt.input, got, tt.want) - } + assert.NewCollecting(t). + Eq(tt.want, got, "sanitizeE2EName(%q) = %q, want", tt.input, got) }) } } @@ -206,9 +185,7 @@ func TestIsPodReady(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got := isPodReady(tt.pod) - if got != tt.want { - t.Errorf("isPodReady() = %v, want %v", got, tt.want) - } + assert.NewCollecting(t).Eq(tt.want, got, "isPodReady()") }) } } @@ -269,9 +246,7 @@ func TestSelectorFromDeployment(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got := selectorFromDeployment(tt.dep) - if got != tt.want { - t.Errorf("selectorFromDeployment() = %q, want %q", got, tt.want) - } + assert.NewCollecting(t).Eq(tt.want, got, "selectorFromDeployment()") }) } } @@ -281,14 +256,11 @@ func TestSelectorFromDeployment(t *testing.T) { // --------------------------------------------------------------------------- func TestInt64Ptr(t *testing.T) { + c := assert.NewCollecting(t) p := int64Ptr(42) - if *p != 42 { - t.Errorf("int64Ptr(42) = %d, want 42", *p) - } + c.Eq(42, *p, "int64Ptr(42)") p2 := int64Ptr(0) - if *p2 != 0 { - t.Errorf("int64Ptr(0) = %d, want 0", *p2) - } + c.Eq(0, *p2, "int64Ptr(0)") } // --------------------------------------------------------------------------- @@ -297,29 +269,20 @@ func TestInt64Ptr(t *testing.T) { func TestDefaultOperatorPreset_FromEnv(t *testing.T) { t.Setenv("OPERATOR_IMG", "my-registry/my-operator:v1.2.3") + c := assert.NewCollecting(t) p := DefaultOperatorPreset() - if p.Image != "my-registry/my-operator:v1.2.3" { - t.Errorf("Image = %q, want %q", p.Image, "my-registry/my-operator:v1.2.3") - } - if p.Namespace != "multigres-operator" { - t.Errorf("Namespace = %q, want %q", p.Namespace, "multigres-operator") - } - if p.DeploymentName != "multigres-operator-controller-manager" { - t.Errorf("DeploymentName = %q, want %q", p.DeploymentName, "multigres-operator-controller-manager") - } - if p.ReadyTimeout != 3*time.Minute { - t.Errorf("ReadyTimeout = %v, want 3m", p.ReadyTimeout) - } + c.Eq("my-registry/my-operator:v1.2.3", p.Image, "Image") + c.Eq("multigres-operator", p.Namespace, "Namespace") + c.Eq("multigres-operator-controller-manager", p.DeploymentName, "DeploymentName") + c.Eq(3*time.Minute, p.ReadyTimeout, "ReadyTimeout") } func TestDefaultOperatorPreset_Fallback(t *testing.T) { t.Setenv("OPERATOR_IMG", "") p := DefaultOperatorPreset() - if p.Image != "ghcr.io/multigres/multigres-operator:dev" { - t.Errorf("Image = %q, want fallback %q", p.Image, "ghcr.io/multigres/multigres-operator:dev") - } + assert.NewCollecting(t).Eq("ghcr.io/multigres/multigres-operator:dev", p.Image, "Image") } // --------------------------------------------------------------------------- @@ -327,14 +290,11 @@ func TestDefaultOperatorPreset_Fallback(t *testing.T) { // --------------------------------------------------------------------------- func TestMultigresImages(t *testing.T) { - if len(MultigresImages) != 4 { - t.Errorf("len(MultigresImages) = %d, want 4", len(MultigresImages)) - } + c := assert.NewCollecting(t) + c.Len(MultigresImages, 4, "len(MultigresImages) = %d, want 4", len(MultigresImages)) // Verify all entries are non-empty for i, img := range MultigresImages { - if img == "" { - t.Errorf("MultigresImages[%d] is empty", i) - } + c.NotEq("", img, "MultigresImages[%d] is empty", i) } } @@ -346,9 +306,7 @@ func TestTestClusterAccessors(t *testing.T) { tc := &TestCluster{ clusterName: "test-cluster-123", } - if tc.ClusterName() != "test-cluster-123" { - t.Errorf("ClusterName() = %q, want %q", tc.ClusterName(), "test-cluster-123") - } + assert.NewCollecting(t).Eq("test-cluster-123", tc.ClusterName(), "ClusterName()") } // --------------------------------------------------------------------------- @@ -382,6 +340,7 @@ func TestPortForwardResultStop(t *testing.T) { // --------------------------------------------------------------------------- func TestOptionsCompose(t *testing.T) { + c := assert.NewCollecting(t) o := defaultE2EOpts() opts := []E2EOption{ WithSequential(), @@ -394,18 +353,10 @@ func TestOptionsCompose(t *testing.T) { opt(o) } - if o.parallel { - t.Error("parallel should be false") - } - if len(o.images) != 2 { - t.Fatalf("len(images) = %d, want 2", len(o.images)) - } - if o.kindConfigPath != "/kind.yaml" { - t.Errorf("kindConfigPath = %q", o.kindConfigPath) - } - if o.clusterWaitTime != 7*time.Minute { - t.Errorf("clusterWaitTime = %v", o.clusterWaitTime) - } + c.False(o.parallel, "parallel should be false") + c.Require().Len(o.images, 2, "len(images) = %d, want 2", len(o.images)) + c.Eq("/kind.yaml", o.kindConfigPath, "kindConfigPath =") + c.Eq(7*time.Minute, o.clusterWaitTime, "clusterWaitTime =") } // --------------------------------------------------------------------------- @@ -413,25 +364,19 @@ func TestOptionsCompose(t *testing.T) { // --------------------------------------------------------------------------- func TestTestCluster_NilAccessors(t *testing.T) { + c := assert.NewCollecting(t) // Accessors should not panic even with nil fields — they just return nil. tc := &TestCluster{ clusterName: "x", } - if tc.Env() != nil { - t.Error("Env() should be nil") - } - if tc.Config() != nil { - t.Error("Config() should be nil") - } - if tc.RESTConfig() != nil { - t.Error("RESTConfig() should be nil") - } - if tc.Clientset() != nil { - t.Error("Clientset() should be nil") - } + c.Nil(tc.Env(), "Env() should be nil") + c.Nil(tc.Config(), "Config() should be nil") + c.Nil(tc.RESTConfig(), "RESTConfig() should be nil") + c.Nil(tc.Clientset(), "Clientset() should be nil") } func TestTestCluster_WithConfig(t *testing.T) { + c := assert.NewCollecting(t) cfg := envconf.New() testEnv := env.NewWithConfig(cfg) tc := &TestCluster{ @@ -439,15 +384,9 @@ func TestTestCluster_WithConfig(t *testing.T) { cfg: cfg, env: testEnv, } - if tc.Config() != cfg { - t.Error("Config() should return the assigned config") - } - if tc.KubeconfigFile() != cfg.KubeconfigFile() { - t.Errorf("KubeconfigFile() = %q, want %q", tc.KubeconfigFile(), cfg.KubeconfigFile()) - } - if tc.Env() == nil { - t.Error("Env() should not be nil") - } + c.Eq(cfg, tc.Config(), "Config() should return the assigned config") + c.Eq(cfg.KubeconfigFile(), tc.KubeconfigFile(), "KubeconfigFile()") + c.NotNil(tc.Env(), "Env() should not be nil") // KlientClient() requires a real kubeconfig — tested via e2e tests only. } @@ -467,9 +406,7 @@ func TestRunFinish(t *testing.T) { return ctx, nil }, }) - if !called { - t.Error("runFinish did not call the func") - } + assert.NewCollecting(t).True(called, "runFinish did not call the func") } func TestRunFinish_Empty(t *testing.T) { @@ -510,9 +447,8 @@ func TestClusterCleanupAction(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got := clusterCleanupAction(tt.policy, tt.failed) - if got != tt.want { - t.Errorf("clusterCleanupAction(%q, %v) = %d, want %d", tt.policy, tt.failed, got, tt.want) - } + assert.NewCollecting(t). + Eq(tt.want, got, "clusterCleanupAction(%q, %v) = %d, want", tt.policy, tt.failed, got) }) } } @@ -557,5 +493,7 @@ func TestMaybeDestroyCluster_OnFailurePassingTest(t *testing.T) { } // Ensure unused imports are consumed. -var _ = os.Getenv -var _ = ptr.To[int32] +var ( + _ = os.Getenv + _ = ptr.To[int32] +) diff --git a/pkg/testutil/envtest_internal_test.go b/pkg/testutil/envtest_internal_test.go index ae523955..16aeac3a 100644 --- a/pkg/testutil/envtest_internal_test.go +++ b/pkg/testutil/envtest_internal_test.go @@ -14,34 +14,32 @@ import ( "k8s.io/client-go/rest" "sigs.k8s.io/controller-runtime/pkg/envtest" "sigs.k8s.io/controller-runtime/pkg/manager" + + "github.com/multigres/testkit/assert" ) // TestCreateEnvtestEnvironment tests environment creation. func TestCreateEnvtestEnvironment(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) testPaths := []string{"/test/path1", "/test/path2"} env := createEnvtestEnvironment(t, testPaths) - if env == nil { - t.Fatal("createEnvtestEnvironment() returned nil") - } + c.Require().NotNil(env, "createEnvtestEnvironment() returned nil") - if len(env.CRDDirectoryPaths) != 2 { - t.Errorf("CRDDirectoryPaths length = %d, want 2", len(env.CRDDirectoryPaths)) - } + c.Len( + env.CRDDirectoryPaths, + 2, + "CRDDirectoryPaths length = %d, want 2", + len(env.CRDDirectoryPaths), + ) - if env.CRDDirectoryPaths[0] != testPaths[0] { - t.Errorf("CRDDirectoryPaths[0] = %s, want %s", env.CRDDirectoryPaths[0], testPaths[0]) - } + c.Eq(testPaths[0], env.CRDDirectoryPaths[0], "CRDDirectoryPaths[0]") - if env.CRDDirectoryPaths[1] != testPaths[1] { - t.Errorf("CRDDirectoryPaths[1] = %s, want %s", env.CRDDirectoryPaths[1], testPaths[1]) - } + c.Eq(testPaths[1], env.CRDDirectoryPaths[1], "CRDDirectoryPaths[1]") - if !env.ErrorIfCRDPathMissing { - t.Error("ErrorIfCRDPathMissing should be true") - } + c.True(env.ErrorIfCRDPathMissing, "ErrorIfCRDPathMissing should be true") } // TestStartEnvtest tests both success and error paths. @@ -73,6 +71,7 @@ func TestStartEnvtest(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mock := &mockTB{TB: t} env := tc.setupFunc(mock) @@ -80,17 +79,11 @@ func TestStartEnvtest(t *testing.T) { cfg := startEnvtest(mock, env) - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + c.Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") if !tc.wantFatal { - if cfg == nil { - t.Error("Config should not be nil on success") - } - if cfg.Host == "" { - t.Error("Config.Host should not be empty on success") - } + c.NotNil(cfg, "Config should not be nil on success") + c.NotEq("", cfg.Host, "Config.Host should not be empty on success") } }) } @@ -117,6 +110,7 @@ func TestCreateEnvtestDir(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mock := &mockTB{TB: t} @@ -125,13 +119,9 @@ func TestCreateEnvtestDir(t *testing.T) { os.RemoveAll(dir) }) - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + c.Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") - if !tc.wantFatal && dir == "" { - t.Error("createEnvtestDir() returned empty string") - } + c.False(!tc.wantFatal && dir == "", "createEnvtestDir() returned empty string") }) } } @@ -184,38 +174,29 @@ func TestWriteKubeconfigFile(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mock := &mockTB{TB: t} path := tc.setupPath(t) writeKubeconfigFile(mock, path, tc.content) - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + c.Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") if !tc.wantFatal { // Verify file was written content, err := os.ReadFile(path) - if err != nil { - t.Fatalf("Failed to read written file: %v", err) - } + c.Require().NoError(err, "Failed to read written file") // Verify content matches - if string(content) != string(tc.content) { - t.Errorf("File content = %q, want %q", string(content), string(tc.content)) - } + c.Eq(string(tc.content), string(content), "File content") // Verify file permissions info, err := os.Stat(path) - if err != nil { - t.Fatalf("Failed to stat file: %v", err) - } + c.Require().NoError(err, "Failed to stat file") expectedMode := os.FileMode(0o600) - if info.Mode().Perm() != expectedMode { - t.Errorf("File permissions = %o, want %o", info.Mode().Perm(), expectedMode) - } + c.Eq(expectedMode, info.Mode().Perm(), "File permissions") } }) } @@ -255,9 +236,7 @@ func TestSetUpClient(t *testing.T) { SetUpClient(mock, cfg, scheme) - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + assert.NewCollecting(t).Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") }) } } @@ -315,9 +294,8 @@ func TestStartManager_Internal(t *testing.T) { scheme := runtime.NewScheme() mgr := SetUpManager(t, cfg, scheme) // Add a runnable that will fail on Start - if err := mgr.Add(&failingRunnable{}); err != nil { - t.Fatalf("Failed to add failing runnable: %v", err) - } + assert.NewAborting(t). + NoError(mgr.Add(&failingRunnable{}), "Failed to add failing runnable") return t.Context(), mgr }, wantFatal: false, @@ -328,6 +306,7 @@ func TestStartManager_Internal(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mock := &mockTB{TB: t} ctx, mgr := tc.setupFunc(mock) @@ -341,13 +320,9 @@ func TestStartManager_Internal(t *testing.T) { <-done } - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + c.Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") - if mock.errorCalled != tc.wantError { - t.Errorf("errorCalled = %v, want %v", mock.errorCalled, tc.wantError) - } + c.Eq(tc.wantError, mock.errorCalled, "errorCalled") }) } } @@ -384,9 +359,7 @@ func TestCleanEnvtest(t *testing.T) { cleanup := cleanEnvtest(mock, env) cleanup() - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + assert.NewCollecting(t).Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") }) } } @@ -433,9 +406,7 @@ func TestSetUpManager(t *testing.T) { SetUpManager(mock, cfg, scheme) - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + assert.NewCollecting(t).Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") }) } } @@ -465,6 +436,7 @@ func TestGenerateKubeconfigFile(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mock := &mockTB{TB: t} @@ -475,13 +447,12 @@ func TestGenerateKubeconfigFile(t *testing.T) { } }) - if mock.fatalCalled != tc.wantFatal { - t.Errorf("fatalCalled = %v, want %v", mock.fatalCalled, tc.wantFatal) - } + c.Eq(tc.wantFatal, mock.fatalCalled, "fatalCalled") - if !tc.wantFatal && kubeconfigPath == "" { - t.Error("kubeconfigPath should not be empty on success") - } + c.False( + !tc.wantFatal && kubeconfigPath == "", + "kubeconfigPath should not be empty on success", + ) }) } } @@ -549,9 +520,8 @@ func TestGetKubeconfigFromUserAdder(t *testing.T) { t.Parallel() _, err := getKubeconfigFromUserAdder(tc.adder) - if (err != nil) != tc.wantError { - t.Errorf("getKubeconfigFromUserAdder() error = %v, wantError %v", err, tc.wantError) - } + assert.NewCollecting(t). + ErrorWhen(tc.wantError, err, "getKubeconfigFromUserAdder() error") }) } } diff --git a/pkg/testutil/envtest_test.go b/pkg/testutil/envtest_test.go index c668e0eb..6daf7f4e 100644 --- a/pkg/testutil/envtest_test.go +++ b/pkg/testutil/envtest_test.go @@ -13,6 +13,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/healthz" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) func TestSetUpEnvTest(t *testing.T) { @@ -37,6 +39,7 @@ func TestSetUpEnvTest(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -44,12 +47,8 @@ func TestSetUpEnvTest(t *testing.T) { cfg := testutil.SetUpEnvtest(t, tc.opts...) mgr := testutil.SetUpManager(t, cfg, scheme) - if err := mgr.AddHealthzCheck("healthz", healthz.Ping); err != nil { - t.Fatalf("Failed to set up health check, %v", err) - } - if err := mgr.AddReadyzCheck("readyz", healthz.Ping); err != nil { - t.Fatalf("Failed to set up ready check, %v", err) - } + c.NoError(mgr.AddHealthzCheck("healthz", healthz.Ping), "Failed to set up health check") + c.NoError(mgr.AddReadyzCheck("readyz", healthz.Ping), "Failed to set up ready check") testutil.StartManager(t, mgr) }) @@ -58,6 +57,7 @@ func TestSetUpEnvTest(t *testing.T) { func TestSetUpClient(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -65,32 +65,22 @@ func TestSetUpClient(t *testing.T) { cfg := testutil.SetUpEnvtest(t) client := testutil.SetUpClient(t, cfg, scheme) - if client == nil { - t.Fatal("SetUpClient() returned nil") - } + c.Require().NotNil(client, "SetUpClient() returned nil") // Verify client works by listing services svcList := &corev1.ServiceList{} - if err := client.List(t.Context(), svcList); err != nil { - t.Errorf("Client.List() failed: %v", err) - } + c.NoError(client.List(t.Context(), svcList), "Client.List() failed") // Create and retrieve a service (tests direct API server access without cache) svc := &corev1.Service{ ObjectMeta: testutil.Obj[corev1.Service]("test-svc", "default").ObjectMeta, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := client.Create(t.Context(), svc); err != nil { - t.Fatalf("Client.Create() failed: %v", err) - } + c.Require().NoError(client.Create(t.Context(), svc), "Client.Create() failed") retrieved := &corev1.Service{} objKey := types.NamespacedName{Name: "test-svc", Namespace: "default"} - if err := client.Get(t.Context(), objKey, retrieved); err != nil { - t.Errorf("Client.Get() failed: %v", err) - } + c.NoError(client.Get(t.Context(), objKey, retrieved), "Client.Get() failed") - if retrieved.Name != "test-svc" { - t.Errorf("Retrieved service Name = %s, want test-svc", retrieved.Name) - } + c.Eq("test-svc", retrieved.Name, "Retrieved service Name") } diff --git a/pkg/testutil/fake_client_test.go b/pkg/testutil/fake_client_test.go index a6fc84bb..b55c22fd 100644 --- a/pkg/testutil/fake_client_test.go +++ b/pkg/testutil/fake_client_test.go @@ -9,6 +9,8 @@ import ( "k8s.io/apimachinery/pkg/runtime" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestFakeClientWithFailures_Get(t *testing.T) { @@ -85,9 +87,7 @@ func TestFakeClientWithFailures_Get(t *testing.T) { result := &corev1.Pod{} err := fakeClient.Get(context.Background(), tc.key, result) - if (err != nil) != tc.wantErr { - t.Errorf("Get() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Get() error") }) } } @@ -151,9 +151,7 @@ func TestFakeClientWithFailures_Create(t *testing.T) { err := fakeClient.Create(context.Background(), tc.obj) - if (err != nil) != tc.wantErr { - t.Errorf("Create() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Create() error") }) } } @@ -200,9 +198,7 @@ func TestFakeClientWithFailures_Update(t *testing.T) { err := fakeClient.Update(context.Background(), pod) - if (err != nil) != tc.wantErr { - t.Errorf("Update() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Update() error") }) } } @@ -255,9 +251,7 @@ func TestFakeClientWithFailures_Delete(t *testing.T) { err := fakeClient.Delete(context.Background(), pod) - if (err != nil) != tc.wantErr { - t.Errorf("Delete() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Delete() error") }) } } @@ -305,9 +299,7 @@ func TestFakeClientWithFailures_StatusUpdate(t *testing.T) { err := fakeClient.Status().Update(context.Background(), pod) - if (err != nil) != tc.wantErr { - t.Errorf("Status().Update() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Status().Update() error") }) } } @@ -357,9 +349,7 @@ func TestFakeClientWithFailures_List(t *testing.T) { podList := &corev1.PodList{} err := fakeClient.List(context.Background(), podList) - if (err != nil) != tc.wantErr { - t.Errorf("List() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "List() error") }) } } @@ -407,9 +397,7 @@ func TestFakeClientWithFailures_Patch(t *testing.T) { patch := client.MergeFrom(pod.DeepCopy()) err := fakeClient.Patch(context.Background(), pod, patch) - if (err != nil) != tc.wantErr { - t.Errorf("Patch() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Patch() error") }) } } @@ -462,9 +450,7 @@ func TestFakeClientWithFailures_DeleteAllOf(t *testing.T) { client.InNamespace("default"), ) - if (err != nil) != tc.wantErr { - t.Errorf("DeleteAllOf() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "DeleteAllOf() error") }) } } @@ -513,9 +499,7 @@ func TestFakeClientWithFailures_StatusPatch(t *testing.T) { patch := client.MergeFrom(pod.DeepCopy()) err := fakeClient.Status().Patch(context.Background(), pod, patch) - if (err != nil) != tc.wantErr { - t.Errorf("Status().Patch() error = %v, wantErr %v", err, tc.wantErr) - } + assert.NewCollecting(t).ErrorWhen(tc.wantErr, err, "Status().Patch() error") }) } } @@ -567,9 +551,8 @@ func TestHelperFunctions_ObjectMatchers(t *testing.T) { fn := tc.setupFn() err := fn(pod) - if err != tc.wantErr { - t.Errorf("Expected error %v, got %v", tc.wantErr, err) - } + assert.NewCollecting(t). + False(err != tc.wantErr, "Expected error %v, got %v", tc.wantErr, err) }) } } @@ -626,9 +609,8 @@ func TestHelperFunctions_KeyMatchers(t *testing.T) { fn := tc.setupFn() err := fn(tc.key) - if err != tc.wantErr { - t.Errorf("Expected error %v, got %v", tc.wantErr, err) - } + assert.NewCollecting(t). + False(err != tc.wantErr, "Expected error %v, got %v", tc.wantErr, err) }) } } @@ -802,9 +784,8 @@ func TestHelperFunctions_AlwaysFail(t *testing.T) { fn := AlwaysFail(tc.wantErr) err := fn(tc.input) - if err != tc.wantErr { - t.Errorf("Expected error %v, got %v", tc.wantErr, err) - } + assert.NewCollecting(t). + False(err != tc.wantErr, "Expected error %v, got %v", tc.wantErr, err) }) } } @@ -816,9 +797,8 @@ func TestHelperFunctions_Panic(t *testing.T) { t.Parallel() defer func() { - if r := recover(); r == nil { - t.Errorf("Expected panic when meta.Accessor fails on nil") - } + assert.NewCollecting(t). + NotNil(recover(), "Expected panic when meta.Accessor fails on nil") }() fn := FailOnObjectName("test", ErrInjected) @@ -829,9 +809,8 @@ func TestHelperFunctions_Panic(t *testing.T) { t.Parallel() defer func() { - if r := recover(); r == nil { - t.Errorf("Expected panic when meta.Accessor fails on nil") - } + assert.NewCollecting(t). + NotNil(recover(), "Expected panic when meta.Accessor fails on nil") }() fn := FailOnNamespace("default", ErrInjected) diff --git a/pkg/testutil/kind_test.go b/pkg/testutil/kind_test.go index 02fb4d2c..2bd0e6c5 100644 --- a/pkg/testutil/kind_test.go +++ b/pkg/testutil/kind_test.go @@ -4,129 +4,103 @@ package testutil import ( "testing" + + "github.com/multigres/testkit/assert" ) func TestDefaultKindConfig_ClusterNameFromEnv(t *testing.T) { t.Setenv("KIND_CLUSTER", "my-custom-cluster") cfg := defaultKindConfig() - if cfg.clusterName != "my-custom-cluster" { - t.Errorf("clusterName = %q, want %q", cfg.clusterName, "my-custom-cluster") - } + assert.NewCollecting(t).Eq("my-custom-cluster", cfg.clusterName, "clusterName") } func TestDefaultKindConfig_ClusterNameFallback(t *testing.T) { t.Setenv("KIND_CLUSTER", "") cfg := defaultKindConfig() - if cfg.clusterName != defaultKindCluster { - t.Errorf("clusterName = %q, want %q", cfg.clusterName, defaultKindCluster) - } + assert.NewCollecting(t).Eq(defaultKindCluster, cfg.clusterName, "clusterName") } func TestDefaultKindConfig_KubectlFromEnv(t *testing.T) { t.Setenv("KUBECTL", "/usr/local/bin/kubectl") cfg := defaultKindConfig() - if cfg.kubectl != "/usr/local/bin/kubectl" { - t.Errorf("kubectl = %q, want %q", cfg.kubectl, "/usr/local/bin/kubectl") - } + assert.NewCollecting(t).Eq("/usr/local/bin/kubectl", cfg.kubectl, "kubectl") } func TestDefaultKindConfig_KubectlFallback(t *testing.T) { t.Setenv("KUBECTL", "") cfg := defaultKindConfig() - if cfg.kubectl != defaultKubectl { - t.Errorf("kubectl = %q, want %q", cfg.kubectl, defaultKubectl) - } + assert.NewCollecting(t).Eq(defaultKubectl, cfg.kubectl, "kubectl") } func TestWithKindCluster(t *testing.T) { cfg := defaultKindConfig() WithKindCluster("override-cluster")(cfg) - if cfg.clusterName != "override-cluster" { - t.Errorf("clusterName = %q, want %q", cfg.clusterName, "override-cluster") - } + assert.NewCollecting(t).Eq("override-cluster", cfg.clusterName, "clusterName") } func TestWithKindKubectl(t *testing.T) { cfg := defaultKindConfig() WithKindKubectl("/opt/bin/kubectl")(cfg) - if cfg.kubectl != "/opt/bin/kubectl" { - t.Errorf("kubectl = %q, want %q", cfg.kubectl, "/opt/bin/kubectl") - } + assert.NewCollecting(t).Eq("/opt/bin/kubectl", cfg.kubectl, "kubectl") } func TestWithKindCRDPaths(t *testing.T) { + c := assert.NewCollecting(t) cfg := defaultKindConfig() WithKindCRDPaths("/path/to/crds", "/another/path")(cfg) - if len(cfg.crdPaths) != 2 { - t.Fatalf("len(crdPaths) = %d, want 2", len(cfg.crdPaths)) - } - if cfg.crdPaths[0] != "/path/to/crds" { - t.Errorf("crdPaths[0] = %q, want %q", cfg.crdPaths[0], "/path/to/crds") - } + c.Require().Len(cfg.crdPaths, 2, "len(crdPaths) = %d, want 2", len(cfg.crdPaths)) + c.Eq("/path/to/crds", cfg.crdPaths[0], "crdPaths[0]") } func TestWithKindCreateCluster(t *testing.T) { + c := assert.NewCollecting(t) cfg := defaultKindConfig() - if cfg.createCluster { - t.Error("createCluster should be false by default") - } + c.False(cfg.createCluster, "createCluster should be false by default") WithKindCreateCluster()(cfg) - if !cfg.createCluster { - t.Error("createCluster should be true after WithKindCreateCluster") - } + c.True(cfg.createCluster, "createCluster should be true after WithKindCreateCluster") } func TestWithKindImages(t *testing.T) { + c := assert.NewCollecting(t) cfg := defaultKindConfig() WithKindImages("img1:latest", "img2:v1")(cfg) - if len(cfg.images) != 2 { - t.Fatalf("len(images) = %d, want 2", len(cfg.images)) - } - if cfg.images[0] != "img1:latest" { - t.Errorf("images[0] = %q, want %q", cfg.images[0], "img1:latest") - } - if cfg.images[1] != "img2:v1" { - t.Errorf("images[1] = %q, want %q", cfg.images[1], "img2:v1") - } + c.Require().Len(cfg.images, 2, "len(images) = %d, want 2", len(cfg.images)) + c.Eq("img1:latest", cfg.images[0], "images[0]") + c.Eq("img2:v1", cfg.images[1], "images[1]") } func TestKindClusterName(t *testing.T) { name := KindClusterName(WithKindCluster("test-cluster")) - if name != "test-cluster" { - t.Errorf("KindClusterName = %q, want %q", name, "test-cluster") - } + assert.NewCollecting(t).Eq("test-cluster", name, "KindClusterName") } func TestKindClusterName_Default(t *testing.T) { t.Setenv("KIND_CLUSTER", "") name := KindClusterName() - if name != defaultKindCluster { - t.Errorf("KindClusterName = %q, want %q", name, defaultKindCluster) - } + assert.NewCollecting(t).Eq(defaultKindCluster, name, "KindClusterName") } func TestRandomSuffix(t *testing.T) { + ck := assert.NewCollecting(t) s1 := randomSuffix() s2 := randomSuffix() - if len(s1) != 8 { - t.Errorf("len(randomSuffix()) = %d, want 8", len(s1)) - } - if s1 == s2 { - t.Errorf("randomSuffix() returned same value twice: %q", s1) - } + ck.Len(s1, 8, "len(randomSuffix()) = %d, want 8", len(s1)) + ck.NotEq(s2, s1, "randomSuffix() returned same value twice") // Check all chars are valid for _, c := range s1 { - if !((c >= 'a' && c <= 'z') || (c >= '0' && c <= '9')) { - t.Errorf("randomSuffix() contains invalid char %q", c) - } + ck.False( + !((c >= 'a' && c <= 'z') || (c >= '0' && c <= '9')), + "randomSuffix() contains invalid char %q", + c, + ) } } diff --git a/pkg/testutil/resource_watcher_cache_internal_test.go b/pkg/testutil/resource_watcher_cache_internal_test.go index 2709c404..2c1ff790 100644 --- a/pkg/testutil/resource_watcher_cache_internal_test.go +++ b/pkg/testutil/resource_watcher_cache_internal_test.go @@ -5,6 +5,8 @@ import ( corev1 "k8s.io/api/core/v1" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) func TestFindLatestEventFor(t *testing.T) { @@ -112,6 +114,7 @@ func TestFindLatestEventFor(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) watcher := &ResourceWatcher{ t: t, @@ -120,23 +123,16 @@ func TestFindLatestEventFor(t *testing.T) { evt := watcher.findLatestEventFor(tc.setupObj()) - if (evt != nil) != tc.wantFound { - t.Fatalf("findLatestEventFor() found=%v, want=%v", evt != nil, tc.wantFound) - } + c.Require(). + Eq(tc.wantFound, (evt != nil), "findLatestEventFor() found=%v, want=", evt != nil) if !tc.wantFound { return } - if evt.Name != tc.wantName { - t.Errorf("findLatestEventFor() Name = %s, want %s", evt.Name, tc.wantName) - } - if evt.Namespace != tc.wantNS { - t.Errorf("findLatestEventFor() Namespace = %s, want %s", evt.Namespace, tc.wantNS) - } - if evt.Kind != tc.wantKind { - t.Errorf("findLatestEventFor() Kind = %s, want %s", evt.Kind, tc.wantKind) - } + c.Eq(tc.wantName, evt.Name, "findLatestEventFor() Name") + c.Eq(tc.wantNS, evt.Namespace, "findLatestEventFor() Namespace") + c.Eq(tc.wantKind, evt.Kind, "findLatestEventFor() Kind") }) } } @@ -182,9 +178,8 @@ func TestFindLatestEvent(t *testing.T) { evt := watcher.findLatestEvent(tc.matchFunc) - if (evt != nil) != tc.wantFound { - t.Errorf("findLatestEvent() found=%v, want=%v", evt != nil, tc.wantFound) - } + assert.NewCollecting(t). + Eq(tc.wantFound, (evt != nil), "findLatestEvent() found=%v, want=", evt != nil) }) } } @@ -214,6 +209,7 @@ func TestCheckLatestEventMatches(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) watcher := &ResourceWatcher{ t: t, @@ -222,16 +218,8 @@ func TestCheckLatestEventMatches(t *testing.T) { matched, diff := watcher.checkLatestEventMatches(tc.setupObj(), nil) - if matched != tc.wantMatched { - t.Errorf("checkLatestEventMatches() matched = %v, want %v", matched, tc.wantMatched) - } - if diff != tc.wantDiff { - t.Errorf( - "checkLatestEventMatches() diff = %s, want %s", - diff, - tc.wantDiff, - ) - } + c.Eq(tc.wantMatched, matched, "checkLatestEventMatches() matched") + c.Eq(tc.wantDiff, diff, "checkLatestEventMatches() diff") }) } } diff --git a/pkg/testutil/resource_watcher_deletion_internal_test.go b/pkg/testutil/resource_watcher_deletion_internal_test.go index 04639ee1..c995005d 100644 --- a/pkg/testutil/resource_watcher_deletion_internal_test.go +++ b/pkg/testutil/resource_watcher_deletion_internal_test.go @@ -12,6 +12,8 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/runtime" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) // TestWaitForDeletion_EmptySlice tests WaitForDeletion with empty slice. @@ -24,9 +26,8 @@ func TestWaitForDeletion_EmptySlice(t *testing.T) { } err := watcher.WaitForDeletion() - if err != nil { - t.Errorf("WaitForDeletion() with empty slice should return nil, got: %v", err) - } + assert.NewCollecting(t). + NoError(err, "WaitForDeletion() with empty slice should return nil, got") } // TestDeletionPredicate_NonMatchingEvents tests deletion predicate with various @@ -78,9 +79,7 @@ func TestDeletionPredicate_NonMatchingEvents(t *testing.T) { // Wait for deletion of non-existent service (times out) err := watcher.WaitForDeletion(Obj[corev1.Service]("svc-target", "default")) - if err == nil { - t.Error("Expected timeout") - } + assert.NewCollecting(t).Error(err, "Expected timeout") } // TestWaitForSingleDeletion_SuccessWithUpdates tests deletion success after @@ -125,7 +124,5 @@ func TestWaitForSingleDeletion_SuccessWithUpdates(t *testing.T) { }() err := watcher.WaitForDeletion(Obj[corev1.Service]("svc-upd-del", "default")) - if err != nil { - t.Errorf("WaitForDeletion() should succeed, got: %v", err) - } + assert.NewCollecting(t).NoError(err, "WaitForDeletion() should succeed, got") } diff --git a/pkg/testutil/resource_watcher_deletion_test.go b/pkg/testutil/resource_watcher_deletion_test.go index 5fdb85e6..5a41c58d 100644 --- a/pkg/testutil/resource_watcher_deletion_test.go +++ b/pkg/testutil/resource_watcher_deletion_test.go @@ -7,7 +7,6 @@ import ( "context" "testing" - "github.com/google/go-cmp/cmp" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -16,6 +15,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) // TestObj tests the generic Obj helper function. @@ -104,9 +105,7 @@ func TestObj(t *testing.T) { t.Run(name, func(t *testing.T) { t.Parallel() - if diff := cmp.Diff(tc.expected, tc.obj); diff != "" { - t.Errorf("Obj mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.expected, tc.obj, "Obj mismatch") }) } } @@ -146,9 +145,7 @@ func TestWaitForDeletion(t *testing.T) { }, assertFunc: func(t *testing.T, watcher *testutil.ResourceWatcher) { err := watcher.WaitForDeletion(testutil.Obj[corev1.Service]("test-svc", "default")) - if err != nil { - t.Errorf("Failed to wait for deletion: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for deletion") }, }, "multiple services deletion": { @@ -184,7 +181,10 @@ func TestWaitForDeletion(t *testing.T) { return watcher.WaitForMatch(svc1, svc2) }, delete: func(ctx context.Context, c client.Client) error { - if err := c.Delete(ctx, testutil.Obj[corev1.Service]("svc-1", "default")); err != nil { + if err := c.Delete( + ctx, + testutil.Obj[corev1.Service]("svc-1", "default"), + ); err != nil { return err } return c.Delete(ctx, testutil.Obj[corev1.Service]("svc-2", "default")) @@ -194,9 +194,7 @@ func TestWaitForDeletion(t *testing.T) { testutil.Obj[corev1.Service]("svc-1", "default"), testutil.Obj[corev1.Service]("svc-2", "default"), ) - if err != nil { - t.Errorf("Failed to wait for multiple deletions: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for multiple deletions") }, }, "mixed resource types deletion": { @@ -248,7 +246,10 @@ func TestWaitForDeletion(t *testing.T) { return watcher.WaitForMatch(svc, deploy) }, delete: func(ctx context.Context, c client.Client) error { - if err := c.Delete(ctx, testutil.Obj[corev1.Service]("my-svc", "default")); err != nil { + if err := c.Delete( + ctx, + testutil.Obj[corev1.Service]("my-svc", "default"), + ); err != nil { return err } return c.Delete(ctx, testutil.Obj[appsv1.Deployment]("my-deploy", "default")) @@ -258,9 +259,7 @@ func TestWaitForDeletion(t *testing.T) { testutil.Obj[corev1.Service]("my-svc", "default"), testutil.Obj[appsv1.Deployment]("my-deploy", "default"), ) - if err != nil { - t.Errorf("Failed to wait for mixed type deletions: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for mixed type deletions") }, }, } @@ -269,6 +268,7 @@ func TestWaitForDeletion(t *testing.T) { name, tc := name, tc t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -279,13 +279,9 @@ func TestWaitForDeletion(t *testing.T) { c := mgr.GetClient() watcher := testutil.NewResourceWatcher(t, ctx, mgr) - if err := tc.setup(ctx, c, watcher); err != nil { - t.Fatalf("Setup failed: %v", err) - } + ck.NoError(tc.setup(ctx, c, watcher), "Setup failed") - if err := tc.delete(ctx, c); err != nil { - t.Fatalf("Delete failed: %v", err) - } + ck.NoError(tc.delete(ctx, c), "Delete failed") tc.assertFunc(t, watcher) }) @@ -296,6 +292,7 @@ func TestWaitForDeletion(t *testing.T) { // existing DELETED event. func TestWaitForDeletion_ExistingDeletedEvent(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -312,32 +309,22 @@ func TestWaitForDeletion_ExistingDeletedEvent(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "to-delete", Namespace: "default"}, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := c.Create(ctx, svc); err != nil { - t.Fatalf("Failed to create Service: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc), "Failed to create Service") // Wait for creation watcher.SetCmpOpts(testutil.IgnoreMetaRuntimeFields(), testutil.IgnoreServiceRuntimeFields()) - if err := watcher.WaitForMatch(svc); err != nil { - t.Fatalf("Failed to wait for Service creation: %v", err) - } + ck.Require().NoError(watcher.WaitForMatch(svc), "Failed to wait for Service creation") // Delete it - if err := c.Delete(ctx, svc); err != nil { - t.Fatalf("Failed to delete Service: %v", err) - } + ck.Require().NoError(c.Delete(ctx, svc), "Failed to delete Service") // Wait for deletion event evt, err := watcher.WaitForEventType("Service", "DELETED") - if err != nil { - t.Fatalf("WaitForEventType() error = %v", err) - } + ck.Require().NoError(err, "WaitForEventType() error =") // Now wait for deletion again - should find existing event err = watcher.WaitForDeletion(testutil.Obj[corev1.Service]("to-delete", "default")) - if err != nil { - t.Errorf("WaitForDeletion() error = %v, want nil (should find existing event)", err) - } + ck.NoError(err, "WaitForDeletion() error") t.Logf("Successfully found existing DELETED event: %+v", evt) } @@ -361,9 +348,7 @@ func TestWaitForDeletion_ContextCanceled(t *testing.T) { // Try to wait for deletion - should fail with watcher stopped err := watcher.WaitForDeletion(testutil.Obj[corev1.Service]("test", "default")) - if err == nil { - t.Error("Expected error when context is canceled") - } + assert.NewCollecting(t).Error(err, "Expected error when context is canceled") t.Logf("Got expected error: %v", err) } @@ -391,6 +376,7 @@ func TestWaitForDeletion_CascadingDelete(t *testing.T) { "See: https://github.com/kubernetes-sigs/controller-runtime/issues/626") t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -405,9 +391,7 @@ func TestWaitForDeletion_CascadingDelete(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "owner", Namespace: "default"}, Data: map[string]string{"key": "value"}, } - if err := c.Create(ctx, owner); err != nil { - t.Fatalf("Failed to create owner: %v", err) - } + ck.Require().NoError(c.Create(ctx, owner), "Failed to create owner") svc := &corev1.Service{ ObjectMeta: metav1.ObjectMeta{ @@ -424,23 +408,18 @@ func TestWaitForDeletion_CascadingDelete(t *testing.T) { }, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := c.Create(ctx, svc); err != nil { - t.Fatalf("Failed to create owned service: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc), "Failed to create owned service") // Wait for service to be created watcher.SetCmpOpts(testutil.IgnoreMetaRuntimeFields(), testutil.IgnoreServiceRuntimeFields()) - if err := watcher.WaitForMatch(svc); err != nil { - t.Fatalf("Failed to wait for service creation: %v", err) - } + ck.Require().NoError(watcher.WaitForMatch(svc), "Failed to wait for service creation") // Delete owner - should cascade to owned service - if err := c.Delete(ctx, owner); err != nil { - t.Fatalf("Failed to delete owner: %v", err) - } + ck.Require().NoError(c.Delete(ctx, owner), "Failed to delete owner") // Wait for cascading deletion - if err := watcher.WaitForDeletion(testutil.Obj[corev1.Service]("owned-svc", "default")); err != nil { - t.Errorf("Cascading deletion failed: %v", err) - } + ck.NoError( + watcher.WaitForDeletion(testutil.Obj[corev1.Service]("owned-svc", "default")), + "Cascading deletion failed", + ) } diff --git a/pkg/testutil/resource_watcher_internal_test.go b/pkg/testutil/resource_watcher_internal_test.go index 51250061..f8418de0 100644 --- a/pkg/testutil/resource_watcher_internal_test.go +++ b/pkg/testutil/resource_watcher_internal_test.go @@ -7,6 +7,8 @@ import ( "github.com/google/go-cmp/cmp" corev1 "k8s.io/api/core/v1" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) // TestSetTimeout verifies SetTimeout updates the timeout field. @@ -21,9 +23,7 @@ func TestSetTimeout(t *testing.T) { newTimeout := 20 * time.Second watcher.SetTimeout(newTimeout) - if watcher.timeout != newTimeout { - t.Errorf("SetTimeout() timeout = %v, want %v", watcher.timeout, newTimeout) - } + assert.NewCollecting(t).Eq(newTimeout, watcher.timeout, "SetTimeout() timeout") } // TestResetTimeout verifies ResetTimeout restores default (5 seconds). @@ -38,9 +38,7 @@ func TestResetTimeout(t *testing.T) { watcher.ResetTimeout() expectedDefault := 5 * time.Second - if watcher.timeout != expectedDefault { - t.Errorf("ResetTimeout() timeout = %v, want %v (default)", watcher.timeout, expectedDefault) - } + assert.NewCollecting(t).Eq(expectedDefault, watcher.timeout, "ResetTimeout() timeout") } // TestSetCmpOpts verifies SetCmpOpts updates the cmpOpts field. @@ -55,9 +53,8 @@ func TestSetCmpOpts(t *testing.T) { newOpts := []cmp.Option{IgnoreMetaRuntimeFields(), IgnoreStatus()} watcher.SetCmpOpts(newOpts...) - if len(watcher.cmpOpts) != 2 { - t.Errorf("SetCmpOpts() cmpOpts length = %d, want 2", len(watcher.cmpOpts)) - } + assert.NewCollecting(t). + Len(watcher.cmpOpts, 2, "SetCmpOpts() cmpOpts length = %d, want 2", len(watcher.cmpOpts)) } // TestResetCmpOpts verifies ResetCmpOpts clears the cmpOpts field. @@ -71,9 +68,7 @@ func TestResetCmpOpts(t *testing.T) { watcher.ResetCmpOpts() - if watcher.cmpOpts != nil { - t.Errorf("ResetCmpOpts() cmpOpts = %v, want nil", watcher.cmpOpts) - } + assert.NewCollecting(t).Nil(watcher.cmpOpts, "ResetCmpOpts() cmpOpts") } // TestWithTimeout verifies WithTimeout option. @@ -89,9 +84,7 @@ func TestWithTimeout(t *testing.T) { opt := WithTimeout(customTimeout) opt(watcher) - if watcher.timeout != customTimeout { - t.Errorf("WithTimeout() set timeout = %v, want %v", watcher.timeout, customTimeout) - } + assert.NewCollecting(t).Eq(customTimeout, watcher.timeout, "WithTimeout() set timeout") } // TestWithCmpOpts verifies WithCmpOpts option. @@ -107,9 +100,8 @@ func TestWithCmpOpts(t *testing.T) { option := WithCmpOpts(opts...) option(watcher) - if len(watcher.cmpOpts) != 2 { - t.Errorf("WithCmpOpts() set cmpOpts length = %d, want 2", len(watcher.cmpOpts)) - } + assert.NewCollecting(t). + Len(watcher.cmpOpts, 2, "WithCmpOpts() set cmpOpts length = %d, want 2", len(watcher.cmpOpts)) } // TestWithExtraResource verifies WithExtraResource option. @@ -126,12 +118,8 @@ func TestWithExtraResource(t *testing.T) { option := WithExtraResource(configMap, secret) option(watcher) - if len(watcher.extraResources) != 2 { - t.Errorf( - "WithExtraResource() set extraResources length = %d, want 2", - len(watcher.extraResources), - ) - } + assert.NewCollecting(t). + Len(watcher.extraResources, 2, "WithExtraResource() set extraResources length = %d, want 2", len(watcher.extraResources)) } // TestWithExtraResource_Duplicates verifies duplicate kinds are handled. @@ -153,12 +141,8 @@ func TestWithExtraResource_Duplicates(t *testing.T) { opt2(watcher) // Both are added to extraResources slice - if len(watcher.extraResources) != 2 { - t.Errorf( - "WithExtraResource() called twice set extraResources length = %d, want 2", - len(watcher.extraResources), - ) - } + assert.NewCollecting(t). + Len(watcher.extraResources, 2, "WithExtraResource() called twice set extraResources length = %d, want 2", len(watcher.extraResources)) // Note: watchResource() deduplicates by checking watchedKinds map, // so the second ConfigMap won't create duplicate event handlers. diff --git a/pkg/testutil/resource_watcher_listener_internal_test.go b/pkg/testutil/resource_watcher_listener_internal_test.go index 5b3a9232..4158acf9 100644 --- a/pkg/testutil/resource_watcher_listener_internal_test.go +++ b/pkg/testutil/resource_watcher_listener_internal_test.go @@ -6,11 +6,14 @@ import ( corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + + "github.com/multigres/testkit/assert" ) // TestSendEvent_ChannelFull tests the default branch when event channel is full. func TestSendEvent_ChannelFull(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) watcher := &ResourceWatcher{ t: t, @@ -29,12 +32,8 @@ func TestSendEvent_ChannelFull(t *testing.T) { // Verify first event is in channel evt := <-watcher.eventCh - if evt.Name != "test" { - t.Errorf("Event in channel has Name = %s, want test", evt.Name) - } - if evt.Type != "ADDED" { - t.Errorf("Event in channel has Type = %s, want ADDED", evt.Type) - } + c.Eq("test", evt.Name, "Event in channel has Name") + c.Eq("ADDED", evt.Type, "Event in channel has Type") // Second event was dropped (channel was full) select { @@ -48,6 +47,7 @@ func TestSendEvent_ChannelFull(t *testing.T) { // TestCollectEvents_SubscriberChannelFull tests that events are dropped when subscriber channel is full. func TestCollectEvents_SubscriberChannelFull(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) ctx := t.Context() @@ -79,16 +79,12 @@ func TestCollectEvents_SubscriberChannelFull(t *testing.T) { eventsCount := len(watcher.events) watcher.mu.RUnlock() - if eventsCount != 2 { - t.Errorf("watcher.events length = %d, want 2", eventsCount) - } + c.Eq(2, eventsCount, "watcher.events length") // Verify subscriber channel only has first event select { case evt := <-subCh: - if evt.Name != "svc1" { - t.Errorf("First event Name = %s, want svc1", evt.Name) - } + c.Eq("svc1", evt.Name, "First event Name") default: t.Error("Expected first event in subscriber channel") } @@ -131,10 +127,6 @@ func TestUnsubscribe_NotFound(t *testing.T) { // Should not panic watcher.unsubscribe(differentCh) - if len(watcher.subscribers) != 1 { - t.Errorf( - "unsubscribe() should not remove other channels, got %d subscribers", - len(watcher.subscribers), - ) - } + assert.NewCollecting(t). + Len(watcher.subscribers, 1, "unsubscribe() should not remove other channels, got %d subscribers", len(watcher.subscribers)) } diff --git a/pkg/testutil/resource_watcher_match_edge_cases_test.go b/pkg/testutil/resource_watcher_match_edge_cases_test.go index d718727a..656e5f30 100644 --- a/pkg/testutil/resource_watcher_match_edge_cases_test.go +++ b/pkg/testutil/resource_watcher_match_edge_cases_test.go @@ -17,6 +17,8 @@ import ( "k8s.io/apimachinery/pkg/util/intstr" "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) var testError = errors.New("test error for coverage") @@ -24,6 +26,7 @@ var testError = errors.New("test error for coverage") // TestVerboseDiffs_ExistingNonMatch tests verbose diff logging when existing event doesn't match. func TestVerboseDiffs_ExistingNonMatch(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -38,9 +41,7 @@ func TestVerboseDiffs_ExistingNonMatch(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "existing-svc", Namespace: "default"}, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := c.Create(ctx, svc); err != nil { - t.Fatalf("Failed to create Service: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc), "Failed to create Service") time.Sleep(100 * time.Millisecond) @@ -65,9 +66,7 @@ func TestVerboseDiffs_ExistingNonMatch(t *testing.T) { watcher.SetCmpOpts(IgnoreMetaRuntimeFields(), IgnoreServiceRuntimeFields()) err := watcher.WaitForMatch(expected) - if err == nil { - t.Error("Expected timeout error") - } + ck.Error(err, "Expected timeout error") } // TestVerboseDiffs_IncomingEvents tests verbose diff logging for incoming events. @@ -99,7 +98,9 @@ func TestVerboseDiffs_IncomingEvents(t *testing.T) { Selector: &metav1.LabelSelector{MatchLabels: map[string]string{"app": "test"}}, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{Labels: map[string]string{"app": "test"}}, - Spec: corev1.PodSpec{Containers: []corev1.Container{{Name: "nginx", Image: "nginx"}}}, + Spec: corev1.PodSpec{ + Containers: []corev1.Container{{Name: "nginx", Image: "nginx"}}, + }, }, }, } @@ -138,9 +139,7 @@ func TestVerboseDiffs_IncomingEvents(t *testing.T) { } err := watcher.WaitForMatch(expected) - if err == nil { - t.Error("Expected timeout") - } + assert.NewCollecting(t).Error(err, "Expected timeout") } // TestWatchResource_AlreadyWatched tests early return when kind already watched. @@ -194,7 +193,5 @@ func TestWaitForEventType_NonMatchingEvents(t *testing.T) { // Wait for DELETED event (will see ADDED but predicate returns false) _, err := watcher.WaitForEventType("Service", "DELETED") - if err == nil { - t.Error("Expected timeout error") - } + assert.NewCollecting(t).Error(err, "Expected timeout error") } diff --git a/pkg/testutil/resource_watcher_match_internal_test.go b/pkg/testutil/resource_watcher_match_internal_test.go index 06eb9cd2..aaf0a84f 100644 --- a/pkg/testutil/resource_watcher_match_internal_test.go +++ b/pkg/testutil/resource_watcher_match_internal_test.go @@ -11,6 +11,8 @@ import ( ctrlcache "sigs.k8s.io/controller-runtime/pkg/cache" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/manager" + + "github.com/multigres/testkit/assert" ) // mockTB implements testing.TB for capturing Fatal/Fatalf and Error/Errorf calls. @@ -94,9 +96,7 @@ func TestWatchResource_Error(t *testing.T) { watcher.watchResource(context.Background(), mockM, &corev1.Service{}) - if !mockTB.fatalCalled { - t.Error("t.Fatalf was not called") - } + assert.NewCollecting(t).True(mockTB.fatalCalled, "t.Fatalf was not called") } // mockInformer implements eventHandlerRegistrar for testing. @@ -141,14 +141,13 @@ func TestAddEventHandlerToInformer_Error(t *testing.T) { watcher.addEventHandlerToInformer(mockInformer, "Service") - if !mockTB.fatalCalled { - t.Error("t.Fatalf was not called") - } + assert.NewCollecting(t).True(mockTB.fatalCalled, "t.Fatalf was not called") } // TestWaitForEvent_Match tests waitForEvent when predicate matches. func TestWaitForEvent_Match(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) ch := make(chan ResourceEvent, 1) @@ -167,12 +166,8 @@ func TestWaitForEvent_Match(t *testing.T) { watcher := &ResourceWatcher{t: t} evt, err := watcher.waitForEvent(ch, deadline, predicate) - if err != nil { - t.Errorf("waitForEvent() error = %v, want nil", err) - } - if evt == nil || evt.Name != "test" { - t.Errorf("waitForEvent() evt.Name = %v, want test", evt) - } + c.NoError(err, "waitForEvent() error") + c.False(evt == nil || evt.Name != "test", "waitForEvent() evt.Name = %v, want test", evt) } // TestWaitForEvent_ChannelClosed tests waitForEvent when the channel is closed. @@ -191,9 +186,7 @@ func TestWaitForEvent_ChannelClosed(t *testing.T) { watcher := &ResourceWatcher{t: t} _, err := watcher.waitForEvent(ch, deadline, predicate) - if !errors.Is(err, context.Canceled) { - t.Errorf("waitForEvent() error = %v, want context.Canceled", err) - } + assert.NewCollecting(t).ErrorIs(err, context.Canceled, "waitForEvent() error") } // TestWaitForEvent_NoMatch tests waitForEvent when predicate never matches. @@ -218,9 +211,7 @@ func TestWaitForEvent_NoMatch(t *testing.T) { watcher := &ResourceWatcher{t: t} _, err := watcher.waitForEvent(ch, deadline, predicate) - if !errors.Is(err, context.DeadlineExceeded) { - t.Errorf("waitForEvent() error = %v, want DeadlineExceeded", err) - } + assert.NewCollecting(t).ErrorIs(err, context.DeadlineExceeded, "waitForEvent() error") } // TestWaitForEvent_TimeoutEdgeCase tests the timeout logic in waitForEvent. @@ -239,9 +230,8 @@ func TestWaitForEvent_TimeoutEdgeCase(t *testing.T) { watcher := &ResourceWatcher{t: t} _, err := watcher.waitForEvent(ch, deadline, predicate) - if !errors.Is(err, context.DeadlineExceeded) { - t.Errorf("waitForEvent() with past deadline should return DeadlineExceeded, got: %v", err) - } + assert.NewCollecting(t). + ErrorIs(err, context.DeadlineExceeded, "waitForEvent() with past deadline should return DeadlineExceeded, got") } // TestWaitForEvent_TimeoutBoundary tests the timeout boundary condition. @@ -266,9 +256,8 @@ func TestWaitForEvent_TimeoutBoundary(t *testing.T) { watcher := &ResourceWatcher{t: t} _, err := watcher.waitForEvent(ch, deadline, predicate) - if !errors.Is(err, context.DeadlineExceeded) { - t.Errorf("waitForEvent() should timeout, got: %v", err) - } + assert.NewCollecting(t). + ErrorIs(err, context.DeadlineExceeded, "waitForEvent() should timeout, got") } // TestWaitForMatch_EmptySlice tests WaitForMatch with empty slice. @@ -281,9 +270,7 @@ func TestWaitForMatch_EmptySlice(t *testing.T) { } err := watcher.WaitForMatch() - if err != nil { - t.Errorf("WaitForMatch() with empty slice should return nil, got: %v", err) - } + assert.NewCollecting(t).NoError(err, "WaitForMatch() with empty slice should return nil, got") } // TestExtractKind tests kind extraction from various object types. @@ -309,9 +296,7 @@ func TestExtractKind(t *testing.T) { t.Parallel() got := extractKind(tc.obj) - if got != tc.want { - t.Errorf("extractKind() = %s, want %s", got, tc.want) - } + assert.NewCollecting(t).Eq(tc.want, got, "extractKind()") }) } } @@ -339,9 +324,7 @@ func TestExtractKind_NoPointer(t *testing.T) { t.Parallel() got := extractKind(tc.obj) - if got != tc.want { - t.Errorf("extractKind() = %s, want %s", got, tc.want) - } + assert.NewCollecting(t).Eq(tc.want, got, "extractKind()") }) } } @@ -353,9 +336,7 @@ func TestExtractKind_NoDot(t *testing.T) { // Create a mock type that when formatted has no dot // In practice, all client.Object types have packages, but we can test the fallback kind := extractKind(&corev1.Node{}) - if kind != "Node" { - t.Errorf("extractKind() = %s, want Node", kind) - } + assert.NewCollecting(t).Eq("Node", kind, "extractKind()") } // TestExtractKind_NonPointer tests extractKind fallback for type without leading *. @@ -366,9 +347,7 @@ func TestExtractKind_NonPointer(t *testing.T) { // by ensuring we handle types correctly pod := corev1.Pod{} kind := extractKind(&pod) - if kind != "Pod" { - t.Errorf("extractKind() = %s, want Pod", kind) - } + assert.NewCollecting(t).Eq("Pod", kind, "extractKind()") } // TestExtractKind_FallbackPath tests the final return when no dot is found. @@ -387,8 +366,6 @@ func TestExtractKind_FallbackPath(t *testing.T) { for _, tc := range tests { got := extractKind(tc.obj) - if got != tc.want { - t.Errorf("extractKind(%T) = %s, want %s", tc.obj, got, tc.want) - } + assert.NewCollecting(t).Eq(tc.want, got, "extractKind(%T) = %s, want", tc.obj, got) } } diff --git a/pkg/testutil/resource_watcher_match_test.go b/pkg/testutil/resource_watcher_match_test.go index 8082029b..97f8f330 100644 --- a/pkg/testutil/resource_watcher_match_test.go +++ b/pkg/testutil/resource_watcher_match_test.go @@ -16,11 +16,14 @@ import ( "k8s.io/apimachinery/pkg/util/intstr" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) // TestResourceWatcher_UnwatchedKinds tests error for unwatched resource kinds. func TestResourceWatcher_UnwatchedKinds(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -35,14 +38,10 @@ func TestResourceWatcher_UnwatchedKinds(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "test", Namespace: "default"}, }) - if err == nil { - t.Error("WaitForMatch() should error for unwatched kind") - } + c.Error(err, "WaitForMatch() should error for unwatched kind") var unwatchedErr *testutil.ErrUnwatchedKinds - if !errors.As(err, &unwatchedErr) { - t.Errorf("Error should be ErrUnwatchedKinds, got: %T", err) - } + c.True(errors.As(err, &unwatchedErr), "Error should be ErrUnwatchedKinds, got: %T", err) if len(unwatchedErr.Kinds) != 1 || unwatchedErr.Kinds[0] != "ConfigMap" { t.Errorf("ErrUnwatchedKinds.Kinds = %v, want [ConfigMap]", unwatchedErr.Kinds) @@ -52,6 +51,7 @@ func TestResourceWatcher_UnwatchedKinds(t *testing.T) { // TestResourceWatcher_WatchDuplicateKind tests that watching the same kind twice doesn't create duplicate handlers. func TestResourceWatcher_WatchDuplicateKind(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -73,27 +73,27 @@ func TestResourceWatcher_WatchDuplicateKind(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "test-cm", Namespace: "default"}, Data: map[string]string{"key": "value"}, } - if err := c.Create(ctx, cm); err != nil { - t.Fatalf("Failed to create ConfigMap: %v", err) - } + ck.Require().NoError(c.Create(ctx, cm), "Failed to create ConfigMap") // Wait for the event watcher.SetCmpOpts(testutil.IgnoreMetaRuntimeFields()) - if err := watcher.WaitForMatch(cm); err != nil { - t.Errorf("Failed to wait for ConfigMap: %v", err) - } + ck.NoError(watcher.WaitForMatch(cm), "Failed to wait for ConfigMap") // Verify we got our ConfigMap event (there may be others from kube-system) // The key test is that duplicate handler registration was prevented by watchResource events := watcher.ForName("test-cm") - if len(events) != 1 { - t.Errorf("Expected 1 event for test-cm, got %d (duplicate handler may have been created)", len(events)) - } + ck.Len( + events, + 1, + "Expected 1 event for test-cm, got %d (duplicate handler may have been created)", + len(events), + ) } // TestResourceWatcher_NonMatchingUpdate tests that updates which don't match expected spec keep waiting. func TestResourceWatcher_NonMatchingUpdate(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -108,9 +108,7 @@ func TestResourceWatcher_NonMatchingUpdate(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "test-svc", Namespace: "default"}, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := c.Create(ctx, svc); err != nil { - t.Fatalf("Failed to create Service: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc), "Failed to create Service") // Start watcher after creation watcher := testutil.NewResourceWatcher(t, ctx, mgr, @@ -137,9 +135,7 @@ func TestResourceWatcher_NonMatchingUpdate(t *testing.T) { // This should timeout because the service has port 80, not 8080 err := watcher.WaitForMatch(expected) - if err == nil { - t.Error("Expected timeout error, got nil") - } + ck.Error(err, "Expected timeout error, got nil") // Error should be a timeout with diff information (we don't check exact message) if !errors.Is(err, context.DeadlineExceeded) { @@ -170,9 +166,7 @@ func TestResourceWatcher_NoEventsTimeout(t *testing.T) { } err := watcher.WaitForMatch(expected) - if err == nil { - t.Error("Expected timeout error, got nil") - } + assert.NewCollecting(t).Error(err, "Expected timeout error, got nil") // We just verify we got an error (timeout), not checking exact message t.Logf("Got expected timeout error: %v", err) @@ -181,6 +175,7 @@ func TestResourceWatcher_NoEventsTimeout(t *testing.T) { // TestWaitForEventType_ExistingEvent tests WaitForEventType finding existing event. func TestWaitForEventType_ExistingEvent(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -197,22 +192,14 @@ func TestWaitForEventType_ExistingEvent(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "test-svc-event", Namespace: "default"}, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := c.Create(ctx, svc); err != nil { - t.Fatalf("Failed to create Service: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc), "Failed to create Service") // Wait for ADDED event evt, err := watcher.WaitForEventType("Service", "ADDED") - if err != nil { - t.Fatalf("WaitForEventType() error = %v", err) - } + ck.Require().NoError(err, "WaitForEventType() error =") - if evt.Type != "ADDED" { - t.Errorf("Event type = %s, want ADDED", evt.Type) - } - if evt.Kind != "Service" { - t.Errorf("Event kind = %s, want Service", evt.Kind) - } + ck.Eq("ADDED", evt.Type, "Event type") + ck.Eq("Service", evt.Kind, "Event kind") } // TestWaitForEventType_Timeout tests WaitForEventType timeout. @@ -232,9 +219,7 @@ func TestWaitForEventType_Timeout(t *testing.T) { // Wait for an event type that won't happen _, err := watcher.WaitForEventType("Service", "DELETED") - if err == nil { - t.Error("Expected timeout error, got nil") - } + assert.NewCollecting(t).Error(err, "Expected timeout error, got nil") t.Logf("Got expected timeout: %v", err) } @@ -261,9 +246,7 @@ func TestWaitForMatch_ContextCanceled(t *testing.T) { Spec: corev1.ServiceSpec{Type: corev1.ServiceTypeClusterIP}, } err := watcher.WaitForMatch(svc) - if err == nil { - t.Error("Expected error when context is canceled") - } + assert.NewCollecting(t).Error(err, "Expected error when context is canceled") t.Logf("Got expected error: %v", err) } diff --git a/pkg/testutil/resource_watcher_test.go b/pkg/testutil/resource_watcher_test.go index e2b7d18f..442674aa 100644 --- a/pkg/testutil/resource_watcher_test.go +++ b/pkg/testutil/resource_watcher_test.go @@ -9,7 +9,6 @@ import ( "testing" "time" - "github.com/google/go-cmp/cmp" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" @@ -19,6 +18,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "github.com/multigres/multigres-operator/pkg/testutil" + + "github.com/multigres/testkit/assert" ) // TestResourceWatcher_BeforeCreation tests that watcher can subscribe to events @@ -72,9 +73,7 @@ func TestResourceWatcher_BeforeCreation(t *testing.T) { watcher.SetCmpOpts(opts...) err := watcher.WaitForMatch(expected) - if err != nil { - t.Errorf("Failed to wait for Service: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for Service") }, }, "statefulset created": { @@ -142,9 +141,7 @@ func TestResourceWatcher_BeforeCreation(t *testing.T) { watcher.SetCmpOpts(opts...) err := watcher.WaitForMatch(expected) - if err != nil { - t.Errorf("Failed to wait for StatefulSet: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for StatefulSet") }, }, "deployment created": { @@ -213,9 +210,7 @@ func TestResourceWatcher_BeforeCreation(t *testing.T) { watcher.SetCmpOpts(opts...) err := watcher.WaitForMatch(expected) - if err != nil { - t.Errorf("Failed to wait for Deployment: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for Deployment") }, }, "multiple unwatched kinds fail immediately": { @@ -238,9 +233,7 @@ func TestResourceWatcher_BeforeCreation(t *testing.T) { return } - if diff := cmp.Diff(want.Kinds, got.Kinds); diff != "" { - t.Errorf("Unwatched kinds mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(want.Kinds, got.Kinds, "Unwatched kinds mismatch") }, }, } @@ -268,9 +261,7 @@ func TestResourceWatcher_BeforeCreation(t *testing.T) { }() // THEN create resources (tests subscription path) - if err := tc.setup(ctx, c); err != nil { - t.Fatalf("Setup failed: %v", err) - } + assert.NewAborting(t).NoError(tc.setup(ctx, c), "Setup failed") // Wait for assertion to complete <-done @@ -341,9 +332,7 @@ func TestResourceWatcher_AfterCreation(t *testing.T) { watcher.SetCmpOpts(opts...) err := watcher.WaitForMatch(expected) - if err != nil { - t.Errorf("Failed to wait for updated Service: %v", err) - } + assert.NewCollecting(t).NoError(err, "Failed to wait for updated Service") }, }, "service deleted after watcher starts": { @@ -377,9 +366,8 @@ func TestResourceWatcher_AfterCreation(t *testing.T) { return } - if evt.Name != "test-svc-delete" { - t.Errorf("Expected test-svc-delete, got %s", evt.Name) - } + assert.NewCollecting(t). + Eq("test-svc-delete", evt.Name, "Expected test-svc-delete, got") t.Logf("Successfully detected Service deletion") }, @@ -390,6 +378,7 @@ func TestResourceWatcher_AfterCreation(t *testing.T) { name, tc := name, tc t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -400,9 +389,7 @@ func TestResourceWatcher_AfterCreation(t *testing.T) { c := mgr.GetClient() // Create resources FIRST - if err := tc.setup(ctx, c); err != nil { - t.Fatalf("Setup failed: %v", err) - } + ck.NoError(tc.setup(ctx, c), "Setup failed") // THEN start watcher (won't see initial creation) watcher := testutil.NewResourceWatcher(t, ctx, mgr) @@ -415,9 +402,7 @@ func TestResourceWatcher_AfterCreation(t *testing.T) { }() // Trigger update (this should be picked up by watcher) - if err := tc.update(ctx, c); err != nil { - t.Fatalf("Update failed: %v", err) - } + ck.NoError(tc.update(ctx, c), "Update failed") // Wait for assertion to complete <-done @@ -428,6 +413,7 @@ func TestResourceWatcher_AfterCreation(t *testing.T) { // TestResourceWatcherEventUtilities tests event inspection utility functions. func TestResourceWatcherEventUtilities(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) scheme := runtime.NewScheme() _ = corev1.AddToScheme(scheme) @@ -458,20 +444,16 @@ func TestResourceWatcherEventUtilities(t *testing.T) { Selector: &metav1.LabelSelector{MatchLabels: map[string]string{"app": "test"}}, Template: corev1.PodTemplateSpec{ ObjectMeta: metav1.ObjectMeta{Labels: map[string]string{"app": "test"}}, - Spec: corev1.PodSpec{Containers: []corev1.Container{{Name: "nginx", Image: "nginx"}}}, + Spec: corev1.PodSpec{ + Containers: []corev1.Container{{Name: "nginx", Image: "nginx"}}, + }, }, }, } - if err := c.Create(ctx, svc1); err != nil { - t.Fatalf("Failed to create svc-1: %v", err) - } - if err := c.Create(ctx, svc2); err != nil { - t.Fatalf("Failed to create svc-2: %v", err) - } - if err := c.Create(ctx, deploy); err != nil { - t.Fatalf("Failed to create deploy-1: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc1), "Failed to create svc-1") + ck.Require().NoError(c.Create(ctx, svc2), "Failed to create svc-2") + ck.Require().NoError(c.Create(ctx, deploy), "Failed to create deploy-1") // Wait for all resources watcher.SetCmpOpts( @@ -481,15 +463,11 @@ func TestResourceWatcherEventUtilities(t *testing.T) { testutil.IgnoreDeploymentSpecDefaults(), testutil.IgnorePodSpecDefaults(), ) - if err := watcher.WaitForMatch(svc1, svc2, deploy); err != nil { - t.Fatalf("Failed to wait for resources: %v", err) - } + ck.Require().NoError(watcher.WaitForMatch(svc1, svc2, deploy), "Failed to wait for resources") // Test Events() returns all collected events events := watcher.Events() - if len(events) < 3 { - t.Errorf("Events() = %d events, want at least 3", len(events)) - } + ck.GreaterOrEqual(3, len(events), "Events()") // Verify Events() contains expected resources foundSvc1, foundSvc2, foundDeploy := false, false, false @@ -504,65 +482,41 @@ func TestResourceWatcherEventUtilities(t *testing.T) { foundDeploy = true } } - if !foundSvc1 { - t.Error("Events() should contain ADDED event for svc-1") - } - if !foundSvc2 { - t.Error("Events() should contain ADDED event for svc-2") - } - if !foundDeploy { - t.Error("Events() should contain ADDED event for deploy-1") - } + ck.True(foundSvc1, "Events() should contain ADDED event for svc-1") + ck.True(foundSvc2, "Events() should contain ADDED event for svc-2") + ck.True(foundDeploy, "Events() should contain ADDED event for deploy-1") // Test ForKind() filters by resource kind svcEvents := watcher.ForKind("Service") - if len(svcEvents) < 2 { - t.Errorf("ForKind(Service) = %d events, want at least 2 (svc-1 and svc-2)", len(svcEvents)) - } + ck.GreaterOrEqual(2, len(svcEvents), "ForKind(Service)") for _, evt := range svcEvents { - if evt.Kind != "Service" { - t.Errorf("ForKind(Service) returned event with Kind = %s", evt.Kind) - } + ck.Eq("Service", evt.Kind, "ForKind(Service) returned event with Kind =") } deployEvents := watcher.ForKind("Deployment") - if len(deployEvents) < 1 { - t.Errorf("ForKind(Deployment) = %d events, want at least 1", len(deployEvents)) - } + ck.GreaterOrEqual(1, len(deployEvents), "ForKind(Deployment)") for _, evt := range deployEvents { - if evt.Kind != "Deployment" { - t.Errorf("ForKind(Deployment) returned event with Kind = %s", evt.Kind) - } + ck.Eq("Deployment", evt.Kind, "ForKind(Deployment) returned event with Kind =") } // Test ForName() filters by resource name svc2Events := watcher.ForName("svc-2") - if len(svc2Events) == 0 { - t.Error("ForName(svc-2) returned no events") - } + ck.NotEmpty(svc2Events, "ForName(svc-2) returned no events") for _, evt := range svc2Events { - if evt.Name != "svc-2" { - t.Errorf("ForName(svc-2) returned event with Name = %s", evt.Name) - } + ck.Eq("svc-2", evt.Name, "ForName(svc-2) returned event with Name =") } // Test Count() matches Events() length count := watcher.Count() - if count != len(events) { - t.Errorf("Count() = %d, len(Events()) = %d, should be equal", count, len(events)) - } + ck.Eq(len(events), count, "Count()") // Test EventCh() provides direct access to event channel ch := watcher.EventCh() - if ch == nil { - t.Fatal("EventCh() returned nil") - } + ck.Require().NotNil(ch, "EventCh() returned nil") // Verify EventCh returns same channel on multiple calls ch2 := watcher.EventCh() - if ch != ch2 { - t.Error("EventCh() should return the same channel instance") - } + ck.Eq(ch2, ch, "EventCh() should return the same channel instance") // Create another resource and verify the event count increases // (can't consume from channel as that would interfere with collectEvents) @@ -572,21 +526,15 @@ func TestResourceWatcherEventUtilities(t *testing.T) { ObjectMeta: metav1.ObjectMeta{Name: "svc-3", Namespace: "default"}, Spec: corev1.ServiceSpec{Ports: []corev1.ServicePort{{Port: 80}}}, } - if err := c.Create(ctx, svc3); err != nil { - t.Fatalf("Failed to create svc-3: %v", err) - } + ck.Require().NoError(c.Create(ctx, svc3), "Failed to create svc-3") // Wait for event to be collected watcher.SetCmpOpts(testutil.IgnoreMetaRuntimeFields(), testutil.IgnoreServiceRuntimeFields()) - if err := watcher.WaitForMatch(svc3); err != nil { - t.Fatalf("Failed to wait for svc-3: %v", err) - } + ck.Require().NoError(watcher.WaitForMatch(svc3), "Failed to wait for svc-3") // Verify event was collected through the channel countAfter := watcher.Count() - if countAfter <= countBefore { - t.Errorf("Count after creating svc-3: %d, should be > %d (events collected via EventCh)", countAfter, countBefore) - } + ck.Greater(countBefore, countAfter, "Count after creating svc-3") // Verify the new event is in Events() foundSvc3 := false @@ -596,9 +544,7 @@ func TestResourceWatcherEventUtilities(t *testing.T) { break } } - if !foundSvc3 { - t.Error("Events() should contain ADDED event for svc-3 (received via EventCh)") - } + ck.True(foundSvc3, "Events() should contain ADDED event for svc-3 (received via EventCh)") } // TestErrUnwatchedKinds_Error tests the Error() method. @@ -610,9 +556,7 @@ func TestErrUnwatchedKinds_Error(t *testing.T) { got := err.Error() want := "the following kinds are not being watched by this ResourceWatcher: [ConfigMap Secret]" - if got != want { - t.Errorf("Error() = %q, want %q", got, want) - } + assert.NewCollecting(t).Eq(want, got, "Error()") } // TestResourceWatcher_ContextCancellation tests behavior when context is canceled. @@ -659,9 +603,7 @@ func TestResourceWatcher_ContextCancellation(t *testing.T) { cancel() // Cancel context to trigger error paths err := tc.testFunc(t, watcher) - if err == nil { - t.Error("Function should error when context is canceled") - } + assert.NewCollecting(t).Error(err, "Function should error when context is canceled") }) } } @@ -688,7 +630,9 @@ func TestResourceWatcher_Timeouts(t *testing.T) { }, "WaitForDeletion timeout": { testFunc: func(t *testing.T, watcher *testutil.ResourceWatcher) error { - return watcher.WaitForDeletion(testutil.Obj[corev1.Service]("nonexistent", "default")) + return watcher.WaitForDeletion( + testutil.Obj[corev1.Service]("nonexistent", "default"), + ) }, }, "WaitForEventType timeout": { @@ -710,9 +654,7 @@ func TestResourceWatcher_Timeouts(t *testing.T) { ) err := tc.testFunc(t, watcher) - if err == nil { - t.Error("Function should timeout") - } + assert.NewCollecting(t).Error(err, "Function should timeout") }) } } diff --git a/pkg/topology/roots_test.go b/pkg/topology/roots_test.go index 4ee69406..7d815175 100644 --- a/pkg/topology/roots_test.go +++ b/pkg/topology/roots_test.go @@ -4,17 +4,18 @@ import ( "strings" "testing" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/utils/ptr" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/certs" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestExternalTopologyKeepsLongRoot(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "cluster-abcdefghijklmnop", @@ -28,9 +29,8 @@ func TestExternalTopologyKeepsLongRoot(t *testing.T) { }, } roots, err := ForCluster(cluster) - require.NoError(t, err) - assert.Equal( - t, + c.Require().NoError(err) + c.EqDeep( "/multigres/namespace-abcdefghijklmnopqrstu/cluster-abcdefghijklmnop", roots.ClusterRoot(), ) @@ -38,6 +38,7 @@ func TestExternalTopologyKeepsLongRoot(t *testing.T) { func TestRootsWithTopologyTLS(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) const namespace = "namespace-abcdefghijklmnopqrstu" const clusterName = "cluster-abcdefghijklmnop" @@ -56,16 +57,17 @@ func TestRootsWithTopologyTLS(t *testing.T) { {"64 byte TLS fallback is unchanged", namespace, clusterName[:21], true, unboundedRoot[:64]}, } { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) roots, err := NewRoots(nil, tc.namespace, tc.clusterName, tc.topoTLS) - require.NoError(t, err) - assert.Equal(t, tc.want, roots.ClusterRoot()) - assert.Equal(t, tc.want+"/global", roots.Global()) + c.Require().NoError(err) + c.EqDeep(tc.want, roots.ClusterRoot()) + c.EqDeep(tc.want+"/global", roots.Global()) cellRoot, err := roots.Cell("zone-a") - require.NoError(t, err) - assert.Equal(t, tc.want+"/zone-a", cellRoot) - assert.True(t, strings.HasPrefix(cellRoot, roots.KeyPrefix())) + c.Require().NoError(err) + c.EqDeep(tc.want+"/zone-a", cellRoot) + c.True(strings.HasPrefix(cellRoot, roots.KeyPrefix())) if tc.topoTLS { - assert.LessOrEqual(t, len(roots.ClusterRoot()), certs.MaxCommonNameBytes) + c.LessOrEqual(certs.MaxCommonNameBytes, len(roots.ClusterRoot())) } }) } @@ -73,12 +75,12 @@ func TestRootsWithTopologyTLS(t *testing.T) { for _, ref := range []string{"~", "namespace-abcdefghijklmnopqrstu", "../multigres-fallback", strings.Repeat("p", 64)} { annotations := map[string]string{metadata.AnnotationProjectRef: ref} plain, err := NewRoots(annotations, namespace, clusterName, false) - require.NoError(t, err) + c.Require().NoError(err) tls, err := NewRoots(annotations, namespace, clusterName, true) - require.NoError(t, err) - assert.Equal(t, plain, tls, "explicit refs must not change with TLS") - assert.NotEqual(t, hashedRoot, tls.ClusterRoot()) - assert.False(t, strings.HasPrefix(hashedRoot+"/global", tls.KeyPrefix())) + c.Require().NoError(err) + c.EqDeep(plain, tls, "explicit refs must not change with TLS") + c.NotEqDeep(hashedRoot, tls.ClusterRoot()) + c.False(strings.HasPrefix(hashedRoot+"/global", tls.KeyPrefix())) } seen := map[string]bool{hashedRoot: true} @@ -90,9 +92,9 @@ func TestRootsWithTopologyTLS(t *testing.T) { {strings.Repeat("n", 63), strings.Repeat("c", 253)}, } { roots, err := NewRoots(nil, pair[0], pair[1], true) - require.NoError(t, err) - assert.LessOrEqual(t, len(roots.ClusterRoot()), certs.MaxCommonNameBytes) - assert.False(t, seen[roots.ClusterRoot()], "fallback identities must be distinct") + c.Require().NoError(err) + c.LessOrEqual(certs.MaxCommonNameBytes, len(roots.ClusterRoot())) + c.False(seen[roots.ClusterRoot()], "fallback identities must be distinct") seen[roots.ClusterRoot()] = true } } @@ -140,47 +142,35 @@ func TestRoots(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) roots, err := NewRoots(tc.annotations, tc.namespace, tc.clusterName, false) if tc.wantErr { - if err == nil { - t.Fatal("expected an error") - } + c.Require().Error(err, "expected an error") return } - if err != nil { - t.Fatalf("NewRoots: %v", err) - } + c.Require().NoError(err, "NewRoots") cell, err := roots.Cell(tc.cellName) - if err != nil { - t.Fatalf("Cell: %v", err) - } - if got := roots.Global(); got != tc.wantGlobal { - t.Errorf("Global() = %q, want %q", got, tc.wantGlobal) - } - if cell != tc.wantCell { - t.Errorf("Cell() = %q, want %q", cell, tc.wantCell) - } - if roots.Global() == cell { - t.Error("global and cell roots must be disjoint") - } + c.Require().NoError(err, "Cell") + c.Eq(tc.wantGlobal, roots.Global(), "Global()") + c.Eq(tc.wantCell, cell, "Cell()") + c.NotEq(cell, roots.Global(), "global and cell roots must be disjoint") }) } } func TestFallbackIdentityIsNamespaceScoped(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) first, err := NewRoots(nil, "tenant-a", "production", false) - if err != nil { - t.Fatal(err) - } + c.NoError(err) second, err := NewRoots(nil, "tenant-b", "production", false) - if err != nil { - t.Fatal(err) - } - if first.Global() == second.Global() { - t.Fatalf("equal cluster names in different namespaces collided at %q", first.Global()) - } + c.NoError(err) + c.NotEq( + second.Global(), + first.Global(), + "equal cluster names in different namespaces collided at", + ) } func TestClusterRootPrefixesEveryRoot(t *testing.T) { @@ -205,23 +195,14 @@ func TestClusterRootPrefixesEveryRoot(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { + c := assert.NewCollecting(t) roots, err := NewRoots(tc.annotations, tc.namespace, tc.clusterName, false) - if err != nil { - t.Fatalf("NewRoots() error = %v", err) - } - if got := roots.ClusterRoot(); got != tc.want { - t.Errorf("ClusterRoot() = %q, want %q", got, tc.want) - } - if want := roots.ClusterRoot() + "/global"; roots.Global() != want { - t.Errorf("Global() = %q, want %q", roots.Global(), want) - } + c.Require().NoError(err, "NewRoots() error =") + c.Eq(tc.want, roots.ClusterRoot(), "ClusterRoot()") + c.Eq(roots.ClusterRoot()+"/global", roots.Global(), "Global()") cell, err := roots.Cell("zone-a") - if err != nil { - t.Fatalf("Cell() error = %v", err) - } - if want := roots.ClusterRoot() + "/zone-a"; cell != want { - t.Errorf("Cell() = %q, want %q", cell, want) - } + c.Require().NoError(err, "Cell() error =") + c.Eq(roots.ClusterRoot()+"/zone-a", cell, "Cell()") }) } } @@ -230,18 +211,15 @@ func TestClusterRootPrefixesEveryRoot(t *testing.T) { // merely starts with the same characters, so authorization grants on the // prefix that includes the separator. func TestKeyPrefixDoesNotReachSiblingClusters(t *testing.T) { + c := assert.NewAborting(t) short, err := NewRoots( map[string]string{metadata.AnnotationProjectRef: "proj_123"}, "supabase", "a", false, ) - if err != nil { - t.Fatalf("NewRoots() error = %v", err) - } + c.NoError(err, "NewRoots() error =") long, err := NewRoots( map[string]string{metadata.AnnotationProjectRef: "proj_1234"}, "supabase", "b", false, ) - if err != nil { - t.Fatalf("NewRoots() error = %v", err) - } + c.NoError(err, "NewRoots() error =") if !strings.HasPrefix(long.ClusterRoot(), short.ClusterRoot()) { t.Fatal("fixtures no longer exercise the sibling prefix hazard") diff --git a/pkg/util/certs/certs_test.go b/pkg/util/certs/certs_test.go index 17b9274e..e38855ea 100644 --- a/pkg/util/certs/certs_test.go +++ b/pkg/util/certs/certs_test.go @@ -14,18 +14,19 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func testScheme(t *testing.T) *runtime.Scheme { t.Helper() scheme := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(scheme); err != nil { - t.Fatalf("AddToScheme() error = %v", err) - } + assert.NewAborting(t).NoError(multigresv1alpha1.AddToScheme(scheme), "AddToScheme() error =") return scheme } func TestBuildDefaultsIssuer(t *testing.T) { + c := assert.NewCollecting(t) owner := &multigresv1alpha1.TopoServer{ ObjectMeta: metav1.ObjectMeta{ Name: "owner", @@ -41,47 +42,30 @@ func TestBuildDefaultsIssuer(t *testing.T) { DNSNames: []any{"example"}, Usages: []any{"server auth"}, }) - if err != nil { - t.Fatalf("Build() error = %v", err) - } + c.Require().NoError(err, "Build() error =") - if cert.GetNamespace() != "supabase" { - t.Errorf("namespace = %q, want supabase", cert.GetNamespace()) - } + c.Eq("supabase", cert.GetNamespace(), "namespace") spec, ok := cert.Object["spec"].(map[string]any) - if !ok { - t.Fatal("spec is not a map") - } + c.Require().True(ok, "spec is not a map") wantIssuerRef := map[string]any{ "name": DefaultIssuerName, "kind": "ClusterIssuer", "group": "cert-manager.io", } - if diff := cmp.Diff(wantIssuerRef, spec["issuerRef"]); diff != "" { - t.Errorf("issuerRef mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff(Duration, spec["duration"]); diff != "" { - t.Errorf("duration mismatch (-want +got):\n%s", diff) - } + c.Eq("", cmp.Diff(wantIssuerRef, spec["issuerRef"]), "issuerRef mismatch (-want +got):\n") + c.Eq("", cmp.Diff(Duration, spec["duration"]), "duration mismatch (-want +got):\n") } func TestTruncateCommonName(t *testing.T) { + c := assert.NewCollecting(t) short := strings.Repeat("a", MaxCommonNameBytes) - if got := TruncateCommonName(short); got != short { - t.Errorf("TruncateCommonName() shortened a name that fits: %q", got) - } + c.Eq(short, TruncateCommonName(short), "TruncateCommonName() shortened a name that fits") long := strings.Repeat("a", MaxCommonNameBytes+40) got := TruncateCommonName(long) - if len(got) > MaxCommonNameBytes { - t.Errorf("TruncateCommonName() = %d bytes, want <= %d", len(got), MaxCommonNameBytes) - } - if got != TruncateCommonName(long) { - t.Error("TruncateCommonName() is not deterministic") - } - if got == TruncateCommonName(long+"b") { - t.Error("TruncateCommonName() collided for different inputs") - } + c.LessOrEqual(MaxCommonNameBytes, len(got), "TruncateCommonName()") + c.Eq(TruncateCommonName(long), got, "TruncateCommonName() is not deterministic") + c.NotEq(TruncateCommonName(long+"b"), got, "TruncateCommonName() collided for different inputs") } func certFixture(t *testing.T, name, secretName string) *unstructured.Unstructured { @@ -100,9 +84,7 @@ func certFixture(t *testing.T, name, secretName string) *unstructured.Unstructur DNSNames: []any{name}, Usages: []any{"server auth"}, }) - if err != nil { - t.Fatalf("Build() error = %v", err) - } + assert.NewAborting(t).NoError(err, "Build() error =") return cert } @@ -118,59 +100,47 @@ func fakeClient(t *testing.T, objs ...client.Object) client.Client { } func TestKeepSets(t *testing.T) { + c := assert.NewCollecting(t) desired := []*unstructured.Unstructured{ certFixture(t, "a", "a-secret"), certFixture(t, "b", "b-secret"), } keepNames, keepSecretNames := KeepSets(desired) - if diff := cmp.Diff( - map[string]struct{}{"a": {}, "b": {}}, keepNames, - ); diff != "" { - t.Errorf("keepNames mismatch (-want +got):\n%s", diff) - } - if diff := cmp.Diff( - map[string]struct{}{"a-secret": {}, "b-secret": {}}, keepSecretNames, - ); diff != "" { - t.Errorf("keepSecretNames mismatch (-want +got):\n%s", diff) - } + c.EqDiff(map[string]struct{}{"a": {}, "b": {}}, keepNames, "keepNames mismatch") + c.EqDiff( + map[string]struct{}{"a-secret": {}, "b-secret": {}}, + keepSecretNames, + "keepSecretNames mismatch", + ) } func TestListTolerantOfMissingCRD(t *testing.T) { + ck := assert.NewCollecting(t) // A scheme without the Certificate type makes the client report no mapping // for the GVK, which is what a cluster without cert-manager looks like. scheme := testScheme(t) c := fake.NewClientBuilder().WithScheme(scheme).Build() got, err := List(context.Background(), c, "supabase") - if err != nil { - t.Fatalf("List() error = %v, want nil when cert-manager is absent", err) - } - if len(got.Items) != 0 { - t.Errorf("got %d Certificates, want 0", len(got.Items)) - } + ck.Require().NoError(err, "List() error") + ck.Empty(got.Items, "got %d Certificates, want 0", len(got.Items)) } func TestFindByNameAndOwnedBy(t *testing.T) { + c := assert.NewCollecting(t) certList := &unstructured.UnstructuredList{} certList.SetGroupVersionKind(GVK) certList.Items = []unstructured.Unstructured{*certFixture(t, "a", "a-secret")} - if got := FindByName(certList, "a"); got == nil { - t.Error("FindByName(a) = nil, want the Certificate") - } - if got := FindByName(certList, "missing"); got != nil { - t.Errorf("FindByName(missing) = %v, want nil", got) - } - if !OwnedBy(&certList.Items[0], "owner-uid") { - t.Error("OwnedBy(owner-uid) = false, want true") - } - if OwnedBy(&certList.Items[0], "other-uid") { - t.Error("OwnedBy(other-uid) = true, want false") - } + c.NotNil(FindByName(certList, "a"), "FindByName(a) = nil, want the Certificate") + c.Nil(FindByName(certList, "missing"), "FindByName(missing)") + c.True(OwnedBy(&certList.Items[0], "owner-uid"), "OwnedBy(owner-uid) = false, want true") + c.False(OwnedBy(&certList.Items[0], "other-uid"), "OwnedBy(other-uid) = true, want false") } func TestPruneDeletesUnwantedCertificatesAndSecrets(t *testing.T) { + ck := assert.NewCollecting(t) stale := certFixture(t, "stale", "stale-secret") kept := certFixture(t, "kept", "kept-secret") unowned := certFixture(t, "unowned", "unowned-secret") @@ -182,16 +152,12 @@ func TestPruneDeletesUnwantedCertificatesAndSecrets(t *testing.T) { c := fakeClient(t, stale, kept, unowned, staleSecret) certList, err := List(context.Background(), c, "supabase") - if err != nil { - t.Fatalf("List() error = %v", err) - } + ck.Require().NoError(err, "List() error =") keepNames, keepSecretNames := KeepSets([]*unstructured.Unstructured{kept}) - if err := Prune( + ck.Require().NoError(Prune( context.Background(), c, "supabase", "owner-uid", certList, keepNames, keepSecretNames, - ); err != nil { - t.Fatalf("Prune() error = %v", err) - } + ), "Prune() error =") for name, wantGone := range map[string]bool{ "stale": true, @@ -205,59 +171,44 @@ func TestPruneDeletesUnwantedCertificatesAndSecrets(t *testing.T) { client.ObjectKey{Namespace: "supabase", Name: name}, got, ) - if wantGone && err == nil { - t.Errorf("Certificate %q still exists, want deleted", name) - } - if !wantGone && err != nil { - t.Errorf("Certificate %q was deleted, want kept: %v", name, err) - } + ck.False(wantGone && err == nil, "Certificate %q still exists, want deleted", name) + ck.False(!wantGone && err != nil, "Certificate %q was deleted, want kept: %v", name, err) } - if err := c.Get( + ck.Error(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "stale-secret"}, &corev1.Secret{}, - ); err == nil { - t.Error("stale Secret still exists, want deleted") - } + ), "stale Secret still exists, want deleted") } func TestApplySkipsUnchangedCertificates(t *testing.T) { + ck := assert.NewCollecting(t) existing := certFixture(t, "a", "a-secret") c := fakeClient(t, existing) certList, err := List(context.Background(), c, "supabase") - if err != nil { - t.Fatalf("List() error = %v", err) - } + ck.Require().NoError(err, "List() error =") before := certList.Items[0].GetResourceVersion() desired := certFixture(t, "a", "a-secret") - if err := Apply( + ck.Require().NoError(Apply( context.Background(), c, certList, "owner-uid", []*unstructured.Unstructured{desired}, - ); err != nil { - t.Fatalf("Apply() error = %v", err) - } + ), "Apply() error =") got := &unstructured.Unstructured{} got.SetGroupVersionKind(GVK) - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "a"}, got, - ); err != nil { - t.Fatalf("Get() error = %v", err) - } - if got.GetResourceVersion() != before { - t.Errorf( - "resourceVersion changed on a no-op apply: %q -> %q", - before, got.GetResourceVersion(), - ) - } + ), "Get() error =") + ck.Eq(before, got.GetResourceVersion(), "resourceVersion changed on a no-op apply") } func TestApplyRejectsForeignCertificate(t *testing.T) { + ck := assert.NewCollecting(t) existing := certFixture(t, "a", "a-secret") existing.SetOwnerReferences([]metav1.OwnerReference{{ APIVersion: "example.com/v1", @@ -268,9 +219,7 @@ func TestApplyRejectsForeignCertificate(t *testing.T) { c := fakeClient(t, existing) certList, err := List(context.Background(), c, "supabase") - if err != nil { - t.Fatalf("List() error = %v", err) - } + ck.Require().NoError(err, "List() error =") before := certList.Items[0].DeepCopy() desired := certFixture(t, "a", "a-secret") @@ -278,154 +227,131 @@ func TestApplyRejectsForeignCertificate(t *testing.T) { context.Background(), c, certList, "owner-uid", []*unstructured.Unstructured{desired}, ) - if err == nil { - t.Fatal("Apply() error = nil, want collision error") - } - if !strings.Contains(err.Error(), "a") || !strings.Contains(err.Error(), "supabase") { - t.Errorf("Apply() error = %q, want namespace and name", err) - } + ck.Require().Error(err, "Apply() error = nil, want collision error") + ck.False( + !strings.Contains(err.Error(), "a") || !strings.Contains(err.Error(), "supabase"), + "Apply() error = %q, want namespace and name", + err, + ) got := &unstructured.Unstructured{} got.SetGroupVersionKind(GVK) - if err := c.Get( + ck.Require().NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "a"}, got, - ); err != nil { - t.Fatalf("foreign Certificate was modified or deleted: %v", err) - } - if diff := cmp.Diff(before.Object, got.Object); diff != "" { - t.Errorf("foreign Certificate changed (-want +got):\n%s", diff) - } + ), "foreign Certificate was modified or deleted") + ck.EqDiff(before.Object, got.Object, "foreign Certificate changed") } func TestApplyCreatesMissingCertificates(t *testing.T) { + ck := assert.NewAborting(t) c := fakeClient(t) certList, err := List(context.Background(), c, "supabase") - if err != nil { - t.Fatalf("List() error = %v", err) - } + ck.NoError(err, "List() error =") desired := certFixture(t, "a", "a-secret") - if err := Apply( + ck.NoError(Apply( context.Background(), c, certList, "owner-uid", []*unstructured.Unstructured{desired}, - ); err != nil { - t.Fatalf("Apply() error = %v", err) - } + ), "Apply() error =") got := &unstructured.Unstructured{} got.SetGroupVersionKind(GVK) - if err := c.Get( + ck.NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "a"}, got, - ); err != nil { - t.Fatalf("expected Certificate to be created, got error %v", err) - } + ), "expected Certificate to be created, got error") } func TestGetAbsentAndMissingCRD(t *testing.T) { t.Run("absent certificate", func(t *testing.T) { + ck := assert.NewCollecting(t) c := fakeClient(t) got, err := Get(context.Background(), c, "supabase", "a") - if err != nil { - t.Fatalf("Get() error = %v, want nil", err) - } - if got != nil { - t.Errorf("Get() = %v, want nil", got) - } + ck.Require().NoError(err, "Get() error") + ck.Nil(got, "Get()") }) t.Run("cert-manager absent", func(t *testing.T) { + ck := assert.NewCollecting(t) // A scheme without the Certificate type is what a cluster with no // cert-manager CRD looks like to the client. c := fake.NewClientBuilder().WithScheme(testScheme(t)).Build() got, err := Get(context.Background(), c, "supabase", "a") - if err != nil { - t.Fatalf("Get() error = %v, want nil when cert-manager is absent", err) - } - if got != nil { - t.Errorf("Get() = %v, want nil", got) - } + ck.Require().NoError(err, "Get() error") + ck.Nil(got, "Get()") }) t.Run("present certificate", func(t *testing.T) { + ck := assert.NewAborting(t) c := fakeClient(t, certFixture(t, "a", "a-secret")) got, err := Get(context.Background(), c, "supabase", "a") - if err != nil { - t.Fatalf("Get() error = %v", err) - } - if got == nil || got.GetName() != "a" { - t.Fatalf("Get() = %v, want the Certificate named a", got) - } + ck.NoError(err, "Get() error =") + ck.False( + got == nil || got.GetName() != "a", + "Get() = %v, want the Certificate named a", + got, + ) }) } func TestDeleteRemovesCertificateAndSecret(t *testing.T) { + ck := assert.NewCollecting(t) cert := certFixture(t, "a", "a-secret") secret := &corev1.Secret{ ObjectMeta: metav1.ObjectMeta{Name: "a-secret", Namespace: "supabase"}, } c := fakeClient(t, cert, secret) - if err := Delete(context.Background(), c, cert); err != nil { - t.Fatalf("Delete() error = %v", err) - } + ck.Require().NoError(Delete(context.Background(), c, cert), "Delete() error =") got := &unstructured.Unstructured{} got.SetGroupVersionKind(GVK) - if err := c.Get( + ck.Error(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "a"}, got, - ); err == nil { - t.Error("Certificate still exists, want deleted") - } - if err := c.Get( + ), "Certificate still exists, want deleted") + ck.Error(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "a-secret"}, &corev1.Secret{}, - ); err == nil { - t.Error("Secret still exists, want deleted") - } + ), "Secret still exists, want deleted") } func TestDeleteIsIdempotent(t *testing.T) { cert := certFixture(t, "a", "a-secret") c := fakeClient(t) - if err := Delete(context.Background(), c, cert); err != nil { - t.Errorf("Delete() on absent objects error = %v, want nil", err) - } + assert.NewCollecting(t). + NoError(Delete(context.Background(), c, cert), "Delete() on absent objects error") } func TestSpecEqual(t *testing.T) { + c := assert.NewCollecting(t) a := certFixture(t, "a", "a-secret") same := certFixture(t, "a", "a-secret") other := certFixture(t, "a", "b-secret") - if !SpecEqual(a, same) { - t.Error("SpecEqual() = false for identical specs") - } - if SpecEqual(a, other) { - t.Error("SpecEqual() = true for differing specs") - } + c.True(SpecEqual(a, same), "SpecEqual() = false for identical specs") + c.False(SpecEqual(a, other), "SpecEqual() = true for differing specs") } func TestApplyOneCreates(t *testing.T) { + ck := assert.NewAborting(t) c := fakeClient(t) - if err := ApplyOne(context.Background(), c, certFixture(t, "a", "a-secret")); err != nil { - t.Fatalf("ApplyOne() error = %v", err) - } + ck.NoError( + ApplyOne(context.Background(), c, certFixture(t, "a", "a-secret")), + "ApplyOne() error =", + ) got := &unstructured.Unstructured{} got.SetGroupVersionKind(GVK) - if err := c.Get( + ck.NoError(c.Get( context.Background(), client.ObjectKey{Namespace: "supabase", Name: "a"}, got, - ); err != nil { - t.Fatalf("expected Certificate to be created, got error %v", err) - } + ), "expected Certificate to be created, got error") } diff --git a/pkg/util/metadata/labels_test.go b/pkg/util/metadata/labels_test.go index 97e3fbf3..bde0e6c5 100644 --- a/pkg/util/metadata/labels_test.go +++ b/pkg/util/metadata/labels_test.go @@ -3,9 +3,9 @@ package metadata_test import ( "testing" - "github.com/google/go-cmp/cmp" - "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestBuildStandardLabels(t *testing.T) { @@ -41,9 +41,7 @@ func TestBuildStandardLabels(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := metadata.BuildStandardLabels(tc.clusterName, tc.componentName) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("BuildStandardLabels() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "BuildStandardLabels() mismatch") }) } } @@ -108,9 +106,7 @@ func TestMergeLabels(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := metadata.MergeLabels(tc.standardLabels, tc.customLabels) - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("MergeLabels() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "MergeLabels() mismatch") }) } } @@ -119,41 +115,34 @@ func TestAddMultigresLabels(t *testing.T) { t.Run("AddCellLabel", func(t *testing.T) { labels := map[string]string{"app.kubernetes.io/name": "multigres"} metadata.AddCellLabel(labels, "zone1") - if labels["multigres.com/cell"] != "zone1" { - t.Errorf("AddCellLabel failed") - } + assert.NewCollecting(t).Eq("zone1", labels["multigres.com/cell"], "AddCellLabel failed") }) t.Run("AddClusterLabel", func(t *testing.T) { labels := map[string]string{"app.kubernetes.io/name": "multigres"} metadata.AddClusterLabel(labels, "prod-cluster") - if labels["multigres.com/cluster"] != "prod-cluster" { - t.Errorf("AddClusterLabel failed") - } + assert.NewCollecting(t). + Eq("prod-cluster", labels["multigres.com/cluster"], "AddClusterLabel failed") }) t.Run("AddShardLabel", func(t *testing.T) { labels := map[string]string{"app.kubernetes.io/name": "multigres"} metadata.AddShardLabel(labels, "shard-0") - if labels["multigres.com/shard"] != "shard-0" { - t.Errorf("AddShardLabel failed") - } + assert.NewCollecting(t).Eq("shard-0", labels["multigres.com/shard"], "AddShardLabel failed") }) t.Run("AddDatabaseLabel", func(t *testing.T) { labels := map[string]string{"app.kubernetes.io/name": "multigres"} metadata.AddDatabaseLabel(labels, "proddb") - if labels["multigres.com/database"] != "proddb" { - t.Errorf("AddDatabaseLabel failed") - } + assert.NewCollecting(t). + Eq("proddb", labels["multigres.com/database"], "AddDatabaseLabel failed") }) t.Run("AddTableGroupLabel", func(t *testing.T) { labels := map[string]string{"app.kubernetes.io/name": "multigres"} metadata.AddTableGroupLabel(labels, "orders") - if labels["multigres.com/tablegroup"] != "orders" { - t.Errorf("AddTableGroupLabel failed") - } + assert.NewCollecting(t). + Eq("orders", labels["multigres.com/tablegroup"], "AddTableGroupLabel failed") }) } @@ -300,9 +289,7 @@ func TestLabelOperations_ComplexScenarios(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := tc.setupFunc() - if diff := cmp.Diff(tc.want, got); diff != "" { - t.Errorf("Label operations mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tc.want, got, "Label operations mismatch") }) } } @@ -349,9 +336,7 @@ func TestAddExtraLabels(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { tt.addFunc(tt.initial) - if diff := cmp.Diff(tt.expected, tt.initial); diff != "" { - t.Errorf("Labels mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(tt.expected, tt.initial, "Labels mismatch") }) } } @@ -370,7 +355,5 @@ func TestGetSelectorLabels(t *testing.T) { } got := metadata.GetSelectorLabels(labels) - if diff := cmp.Diff(want, got); diff != "" { - t.Errorf("GetSelectorLabels() mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t).EqDiff(want, got, "GetSelectorLabels() mismatch") } diff --git a/pkg/util/metadata/project_ref_test.go b/pkg/util/metadata/project_ref_test.go index 9b234b50..1b9392cf 100644 --- a/pkg/util/metadata/project_ref_test.go +++ b/pkg/util/metadata/project_ref_test.go @@ -4,6 +4,8 @@ import ( "testing" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) func TestResolveProjectRef(t *testing.T) { @@ -40,9 +42,7 @@ func TestResolveProjectRef(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { got := metadata.ResolveProjectRef(tc.annotations, tc.clusterName) - if got != tc.want { - t.Fatalf("ResolveProjectRef() = %q, want %q", got, tc.want) - } + assert.NewAborting(t).Eq(tc.want, got, "ResolveProjectRef()") }) } } diff --git a/pkg/util/name/name_test.go b/pkg/util/name/name_test.go index 9ecdf0cc..c7820f16 100644 --- a/pkg/util/name/name_test.go +++ b/pkg/util/name/name_test.go @@ -19,6 +19,8 @@ package name import ( "strings" "testing" + + "github.com/multigres/testkit/assert" ) // TestJoin checks determinism and uniqueness. @@ -74,13 +76,13 @@ func TestJoin(t *testing.T) { }, } for _, test := range table { - if got := JoinWithConstraints( + got := JoinWithConstraints( DefaultConstraints, test.a...) == JoinWithConstraints( DefaultConstraints, - test.b...); got != test.shouldEqual { - t.Errorf("JoinWithConstraints: %s: got %v; want %v", test.name, got, test.shouldEqual) - } + test.b...) + assert.NewCollecting(t). + Eq(test.shouldEqual, got, "JoinWithConstraints: %s: got %v; want", test.name, got) } } @@ -88,9 +90,8 @@ func TestJoin(t *testing.T) { func TestJoinHash(t *testing.T) { parts := []string{"hello", "world"} want := "hello-world-344ce285" - if got := JoinWithConstraints(DefaultConstraints, parts...); got != want { - t.Fatalf("JoinWithConstraints(%v) = %q, want %q", parts, got, want) - } + got := JoinWithConstraints(DefaultConstraints, parts...) + assert.NewAborting(t).Eq(want, got, "JoinWithConstraints(%v) = %q, want", parts, got) } func TestJoinWithConstraints(t *testing.T) { @@ -145,6 +146,7 @@ func TestJoinWithConstraints(t *testing.T) { // TestJoinWithConstraintsMaxLength checks that values are truncated to fit // within the max length. func TestJoinWithConstraintsMaxLength(t *testing.T) { + c := assert.NewCollecting(t) cons := Constraints{ MaxLength: 25, ValidFirstChar: isLowercaseAlphanumeric, @@ -152,18 +154,14 @@ func TestJoinWithConstraintsMaxLength(t *testing.T) { // The total length after truncation should be equal to MaxLength. out := JoinWithConstraints(cons, strings.Repeat("a", 20), strings.Repeat("b", 20)) - if len(out) != cons.MaxLength { - t.Errorf("len(%q) = %v; want %v", out, len(out), cons.MaxLength) - } + c.Len(out, cons.MaxLength, "len(%q) = %v; want", out, len(out)) // The outputs should still be unique thanks to the hash suffix, // even if the truncated portion is the same because the difference between // inputs is at the end that gets cut off. out1 := JoinWithConstraints(cons, strings.Repeat("a", 20), strings.Repeat("b", 100)+"1") out2 := JoinWithConstraints(cons, strings.Repeat("a", 20), strings.Repeat("b", 100)+"2") - if out1 == out2 { - t.Errorf("got same output for two different inputs: %v", out1) - } + c.NotEq(out2, out1, "got same output for two different inputs") } // TestJoinWithConstraintsTransform checks that outputs are still @@ -180,9 +178,7 @@ func TestJoinWithConstraintsTransform(t *testing.T) { // transformation. out1 := JoinWithConstraints(cons, "disallowed_symbol") out2 := JoinWithConstraints(cons, "disallowed/symbol") - if out1 == out2 { - t.Errorf("got same output for two different inputs: %v", out1) - } + assert.NewCollecting(t).NotEq(out2, out1, "got same output for two different inputs") } // TestCollisionPrevention verifies that the naming scheme prevents the collision @@ -195,9 +191,8 @@ func TestCollisionPrevention(t *testing.T) { name2 := JoinWithConstraints(DefaultConstraints, "production", "db-app", "sales") // These should produce different names despite appearing identical before hashing - if name1 == name2 { - t.Errorf("collision detected: both scenarios produced the same name %q", name1) - } + assert.NewCollecting(t). + NotEq(name2, name1, "collision detected: both scenarios produced the same name") // Verify both start with similar visible parts but have different hashes t.Logf("Scenario 1 name: %s", name1) @@ -250,9 +245,7 @@ func TestMultigresResourceNaming(t *testing.T) { // Check that it has a hash at the end (8 hex chars after last hyphen) parts := strings.Split(got, "-") lastPart := parts[len(parts)-1] - if len(lastPart) != hashLength { - t.Errorf("expected hash suffix of length %d, got %q", hashLength, lastPart) - } + assert.NewCollecting(t).Len(lastPart, hashLength, "expected hash suffix of length") } t.Logf("Generated name: %s", got) @@ -262,17 +255,24 @@ func TestMultigresResourceNaming(t *testing.T) { // TestServiceConstraints verifies that service constraints work correctly. func TestServiceConstraints(t *testing.T) { + c := assert.NewCollecting(t) parts := []string{"production", "gateway"} got := JoinWithConstraints(ServiceConstraints, parts...) - if len(got) > 63 { - t.Errorf("service name %q exceeds 63 character limit (got %d)", got, len(got)) - } + c.LessOrEqual( + 63, + len(got), + "service name %q exceeds 63 character limit (got %d)", + got, + len(got), + ) // Should start with a lowercase letter - if !isLowercaseLetter(rune(got[0])) { - t.Errorf("service name %q does not start with lowercase letter", got) - } + c.True( + isLowercaseLetter(rune(got[0])), + "service name %q does not start with lowercase letter", + got, + ) t.Logf("Service name: %s (length: %d)", got, len(got)) } @@ -280,9 +280,7 @@ func TestServiceConstraints(t *testing.T) { // TestInvalidConstraintsPanic verifies that invalid constraints cause a panic. func TestInvalidConstraintsPanic(t *testing.T) { defer func() { - if r := recover(); r == nil { - t.Errorf("expected panic for invalid constraints") - } + assert.NewCollecting(t).NotNil(recover(), "expected panic for invalid constraints") }() invalidCons := Constraints{ @@ -296,7 +294,5 @@ func TestInvalidConstraintsPanic(t *testing.T) { // TestEmptyParts verifies handling of empty input. func TestEmptyParts(t *testing.T) { got := JoinWithConstraints(DefaultConstraints) - if got != "" { - t.Errorf("expected empty string for empty parts, got %q", got) - } + assert.NewCollecting(t).Eq("", got, "expected empty string for empty parts, got") } diff --git a/pkg/util/pvc/orphan_test.go b/pkg/util/pvc/orphan_test.go index 481e347f..55b72c01 100644 --- a/pkg/util/pvc/orphan_test.go +++ b/pkg/util/pvc/orphan_test.go @@ -6,7 +6,6 @@ import ( "time" "github.com/go-logr/logr" - "github.com/stretchr/testify/require" corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/types" @@ -15,6 +14,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" "github.com/multigres/multigres-operator/pkg/util/metadata" + + "github.com/multigres/testkit/assert" ) const ( @@ -46,51 +47,55 @@ func newClient(t *testing.T, objs ...client.Object) client.Client { func TestMarkOrphan_AddsLabelAndStripsOwnerRef(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) pvc := newPVC("a", shardUID, otherUID) c := newClient(t, pvc) now := time.Date(2026, 5, 19, 12, 0, 0, 0, time.UTC) - require.NoError(t, MarkOrphan(context.Background(), c, pvc, shardUID, now)) + ck.NoError(MarkOrphan(context.Background(), c, pvc, shardUID, now)) got := &corev1.PersistentVolumeClaim{} - require.NoError(t, c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) - require.Equal(t, "2026-05-19T12-00-00Z", got.Labels[metadata.LabelOrphan]) - require.Len(t, got.OwnerReferences, 1) - require.Equal(t, otherUID, got.OwnerReferences[0].UID) + ck.NoError(c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) + ck.EqDeep("2026-05-19T12-00-00Z", got.Labels[metadata.LabelOrphan]) + ck.Len(got.OwnerReferences, 1) + ck.EqDeep(otherUID, got.OwnerReferences[0].UID) } func TestMarkOrphan_Idempotent(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) pvc := newPVC("a") pvc.Labels = map[string]string{metadata.LabelOrphan: "2026-05-01T00-00-00Z"} c := newClient(t, pvc) now := time.Date(2026, 5, 19, 12, 0, 0, 0, time.UTC) - require.NoError(t, MarkOrphan(context.Background(), c, pvc, "", now)) + ck.NoError(MarkOrphan(context.Background(), c, pvc, "", now)) got := &corev1.PersistentVolumeClaim{} - require.NoError(t, c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) + ck.NoError(c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) // Existing label preserved — retention is measured from the original event. - require.Equal(t, "2026-05-01T00-00-00Z", got.Labels[metadata.LabelOrphan]) + ck.EqDeep("2026-05-01T00-00-00Z", got.Labels[metadata.LabelOrphan]) } func TestMarkOrphan_NoOwnerUIDLeavesRefs(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) pvc := newPVC("a", shardUID) c := newClient(t, pvc) - require.NoError(t, MarkOrphan(context.Background(), c, pvc, "", time.Now())) + ck.NoError(MarkOrphan(context.Background(), c, pvc, "", time.Now())) got := &corev1.PersistentVolumeClaim{} - require.NoError(t, c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) - require.Len(t, got.OwnerReferences, 1) + ck.NoError(c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) + ck.Len(got.OwnerReferences, 1) } func TestClearOrphan_RemovesLabel(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) logger := logr.Discard() pvc := newPVC("a") @@ -100,38 +105,40 @@ func TestClearOrphan_RemovesLabel(t *testing.T) { } c := newClient(t, pvc) - require.NoError(t, ClearOrphan(context.Background(), logger, c, pvc)) + ck.NoError(ClearOrphan(context.Background(), logger, c, pvc)) got := &corev1.PersistentVolumeClaim{} - require.NoError(t, c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) + ck.NoError(c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) _, hasOrphan := got.Labels[metadata.LabelOrphan] - require.False(t, hasOrphan) + ck.False(hasOrphan) // other labels should remain untouched. - require.Equal(t, metadata.ManagedByMultigres, got.Labels[metadata.LabelAppManagedBy]) + ck.EqDeep(metadata.ManagedByMultigres, got.Labels[metadata.LabelAppManagedBy]) } func TestClearOrphan_NoLabelIsNoOp(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) logger := logr.Discard() pvc := newPVC("a") pvc.Labels = map[string]string{metadata.LabelAppManagedBy: metadata.ManagedByMultigres} c := newClient(t, pvc) - require.NoError(t, ClearOrphan(context.Background(), logger, c, pvc)) + ck.NoError(ClearOrphan(context.Background(), logger, c, pvc)) got := &corev1.PersistentVolumeClaim{} - require.NoError(t, c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) - require.Equal(t, metadata.ManagedByMultigres, got.Labels[metadata.LabelAppManagedBy]) + ck.NoError(c.Get(context.Background(), client.ObjectKeyFromObject(pvc), got)) + ck.EqDeep(metadata.ManagedByMultigres, got.Labels[metadata.LabelAppManagedBy]) } func TestHasOrphanLabel(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) - require.False(t, HasOrphanLabel(nil)) - require.False(t, HasOrphanLabel(newPVC("a"))) + c.False(HasOrphanLabel(nil)) + c.False(HasOrphanLabel(newPVC("a"))) labeled := newPVC("a") labeled.Labels = map[string]string{metadata.LabelOrphan: "x"} - require.True(t, HasOrphanLabel(labeled)) + c.True(HasOrphanLabel(labeled)) } diff --git a/pkg/util/pvc/retention_test.go b/pkg/util/pvc/retention_test.go index 8cd2b9d6..cde8e03b 100644 --- a/pkg/util/pvc/retention_test.go +++ b/pkg/util/pvc/retention_test.go @@ -6,6 +6,8 @@ import ( appsv1 "k8s.io/api/apps/v1" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestBuildRetentionPolicy(t *testing.T) { @@ -58,13 +60,10 @@ func TestBuildRetentionPolicy(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) got := BuildRetentionPolicy(tc.policy) - if got.WhenDeleted != tc.wantDeleted { - t.Errorf("WhenDeleted = %q, want %q", got.WhenDeleted, tc.wantDeleted) - } - if got.WhenScaled != tc.wantScaled { - t.Errorf("WhenScaled = %q, want %q", got.WhenScaled, tc.wantScaled) - } + c.Eq(tc.wantDeleted, got.WhenDeleted, "WhenDeleted") + c.Eq(tc.wantScaled, got.WhenScaled, "WhenScaled") }) } } diff --git a/pkg/util/status/conditions_test.go b/pkg/util/status/conditions_test.go index 48fc28c7..0ee6ad0c 100644 --- a/pkg/util/status/conditions_test.go +++ b/pkg/util/status/conditions_test.go @@ -4,6 +4,8 @@ import ( "testing" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + + "github.com/multigres/testkit/assert" ) func TestSetCondition(t *testing.T) { @@ -91,32 +93,28 @@ func TestSetCondition(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) conditions := make([]metav1.Condition, len(tc.existing)) copy(conditions, tc.existing) SetCondition(&conditions, tc.condition) - if len(conditions) != len(tc.want) { - t.Fatalf("got %d conditions, want %d", len(conditions), len(tc.want)) - } + c.Require().Len(conditions, len(tc.want), "got %d conditions, want", len(conditions)) for i, got := range conditions { want := tc.want[i] - if got.Type != want.Type { - t.Errorf("[%d] type = %s, want %s", i, got.Type, want.Type) - } - if got.Status != want.Status { - t.Errorf("[%d] status = %s, want %s", i, got.Status, want.Status) - } - if got.Reason != want.Reason { - t.Errorf("[%d] reason = %s, want %s", i, got.Reason, want.Reason) - } + c.Eq(want.Type, got.Type, "[%d] type = %s, want", i, got.Type) + c.Eq(want.Status, got.Status, "[%d] status = %s, want", i, got.Status) + c.Eq(want.Reason, got.Reason, "[%d] reason = %s, want", i, got.Reason) if want.LastTransitionTime.IsZero() && !got.LastTransitionTime.IsZero() { // The "preserves transition time" case: original was zero, // SetCondition should have kept the existing (zero) time. - if name == "preserves transition time when status unchanged" { - t.Errorf("[%d] expected preserved zero transition time, got %v", - i, got.LastTransitionTime) - } + c.NotEq( + "preserves transition time when status unchanged", + name, + "[%d] expected preserved zero transition time, got %v", + i, + got.LastTransitionTime, + ) } } }) @@ -169,9 +167,8 @@ func TestIsConditionTrue(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() - if got := IsConditionTrue(tc.conditions, tc.condType); got != tc.want { - t.Errorf("IsConditionTrue() = %v, want %v", got, tc.want) - } + assert.NewCollecting(t). + Eq(tc.want, IsConditionTrue(tc.conditions, tc.condType), "IsConditionTrue()") }) } } @@ -222,9 +219,8 @@ func TestIsConditionFalse(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() - if got := IsConditionFalse(tc.conditions, tc.condType); got != tc.want { - t.Errorf("IsConditionFalse() = %v, want %v", got, tc.want) - } + assert.NewCollecting(t). + Eq(tc.want, IsConditionFalse(tc.conditions, tc.condType), "IsConditionFalse()") }) } } diff --git a/pkg/util/status/phase_test.go b/pkg/util/status/phase_test.go index 95fe442a..6fbc2a91 100644 --- a/pkg/util/status/phase_test.go +++ b/pkg/util/status/phase_test.go @@ -7,6 +7,8 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestComputePhase(t *testing.T) { @@ -44,9 +46,7 @@ func TestComputePhase(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := ComputePhase(tt.ready, tt.total); got != tt.want { - t.Errorf("ComputePhase() = %v, want %v", got, tt.want) - } + assert.NewCollecting(t).Eq(tt.want, ComputePhase(tt.ready, tt.total), "ComputePhase()") }) } } @@ -188,9 +188,7 @@ func TestIsCrashLooping(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := IsCrashLooping(&tt.pod); got != tt.want { - t.Errorf("IsCrashLooping() = %v, want %v", got, tt.want) - } + assert.NewCollecting(t).Eq(tt.want, IsCrashLooping(&tt.pod), "IsCrashLooping()") }) } } @@ -239,9 +237,7 @@ func TestAnyCrashLooping(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := AnyCrashLooping(tt.pods); got != tt.want { - t.Errorf("AnyCrashLooping() = %v, want %v", got, tt.want) - } + assert.NewCollecting(t).Eq(tt.want, AnyCrashLooping(tt.pods), "AnyCrashLooping()") }) } } @@ -333,13 +329,10 @@ func TestComputeWorkloadPhase(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) result := ComputeWorkloadPhase(tt.input) - if result.Phase != tt.wantPhase { - t.Errorf("ComputeWorkloadPhase() phase = %v, want %v", result.Phase, tt.wantPhase) - } - if result.Message == "" { - t.Error("ComputeWorkloadPhase() returned empty message") - } + c.Eq(tt.wantPhase, result.Phase, "ComputeWorkloadPhase() phase") + c.NotEq("", result.Message, "ComputeWorkloadPhase() returned empty message") }) } } diff --git a/pkg/webhook/cel_validation_test.go b/pkg/webhook/cel_validation_test.go index ceb51032..5956e62c 100644 --- a/pkg/webhook/cel_validation_test.go +++ b/pkg/webhook/cel_validation_test.go @@ -15,19 +15,24 @@ import ( "k8s.io/client-go/rest" "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) func getPrivilegedClient(t *testing.T) client.Client { t.Helper() - if TestCfg == nil { - t.Fatal("TestCfg is nil") - } + ck := assert.NewAborting(t) + ck.NotNil(TestCfg, "TestCfg is nil") // Ensure the impersonated identity has permissions (cluster-admin) binding := &rbacv1.ClusterRoleBinding{ ObjectMeta: metav1.ObjectMeta{Name: "operator-admin-binding"}, Subjects: []rbacv1.Subject{ - {Kind: "User", Name: "system:serviceaccount:default:multigres-operator", APIGroup: "rbac.authorization.k8s.io"}, + { + Kind: "User", + Name: "system:serviceaccount:default:multigres-operator", + APIGroup: "rbac.authorization.k8s.io", + }, }, RoleRef: rbacv1.RoleRef{ Kind: "ClusterRole", @@ -38,9 +43,7 @@ func getPrivilegedClient(t *testing.T) client.Client { // Use the existing admin client to create the binding if err := k8sClient.Create(context.Background(), binding); err != nil { // Ignore if already exists, otherwise fail - if client.IgnoreAlreadyExists(err) != nil { - t.Fatalf("Failed to create ClusterRoleBinding: %v", err) - } + ck.NoError(client.IgnoreAlreadyExists(err), "Failed to create ClusterRoleBinding: %v", err) } config := *TestCfg @@ -48,9 +51,7 @@ func getPrivilegedClient(t *testing.T) client.Client { UserName: "system:serviceaccount:default:multigres-operator", } c, err := client.New(&config, client.Options{Scheme: k8sClient.Scheme()}) - if err != nil { - t.Fatalf("Failed to create privileged client: %v", err) - } + ck.NoError(err, "Failed to create privileged client") return c } @@ -92,10 +93,18 @@ func TestCEL_MultigresCluster(t *testing.T) { Name: "invalid-cell", ZoneID: "use1-az1", Spec: &multigresv1alpha1.CellInlineSpec{ - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(1)), + }, + }, }, Overrides: &multigresv1alpha1.CellOverrides{ - Multigateway: &multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(2))}}, + Multigateway: &multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(2)), + }, + }, }, }, }, @@ -117,7 +126,11 @@ func TestCEL_MultigresCluster(t *testing.T) { ZoneID: "use1-az1", CellTemplate: "some-template", Spec: &multigresv1alpha1.CellInlineSpec{ - Multigateway: multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + Multigateway: multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{ + Replicas: ptr.To(int32(1)), + }, + }, }, }, }, @@ -231,16 +244,13 @@ func TestCEL_MultigresCluster(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) setTestPostgresPasswordSecretRef(tc.cluster) err := k8sClient.Create(ctx, tc.cluster) - if err == nil { - t.Fatal("Expected error, got nil") - } + c.Require().Error(err, "Expected error, got nil") // In envtest, CEL errors usually appear in the error string if tc.expectError != "" { - if !strings.Contains(err.Error(), tc.expectError) { - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) - } + c.StrContains(err.Error(), tc.expectError, "Expected error message to contain") } }) } @@ -263,8 +273,10 @@ func TestCEL_TopoServer(t *testing.T) { }, Spec: multigresv1alpha1.MultigresClusterSpec{ GlobalTopoServer: &multigresv1alpha1.GlobalTopoServerSpec{ - Etcd: &multigresv1alpha1.EtcdSpec{Replicas: ptr.To(int32(1))}, - External: &multigresv1alpha1.ExternalTopoServerSpec{Endpoints: []multigresv1alpha1.EndpointUrl{"http://etcd:2379"}}, + Etcd: &multigresv1alpha1.EtcdSpec{Replicas: ptr.To(int32(1))}, + External: &multigresv1alpha1.ExternalTopoServerSpec{ + Endpoints: []multigresv1alpha1.EndpointUrl{"http://etcd:2379"}, + }, }, Cells: []multigresv1alpha1.CellConfig{}, // Empty for simplicity }, @@ -275,15 +287,12 @@ func TestCEL_TopoServer(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) setTestPostgresPasswordSecretRef(tc.cluster) err := k8sClient.Create(ctx, tc.cluster) - if err == nil { - t.Fatal("Expected error, got nil") - } + c.Require().Error(err, "Expected error, got nil") if tc.expectError != "" { - if !strings.Contains(err.Error(), tc.expectError) { - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) - } + c.StrContains(err.Error(), tc.expectError, "Expected error message to contain") } }) } @@ -343,15 +352,12 @@ func TestCEL_Limits(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) setTestPostgresPasswordSecretRef(tc.cluster) err := k8sClient.Create(ctx, tc.cluster) - if err == nil { - t.Fatal("Expected error, got nil") - } + c.Require().Error(err, "Expected error, got nil") if tc.expectError != "" { - if !strings.Contains(err.Error(), tc.expectError) { - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) - } + c.StrContains(err.Error(), tc.expectError, "Expected error message to contain") } }) } @@ -411,15 +417,12 @@ func TestCEL_Database(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) setTestPostgresPasswordSecretRef(tc.cluster) err := k8sClient.Create(ctx, tc.cluster) - if err == nil { - t.Fatal("Expected error, got nil") - } + c.Require().Error(err, "Expected error, got nil") if tc.expectError != "" { - if !strings.Contains(err.Error(), tc.expectError) { - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) - } + c.StrContains(err.Error(), tc.expectError, "Expected error message to contain") } }) } @@ -466,15 +469,12 @@ func TestCEL_Multiadmin(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewCollecting(t) setTestPostgresPasswordSecretRef(tc.cluster) err := k8sClient.Create(ctx, tc.cluster) - if err == nil { - t.Fatal("Expected error, got nil") - } + c.Require().Error(err, "Expected error, got nil") if tc.expectError != "" { - if !strings.Contains(err.Error(), tc.expectError) { - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) - } + c.StrContains(err.Error(), tc.expectError, "Expected error message to contain") } }) } @@ -482,6 +482,7 @@ func TestCEL_Multiadmin(t *testing.T) { func TestCEL_ShardImmutability(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) shardName := "immutable-shard" shard := &multigresv1alpha1.Shard{ @@ -519,38 +520,35 @@ func TestCEL_ShardImmutability(t *testing.T) { // Create privClient := getPrivilegedClient(t) setTestShardPostgresPasswordSecretRef(shard) - if err := privClient.Create(ctx, shard); err != nil { - t.Fatalf("Failed to create Shard: %v", err) - } + c.Require().NoError(privClient.Create(ctx, shard), "Failed to create Shard") // Try to update immutable fields (Validation removed by user, so updates should succeed) toUpdate := shard.DeepCopy() toUpdate.Spec.DatabaseName = "other-db" - if err := privClient.Update(ctx, toUpdate); err != nil { - t.Errorf("Expected success when updating previously immutable databaseName, got error: %v", err) - } + c.NoError( + privClient.Update(ctx, toUpdate), + "Expected success when updating previously immutable databaseName, got error", + ) // Refetch to avoid conflict - if err := privClient.Get(ctx, client.ObjectKeyFromObject(shard), shard); err != nil { - t.Fatal(err) - } + c.Require().NoError(privClient.Get(ctx, client.ObjectKeyFromObject(shard), shard)) toUpdate = shard.DeepCopy() toUpdate.Spec.TableGroupName = "other-tg" - if err := privClient.Update(ctx, toUpdate); err != nil { - t.Errorf("Expected success when updating previously immutable tableGroupName, got error: %v", err) - } + c.NoError( + privClient.Update(ctx, toUpdate), + "Expected success when updating previously immutable tableGroupName, got error", + ) // Refetch to avoid conflict - if err := privClient.Get(ctx, client.ObjectKeyFromObject(shard), shard); err != nil { - t.Fatal(err) - } + c.Require().NoError(privClient.Get(ctx, client.ObjectKeyFromObject(shard), shard)) toUpdate = shard.DeepCopy() toUpdate.Spec.ShardName = "1" - if err := privClient.Update(ctx, toUpdate); err != nil { - t.Errorf("Expected success when updating previously immutable shardName, got error: %v", err) - } + c.NoError( + privClient.Update(ctx, toUpdate), + "Expected success when updating previously immutable shardName, got error", + ) } func TestCEL_ExtendedValidation(t *testing.T) { @@ -564,11 +562,24 @@ func TestCEL_ExtendedValidation(t *testing.T) { { name: "Invalid Duplicate Cells in Multiorch", shard: &multigresv1alpha1.Shard{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-duplicate-multiorch-cells", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-duplicate-multiorch-cells", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, Multiorch: multigresv1alpha1.MultiorchSpec{ Cells: []multigresv1alpha1.CellName{"cell-1", "cell-1"}, // Duplicate }, @@ -582,12 +593,25 @@ func TestCEL_ExtendedValidation(t *testing.T) { { name: "Invalid Duplicate Cells in Pool", shard: &multigresv1alpha1.Shard{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-duplicate-pool-cells", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-duplicate-pool-cells", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, - Multiorch: multigresv1alpha1.MultiorchSpec{}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, + Multiorch: multigresv1alpha1.MultiorchSpec{}, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ "rw": { Cells: []multigresv1alpha1.CellName{"cell-1", "cell-1"}, // Duplicate @@ -639,12 +663,25 @@ func TestCEL_ExtendedValidation(t *testing.T) { { name: "Invalid ReplicasPerCell Limit", shard: &multigresv1alpha1.Shard{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-replicas-per-cell-limit", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-replicas-per-cell-limit", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, - Multiorch: multigresv1alpha1.MultiorchSpec{}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, + Multiorch: multigresv1alpha1.MultiorchSpec{}, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ "rw": { Cells: []multigresv1alpha1.CellName{"cell-1"}, @@ -658,12 +695,25 @@ func TestCEL_ExtendedValidation(t *testing.T) { { name: "Invalid ReplicasPerCell Zero", shard: &multigresv1alpha1.Shard{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-replicas-per-cell-zero", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-replicas-per-cell-zero", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, - Multiorch: multigresv1alpha1.MultiorchSpec{}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, + Multiorch: multigresv1alpha1.MultiorchSpec{}, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ "rw": { Cells: []multigresv1alpha1.CellName{"cell-1"}, @@ -681,13 +731,15 @@ func TestCEL_ExtendedValidation(t *testing.T) { privClient := getPrivilegedClient(t) setTestObjectPostgresPasswordSecretRef(tc.shard) err := privClient.Create(ctx, tc.shard) - if err == nil { - t.Fatal("Expected error, got nil") - } + assert.NewAborting(t).Error(err, "Expected error, got nil") if tc.expectError != "" { if !strings.Contains(err.Error(), tc.expectError) { t.Logf("ACTUAL ERROR: %v", err) - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) + t.Errorf( + "Expected error message to contain %q, got %q", + tc.expectError, + err.Error(), + ) } } }) @@ -726,28 +778,31 @@ func TestCEL_PoolRuntimeIdentity(t *testing.T) { } } - t.Run("allows different PGDATA UIDs at API layer because PoolSpec may be partial", func(t *testing.T) { - shard := newShard("cel-runtime-uid-mismatch", ptr.To(int64(1000)), ptr.To(int64(3000))) - setTestShardPostgresPasswordSecretRef(shard) - if err := getPrivilegedClient(t).Create(ctx, shard); err != nil { - t.Fatalf("expected API validation to allow partial PoolSpec identity, got %v", err) - } - }) - - t.Run("allows multipooler-only UID at API layer because PoolSpec may be partial", func(t *testing.T) { - shard := newShard("cel-runtime-multipooler-only", nil, ptr.To(int64(1000))) - setTestShardPostgresPasswordSecretRef(shard) - if err := getPrivilegedClient(t).Create(ctx, shard); err != nil { - t.Fatalf("expected API validation to allow partial PoolSpec identity, got %v", err) - } - }) + t.Run( + "allows different PGDATA UIDs at API layer because PoolSpec may be partial", + func(t *testing.T) { + shard := newShard("cel-runtime-uid-mismatch", ptr.To(int64(1000)), ptr.To(int64(3000))) + setTestShardPostgresPasswordSecretRef(shard) + assert.NewAborting(t). + NoError(getPrivilegedClient(t).Create(ctx, shard), "expected API validation to allow partial PoolSpec identity, got") + }, + ) + + t.Run( + "allows multipooler-only UID at API layer because PoolSpec may be partial", + func(t *testing.T) { + shard := newShard("cel-runtime-multipooler-only", nil, ptr.To(int64(1000))) + setTestShardPostgresPasswordSecretRef(shard) + assert.NewAborting(t). + NoError(getPrivilegedClient(t).Create(ctx, shard), "expected API validation to allow partial PoolSpec identity, got") + }, + ) t.Run("accepts root primary and filesystem groups", func(t *testing.T) { shard := newShard("cel-runtime-root-group", ptr.To(int64(1000)), ptr.To(int64(1000))) setTestShardPostgresPasswordSecretRef(shard) - if err := getPrivilegedClient(t).Create(ctx, shard); err != nil { - t.Fatalf("expected group ID 0 to be accepted, got %v", err) - } + assert.NewAborting(t). + NoError(getPrivilegedClient(t).Create(ctx, shard), "expected group ID 0 to be accepted, got") }) } @@ -763,14 +818,33 @@ func TestCEL_S3ServiceAccountName(t *testing.T) { { name: "Rejected: serviceAccountName with credentialsSecret", shard: &multigresv1alpha1.Shard{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-s3-sa-creds-conflict", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-s3-sa-creds-conflict", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, - Multiorch: multigresv1alpha1.MultiorchSpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, + Multiorch: multigresv1alpha1.MultiorchSpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, + }, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ - "rw": {Type: "readWrite", Cells: []multigresv1alpha1.CellName{"cell-1"}, ReplicasPerCell: ptr.To(int32(1))}, + "rw": { + Type: "readWrite", + Cells: []multigresv1alpha1.CellName{"cell-1"}, + ReplicasPerCell: ptr.To(int32(1)), + }, }, Backup: &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeS3, @@ -788,14 +862,33 @@ func TestCEL_S3ServiceAccountName(t *testing.T) { { name: "Rejected: serviceAccountName with useEnvCredentials", shard: &multigresv1alpha1.Shard{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-s3-sa-envcreds-conflict", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-s3-sa-envcreds-conflict", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, - Multiorch: multigresv1alpha1.MultiorchSpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, + Multiorch: multigresv1alpha1.MultiorchSpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, + }, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ - "rw": {Type: "readWrite", Cells: []multigresv1alpha1.CellName{"cell-1"}, ReplicasPerCell: ptr.To(int32(1))}, + "rw": { + Type: "readWrite", + Cells: []multigresv1alpha1.CellName{"cell-1"}, + ReplicasPerCell: ptr.To(int32(1)), + }, }, Backup: &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeS3, @@ -815,12 +908,28 @@ func TestCEL_S3ServiceAccountName(t *testing.T) { shard: &multigresv1alpha1.Shard{ ObjectMeta: metav1.ObjectMeta{Name: "cel-s3-sa-only", Namespace: testNamespace}, Spec: multigresv1alpha1.ShardSpec{ - DatabaseName: "postgres", TableGroupName: "default", ShardName: "0", - Images: multigresv1alpha1.ShardImages{Postgres: "p", Multiorch: "o", Multipooler: "po"}, - GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{Address: "etcd", RootPath: "/", Implementation: "etcd"}, - Multiorch: multigresv1alpha1.MultiorchSpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + DatabaseName: "postgres", + TableGroupName: "default", + ShardName: "0", + Images: multigresv1alpha1.ShardImages{ + Postgres: "p", + Multiorch: "o", + Multipooler: "po", + }, + GlobalTopoServer: multigresv1alpha1.GlobalTopoServerRef{ + Address: "etcd", + RootPath: "/", + Implementation: "etcd", + }, + Multiorch: multigresv1alpha1.MultiorchSpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, + }, Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ - "rw": {Type: "readWrite", Cells: []multigresv1alpha1.CellName{"cell-1"}, ReplicasPerCell: ptr.To(int32(1))}, + "rw": { + Type: "readWrite", + Cells: []multigresv1alpha1.CellName{"cell-1"}, + ReplicasPerCell: ptr.To(int32(1)), + }, }, Backup: &multigresv1alpha1.BackupConfig{ Type: multigresv1alpha1.BackupTypeS3, @@ -839,20 +948,21 @@ func TestCEL_S3ServiceAccountName(t *testing.T) { privClient := getPrivilegedClient(t) for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { + c := assert.NewAborting(t) setTestObjectPostgresPasswordSecretRef(tc.shard) err := privClient.Create(ctx, tc.shard) if tc.shouldPass { - if err != nil { - t.Fatalf("Expected success, got error: %v", err) - } + c.NoError(err, "Expected success, got error") return } - if err == nil { - t.Fatal("Expected error, got nil") - } + c.Error(err, "Expected error, got nil") if tc.expectError != "" && !strings.Contains(err.Error(), tc.expectError) { t.Logf("ACTUAL ERROR: %v", err) - t.Errorf("Expected error message to contain %q, got %q", tc.expectError, err.Error()) + t.Errorf( + "Expected error message to contain %q, got %q", + tc.expectError, + err.Error(), + ) } }) } @@ -860,9 +970,13 @@ func TestCEL_S3ServiceAccountName(t *testing.T) { func TestCEL_StatelessReplicasLimit(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ - ObjectMeta: metav1.ObjectMeta{Name: "cel-stateless-replicas-limit", Namespace: testNamespace}, + ObjectMeta: metav1.ObjectMeta{ + Name: "cel-stateless-replicas-limit", + Namespace: testNamespace, + }, Spec: multigresv1alpha1.MultigresClusterSpec{ Multiadmin: &multigresv1alpha1.MultiadminConfig{ Spec: &multigresv1alpha1.StatelessSpec{ @@ -875,10 +989,11 @@ func TestCEL_StatelessReplicasLimit(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(ctx, cluster) - if err == nil { - t.Fatal("Expected error, got nil") - } - if !strings.Contains(err.Error(), "less than or equal to 128") { - t.Errorf("Expected error to contain 'less than or equal to 128', got %v", err) - } + c.Require().Error(err, "Expected error, got nil") + c.StrContains( + err.Error(), + "less than or equal to 128", + "Expected error to contain 'less than or equal to 128', got %v", + err, + ) } diff --git a/pkg/webhook/etcd_maintenance_validation_test.go b/pkg/webhook/etcd_maintenance_validation_test.go index 1701d761..e80f3d78 100644 --- a/pkg/webhook/etcd_maintenance_validation_test.go +++ b/pkg/webhook/etcd_maintenance_validation_test.go @@ -8,6 +8,8 @@ import ( "k8s.io/apimachinery/pkg/apis/meta/v1/unstructured" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) func TestCEL_EtcdMaintenance(t *testing.T) { @@ -30,6 +32,7 @@ func TestCEL_EtcdMaintenance(t *testing.T) { {"large-quota", map[string]any{"quotaBackendBytes": int64(9 << 30)}, "less than or equal"}, } { t.Run(tc.name, func(t *testing.T) { + ck := assert.NewCollecting(t) obj := &unstructured.Unstructured{ Object: map[string]any{ "apiVersion": "multigres.com/v1alpha1", @@ -45,20 +48,15 @@ func TestCEL_EtcdMaintenance(t *testing.T) { } err := c.Create(t.Context(), obj) if tc.wantError != "" { - if err == nil || !strings.Contains(err.Error(), tc.wantError) { - t.Fatalf("expected %q, got %v", tc.wantError, err) - } + ck.Require(). + False(err == nil || !strings.Contains(err.Error(), tc.wantError), "expected %q, got %v", tc.wantError, err) return } - if err != nil { - t.Fatal(err) - } + ck.Require().NoError(err) t.Cleanup(func() { _ = c.Delete(t.Context(), obj) }) stored := &unstructured.Unstructured{} stored.SetGroupVersionKind(obj.GroupVersionKind()) - if err := c.Get(t.Context(), client.ObjectKeyFromObject(obj), stored); err != nil { - t.Fatal(err) - } + ck.Require().NoError(c.Get(t.Context(), client.ObjectKeyFromObject(obj), stored)) for key, want := range tc.config { got, found, err := unstructured.NestedFieldNoCopy( stored.Object, @@ -67,9 +65,14 @@ func TestCEL_EtcdMaintenance(t *testing.T) { "maintenance", key, ) - if err != nil || !found || got != want { - t.Errorf("%s: got %v, want %v (err %v)", key, got, want, err) - } + ck.False( + err != nil || !found || got != want, + "%s: got %v, want %v (err %v)", + key, + got, + want, + err, + ) } }) } diff --git a/pkg/webhook/handlers/defaulter_test.go b/pkg/webhook/handlers/defaulter_test.go index fe897059..1dd3290f 100644 --- a/pkg/webhook/handlers/defaulter_test.go +++ b/pkg/webhook/handlers/defaulter_test.go @@ -1,7 +1,6 @@ package handlers import ( - "strings" "testing" "github.com/google/go-cmp/cmp" @@ -17,6 +16,8 @@ import ( "k8s.io/utils/ptr" "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" + + "github.com/multigres/testkit/assert" ) func TestMultigresClusterDefaulter_Handle(t *testing.T) { @@ -191,9 +192,8 @@ func TestMultigresClusterDefaulter_Handle(t *testing.T) { }, }, } - if diff := cmp.Diff(want, &cluster.Spec, cmpopts.EquateEmpty()); diff != "" { - t.Errorf("Cluster mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t). + EqDiffOpts(want, &cluster.Spec, []cmp.Option{cmpopts.EquateEmpty()}, "Cluster mismatch") }, }, "Happy Path: Fallbacks -> Promotes to Explicit": { @@ -489,9 +489,8 @@ func TestMultigresClusterDefaulter_Handle(t *testing.T) { }, }, } - if diff := cmp.Diff(want, cluster, cmpopts.EquateEmpty()); diff != "" { - t.Errorf("Cluster mismatch (-want +got):\n%s", diff) - } + assert.NewCollecting(t). + EqDiffOpts(want, cluster, []cmp.Option{cmpopts.EquateEmpty()}, "Cluster mismatch") }, }, } @@ -499,6 +498,7 @@ func TestMultigresClusterDefaulter_Handle(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) var res *resolver.Resolver if !tc.nilResolver { @@ -528,12 +528,14 @@ func TestMultigresClusterDefaulter_Handle(t *testing.T) { err := defaulter.Default(t.Context(), obj) if tc.expectError != "" { - if err == nil { - t.Fatalf("Expected error containing %q, got nil", tc.expectError) - } - if !strings.Contains(err.Error(), tc.expectError) { - t.Fatalf("Expected error containing %q, got: %v", tc.expectError, err) - } + ck.Error(err, "Expected error containing %q, got nil", tc.expectError) + ck.StrContains( + err.Error(), + tc.expectError, + "Expected error containing %q, got: %v", + tc.expectError, + err, + ) } else if err != nil { t.Fatalf("Unexpected error: %v", err) } diff --git a/pkg/webhook/handlers/validator_test.go b/pkg/webhook/handlers/validator_test.go index bf5aa485..47c56d21 100644 --- a/pkg/webhook/handlers/validator_test.go +++ b/pkg/webhook/handlers/validator_test.go @@ -21,6 +21,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" "sigs.k8s.io/controller-runtime/pkg/webhook/admission" + + "github.com/multigres/testkit/assert" ) // setupScheme creates a new scheme with all required types registered @@ -514,6 +516,7 @@ func TestMultigresClusterValidator(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) // Default existing objects if nil existing := tc.existing @@ -563,20 +566,15 @@ func TestMultigresClusterValidator(t *testing.T) { warnings, err = validator.ValidateDelete(t.Context(), tc.object) } - if tc.wantAllowed && err != nil { - t.Fatalf("Expected allowed, got error: %v", err) - } + c.Require().False(tc.wantAllowed && err != nil, "Expected allowed, got error: %v", err) if !tc.wantAllowed { - if err == nil { - t.Fatal("Expected error, got nil") - } - if tc.wantMessage != "" && !strings.Contains(err.Error(), tc.wantMessage) { - t.Errorf( - "Expected error message containing '%s', got '%v'", - tc.wantMessage, - err, - ) - } + c.Require().Error(err, "Expected error, got nil") + c.False( + tc.wantMessage != "" && !strings.Contains(err.Error(), tc.wantMessage), + "Expected error message containing '%s', got '%v'", + tc.wantMessage, + err, + ) } // Check Warnings @@ -592,13 +590,12 @@ func TestMultigresClusterValidator(t *testing.T) { break } } - if !found { - t.Errorf( - "Expected warning containing '%s', got warnings: %v", - want, - warnings, - ) - } + c.True( + found, + "Expected warning containing '%s', got warnings: %v", + want, + warnings, + ) } } } @@ -621,9 +618,8 @@ func TestMultigresClusterValidator_WrongType(t *testing.T) { t.Parallel() validator := NewMultigresClusterValidator(fake.NewClientBuilder().Build()) _, err := validator.ValidateCreate(t.Context(), &TrulyOnlyRuntimeObject{}) - if err == nil || !strings.Contains(err.Error(), "expected MultigresCluster") { - t.Errorf("Expected wrong type error, got: %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "expected MultigresCluster"), "Expected wrong type error, got: %v", err) } func TestTemplateValidator(t *testing.T) { @@ -791,6 +787,7 @@ func TestTemplateValidator(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) objs := make([]client.Object, len(tc.existing)) for i, obj := range tc.existing { @@ -836,27 +833,21 @@ func TestTemplateValidator(t *testing.T) { _, err = validator.ValidateDelete(t.Context(), obj) } if method != "Delete" { - if err != nil { - t.Errorf("%s: Expected nil error, got %v", method, err) - } + c.NoError(err, "%s: Expected nil error, got", method) continue } // For Delete - if tc.wantAllowed && err != nil { - t.Fatalf("Delete: Expected allowed, got error: %v", err) - } + c.Require(). + False(tc.wantAllowed && err != nil, "Delete: Expected allowed, got error: %v", err) if !tc.wantAllowed { - if err == nil { - t.Fatal("Delete: Expected error, got nil") - } - if tc.wantMessage != "" && !strings.Contains(err.Error(), tc.wantMessage) { - t.Errorf( - "Delete: Expected error message containing '%s', got '%v'", - tc.wantMessage, - err, - ) - } + c.Require().Error(err, "Delete: Expected error, got nil") + c.False( + tc.wantMessage != "" && !strings.Contains(err.Error(), tc.wantMessage), + "Delete: Expected error message containing '%s', got '%v'", + tc.wantMessage, + err, + ) } } }) @@ -917,22 +908,19 @@ func TestTemplateValidator_ShardTemplatePoolNames(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) validator := NewTemplateValidator(fakeClient, "ShardTemplate") if name == "Non-ShardTemplate skipped" { // CellTemplate with Kind != ShardTemplate should skip validation v := NewTemplateValidator(fakeClient, "CellTemplate") _, err := v.ValidateCreate(t.Context(), &multigresv1alpha1.CellTemplate{}) - if err != nil { - t.Fatalf("Expected nil for non-ShardTemplate, got %v", err) - } + c.Require().NoError(err, "Expected nil for non-ShardTemplate, got") // ShardTemplate validator with wrong object type v2 := NewTemplateValidator(fakeClient, "ShardTemplate") _, err2 := v2.ValidateCreate(t.Context(), &multigresv1alpha1.CellTemplate{}) - if err2 != nil { - t.Fatalf("Expected nil for wrong object type, got %v", err2) - } + c.Require().NoError(err2, "Expected nil for wrong object type, got") return } @@ -959,18 +947,15 @@ func TestTemplateValidator_ShardTemplatePoolNames(t *testing.T) { } if tc.wantErr == "" { - if err != nil { - t.Errorf("%s: expected nil error, got %v", method, err) - } + c.NoError(err, "%s: expected nil error, got", method) } else { - if err == nil || !strings.Contains(err.Error(), tc.wantErr) { - t.Errorf( - "%s: expected error containing '%s', got %v", - method, - tc.wantErr, - err, - ) - } + c.False( + err == nil || !strings.Contains(err.Error(), tc.wantErr), + "%s: expected error containing '%s', got %v", + method, + tc.wantErr, + err, + ) } } }) @@ -1039,9 +1024,8 @@ func TestTemplateValidator_CellTemplateBuffer(t *testing.T) { // window 30s merges with no override: binary maxFailoverDuration // default (20s) applies at startup, so the update must be rejected. _, err := v.ValidateUpdate(t.Context(), stored, windowOnly) - if err == nil || !strings.Contains(err.Error(), "cell 'z1'") { - t.Errorf("expected merged validation error naming the cell, got %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "cell 'z1'"), "expected merged validation error naming the cell, got %v", err) }) t.Run("metadata-only update never gated on consumer state", func(t *testing.T) { @@ -1054,9 +1038,8 @@ func TestTemplateValidator_CellTemplateBuffer(t *testing.T) { v := NewTemplateValidator(c, "CellTemplate") relabeled := windowOnly.DeepCopy() relabeled.Annotations = map[string]string{"touched": "true"} - if _, err := v.ValidateUpdate(t.Context(), windowOnly, relabeled); err != nil { - t.Errorf("metadata-only update must be accepted, got %v", err) - } + _, err := v.ValidateUpdate(t.Context(), windowOnly, relabeled) + assert.NewCollecting(t).NoError(err, "metadata-only update must be accepted, got") }) t.Run("pre-existing breakage does not wedge template writes", func(t *testing.T) { @@ -1081,9 +1064,9 @@ func TestTemplateValidator_CellTemplateBuffer(t *testing.T) { updated := tpl(&multigresv1alpha1.GatewayBufferConfig{ MaxFailoverDuration: &metav1.Duration{Duration: 20 * time.Second}, }) - if _, err := v.ValidateUpdate(t.Context(), stored, updated); err != nil { - t.Errorf("write leaving a pre-broken consumer equally broken must pass, got %v", err) - } + _, err := v.ValidateUpdate(t.Context(), stored, updated) + assert.NewCollecting(t). + NoError(err, "write leaving a pre-broken consumer equally broken must pass, got") }) t.Run("terminating consumers do not gate template writes", func(t *testing.T) { @@ -1105,9 +1088,8 @@ func TestTemplateValidator_CellTemplateBuffer(t *testing.T) { c := fake.NewClientBuilder().WithScheme(scheme). WithObjects(terminating, stored).Build() v := NewTemplateValidator(c, "CellTemplate") - if _, err := v.ValidateUpdate(t.Context(), stored, windowOnly); err != nil { - t.Errorf("terminating consumer must not gate the write, got %v", err) - } + _, err := v.ValidateUpdate(t.Context(), stored, windowOnly) + assert.NewCollecting(t).NoError(err, "terminating consumer must not gate the write, got") }) t.Run( @@ -1235,9 +1217,9 @@ func TestTemplateValidator_CellTemplateBuffer(t *testing.T) { } c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(independent, stored).Build() v := NewTemplateValidator(c, "CellTemplate") - if _, err := v.ValidateDelete(t.Context(), stored); err != nil { - t.Errorf("deleting 'default' must be allowed when consumers stay valid, got %v", err) - } + _, err := v.ValidateDelete(t.Context(), stored) + assert.NewCollecting(t). + NoError(err, "deleting 'default' must be allowed when consumers stay valid, got") }) t.Run("update accepted when consumer override completes the config", func(t *testing.T) { @@ -1252,9 +1234,9 @@ func TestTemplateValidator_CellTemplateBuffer(t *testing.T) { }, )).Build() v := NewTemplateValidator(c, "CellTemplate") - if _, err := v.ValidateUpdate(t.Context(), nil, windowOnly); err != nil { - t.Errorf("override raises maxFailoverDuration, update must be accepted, got %v", err) - } + _, err := v.ValidateUpdate(t.Context(), nil, windowOnly) + assert.NewCollecting(t). + NoError(err, "override raises maxFailoverDuration, update must be accepted, got") }) } @@ -1312,6 +1294,7 @@ func TestChildResourceValidator(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) // Create context with admission request ctx := t.Context() @@ -1346,20 +1329,15 @@ func TestChildResourceValidator(t *testing.T) { _, err = validator.ValidateDelete(ctx, obj) } - if tc.wantAllowed && err != nil { - t.Fatalf("Expected allowed, got error: %v", err) - } + c.Require().False(tc.wantAllowed && err != nil, "Expected allowed, got error: %v", err) if !tc.wantAllowed { - if err == nil { - t.Fatal("Expected error, got nil") - } - if tc.wantMessage != "" && !strings.Contains(err.Error(), tc.wantMessage) { - t.Errorf( - "Expected error message containing '%s', got '%v'", - tc.wantMessage, - err, - ) - } + c.Require().Error(err, "Expected error, got nil") + c.False( + tc.wantMessage != "" && !strings.Contains(err.Error(), tc.wantMessage), + "Expected error message containing '%s', got '%v'", + tc.wantMessage, + err, + ) } }) } @@ -1367,9 +1345,7 @@ func TestChildResourceValidator(t *testing.T) { t.Run("Wrong Type", func(t *testing.T) { t.Parallel() _, err := validator.ValidateCreate(t.Context(), &TrulyOnlyRuntimeObject{}) - if err == nil { - t.Error("Expected error for wrong type, got nil") - } + assert.NewCollecting(t).Error(err, "Expected error for wrong type, got nil") }) } @@ -1405,22 +1381,22 @@ func TestValidateNoStorageShrink(t *testing.T) { oldObj := makeCluster("10Gi") newObj := makeCluster("20Gi") _, err := validateNoStorageShrink(oldObj, newObj) - if err != nil { - t.Fatalf("expected no error for storage grow, got: %v", err) - } + assert.NewAborting(t).NoError(err, "expected no error for storage grow, got") }) t.Run("rejects storage shrink", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) oldObj := makeCluster("20Gi") newObj := makeCluster("10Gi") _, err := validateNoStorageShrink(oldObj, newObj) - if err == nil { - t.Fatal("expected error for storage shrink, got nil") - } - if !strings.Contains(err.Error(), "cannot be decreased") { - t.Errorf("expected 'cannot be decreased' error, got: %v", err) - } + c.Require().Error(err, "expected error for storage shrink, got nil") + c.StrContains( + err.Error(), + "cannot be decreased", + "expected 'cannot be decreased' error, got: %v", + err, + ) }) t.Run("no-op when sizes equal", func(t *testing.T) { @@ -1428,34 +1404,27 @@ func TestValidateNoStorageShrink(t *testing.T) { oldObj := makeCluster("10Gi") newObj := makeCluster("10Gi") _, err := validateNoStorageShrink(oldObj, newObj) - if err != nil { - t.Fatalf("expected no error for equal sizes, got: %v", err) - } + assert.NewAborting(t).NoError(err, "expected no error for equal sizes, got") }) t.Run("ignores non-MultigresCluster objects", func(t *testing.T) { t.Parallel() _, err := validateNoStorageShrink(&TrulyOnlyRuntimeObject{}, &TrulyOnlyRuntimeObject{}) - if err != nil { - t.Fatalf("expected no error for wrong types, got: %v", err) - } + assert.NewAborting(t).NoError(err, "expected no error for wrong types, got") }) t.Run("ignores parse errors", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) oldObj := makeCluster("invalidQty") newObj := makeCluster("10Gi") _, err := validateNoStorageShrink(oldObj, newObj) - if err != nil { - t.Fatalf("expected no error when parsing fails, got: %v", err) - } + c.NoError(err, "expected no error when parsing fails, got") oldObj2 := makeCluster("10Gi") newObj2 := makeCluster("invalidQty") _, err2 := validateNoStorageShrink(oldObj2, newObj2) - if err2 != nil { - t.Fatalf("expected no error when parsing fails, got: %v", err2) - } + c.NoError(err2, "expected no error when parsing fails, got") }) t.Run("collects from shard overrides", func(t *testing.T) { @@ -1479,9 +1448,8 @@ func TestValidateNoStorageShrink(t *testing.T) { }, } sizes := collectPoolStorageSizes(obj) - if sizes["db1/tg1/s1/pool1"] != "42Gi" { - t.Errorf("expected 42Gi, got %v", sizes) - } + assert.NewCollecting(t). + Eq("42Gi", sizes["db1/tg1/s1/pool1"], "expected 42Gi, got %v", sizes) }) } @@ -1509,9 +1477,7 @@ func TestValidateEtcdReplicasImmutable(t *testing.T) { oldObj := makeCluster(ptr.To(int32(5)), false) newObj := makeCluster(ptr.To(int32(5)), false) _, err := validateEtcdReplicasImmutable(oldObj, newObj) - if err != nil { - t.Fatalf("expected nil error, got %v", err) - } + assert.NewAborting(t).NoError(err, "expected nil error, got") }) t.Run("rejects changed replica counts", func(t *testing.T) { @@ -1519,9 +1485,8 @@ func TestValidateEtcdReplicasImmutable(t *testing.T) { oldObj := makeCluster(ptr.To(int32(3)), false) newObj := makeCluster(ptr.To(int32(5)), false) _, err := validateEtcdReplicasImmutable(oldObj, newObj) - if err == nil || !strings.Contains(err.Error(), "etcd uses static bootstrap") { - t.Fatalf("expected immutable error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "etcd uses static bootstrap"), "expected immutable error, got %v", err) }) t.Run("allows transitions to or from external (0 replicas)", func(t *testing.T) { @@ -1530,21 +1495,16 @@ func TestValidateEtcdReplicasImmutable(t *testing.T) { oldObj := makeCluster(nil, true) newObj := makeCluster(ptr.To(int32(3)), false) _, err := validateEtcdReplicasImmutable(oldObj, newObj) - if err != nil { - t.Fatalf("expected nil error, got %v", err) - } + assert.NewAborting(t).NoError(err, "expected nil error, got") }) t.Run("ignores wrong types", func(t *testing.T) { t.Parallel() + c := assert.NewAborting(t) _, err := validateEtcdReplicasImmutable(&TrulyOnlyRuntimeObject{}, makeCluster(nil, false)) - if err != nil { - t.Fatalf("expected nil error, got %v", err) - } + c.NoError(err, "expected nil error, got") _, err = validateEtcdReplicasImmutable(makeCluster(nil, false), &TrulyOnlyRuntimeObject{}) - if err != nil { - t.Fatalf("expected nil error, got %v", err) - } + c.NoError(err, "expected nil error, got") }) } @@ -1585,9 +1545,8 @@ func TestMultigresClusterValidator_ValidateUpdate(t *testing.T) { t.Run("bubbles up base validation error", func(t *testing.T) { t.Parallel() _, err := validator.ValidateUpdate(t.Context(), baseCluster, baseCluster) - if err == nil || !strings.Contains(err.Error(), "not found") { - t.Fatalf("expected validation error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "not found"), "expected validation error, got %v", err) }) validCluster := &multigresv1alpha1.MultigresCluster{ @@ -1629,9 +1588,8 @@ func TestMultigresClusterValidator_ValidateUpdate(t *testing.T) { t.Run("bubbles up shrink error", func(t *testing.T) { t.Parallel() _, err := validator.ValidateUpdate(t.Context(), largeCluster, shrunkCluster) - if err == nil || !strings.Contains(err.Error(), "shrink is not supported") { - t.Fatalf("expected shrink error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "shrink is not supported"), "expected shrink error, got %v", err) }) largeEtcdCluster := validCluster.DeepCopy() @@ -1646,9 +1604,8 @@ func TestMultigresClusterValidator_ValidateUpdate(t *testing.T) { t.Run("bubbles up etcd error", func(t *testing.T) { t.Parallel() _, err := validator.ValidateUpdate(t.Context(), largeEtcdCluster, smallEtcdCluster) - if err == nil || !strings.Contains(err.Error(), "etcd uses static bootstrap") { - t.Fatalf("expected etcd error, got %v", err) - } + assert.NewAborting(t). + False(err == nil || !strings.Contains(err.Error(), "etcd uses static bootstrap"), "expected etcd error, got %v", err) }) } @@ -1695,18 +1652,19 @@ func TestValidatePostgresConfig(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) err := validatePostgresConfig(tc.cluster) if tc.wantErr == "" { - if err != nil { - t.Errorf("expected nil, got %v", err) - } + c.NoError(err, "expected nil, got") } else if err == nil || !strings.Contains(err.Error(), tc.wantErr) { t.Errorf("expected error containing %q, got %v", tc.wantErr, err) } // The error must identify the shard location. - if tc.wantErr != "" && err != nil && !strings.Contains(err.Error(), "shard") { - t.Errorf("error should name the shard, got %v", err) - } + c.False( + tc.wantErr != "" && err != nil && !strings.Contains(err.Error(), "shard"), + "error should name the shard, got %v", + err, + ) }) } } diff --git a/pkg/webhook/integration_test.go b/pkg/webhook/integration_test.go index f1694830..db3383a3 100644 --- a/pkg/webhook/integration_test.go +++ b/pkg/webhook/integration_test.go @@ -9,7 +9,6 @@ import ( "fmt" "os" "path/filepath" - "strings" "testing" "time" @@ -33,6 +32,8 @@ import ( "github.com/multigres/multigres-operator/pkg/resolver" "github.com/multigres/multigres-operator/pkg/util/metadata" multigreswebhook "github.com/multigres/multigres-operator/pkg/webhook" + + "github.com/multigres/testkit/assert" ) const ( @@ -171,7 +172,9 @@ func createDefaults(c client.Client) error { &multigresv1alpha1.CellTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "default", Namespace: testNamespace}, Spec: multigresv1alpha1.CellTemplateSpec{ - Multigateway: &multigresv1alpha1.MultigatewaySpec{StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}}, + Multigateway: &multigresv1alpha1.MultigatewaySpec{ + StatelessSpec: multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, + }, }, }, &multigresv1alpha1.ShardTemplate{ @@ -180,8 +183,10 @@ func createDefaults(c client.Client) error { }, &storagev1.StorageClass{ ObjectMeta: metav1.ObjectMeta{ - Name: "standard", - Annotations: map[string]string{"storageclass.kubernetes.io/is-default-class": "true"}, + Name: "standard", + Annotations: map[string]string{ + "storageclass.kubernetes.io/is-default-class": "true", + }, }, Provisioner: "k8s.io/fake", }, @@ -209,7 +214,11 @@ func waitForClusterList(t *testing.T, c client.Client, clusterName string) { t.Fatalf("Timeout waiting for cluster '%s' to appear in List() cache", clusterName) case <-ticker.C: clusters := &multigresv1alpha1.MultigresClusterList{} - if err := c.List(context.Background(), clusters, client.InNamespace(testNamespace)); err != nil { + if err := c.List( + context.Background(), + clusters, + client.InNamespace(testNamespace), + ); err != nil { continue } for _, item := range clusters.Items { @@ -247,6 +256,7 @@ func setTestShardPostgresPasswordSecretRef(shard *multigresv1alpha1.Shard) { func TestWebhook_Mutation(t *testing.T) { t.Run("Should Inject System Catalog and Defaults", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "mutation-test", @@ -259,31 +269,34 @@ func TestWebhook_Mutation(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.Require().NoError(k8sClient.Create(ctx, cluster), "Failed to create cluster") fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatalf("Failed to get cluster: %v", err) - } + c.Require(). + NoError(k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched), "Failed to get cluster") if len(fetched.Spec.Databases) == 0 { t.Error("Webhook failed to inject 'postgres' database") } else { db := fetched.Spec.Databases[0] - if db.Name != "postgres" || !db.Default { - t.Errorf("System database incorrect. Got Name=%s Default=%v", db.Name, db.Default) - } - } - - if fetched.Spec.TemplateDefaults.CoreTemplate != "default" { - t.Errorf("Expected CoreTemplate to be promoted to 'default', got %q", fetched.Spec.TemplateDefaults.CoreTemplate) - } - - if fetched.Spec.Multiadmin != nil { - t.Error("Expected spec.multiadmin to be nil (preserved dynamic link to template, no overrides provided)") - } + c.False( + db.Name != "postgres" || !db.Default, + "System database incorrect. Got Name=%s Default=%v", + db.Name, + db.Default, + ) + } + + c.Eq( + "default", + fetched.Spec.TemplateDefaults.CoreTemplate, + "Expected CoreTemplate to be promoted to 'default', got", + ) + + c.Nil( + fetched.Spec.Multiadmin, + "Expected spec.multiadmin to be nil (preserved dynamic link to template, no overrides provided)", + ) }) } @@ -304,12 +317,12 @@ func TestWebhook_Validation(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err == nil { - t.Fatal("Expected error creating cluster with missing template, got nil") - } + assert.NewAborting(t). + Error(k8sClient.Create(ctx, cluster), "Expected error creating cluster with missing template, got nil") }) t.Run("Should Reject Unknown Postgres Parameter", func(t *testing.T) { + c := assert.NewAborting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "bad-guc", Namespace: testNamespace}, Spec: multigresv1alpha1.MultigresClusterSpec{ @@ -336,24 +349,24 @@ func TestWebhook_Validation(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(ctx, cluster) - if err == nil { - t.Fatal("expected rejection for unknown postgres parameter") - } - if !strings.Contains(err.Error(), "unknown parameter") { - t.Fatalf("expected 'unknown parameter' error, got: %v", err) - } + c.Error(err, "expected rejection for unknown postgres parameter") + c.StrContains( + err.Error(), + "unknown parameter", + "expected 'unknown parameter' error, got: %v", + err, + ) }) } func TestWebhook_TemplateProtection(t *testing.T) { t.Run("Should Prevent Deleting In-Use Template", func(t *testing.T) { + c := assert.NewAborting(t) tpl := &multigresv1alpha1.CoreTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "production-core", Namespace: testNamespace}, Spec: multigresv1alpha1.CoreTemplateSpec{}, } - if err := k8sClient.Create(ctx, tpl); err != nil { - t.Fatalf("Failed to create template: %v", err) - } + c.NoError(k8sClient.Create(ctx, tpl), "Failed to create template") cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ @@ -369,9 +382,7 @@ func TestWebhook_TemplateProtection(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.NoError(k8sClient.Create(ctx, cluster), "Failed to create cluster") // CRITICAL: Use cachedClient here. // We must wait until the Webhook's internal cache sees the cluster. @@ -379,9 +390,7 @@ func TestWebhook_TemplateProtection(t *testing.T) { // would see 0 clusters and allow the delete. waitForClusterList(t, cachedClient, "prod-cluster") - if err := k8sClient.Delete(ctx, tpl); err == nil { - t.Fatal("Expected error deleting in-use template, got nil") - } + c.Error(k8sClient.Delete(ctx, tpl), "Expected error deleting in-use template, got nil") }) } @@ -394,9 +403,8 @@ func TestWebhook_ChildResourceProtection(t *testing.T) { }, } - if err := k8sClient.Create(ctx, cell); err == nil { - t.Fatal("Expected error creating Child Resource directly, got nil") - } + assert.NewAborting(t). + Error(k8sClient.Create(ctx, cell), "Expected error creating Child Resource directly, got nil") }) } @@ -406,6 +414,7 @@ func TestWebhook_ChildResourceProtection(t *testing.T) { func TestWebhook_CellAppendOnly(t *testing.T) { t.Run("Should Reject Cell Removal", func(t *testing.T) { + c := assert.NewAborting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "cell-removal-test", @@ -421,28 +430,24 @@ func TestWebhook_CellAppendOnly(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.NoError(k8sClient.Create(ctx, cluster), "Failed to create cluster") fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatalf("Failed to get cluster: %v", err) - } + c.NoError( + k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched), + "Failed to get cluster", + ) fetched.Spec.Cells = []multigresv1alpha1.CellConfig{ {Name: "cell-a", ZoneID: "use1-az1"}, } err := k8sClient.Update(ctx, fetched) - if err == nil { - t.Fatal("Expected error removing a cell, got nil") - } - if !strings.Contains(err.Error(), "Append-Only") { - t.Fatalf("Expected 'Append-Only' in error, got: %v", err) - } + c.Error(err, "Expected error removing a cell, got nil") + c.StrContains(err.Error(), "Append-Only", "Expected 'Append-Only' in error, got: %v", err) }) t.Run("Should Reject Cell Rename", func(t *testing.T) { + c := assert.NewAborting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "cell-rename-test", @@ -457,29 +462,25 @@ func TestWebhook_CellAppendOnly(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.NoError(k8sClient.Create(ctx, cluster), "Failed to create cluster") fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatalf("Failed to get cluster: %v", err) - } + c.NoError( + k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched), + "Failed to get cluster", + ) // Replace the cell with a different name — effectively a rename fetched.Spec.Cells = []multigresv1alpha1.CellConfig{ {Name: "renamed-cell", ZoneID: "use1-az1"}, } err := k8sClient.Update(ctx, fetched) - if err == nil { - t.Fatal("Expected error renaming a cell, got nil") - } - if !strings.Contains(err.Error(), "Append-Only") { - t.Fatalf("Expected 'Append-Only' in error, got: %v", err) - } + c.Error(err, "Expected error renaming a cell, got nil") + c.StrContains(err.Error(), "Append-Only", "Expected 'Append-Only' in error, got: %v", err) }) t.Run("Should Allow Adding New Cells", func(t *testing.T) { + c := assert.NewAborting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ Name: "cell-add-test", @@ -494,21 +495,18 @@ func TestWebhook_CellAppendOnly(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.NoError(k8sClient.Create(ctx, cluster), "Failed to create cluster") fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatalf("Failed to get cluster: %v", err) - } + c.NoError( + k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched), + "Failed to get cluster", + ) fetched.Spec.Cells = append(fetched.Spec.Cells, multigresv1alpha1.CellConfig{ Name: "cell-y", ZoneID: "use1-az2", }) - if err := k8sClient.Update(ctx, fetched); err != nil { - t.Fatalf("Expected appending a cell to succeed, got: %v", err) - } + c.NoError(k8sClient.Update(ctx, fetched), "Expected appending a cell to succeed, got") }) } @@ -518,6 +516,7 @@ func TestWebhook_CellAppendOnly(t *testing.T) { func TestWebhook_OverridePrecedence(t *testing.T) { t.Run("Inline Spec Should Override Template", func(t *testing.T) { + c := assert.NewAborting(t) tplName := "base-template" tpl := &multigresv1alpha1.CoreTemplate{ ObjectMeta: metav1.ObjectMeta{Name: tplName, Namespace: testNamespace}, @@ -525,9 +524,7 @@ func TestWebhook_OverridePrecedence(t *testing.T) { Multiadmin: &multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(1))}, }, } - if err := k8sClient.Create(ctx, tpl); err != nil { - t.Fatalf("Failed to create template: %v", err) - } + c.NoError(k8sClient.Create(ctx, tpl), "Failed to create template") cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "override-test", Namespace: testNamespace}, @@ -542,32 +539,31 @@ func TestWebhook_OverridePrecedence(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Failed to create cluster: %v", err) - } + c.NoError(k8sClient.Create(ctx, cluster), "Failed to create cluster") fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatal(err) - } + c.NoError(k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched)) - if fetched.Spec.Multiadmin.Spec.Replicas == nil || *fetched.Spec.Multiadmin.Spec.Replicas != 3 { - t.Errorf("Expected inline overrides (3) to win over template (1), got: %v", fetched.Spec.Multiadmin.Spec.Replicas) + if fetched.Spec.Multiadmin.Spec.Replicas == nil || + *fetched.Spec.Multiadmin.Spec.Replicas != 3 { + t.Errorf( + "Expected inline overrides (3) to win over template (1), got: %v", + fetched.Spec.Multiadmin.Spec.Replicas, + ) } }) } func TestWebhook_SpecificRefPrecedence(t *testing.T) { t.Run("Specific TemplateRef Should NOT be Expanded (Spec Conflict)", func(t *testing.T) { + c := assert.NewCollecting(t) specTpl := &multigresv1alpha1.CoreTemplate{ ObjectMeta: metav1.ObjectMeta{Name: "specific-large", Namespace: testNamespace}, Spec: multigresv1alpha1.CoreTemplateSpec{ Multiadmin: &multigresv1alpha1.StatelessSpec{Replicas: ptr.To(int32(5))}, }, } - if err := k8sClient.Create(ctx, specTpl); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(ctx, specTpl)) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "ref-precedence-test", Namespace: testNamespace}, @@ -579,26 +575,26 @@ func TestWebhook_SpecificRefPrecedence(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(ctx, cluster)) fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatal(err) - } - - if fetched.Spec.Multiadmin.Spec != nil { - t.Errorf("Expected Multiadmin.Spec to be nil when TemplateRef is set, but got: %v", fetched.Spec.Multiadmin.Spec) - } - if fetched.Spec.Multiadmin.TemplateRef != "specific-large" { - t.Errorf("Expected TemplateRef to be preserved") - } + c.Require().NoError(k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched)) + + c.Nil( + fetched.Spec.Multiadmin.Spec, + "Expected Multiadmin.Spec to be nil when TemplateRef is set, but got", + ) + c.Eq( + "specific-large", + fetched.Spec.Multiadmin.TemplateRef, + "Expected TemplateRef to be preserved", + ) }) } func TestWebhook_SystemCatalogIdempotency(t *testing.T) { t.Run("Should Not Duplicate Existing System Catalog", func(t *testing.T) { + c := assert.NewCollecting(t) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{Name: "idempotency-test", Namespace: testNamespace}, Spec: multigresv1alpha1.MultigresClusterSpec{ @@ -617,35 +613,26 @@ func TestWebhook_SystemCatalogIdempotency(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Create(ctx, cluster)) fetched := &multigresv1alpha1.MultigresCluster{} - if err := k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched); err != nil { - t.Fatal(err) - } + c.Require().NoError(k8sClient.Get(ctx, client.ObjectKeyFromObject(cluster), fetched)) - if len(fetched.Spec.Databases) != 1 { - t.Errorf("Expected 1 database, got %d", len(fetched.Spec.Databases)) - } + c.Len(fetched.Spec.Databases, 1, "Expected 1 database, got %d", len(fetched.Spec.Databases)) tgList := fetched.Spec.Databases[0].TableGroups - if len(tgList) != 1 { - t.Errorf("Expected 1 tablegroup, got %d", len(tgList)) - } + c.Len(tgList, 1, "Expected 1 tablegroup, got %d", len(tgList)) }) } func TestWebhook_DeepTemplateProtection(t *testing.T) { t.Run("Should Protect Deeply Nested ShardTemplate", func(t *testing.T) { + c := assert.NewAborting(t) stName := "sensitive-shard-tpl" st := &multigresv1alpha1.ShardTemplate{ ObjectMeta: metav1.ObjectMeta{Name: stName, Namespace: testNamespace}, Spec: multigresv1alpha1.ShardTemplateSpec{}, } - if err := k8sClient.Create(ctx, st); err != nil { - t.Fatal(err) - } + c.NoError(k8sClient.Create(ctx, st)) cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ @@ -676,29 +663,22 @@ func TestWebhook_DeepTemplateProtection(t *testing.T) { }, } setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatal(err) - } + c.NoError(k8sClient.Create(ctx, cluster)) waitForClusterList(t, cachedClient, "deep-ref-cluster") - if err := k8sClient.Delete(ctx, st); err == nil { - t.Fatal("Expected error deleting in-use ShardTemplate, got nil") - } + c.Error(k8sClient.Delete(ctx, st), "Expected error deleting in-use ShardTemplate, got nil") }) } func TestWebhook_StorageClassValidation(t *testing.T) { t.Run("Should Reject When No Default SC and No Explicit Class", func(t *testing.T) { + c := assert.NewAborting(t) // Remove the default StorageClass annotation sc := &storagev1.StorageClass{} - if err := k8sClient.Get(ctx, client.ObjectKey{Name: "standard"}, sc); err != nil { - t.Fatalf("Failed to get SC: %v", err) - } + c.NoError(k8sClient.Get(ctx, client.ObjectKey{Name: "standard"}, sc), "Failed to get SC") sc.Annotations["storageclass.kubernetes.io/is-default-class"] = "false" - if err := k8sClient.Update(ctx, sc); err != nil { - t.Fatalf("Failed to update SC: %v", err) - } + c.NoError(k8sClient.Update(ctx, sc), "Failed to update SC") cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ @@ -713,30 +693,26 @@ func TestWebhook_StorageClassValidation(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) err := k8sClient.Create(ctx, cluster) - if err == nil { - t.Fatal("Expected rejection due to missing default StorageClass, got nil") - } - if !strings.Contains(err.Error(), "no default StorageClass found") { - t.Fatalf("Expected 'no default StorageClass found' in error, got: %v", err) - } + c.Error(err, "Expected rejection due to missing default StorageClass, got nil") + c.StrContains( + err.Error(), + "no default StorageClass found", + "Expected 'no default StorageClass found' in error, got: %v", + err, + ) // Restore the default SC for subsequent tests sc.Annotations["storageclass.kubernetes.io/is-default-class"] = "true" - if err := k8sClient.Update(ctx, sc); err != nil { - t.Fatalf("Failed to restore SC: %v", err) - } + c.NoError(k8sClient.Update(ctx, sc), "Failed to restore SC") }) t.Run("Should Accept With Explicit Class Even Without Default SC", func(t *testing.T) { + c := assert.NewAborting(t) // Remove the default StorageClass annotation sc := &storagev1.StorageClass{} - if err := k8sClient.Get(ctx, client.ObjectKey{Name: "standard"}, sc); err != nil { - t.Fatalf("Failed to get SC: %v", err) - } + c.NoError(k8sClient.Get(ctx, client.ObjectKey{Name: "standard"}, sc), "Failed to get SC") sc.Annotations["storageclass.kubernetes.io/is-default-class"] = "false" - if err := k8sClient.Update(ctx, sc); err != nil { - t.Fatalf("Failed to update SC: %v", err) - } + c.NoError(k8sClient.Update(ctx, sc), "Failed to update SC") cluster := &multigresv1alpha1.MultigresCluster{ ObjectMeta: metav1.ObjectMeta{ @@ -767,8 +743,11 @@ func TestWebhook_StorageClassValidation(t *testing.T) { Spec: &multigresv1alpha1.ShardInlineSpec{ Pools: map[multigresv1alpha1.PoolName]multigresv1alpha1.PoolSpec{ "default": { - Type: "readWrite", - Storage: multigresv1alpha1.StorageSpec{Class: "manual", Size: "10Gi"}, + Type: "readWrite", + Storage: multigresv1alpha1.StorageSpec{ + Class: "manual", + Size: "10Gi", + }, }, }, }, @@ -780,14 +759,13 @@ func TestWebhook_StorageClassValidation(t *testing.T) { setTestPostgresPasswordSecretRef(cluster) - if err := k8sClient.Create(ctx, cluster); err != nil { - t.Fatalf("Expected acceptance with explicit storage class, got: %v", err) - } + c.NoError( + k8sClient.Create(ctx, cluster), + "Expected acceptance with explicit storage class, got", + ) // Restore the default SC sc.Annotations["storageclass.kubernetes.io/is-default-class"] = "true" - if err := k8sClient.Update(ctx, sc); err != nil { - t.Fatalf("Failed to restore SC: %v", err) - } + c.NoError(k8sClient.Update(ctx, sc), "Failed to restore SC") }) } diff --git a/pkg/webhook/pki_test.go b/pkg/webhook/pki_test.go index 8a803458..71d2b9c5 100644 --- a/pkg/webhook/pki_test.go +++ b/pkg/webhook/pki_test.go @@ -15,17 +15,16 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/client/fake" "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + "github.com/multigres/testkit/assert" ) func pkiScheme(tb testing.TB) *runtime.Scheme { tb.Helper() + c := assert.NewAborting(tb) s := runtime.NewScheme() - if err := admissionregistrationv1.AddToScheme(s); err != nil { - tb.Fatal(err) - } - if err := appsv1.AddToScheme(s); err != nil { - tb.Fatal(err) - } + c.NoError(admissionregistrationv1.AddToScheme(s)) + c.NoError(appsv1.AddToScheme(s)) return s } @@ -40,6 +39,7 @@ func TestPatchWebhookCABundle(t *testing.T) { t.Run("Patches Both Webhook Configs via SSA", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mutating := &admissionregistrationv1.MutatingWebhookConfiguration{ ObjectMeta: metav1.ObjectMeta{Name: MutatingWebhookName}, @@ -69,57 +69,44 @@ func TestPatchWebhookCABundle(t *testing.T) { WithObjects(mutating, validating). Build() - if err := PatchWebhookCABundle(context.Background(), cl, caBundle); err != nil { - t.Fatalf("unexpected error: %v", err) - } + c.Require(). + NoError(PatchWebhookCABundle(context.Background(), cl, caBundle), "unexpected error") // Verify mutating: caBundle + annotation got := &admissionregistrationv1.MutatingWebhookConfiguration{} - if err := cl.Get( + c.Require().NoError(cl.Get( context.Background(), client.ObjectKeyFromObject(mutating), got, - ); err != nil { - t.Fatal(err) - } - if string(got.Webhooks[0].ClientConfig.CABundle) != string(caBundle) { - t.Errorf( - "mutating CABundle = %q, want %q", - got.Webhooks[0].ClientConfig.CABundle, - caBundle, - ) - } - if got.Annotations[CertStrategyAnnotation] != CertStrategySelfSigned { - t.Errorf( - "mutating annotation = %q, want %q", - got.Annotations[CertStrategyAnnotation], - CertStrategySelfSigned, - ) - } + )) + c.Eq( + string(caBundle), + string(got.Webhooks[0].ClientConfig.CABundle), + "mutating CABundle = %q, want %q", + got.Webhooks[0].ClientConfig.CABundle, + caBundle, + ) + c.Eq(CertStrategySelfSigned, got.Annotations[CertStrategyAnnotation], "mutating annotation") // Verify validating: caBundle + annotation gotV := &admissionregistrationv1.ValidatingWebhookConfiguration{} - if err := cl.Get( + c.Require().NoError(cl.Get( context.Background(), client.ObjectKeyFromObject(validating), gotV, - ); err != nil { - t.Fatal(err) - } - if string(gotV.Webhooks[0].ClientConfig.CABundle) != string(caBundle) { - t.Errorf( - "validating CABundle = %q, want %q", - gotV.Webhooks[0].ClientConfig.CABundle, - caBundle, - ) - } - if gotV.Annotations[CertStrategyAnnotation] != CertStrategySelfSigned { - t.Errorf( - "validating annotation = %q, want %q", - gotV.Annotations[CertStrategyAnnotation], - CertStrategySelfSigned, - ) - } + )) + c.Eq( + string(caBundle), + string(gotV.Webhooks[0].ClientConfig.CABundle), + "validating CABundle = %q, want %q", + gotV.Webhooks[0].ClientConfig.CABundle, + caBundle, + ) + c.Eq( + CertStrategySelfSigned, + gotV.Annotations[CertStrategyAnnotation], + "validating annotation", + ) }) t.Run("Tolerates NotFound", func(t *testing.T) { @@ -127,9 +114,8 @@ func TestPatchWebhookCABundle(t *testing.T) { cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).Build() - if err := PatchWebhookCABundle(context.Background(), cl, caBundle); err != nil { - t.Fatalf("expected no error for missing configs, got: %v", err) - } + assert.NewAborting(t). + NoError(PatchWebhookCABundle(context.Background(), cl, caBundle), "expected no error for missing configs, got") }) t.Run("Error: Mutating Get Failure", func(t *testing.T) { @@ -148,9 +134,8 @@ func TestPatchWebhookCABundle(t *testing.T) { Build() err := PatchWebhookCABundle(context.Background(), cl, caBundle) - if err == nil || !strings.Contains(err.Error(), "failed to get mutating webhook config") { - t.Errorf("expected mutating get error, got: %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to get mutating webhook config"), "expected mutating get error, got: %v", err) }) t.Run("Error: Mutating Patch Failure", func(t *testing.T) { @@ -182,9 +167,8 @@ func TestPatchWebhookCABundle(t *testing.T) { Build() err := PatchWebhookCABundle(context.Background(), cl, caBundle) - if err == nil || !strings.Contains(err.Error(), "patch fail") { - t.Errorf("expected patch error, got: %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "patch fail"), "expected patch error, got: %v", err) }) t.Run("Error: Validating Get Failure", func(t *testing.T) { @@ -203,9 +187,8 @@ func TestPatchWebhookCABundle(t *testing.T) { Build() err := PatchWebhookCABundle(context.Background(), cl, caBundle) - if err == nil || !strings.Contains(err.Error(), "failed to get validating webhook config") { - t.Errorf("expected validating get error, got: %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to get validating webhook config"), "expected validating get error, got: %v", err) }) t.Run("Error: Validating Patch Failure", func(t *testing.T) { @@ -237,9 +220,8 @@ func TestPatchWebhookCABundle(t *testing.T) { Build() err := PatchWebhookCABundle(context.Background(), cl, caBundle) - if err == nil || !strings.Contains(err.Error(), "patch fail") { - t.Errorf("expected patch error, got: %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "patch fail"), "expected patch error, got: %v", err) }) t.Run("Skips Patching When No Webhooks", func(t *testing.T) { @@ -259,9 +241,8 @@ func TestPatchWebhookCABundle(t *testing.T) { WithObjects(mutating, validating). Build() - if err := PatchWebhookCABundle(context.Background(), cl, caBundle); err != nil { - t.Fatalf("unexpected error: %v", err) - } + assert.NewAborting(t). + NoError(PatchWebhookCABundle(context.Background(), cl, caBundle), "unexpected error") }) } @@ -279,9 +260,8 @@ func TestHasCertAnnotation(t *testing.T) { } cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).WithObjects(mutating).Build() - if !HasCertAnnotation(context.Background(), cl) { - t.Error("expected true when mutating has annotation") - } + assert.NewCollecting(t). + True(HasCertAnnotation(context.Background(), cl), "expected true when mutating has annotation") }) t.Run("True When Validating Has Annotation", func(t *testing.T) { @@ -295,9 +275,8 @@ func TestHasCertAnnotation(t *testing.T) { } cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).WithObjects(validating).Build() - if !HasCertAnnotation(context.Background(), cl) { - t.Error("expected true when validating has annotation") - } + assert.NewCollecting(t). + True(HasCertAnnotation(context.Background(), cl), "expected true when validating has annotation") }) t.Run("False When No Annotation", func(t *testing.T) { @@ -308,18 +287,16 @@ func TestHasCertAnnotation(t *testing.T) { } cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).WithObjects(mutating).Build() - if HasCertAnnotation(context.Background(), cl) { - t.Error("expected false when no annotation") - } + assert.NewCollecting(t). + False(HasCertAnnotation(context.Background(), cl), "expected false when no annotation") }) t.Run("False When Configs Missing", func(t *testing.T) { t.Parallel() cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).Build() - if HasCertAnnotation(context.Background(), cl) { - t.Error("expected false when configs don't exist") - } + assert.NewCollecting(t). + False(HasCertAnnotation(context.Background(), cl), "expected false when configs don't exist") }) } @@ -331,6 +308,7 @@ func TestFindOperatorDeployment(t *testing.T) { t.Run("Found by Labels", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) dep := &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ @@ -342,16 +320,17 @@ func TestFindOperatorDeployment(t *testing.T) { cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).WithObjects(dep).Build() got, err := FindOperatorDeployment(context.Background(), cl, namespace, labels, "") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got == nil || got.Name != "my-operator" { - t.Errorf("expected deployment 'my-operator', got %v", got) - } + c.Require().NoError(err, "unexpected error") + c.False( + got == nil || got.Name != "my-operator", + "expected deployment 'my-operator', got %v", + got, + ) }) t.Run("Found by Name", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) dep := &appsv1.Deployment{ ObjectMeta: metav1.ObjectMeta{ @@ -368,42 +347,37 @@ func TestFindOperatorDeployment(t *testing.T) { nil, "explicit-name", ) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got == nil || got.Name != "explicit-name" { - t.Errorf("expected deployment 'explicit-name', got %v", got) - } + c.Require().NoError(err, "unexpected error") + c.False( + got == nil || got.Name != "explicit-name", + "expected deployment 'explicit-name', got %v", + got, + ) }) t.Run("Not Found Returns nil", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).Build() got, err := FindOperatorDeployment(context.Background(), cl, namespace, nil, "nonexistent") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got != nil { - t.Errorf("expected nil, got %v", got) - } + c.Require().NoError(err, "unexpected error") + c.Nil(got, "expected nil, got") }) t.Run("No Labels No Name Returns nil", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).Build() got, err := FindOperatorDeployment(context.Background(), cl, namespace, nil, "") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got != nil { - t.Errorf("expected nil, got %v", got) - } + c.Require().NoError(err, "unexpected error") + c.Nil(got, "expected nil, got") }) t.Run("Multiple Matches Picks Oldest", func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) older := metav1.NewTime(time.Now().Add(-1 * time.Hour)) newer := metav1.NewTime(time.Now()) @@ -426,12 +400,12 @@ func TestFindOperatorDeployment(t *testing.T) { cl := fake.NewClientBuilder().WithScheme(pkiScheme(t)).WithObjects(dep1, dep2).Build() got, err := FindOperatorDeployment(context.Background(), cl, namespace, labels, "") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if got == nil || got.Name != "op-older" { - t.Errorf("expected oldest deployment 'op-older', got %v", got) - } + c.Require().NoError(err, "unexpected error") + c.False( + got == nil || got.Name != "op-older", + "expected oldest deployment 'op-older', got %v", + got, + ) }) t.Run("Error: List Failure", func(t *testing.T) { @@ -447,9 +421,8 @@ func TestFindOperatorDeployment(t *testing.T) { Build() _, err := FindOperatorDeployment(context.Background(), cl, namespace, labels, "") - if err == nil || !strings.Contains(err.Error(), "failed to list deployments by labels") { - t.Errorf("expected list error, got: %v", err) - } + assert.NewCollecting(t). + False(err == nil || !strings.Contains(err.Error(), "failed to list deployments by labels"), "expected list error, got: %v", err) }) t.Run("Error: Get by Name Failure", func(t *testing.T) { @@ -465,9 +438,10 @@ func TestFindOperatorDeployment(t *testing.T) { Build() _, err := FindOperatorDeployment(context.Background(), cl, namespace, nil, "some-name") - if err == nil || - !strings.Contains(err.Error(), "failed to get operator deployment by name") { - t.Errorf("expected get error, got: %v", err) - } + assert.NewCollecting(t).False(err == nil || + !strings.Contains( + err.Error(), + "failed to get operator deployment by name", + ), "expected get error, got: %v", err) }) } diff --git a/pkg/webhook/setup_test.go b/pkg/webhook/setup_test.go index 73083e37..3aa8df6e 100644 --- a/pkg/webhook/setup_test.go +++ b/pkg/webhook/setup_test.go @@ -15,6 +15,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client/fake" "sigs.k8s.io/controller-runtime/pkg/manager" "sigs.k8s.io/controller-runtime/pkg/webhook" + + "github.com/multigres/testkit/assert" ) // mockManager implements manager.Manager for testing. @@ -66,9 +68,7 @@ func (s *mockServer) WebhookMux() *http.ServeMux { return http.N func setupTestDeps(tb testing.TB) (*runtime.Scheme, client.Client) { tb.Helper() s := runtime.NewScheme() - if err := multigresv1alpha1.AddToScheme(s); err != nil { - tb.Fatalf("Failed to add scheme: %v", err) - } + assert.NewAborting(tb).NoError(multigresv1alpha1.AddToScheme(s), "Failed to add scheme") c := fake.NewClientBuilder().WithScheme(s).Build() return s, c } @@ -207,24 +207,24 @@ func TestSetup(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) mgr := tc.mgrFunc(t) err := Setup(mgr, tc.resolver, tc.opts) if tc.expectError != "" { - if err == nil { - t.Fatalf("Expected error containing %q, got nil", tc.expectError) - } - if diff := cmp.Diff( + c.Require().Error(err, "Expected error containing %q, got nil", tc.expectError) + diff := cmp.Diff( true, strings.Contains(err.Error(), tc.expectError), - ); diff != "" { - t.Errorf( - "Error message mismatch (-got +want matching check):\n%s\nGot error: %v", - diff, - err, - ) - } + ) + c.Eq( + "", + diff, + "Error message mismatch (-got +want matching check):\n%s\nGot error: %v", + diff, + err, + ) } else if err != nil { t.Errorf("Unexpected error: %v", err) } diff --git a/test/e2e/dedicated/deletion/deletion_test.go b/test/e2e/dedicated/deletion/deletion_test.go index 712abd86..e8a082ea 100644 --- a/test/e2e/dedicated/deletion/deletion_test.go +++ b/test/e2e/dedicated/deletion/deletion_test.go @@ -16,16 +16,17 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestClusterDeletion verifies that deleting a MultigresCluster triggers // cascading deletion of all child resources (CRDs and Kubernetes resources). func TestClusterDeletion(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Load and create the minimal sample. @@ -36,9 +37,7 @@ func TestClusterDeletion(t *testing.T) { WhenDeleted: multigresv1alpha1.RetainPVCRetentionPolicy, WhenScaled: multigresv1alpha1.RetainPVCRetentionPolicy, } - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for full provisioning. cluster.WaitForAllPodsReady(t, ns) @@ -46,20 +45,21 @@ func TestClusterDeletion(t *testing.T) { // Delete the cluster. clusterKey := client.ObjectKeyFromObject(cr) - if err := c.Delete(ctx, cr); err != nil { - t.Fatalf("delete MultigresCluster: %v", err) - } + ck.NoError(c.Delete(ctx, cr), "delete MultigresCluster") // Wait for the MultigresCluster object to disappear. pollCtx, cancel := context.WithTimeout(ctx, 5*time.Minute) defer cancel() - err = wait.PollUntilContextCancel(pollCtx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - err := c.Get(ctx, clusterKey, &multigresv1alpha1.MultigresCluster{}) - return apierrors.IsNotFound(err), nil - }) - if err != nil { - t.Fatalf("MultigresCluster not deleted: %v", err) - } + err = wait.PollUntilContextCancel( + pollCtx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + err := c.Get(ctx, clusterKey, &multigresv1alpha1.MultigresCluster{}) + return apierrors.IsNotFound(err), nil + }, + ) + ck.NoError(err, "MultigresCluster not deleted") // Verify child CRDs are cleaned up. framework.WaitForEmpty(t, c, ns, @@ -112,16 +112,16 @@ func TestClusterDeletion(t *testing.T) { // Verify data (non-topo) PVCs are retained under the Retain policy forced // above, only topo PVCs are force-deleted by cluster cleanup. pvcList := &corev1.PersistentVolumeClaimList{} - if err := c.List(ctx, pvcList, client.InNamespace(ns)); err != nil { - t.Fatalf("list PVCs: %v", err) - } + ck.NoError(c.List(ctx, pvcList, client.InNamespace(ns)), "list PVCs") nonTopoCount := 0 for _, pvc := range pvcList.Items { if !strings.Contains(pvc.Name, "-topo-") { nonTopoCount++ } } - if nonTopoCount == 0 { - t.Fatalf("expected at least one non-topo (data) PVC to survive cluster deletion under Retain, found none") - } + ck.NotEq( + 0, + nonTopoCount, + "expected at least one non-topo (data) PVC to survive cluster deletion under Retain, found none", + ) } diff --git a/test/e2e/dedicated/inline/inline_test.go b/test/e2e/dedicated/inline/inline_test.go index 3a8d415d..d59942fd 100644 --- a/test/e2e/dedicated/inline/inline_test.go +++ b/test/e2e/dedicated/inline/inline_test.go @@ -8,23 +8,22 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestInlineCluster applies config/samples/no-templates.yaml and verifies the // full resource tree is provisioned, all pods become ready, and psql SELECT 1 // succeeds through the multigateway. func TestInlineCluster(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") // Load and apply the inline (no-templates) sample. cr := framework.MustLoadCluster("config/samples/no-templates.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") // Verify child CRDs. framework.WaitForCRDCount(t, c, ns, diff --git a/test/e2e/dedicated/minimal/minimal_test.go b/test/e2e/dedicated/minimal/minimal_test.go index a9e2803b..b7498a57 100644 --- a/test/e2e/dedicated/minimal/minimal_test.go +++ b/test/e2e/dedicated/minimal/minimal_test.go @@ -8,23 +8,22 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestMinimalCluster applies the equivalent of config/samples/minimal.yaml and // verifies the full resource tree is provisioned, all pods become ready, and // psql SELECT 1 succeeds through the multigateway. func TestMinimalCluster(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") // Load and apply the minimal sample. cr := framework.MustLoadCluster("config/samples/minimal.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") // Verify child CRDs. framework.WaitForCRDCount(t, c, ns, diff --git a/test/e2e/dedicated/templated/templated_test.go b/test/e2e/dedicated/templated/templated_test.go index 0c9cb6c4..40f41053 100644 --- a/test/e2e/dedicated/templated/templated_test.go +++ b/test/e2e/dedicated/templated/templated_test.go @@ -8,38 +8,31 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestTemplatedCluster applies the template CRs from config/samples/templates/ // and the templated cluster CR from config/samples/templated-cluster.yaml, then // verifies the full resource tree, pod health, and psql connectivity. func TestTemplatedCluster(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Create templates first — the cluster CR references them. coreTmpl := framework.MustLoadCoreTemplate("config/samples/templates/core.yaml", ns) - if err := c.Create(ctx, coreTmpl); err != nil { - t.Fatalf("create CoreTemplate: %v", err) - } + ck.NoError(c.Create(ctx, coreTmpl), "create CoreTemplate") cellTmpl := framework.MustLoadCellTemplate("config/samples/templates/cell.yaml", ns) - if err := c.Create(ctx, cellTmpl); err != nil { - t.Fatalf("create CellTemplate: %v", err) - } + ck.NoError(c.Create(ctx, cellTmpl), "create CellTemplate") shardTmpl := framework.MustLoadShardTemplate("config/samples/templates/shard.yaml", ns) - if err := c.Create(ctx, shardTmpl); err != nil { - t.Fatalf("create ShardTemplate: %v", err) - } + ck.NoError(c.Create(ctx, shardTmpl), "create ShardTemplate") // Create the cluster referencing the templates. cr := framework.MustLoadCluster("config/samples/templated-cluster.yaml", ns) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") // Verify child CRDs. framework.WaitForCRDCount(t, c, ns, diff --git a/test/e2e/framework/diagnostics_test.go b/test/e2e/framework/diagnostics_test.go index d2d1a255..af36c67e 100644 --- a/test/e2e/framework/diagnostics_test.go +++ b/test/e2e/framework/diagnostics_test.go @@ -5,6 +5,8 @@ package framework import ( "bytes" "testing" + + "github.com/multigres/testkit/assert" ) func TestSafeDiagnosticName(t *testing.T) { @@ -22,15 +24,16 @@ func TestSafeDiagnosticName(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - if got := safeDiagnosticName(tt.input); got != tt.want { - t.Fatalf("safeDiagnosticName(%q) = %q, want %q", tt.input, got, tt.want) - } + got := safeDiagnosticName(tt.input) + assert.NewAborting(t). + Eq(tt.want, got, "safeDiagnosticName(%q) = %q, want", tt.input, got) }) } } func TestResourceStatusOutputExcludesSpecAndAnnotations(t *testing.T) { t.Parallel() + c := assert.NewCollecting(t) input := []byte(`{ "items": [{ "apiVersion": "multigres.com/v1alpha1", @@ -46,20 +49,20 @@ func TestResourceStatusOutputExcludesSpecAndAnnotations(t *testing.T) { }] }`) output, err := resourceStatusOutput(input) - if err != nil { - t.Fatalf("resourceStatusOutput: %v", err) - } + c.Require().NoError(err, "resourceStatusOutput") for _, excluded := range [][]byte{ []byte(`"spec"`), []byte(`"annotations"`), []byte("do-not-copy"), []byte("private"), } { - if bytes.Contains(output, excluded) { - t.Errorf("output contains excluded context %q: %s", excluded, output) - } - } - if !bytes.Contains(output, []byte(`"pooler-0": "PRIMARY"`)) { - t.Fatalf("output does not contain pod role status: %s", output) + c.False( + bytes.Contains(output, excluded), + "output contains excluded context %q: %s", + excluded, + output, + ) } + c.Require(). + True(bytes.Contains(output, []byte(`"pooler-0": "PRIMARY"`)), "output does not contain pod role status: %s", output) } diff --git a/test/e2e/framework/fixtures_test.go b/test/e2e/framework/fixtures_test.go index a4b6248b..d70a4aeb 100644 --- a/test/e2e/framework/fixtures_test.go +++ b/test/e2e/framework/fixtures_test.go @@ -13,7 +13,6 @@ import ( "github.com/multigres/multigres/go/common/consensus" clustermetadata "github.com/multigres/multigres/go/pb/clustermetadata" - "github.com/stretchr/testify/require" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/runtime" "k8s.io/apimachinery/pkg/runtime/serializer" @@ -23,24 +22,28 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/resolver" + + "github.com/multigres/testkit/assert" ) // Check every shared fixture and every sample consumed by the shared/dedicated // suites without starting Kubernetes. This tests resolved shard-wide capacity, // including template defaults and pool/cell placement, not replicas per pool. func TestE2EFixturesSupportBootstrap(t *testing.T) { + ck := assert.NewAborting(t) root, err := repoRoot() - require.NoError(t, err) + ck.NoError(err) scheme := runtime.NewScheme() - require.NoError(t, multigresv1alpha1.AddToScheme(scheme)) - require.NoError(t, corev1.AddToScheme(scheme)) + ck.NoError(multigresv1alpha1.AddToScheme(scheme)) + ck.NoError(corev1.AddToScheme(scheme)) strict := serializer.NewCodecFactory(scheme, serializer.EnableStrict).UniversalDeserializer() decode := func(t *testing.T, path string) []runtime.Object { t.Helper() + c := assert.NewAborting(t) data, err := os.ReadFile( path, ) // #nosec G304 -- Only repository fixture paths enumerated below are read. - require.NoError(t, err) + c.NoError(err) documents := utilyaml.NewYAMLOrJSONDecoder(bytes.NewReader(data), 4096) var objects []runtime.Object for { @@ -48,13 +51,13 @@ func TestE2EFixturesSupportBootstrap(t *testing.T) { if err := documents.Decode(&raw); err == io.EOF { break } else { - require.NoError(t, err) + c.NoError(err) } if len(raw.Raw) == 0 { continue } object, _, err := strict.Decode(raw.Raw, nil, nil) - require.NoError(t, err, "%s must use current API fields", path) + c.NoError(err, "%s must use current API fields", path) objects = append(objects, object) } return objects @@ -62,30 +65,31 @@ func TestE2EFixturesSupportBootstrap(t *testing.T) { var templates []client.Object for _, directory := range []string{"test/e2e/fixtures/templates", "config/samples/templates"} { paths, err := filepath.Glob(filepath.Join(root, directory, "*.yaml")) - require.NoError(t, err) + ck.NoError(err) for _, path := range paths { objects := decode(t, path) - require.Len(t, objects, 1) + ck.Len(objects, 1) template := objects[0].(client.Object) template.SetNamespace("test") templates = append(templates, template) } } paths, err := filepath.Glob(filepath.Join(root, "test/e2e/fixtures/*.yaml")) - require.NoError(t, err) + ck.NoError(err) for _, sample := range []string{"minimal.yaml", "no-templates.yaml", "templated-cluster.yaml"} { paths = append(paths, filepath.Join(root, "config/samples", sample)) } for _, path := range paths { relative, err := filepath.Rel(root, path) - require.NoError(t, err) + ck.NoError(err) t.Run(relative, func(t *testing.T) { + ck := assert.NewAborting(t) decode(t, path) cluster := MustLoadCluster(relative, "test") c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(templates...).Build() r := resolver.NewResolver(c, "test") _, err = r.PopulateClusterDefaults(context.Background(), cluster) - require.NoError(t, err) + ck.NoError(err) var cells []multigresv1alpha1.CellName for _, cell := range cluster.Spec.Cells { cells = append(cells, cell.Name) @@ -96,9 +100,9 @@ func TestE2EFixturesSupportBootstrap(t *testing.T) { policyName = cluster.Spec.DurabilityPolicy } policyProto, err := consensus.ParseUserSpecifiedDurabilityPolicy(policyName) - require.NoError(t, err) + ck.NoError(err) policy, err := consensus.NewPolicyFromProto(policyProto) - require.NoError(t, err) + ck.NoError(err) for _, group := range database.TableGroups { for _, shard := range group.Shards { if shard.ShardTemplate == "" { @@ -111,7 +115,7 @@ func TestE2EFixturesSupportBootstrap(t *testing.T) { AllCellNames: cells, MaterializeCellDefaults: true, }, ) - require.NoError(t, err) + ck.NoError(err) var cohort []*clustermetadata.ID for name, pool := range resolved.Pools { for _, cell := range pool.Cells { @@ -126,8 +130,7 @@ func TestE2EFixturesSupportBootstrap(t *testing.T) { } } } - require.True( - t, + ck.True( consensus.CohortSurvivesAnyMemberLoss(policy, cohort), "%s/%s/%s: %d poolers cannot safely bootstrap %s", database.Name, diff --git a/test/e2e/framework/helpers_test.go b/test/e2e/framework/helpers_test.go index 7fa3a423..05eff42a 100644 --- a/test/e2e/framework/helpers_test.go +++ b/test/e2e/framework/helpers_test.go @@ -5,43 +5,38 @@ package framework import ( "testing" - "github.com/stretchr/testify/require" - multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestMinimalFixtureUsesFailureSafeBootstrapCohort(t *testing.T) { + c := assert.NewAborting(t) cluster := MustLoadCluster("config/samples/minimal.yaml", "test") pool := cluster.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] - if pool.ReplicasPerCell == nil { - t.Fatal("synthetic default pool replicasPerCell is nil") - } - if got, want := *pool.ReplicasPerCell, int32(3); got != want { - t.Fatalf("synthetic default pool replicasPerCell = %d, want %d", got, want) - } + c.NotNil(pool.ReplicasPerCell, "synthetic default pool replicasPerCell is nil") + got, want := *pool.ReplicasPerCell, int32(3) + c.Eq(want, got, "synthetic default pool replicasPerCell") } func TestTemplatedFixturePreservesReferences(t *testing.T) { + c := assert.NewAborting(t) cluster := MustLoadCluster("test/e2e/fixtures/templated.yaml", "test") WithCIResources(&cluster.Spec) // Callers may apply resources more than once. - require.Equal( - t, - multigresv1alpha1.TemplateRef("e2e-core"), - cluster.Spec.TemplateDefaults.CoreTemplate, - ) - require.Nil(t, cluster.Spec.GlobalTopoServer) - require.Nil(t, cluster.Spec.Multiadmin) - require.Nil(t, cluster.Spec.MultiadminWeb) - require.Equal(t, multigresv1alpha1.TemplateRef("e2e-cell"), cluster.Spec.Cells[0].CellTemplate) - require.Nil(t, cluster.Spec.Cells[0].Spec) + c.EqDeep(multigresv1alpha1.TemplateRef("e2e-core"), cluster.Spec.TemplateDefaults.CoreTemplate) + c.Nil(cluster.Spec.GlobalTopoServer) + c.Nil(cluster.Spec.Multiadmin) + c.Nil(cluster.Spec.MultiadminWeb) + c.EqDeep(multigresv1alpha1.TemplateRef("e2e-cell"), cluster.Spec.Cells[0].CellTemplate) + c.Nil(cluster.Spec.Cells[0].Spec) shard := cluster.Spec.Databases[0].TableGroups[0].Shards[0] - require.Equal(t, multigresv1alpha1.TemplateRef("e2e-shard"), shard.ShardTemplate) - require.Nil(t, shard.Spec) + c.EqDeep(multigresv1alpha1.TemplateRef("e2e-shard"), shard.ShardTemplate) + c.Nil(shard.Spec) template := MustLoadShardTemplate("test/e2e/fixtures/templates/shard.yaml", "test") pool := template.Spec.Pools["default"] - require.NotNil(t, pool.ReplicasPerCell) - require.Equal(t, int32(3), *pool.ReplicasPerCell) + c.NotNil(pool.ReplicasPerCell) + c.EqDeep(int32(3), *pool.ReplicasPerCell) } func TestWithCIResourcesPreservesTemplateConfiguration(t *testing.T) { @@ -51,6 +46,7 @@ func TestWithCIResourcesPreservesTemplateConfiguration(t *testing.T) { name = "template defaults" } t.Run(name, func(t *testing.T) { + c := assert.NewAborting(t) spec := multigresv1alpha1.MultigresClusterSpec{ GlobalTopoServer: &multigresv1alpha1.GlobalTopoServerSpec{TemplateRef: "core"}, Multiadmin: &multigresv1alpha1.MultiadminConfig{TemplateRef: "core"}, @@ -82,18 +78,21 @@ func TestWithCIResourcesPreservesTemplateConfiguration(t *testing.T) { before := spec.DeepCopy() WithCIResources(&spec) WithCIResources(&spec) - require.Equal(t, before, &spec) + c.EqDeep(before, &spec) if defaults { spec.Databases[0].TableGroups[0].Shards = nil WithCIResources(&spec) - require.Nil(t, spec.Databases[0].TableGroups[0].Shards[0].Spec, - "a synthesized shard must still inherit its default template") + c.Nil( + spec.Databases[0].TableGroups[0].Shards[0].Spec, + "a synthesized shard must still inherit its default template", + ) } }) } } func TestWithCIResourcesPreservesOverridesAndExternalTopo(t *testing.T) { + c := assert.NewAborting(t) spec := multigresv1alpha1.MultigresClusterSpec{ GlobalTopoServer: &multigresv1alpha1.GlobalTopoServerSpec{ External: &multigresv1alpha1.ExternalTopoServerSpec{ @@ -115,19 +114,19 @@ func TestWithCIResourcesPreservesOverridesAndExternalTopo(t *testing.T) { }, } WithCIResources(&spec) - require.Nil(t, spec.GlobalTopoServer.Etcd) - require.Equal( - t, + c.Nil(spec.GlobalTopoServer.Etcd) + c.EqDeep( []multigresv1alpha1.EndpointUrl{"http://etcd:2379"}, spec.GlobalTopoServer.External.Endpoints, ) - require.Nil(t, spec.Cells[0].Spec) - require.NotNil(t, spec.Cells[0].Overrides) - require.Nil(t, spec.Databases[0].TableGroups[0].Shards[0].Spec) - require.NotNil(t, spec.Databases[0].TableGroups[0].Shards[0].Overrides) + c.Nil(spec.Cells[0].Spec) + c.NotNil(spec.Cells[0].Overrides) + c.Nil(spec.Databases[0].TableGroups[0].Shards[0].Spec) + c.NotNil(spec.Databases[0].TableGroups[0].Shards[0].Overrides) } func TestWithCIResourcesPreservesInlineConfiguration(t *testing.T) { + c := assert.NewAborting(t) cluster := MustLoadCluster("config/samples/no-templates.yaml", "test") spec := &cluster.Spec before := spec.DeepCopy() @@ -138,9 +137,9 @@ func TestWithCIResourcesPreservesInlineConfiguration(t *testing.T) { ShardTemplate: "shard", } WithCIResources(spec) - require.Equal(t, before.GlobalTopoServer, spec.GlobalTopoServer) - require.Equal(t, before.Multiadmin, spec.Multiadmin) - require.Equal(t, before.MultiadminWeb, spec.MultiadminWeb) - require.Equal(t, before.Cells, spec.Cells) - require.Equal(t, before.Databases, spec.Databases) + c.EqDeep(before.GlobalTopoServer, spec.GlobalTopoServer) + c.EqDeep(before.Multiadmin, spec.Multiadmin) + c.EqDeep(before.MultiadminWeb, spec.MultiadminWeb) + c.EqDeep(before.Cells, spec.Cells) + c.EqDeep(before.Databases, spec.Databases) } diff --git a/test/e2e/framework/image_overrides_test.go b/test/e2e/framework/image_overrides_test.go index 309535f3..1d7face7 100644 --- a/test/e2e/framework/image_overrides_test.go +++ b/test/e2e/framework/image_overrides_test.go @@ -8,27 +8,29 @@ import ( "testing" multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" + + "github.com/multigres/testkit/assert" ) func TestApplyImageOverrides(t *testing.T) { t.Setenv(postgresImageEnv, "example.test/postgres:custom") t.Setenv(multigatewayImageEnv, "example.test/multigres:gateway") + c := assert.NewAborting(t) cluster := &multigresv1alpha1.MultigresCluster{} applyImageOverrides(cluster) - if got := string(cluster.Spec.Images.Postgres); got != "example.test/postgres:custom" { - t.Fatalf("postgres image = %q", got) - } - if got := string(cluster.Spec.Images.Multigateway); got != "example.test/multigres:gateway" { - t.Fatalf("multigateway image = %q", got) - } - if cluster.Spec.Images.Multiadmin != "" { - t.Fatalf("unset multiadmin override unexpectedly changed to %q", cluster.Spec.Images.Multiadmin) - } + c.Eq("example.test/postgres:custom", string(cluster.Spec.Images.Postgres), "postgres image =") + c.Eq( + "example.test/multigres:gateway", + string(cluster.Spec.Images.Multigateway), + "multigateway image =", + ) + c.Eq("", cluster.Spec.Images.Multiadmin, "unset multiadmin override unexpectedly changed to") } func TestRuntimeImagesUsesOverridesAndDeduplicates(t *testing.T) { + c := assert.NewAborting(t) const multigresNightly = "ghcr.io/multigres/multigres:nightly-sha-abcdef0" t.Setenv(multiadminImageEnv, multigresNightly) t.Setenv(multiorchImageEnv, multigresNightly) @@ -36,12 +38,13 @@ func TestRuntimeImagesUsesOverridesAndDeduplicates(t *testing.T) { t.Setenv(multigatewayImageEnv, multigresNightly) images := runtimeImages() - if got := count(images, multigresNightly); got != 1 { - t.Fatalf("nightly multigres image occurs %d times in %v", got, images) - } - if slices.Contains(images, multigresv1alpha1.DefaultMultiadminImage) { - t.Fatalf("default multigres image retained despite complete override: %v", images) - } + got := count(images, multigresNightly) + c.Eq(1, got, "nightly multigres image occurs %d times in %v", got, images) + c.NotContains( + images, + multigresv1alpha1.DefaultMultiadminImage, + "default multigres image retained despite complete override", + ) } func count(values []string, target string) int { @@ -88,9 +91,7 @@ func TestRuntimeImagesUsesCommittedDefaults(t *testing.T) { want = slices.Compact(want) got := runtimeImages() slices.Sort(got) - if !slices.Equal(got, want) { - t.Fatalf("runtimeImages() = %v, want %v", got, want) - } + assert.NewAborting(t).EqDiff(want, got, "runtimeImages()") }) } } diff --git a/test/e2e/framework/images_test.go b/test/e2e/framework/images_test.go index be2dc04a..ac424d22 100644 --- a/test/e2e/framework/images_test.go +++ b/test/e2e/framework/images_test.go @@ -9,8 +9,7 @@ import ( "strings" "testing" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" + "github.com/multigres/testkit/assert" ) func TestLoadDigestImages(t *testing.T) { @@ -52,14 +51,15 @@ func TestLoadDigestImages(t *testing.T) { {name: "no images", empty: true}, } { t.Run(tt.name, func(t *testing.T) { + c := assert.NewCollecting(t) dir := t.TempDir() logPath := filepath.Join(dir, "commands") - require.NoError(t, os.WriteFile(logPath, nil, 0o600)) + c.Require().NoError(os.WriteFile(logPath, nil, 0o600)) // #nosec G306 -- Owner-only executable test fixture in t.TempDir. - require.NoError(t, os.WriteFile(filepath.Join(dir, "kind"), []byte( + c.Require().NoError(os.WriteFile(filepath.Join(dir, "kind"), []byte( "#!/bin/sh\nprintf 'node-a\\nnode-b\\n'\n"), 0o700)) // #nosec G306 -- Owner-only executable test fixture in t.TempDir. - require.NoError(t, os.WriteFile(filepath.Join(dir, "docker"), []byte(`#!/bin/sh + c.Require().NoError(os.WriteFile(filepath.Join(dir, "docker"), []byte(`#!/bin/sh printf '%s\n' "$*" >> "$COMMAND_LOG" case "$4" in inspecti) [ "$CACHE_HIT" = true ] ;; @@ -82,16 +82,16 @@ esac } err := LoadImages(context.Background(), "test-cluster", images) if tt.wantError { - require.ErrorContains(t, err, image) - assert.ErrorContains(t, err, "node-a") - assert.ErrorContains(t, err, "registry unavailable") + c.Require().ErrorContains(err, image) + c.ErrorContains(err, "node-a") + c.ErrorContains(err, "registry unavailable") } else { - require.NoError(t, err) + c.Require().NoError(err) } // #nosec G304 -- Read only the command log created in t.TempDir above. calls, err := os.ReadFile(logPath) - require.NoError(t, err) - assert.Equal(t, strings.Join(tt.wantCalls, "\n"), strings.TrimSpace(string(calls))) + c.Require().NoError(err) + c.EqDeep(strings.Join(tt.wantCalls, "\n"), strings.TrimSpace(string(calls))) }) } } diff --git a/test/e2e/shared/deletion/deletion_test.go b/test/e2e/shared/deletion/deletion_test.go index f43a9107..1e7b40ba 100644 --- a/test/e2e/shared/deletion/deletion_test.go +++ b/test/e2e/shared/deletion/deletion_test.go @@ -18,25 +18,24 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestClusterDeletion verifies that deleting a MultigresCluster triggers // cascading deletion of all child resources (CRDs and Kubernetes resources). func TestClusterDeletion(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Load and create the minimal sample. cr := framework.MustLoadCluster("config/samples/minimal.yaml", ns) cr.Name = "delete-me" // distinct name for clarity in logs - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for full provisioning. cluster.WaitForAllPodsReady(t, ns) @@ -44,20 +43,21 @@ func TestClusterDeletion(t *testing.T) { // Delete the cluster. clusterKey := client.ObjectKeyFromObject(cr) - if err := c.Delete(ctx, cr); err != nil { - t.Fatalf("delete MultigresCluster: %v", err) - } + ck.NoError(c.Delete(ctx, cr), "delete MultigresCluster") // Wait for the MultigresCluster object to disappear. pollCtx, cancel := context.WithTimeout(ctx, 5*time.Minute) defer cancel() - err = wait.PollUntilContextCancel(pollCtx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - err := c.Get(ctx, clusterKey, &multigresv1alpha1.MultigresCluster{}) - return apierrors.IsNotFound(err), nil - }) - if err != nil { - t.Fatalf("MultigresCluster not deleted: %v", err) - } + err = wait.PollUntilContextCancel( + pollCtx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + err := c.Get(ctx, clusterKey, &multigresv1alpha1.MultigresCluster{}) + return apierrors.IsNotFound(err), nil + }, + ) + ck.NoError(err, "MultigresCluster not deleted") // Verify child CRDs are cleaned up. framework.WaitForEmpty(t, c, ns, @@ -95,11 +95,10 @@ func TestClusterDeletion(t *testing.T) { // StatefulSet no longer exist by the time cluster deletion starts. func TestClusterDeletionAfterSwitchingToExternalTopo(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.Require().NoError(err, "create CR client") ctx := context.Background() cr := framework.MustLoadCluster("config/samples/minimal.yaml", ns) @@ -108,22 +107,19 @@ func TestClusterDeletionAfterSwitchingToExternalTopo(t *testing.T) { WhenDeleted: multigresv1alpha1.RetainPVCRetentionPolicy, WhenScaled: multigresv1alpha1.RetainPVCRetentionPolicy, } - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.Require().NoError(c.Create(ctx, cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) topoStatefulSet := framework.WaitForStatefulSet(t, c, ns, "etcd") topoPVCNames := waitForBoundTopoPVCs(t, c, ns, cr.Name) nonTopoPVCNames := listNonTopoPVCNames(t, c, ns, topoPVCNames) - if len(nonTopoPVCNames) == 0 { - t.Fatal("expected at least one non-topo PVC to verify the Retain policy") - } + ck.Require(). + NotEmpty(nonTopoPVCNames, "expected at least one non-topo PVC to verify the Retain policy") // Patch directly instead of using framework.PatchCluster (the external // endpoint is intentionally unreachable). - if err := c.Patch(ctx, cr, client.RawPatch(types.MergePatchType, []byte(`{ + ck.Require().NoError(c.Patch(ctx, cr, client.RawPatch(types.MergePatchType, []byte(`{ "spec": { "globalTopoServer": { "etcd": null, @@ -133,9 +129,7 @@ func TestClusterDeletionAfterSwitchingToExternalTopo(t *testing.T) { } } } - }`))); err != nil { - t.Fatalf("switch cluster to external topology: %v", err) - } + }`))), "switch cluster to external topology") waitForObjectsDeleted(t, c, 3*time.Minute, objectToDelete{ @@ -153,22 +147,17 @@ func TestClusterDeletionAfterSwitchingToExternalTopo(t *testing.T) { // Retain must leave the old managed-topology PVCs behind during the switch. for _, name := range topoPVCNames { pvc := &corev1.PersistentVolumeClaim{} - if err := c.Get(ctx, client.ObjectKey{Namespace: ns, Name: name}, pvc); err != nil { - t.Fatalf("expected retained topo PVC %q after switching to external topo: %v", name, err) - } - if !pvc.DeletionTimestamp.IsZero() { - t.Fatalf("topo PVC %q is unexpectedly terminating after switching to external topo", name) - } + ck.Require(). + NoError(c.Get(ctx, client.ObjectKey{Namespace: ns, Name: name}, pvc), "expected retained topo PVC %q after switching to external topo", name) + ck.Require(). + True(pvc.DeletionTimestamp.IsZero(), "topo PVC %q is unexpectedly terminating after switching to external topo", name) } liveCluster := &multigresv1alpha1.MultigresCluster{} clusterKey := client.ObjectKeyFromObject(cr) - if err := c.Get(ctx, clusterKey, liveCluster); err != nil { - t.Fatalf("get MultigresCluster before deletion: %v", err) - } - if err := c.Delete(ctx, liveCluster); err != nil { - t.Fatalf("delete MultigresCluster: %v", err) - } + ck.Require(). + NoError(c.Get(ctx, clusterKey, liveCluster), "get MultigresCluster before deletion") + ck.Require().NoError(c.Delete(ctx, liveCluster), "delete MultigresCluster") objects := []objectToDelete{{ key: clusterKey, @@ -200,9 +189,11 @@ func TestClusterDeletionAfterSwitchingToExternalTopo(t *testing.T) { t.Errorf("expected retained non-topo PVC %q after cluster deletion: %v", name, err) continue } - if !pvc.DeletionTimestamp.IsZero() { - t.Errorf("retained non-topo PVC %q is unexpectedly terminating", name) - } + ck.True( + pvc.DeletionTimestamp.IsZero(), + "retained non-topo PVC %q is unexpectedly terminating", + name, + ) } } @@ -216,35 +207,38 @@ func waitForBoundTopoPVCs( var names []string pollCtx, cancel := context.WithTimeout(context.Background(), 3*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(pollCtx, 2*time.Second, true, func(ctx context.Context) (bool, error) { - list := &corev1.PersistentVolumeClaimList{} - if err := c.List( - ctx, - list, - client.InNamespace(ns), - client.MatchingLabels{ - metadata.LabelMultigresCluster: clusterName, - metadata.LabelAppComponent: metadata.ComponentTopoServer, - }, - ); err != nil { - return false, nil - } - if len(list.Items) == 0 { - return false, nil - } - currentNames := make([]string, 0, len(list.Items)) - for _, pvc := range list.Items { - if pvc.Status.Phase != corev1.ClaimBound { + err := wait.PollUntilContextCancel( + pollCtx, + 2*time.Second, + true, + func(ctx context.Context) (bool, error) { + list := &corev1.PersistentVolumeClaimList{} + if err := c.List( + ctx, + list, + client.InNamespace(ns), + client.MatchingLabels{ + metadata.LabelMultigresCluster: clusterName, + metadata.LabelAppComponent: metadata.ComponentTopoServer, + }, + ); err != nil { return false, nil } - currentNames = append(currentNames, pvc.Name) - } - names = currentNames - return true, nil - }) - if err != nil { - t.Fatalf("timed out waiting for bound managed-topology PVCs: %v", err) - } + if len(list.Items) == 0 { + return false, nil + } + currentNames := make([]string, 0, len(list.Items)) + for _, pvc := range list.Items { + if pvc.Status.Phase != corev1.ClaimBound { + return false, nil + } + currentNames = append(currentNames, pvc.Name) + } + names = currentNames + return true, nil + }, + ) + assert.NewAborting(t).NoError(err, "timed out waiting for bound managed-topology PVCs") return names } @@ -261,9 +255,8 @@ func listNonTopoPVCNames( } list := &corev1.PersistentVolumeClaimList{} - if err := c.List(context.Background(), list, client.InNamespace(ns)); err != nil { - t.Fatalf("list PVCs: %v", err) - } + assert.NewAborting(t). + NoError(c.List(context.Background(), list, client.InNamespace(ns)), "list PVCs") var names []string for _, pvc := range list.Items { if _, isTopo := topoPVCs[pvc.Name]; !isTopo { @@ -288,14 +281,25 @@ func waitForObjectsDeleted( t.Helper() pollCtx, cancel := context.WithTimeout(context.Background(), timeout) defer cancel() - err := wait.PollUntilContextCancel(pollCtx, 2*time.Second, true, func(ctx context.Context) (bool, error) { - for _, object := range objects { - if err := c.Get(ctx, object.key, object.object.DeepCopyObject().(client.Object)); !apierrors.IsNotFound(err) { - return false, nil + err := wait.PollUntilContextCancel( + pollCtx, + 2*time.Second, + true, + func(ctx context.Context) (bool, error) { + for _, object := range objects { + if err := c.Get( + ctx, + object.key, + object.object.DeepCopyObject().(client.Object), + ); !apierrors.IsNotFound( + err, + ) { + return false, nil + } } - } - return true, nil - }) + return true, nil + }, + ) if err != nil { pending := make([]string, 0, len(objects)) for _, object := range objects { diff --git a/test/e2e/shared/drain/drain_test.go b/test/e2e/shared/drain/drain_test.go index 5406a92c..8d71b56a 100644 --- a/test/e2e/shared/drain/drain_test.go +++ b/test/e2e/shared/drain/drain_test.go @@ -26,6 +26,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) var cluster *framework.Cluster @@ -42,11 +44,10 @@ func TestMain(m *testing.M) { // TestExternalPoolerDeletion verifies recovery after replica and primary deletion. func TestExternalPoolerDeletion(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) @@ -55,9 +56,7 @@ func TestExternalPoolerDeletion(t *testing.T) { pool := cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] pool.ReplicasPerCell = &replicas cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] = pool - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") t.Log("waiting for the initial three-pooler shard to become healthy") cluster.WaitForAllPodsReady(t, ns) @@ -68,27 +67,22 @@ func TestExternalPoolerDeletion(t *testing.T) { deleteAndWaitForReplacement(t, c, replica) shard = waitForHealthyShard(t, c, ns, cr.Name) primaryAfterReplica, _ := waitForPrimaryAndReplica(t, c, ns, cr.Name, shard) - if primaryAfterReplica.Name != primary.Name { - t.Fatalf("replica deletion changed primary from %q to %q", primary.Name, primaryAfterReplica.Name) - } + ck.Eq(primary.Name, primaryAfterReplica.Name, "replica deletion changed primary from") t.Logf("deleting primary %q", primaryAfterReplica.Name) deleteAndWaitForReplacement(t, c, primaryAfterReplica) shard = waitForHealthyShard(t, c, ns, cr.Name) primaryAfterFailover, _ := waitForPrimaryAndReplica(t, c, ns, cr.Name, shard) - if primaryAfterFailover.Name == primaryAfterReplica.Name { - t.Fatalf("primary %q was not replaced after deletion", primaryAfterReplica.Name) - } + ck.NotEq(primaryAfterReplica.Name, primaryAfterFailover.Name, "primary") t.Logf("failover completed with primary %q", primaryAfterFailover.Name) } // TestGracefulScaleDown verifies that a pool can scale down while remaining healthy. func TestGracefulScaleDown(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) @@ -97,9 +91,7 @@ func TestGracefulScaleDown(t *testing.T) { pool := cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] pool.ReplicasPerCell = &replicas cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] = pool - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") t.Log("waiting for the initial three-pooler shard to become healthy") cluster.WaitForAllPodsReady(t, ns) @@ -108,23 +100,17 @@ func TestGracefulScaleDown(t *testing.T) { poolPodsBeforeScaleDown := listPoolPods(t, c, ns, cr.Name) t.Log("scaling the pool from three poolers to two") - if err := c.Get(ctx, client.ObjectKeyFromObject(cr), cr); err != nil { - t.Fatalf("get MultigresCluster: %v", err) - } + ck.NoError(c.Get(ctx, client.ObjectKeyFromObject(cr), cr), "get MultigresCluster") replicas = 2 pool = cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] pool.ReplicasPerCell = &replicas cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] = pool - if err := c.Update(ctx, cr); err != nil { - t.Fatalf("scale down MultigresCluster: %v", err) - } + ck.NoError(c.Update(ctx, cr), "scale down MultigresCluster") waitForPoolPodCount(t, c, ns, cr.Name, 2) shard = waitForHealthyShard(t, c, ns, cr.Name) primaryAfterScaleDown, _ := waitForPrimaryAndReplica(t, c, ns, cr.Name, shard) - if primaryAfterScaleDown.Name != primary.Name { - t.Fatalf("scale-down changed primary from %q to %q", primary.Name, primaryAfterScaleDown.Name) - } + ck.Eq(primary.Name, primaryAfterScaleDown.Name, "scale-down changed primary from") removed := removedPoolPod(t, poolPodsBeforeScaleDown, listPoolPods(t, c, ns, cr.Name)) probe := newMultiadminProbe(t, c, ns, cr.Name, primaryAfterScaleDown) t.Cleanup(func() { _ = c.Delete(context.Background(), probe) }) @@ -141,24 +127,27 @@ func waitForHealthyShard( var found *multigresv1alpha1.Shard ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(ctx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - shards := &multigresv1alpha1.ShardList{} - if err := c.List(ctx, shards, - client.InNamespace(namespace), - client.MatchingLabels{metadata.LabelMultigresCluster: clusterName}, - ); err != nil || len(shards.Items) != 1 { - return false, nil - } - shard := &shards.Items[0] - if shard.Status.Phase != multigresv1alpha1.PhaseHealthy || !shard.Status.OrchReady { - return false, nil - } - found = shard - return true, nil - }) - if err != nil { - t.Fatalf("timed out waiting for healthy shard: %v", err) - } + err := wait.PollUntilContextCancel( + ctx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + shards := &multigresv1alpha1.ShardList{} + if err := c.List(ctx, shards, + client.InNamespace(namespace), + client.MatchingLabels{metadata.LabelMultigresCluster: clusterName}, + ); err != nil || len(shards.Items) != 1 { + return false, nil + } + shard := &shards.Items[0] + if shard.Status.Phase != multigresv1alpha1.PhaseHealthy || !shard.Status.OrchReady { + return false, nil + } + found = shard + return true, nil + }, + ) + assert.NewAborting(t).NoError(err, "timed out waiting for healthy shard") return found } @@ -172,103 +161,112 @@ func waitForPrimaryAndReplica( var primary, replica *corev1.Pod ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(ctx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - freshShard := &multigresv1alpha1.Shard{} - if err := c.Get(ctx, client.ObjectKeyFromObject(shard), freshShard); err != nil { - return false, nil - } - pods := &corev1.PodList{} - if err := c.List(ctx, pods, - client.InNamespace(namespace), - client.MatchingLabels{ - metadata.LabelMultigresCluster: clusterName, - metadata.LabelMultigresPool: "default", - }, - ); err != nil { - return false, nil - } - primary, replica = nil, nil - for i := range pods.Items { - pod := &pods.Items[i] - if !podReady(pod) { + err := wait.PollUntilContextCancel( + ctx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + freshShard := &multigresv1alpha1.Shard{} + if err := c.Get(ctx, client.ObjectKeyFromObject(shard), freshShard); err != nil { return false, nil } - switch freshShard.Status.PodRoles[pod.Name] { - case "PRIMARY": - primary = pod.DeepCopy() - case "REPLICA": - if replica == nil { - replica = pod.DeepCopy() + pods := &corev1.PodList{} + if err := c.List(ctx, pods, + client.InNamespace(namespace), + client.MatchingLabels{ + metadata.LabelMultigresCluster: clusterName, + metadata.LabelMultigresPool: "default", + }, + ); err != nil { + return false, nil + } + primary, replica = nil, nil + for i := range pods.Items { + pod := &pods.Items[i] + if !podReady(pod) { + return false, nil + } + switch freshShard.Status.PodRoles[pod.Name] { + case "PRIMARY": + primary = pod.DeepCopy() + case "REPLICA": + if replica == nil { + replica = pod.DeepCopy() + } } } - } - return primary != nil && replica != nil, nil - }) - if err != nil { - t.Fatalf("timed out waiting for primary and replica: %v", err) - } + return primary != nil && replica != nil, nil + }, + ) + assert.NewAborting(t).NoError(err, "timed out waiting for primary and replica") return primary, replica } func deleteAndWaitForReplacement(t testing.TB, c client.Client, pod *corev1.Pod) { t.Helper() + ck := assert.NewAborting(t) ctx := context.Background() - if err := c.Delete(ctx, pod); err != nil { - t.Fatalf("delete pod %q: %v", pod.Name, err) - } + ck.NoError(c.Delete(ctx, pod), "delete pod %q", pod.Name) key := types.NamespacedName{Namespace: pod.Namespace, Name: pod.Name} waitCtx, cancel := context.WithTimeout(ctx, 5*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(waitCtx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - replacement := &corev1.Pod{} - if err := c.Get(ctx, key, replacement); err != nil { - if apierrors.IsNotFound(err) { - return false, nil + err := wait.PollUntilContextCancel( + waitCtx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + replacement := &corev1.Pod{} + if err := c.Get(ctx, key, replacement); err != nil { + if apierrors.IsNotFound(err) { + return false, nil + } + return false, err } - return false, err - } - return replacement.UID != pod.UID && replacement.DeletionTimestamp.IsZero() && podReady(replacement), nil - }) - if err != nil { - t.Fatalf("timed out waiting for replacement of pod %q: %v", pod.Name, err) - } + return replacement.UID != pod.UID && replacement.DeletionTimestamp.IsZero() && + podReady(replacement), nil + }, + ) + ck.NoError(err, "timed out waiting for replacement of pod %q", pod.Name) } func waitForPoolPodCount(t testing.TB, c client.Client, namespace, clusterName string, want int) { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(ctx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - pods := &corev1.PodList{} - if err := c.List(ctx, pods, - client.InNamespace(namespace), - client.MatchingLabels{ - metadata.LabelMultigresCluster: clusterName, - metadata.LabelMultigresPool: "default", - }, - ); err != nil { - return false, err - } - if len(pods.Items) != want { - return false, nil - } - for i := range pods.Items { - if !pods.Items[i].DeletionTimestamp.IsZero() || !podReady(&pods.Items[i]) { + err := wait.PollUntilContextCancel( + ctx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + pods := &corev1.PodList{} + if err := c.List(ctx, pods, + client.InNamespace(namespace), + client.MatchingLabels{ + metadata.LabelMultigresCluster: clusterName, + metadata.LabelMultigresPool: "default", + }, + ); err != nil { + return false, err + } + if len(pods.Items) != want { return false, nil } - } - return true, nil - }) - if err != nil { - t.Fatalf("timed out waiting for %d ready pool pods: %v", want, err) - } + for i := range pods.Items { + if !pods.Items[i].DeletionTimestamp.IsZero() || !podReady(&pods.Items[i]) { + return false, nil + } + } + return true, nil + }, + ) + assert.NewAborting(t).NoError(err, "timed out waiting for %d ready pool pods", want) } func listPoolPods(t testing.TB, c client.Client, namespace, clusterName string) []*corev1.Pod { t.Helper() pods := &corev1.PodList{} - if err := c.List( + assert.NewAborting(t).NoError(c.List( context.Background(), pods, client.InNamespace(namespace), @@ -276,9 +274,7 @@ func listPoolPods(t testing.TB, c client.Client, namespace, clusterName string) metadata.LabelMultigresCluster: clusterName, metadata.LabelMultigresPool: "default", }, - ); err != nil { - t.Fatalf("list pool pods: %v", err) - } + ), "list pool pods") result := make([]*corev1.Pod, 0, len(pods.Items)) for i := range pods.Items { result = append(result, pods.Items[i].DeepCopy()) @@ -299,9 +295,7 @@ func removedPoolPod(t testing.TB, before, after []*corev1.Pod) *corev1.Pod { removed = append(removed, pod) } } - if len(removed) != 1 { - t.Fatalf("removed pool pods = %d, want 1", len(removed)) - } + assert.NewAborting(t).Len(removed, 1, "removed pool pods = %d, want 1", len(removed)) return removed[0] } @@ -315,25 +309,32 @@ func waitForCohortRemoval( ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(ctx, 2*time.Second, true, func(ctx context.Context) (bool, error) { - status, err := cohortProbeStatus(ctx, namespace, probeName) - if err != nil { - return false, nil - } - members := status.GetConsensusStatus().GetCurrentPosition().GetPosition().GetDecision().GetCohortMembers() - if len(members) != 2 { - return false, nil - } - for _, member := range members { - if member.GetName() == removedID { + err := wait.PollUntilContextCancel( + ctx, + 2*time.Second, + true, + func(ctx context.Context) (bool, error) { + status, err := cohortProbeStatus(ctx, namespace, probeName) + if err != nil { return false, nil } - } - return true, nil - }) - if err != nil { - t.Fatalf("cohort still contains removed pooler %q: %v", removedID, err) - } + members := status.GetConsensusStatus(). + GetCurrentPosition(). + GetPosition(). + GetDecision(). + GetCohortMembers() + if len(members) != 2 { + return false, nil + } + for _, member := range members { + if member.GetName() == removedID { + return false, nil + } + } + return true, nil + }, + ) + assert.NewAborting(t).NoError(err, "cohort still contains removed pooler %q", removedID) } func waitForPoolerShutdown(t testing.TB, namespace, probeName string, removed *corev1.Pod) { @@ -342,21 +343,26 @@ func waitForPoolerShutdown(t testing.TB, namespace, probeName string, removed *c ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(ctx, 2*time.Second, true, func(ctx context.Context) (bool, error) { - poolers, err := poolersProbeStatus(ctx, namespace, probeName) - if err != nil { - return false, nil - } - for _, pooler := range poolers.GetPoolers() { - if pooler.GetId().GetName() == removedID { - return pooler.GetLifecycleStatus().GetStatus() == clustermetadatapb.PoolerLifecycleStatus_LIFECYCLE_SHUTDOWN, nil + err := wait.PollUntilContextCancel( + ctx, + 2*time.Second, + true, + func(ctx context.Context) (bool, error) { + poolers, err := poolersProbeStatus(ctx, namespace, probeName) + if err != nil { + return false, nil } - } - return false, nil - }) - if err != nil { - t.Fatalf("pooler %q did not reach LIFECYCLE_SHUTDOWN: %v", removedID, err) - } + for _, pooler := range poolers.GetPoolers() { + if pooler.GetId().GetName() == removedID { + return pooler.GetLifecycleStatus(). + GetStatus() == + clustermetadatapb.PoolerLifecycleStatus_LIFECYCLE_SHUTDOWN, nil + } + } + return false, nil + }, + ) + assert.NewAborting(t).NoError(err, "pooler %q did not reach LIFECYCLE_SHUTDOWN", removedID) } func newMultiadminProbe( @@ -403,9 +409,7 @@ done`, body, host, host, host) }}, }, } - if err := c.Create(context.Background(), probe); err != nil { - t.Fatalf("create cohort probe: %v", err) - } + assert.NewAborting(t).NoError(c.Create(context.Background(), probe), "create cohort probe") return probe } @@ -439,7 +443,10 @@ func poolerServiceID(t testing.TB, pod *corev1.Pod) string { return "" } -func cohortProbeStatus(ctx context.Context, namespace, podName string) (*multiadminpb.GetPoolerStatusResponse, error) { +func cohortProbeStatus( + ctx context.Context, + namespace, podName string, +) (*multiadminpb.GetPoolerStatusResponse, error) { output, err := multiadminProbeOutput(ctx, namespace, podName) if err != nil { return nil, err @@ -455,7 +462,10 @@ func cohortProbeStatus(ctx context.Context, namespace, podName string) (*multiad return status, nil } -func poolersProbeStatus(ctx context.Context, namespace, podName string) (*multiadminpb.GetPoolersResponse, error) { +func poolersProbeStatus( + ctx context.Context, + namespace, podName string, +) (*multiadminpb.GetPoolersResponse, error) { output, err := multiadminProbeOutput(ctx, namespace, podName) if err != nil { return nil, err @@ -472,7 +482,10 @@ func poolersProbeStatus(ctx context.Context, namespace, podName string) (*multia } func multiadminProbeOutput(ctx context.Context, namespace, podName string) (string, error) { - stream, err := cluster.Clientset.CoreV1().Pods(namespace).GetLogs(podName, &corev1.PodLogOptions{}).Stream(ctx) + stream, err := cluster.Clientset.CoreV1(). + Pods(namespace). + GetLogs(podName, &corev1.PodLogOptions{}). + Stream(ctx) if err != nil { return "", err } @@ -493,7 +506,9 @@ func parseProbeResponse(output, startMarker, endMarker string) ([]byte, error) { continue } httpResponse, err := http.ReadResponse( - bufio.NewReader(strings.NewReader(strings.TrimSpace(response[start+len(startMarker):]))), + bufio.NewReader( + strings.NewReader(strings.TrimSpace(response[start+len(startMarker):])), + ), &http.Request{Method: http.MethodPost}, ) if err != nil { diff --git a/test/e2e/shared/inline/inline_test.go b/test/e2e/shared/inline/inline_test.go index 171ca2b4..6569705e 100644 --- a/test/e2e/shared/inline/inline_test.go +++ b/test/e2e/shared/inline/inline_test.go @@ -8,6 +8,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestInlineCluster applies config/samples/no-templates.yaml and verifies the @@ -15,17 +17,14 @@ import ( // succeeds through the multigateway. func TestInlineCluster(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") // Load and apply the inline (no-templates) sample. cr := framework.MustLoadCluster("config/samples/no-templates.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") // Verify child CRDs. framework.WaitForCRDCount(t, c, ns, diff --git a/test/e2e/shared/minimal/minimal_test.go b/test/e2e/shared/minimal/minimal_test.go index d3073956..5c576a42 100644 --- a/test/e2e/shared/minimal/minimal_test.go +++ b/test/e2e/shared/minimal/minimal_test.go @@ -8,6 +8,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestMinimalCluster applies the equivalent of config/samples/minimal.yaml and @@ -15,17 +17,14 @@ import ( // psql SELECT 1 succeeds through the multigateway. func TestMinimalCluster(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") // Load and apply the minimal sample. cr := framework.MustLoadCluster("config/samples/minimal.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") // Verify child CRDs. framework.WaitForCRDCount(t, c, ns, diff --git a/test/e2e/shared/postgresconfig/baseline_over_ref_test.go b/test/e2e/shared/postgresconfig/baseline_over_ref_test.go index 5115dfe1..a2267857 100644 --- a/test/e2e/shared/postgresconfig/baseline_over_ref_test.go +++ b/test/e2e/shared/postgresconfig/baseline_over_ref_test.go @@ -12,6 +12,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestBaselineWinsOverRef verifies end-to-end that the operator's own @@ -33,11 +35,10 @@ import ( // still applies, proving the fix layers the ref UNDER the baseline rather than // ignoring it. func TestBaselineWinsOverRef(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Deprecated ref (key "postgresql.conf", like the real project): sets a @@ -52,9 +53,7 @@ func TestBaselineWinsOverRef(t *testing.T) { "postgresql.conf": "effective_cache_size = '999MB'\nseq_page_cost = '2.5'", }, } - if err := c.Create(ctx, refCM); err != nil { - t.Fatalf("create ref ConfigMap: %v", err) - } + ck.NoError(c.Create(ctx, refCM), "create ref ConfigMap") // Create with the ref set and NO inline postgresConfig, and pin the pool // memory to 512Mi so the operator's sized effective_cache_size is a @@ -73,9 +72,7 @@ func TestBaselineWinsOverRef(t *testing.T) { pool.Postgres.Resources.Limits[corev1.ResourceMemory] = resource.MustParse("512Mi") shard.Spec.Pools[name] = pool } - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") framework.WaitForPod(t, c, ns, "postgres") cluster.WaitForAllPodsReady(t, ns) diff --git a/test/e2e/shared/postgresconfig/helpers_test.go b/test/e2e/shared/postgresconfig/helpers_test.go index 3e92e6c9..64eaa04d 100644 --- a/test/e2e/shared/postgresconfig/helpers_test.go +++ b/test/e2e/shared/postgresconfig/helpers_test.go @@ -9,18 +9,23 @@ import ( corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/types" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) // poolPodUIDs returns the current pool pods (component shard-pool) in ns keyed by // name -> UID. Tests snapshot it before and after a config change to tell a reload // (stable UIDs) apart from a restart (recreated pods, new UIDs). -func poolPodUIDs(t *testing.T, ctx context.Context, c client.Client, ns string) map[string]types.UID { +func poolPodUIDs( + t *testing.T, + ctx context.Context, + c client.Client, + ns string, +) map[string]types.UID { t.Helper() pods := &corev1.PodList{} - if err := c.List(ctx, pods, client.InNamespace(ns), - client.MatchingLabels{"app.kubernetes.io/component": "shard-pool"}); err != nil { - t.Fatalf("list pool pods: %v", err) - } + assert.NewAborting(t).NoError(c.List(ctx, pods, client.InNamespace(ns), + client.MatchingLabels{"app.kubernetes.io/component": "shard-pool"}), "list pool pods") uids := map[string]types.UID{} for i := range pods.Items { uids[pods.Items[i].Name] = pods.Items[i].UID diff --git a/test/e2e/shared/postgresconfig/logconnections_test.go b/test/e2e/shared/postgresconfig/logconnections_test.go index 57c491f8..4b9be527 100644 --- a/test/e2e/shared/postgresconfig/logconnections_test.go +++ b/test/e2e/shared/postgresconfig/logconnections_test.go @@ -7,6 +7,8 @@ import ( "testing" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestLogConnectionsTakesEffect is a regression test for a postgresql.conf change @@ -25,18 +27,15 @@ import ( // This test flips log_connections on against a running cluster and asserts the // effective value read back through the gateway actually becomes "on". func TestLogConnectionsTakesEffect(t *testing.T) { + ck := assert.NewCollecting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.Require().NoError(err, "create CR client") ctx := context.Background() cr := framework.MustLoadCluster("config/samples/no-templates.yaml", ns) framework.WithCIResources(&cr.Spec) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.Require().NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for Postgres to come up and serve queries through the gateway. framework.WaitForPod(t, c, ns, "postgres") @@ -52,9 +51,7 @@ func TestLogConnectionsTakesEffect(t *testing.T) { // the initial config still converging. framework.WaitForShardConfigSettled(t, c, ns) before := poolPodUIDs(t, ctx, c, ns) - if len(before) == 0 { - t.Fatal("no pool pods found before enabling log_connections") - } + ck.Require().NotEmpty(before, "no pool pods found before enabling log_connections") // Turn log_connections on via the inline spec.postgresConfig map. Send the full // databases array with the one changed value: JSON merge-patch replaces arrays @@ -86,7 +83,10 @@ func TestLogConnectionsTakesEffect(t *testing.T) { break } } - if !recreated { - t.Errorf("expected pool pods to be recreated to apply log_connections, but UIDs were unchanged: before=%v after=%v", before, after) - } + ck.True( + recreated, + "expected pool pods to be recreated to apply log_connections, but UIDs were unchanged: before=%v after=%v", + before, + after, + ) } diff --git a/test/e2e/shared/postgresconfig/postgresconfig_test.go b/test/e2e/shared/postgresconfig/postgresconfig_test.go index 6a111d9c..e1f46dba 100644 --- a/test/e2e/shared/postgresconfig/postgresconfig_test.go +++ b/test/e2e/shared/postgresconfig/postgresconfig_test.go @@ -14,6 +14,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestPostgresConfigManagement exercises the whole PostgreSQL-config feature @@ -25,11 +27,10 @@ import ( // through the gateway, so it proves the setting took effect in the running // server — not merely that a ConfigMap has the right text. func TestPostgresConfigManagement(t *testing.T) { + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Legacy postgresConfigRef ConfigMap: sets seq_page_cost (a key the operator's @@ -42,14 +43,10 @@ func TestPostgresConfigManagement(t *testing.T) { "custom.conf": "seq_page_cost = '2.5'\nwork_mem = '64MB'", }, } - if err := c.Create(ctx, refCM); err != nil { - t.Fatalf("create ref ConfigMap: %v", err) - } + ck.NoError(c.Create(ctx, refCM), "create ref ConfigMap") cr := configCluster(ns) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for Postgres to come up and serve queries. framework.WaitForPod(t, c, ns, "postgres") @@ -86,19 +83,16 @@ func TestPostgresConfigManagement(t *testing.T) { }) t.Run("operator renders a per-shard ConfigMap", func(t *testing.T) { + ck := assert.NewCollecting(t) cms := &corev1.ConfigMapList{} - if err := c.List(ctx, cms, client.InNamespace(ns)); err != nil { - t.Fatalf("list ConfigMaps: %v", err) - } + ck.Require().NoError(c.List(ctx, cms, client.InNamespace(ns)), "list ConfigMaps") found := false for i := range cms.Items { if strings.HasSuffix(cms.Items[i].Name, "-postgres-config") { found = true } } - if !found { - t.Error("expected an operator-rendered -postgres-config ConfigMap") - } + ck.True(found, "expected an operator-rendered -postgres-config ConfigMap") }) // A reload-safe change (work_mem, user context) must take effect in the @@ -111,15 +105,14 @@ func TestPostgresConfigManagement(t *testing.T) { // the in-place reload is exercised without a concurrent rolling restart // recreating the pods out from under it. t.Run("reload-safe change applies without recreating pods", func(t *testing.T) { + ck := assert.NewCollecting(t) // Ensure the initial config rollout has fully settled before changing // work_mem, so the in-place reload runs on a stable cluster rather than // racing the primary-last restart still converging the initial config. framework.WaitForShardConfigSettled(t, c, ns) before := poolPodUIDs(t, ctx, c, ns) - if len(before) == 0 { - t.Fatal("no pool pods found before reload-safe change") - } + ck.Require().NotEmpty(before, "no pool pods found before reload-safe change") live := framework.GetCluster(t, c, ns, cr.Name) live.Spec.Databases[0].TableGroups[0].Shards[0].Spec.PostgresConfig["work_mem"] = "32MB" @@ -135,13 +128,20 @@ func TestPostgresConfigManagement(t *testing.T) { // The pods must be the very same objects — a reload-safe change must not // recreate them. Stable UIDs prove Postgres was never restarted. after := poolPodUIDs(t, ctx, c, ns) - if len(after) != len(before) { - t.Errorf("pool pod set changed across a reload-safe change: before=%v after=%v", before, after) - } + ck.Len( + after, + len(before), + "pool pod set changed across a reload-safe change: before=%v after=", + before, + ) for name, uid := range before { - if after[name] != uid { - t.Errorf("pool pod %s was recreated by a reload-safe change: UID %s -> %s", name, uid, after[name]) - } + ck.Eq( + uid, + after[name], + "pool pod %s was recreated by a reload-safe change: UID %s ->", + name, + uid, + ) } }) @@ -161,6 +161,7 @@ func TestPostgresConfigManagement(t *testing.T) { // PostgreSQL default (0.01). Runs BEFORE the restart subtest, on a settled // cluster, so the in-place reload is not raced by a concurrent rolling restart. t.Run("removing a reload-safe setting reverts it in place", func(t *testing.T) { + ck := assert.NewCollecting(t) framework.WaitForShardConfigSettled(t, c, ns) // Add cpu_tuple_cost via the inline map and confirm it reloads in. @@ -172,16 +173,17 @@ func TestPostgresConfigManagement(t *testing.T) { framework.WaitForPsqlValue(t, cluster, ns, gw, "SHOW cpu_tuple_cost", "0.05") before := poolPodUIDs(t, ctx, c, ns) - if len(before) == 0 { - t.Fatal("no pool pods found before the removal") - } + ck.Require().NotEmpty(before, "no pool pods found before the removal") // Now REMOVE it entirely; it must revert to the default (0.01) in place. // This is the marker's job: the removal leaves every other expected setting // unchanged, so only the moved marker forces the pooler to wait for the // kubelet-synced file before confirming the reload. live = framework.GetCluster(t, c, ns, cr.Name) - delete(live.Spec.Databases[0].TableGroups[0].Shards[0].Spec.PostgresConfig, "cpu_tuple_cost") + delete( + live.Spec.Databases[0].TableGroups[0].Shards[0].Spec.PostgresConfig, + "cpu_tuple_cost", + ) framework.PatchCluster(t, c, live, framework.MustMarshal(map[string]any{ "spec": map[string]any{"databases": live.Spec.Databases}, })) @@ -189,13 +191,20 @@ func TestPostgresConfigManagement(t *testing.T) { // The removal was applied by an in-place reload — no pod recreation. after := poolPodUIDs(t, ctx, c, ns) - if len(after) != len(before) { - t.Errorf("pool pod set changed across a reload-safe removal: before=%v after=%v", before, after) - } + ck.Len( + after, + len(before), + "pool pod set changed across a reload-safe removal: before=%v after=", + before, + ) for name, uid := range before { - if after[name] != uid { - t.Errorf("pool pod %s was recreated by a reload-safe removal: UID %s -> %s", name, uid, after[name]) - } + ck.Eq( + uid, + after[name], + "pool pod %s was recreated by a reload-safe removal: UID %s ->", + name, + uid, + ) } }) diff --git a/test/e2e/shared/scaling/scaling_test.go b/test/e2e/shared/scaling/scaling_test.go index af06d1f6..cde5d9f5 100644 --- a/test/e2e/shared/scaling/scaling_test.go +++ b/test/e2e/shared/scaling/scaling_test.go @@ -9,6 +9,8 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestStatelessScaling verifies that scaling stateless components (multiadmin, @@ -22,16 +24,13 @@ func TestStatelessScaling(t *testing.T) { func testScaleMultiadmin(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Verify initial: 1 multiadmin replica. @@ -53,16 +52,13 @@ func testScaleMultiadmin(t *testing.T) { func testScaleMultigateway(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Verify initial: 1 multigateway replica. @@ -91,16 +87,13 @@ func testScaleMultigateway(t *testing.T) { func testLargeScaleMultiadmin(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Scale 1 → 5. @@ -120,16 +113,13 @@ func testLargeScaleMultiadmin(t *testing.T) { func testLargeScaleMultigateway(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Scale 1 → 5. diff --git a/test/e2e/shared/templated/templated_test.go b/test/e2e/shared/templated/templated_test.go index 4a1075fa..e0b3fe7b 100644 --- a/test/e2e/shared/templated/templated_test.go +++ b/test/e2e/shared/templated/templated_test.go @@ -8,6 +8,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestTemplatedCluster applies the template CRs from config/samples/templates/ @@ -15,32 +17,23 @@ import ( // verifies the full resource tree, pod health, and psql connectivity. func TestTemplatedCluster(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Create templates first — the cluster CR references them. coreTmpl := framework.MustLoadCoreTemplate("config/samples/templates/core.yaml", ns) - if err := c.Create(ctx, coreTmpl); err != nil { - t.Fatalf("create CoreTemplate: %v", err) - } + ck.NoError(c.Create(ctx, coreTmpl), "create CoreTemplate") cellTmpl := framework.MustLoadCellTemplate("config/samples/templates/cell.yaml", ns) - if err := c.Create(ctx, cellTmpl); err != nil { - t.Fatalf("create CellTemplate: %v", err) - } + ck.NoError(c.Create(ctx, cellTmpl), "create CellTemplate") shardTmpl := framework.MustLoadShardTemplate("config/samples/templates/shard.yaml", ns) - if err := c.Create(ctx, shardTmpl); err != nil { - t.Fatalf("create ShardTemplate: %v", err) - } + ck.NoError(c.Create(ctx, shardTmpl), "create ShardTemplate") // Create the cluster referencing the templates. cr := framework.MustLoadCluster("config/samples/templated-cluster.yaml", ns) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") // Verify child CRDs. framework.WaitForCRDCount(t, c, ns, diff --git a/test/e2e/shared/templates/templates_test.go b/test/e2e/shared/templates/templates_test.go index bd417e18..82e31321 100644 --- a/test/e2e/shared/templates/templates_test.go +++ b/test/e2e/shared/templates/templates_test.go @@ -6,11 +6,11 @@ import ( "context" "testing" - "github.com/stretchr/testify/require" - multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/multigres/testkit/assert" ) // TestTemplatePropagation verifies that template values propagate correctly @@ -23,32 +23,23 @@ func TestTemplatePropagation(t *testing.T) { func testVerifyPropagation(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.Require().NoError(err, "create CR client") ctx := context.Background() // Create templates first. coreTmpl := framework.MustLoadCoreTemplate("test/e2e/fixtures/templates/core.yaml", ns) - if err := c.Create(ctx, coreTmpl); err != nil { - t.Fatalf("create CoreTemplate: %v", err) - } + ck.Require().NoError(c.Create(ctx, coreTmpl), "create CoreTemplate") cellTmpl := framework.MustLoadCellTemplate("test/e2e/fixtures/templates/cell.yaml", ns) - if err := c.Create(ctx, cellTmpl); err != nil { - t.Fatalf("create CellTemplate: %v", err) - } + ck.Require().NoError(c.Create(ctx, cellTmpl), "create CellTemplate") shardTmpl := framework.MustLoadShardTemplate("test/e2e/fixtures/templates/shard.yaml", ns) - if err := c.Create(ctx, shardTmpl); err != nil { - t.Fatalf("create ShardTemplate: %v", err) - } + ck.Require().NoError(c.Create(ctx, shardTmpl), "create ShardTemplate") // Create cluster referencing the templates. cr := framework.MustLoadCluster("test/e2e/fixtures/templated.yaml", ns) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.Require().NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for Shard CRD to be created. framework.WaitForCRDCount(t, c, ns, @@ -59,21 +50,21 @@ func testVerifyPropagation(t *testing.T) { // Verify Shard inherited values from templates. shards := &multigresv1alpha1.ShardList{} - if err := c.List(ctx, shards, client.InNamespace(ns)); err != nil { - t.Fatalf("list Shards: %v", err) - } + ck.Require().NoError(c.List(ctx, shards, client.InNamespace(ns)), "list Shards") shard := shards.Items[0] - require.Len(t, shard.Spec.Pools, len(shardTmpl.Spec.Pools)) + ck.Require().Len(shard.Spec.Pools, len(shardTmpl.Spec.Pools)) // Pool storage from ShardTemplate should be 1Gi. for poolName, pool := range shard.Spec.Pools { - require.Equal( - t, shardTmpl.Spec.Pools[poolName].ReplicasPerCell, pool.ReplicasPerCell, - "pool %s must inherit the template's failure-safe replica count", poolName, + ck.Require(). + EqDeep(shardTmpl.Spec.Pools[poolName].ReplicasPerCell, pool.ReplicasPerCell, "pool %s must inherit the template's failure-safe replica count", poolName) + ck.Eq( + "1Gi", + pool.Storage.Size, + "pool %s storage = %s, want 1Gi (from ShardTemplate)", + poolName, + pool.Storage.Size, ) - if pool.Storage.Size != "1Gi" { - t.Errorf("pool %s storage = %s, want 1Gi (from ShardTemplate)", poolName, pool.Storage.Size) - } } // Wait for all pods to come up. @@ -81,46 +72,30 @@ func testVerifyPropagation(t *testing.T) { // Check actual resolution, not values that could also come from defaults. live := framework.GetCluster(t, c, ns, cr.Name) - require.NotNil(t, live.Status.ResolvedTemplates) - require.ElementsMatch( - t, - []multigresv1alpha1.TemplateRef{"e2e-core"}, - live.Status.ResolvedTemplates.CoreTemplates, - ) - require.ElementsMatch( - t, - []multigresv1alpha1.TemplateRef{"e2e-cell"}, - live.Status.ResolvedTemplates.CellTemplates, - ) - require.ElementsMatch( - t, - []multigresv1alpha1.TemplateRef{"e2e-shard"}, - live.Status.ResolvedTemplates.ShardTemplates, - ) + ck.Require().NotNil(live.Status.ResolvedTemplates) + ck.Require(). + ElementsMatch([]multigresv1alpha1.TemplateRef{"e2e-core"}, live.Status.ResolvedTemplates.CoreTemplates) + ck.Require(). + ElementsMatch([]multigresv1alpha1.TemplateRef{"e2e-cell"}, live.Status.ResolvedTemplates.CellTemplates) + ck.Require(). + ElementsMatch([]multigresv1alpha1.TemplateRef{"e2e-shard"}, live.Status.ResolvedTemplates.ShardTemplates) } func testPartialOverride(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") ctx := context.Background() // Create templates. coreTmpl := framework.MustLoadCoreTemplate("test/e2e/fixtures/templates/core.yaml", ns) - if err := c.Create(ctx, coreTmpl); err != nil { - t.Fatalf("create CoreTemplate: %v", err) - } + ck.NoError(c.Create(ctx, coreTmpl), "create CoreTemplate") cellTmpl := framework.MustLoadCellTemplate("test/e2e/fixtures/templates/cell.yaml", ns) - if err := c.Create(ctx, cellTmpl); err != nil { - t.Fatalf("create CellTemplate: %v", err) - } + ck.NoError(c.Create(ctx, cellTmpl), "create CellTemplate") shardTmpl := framework.MustLoadShardTemplate("test/e2e/fixtures/templates/shard.yaml", ns) - if err := c.Create(ctx, shardTmpl); err != nil { - t.Fatalf("create ShardTemplate: %v", err) - } + ck.NoError(c.Create(ctx, shardTmpl), "create ShardTemplate") // Create cluster with an inline override on multiadmin replicas. cr := framework.MustLoadCluster("test/e2e/fixtures/templated.yaml", ns) @@ -130,9 +105,7 @@ func testPartialOverride(t *testing.T) { }, } framework.WithCIResources(&cr.Spec) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for multiadmin to have 2 replicas (override wins over template's 1). framework.WaitForDeploymentReplicas(t, c, ns, "multiadmin", 2) @@ -141,32 +114,23 @@ func testPartialOverride(t *testing.T) { func testPVCDeletionPolicyInheritance(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.Require().NoError(err, "create CR client") ctx := context.Background() // Create templates — ShardTemplate has pvcDeletionPolicy: Delete/Delete. coreTmpl := framework.MustLoadCoreTemplate("test/e2e/fixtures/templates/core.yaml", ns) - if err := c.Create(ctx, coreTmpl); err != nil { - t.Fatalf("create CoreTemplate: %v", err) - } + ck.Require().NoError(c.Create(ctx, coreTmpl), "create CoreTemplate") cellTmpl := framework.MustLoadCellTemplate("test/e2e/fixtures/templates/cell.yaml", ns) - if err := c.Create(ctx, cellTmpl); err != nil { - t.Fatalf("create CellTemplate: %v", err) - } + ck.Require().NoError(c.Create(ctx, cellTmpl), "create CellTemplate") shardTmpl := framework.MustLoadShardTemplate("test/e2e/fixtures/templates/shard.yaml", ns) - if err := c.Create(ctx, shardTmpl); err != nil { - t.Fatalf("create ShardTemplate: %v", err) - } + ck.Require().NoError(c.Create(ctx, shardTmpl), "create ShardTemplate") // Create cluster referencing templates. cr := framework.MustLoadCluster("test/e2e/fixtures/templated.yaml", ns) - if err := c.Create(ctx, cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.Require().NoError(c.Create(ctx, cr), "create MultigresCluster") // Wait for Shard to be created. framework.WaitForCRDCount(t, c, ns, @@ -177,19 +141,12 @@ func testPVCDeletionPolicyInheritance(t *testing.T) { // Verify Shard inherited the PVC deletion policy from ShardTemplate. shards := &multigresv1alpha1.ShardList{} - if err := c.List(ctx, shards, client.InNamespace(ns)); err != nil { - t.Fatalf("list Shards: %v", err) - } + ck.Require().NoError(c.List(ctx, shards, client.InNamespace(ns)), "list Shards") shard := shards.Items[0] - if shard.Spec.PVCDeletionPolicy == nil { - t.Fatal("Shard PVCDeletionPolicy is nil, expected inheritance from ShardTemplate") - } - if shard.Spec.PVCDeletionPolicy.WhenDeleted != "Delete" { - t.Errorf("PVCDeletionPolicy.WhenDeleted = %s, want Delete", shard.Spec.PVCDeletionPolicy.WhenDeleted) - } - if shard.Spec.PVCDeletionPolicy.WhenScaled != "Delete" { - t.Errorf("PVCDeletionPolicy.WhenScaled = %s, want Delete", shard.Spec.PVCDeletionPolicy.WhenScaled) - } + ck.Require(). + NotNil(shard.Spec.PVCDeletionPolicy, "Shard PVCDeletionPolicy is nil, expected inheritance from ShardTemplate") + ck.Eq("Delete", shard.Spec.PVCDeletionPolicy.WhenDeleted, "PVCDeletionPolicy.WhenDeleted") + ck.Eq("Delete", shard.Spec.PVCDeletionPolicy.WhenScaled, "PVCDeletionPolicy.WhenScaled") } func int32Ptr(i int32) *int32 { return &i } diff --git a/test/e2e/shared/topotls/topotls_test.go b/test/e2e/shared/topotls/topotls_test.go index dc88232a..fccb73c9 100644 --- a/test/e2e/shared/topotls/topotls_test.go +++ b/test/e2e/shared/topotls/topotls_test.go @@ -13,6 +13,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestTopoTLSCluster brings up a cluster with topoTLS enabled and verifies it @@ -23,6 +25,7 @@ import ( // anywhere in that path would leave the cluster unable to become ready. func TestTopoTLSCluster(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) // cert-manager and the shared CA the operator signs the etcd serving // certificate and the client credential with. @@ -30,14 +33,10 @@ func TestTopoTLSCluster(t *testing.T) { ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/topo-tls.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") // The child resource tree still forms with topology TLS on. framework.WaitForCRDCount(t, c, ns, @@ -76,14 +75,18 @@ func waitForTopologyReady(t *testing.T, c client.Client, ns, name string) { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute) defer cancel() - err := wait.PollUntilContextCancel(ctx, 5*time.Second, true, func(ctx context.Context) (bool, error) { - cluster := &multigresv1alpha1.MultigresCluster{} - if err := c.Get(ctx, client.ObjectKey{Namespace: ns, Name: name}, cluster); err != nil { - return false, nil - } - return meta.IsStatusConditionTrue(cluster.Status.Conditions, "TopologyReady"), nil - }) - if err != nil { - t.Fatalf("timed out waiting for TopologyReady on cluster %s/%s: %v", ns, name, err) - } + err := wait.PollUntilContextCancel( + ctx, + 5*time.Second, + true, + func(ctx context.Context) (bool, error) { + cluster := &multigresv1alpha1.MultigresCluster{} + if err := c.Get(ctx, client.ObjectKey{Namespace: ns, Name: name}, cluster); err != nil { + return false, nil + } + return meta.IsStatusConditionTrue(cluster.Status.Conditions, "TopologyReady"), nil + }, + ) + assert.NewAborting(t). + NoError(err, "timed out waiting for TopologyReady on cluster %s/%s", ns, name) } diff --git a/test/e2e/shared/verification/verification_test.go b/test/e2e/shared/verification/verification_test.go index bc0fe7f6..78e813a5 100644 --- a/test/e2e/shared/verification/verification_test.go +++ b/test/e2e/shared/verification/verification_test.go @@ -9,7 +9,6 @@ import ( "testing" "time" - "github.com/stretchr/testify/require" corev1 "k8s.io/api/core/v1" policyv1 "k8s.io/api/policy/v1" apiresource "k8s.io/apimachinery/pkg/api/resource" @@ -24,6 +23,8 @@ import ( shardcontroller "github.com/multigres/multigres-operator/pkg/resource-handler/controller/shard" "github.com/multigres/multigres-operator/pkg/util/metadata" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestResourceVerification verifies that the operator creates the expected @@ -36,12 +37,11 @@ func TestResourceVerification(t *testing.T) { } func testMultiCellFilesystemBackup(t *testing.T) { + ck := assert.NewCollecting(t) ns := cluster.CreateNamespace(t) createStaticRWXVolume(t, ns) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.Require().NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) cr.Name = "multi-cell-fs-backup" @@ -67,83 +67,83 @@ func testMultiCellFilesystemBackup(t *testing.T) { pool.ReplicasPerCell = &replicas cr.Spec.Databases[0].TableGroups[0].Shards[0].Spec.Pools["default"] = pool - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.Require().NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) claims := &corev1.PersistentVolumeClaimList{} - if err := c.List(context.Background(), claims, + ck.Require().NoError(c.List(context.Background(), claims, client.InNamespace(ns), client.MatchingLabels{ metadata.LabelMultigresCluster: cr.Name, metadata.LabelMultigresDatabase: string(cr.Spec.Databases[0].Name), metadata.LabelMultigresTableGroup: string(cr.Spec.Databases[0].TableGroups[0].Name), - metadata.LabelMultigresShard: string(cr.Spec.Databases[0].TableGroups[0].Shards[0].Name), + metadata.LabelMultigresShard: string( + cr.Spec.Databases[0].TableGroups[0].Shards[0].Name, + ), }, - ); err != nil { - t.Fatalf("list backup PVCs: %v", err) - } + ), "list backup PVCs") var backupClaims []corev1.PersistentVolumeClaim for _, candidate := range claims.Items { if candidate.Labels[metadata.LabelMultigresPool] == "" { backupClaims = append(backupClaims, candidate) } } - if len(backupClaims) != 1 { - t.Fatalf("backup PVC count = %d, want 1", len(backupClaims)) - } + ck.Require().Len(backupClaims, 1, "backup PVC count = %d, want 1", len(backupClaims)) claim := &backupClaims[0] if len(claim.Spec.AccessModes) != 1 || claim.Spec.AccessModes[0] != corev1.ReadWriteMany { t.Fatalf("backup PVC access modes = %v, want [ReadWriteMany]", claim.Spec.AccessModes) } pods := &corev1.PodList{} - if err := c.List(context.Background(), pods, + ck.Require().NoError(c.List(context.Background(), pods, client.InNamespace(ns), client.MatchingLabels{ metadata.LabelMultigresCluster: cr.Name, metadata.LabelMultigresPool: "default", }, - ); err != nil { - t.Fatalf("list pooler pods: %v", err) - } - if len(pods.Items) != 4 { - t.Fatalf("pooler pod count = %d, want 4", len(pods.Items)) - } + ), "list pooler pods") + ck.Require().Len(pods.Items, 4, "pooler pod count = %d, want 4", len(pods.Items)) for _, pod := range pods.Items { var mountedClaim string for _, volume := range pod.Spec.Volumes { - if volume.Name == shardcontroller.BackupVolumeName && volume.PersistentVolumeClaim != nil { + if volume.Name == shardcontroller.BackupVolumeName && + volume.PersistentVolumeClaim != nil { mountedClaim = volume.PersistentVolumeClaim.ClaimName break } } - if mountedClaim != claim.Name { - t.Errorf("pod %q mounts backup claim %q, want %q", pod.Name, mountedClaim, claim.Name) - } + ck.Eq( + claim.Name, + mountedClaim, + "pod %q mounts backup claim %q, want", + pod.Name, + mountedClaim, + ) } ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) defer cancel() - err = wait.PollUntilContextCancel(ctx, 3*time.Second, true, func(ctx context.Context) (bool, error) { - shards := &multigresv1alpha1.ShardList{} - if err := c.List(ctx, shards, - client.InNamespace(ns), - client.MatchingLabels{metadata.LabelMultigresCluster: cr.Name}, - ); err != nil || len(shards.Items) != 1 { - return false, nil - } - for _, role := range shards.Items[0].Status.PodRoles { - if role == "PRIMARY" { - return true, nil + err = wait.PollUntilContextCancel( + ctx, + 3*time.Second, + true, + func(ctx context.Context) (bool, error) { + shards := &multigresv1alpha1.ShardList{} + if err := c.List(ctx, shards, + client.InNamespace(ns), + client.MatchingLabels{metadata.LabelMultigresCluster: cr.Name}, + ); err != nil || len(shards.Items) != 1 { + return false, nil } - } - return false, nil - }) - if err != nil { - t.Fatalf("timed out waiting for bootstrap to elect a primary: %v", err) - } + for _, role := range shards.Items[0].Status.PodRoles { + if role == "PRIMARY" { + return true, nil + } + } + return false, nil + }, + ) + ck.Require().NoError(err, "timed out waiting for bootstrap to elect a primary") } // createStaticRWXVolume supplies the claim used by this test. A Kind cluster @@ -156,38 +156,42 @@ func createStaticRWXVolume(t *testing.T, namespace string) { hostPath := "/var/local/multigres-e2e-rwx/" + namespace prepareStaticRWXHostPath(t, hostPath) - _, err := cluster.Clientset.CoreV1().PersistentVolumes().Create(context.Background(), &corev1.PersistentVolume{ - ObjectMeta: metav1.ObjectMeta{Name: name}, - Spec: corev1.PersistentVolumeSpec{ - Capacity: corev1.ResourceList{ - corev1.ResourceStorage: apiresource.MustParse("1Gi"), - }, - AccessModes: []corev1.PersistentVolumeAccessMode{corev1.ReadWriteMany}, - PersistentVolumeReclaimPolicy: corev1.PersistentVolumeReclaimRetain, - StorageClassName: "e2e-rwx", - NodeAffinity: &corev1.VolumeNodeAffinity{ - Required: &corev1.NodeSelector{ - NodeSelectorTerms: []corev1.NodeSelectorTerm{{ - MatchExpressions: []corev1.NodeSelectorRequirement{{ - Key: "node-role.kubernetes.io/control-plane", - Operator: corev1.NodeSelectorOpExists, + _, err := cluster.Clientset.CoreV1(). + PersistentVolumes(). + Create(context.Background(), &corev1.PersistentVolume{ + ObjectMeta: metav1.ObjectMeta{Name: name}, + Spec: corev1.PersistentVolumeSpec{ + Capacity: corev1.ResourceList{ + corev1.ResourceStorage: apiresource.MustParse("1Gi"), + }, + AccessModes: []corev1.PersistentVolumeAccessMode{ + corev1.ReadWriteMany, + }, + PersistentVolumeReclaimPolicy: corev1.PersistentVolumeReclaimRetain, + StorageClassName: "e2e-rwx", + NodeAffinity: &corev1.VolumeNodeAffinity{ + Required: &corev1.NodeSelector{ + NodeSelectorTerms: []corev1.NodeSelectorTerm{{ + MatchExpressions: []corev1.NodeSelectorRequirement{{ + Key: "node-role.kubernetes.io/control-plane", + Operator: corev1.NodeSelectorOpExists, + }}, }}, - }}, + }, }, - }, - PersistentVolumeSource: corev1.PersistentVolumeSource{ - HostPath: &corev1.HostPathVolumeSource{ - Path: hostPath, - Type: ptr.To(corev1.HostPathDirectoryOrCreate), + PersistentVolumeSource: corev1.PersistentVolumeSource{ + HostPath: &corev1.HostPathVolumeSource{ + Path: hostPath, + Type: ptr.To(corev1.HostPathDirectoryOrCreate), + }, }, }, - }, - }, metav1.CreateOptions{}) - if err != nil { - t.Fatalf("create static RWX volume: %v", err) - } + }, metav1.CreateOptions{}) + assert.NewAborting(t).NoError(err, "create static RWX volume") t.Cleanup(func() { - _ = cluster.Clientset.CoreV1().PersistentVolumes().Delete(context.Background(), name, metav1.DeleteOptions{}) + _ = cluster.Clientset.CoreV1(). + PersistentVolumes(). + Delete(context.Background(), name, metav1.DeleteOptions{}) }) } @@ -196,20 +200,18 @@ func prepareStaticRWXHostPath(t *testing.T, path string) { node := cluster.Name + "-control-plane" for _, args := range [][]string{{"mkdir", "-p", path}, {"chmod", "0777", path}} { - output, err := exec.CommandContext(context.Background(), "docker", append([]string{"exec", node}, args...)...).CombinedOutput() - if err != nil { - t.Fatalf("%s static RWX hostPath: %v\n%s", args[0], err, output) - } + output, err := exec.CommandContext(context.Background(), "docker", append([]string{"exec", node}, args...)...). + CombinedOutput() + assert.NewAborting(t).NoError(err, "%s static RWX hostPath: %v\n%s", args[0], err, output) } } func testPDB(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) // Four members across two pools and cells must share one shard-wide budget. @@ -222,20 +224,18 @@ func testPDB(t *testing.T) { extra.ReplicasPerCell = ptr.To(int32(1)) extra.Cells = []multigresv1alpha1.CellName{"zone-b"} pools["extra"] = *extra - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") poolLabels := client.MatchingLabels{ metadata.LabelAppComponent: shardcontroller.PoolComponentName, } framework.WaitForPodCount(t, c, ns, poolLabels, 4, "poolers across both pools") cluster.WaitForAllPodsReady(t, ns) shards := &multigresv1alpha1.ShardList{} - require.NoError(t, c.List(context.Background(), shards, client.InNamespace(ns))) - require.Len(t, shards.Items, 1) + ck.NoError(c.List(context.Background(), shards, client.InNamespace(ns))) + ck.Len(shards.Items, 1) poolers := &corev1.PodList{} - require.NoError(t, c.List(context.Background(), poolers, client.InNamespace(ns), poolLabels)) - require.Len(t, poolers.Items, 4) + ck.NoError(c.List(context.Background(), poolers, client.InNamespace(ns), poolLabels)) + ck.Len(poolers.Items, 4) pdbs := framework.ListPDBs(t, c, ns) var shardPDBs []policyv1.PodDisruptionBudget @@ -244,58 +244,51 @@ func testPDB(t *testing.T) { shardPDBs = append(shardPDBs, pdb) } } - require.Len(t, shardPDBs, 1) + ck.Len(shardPDBs, 1) pdb := &shardPDBs[0] minimum := intstr.FromInt32(3) - require.Equal(t, &minimum, pdb.Spec.MinAvailable) - require.Nil(t, pdb.Spec.MaxUnavailable) - require.NotNil(t, pdb.Spec.Selector) + ck.EqDeep(&minimum, pdb.Spec.MinAvailable) + ck.Nil(pdb.Spec.MaxUnavailable) + ck.NotNil(pdb.Spec.Selector) selector, err := metav1.LabelSelectorAsSelector(pdb.Spec.Selector) - require.NoError(t, err) + ck.NoError(err) for _, pod := range poolers.Items { - require.True( - t, selector.Matches(labels.Set(pod.Labels)), "shard PDB must cover %s", pod.Name, - ) + ck.True(selector.Matches(labels.Set(pod.Labels)), "shard PDB must cover %s", pod.Name) matches := 0 for _, candidate := range pdbs { selector, err := metav1.LabelSelectorAsSelector(candidate.Spec.Selector) - require.NoError(t, err) + ck.NoError(err) if selector.Matches(labels.Set(pod.Labels)) { matches++ } } - require.Equal(t, 1, matches, "pooler %s must not match overlapping PDBs", pod.Name) + ck.EqDeep(1, matches, "pooler %s must not match overlapping PDBs", pod.Name) } // Verify the Kubernetes disruption controller agrees with the desired budget. - require.Eventually(t, func() bool { + ck.EventuallyTrue(time.Minute, time.Second, func() bool { if err := c.Get(context.Background(), client.ObjectKeyFromObject(pdb), pdb); err != nil { return false } return pdb.Status.ObservedGeneration == pdb.Generation && pdb.Status.CurrentHealthy == 4 && pdb.Status.DesiredHealthy == 3 && pdb.Status.DisruptionsAllowed == 1 - }, time.Minute, time.Second) + }) } func testMultiadminWeb(t *testing.T) { t.Parallel() + ck := assert.NewCollecting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.Require().NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.Require().NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Verify multiadminweb deployment exists (container name has a hyphen). dep := framework.WaitForDeployment(t, c, ns, "multiadmin-web") - if dep.Status.ReadyReplicas < 1 { - t.Errorf("multiadmin-web has %d ready replicas, want >= 1", dep.Status.ReadyReplicas) - } + ck.GreaterOrEqual(1, dep.Status.ReadyReplicas, "multiadmin-web has") // Verify multiadminweb service exists. framework.WaitForService(t, c, ns, "http", 18100) @@ -303,24 +296,19 @@ func testMultiadminWeb(t *testing.T) { func testLogLevels(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") cr := framework.MustLoadCluster("test/e2e/fixtures/log-levels.yaml", ns) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Check that pods have the expected --log-level settings. ctx := context.Background() pods := &corev1.PodList{} - if err := c.List(ctx, pods, client.InNamespace(ns)); err != nil { - t.Fatalf("list pods: %v", err) - } + ck.NoError(c.List(ctx, pods, client.InNamespace(ns)), "list pods") expectedLevels := map[string]string{ "multipooler": "warn", diff --git a/test/e2e/shared/webhook/webhook_test.go b/test/e2e/shared/webhook/webhook_test.go index a3f80ce8..b56558e6 100644 --- a/test/e2e/shared/webhook/webhook_test.go +++ b/test/e2e/shared/webhook/webhook_test.go @@ -8,6 +8,8 @@ import ( multigresv1alpha1 "github.com/multigres/multigres-operator/api/v1alpha1" "github.com/multigres/multigres-operator/test/e2e/framework" + + "github.com/multigres/testkit/assert" ) // TestWebhookRejections verifies that the admission rules correctly reject @@ -24,11 +26,10 @@ func TestWebhookRejections(t *testing.T) { func testRemoveCell(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") // Create cluster with 2 cells so we can try removing one. cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) @@ -36,9 +37,7 @@ func testRemoveCell(t *testing.T) { Name: "zone-b", Region: "us-central1", }) - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Get live CR and remove the second cell. @@ -50,11 +49,10 @@ func testRemoveCell(t *testing.T) { func testRemovePool(t *testing.T) { t.Parallel() + ck := assert.NewAborting(t) ns := cluster.CreateNamespace(t) c, err := cluster.CRClient() - if err != nil { - t.Fatalf("create CR client: %v", err) - } + ck.NoError(err, "create CR client") // Create cluster with 2 pools. cr := framework.MustLoadCluster("test/e2e/fixtures/base.yaml", ns) @@ -67,9 +65,7 @@ func testRemovePool(t *testing.T) { Size: "1Gi", }, } - if err := c.Create(context.Background(), cr); err != nil { - t.Fatalf("create MultigresCluster: %v", err) - } + ck.NoError(c.Create(context.Background(), cr), "create MultigresCluster") cluster.WaitForAllPodsReady(t, ns) // Get live CR and remove the extra pool.