diff --git a/.github/workflows/build-and-release.yaml b/.github/workflows/build-and-release.yaml index aedf85a0a..9ceb2873e 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 000000000..ce840c222 --- /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 diff --git a/.golangci.toml b/.golangci.toml index cf89d7278..ce4a124a9 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 3c215ac3c..32a24b310 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 235d2b31d..2c2746222 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 @@ -278,22 +278,65 @@ 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/... + +# 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)" \ - 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" @@ -679,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 @@ -695,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/api/v1alpha1/etcd_maintenance_test.go b/api/v1alpha1/etcd_maintenance_test.go index f4ddd9858..214f518ff 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 4f598b2a6..0522f963c 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 b20256662..03ffb3776 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 ea56ffb3a..d225e68ee 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 19e23d9ad..d2b7b2bad 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 afd81bb42..0515c1356 100644 --- a/go.mod +++ b/go.mod @@ -1,28 +1,30 @@ module github.com/multigres/multigres-operator -go 1.26.6 +go 1.27.1 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.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 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 ( @@ -98,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 @@ -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 f263bbdfc..85359d114 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.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= @@ -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/pkg/cert/generator_test.go b/pkg/cert/generator_test.go index 3afda7f66..60a3d62f6 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 ae566218c..660169d3a 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 612d2d457..08105d0bf 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 dd29b6546..6bf633e55 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 d99080267..f0b8d4fa5 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 f5cc5063d..8a71f94db 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 5aed7a7d7..431cc7fac 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 e22e57dc8..03db82c85 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 663f75bfd..6b5193bbe 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 7ca51c944..ff03bcf37 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 b439a6552..3a737dec3 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 073fc335c..5c87191a1 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 0330b5d39..ccdbacddf 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 f27eeecc1..2ef5bf022 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 a95a8b230..2462a109c 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 2b70fe5eb..ea581330e 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 c906d2f27..459080971 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 7a4c02e53..4f196c842 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 6d0764c69..d95911301 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 dab8c95e9..dcd95c306 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 51cdbc4c8..50d6fc95d 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 d2b543db5..e3c368648 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 997e15f39..6a95e7cdb 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 99ba1d92e..c3d9e931e 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 797c05490..a3caa1691 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 95146a4c5..cfbaf6edc 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 30fa87622..48d0acd54 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 55156c343..be39982b2 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 a15a7a66c..210ef9971 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 75fef08e8..cd0ddbe99 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 82547ba13..df295bba3 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 3f399fc0c..9e75e50bd 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 c3cd7c2e7..77131c7c0 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 1de4ed8b6..ae3512014 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.go b/pkg/data-handler/posture/posture.go index 049b8cebc..8d0a83b6b 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/data-handler/posture/posture_test.go b/pkg/data-handler/posture/posture_test.go index 59423e455..9800fe529 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 18878789b..b1f4392da 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 462445855..106b4e9ae 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 0c91b3645..efd540fe0 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 cb95ee85f..8edf4000f 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 f0a9a132b..151d4d222 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 c7bf2d641..158959a2d 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 2a27ce959..b3aefb598 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 0b1991815..e851dcd5f 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 2d6fb9949..d318cf358 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 071faa463..ce38ecf97 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 36c33a108..ab16dc87e 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 edbce9f1e..95572c26f 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 e5d1c7c06..d93087b10 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 772936160..b2c411555 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 770088954..1d6f8ff68 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 e2ade030f..043a501cd 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 50118645b..e83f00773 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 780aab570..8ca89b8a6 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 a038002d9..cd0943586 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 b1e20fc5e..b1dc5a5b2 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 48f4ddcd4..ee78f4951 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 78ca228a3..6c127727a 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 f7b9164b0..b193b17ec 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 a6595fdff..7fed0a453 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 7d71d7b14..f16aae098 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 9e5ae9fa4..2af58be4a 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 3162bfa6b..dda557f6b 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 eefef580f..2fd15bd13 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 164c8f9ca..328658435 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 cc56a34c3..9a98f1fc4 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 85a1336e4..0d29c2a24 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 1c28eecac..e306ae556 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 10f006cdb..2e9276d97 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 cea08c4b1..cd7ad29b6 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 730dd7953..8b9cb34db 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 f04fd51ea..f84a5a9b7 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 b17725026..1af120e62 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 48e134687..81a5028dd 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 0fbe50f16..db9a4c85c 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 89e65b01b..21e16f6eb 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 bcd294077..3cb855c6f 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 d67abae68..f7d540da8 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.go b/pkg/resource-handler/controller/shard/reconcile_data_plane.go index 57ac03014..7bb1bd499 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 @@ -205,8 +211,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 +232,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") @@ -270,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, @@ -317,7 +336,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") } } @@ -331,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" @@ -343,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, @@ -374,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, @@ -436,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, @@ -490,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, @@ -498,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( @@ -518,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( @@ -641,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 28e2a9185..0ced1c72e 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 } @@ -114,6 +111,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 +124,7 @@ func postureTestPod() *corev1.Pod { metadata.LabelMultigresDatabase: "database", metadata.LabelMultigresTableGroup: "table-group", metadata.LabelMultigresShard: "0", + metadata.LabelAppComponent: PoolComponentName, }, }} } @@ -134,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{{ @@ -163,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 && @@ -191,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, @@ -201,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() }() @@ -223,30 +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) - } - if retryAfter != 0 { - t.Error("second inconsistent posture observation requested another debounce requeue") - } + 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. + 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", @@ -256,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, @@ -272,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", @@ -293,27 +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) - } - if retryAfter != 0 { - t.Error("second incomplete posture observation requested another debounce requeue") - } + 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. + 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() @@ -342,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 @@ -392,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( @@ -453,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) { @@ -464,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) { @@ -482,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, @@ -511,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 @@ -582,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( @@ -619,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, @@ -630,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, @@ -666,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) { @@ -713,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.go b/pkg/resource-handler/controller/shard/reconcile_deletion.go index 4a949499c..13dbaa6b9 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) { @@ -342,7 +346,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_deletion_internal_test.go b/pkg/resource-handler/controller/shard/reconcile_deletion_internal_test.go index 30c669994..e20fd3acd 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 d56eb92fa..8596a9538 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 2237b3686..ca6d63f3c 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.go b/pkg/resource-handler/controller/shard/reconcile_readiness.go index fe6886f2b..71234eae2 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/pkg/resource-handler/controller/shard/reconcile_readiness_test.go b/pkg/resource-handler/controller/shard/reconcile_readiness_test.go index c827da415..3518d6b71 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 69b7e5e73..bb04d5c94 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 new file mode 100644 index 000000000..fd284815c --- /dev/null +++ b/pkg/resource-handler/controller/shard/registration_requeue_test.go @@ -0,0 +1,819 @@ +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" + + "github.com/multigres/testkit/assert" +) + +// 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) + assert.NewAborting(t). + False(got < tc.min || got > tc.max, "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 { + 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) + } + } +} + +// 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 + } + assert.NewAborting(t).GreaterOrEqual(10, len(seen), "only") +} + +// 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) + } + + assert.NewAborting(t).Eq(0, len(r.postureStrikes), "a thousand shards seen and settled left") +} + +// 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() + c := assert.NewAborting(t) + + r := &ShardReconciler{} + s := shardNamed("ns", "shard-0") + + for range 5 { + r.recordPostureObservation(s, true) + } + 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. + 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++ { + 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. + 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 +// 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) + } + + assert.NewAborting(t).Eq(0, len(r.notConvergedSince), "a thousand shards seen and settled left") +} + +// 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() + c := assert.NewAborting(t) + + clk := &fakeClock{t: time.Unix(1_700_000_000, 0)} + r := &ShardReconciler{Clock: clk.now} + s := shardNamed("ns", "shard-0") + + c.Eq(0, r.recordNotConverged(s, true), "first not-converged observation reported elapsed") + clk.advance(37 * time.Second) + 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. + 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") + r.recordPostureObservation(s, true) + r.recordNotConverged(s, true) + + r.forgetStrikes(s.Namespace, s.Name) + + c.Eq(0, len(r.postureStrikes), "posture strikes") + c.Eq(0, len(r.notConvergedSince), "not-converged-since") +} + +// 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) { + ck := assert.NewCollecting(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()} + + _, 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) + } + _, ok := r.notConvergedSince[key] + ck.False(ok, "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) { + ck := assert.NewCollecting(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"}, + } + _, 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) + } + _, 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, +// 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} + 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)) + 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} + 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{ + 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) { + 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)) + c.Require().NoError(err, "BuildMultiorchDeployment() error =") + 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) + 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") + } + _, 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) + } +} + +// 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) { + 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() + 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) + 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) + 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, + ) + } +} + +// 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) { + 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"} + c.Require().NoError(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), "register pooler") + 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) + 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) + 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) + } + } +} + +// 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) { + 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 := 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) + 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) + 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) + } + } +} + +// 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) { + 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() + 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) + c.Require().NoError(err, "reconcilePosture() error =") + min, max := wantDelayRange(tc.elapsed) + c.False( + delay < min || delay > max, + "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) { + ck := 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() + 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) + ck.Require().NoError(err, "reconcilePosture() error =") + min, max := wantDelayRange(tc.elapsed) + 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{} + 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", + 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. + ck.Require().NoError(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), "promote pooler-0 in topology") + 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) + 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) + _, 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", + shard.Status.Conditions, + ) + } + + for _, pod := range []*corev1.Pod{pod0, pod1} { + got := &corev1.Pod{} + 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", + pod.Name, condition) + } + } +} diff --git a/pkg/resource-handler/controller/shard/reload_internal_test.go b/pkg/resource-handler/controller/shard/reload_internal_test.go index 1edbea829..fc3cd019a 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 2f4201b34..4acfeda62 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.go b/pkg/resource-handler/controller/shard/shard_controller.go index d1229ef4d..a36c2a7e3 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/pkg/resource-handler/controller/shard/shard_controller_internal_test.go b/pkg/resource-handler/controller/shard/shard_controller_internal_test.go index 344f3ce75..c774a331f 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 ea822b4aa..e5a8e31fa 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 76cd79ea9..2934df6b8 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 70137e361..3c5389eef 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 4e3f6d40e..eae910adc 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 f6918cff1..4ec6cabe3 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 9a50bc285..5d4615049 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 f9351a66a..8748d92f8 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 96fa332ce..5cc42ecaa 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 595b57f83..3c58d5123 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 3316d8e77..8d91674c2 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 424b4fc58..8e0079161 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 d959add02..cb317a342 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 83123b0eb..aef58fd1a 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 8db580df8..770aeb86f 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 5b186335c..934e5aa5e 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 76db5eaac..6a267609c 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 2b084b7b3..b3e424719 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 d14629c8e..e23906062 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 c5d0a41c0..1e553f1d8 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 47db4eceb..84ab3ca24 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 9df184472..4a236dd6c 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 3d046b1b2..f007149a1 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 71cd46cc8..fe5d81764 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 ae5239552..16aeac3aa 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 c668e0eb4..6daf7f4e2 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 a6fc84bb9..b55c22fd5 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 02fb4d2c2..2bd0e6c5b 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 2709c404a..2c1ff7905 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 04639ee1e..c995005d4 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 5fdb85e62..5a41c58db 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 51250061f..f8418de0e 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 5b3a92323..4158acf92 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 d718727a1..656e5f305 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 06eb9cd2d..aaf0a84f1 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 8082029bb..97f8f3300 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 e2b7d18fe..442674aa8 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 4ee69406b..7d8151759 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 17b9274ef..e38855eaa 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 97e3fbf3f..bde0e6c52 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 9b234b502..1b9392cff 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 9ecdf0cc5..c7820f160 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 481e347f5..55b72c01b 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 8cd2b9d6f..cde8e03be 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 48fc28c73..0ee6ad0cb 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 95fe442a0..6fbc2a916 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 ceb51032d..5956e62ce 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 1701d7619..e80f3d784 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 fe897059f..1dd3290f2 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 bf5aa4857..47c56d216 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 f16948302..db3383a3b 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 8a8034584..71d2b9c54 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 73083e37a..3aa8df6e1 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 712abd867..e8a082eab 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 3a8d415da..d59942fd7 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 a9e2803bd..b7498a57f 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 0c9cb6c4b..40f410539 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 d2d1a2553..af36c67e7 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 a4b6248b5..d70a4aebf 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 7fa3a4235..05eff42a0 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 309535f33..1d7face7d 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 be2dc04ab..ac424d224 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 f43a91076..1e7b40ba5 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 5406a92c0..8d71b56a0 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 171ca2b4e..6569705e9 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 d30739564..5c576a42f 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 5115dfe1f..a2267857d 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 3e92e6c93..64eaa04d2 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 57c491f89..4b9be527e 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 6a111d9c2..e1f46dba8 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 af06d1f63..cde5d9f51 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 4a1075fa6..e0b3fe7b0 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 bd417e181..82e313219 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 dc88232a7..fccb73c99 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 bc0fe7f6b..78e813a59 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 a3f80ce8e..b56558e6a 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. diff --git a/test/suite/case.go b/test/suite/case.go new file mode 100644 index 000000000..f43e50918 --- /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 000000000..6b6691ffe --- /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 000000000..440b017ad --- /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 000000000..f68746eda --- /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 000000000..5112c0c6d --- /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 000000000..3be6874b1 --- /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 000000000..214d750e7 --- /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 000000000..2d2a872dd --- /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 000000000..d05787639 --- /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 000000000..7b233ab8f --- /dev/null +++ b/test/suite/scenario_race_test.go @@ -0,0 +1,325 @@ +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) + conflicts := c.fieldOwnershipConflicts(key) + 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) { + 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 000000000..2ae0c28e2 --- /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 000000000..d494a34dd --- /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 000000000..dedc79ef4 --- /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 000000000..f4718deb9 --- /dev/null +++ b/test/suite/scenario_thrash_test.go @@ -0,0 +1,411 @@ +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 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) + 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) + } + + // 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. + // + // 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) + + // 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) +} + +// 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 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. 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) + 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 asserts. + release() + + 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 + // 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 fmt.Errorf("got %+v", members) + } + return nil + }, + ) +} + +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 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{} + 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 000000000..74a5986db --- /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 000000000..6cd6311de --- /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 000000000..428d78a3e --- /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 000000000..2b8c8667f --- /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 000000000..c0c3211a5 --- /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 000000000..622c0d019 --- /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 +)