diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index f55d84816..e274ac323 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -210,6 +210,22 @@ jobs: - name: Validate exact pure runtime test group run: timeout --signal=TERM --kill-after=10s 300s python3 hack/scripts/runtime-miri.py ${{ matrix.group }} + racer-envtest: + name: Racer API Server Integration Tests + runs-on: ubuntu-24.04 + timeout-minutes: 15 + steps: + - name: Checkout + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + + - name: Set up Go + uses: actions/setup-go@b7ad1dad31e06c5925ef5d2fc7ad053ef454303e # v7.0.0 + with: + go-version-file: go.mod + cache-dependency-path: go.sum + + - name: Provision envtest and run controller API-server race tests + run: timeout --signal=TERM --kill-after=10s 300s make racer-envtest-ci # ---------- Orca Integration Tests ---------- # Spins up Garage and Azurite via testcontainers-go and runs the diff --git a/.github/workflows/nightly.yaml b/.github/workflows/nightly.yaml index cbc6045cb..4f3a893c2 100644 --- a/.github/workflows/nightly.yaml +++ b/.github/workflows/nightly.yaml @@ -291,6 +291,8 @@ jobs: file: images/unbounded-operator/Containerfile - name: gantry file: images/gantry/Containerfile + - name: racer-controller + file: images/racer-controller/Containerfile - name: orca file: images/orca/Containerfile - name: inventory-aggregator diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index 1478dd7cc..f885bc821 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -853,6 +853,10 @@ jobs: file: images/gantry/Containerfile platforms: linux/amd64,linux/arm64 trivy: true + - name: racer-controller + file: images/racer-controller/Containerfile + platforms: linux/amd64,linux/arm64 + trivy: true - name: host-ubuntu2404 file: images/host-ubuntu2404/Containerfile platforms: linux/amd64,linux/arm64 diff --git a/Makefile b/Makefile index 52f46b88f..ed9bc478e 100644 --- a/Makefile +++ b/Makefile @@ -5,6 +5,9 @@ GOBUILD=$(GOCMD) build GOTEST=$(GOCMD) test GOMOD=$(GOCMD) mod GOLINT=golangci-lint run -c .golangci.yaml +ENVTEST_K8S_VERSION ?= 1.37.0 +SETUP_ENVTEST_VERSION ?= v0.25.2-0.20260923145615-d837464d41be +SETUP_ENVTEST = $(CURDIR)/bin/setup-envtest-$(SETUP_ENVTEST_VERSION) GO_PACKAGE_PATTERNS=./api/... ./cmd/... ./deploy/... ./e2e/... ./hack/... ./internal/... ./pkg/... # e2e packages hold nothing but files behind the e2e build tag, so `go list` # needs the tag to see them at all. Without it they are silently skipped by @@ -356,6 +359,14 @@ help: ## Show this help @echo "Documentation:" @echo " docs-serve Start local Hugo dev server" @echo "" + @echo "Racer Controller:" + @echo " racer-controller Test and build the Go controller" + @echo " racer-controller-build Build the Go controller without lint/test" + @echo " racer-test Lint and race-test the controller and deployment contracts" + @echo " racer-envtest Run controller API-server tests with KUBEBUILDER_ASSETS" + @echo " racer-envtest-ci Provision pinned assets and run controller API-server tests" + @echo " racer-generate Generate Racer deepcopy and CRD artifacts" + @echo "" @echo "Common variables (override with VAR=value):" @echo " VERSION=$(VERSION)" @echo " GIT_COMMIT=$(GIT_COMMIT)" @@ -541,6 +552,37 @@ e2e-gantry: $(HELM) ## Run the kind-based Gantry e2e suite e2e-playpen: ## Run the kind-based playpen e2e suite $(GOTEST) -tags=e2e ./e2e/playpen -v -timeout=10m +.PHONY: racer-controller racer-controller-build racer-test racer-server-test racer-envtest racer-envtest-ci racer-generate +racer-controller: racer-server-test racer-controller-build ## Test and build the Racer controller + +racer-controller-build: ## Build the Racer controller without lint/test + @mkdir -p bin + timeout --signal=TERM --kill-after=10s 300s $(GOBUILD) -trimpath -ldflags '$(STAMP_LDFLAGS)' -o bin/racer-controller ./cmd/racer-controller + +racer-server-test: ## Lint and race-test the Racer server + timeout --signal=TERM --kill-after=10s 300s $(GOLINT) ./api/racer/... ./internal/racer/... ./cmd/racer-controller/... + timeout --signal=TERM --kill-after=10s 300s $(GOTEST) -timeout=5m -race ./api/racer/... ./internal/racer/... ./cmd/racer-controller/... + +racer-test: racer-server-test ## Check the Racer controller + +$(SETUP_ENVTEST): + @mkdir -p bin tmp/envtest-tools + TMPDIR="$(CURDIR)/tmp/envtest-tools" GOBIN="$(CURDIR)/bin" timeout --signal=TERM --kill-after=10s 300s $(GOCMD) install sigs.k8s.io/controller-runtime/tools/setup-envtest@$(SETUP_ENVTEST_VERSION) + mv bin/setup-envtest "$(SETUP_ENVTEST)" + +racer-envtest-ci: $(SETUP_ENVTEST) ## Provision pinned local API-server assets and require Racer envtest + @mkdir -p tmp/racer-envtest + @assets=$$(TMPDIR="$(CURDIR)/tmp/racer-envtest" timeout --signal=TERM --kill-after=10s 300s "$(SETUP_ENVTEST)" use $(ENVTEST_K8S_VERSION) --bin-dir "$(CURDIR)/bin/envtest" -p path) && \ + $(MAKE) racer-envtest KUBEBUILDER_ASSETS="$$assets" + +racer-envtest: ## Run real API-server, manager election, TLS and crash-recovery tests + @test -n "$(KUBEBUILDER_ASSETS)" || { echo "Set KUBEBUILDER_ASSETS to repository-local envtest binaries"; exit 1; } + @mkdir -p tmp/racer-envtest + TMPDIR="$(CURDIR)/tmp/racer-envtest" KUBEBUILDER_ASSETS="$(KUBEBUILDER_ASSETS)" timeout --signal=TERM --kill-after=10s 300s $(GOTEST) -race ./internal/racer ./internal/racer/authority -run '^TestEnvtest' -count=1 -v -timeout=5m + +racer-generate: ## Generate Racer deepcopy and CRD artifacts + timeout --signal=TERM --kill-after=10s 300s $(GOCMD) generate ./api/racer/v1alpha1 + build: machina-manifests token-refresher-manifests machine-ops-manifests playpen-manifests net-manifests unbounded-operator-manifests gantry-manifests ## Build all Go packages $(GOBUILD) ./... diff --git a/NOTICE b/NOTICE index 4b42c7b9f..d1435539f 100644 --- a/NOTICE +++ b/NOTICE @@ -389,6 +389,13 @@ notices: license: - name: Apache License, Version 2.0 link: https://github.com/google/renameio/blob/v2.0.2/LICENSE + - dependency: github.com/google/uuid + ecosystem: go + copyright: + - Copyright (c) 2009,2014 Google Inc. All rights reserved. + license: + - name: BSD 3-Clause License + link: https://github.com/google/uuid/blob/v1.6.0/LICENSE - dependency: github.com/insomniacslk/dhcp ecosystem: go copyright: diff --git a/api/racer/v1alpha1/clustercache_test.go b/api/racer/v1alpha1/clustercache_test.go new file mode 100644 index 000000000..2facfce18 --- /dev/null +++ b/api/racer/v1alpha1/clustercache_test.go @@ -0,0 +1,43 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package v1alpha1 + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" +) + +func TestClusterCacheRegistrationAndJSON(t *testing.T) { + scheme := runtime.NewScheme() + require.NoError(t, AddToScheme(scheme)) + + for kind, expected := range map[string]runtime.Object{ + "ClusterCache": &ClusterCache{}, + "ClusterCacheList": &ClusterCacheList{}, + } { + object, err := scheme.New(GroupVersion.WithKind(kind)) + require.NoError(t, err) + require.IsType(t, expected, object) + } + + cache := &ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache", Labels: map[string]string{"test": "original"}}} + data, err := json.Marshal(cache) + require.NoError(t, err) + + var object map[string]any + require.NoError(t, json.Unmarshal(data, &object)) + require.NotContains(t, object, "spec") + + copy := cache.DeepCopy() + copy.Labels["test"] = "copy" + require.Equal(t, "original", cache.Labels["test"]) + list := &ClusterCacheList{Items: []ClusterCache{*cache}} + listCopy := list.DeepCopy() + listCopy.Items[0].Labels["test"] = "list-copy" + require.Equal(t, "original", list.Items[0].Labels["test"]) +} diff --git a/api/racer/v1alpha1/clustercache_types.go b/api/racer/v1alpha1/clustercache_types.go new file mode 100644 index 000000000..9cedeff94 --- /dev/null +++ b/api/racer/v1alpha1/clustercache_types.go @@ -0,0 +1,28 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package v1alpha1 + +import metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + +// ClusterCache names a cache. Its Kubernetes UID is its wire identity. +// Socket paths are derived from metadata.name. Client access is controlled by pod volume mounts. +// The name limit keeps /run/racer//origin/socket within Linux sockaddr_un. +// Each DNS label is limited to 63 characters to match the wire contract; +// Kubernetes metadata validation enforces DNS subdomain spelling. +// The bounded regex checks label lengths without a costly CEL split/all loop. +// +kubebuilder:object:root=true +// +kubebuilder:resource:scope=Cluster,shortName=ccache +// +kubebuilder:validation:XValidation:rule="size(self.metadata.name) <= 82",message="name must fit the canonical Unix socket path" +// +kubebuilder:validation:XValidation:rule="self.metadata.name.matches('^[^.]{1,63}([.][^.]{1,63})*$')",message="each name label must be at most 63 characters" +type ClusterCache struct { + metav1.TypeMeta `json:",inline"` + metav1.ObjectMeta `json:"metadata,omitempty"` +} + +// +kubebuilder:object:root=true +type ClusterCacheList struct { + metav1.TypeMeta `json:",inline"` + metav1.ListMeta `json:"metadata,omitempty"` + Items []ClusterCache `json:"items"` +} diff --git a/api/racer/v1alpha1/crd/racer.unbounded-cloud.io_clustercaches.yaml b/api/racer/v1alpha1/crd/racer.unbounded-cloud.io_clustercaches.yaml new file mode 100644 index 000000000..8c8b4d073 --- /dev/null +++ b/api/racer/v1alpha1/crd/racer.unbounded-cloud.io_clustercaches.yaml @@ -0,0 +1,56 @@ +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 +--- +apiVersion: apiextensions.k8s.io/v1 +kind: CustomResourceDefinition +metadata: + annotations: + controller-gen.kubebuilder.io/version: v0.20.1 + name: clustercaches.racer.unbounded-cloud.io +spec: + group: racer.unbounded-cloud.io + names: + kind: ClusterCache + listKind: ClusterCacheList + plural: clustercaches + shortNames: + - ccache + singular: clustercache + scope: Cluster + versions: + - name: v1alpha1 + schema: + openAPIV3Schema: + description: |- + ClusterCache names a cache. Its Kubernetes UID is its wire identity. + Socket paths are derived from metadata.name. Client access is controlled by pod volume mounts. + The name limit keeps /run/racer//origin/socket within Linux sockaddr_un. + Each DNS label is limited to 63 characters to match the wire contract; + Kubernetes metadata validation enforces DNS subdomain spelling. + The bounded regex checks label lengths without a costly CEL split/all loop. + properties: + apiVersion: + description: |- + APIVersion defines the versioned schema of this representation of an object. + Servers should convert recognized schemas to the latest internal value, and + may reject unrecognized values. + More info: https://git.k8s.io/community/contributors/devel/sig-architecture/api-conventions.md#resources + type: string + kind: + description: |- + Kind is a string value representing the REST resource this object represents. + Servers may infer this from the endpoint the client submits requests to. + Cannot be updated. + In CamelCase. + More info: https://git.k8s.io/community/contributors/devel/sig-architecture/api-conventions.md#types-kinds + type: string + metadata: + type: object + type: object + x-kubernetes-validations: + - message: name must fit the canonical Unix socket path + rule: size(self.metadata.name) <= 82 + - message: each name label must be at most 63 characters + rule: self.metadata.name.matches('^[^.]{1,63}([.][^.]{1,63})*$') + served: true + storage: true diff --git a/api/racer/v1alpha1/groupversion_info.go b/api/racer/v1alpha1/groupversion_info.go new file mode 100644 index 000000000..60363f8ef --- /dev/null +++ b/api/racer/v1alpha1/groupversion_info.go @@ -0,0 +1,9 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// +kubebuilder:object:generate=true +// +groupName=racer.unbounded-cloud.io +package v1alpha1 + +//go:generate go tool controller-gen object:headerFile=../../../hack/boilerplate.go.txt paths=. +//go:generate go tool controller-gen crd:headerFile=../../../hack/boilerplate.yaml.txt paths=. output:crd:dir=crd diff --git a/api/racer/v1alpha1/register.go b/api/racer/v1alpha1/register.go new file mode 100644 index 000000000..18988755b --- /dev/null +++ b/api/racer/v1alpha1/register.go @@ -0,0 +1,23 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package v1alpha1 + +import ( + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/runtime/schema" +) + +const GroupName = "racer.unbounded-cloud.io" + +var ( + GroupVersion = schema.GroupVersion{Group: GroupName, Version: "v1alpha1"} + SchemeBuilder = runtime.NewSchemeBuilder(func(s *runtime.Scheme) error { + s.AddKnownTypes(GroupVersion, &ClusterCache{}, &ClusterCacheList{}) + metav1.AddToGroupVersion(s, GroupVersion) + + return nil + }) + AddToScheme = SchemeBuilder.AddToScheme +) diff --git a/api/racer/v1alpha1/zz_generated.deepcopy.go b/api/racer/v1alpha1/zz_generated.deepcopy.go new file mode 100644 index 000000000..a2d217550 --- /dev/null +++ b/api/racer/v1alpha1/zz_generated.deepcopy.go @@ -0,0 +1,69 @@ +//go:build !ignore_autogenerated + +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// Code generated by controller-gen. DO NOT EDIT. + +package v1alpha1 + +import ( + "k8s.io/apimachinery/pkg/runtime" +) + +// DeepCopyInto is an autogenerated deepcopy function, copying the receiver, writing into out. in must be non-nil. +func (in *ClusterCache) DeepCopyInto(out *ClusterCache) { + *out = *in + out.TypeMeta = in.TypeMeta + in.ObjectMeta.DeepCopyInto(&out.ObjectMeta) +} + +// DeepCopy is an autogenerated deepcopy function, copying the receiver, creating a new ClusterCache. +func (in *ClusterCache) DeepCopy() *ClusterCache { + if in == nil { + return nil + } + out := new(ClusterCache) + in.DeepCopyInto(out) + return out +} + +// DeepCopyObject is an autogenerated deepcopy function, copying the receiver, creating a new runtime.Object. +func (in *ClusterCache) DeepCopyObject() runtime.Object { + if c := in.DeepCopy(); c != nil { + return c + } + return nil +} + +// DeepCopyInto is an autogenerated deepcopy function, copying the receiver, writing into out. in must be non-nil. +func (in *ClusterCacheList) DeepCopyInto(out *ClusterCacheList) { + *out = *in + out.TypeMeta = in.TypeMeta + in.ListMeta.DeepCopyInto(&out.ListMeta) + if in.Items != nil { + in, out := &in.Items, &out.Items + *out = make([]ClusterCache, len(*in)) + for i := range *in { + (*in)[i].DeepCopyInto(&(*out)[i]) + } + } +} + +// DeepCopy is an autogenerated deepcopy function, copying the receiver, creating a new ClusterCacheList. +func (in *ClusterCacheList) DeepCopy() *ClusterCacheList { + if in == nil { + return nil + } + out := new(ClusterCacheList) + in.DeepCopyInto(out) + return out +} + +// DeepCopyObject is an autogenerated deepcopy function, copying the receiver, creating a new runtime.Object. +func (in *ClusterCacheList) DeepCopyObject() runtime.Object { + if c := in.DeepCopy(); c != nil { + return c + } + return nil +} diff --git a/cmd/racer-controller/README.md b/cmd/racer-controller/README.md index 9d8663f05..ca5c3c217 100644 --- a/cmd/racer-controller/README.md +++ b/cmd/racer-controller/README.md @@ -1,6 +1,31 @@ # Racer Controller -Placeholder for the replacement Racer controller implementation. +The Racer controller keeps track of Racer nodes and named caches in Kubernetes. -The previous Racer implementation has been removed. The replacement is not -included in this branch. +## ClusterCache + +A `ClusterCache` gives a Racer cache a name. It applies to the whole cluster, +not a single namespace, and has no `spec` fields to configure. + +```yaml +apiVersion: racer.unbounded-cloud.io/v1alpha1 +kind: ClusterCache +metadata: + name: images +``` + +Use a lowercase DNS name, such as `images` or `team.images`. Names can be up to +82 characters long, with at most 63 characters in each dot-separated part. + +List caches with: + +```sh +kubectl get clustercaches +``` + +Applications access a cache through its client socket at +`/run/racer//client/socket`. Pod volume mounts control access; creating a +`ClusterCache` does not grant every Pod access to it. + +Deleting and recreating a `ClusterCache` creates a new cache identity, even if +you reuse the name. diff --git a/cmd/racer-controller/main.go b/cmd/racer-controller/main.go new file mode 100644 index 000000000..4e093621b --- /dev/null +++ b/cmd/racer-controller/main.go @@ -0,0 +1,41 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "fmt" + "os" + + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/log/zap" + + "github.com/Azure/unbounded/internal/racer" + "github.com/Azure/unbounded/internal/version" +) + +func main() { + ctrl.SetLogger(zap.New()) + + if err := run(); err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } +} + +func run() error { + if len(os.Args) != 1 { + return fmt.Errorf("usage: racer-controller") + } + + ctrl.Log.Info("starting racer-controller", "version", version.String()) + + cfg, err := racer.LoadConfig() + if err != nil { + return err + } + + ctx := ctrl.SetupSignalHandler() + + return racer.Run(ctx, cfg) +} diff --git a/cmd/racer-controller/main_test.go b/cmd/racer-controller/main_test.go new file mode 100644 index 000000000..1d2e977f8 --- /dev/null +++ b/cmd/racer-controller/main_test.go @@ -0,0 +1,46 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "bytes" + "os" + "testing" + + "github.com/stretchr/testify/require" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/log/zap" + + "github.com/Azure/unbounded/internal/version" +) + +func TestRejectInitializeCLI(t *testing.T) { + before := os.Args + + t.Cleanup(func() { os.Args = before }) + + for _, args := range [][]string{{"racer-controller", "initialize"}, {"racer-controller", "other"}, {"racer-controller", "initialize", "extra"}} { + os.Args = args + if err := run(); err == nil || err.Error() != "usage: racer-controller" { + t.Fatalf("args %v: %v", args, err) + } + } +} + +func TestStartupLogsBuildMetadataBeforeConfigurationFailure(t *testing.T) { + before := os.Args + + t.Cleanup(func() { os.Args = before }) + + os.Args = []string{"racer-controller"} + + t.Setenv("RACER_CLUSTER_ID", "invalid") + + var output bytes.Buffer + ctrl.SetLogger(zap.New(zap.WriteTo(&output))) + t.Cleanup(func() { ctrl.SetLogger(zap.New()) }) + require.Error(t, run()) + require.Contains(t, output.String(), "starting racer-controller") + require.Contains(t, output.String(), version.String()) +} diff --git a/go.mod b/go.mod index 535bd851e..5b16cceb5 100644 --- a/go.mod +++ b/go.mod @@ -54,6 +54,7 @@ require ( github.com/google/go-tpm-tools v0.4.10 github.com/google/licensecheck v0.3.1 github.com/google/renameio/v2 v2.0.2 + github.com/google/uuid v1.6.0 github.com/insomniacslk/dhcp v0.0.0-20260220084031-5adc3eb26f91 github.com/ipfs/go-cid v0.6.2 github.com/lib/pq v1.12.3 @@ -183,7 +184,6 @@ require ( github.com/google/go-cmp v0.7.0 // indirect github.com/google/go-querystring v1.1.0 // indirect github.com/google/gopacket v1.1.19 // indirect - github.com/google/uuid v1.6.0 // indirect github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674 // indirect github.com/hashicorp/golang-lru v1.0.2 // indirect github.com/huandu/xstrings v1.5.0 // indirect diff --git a/images/racer-controller/Containerfile b/images/racer-controller/Containerfile new file mode 100644 index 000000000..34ed49113 --- /dev/null +++ b/images/racer-controller/Containerfile @@ -0,0 +1,32 @@ +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +FROM --platform=$BUILDPLATFORM docker.io/library/golang:1.26.8-trixie AS builder +WORKDIR /src +COPY go.mod go.sum ./ +RUN go mod download +COPY . . +ARG TARGETOS +ARG TARGETARCH +ARG VERSION=dev +ARG GIT_COMMIT=unknown +ARG BUILD_TIME=unknown +RUN CGO_ENABLED=0 GOOS=${TARGETOS} GOARCH=${TARGETARCH} \ + go build -trimpath \ + -ldflags "-X github.com/Azure/unbounded/internal/version.Version=${VERSION} -X github.com/Azure/unbounded/internal/version.GitCommit=${GIT_COMMIT} -X github.com/Azure/unbounded/internal/version.BuildTime=${BUILD_TIME}" \ + -o /out/racer-controller ./cmd/racer-controller + +FROM docker.io/library/debian:trixie-slim +ARG VERSION=dev +ARG GIT_COMMIT=unknown +LABEL org.opencontainers.image.title="racer-controller" \ + org.opencontainers.image.licenses="Apache-2.0" \ + org.opencontainers.image.version="${VERSION}" \ + org.opencontainers.image.revision="${GIT_COMMIT}" +RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates \ + && rm -rf /var/lib/apt/lists/* +COPY --from=builder /out/racer-controller /usr/local/bin/racer-controller +COPY LICENSE NOTICE /usr/share/licenses/racer-controller/ +USER 65532:65532 +EXPOSE 8443 8080 8081 +ENTRYPOINT ["/usr/local/bin/racer-controller"] diff --git a/internal/racer/authority/admission_test.go b/internal/racer/authority/admission_test.go new file mode 100644 index 000000000..98a993bc7 --- /dev/null +++ b/internal/racer/authority/admission_test.go @@ -0,0 +1,175 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "context" + "io" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + + "github.com/Azure/unbounded/internal/racer/wire" +) + +type admissionBoundaryWriter struct { + calls int + after func() +} + +type admissionTestKey struct{} + +func (w *admissionBoundaryWriter) Write(b []byte) (int, error) { + w.calls++ + if w.calls == 1 { + w.after() + } + + return len(b), nil +} + +func TestAdmissionPublicationEpochAndOriginalDeadline(t *testing.T) { + for _, action := range []string{"advance", "suspend recover", "process cancel", "expiry"} { + t.Run(action, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + + process, stopProcess := context.WithCancel(t.Context()) + defer stopProcess() + // Start a process-bound epoch and use a large immutable response to + // exercise checks after the first chunk, before callback scheduling. + a.publications.Suspend() + a.BindProcess(process) + next := *a.publications.current + next.encoded = strings.Repeat("x", 64*1024) + next.record.Sequence++ + require.NoError(t, a.publications.Install(&next)) + handle, err := a.Current() + require.NoError(t, err) + guard, stop, err := handle.Admit(t.Context()) + require.NoError(t, err) + + defer stop() + + deadline, _ := guard.Context().Deadline() + // Ordinary derived contexts must not hide admission or revocation. + ctx, cancel := context.WithCancel(context.WithValue(guard.Context(), admissionTestKey{}, "value")) + defer cancel() + + response := handle.ForBase(0, "") + writer := &admissionBoundaryWriter{after: func() { + advanced := *a.publications.current + advanced.record.Sequence++ + require.NoError(t, a.publications.Install(&advanced)) + _, _, err := handle.Admit(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + + switch action { + case "suspend recover": + a.publications.Suspend() + require.NoError(t, a.publications.Install(&advanced)) + case "process cancel": + stopProcess() + case "expiry": + time.Sleep(time.Until(deadline)) + } + }} + + n, err := response.WriteTo(ctx, guard, writer) + if action == "advance" { + require.NoError(t, err) + require.EqualValues(t, 64*1024, n) + require.Equal(t, 2, writer.calls) + require.NoError(t, guard.Check(ctx)) + time.Sleep(time.Until(deadline)) + require.ErrorIs(t, guard.Check(ctx), context.DeadlineExceeded) + } else { + require.Error(t, err) + require.EqualValues(t, 32*1024, n) + require.Equal(t, 1, writer.calls) + } + + got, _ := guard.Context().Deadline() + require.Equal(t, deadline, got) + // Check is synchronous; cancellation notification may arrive later. + synctest.Wait() + require.Error(t, guard.Context().Err()) + + select { + case <-guard.Context().Done(): + default: + t.Fatal("context Err returned before Done closed") + } + }) + }) + } +} + +func TestAdmissionRejectsMissingWrongImageAndExpiredTrust(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + old, err := a.Current() + require.NoError(t, err) + guard, stop, err := old.Admit(t.Context()) + require.NoError(t, err) + + defer stop() + + for _, missing := range []*Admission{nil, {}} { + _, err = old.ForBase(0, "").WriteTo(t.Context(), missing, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + } + + next := *a.publications.current + next.record.Sequence++ + require.NoError(t, a.publications.Install(&next)) + current, err := a.Current() + require.NoError(t, err) + _, err = current.ForBase(0, "").WriteTo(t.Context(), guard, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + trust, stopTrust, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stopTrust() + + combined, stopCombined, err := current.AdmitWithTrust(t.Context(), trust) + require.NoError(t, err) + + defer stopCombined() + + trustDeadline, _ := trust.Context().Deadline() + combinedDeadline, _ := combined.Context().Deadline() + require.False(t, combinedDeadline.After(trustDeadline)) + a.trust.invalidate() + require.ErrorIs(t, combined.Check(t.Context()), context.Canceled) + _, _, err = current.AdmitWithTrust(t.Context(), trust) + require.ErrorIs(t, err, context.Canceled) + }) +} + +func TestTrustAdmissionProcessCancellation(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + + process, cancel := context.WithCancel(t.Context()) + defer cancel() + + a.BindProcess(process) + guard, stop, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stop() + + cancel() + require.ErrorIs(t, guard.Check(t.Context()), context.Canceled) + _, _, err = a.AdmitTrust(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + <-guard.Context().Done() + require.ErrorIs(t, guard.Context().Err(), context.Canceled) +} diff --git a/internal/racer/authority/authority.go b/internal/racer/authority/authority.go new file mode 100644 index 000000000..46265da3b --- /dev/null +++ b/internal/racer/authority/authority.go @@ -0,0 +1,1910 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "bytes" + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "math/big" + "net/http" + "net/url" + "regexp" + "slices" + "strings" + "sync" + "time" + + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/util/validation" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +// Authority owns durable credential and publication operations and their shared +// admission gate. No operation exposes installation proofs or issuer material. +// Its stores, installation proofs, signing material, and gate are private. +type Authority struct { + config Config + reader client.Reader + client client.Writer + gate *catalogGate + publications *publicationStore + trust *trustStore + publisher *publisher + credentials *credentials + bootstrap *bootstrap + accepted AcceptedMembers +} + +// New composes only: no I/O, cryptography, or goroutines. Config is copied. +// Dependencies must provide authoritative reads independently of discovery caches. +func New(cfg Config, deps Dependencies) *Authority { + cfg = cfg.effective() + c, reader := deps.Writer, deps.Reader + a := &Authority{config: cfg, reader: reader, client: c, gate: newCatalogGate(), publications: newPublications(), trust: &trustStore{maxAge: cfg.SnapshotMaxAge}} + a.publications.maxAge = cfg.SnapshotMaxAge + a.accepted = make(AcceptedMembers) + a.publisher = &publisher{Writer: c, APIReader: reader, Config: cfg, Publications: a.publications, Trust: a.trust} + a.credentials = &credentials{Writer: c, APIReader: reader, Config: cfg, Trust: a.trust, Now: deps.Now} + a.bootstrap = &bootstrap{Client: c, APIReader: reader, Config: cfg, Issuer: &issuer{APIReader: reader, Config: cfg, Trust: a.trust, CatalogGate: a.gate, Now: deps.Now}} + a.bootstrap.owner = a + + return a +} + +// Recover validates permanent installation state without granting serving rights. +func (a *Authority) Recover(ctx context.Context, writer client.Writer) error { + ctx, cancel := context.WithTimeout(ctx, 30*time.Second) + defer cancel() + + if err := a.config.Validate(); err != nil { + return err + } + + if err := a.gate.Acquire(ctx); err != nil { + return err + } + defer a.gate.Release() + + if err := ensureInstalled(ctx, writer, a.reader, a.config); err != nil { + return fmt.Errorf("ensure Racer installation: %w", err) + } + + if _, _, err := readVersion(ctx, a.reader, a.config); err != nil { + return fmt.Errorf("recover Racer installation: %w", err) + } + + return nil +} + +// PublicationHandle is a response-only view. It cannot be installed or advanced. +type PublicationHandle struct { + image *committedPublication + owner *Authority +} + +func (p *PublicationHandle) Sequence() wire.Sequence { + if p == nil || p.image == nil { + return 0 + } + + return p.image.record.Sequence +} + +func (p *PublicationHandle) ForBase(sequence wire.Sequence, hash string) Response { + if p == nil || p.image == nil || p.owner == nil { + return Response{} + } + + return Response{response: p.image.ForBase(sequence, hash), owner: p.owner, image: p.image} +} + +func (p *PublicationHandle) Admit(ctx context.Context) (*Admission, context.CancelFunc, error) { + if p == nil || p.owner == nil || p.image == nil { + return nil, nil, wire.Unavailable + } + + guard, cancel, err := p.image.admit(ctx) + if err != nil { + return nil, nil, err + } + + guard.owner = p.owner + + return guard, cancel, nil +} + +// AdmitWithTrust binds publication delivery to the prior trust admission. +func (p *PublicationHandle) AdmitWithTrust(window context.Context, trust *Admission) (*Admission, context.CancelFunc, error) { + if trust == nil || p == nil || trust.owner == nil || trust.owner != p.owner || !trust.trust { + return nil, nil, wire.Forbidden + } + + if err := trust.Check(window); err != nil { + return nil, nil, err + } + + deadline, _ := trust.Context().Deadline() + window, stopWindow := context.WithDeadline(window, deadline) + + guard, cancel, err := p.Admit(window) + if err != nil { + stopWindow() + return nil, nil, err + } + + guard.parent = trust + stop := context.AfterFunc(trust.Context(), cancel) + + return guard, func() { stop(); cancel(); stopWindow() }, nil +} + +func (a *Authority) TrustReady() error { _, err := a.trust.pool(); return err } +func (a *Authority) PublicationReady() error { _, err := a.publications.Current(); return err } +func (a *Authority) BindProcess(ctx context.Context) { + a.publications.bindProcess(ctx) + a.trust.mu.Lock() + defer a.trust.mu.Unlock() + + a.trust.process = ctx +} + +func (a *Authority) Current() (*PublicationHandle, error) { + p, err := a.publications.Current() + if err != nil { + return nil, err + } + + return &PublicationHandle{image: p, owner: a}, nil +} + +func (a *Authority) CurrentAndSubscribe() (*PublicationHandle, <-chan struct{}, error) { + p, changed, err := a.publications.CurrentAndSubscribe() + if err != nil { + return nil, changed, err + } + + return &PublicationHandle{image: p, owner: a}, changed, nil +} + +func (a *Authority) Wait(ctx context.Context, identity NodeIdentity, after *wire.Sequence) (*PublicationHandle, error) { + if identity.owner != a || a == nil { + return nil, wire.Unauthenticated + } + + p, err := a.publications.Wait(ctx, identity, after) + if err != nil || p == nil { + return nil, err + } + + return &PublicationHandle{image: p, owner: a}, nil +} + +func (a *Authority) AdmitTrust(ctx context.Context) (*Admission, context.CancelFunc, error) { + a.trust.mu.RLock() + defer a.trust.mu.RUnlock() + + guard, cancel, err := a.trust.admitLocked(ctx) + if err != nil { + return nil, nil, err + } + + guard.owner = a + + return guard, cancel, nil +} + +// TrustPool returns an independent pool: TLS callers cannot mutate local trust. +func (a *Authority) TrustPool() (*x509.CertPool, error) { + p, err := a.trust.pool() + if err != nil { + return nil, err + } + + return p.Clone(), nil +} + +func (a *Authority) AuthenticateCertificate(ctx context.Context, state *tls.ConnectionState) (NodeIdentity, error) { + identity, err := authenticateCertificate(ctx, a.trust, a.config, state) + if err != nil { + return NodeIdentity{}, err + } + + identity.owner = a + + return identity, nil +} + +func (a *Authority) Authenticate(ctx context.Context, request *http.Request) (NodeIdentity, error) { + return a.bootstrap.Authenticate(ctx, request) +} + +func (a *Authority) Issue(ctx context.Context, identity NodeIdentity, request wire.BootstrapRequest) ([]byte, error) { + if identity.owner != a || a == nil || !identity.bearer { + return nil, wire.Unauthenticated + } + + return a.bootstrap.Issuer.Issue(ctx, identity, request) +} + +func (a *Authority) EnrollWithHint(ctx context.Context, request *http.Request, body wire.BootstrapRequest) ([]byte, EnrollmentHint, error) { + return a.bootstrap.enroll(ctx, request, body) +} + +// KeyringHandle is a comparable opaque view, suitable for detecting replacement +// across authentication without exposing the accepted encoding or install API. +type KeyringHandle struct { + image *acceptedKeyring + owner *Authority + epoch context.Context +} + +func (k KeyringHandle) Generation() wire.Generation { + if k.image == nil { + return 0 + } + + return k.image.generation +} + +func (k KeyringHandle) Response() Response { + if k.image == nil || k.owner == nil { + return Response{} + } + + return Response{response: publicationResponse{encoded: k.image.encoded}, owner: k.owner, trust: true, epoch: k.epoch, bundle: k.image} +} + +// Response exposes bounded writing, never an installation proof or mutable bytes. +type Response struct { + response publicationResponse + owner *Authority + image *committedPublication + trust bool + epoch context.Context + bundle *acceptedKeyring +} + +// Admission is an opaque serving capability, separate from request cancellation. +// Context returns an ordinary context for transport cleanup; Check synchronously +// checks revocation and the original deadline before and after writes and flushes. +type Admission struct { + ctx context.Context + owner *Authority + image *committedPublication + trust bool + epoch context.Context + bundle *acceptedKeyring + parent *Admission + process context.Context +} + +func (g *Admission) Context() context.Context { return g.ctx } + +func (g *Admission) Check(ctx context.Context) error { + if g == nil || g.ctx == nil || g.epoch == nil { + return wire.Forbidden + } + + if err := g.epoch.Err(); err != nil { + return err + } + + if g.process != nil { + if err := g.process.Err(); err != nil { + return err + } + } + + if deadline, ok := g.ctx.Deadline(); ok && !time.Now().Before(deadline) { + return context.DeadlineExceeded + } + + if g.parent != nil { + if err := g.parent.Check(ctx); err != nil { + return err + } + } + + if err := g.ctx.Err(); err != nil { + return err + } + + if deadline, ok := ctx.Deadline(); ok && !time.Now().Before(deadline) { + return context.DeadlineExceeded + } + + return ctx.Err() +} + +func newAdmission(parent, epoch context.Context, expiry time.Time) (*Admission, context.CancelFunc) { + ctx, cancel := context.WithDeadline(parent, expiry) + stop := context.AfterFunc(epoch, cancel) + + return &Admission{ctx: ctx, epoch: epoch}, func() { stop(); cancel() } +} + +func (Response) String() string { return "" } +func (Response) GoString() string { return "" } + +func (r Response) WriteTo(ctx context.Context, guard *Admission, w io.Writer) (int64, error) { + if guard == nil || r.owner == nil || guard.owner != r.owner || r.trust && (!guard.trust || r.epoch == nil || r.epoch != guard.epoch || r.bundle == nil || r.bundle != guard.bundle) || r.image != nil && guard.image != r.image { + return 0, wire.Forbidden + } + + if err := guard.Check(ctx); err != nil { + return 0, err + } + + n, err := r.response.writeTo(ctx, admissionWriter{ctx: ctx, guard: guard, writer: w}) + if err == nil { + err = guard.Check(ctx) + } + + return n, err +} + +type admissionWriter struct { + ctx context.Context + guard *Admission + writer io.Writer +} + +func (w admissionWriter) Write(b []byte) (int, error) { + if err := w.guard.Check(w.ctx); err != nil { + return 0, err + } + + n, err := w.writer.Write(b) + if err == nil { + err = w.guard.Check(w.ctx) + } + + return n, err +} + +func (a *Authority) WaitKeyring(ctx context.Context, after *wire.Generation) (*KeyringHandle, error) { + k, err := a.trust.waitKeyring(ctx, after) + if err != nil || k == nil { + return nil, err + } + + a.trust.mu.RLock() + defer a.trust.mu.RUnlock() + + if k != a.trust.bundle { + return nil, wire.Unavailable + } + + return &KeyringHandle{image: k, owner: a, epoch: a.trust.authority}, nil +} + +func (a *Authority) Keyring() (KeyringHandle, error) { + a.trust.mu.RLock() + defer a.trust.mu.RUnlock() + + if a.trust.bundle == nil || a.trust.authority == nil || a.trust.authority.Err() != nil || time.Since(a.trust.confirmed) >= a.trust.maxAge { + return KeyringHandle{}, wire.Unavailable + } + + return KeyringHandle{image: a.trust.bundle, owner: a, epoch: a.trust.authority}, nil +} + +// Config is copied by New. It contains authority policy, not listener, manager, +// filesystem, network replication, or HTTP admission settings. +type Config struct { + Cluster wire.ClusterID + Namespace string + DataplaneServiceAccount string + ControllerServiceAccount string + DaemonSetName string + CredentialsSecretName string + VersionConfigMapName string + InstallationConfigMapName string + Rotation RotationPolicy + CertificateLifetime time.Duration + SnapshotMaxAge time.Duration + MaxTokenBytes int +} + +// Dependencies are captured at construction. Reader must bypass informer caches. +// Writer needs only ordinary Kubernetes writes, including TokenReview creation. +type Dependencies struct { + Reader client.Reader + Writer client.Writer + Now func() time.Time +} + +func (c Config) effective() Config { + if c.CertificateLifetime == 0 { + c.CertificateLifetime = wire.CertificateLifetime + } + + if c.SnapshotMaxAge == 0 { + c.SnapshotMaxAge = 30 * time.Second + } + + return c +} + +func (c Config) Validate() error { + c = c.effective() + if !wire.ValidUUID(string(c.Cluster)) || len(validation.IsDNS1123Label(c.Namespace)) != 0 || c.SnapshotMaxAge < time.Second { + return wire.InvalidRequest + } + + for _, name := range []string{c.VersionConfigMapName, c.InstallationConfigMapName, c.DaemonSetName, c.CredentialsSecretName, c.DataplaneServiceAccount} { + if len(validation.IsDNS1123Subdomain(name)) != 0 { + return fmt.Errorf("resource name: %w", wire.InvalidRequest) + } + } + + if c.VersionConfigMapName == c.InstallationConfigMapName { + return wire.InvalidRequest + } + + lifetime := c.CertificateLifetime + if lifetime < 2*time.Minute || lifetime > wire.CertificateLifetime || lifetime%time.Second != 0 { + return wire.InvalidRequest + } + + if c.Rotation.PrepareFor <= 0 || c.Rotation.Interval < c.Rotation.PrepareFor || c.Rotation.RetainFor < lifetime || c.Rotation.Interval > 365*24*time.Hour || c.Rotation.RetainFor > 365*24*time.Hour { + return wire.InvalidRequest + } + + return nil +} + +const ReplicationAudience = "racer-controller-replication" + +func (a *Authority) Observe(ctx context.Context) error { + if err := a.gate.Acquire(ctx); err != nil { + return err + } + defer a.gate.Release() + + state, err := loadSigning(ctx, a.reader, a.config, time.Now()) + if err == nil { + err = a.trust.install(ctx, state.roots, state.bundle) + } + + if err == nil { + _, record, readErr := readVersion(ctx, a.reader, a.config) + + err = readErr + if err == nil { + err = a.publications.confirm(record) + } + } + + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return err + } + + if shouldInvalidateTrust(err) { + a.trust.invalidate() + a.publications.Suspend() + } + + return err +} + +func (a *Authority) AcceptReplica(ctx, process context.Context, image wire.Publication) error { + encoded, err := wire.EncodePublication(image) + if err != nil { + return err + } + + content, membership, err := wire.ContentHashes(image) + if err != nil { + return err + } + + want := versionRecord{Cluster: image.Cluster, Sequence: image.Sequence, MembershipVersion: image.MembershipVersion, ContentHash: content, MembershipHash: membership} + + if err := a.gate.Acquire(ctx); err != nil { + return err + } + defer a.gate.Release() + + _, record, err := readVersion(ctx, a.reader, a.config) + if err != nil { + if shouldInvalidateTrust(err) { + a.publications.Suspend() + a.trust.invalidate() + } + + return err + } + + if err := a.publications.confirm(record); err != nil { + a.publications.Suspend() + return err + } + + if record != want { + return wire.Unavailable + } + + return a.publications.Install(&committedPublication{owner: a.publications, record: record, encoded: string(encoded), leadership: process}) +} + +type ReplicaIdentity struct { + owner *Authority + uid string + expires time.Time +} + +func (i ReplicaIdentity) UID() string { return i.uid } +func (i ReplicaIdentity) Expires() time.Time { return i.expires } + +func (a *Authority) AuthenticateReplica(ctx context.Context, request *http.Request) (ReplicaIdentity, error) { + status, token, err := reviewBearer(ctx, a.client, request, ReplicationAudience, 0) + if err != nil { + return ReplicaIdentity{}, err + } + + if status.User.Username != "system:serviceaccount:"+a.config.Namespace+":"+a.config.ControllerServiceAccount { + return ReplicaIdentity{}, wire.Forbidden + } + + name, uid := singleExtra(status.User, "pod-name"), singleExtra(status.User, "pod-uid") + if name == "" || uid == "" || status.User.UID == "" { + return ReplicaIdentity{}, wire.Unauthenticated + } + + var pod corev1.Pod + if err := a.reader.Get(ctx, client.ObjectKey{Namespace: a.config.Namespace, Name: name}, &pod); err != nil { + return ReplicaIdentity{}, authorizationError(err) + } + + if !controllerPod(a.config, &pod) || string(pod.UID) != uid { + return ReplicaIdentity{}, wire.Forbidden + } + + var sa corev1.ServiceAccount + if err := a.reader.Get(ctx, client.ObjectKey{Namespace: a.config.Namespace, Name: a.config.ControllerServiceAccount}, &sa); err != nil { + return ReplicaIdentity{}, authorizationError(err) + } + + if string(sa.UID) != status.User.UID || sa.DeletionTimestamp != nil { + return ReplicaIdentity{}, wire.Forbidden + } + + expires, err := tokenExpiration(token) + if err != nil { + return ReplicaIdentity{}, err + } + + return ReplicaIdentity{owner: a, uid: uid, expires: expires}, nil +} + +func controllerPod(cfg Config, pod *corev1.Pod) bool { + return members.ControllerPod(pod, cfg.Namespace, cfg.ControllerServiceAccount) +} + +type publisher struct { + client.Writer + APIReader client.Reader + Config Config + Publications *publicationStore + Trust *trustStore +} + +func (r *publisher) suspendInvalidAuthority(err error) { + if shouldInvalidateTrust(err) { + r.Publications.Suspend() + r.Trust.invalidate() + } +} + +type ( + AcceptedMembers = members.History + TopologyHints struct { + Nodes corev1.NodeList + Members members.History + } + TopologyObservation struct { + Nodes corev1.NodeList + Input members.Input + Catalog []wire.CacheDefinition + } +) + +// PublishTopology performs discovery under private admission, then durable CAS and +// installation. Only this successful operation advances publisher history. +func (a *Authority) PublishTopology(ctx context.Context, observe func(context.Context) (TopologyObservation, error)) (TopologyHints, error) { + r := a.publisher + cfg := r.Config + + if err := a.gate.Acquire(ctx); err != nil { + return TopologyHints{}, err + } + defer a.gate.Release() + + if err := ctx.Err(); err != nil { + return TopologyHints{}, err + } + + cm, previous, err := readVersion(ctx, r.APIReader, cfg) + if err != nil { + r.suspendInvalidAuthority(err) + return TopologyHints{}, err + } + + if err := r.Publications.confirm(previous); err != nil { + return TopologyHints{}, err + } + + if observe == nil { + return TopologyHints{}, wire.InvalidRequest + } + + observation, err := observe(ctx) + if err != nil { + return TopologyHints{}, err + } + + catalog, err := r.keyedCatalog(ctx, cm, observation.Catalog) + if err != nil { + return TopologyHints{}, err + } + + result, err := members.Reconcile(observation.Input, a.accepted) + if err != nil { + return TopologyHints{}, err + } + + for _, d := range result.Diagnostics { + ctrl.LoggerFrom(ctx).Info("membership input rejected", "object", d.Object, "field", d.Field, "reason", d.Reason) + } + + prepared, err := r.Publications.Prepare(previous, cm.ResourceVersion, result.Members, catalog) + if err != nil { + return TopologyHints{}, err + } + + committed, err := r.CommitVersion(ctx, prepared) + if err != nil { + return TopologyHints{}, err + } + + if err := ctx.Err(); err != nil { + return TopologyHints{}, err + } + + if err := r.Publications.Install(committed); err != nil { + return TopologyHints{}, err + } + + a.accepted = result.Members + + return TopologyHints{Nodes: *observation.Nodes.DeepCopy(), Members: cloneAccepted(result.Members)}, nil +} + +func (r *publisher) keyedCatalog(ctx context.Context, version *corev1.ConfigMap, catalog []wire.CacheDefinition) ([]wire.CacheDefinition, error) { + claim := version.Annotations[credentialClaim] + if claim == "" { + return nil, nil + } + + credentials, err := readBoundCredentials(ctx, r.APIReader, r.Config, claim, version) + if err == nil { + err = r.Trust.validateReplay(credentials.bundle) + } + + if err != nil { + r.suspendInvalidAuthority(err) + return nil, err + } + + keyed := keyedCaches(credentials.bundle) + + accepted := make([]wire.CacheDefinition, 0, len(catalog)) + for _, cache := range catalog { + if keyed[cache.ID] { + accepted = append(accepted, cache) + } + } + + return accepted, nil +} + +func cloneAccepted(accepted AcceptedMembers) AcceptedMembers { + copy := make(AcceptedMembers, len(accepted)) + for id, member := range accepted { + member.RDMANICs = slices.Clone(member.RDMANICs) + for i := range member.RDMANICs { + if numa := member.RDMANICs[i].NUMANode; numa != nil { + copy := *numa + member.RDMANICs[i].NUMANode = © + } + } + + copy[id] = member + } + + return copy +} + +const ( + enrolledSharesAnnotation = members.EnrolledSharesAnnotation + enrolledRDMANICsAnnotation = members.EnrolledRDMANICsAnnotation + admittedMemberAnnotation = members.AdmittedMemberAnnotation +) + +// NodeIdentity is verified output, never populated from an untrusted request. +type NodeIdentity struct { + owner *Authority + bearer bool + nodeName string + cluster wire.ClusterID + node wire.NodeID + expires time.Time +} + +func (i NodeIdentity) Node() wire.NodeID { return i.node } +func (i NodeIdentity) Cluster() wire.ClusterID { return i.cluster } +func (i NodeIdentity) Expires() time.Time { return i.expires } + +// issuer accesses a controller-only Secret. Its private key is never projected +// into dataplane Pods or included in a response or diagnostic. +type issuer struct { + APIReader client.Reader + Config Config + Trust *trustStore + CatalogGate *catalogGate + Now func() time.Time +} + +type signingMaterial struct { + Certificate []byte `json:"certificate"` + PrivateKey []byte `json:"private_key"` +} + +type issuerMaterial struct { + Keys map[string]signingMaterial `json:"keys"` +} + +func (signingMaterial) String() string { return "" } +func (signingMaterial) GoString() string { return "" } +func (issuerMaterial) String() string { return "" } +func (issuerMaterial) GoString() string { return "" } + +func serialNumber() (*big.Int, error) { + n, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128)) + if err != nil { + return nil, err + } + + return n.Add(n, big.NewInt(1)), nil +} + +// Backdate validity starts for modest clock skew without extending expiration. +const certificateClockSkew = time.Minute + +func generateIssuer(now time.Time, cfg Config) ([]byte, []byte, error) { + cfg = cfg.effective() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + + serial, err := serialNumber() + if err != nil { + return nil, nil, err + } + + template := &x509.Certificate{SerialNumber: serial, Subject: pkix.Name{CommonName: "Racer " + string(cfg.Cluster)}, NotBefore: now.Add(-certificateClockSkew), NotAfter: now.Add(cfg.Rotation.Interval + cfg.Rotation.PrepareFor + cfg.Rotation.RetainFor + 2*cfg.CertificateLifetime), IsCA: true, BasicConstraintsValid: true, MaxPathLenZero: true, KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign} + + cert, err := x509.CreateCertificate(rand.Reader, template, template, pub, key) + if err != nil { + return nil, nil, err + } + + if len(cert) > reservedRootBytes { + return nil, nil, wire.TooLarge + } + + encoded, err := x509.MarshalPKCS8PrivateKey(key) + + return cert, encoded, err +} + +func parseSigning(m signingMaterial) (*x509.Certificate, ed25519.PrivateKey, error) { + cert, err := x509.ParseCertificate(m.Certificate) + if err != nil { + return nil, nil, wire.Unavailable + } + + private, err := x509.ParsePKCS8PrivateKey(m.PrivateKey) + if err != nil { + return nil, nil, wire.Unavailable + } + + key, ok := private.(ed25519.PrivateKey) + + pub, publicOK := cert.PublicKey.(ed25519.PublicKey) + if !ok || !publicOK || !pub.Equal(key.Public()) || !cert.IsCA || !cert.BasicConstraintsValid || cert.KeyUsage&x509.KeyUsageCertSign == 0 || cert.CheckSignatureFrom(cert) != nil { + return nil, nil, wire.Unavailable + } + + return cert, key, nil +} + +type parsedSigning struct { + certificate *x509.Certificate + key ed25519.PrivateKey +} + +type signingState struct { + certificate *x509.Certificate + key ed25519.PrivateKey + roots *x509.CertPool + bundle wire.KeyringBundle +} + +func loadSigning(ctx context.Context, reader client.Reader, cfg Config, now time.Time) (signingState, error) { + cfg = cfg.effective() + + if err := ctx.Err(); err != nil { + return signingState{}, err + } + + if err := cfg.Validate(); err != nil { + return signingState{}, err + } + + version, _, err := readVersion(ctx, reader, cfg) + if err != nil { + return signingState{}, err + } + + claim := version.Annotations[credentialClaim] + if !validCredentialClaim(cfg, claim) { + return signingState{}, wire.Unavailable + } + + credentials, err := readBoundCredentials(ctx, reader, cfg, claim, version) + if err != nil { + return signingState{}, err + } + + active := credentials.signing[credentials.rotation.ActiveIssuer] + + cert, key := active.certificate, active.key + if now.Before(cert.NotBefore) || now.Add(cfg.CertificateLifetime).After(cert.NotAfter) { + return signingState{}, wire.Unavailable + } + + roots := x509.NewCertPool() + + for _, der := range credentials.bundle.PeerTrustRoots { + root := credentials.signing[rootID(der)].certificate + if !now.Before(root.NotBefore) && now.Before(root.NotAfter) { + roots.AddCert(root) + } + } + + if err := ctx.Err(); err != nil { + return signingState{}, err + } + + return signingState{certificate: cert, key: key, roots: roots, bundle: credentials.bundle}, nil +} + +// Issuance can also observe invalid durable authority. It may withdraw trust, +// but only controller reconciliation can install or restore serving trust. +func (i *issuer) loadSigning(ctx context.Context, now time.Time) (signingState, error) { + // Serialize observations with controller installation so an in-flight valid + // read cannot restore trust after another operation observes invalidity. + if err := i.CatalogGate.Acquire(ctx); err != nil { + return signingState{}, err + } + defer i.CatalogGate.Release() + + state, err := loadSigning(ctx, i.APIReader, i.Config, now) + if err == nil { + err = i.Trust.validateReplay(state.bundle) + } + + // Issuance only reads authority. A request ending supplies no invalid evidence + // and cannot make a write ambiguous. Inspect the returned error, not ctx.Err: + // validation failures observed alongside cancellation must still revoke trust. + if shouldInvalidateTrust(err) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) { + i.Trust.invalidate() + } + + return state, err +} + +// Issue accepts only the identity returned by token authentication. CSR names, +// extensions and requested usages are discarded. Enrollment is correlation only. +// It returns an owned, validated JSON response within the bootstrap wire bound. +func (i *issuer) Issue(ctx context.Context, identity NodeIdentity, request wire.BootstrapRequest) ([]byte, error) { + cfg := i.Config + + if err := ctx.Err(); err != nil { + return nil, err + } + + now := credentialTime(i.Now) + if identity.cluster != cfg.Cluster || !wire.ValidUUID(string(identity.node)) || !identity.expires.After(now) { + return nil, wire.Forbidden + } + + if request.Cluster != identity.cluster { + return nil, wire.Forbidden + } + + if err := wire.ValidateBootstrapRequest(request); err != nil { + return nil, err + } + + pub, err := bootstrapPublicKey(request.CSRDER) + if err != nil { + return nil, err + } + + state, err := i.loadSigning(ctx, now) + if err != nil { + return nil, err + } + + serial, err := serialNumber() + if err != nil { + return nil, err + } + + uri := &url.URL{Scheme: "spiffe", Host: string(identity.cluster), Path: "/node/" + string(identity.node)} + template := &x509.Certificate{SerialNumber: serial, NotBefore: now.Add(-certificateClockSkew), NotAfter: now.Add(cfg.CertificateLifetime), BasicConstraintsValid: true, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, URIs: []*url.URL{uri}} + + if err := ctx.Err(); err != nil { + return nil, err + } + + leaf, err := x509.CreateCertificate(rand.Reader, template, state.certificate, pub, state.key) + if err != nil { + return nil, wire.Unavailable + } + + response := wire.BootstrapResponse{SchemaVersion: wire.SchemaVersion, Cluster: identity.cluster, Node: identity.node, Enrollment: request.Enrollment, CertificateChain: [][]byte{leaf, state.certificate.Raw}} + + encoded, err := wire.EncodeBootstrap(response) + if err != nil { + return nil, err + } + + if err := ctx.Err(); err != nil { + return nil, err + } + + return encoded, nil +} + +func bootstrapPublicKey(der []byte) (ed25519.PublicKey, error) { + csr, err := x509.ParseCertificateRequest(der) + if err != nil || csr.CheckSignature() != nil { + return nil, wire.InvalidRequest + } + + pub, ok := csr.PublicKey.(ed25519.PublicKey) + if !ok { + return nil, wire.InvalidRequest + } + + return pub, nil +} + +// authenticateCertificate requires a verified chain, the client-auth usage, +// cluster-scoped Node URI SAN, and current validity against local trust. Recheck on +// every poll: an existing TLS connection must not bypass certificate expiry. +// Membership and Kubernetes workload state are not certificate authorization. +func authenticateCertificate(ctx context.Context, trust *trustStore, cfg Config, state *tls.ConnectionState) (NodeIdentity, error) { + if err := ctx.Err(); err != nil { + return NodeIdentity{}, err + } + + if state == nil || !state.HandshakeComplete || len(state.VerifiedChains) == 0 || len(state.PeerCertificates) == 0 { + return NodeIdentity{}, wire.Unauthenticated + } + + now := time.Now() + leaf := state.PeerCertificates[0] + + node, err := certificateNode(leaf, cfg.Cluster, now) + if err != nil { + return NodeIdentity{}, err + } + + roots, err := trust.pool() + if err != nil { + return NodeIdentity{}, wire.Unavailable + } + + intermediates := x509.NewCertPool() + for _, cert := range state.PeerCertificates[1:] { + intermediates.AddCert(cert) + } + + chains, err := leaf.Verify(x509.VerifyOptions{Roots: roots, Intermediates: intermediates, CurrentTime: now, KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}}) + if err != nil { + return NodeIdentity{}, wire.Unauthenticated + } + + expires := leaf.NotAfter + for _, cert := range chains[0] { + if cert.NotAfter.Before(expires) { + expires = cert.NotAfter + } + } + + if err := ctx.Err(); err != nil { + return NodeIdentity{}, err + } + + if !time.Now().Before(expires) { + return NodeIdentity{}, wire.Unauthenticated + } + + return NodeIdentity{cluster: cfg.Cluster, node: node, expires: expires}, nil +} + +func certificateNode(leaf *x509.Certificate, cluster wire.ClusterID, now time.Time) (wire.NodeID, error) { + if !validClientCertificate(leaf, now) || len(leaf.URIs) != 1 { + return "", wire.Unauthenticated + } + + uri := leaf.URIs[0] + + node := wire.NodeID(strings.TrimPrefix(uri.Path, "/node/")) + if !wire.ValidUUID(string(node)) || uri.String() != "spiffe://"+uri.Host+"/node/"+string(node) { + return "", wire.Unauthenticated + } + + if uri.Host != string(cluster) { + return "", wire.Forbidden + } + + return node, nil +} + +func validClientCertificate(leaf *x509.Certificate, now time.Time) bool { + _, ed25519Key := leaf.PublicKey.(ed25519.PublicKey) + + return ed25519Key && !leaf.IsCA && leaf.KeyUsage == x509.KeyUsageDigitalSignature && + len(leaf.ExtKeyUsage) == 1 && leaf.ExtKeyUsage[0] == x509.ExtKeyUsageClientAuth && + len(leaf.UnknownExtKeyUsage) == 0 && !now.Before(leaf.NotBefore) && now.Before(leaf.NotAfter) +} + +type bootstrap struct { + owner *Authority + Client client.Writer + APIReader client.Reader + Config Config + Issuer *issuer +} + +// Authenticate performs TokenReview for racer-control, checks the live bound Pod +// UID and authorized ServiceAccount/workload, and resolves its assigned Node UID. +// Token contents, CSR contents, and requested names are not authority on their own. +func (b *bootstrap) Authenticate(ctx context.Context, r *http.Request) (NodeIdentity, error) { + cfg := b.Config + + if err := ctx.Err(); err != nil { + return NodeIdentity{}, err + } + + if b.Client == nil || b.APIReader == nil { + return NodeIdentity{}, wire.Unavailable + } + + status, token, err := reviewBearer(ctx, b.Client, r, wire.TokenAudience, cfg.MaxTokenBytes) + if err != nil { + return NodeIdentity{}, err + } + + if status.User.Username != "system:serviceaccount:"+cfg.Namespace+":"+cfg.DataplaneServiceAccount { + return NodeIdentity{}, wire.Forbidden + } + // TokenReview authenticates the token. Its JWT expiration is used only to + // shorten authorization, never to establish identity or extend validity. + expires, err := tokenExpiration(token) + if err != nil { + return NodeIdentity{}, err + } + + node, err := b.authorizeTokenBinding(ctx, cfg, status.User) + if err != nil { + return NodeIdentity{}, err + } + + if err := ctx.Err(); err != nil { + return NodeIdentity{}, err + } + + if !time.Now().Before(expires) { + return NodeIdentity{}, wire.Unauthenticated + } + + return NodeIdentity{owner: b.owner, bearer: true, cluster: cfg.Cluster, node: wire.NodeID(node.UID), nodeName: node.Name, expires: expires}, nil +} + +func (b *bootstrap) authorizeTokenBinding(ctx context.Context, cfg Config, user authv1.UserInfo) (*corev1.Node, error) { + podName, podUID := singleExtra(user, "pod-name"), singleExtra(user, "pod-uid") + if podName == "" || podUID == "" || user.UID == "" { + return nil, wire.Unauthenticated + } + + var pod corev1.Pod + if err := b.APIReader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: podName}, &pod); err != nil { + return nil, authorizationError(err) + } + + if string(pod.UID) != podUID { + return nil, wire.Forbidden + } + + if err := authorizePod(ctx, b.APIReader, cfg, &pod, user.UID); err != nil { + return nil, err + } + + var node corev1.Node + if err := b.APIReader.Get(ctx, client.ObjectKey{Name: pod.Spec.NodeName}, &node); err != nil { + return nil, authorizationError(err) + } + + if !authorizedNode(&node) { + return nil, wire.Forbidden + } + // Require unambiguous current node bindings as well as the live Pod binding. + for key, want := range map[string]string{"node-name": node.Name, "node-uid": string(node.UID)} { + if singleExtra(user, key) != want { + return nil, wire.Forbidden + } + } + + return &node, nil +} + +func singleExtra(user authv1.UserInfo, key string) string { + values := user.Extra["authentication.kubernetes.io/"+key] + if len(values) != 1 { + return "" + } + + return values[0] +} + +// reviewBearer shares only token parsing and API authentication. Callers retain +// their distinct workload, service-account, binding, and expiration policies. +func reviewBearer(ctx context.Context, c client.Writer, r *http.Request, audience string, maxBytes int) (authv1.TokenReviewStatus, string, error) { + values := r.Header.Values("Authorization") + if len(values) != 1 { + return authv1.TokenReviewStatus{}, "", wire.Unauthenticated + } + + scheme, token, ok := strings.Cut(values[0], " ") + if !ok || !strings.EqualFold(scheme, "Bearer") || token == "" || strings.ContainsAny(token, " \t\r\n,") || maxBytes > 0 && len(token) > maxBytes { + return authv1.TokenReviewStatus{}, "", wire.Unauthenticated + } + + review := &authv1.TokenReview{Spec: authv1.TokenReviewSpec{Token: token, Audiences: []string{audience}}} + if c == nil || c.Create(ctx, review) != nil { + return authv1.TokenReviewStatus{}, "", wire.Unavailable + } + + status := review.Status + if !status.Authenticated || status.Error != "" || !slices.Contains(status.Audiences, audience) { + return authv1.TokenReviewStatus{}, "", wire.Unauthenticated + } + + return status, token, nil +} + +func tokenExpiration(token string) (time.Time, error) { + parts := strings.Split(token, ".") + if len(parts) != 3 { + return time.Time{}, wire.Unauthenticated + } + + payload, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return time.Time{}, wire.Unauthenticated + } + + var claims struct { + Expiration int64 `json:"exp"` + } + if json.Unmarshal(payload, &claims) != nil || claims.Expiration <= 0 { + return time.Time{}, wire.Unauthenticated + } + + expires := time.Unix(claims.Expiration, 0) + if !time.Now().Before(expires) { + return time.Time{}, wire.Unauthenticated + } + + return expires, nil +} + +// enroll validates CSR proof of possession and binds the issued identity to the +// token, not caller-provided SANs. Every issuance uses a token, including renewal. +// Retries correlate by enrollment ID; there is no persistent receipt ledger. +// The returned bytes are the issuer's bounded, validated JSON response. +func (b *bootstrap) enroll(ctx context.Context, r *http.Request, request wire.BootstrapRequest) ([]byte, EnrollmentHint, error) { + identity, err := b.Authenticate(ctx, r) + if err != nil { + return nil, EnrollmentHint{}, err + } + + if b.Issuer == nil { + return nil, EnrollmentHint{}, wire.Unavailable + } + + ctx, cancel := context.WithDeadline(ctx, identity.expires) + defer cancel() + + response, err := b.Issuer.Issue(ctx, identity, request) + if err != nil { + return nil, EnrollmentHint{}, err + } + // Resolve the same live UID again before persisting an authenticated proposal. + // Both annotations are proposals only; explicit administrator values win. + var live corev1.Node + if err := b.APIReader.Get(ctx, client.ObjectKey{Name: identity.nodeName}, &live); err != nil { + return nil, EnrollmentHint{}, err + } + + node := &live + if wire.NodeID(node.UID) == identity.node { + if !authorizedNode(node) { + return nil, EnrollmentHint{}, wire.Forbidden + } + + response, err = bootstrapBlockDevices(ctx, response, node) + if err != nil { + return nil, EnrollmentHint{}, err + } + + return response, EnrollmentHint{Node: *node, Shares: request.Shares, RDMANICs: wire.CanonicalRDMANICs(request.RDMANICs), Expires: identity.expires}, nil + } + + return nil, EnrollmentHint{}, wire.Forbidden +} + +func bootstrapBlockDevices(ctx context.Context, response []byte, node *corev1.Node) ([]byte, error) { + pattern := node.Annotations[wire.BlockDevicesAnnotation] + if pattern == "" { + return response, nil + } + + reason := "" + if len(pattern) > 1024 { + reason = "exceeds 1024 bytes" + } else if _, err := regexp.Compile(pattern); err != nil { + reason = "invalid regular expression" + } + + if reason != "" { + ctrl.LoggerFrom(ctx).Info("warning: ignoring block device annotation; using file-backed storage", "node", node.Name, "annotation", wire.BlockDevicesAnnotation, "reason", reason) + return response, nil + } + + decoded, err := wire.DecodeBootstrapResponse(bytes.NewReader(response)) + if err != nil { + return nil, err + } + + // The dataplane matches /dev/disk/by-id basenames at process startup only. + decoded.BlockDevices = pattern + + return wire.EncodeBootstrap(decoded) +} + +// EnrollmentHint is a detached UID/resource-version-bound proposal, not authority. +// The root adapter persists it with optimistic concurrency after gate release. +type EnrollmentHint struct { + Node corev1.Node + Shares uint32 + RDMANICs []wire.RDMANIC + Expires time.Time +} + +func authorizationError(err error) error { + if apierrors.IsNotFound(err) { + return wire.Forbidden + } + + return wire.Unavailable +} + +func authorizedNode(node *corev1.Node) bool { + _, excluded := node.Labels[wire.ExclusionLabel] + return node.Name != "" && wire.ValidUUID(string(node.UID)) && node.DeletionTimestamp == nil && !excluded +} + +// The configured namespace/name designate the managed workload. A Pod must be +// controlled by that exact current DaemonSet UID, not just carry matching labels. +func managedPod(cfg Config, pod *corev1.Pod) bool { + if pod.Namespace != cfg.Namespace || pod.UID == "" || pod.DeletionTimestamp != nil || + pod.Spec.NodeName == "" || pod.Spec.ServiceAccountName != cfg.DataplaneServiceAccount || + pod.Status.Phase == corev1.PodSucceeded || pod.Status.Phase == corev1.PodFailed { + return false + } + + owner := metav1.GetControllerOf(pod) + if owner == nil || owner.APIVersion != "apps/v1" || owner.Kind != "DaemonSet" || + owner.Name != cfg.DaemonSetName || owner.UID == "" { + return false + } + + return true +} + +func authorizePod(ctx context.Context, reader client.Reader, cfg Config, pod *corev1.Pod, serviceAccountUID string) error { + if !managedPod(cfg, pod) { + return wire.Forbidden + } + + ownership, err := members.ReadWorkloadIdentities(ctx, reader, cfg.Namespace, cfg.DaemonSetName) + if err != nil { + return authorizationError(err) + } + + if !ownership.Owns(pod) { + return wire.Forbidden + } + + var sa corev1.ServiceAccount + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.DataplaneServiceAccount}, &sa); err != nil { + return authorizationError(err) + } + + if sa.UID == "" || sa.DeletionTimestamp != nil || serviceAccountUID != "" && string(sa.UID) != serviceAccountUID { + return wire.Forbidden + } + + return ctx.Err() +} + +// versionRecord is the only persisted topology bookkeeping. Hashes cover +// canonical content excluding counters. No member history or publication bytes +// are stored. ResourceVersion CAS must precede installing a publication. +type versionRecord struct { + Cluster wire.ClusterID `json:"cluster"` + Sequence wire.Sequence `json:"sequence,string"` + MembershipVersion wire.MembershipVersion `json:"membership_version,string"` + ContentHash string `json:"content_hash"` + MembershipHash string `json:"membership_hash"` +} + +// preparedPublication owns encoded candidate bytes. Publishers require CAS; +// replicas require canonical validation and an authoritative durable confirmation. +type preparedPublication struct { + owner *publicationStore + previous versionRecord + resourceVersion string + record versionRecord + encoded string + delta string + deltaBase string + deltaSequence wire.Sequence +} + +type committedPublication struct { + owner *publicationStore + record versionRecord + encoded string + delta string + deltaBase string + deltaSequence wire.Sequence + leadership context.Context + authority context.Context +} + +// publicationResponse is a response-only view, never installable state. The +// caller owns one admitted image context through both write and flush. +type publicationResponse struct{ encoded string } + +func (p publicationResponse) writeTo(ctx context.Context, w io.Writer) (int64, error) { + var written int64 + + for remaining := p.encoded; remaining != ""; { + if err := ctx.Err(); err != nil { + return written, err + } + // ResponseWriter need not implement StringWriter. Limit conversion scratch + // to 32 KiB rather than allocating a full publication for every response. + chunk := remaining[:min(len(remaining), 32*1024)] + + n, err := w.Write([]byte(chunk)) + if err == nil { + err = ctx.Err() + } + + written += int64(n) + if err != nil { + return written, err + } + + if n != len(chunk) { + return written, io.ErrShortWrite + } + + remaining = remaining[n:] + } + + return written, ctx.Err() +} + +// admit pins one response to its image's revocable authority and the +// freshness deadline at write admission. Later confirmations cannot extend an +// in-flight response. Suspension revokes it; ordinary advancement does not. +func (p *committedPublication) admit(parent context.Context) (*Admission, context.CancelFunc, error) { + if p == nil || p.owner == nil { + return nil, nil, wire.Unavailable + } + + owner := p.owner + owner.mu.Lock() + defer owner.mu.Unlock() + + if _, err := owner.currentLocked(); err != nil { + return nil, nil, err + } + + if p != owner.current || p.authority == nil || p.authority.Err() != nil { + return nil, nil, wire.Unavailable + } + + guard, cancel := newAdmission(parent, p.authority, owner.confirmed.Add(owner.maxAge)) + guard.image = p + + return guard, cancel, nil +} + +// publicationStore owns only the current immutable publication and one broadcast +// notification. Older state belongs to dataplanes; poll admission belongs to Server. +type publicationStore struct { + mu sync.Mutex + current *committedPublication + changed chan struct{} + suspended bool + process context.Context + confirmed time.Time + maxAge time.Duration + observed versionRecord + revoke context.CancelFunc + epoch context.Context +} + +func newPublications() *publicationStore { + return &publicationStore{changed: make(chan struct{}), maxAge: 30 * time.Second} +} + +func (p *publicationStore) bindProcess(ctx context.Context) { + p.mu.Lock() + defer p.mu.Unlock() + + p.process = ctx +} + +func (p *publicationStore) notifyLocked() { close(p.changed); p.changed = make(chan struct{}) } + +func (p *publicationStore) Prepare(previous versionRecord, resourceVersion string, members AcceptedMembers, caches []wire.CacheDefinition) (*preparedPublication, error) { + if !previous.valid() || resourceVersion == "" { + return nil, wire.InvalidRequest + } + + v := wire.Publication{SchemaVersion: wire.SchemaVersion, Cluster: previous.Cluster, Caches: caches, Members: make([]wire.Member, 0, len(members))} + for id, member := range members { + if id != member.Node { + return nil, wire.InvalidRequest + } + + v.Members = append(v.Members, member) + } + + candidate, err := wire.NewCanonicalCandidate(v) + if err != nil { + return nil, err + } + + content, membership, err := candidate.ContentHashes() + if err != nil { + return nil, err + } + + record := previous + if content != previous.ContentHash { + if record.Sequence == ^wire.Sequence(0) { + return nil, wire.Unavailable + } + + record.Sequence++ + } + + if membership != previous.MembershipHash { + if content == previous.ContentHash || record.MembershipVersion == ^wire.MembershipVersion(0) { + return nil, wire.Unavailable + } + + record.MembershipVersion++ + } + + record.ContentHash, record.MembershipHash = content, membership + + encoded, err := candidate.EncodePublication(record.Sequence, record.MembershipVersion) + if err != nil { + return nil, err + } + + prepared := &preparedPublication{owner: p, previous: previous, resourceVersion: resourceVersion, record: record, encoded: string(encoded)} + v.Sequence, v.MembershipVersion = record.Sequence, record.MembershipVersion + p.prepareDelta(prepared, v) + + return prepared, nil +} + +func (p *publicationStore) prepareDelta(prepared *preparedPublication, v wire.Publication) { + previous, record := prepared.previous, prepared.record + // Capture immutable base under the lock, then diff/encode entirely outside it. + if current, err := p.Current(); err == nil && current.record == previous && record.Sequence > previous.Sequence { + if base, err := wire.DecodePublication(strings.NewReader(current.encoded)); err == nil { + if delta, err := wire.EncodeDelta(base, v); err == nil && len(delta) < len(prepared.encoded) { + prepared.delta, prepared.deltaBase = string(delta), previous.ContentHash + prepared.deltaSequence = previous.Sequence + } + } + } +} + +// CommitVersion mints publisher installable state after a resource-version CAS. +// Even unchanged content is CAS-confirmed; its counters and bytes remain identical. +func (r *publisher) CommitVersion(ctx context.Context, p *preparedPublication) (*committedPublication, error) { + cfg := r.Config + + if err := ctx.Err(); err != nil { + return nil, err + } + + if p == nil || p.owner != r.Publications || p.previous.Cluster != cfg.Cluster { + return nil, wire.InvalidRequest + } + + cm, previous, err := readVersion(ctx, r.APIReader, cfg) + if err != nil { + r.suspendInvalidAuthority(err) + return nil, err + } + + if err := r.Publications.confirm(previous); err != nil { + return nil, err + } + + if cm.ResourceVersion != p.resourceVersion || previous != p.previous { + return nil, apierrors.NewConflict(corev1.Resource("configmaps"), cm.Name, wire.Conflict) + } + + cm.Data = versionData(p.record) + + if err := ctx.Err(); err != nil { + return nil, err + } + + if err := r.Update(ctx, cm); err != nil { + return nil, err + } + + if err := ctx.Err(); err != nil { + return nil, err + } + + return &committedPublication{owner: p.owner, record: p.record, encoded: p.encoded, delta: p.delta, deltaBase: p.deltaBase, deltaSequence: p.deltaSequence, leadership: ctx}, nil +} + +// ForBase returns a shared bounded delta only for the exact authenticated cursor. +// Coalesced/skipped updates and controller restarts automatically use the full image. +func (p *committedPublication) ForBase(sequence wire.Sequence, hash string) publicationResponse { + if sequence == 0 || sequence != p.deltaSequence || hash == "" || hash != p.deltaBase || p.delta == "" { + return publicationResponse{encoded: p.encoded} + } + + return publicationResponse{encoded: p.delta} +} + +func (p *publicationStore) Install(next *committedPublication) error { + p.mu.Lock() + defer p.mu.Unlock() + + if next == nil || next.owner != p || next.leadership == nil || !next.record.valid() || next.encoded == "" { + return wire.InvalidRequest + } + + if err := next.leadership.Err(); err != nil { + return err + } + + if err := p.observeLocked(next.record); err != nil { + return err + } + + if p.process != nil { + copy := *next + copy.leadership = p.process + next = © + } + + if current := p.current; current != nil { + // observeLocked is the sole counter/hash high-water guard. Installed + // state can only lag that observation, never exceed it. + if next.record.Sequence == current.record.Sequence { + return p.reconfirmLocked(next) + } + } + + p.current = p.authorizeLocked(next) + p.confirmed = time.Now() + p.suspended = false + p.notifyLocked() + + return nil +} + +func (p *publicationStore) reconfirmLocked(next *committedPublication) error { + current := p.current + if next.record != current.record || next.encoded != current.encoded { + return wire.Conflict + } + + p.confirmed = time.Now() + if current.leadership.Err() != nil || p.suspended { + p.current = p.authorizeLocked(next) + } + + if p.suspended { + p.suspended = false + p.notifyLocked() + } + + return nil +} + +func (p *publicationStore) authorizeLocked(next *committedPublication) *committedPublication { + if p.epoch == nil || p.epoch.Err() != nil { + p.epoch, p.revoke = context.WithCancel(next.leadership) + } + + copy := *next + copy.authority = p.epoch + + return © +} + +// confirm never promotes a hash to an image. It only renews the freshness of an +// already validated image when all durable counters and hashes still match. +func (p *publicationStore) confirm(record versionRecord) error { + p.mu.Lock() + defer p.mu.Unlock() + + if err := p.observeLocked(record); err != nil { + return err + } + + if p.current != nil && record == p.current.record && !p.suspended { + p.confirmed = time.Now() + } + + return nil +} + +// observed is independent of installed bytes, including before the first image. +// Suspension never forgets this high-water mark. Skipped versions may return to +// earlier hashes, but counters and unchanged membership versions must agree. +func (p *publicationStore) observeLocked(record versionRecord) error { + if !record.valid() { + p.suspendLocked() + return wire.Conflict + } + + if old := p.observed; old.Sequence != 0 { + rollback := record.Cluster != old.Cluster || record.Sequence < old.Sequence || record.MembershipVersion < old.MembershipVersion + conflictingReplay := record.Sequence == old.Sequence && record != old + changedMembership := record.MembershipVersion == old.MembershipVersion && record.MembershipHash != old.MembershipHash + // Check progression only after ruling out rollback, before subtracting + // unsigned counters. Each membership change requires a publication change. + if rollback || conflictingReplay || changedMembership || uint64(record.MembershipVersion-old.MembershipVersion) > uint64(record.Sequence-old.Sequence) { + p.suspendLocked() + return wire.Conflict + } + } + + p.observed = record + + return nil +} + +func (p *publicationStore) Current() (*committedPublication, error) { + p.mu.Lock() + defer p.mu.Unlock() + + return p.currentLocked() +} + +// CurrentAndSubscribe atomically reads the current publication and subscribes to +// changes, including when unavailable. The channel closes on install or suspension; +// leadership cancellation must be observed separately by the caller. +func (p *publicationStore) CurrentAndSubscribe() (*committedPublication, <-chan struct{}, error) { + p.mu.Lock() + defer p.mu.Unlock() + + current, err := p.currentLocked() + + return current, p.changed, err +} + +func (p *publicationStore) currentLocked() (*committedPublication, error) { + if p.current == nil || p.suspended || time.Since(p.confirmed) >= p.maxAge { + return nil, wire.Unavailable + } + + if err := p.current.leadership.Err(); err != nil { + return nil, err + } + + return p.current, nil +} + +// Suspend withdraws readiness and wakes polls after failure to validate durable +// authority. Keep bytes/history for a later successful CAS, never serve them until +// then. Invalid desired inputs alone do not suspend the last valid publication. +func (p *publicationStore) Suspend() { + p.mu.Lock() + defer p.mu.Unlock() + + p.suspendLocked() +} + +func (p *publicationStore) suspendLocked() { + if p.revoke != nil { + p.revoke() + } + + if !p.suspended { + p.suspended = true + p.notifyLocked() + } +} + +// Wait rejects invalid identities/cursors and honors context cancellation and +// certificate expiration. It shares publication bytes and broadcast notifications; +// callers own admission for the full response lifetime, including writes and flush. +func (p *publicationStore) Wait(ctx context.Context, identity NodeIdentity, after *wire.Sequence) (*committedPublication, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + + if !wire.ValidUUID(string(identity.node)) || !time.Now().Before(identity.expires) { + return nil, wire.Unauthenticated + } + + current, changed, err := p.CurrentAndSubscribe() + if err != nil { + return nil, err + } + + if identity.cluster != current.record.Cluster { + return nil, wire.Forbidden + } + + if after != nil && *after == 0 { + return nil, wire.Conflict + } + + if after != nil && *after > current.record.Sequence { + return nil, wire.Unavailable + } + + if err := ctx.Err(); err != nil { + return nil, err + } + + if !time.Now().Before(identity.expires) { + return nil, wire.Unauthenticated + } + + if after == nil || current.record.Sequence > *after { + if err := current.leadership.Err(); err != nil { + return nil, err + } + + return current, nil + } + + return p.waitForPublication(ctx, identity, *after, current, changed) +} + +func (p *publicationStore) waitForPublication(ctx context.Context, identity NodeIdentity, after wire.Sequence, current *committedPublication, changed <-chan struct{}) (*committedPublication, error) { + timer := time.NewTimer(wire.PollWait) + defer timer.Stop() + + expiration := time.NewTimer(time.Until(identity.expires)) + defer expiration.Stop() + + freshness := time.NewTicker(min(p.maxAge, time.Second)) + defer freshness.Stop() + + for { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-current.leadership.Done(): + return nil, current.leadership.Err() + case <-expiration.C: + return nil, wire.Unauthenticated + case <-freshness.C: + if _, err := p.Current(); err != nil { + return nil, err + } + + continue + case <-timer.C: + return nil, pollTimeoutError(ctx, current.leadership, identity.expires) + case <-changed: + } + + if err := ctx.Err(); err != nil { + return nil, err + } + + if !time.Now().Before(identity.expires) { + return nil, wire.Unauthenticated + } + + var err error + + current, changed, err = p.CurrentAndSubscribe() + if err != nil { + return nil, err + } + + if current.record.Sequence > after { + return current, nil + } + } +} + +func pollTimeoutError(ctx, leadership context.Context, expires time.Time) error { + if err := ctx.Err(); err != nil { + return err + } + + if err := leadership.Err(); err != nil { + return err + } + + if !time.Now().Before(expires) { + return wire.Unauthenticated + } + + return nil +} diff --git a/internal/racer/authority/authority_test.go b/internal/racer/authority/authority_test.go new file mode 100644 index 000000000..95def338b --- /dev/null +++ b/internal/racer/authority/authority_test.go @@ -0,0 +1,1900 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "bytes" + "context" + "crypto/ecdsa" + "crypto/ed25519" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "net/http" + "net/http/httptest" + "net/url" + "os" + "reflect" + "slices" + "strings" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/go-logr/logr/funcr" + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + clientgoscheme "k8s.io/client-go/kubernetes/scheme" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + "sigs.k8s.io/controller-runtime/pkg/envtest" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestAuthorityOperationsSharePrivateGate(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + + _, err := a.PublishTopology(ctx, func(ctx context.Context) (TopologyObservation, error) { + blocked, stop := context.WithCancel(ctx) + stop() + + _, err := a.ReconcileCredentials(blocked) + require.ErrorIs(t, err, context.Canceled) + require.ErrorIs(t, a.Observe(blocked), context.Canceled) + _, err = a.TrustPool() + require.NoError(t, err, "failed admission must not invalidate trust") + + select { + case <-a.gate.token: + t.Fatal("discovery callback escaped operation gate") + default: + } + + return f.a.Topology.observeTopology(ctx) + }) + require.NoError(t, err) + require.NoError(t, a.Observe(ctx), "returned hints must not retain gate") +} + +func TestAuthorityPublicationHistoryIsOperationOwned(t *testing.T) { + f := newServingFixture(t) + a, r := f.a.authority, f.a.Topology + before := cloneAccepted(a.accepted) + update, err := a.PublishTopology(t.Context(), r.observeTopology) + require.NoError(t, err) + + member := update.Members[testNodeUID] + member.Shares = 999 + update.Members[testNodeUID] = member + + require.Equal(t, before, a.accepted, "returned hints alias authority history") + + update, err = a.PublishTopology(t.Context(), r.observeTopology) + require.NoError(t, err) + delete(update.Members, testNodeUID) + require.Equal(t, before, a.accepted, "annotation hints alias authority history") + + var node corev1.Node + require.NoError(t, r.Get(t.Context(), client.ObjectKey{Name: "worker"}, &node)) + node.Annotations = map[string]string{wire.SharesAnnotation: "7"} + require.NoError(t, r.Update(t.Context(), &node)) + base := r.Client.(client.WithWatch) + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + return apierrors.NewConflict(corev1.Resource("configmaps"), "version", wire.Conflict) + }}) + a.publisher.Writer = r.Client + _, err = a.PublishTopology(t.Context(), r.observeTopology) + require.True(t, apierrors.IsConflict(err)) + require.Equal(t, before, a.accepted, "failed CAS advanced history") + + r.Client = base + a.publisher.Writer = base + _, err = a.PublishTopology(t.Context(), r.observeTopology) + require.NoError(t, err) + require.EqualValues(t, 7, a.accepted[testNodeUID].Shares) +} + +func TestAuthorityHandlesRetainRevocationSemantics(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + image, err := a.Current() + require.NoError(t, err) + write, stop, err := image.Admit(t.Context()) + require.NoError(t, err) + + defer stop() + + trust, stopTrust, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stopTrust() + + deadline, _ := trust.Context().Deadline() + + require.NoError(t, a.Observe(t.Context())) + + nextDeadline, _ := trust.Context().Deadline() + require.Equal(t, deadline, nextDeadline) + require.NoError(t, trust.Check(t.Context())) + + var node corev1.Node + require.NoError(t, f.a.Topology.Get(t.Context(), client.ObjectKey{Name: "worker"}, &node)) + node.Annotations = map[string]string{wire.SharesAnnotation: "7"} + require.NoError(t, f.a.Topology.Update(t.Context(), &node)) + _, err = a.PublishTopology(t.Context(), f.a.Topology.observeTopology) + require.NoError(t, err) + require.NoError(t, write.Check(t.Context()), "replacement must preserve admitted image") + _, _, err = image.Admit(t.Context()) + require.ErrorIs(t, err, wire.Unavailable, "old image cannot admit new responses") + require.NoError(t, trust.Check(t.Context()), "publication replacement is not trust invalidation") + + secret := &corev1.Secret{} + require.NoError(t, f.a.Topology.Get(t.Context(), client.ObjectKey{Namespace: a.config.Namespace, Name: a.config.CredentialsSecretName}, secret)) + require.NoError(t, f.a.Topology.Delete(t.Context(), secret)) + require.Error(t, a.Observe(t.Context())) + require.ErrorIs(t, trust.Check(t.Context()), context.Canceled) + require.ErrorIs(t, write.Check(t.Context()), context.Canceled) +} + +func TestAuthorityObservationFailureDoesNotPublish(t *testing.T) { + r := initializedTopology(t) + a := r.authority + boom := errors.New("discovery failed") + _, err := a.PublishTopology(t.Context(), func(context.Context) (TopologyObservation, error) { return TopologyObservation{}, boom }) + require.ErrorIs(t, err, boom) + _, err = a.Current() + require.ErrorIs(t, err, wire.Unavailable) + require.Empty(t, a.accepted) + _, err = a.PublishTopology(t.Context(), r.observeTopology) + require.NoError(t, err, "failed callback did not release gate") +} + +func TestAuthorityBlockedOperationsHonorCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + results := make(chan error, 4) + _, err := a.PublishTopology(t.Context(), func(ctx context.Context) (TopologyObservation, error) { + blocked, cancelBlocked := context.WithCancel(ctx) + defer cancelBlocked() + + go func() { _, err := a.ReconcileCredentials(blocked); results <- err }() + go func() { results <- a.Observe(blocked) }() + go func() { + _, err := a.PublishTopology(blocked, func(context.Context) (TopologyObservation, error) { + t.Error("blocked publication entered discovery") + return TopologyObservation{}, nil + }) + results <- err + }() + go func() { + _, err := a.Issue(blocked, NodeIdentity{owner: a, bearer: true, cluster: a.config.Cluster, node: testNodeUID, expires: time.Now().Add(time.Hour)}, f.request) + results <- err + }() + + synctest.Wait() + + select { + case err := <-results: + t.Fatalf("operation bypassed held gate: %v", err) + default: + } + + cancelBlocked() + + for range 4 { + require.ErrorIs(t, <-results, context.Canceled) + } + + return f.a.Topology.observeTopology(ctx) + }) + require.NoError(t, err) + require.NoError(t, a.TrustReady(), "canceled waiters invalidated accepted trust") + require.NoError(t, a.Observe(ctx)) + }) +} + +func TestAuthorityConstructorCopiesConfigWithoutIO(t *testing.T) { + cfg := testConfig(t) + a := New(cfg, Dependencies{}) + cfg.Cluster = "" + require.NotEqual(t, cfg.Cluster, a.config.Cluster) + require.ErrorIs(t, a.PublicationReady(), wire.Unavailable) + require.ErrorIs(t, a.TrustReady(), wire.Unavailable) + require.Empty(t, a.accepted) +} + +// Fixtures inject faults into real Authority operations without controller policy. +type topologyFixture struct { + client.Client + APIReader client.Reader + Config Config + Publications *Publications + Trust *Trust + authority *Authority +} +type credentialsFixture struct { + client.Client + APIReader client.Reader + Config Config + Trust *Trust + Now func() time.Time + authority *Authority +} +type Application struct { + Topology *topologyFixture + Keyring *credentialsFixture + Server *fixtureServer + authority *Authority +} +type fixtureServer struct { + Bootstrap *Bootstrap + Trust *Trust + Publications *Publications + Config Config +} + +func Assemble(cfg Config, c client.Client, reader client.Reader) *Application { + a := New(cfg, Dependencies{Writer: c, Reader: reader}) + + return &Application{ + authority: a, + Topology: &topologyFixture{Client: c, APIReader: reader, Config: cfg, Publications: a.publications, Trust: a.trust, authority: a}, + Keyring: &credentialsFixture{Client: c, APIReader: reader, Config: cfg, Trust: a.trust, authority: a}, + Server: &fixtureServer{Bootstrap: a.bootstrap, Trust: a.trust, Publications: a.publications, Config: cfg}, + } +} + +func (a *Application) Recover(ctx context.Context, writer client.Writer) error { + return a.authority.Recover(ctx, writer) +} + +func (r *topologyFixture) operations() *Authority { + r.authority.publisher.Writer = r.Client + r.authority.publisher.APIReader = r.APIReader + r.authority.publisher.Config = r.Config.effective() + + return r.authority +} + +func (r *topologyFixture) CommitVersion(ctx context.Context, p *PreparedPublication) (*CommittedPublication, error) { + return r.operations().publisher.CommitVersion(ctx, p) +} + +func (r *topologyFixture) observeTopology(ctx context.Context) (TopologyObservation, error) { + var nodes corev1.NodeList + if err := r.List(ctx, &nodes); err != nil { + return TopologyObservation{}, err + } + + var caches racerv1.ClusterCacheList + if err := r.APIReader.List(ctx, &caches); err != nil { + return TopologyObservation{}, err + } + + catalog, err := BuildCatalog(caches.Items) + if err != nil { + return TopologyObservation{}, err + } + + ids, err := members.ReadWorkloadIdentities(ctx, r.APIReader, r.Config.Namespace, r.Config.DaemonSetName) + if err != nil { + return TopologyObservation{}, err + } + + pods := map[string][]corev1.Pod{} + + for _, node := range nodes.Items { + var list corev1.PodList + if err := r.List(ctx, &list, client.InNamespace(r.Config.Namespace), client.MatchingFields{podNodeIndex: node.Name}); err != nil { + return TopologyObservation{}, err + } + + pods[node.Name] = list.Items + } + + return TopologyObservation{Nodes: nodes, Catalog: catalog, Input: members.Input{Nodes: nodes.Items, PodsByNode: pods, Ownership: ids, PeerPort: 8082}}, nil +} + +func (r *credentialsFixture) operations() *Authority { + a := r.authority + a.credentials.Writer, a.credentials.APIReader, a.credentials.Config, a.credentials.Now = r.Client, r.APIReader, r.Config.effective(), r.Now + + return a +} + +const podNodeIndex = "spec.nodeName" + +func podNodeKeys(obj client.Object) []string { + pod, ok := obj.(*corev1.Pod) + if !ok || pod.Spec.NodeName == "" { + return nil + } + + return []string{pod.Spec.NodeName} +} + +func LoadConfig() (Config, error) { + return Config{Cluster: "22222222-2222-4222-8222-222222222222", Namespace: "racer", DataplaneServiceAccount: "racer-dataplane", ControllerServiceAccount: "racer-controller", DaemonSetName: "racer-dataplane", CredentialsSecretName: "racer-credentials", VersionConfigMapName: "racer-version", InstallationConfigMapName: "racer-installation", Rotation: RotationPolicy{Interval: 24 * time.Hour, PrepareFor: time.Hour, RetainFor: 48 * time.Hour}, CertificateLifetime: wire.CertificateLifetime, SnapshotMaxAge: 30 * time.Second, MaxTokenBytes: 16384}, nil +} + +const ( + testNodeUID = "11111111-1111-4111-8111-111111111111" + testOtherUID = "22222222-2222-4222-8222-222222222222" +) + +// Legacy names exist exclusively inside whitebox tests, never the package API. +type ( + Trust = trustStore + Publications = publicationStore + Issuer = issuer + Bootstrap = bootstrap + PreparedPublication = preparedPublication + CommittedPublication = committedPublication + RotationState = rotationState + VersionRecord = versionRecord +) + +func NewPublications() *publicationStore { return newPublications() } + +var AuthenticateCertificate = authenticateCertificate + +func (r *credentialsFixture) now() time.Time { return credentialTime(r.Now) } + +func catalogCache(name string, uid types.UID) racerv1.ClusterCache { + return racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: name, UID: uid}} +} + +const ( + testDaemonSetUID types.UID = "33333333-3333-4333-8333-333333333333" + DataplaneDaemonSetName string = "racer-dataplane" +) + +func memberNode() corev1.Node { + return corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "node-a", UID: testNodeUID}} +} + +func memberPod(uid types.UID, created int64, ip string) corev1.Pod { + controller := true + return corev1.Pod{ObjectMeta: metav1.ObjectMeta{Name: "racer-" + string(uid), UID: uid, Namespace: "racer", CreationTimestamp: metav1.NewTime(time.Unix(created, 0)), OwnerReferences: []metav1.OwnerReference{{APIVersion: "apps/v1", Kind: "DaemonSet", Name: DataplaneDaemonSetName, UID: testDaemonSetUID, Controller: &controller}}}, Spec: corev1.PodSpec{NodeName: "node-a"}, Status: corev1.PodStatus{PodIP: ip}} +} + +func integrationInstallation(t *testing.T, c client.Client, namespace string) *Application { + t.Helper() + cfg := testConfig(t) + + cfg.Namespace = namespace + for _, obj := range []client.Object{&corev1.Namespace{ObjectMeta: metav1.ObjectMeta{Name: namespace}}, &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: namespace, Name: cfg.InstallationConfigMapName}, Data: map[string]string{"cluster": string(cfg.Cluster), "version_configmap": cfg.VersionConfigMapName, "state": "fresh", markerInitializationProtocol: stagedInitialization}}} { + if err := c.Create(t.Context(), obj); err != nil { + t.Fatal(err) + } + } + + return Assemble(cfg, c, c) +} + +type servingFixture struct { + a *Application + request wire.BootstrapRequest + ctx context.Context +} + +func newServingFixture(t *testing.T) *servingFixture { + t.Helper() + + node := memberNode() + node.Name = "worker" + pod := memberPod("pod-uid", 1, "192.0.2.1") + pod.Spec.NodeName = node.Name + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: DataplaneDaemonSetName, UID: testDaemonSetUID}} + r := initializedTopology(t, &node, &pod, ds) + a := Assemble(r.Config, r.Client, r.APIReader) + runKeys(t, a.Keyring) + reconcileTopology(t, a.Topology, t.Context()) + + _, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + csr, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{}, key) + if err != nil { + t.Fatal(err) + } + + return &servingFixture{a: a, ctx: t.Context(), request: wire.BootstrapRequest{SchemaVersion: 1, Cluster: r.Config.Cluster, Enrollment: testOtherUID, CSRDER: csr, Shares: wire.DefaultShares}} +} + +func trustReady(trust *Trust) bool { + _, err := trust.pool() + return err == nil +} + +func rejectWrites(t *testing.T, base client.WithWatch) client.WithWatch { + t.Helper() + + return interceptor.NewClient(base, interceptor.Funcs{ + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + t.Fatal("unexpected Create of committed authority") + return nil + }, + Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + t.Fatal("unexpected Update of committed authority") + return nil + }, + }) +} + +func BuildCatalog(caches []racerv1.ClusterCache) ([]wire.CacheDefinition, error) { + return members.BuildCatalog(caches) +} + +func TestEnvtestAuthority(t *testing.T) { + assets := os.Getenv("KUBEBUILDER_ASSETS") + if assets == "" { + t.Skip("set KUBEBUILDER_ASSETS to run the real API-server integration suite") + } + + scheme := runtime.NewScheme() + for _, add := range []func(*runtime.Scheme) error{clientgoscheme.AddToScheme, racerv1.AddToScheme} { + if err := add(scheme); err != nil { + t.Fatal(err) + } + } + + environment := &envtest.Environment{BinaryAssetsDirectory: assets, CRDDirectoryPaths: []string{"../../../api/racer/v1alpha1/crd"}, ErrorIfCRDPathMissing: true} + + rc, err := environment.Start() + if err != nil { + t.Fatal(err) + } + + t.Cleanup(func() { + if err := environment.Stop(); err != nil { + t.Error(err) + } + }) + + c, err := client.NewWithWatch(rc, client.Options{Scheme: scheme}) + if err != nil { + t.Fatal(err) + } + + t.Run("staged-initialization", func(t *testing.T) { integrationStagedInitialization(t, c) }) + t.Run("catalog-capacity", func(t *testing.T) { integrationCatalogCapacity(t, c) }) +} + +func TestPublicKeyringRotationPinsWriteAdmission(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + cfg := Config{Cluster: "11111111-1111-4111-8111-111111111111", Namespace: "racer", DataplaneServiceAccount: "racer-dataplane", DaemonSetName: "racer-dataplane", CredentialsSecretName: "racer-credentials", VersionConfigMapName: "racer-version", InstallationConfigMapName: "racer-installation", SnapshotMaxAge: 5 * time.Second, Rotation: RotationPolicy{Interval: 24 * time.Hour, PrepareFor: time.Hour, RetainFor: 48 * time.Hour}} + + scheme := runtime.NewScheme() + require.NoError(t, corev1.AddToScheme(scheme)) + require.NoError(t, racerv1.AddToScheme(scheme)) + + marker := &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName, UID: "installation"}, Data: map[string]string{"cluster": string(cfg.Cluster), "version_configmap": cfg.VersionConfigMapName, "state": "fresh", markerInitializationProtocol: stagedInitialization}} + c := stagedFakeClient(fake.NewClientBuilder().WithScheme(scheme).WithObjects(marker).Build()) + now := time.Now() + + a := New(cfg, Dependencies{Reader: c, Writer: c, Now: func() time.Time { return now }}) + require.NoError(t, a.Recover(t.Context(), c)) + _, err := a.ReconcileCredentials(t.Context()) + require.NoError(t, err) + + old, err := a.Keyring() + require.NoError(t, err) + + admitted, stop, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stop() + + deadline, _ := admitted.Context().Deadline() + + var before bytes.Buffer + + _, err = old.Response().WriteTo(t.Context(), admitted, &before) + require.NoError(t, err) + + time.Sleep(3 * time.Second) + + now = now.Add(cfg.Rotation.Interval - cfg.Rotation.PrepareFor) + + _, err = a.ReconcileCredentials(t.Context()) + require.NoError(t, err) + + current, err := a.Keyring() + require.NoError(t, err) + require.Greater(t, current.Generation(), old.Generation(), "ordinary rotation did not advance bundle") + + fresh, stopFresh, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stopFresh() + + _, err = old.Response().WriteTo(t.Context(), fresh, io.Discard) + require.ErrorIs(t, err, wire.Forbidden, "superseded handle borrowed fresh admission") + _, err = current.Response().WriteTo(t.Context(), fresh, io.Discard) + require.NoError(t, err) + + var after bytes.Buffer + + _, err = old.Response().WriteTo(t.Context(), admitted, &after) + require.NoError(t, err, "rotation revoked admitted write") + require.Equal(t, before.Bytes(), after.Bytes(), "admitted encoding changed") + + got, _ := admitted.Context().Deadline() + require.Equal(t, deadline, got, "rotation extended admitted deadline") + + time.Sleep(2 * time.Second) + + _, err = old.Response().WriteTo(t.Context(), admitted, io.Discard) + require.ErrorIs(t, err, context.DeadlineExceeded, "old write outlived pinned freshness") + }) +} + +func TestPublicAuthorityScaffoldAndOpaqueValues(t *testing.T) { + a := New(Config{}, Dependencies{}) + if !errors.Is(a.TrustReady(), wire.Unavailable) || !errors.Is(a.PublicationReady(), wire.Unavailable) { + t.Fatal("constructor granted authority") + } + + if !errors.Is(a.Recover(t.Context(), nil), wire.InvalidRequest) { + t.Fatal("zero configuration reached I/O") + } + + if _, err := a.Issue(t.Context(), NodeIdentity{}, wire.BootstrapRequest{}); !errors.Is(err, wire.Unauthenticated) { + t.Fatal("zero identity accepted", err) + } + + if _, err := a.Wait(t.Context(), NodeIdentity{}, nil); !errors.Is(err, wire.Unauthenticated) { + t.Fatal("zero poll identity accepted", err) + } + + var handle PublicationHandle + if _, _, err := handle.Admit(t.Context()); !errors.Is(err, wire.Unavailable) { + t.Fatal("zero publication handle accepted", err) + } + + if _, err := handle.ForBase(0, "").WriteTo(context.Background(), nil, io.Discard); !errors.Is(err, wire.Forbidden) { + t.Fatal("unguarded response accepted", err) + } + + for _, value := range []any{NodeIdentity{}, ReplicaIdentity{}, PublicationHandle{}, KeyringHandle{}, Response{}, Admission{}, *a} { + typeOf := reflect.TypeOf(value) + for i := range typeOf.NumField() { + if typeOf.Field(i).IsExported() { + t.Fatalf("%s exposes mutable field %s", typeOf.Name(), typeOf.Field(i).Name) + } + } + } +} + +func TestIdentityAndServingHandleProvenance(t *testing.T) { + one, two := newServingFixture(t), newServingFixture(t) + + a, b := one.a.authority, two.a.authority + _, err := a.Issue(t.Context(), NodeIdentity{owner: a, cluster: a.config.Cluster, node: testNodeUID, expires: time.Now().Add(time.Hour)}, one.request) + require.ErrorIs(t, err, wire.Unauthenticated, "certificate identity cannot authorize token-only issuance") + + for _, identity := range []NodeIdentity{{}, {owner: b, cluster: a.config.Cluster, node: testNodeUID, expires: time.Now().Add(time.Hour)}} { + _, err := a.Issue(t.Context(), identity, one.request) + require.ErrorIs(t, err, wire.Unauthenticated) + _, err = a.Wait(t.Context(), identity, nil) + require.ErrorIs(t, err, wire.Unauthenticated) + } + + p, err := a.Current() + require.NoError(t, err) + other, stopOther, err := b.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stopOther() + + _, _, err = p.AdmitWithTrust(t.Context(), other) + require.ErrorIs(t, err, wire.Forbidden) + _, err = p.ForBase(0, "").WriteTo(t.Context(), other, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + + var zero PublicationHandle + + _, _, err = zero.Admit(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + _, err = zero.ForBase(0, "").WriteTo(t.Context(), nil, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + keyring, err := a.Keyring() + require.NoError(t, err) + _, err = keyring.Response().WriteTo(t.Context(), other, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + _, err = keyring.Response().WriteTo(t.Context(), nil, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + + var empty KeyringHandle + + _, err = empty.Response().WriteTo(t.Context(), other, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) +} + +func TestKeyringHandleCannotBorrowRecoveredTrust(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + legacy, stopLegacy, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stopLegacy() + + old, err := a.Keyring() + require.NoError(t, err) + guard, stop, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stop() + + a.trust.invalidate() + require.ErrorIs(t, legacy.Check(t.Context()), context.Canceled) + require.NoError(t, a.Observe(t.Context())) + require.ErrorIs(t, guard.Check(t.Context()), context.Canceled) + fresh, stopFresh, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stopFresh() + + _, err = old.Response().WriteTo(t.Context(), fresh, io.Discard) + require.ErrorIs(t, err, wire.Forbidden) + current, err := a.Keyring() + require.NoError(t, err) + _, err = current.Response().WriteTo(t.Context(), fresh, io.Discard) + require.NoError(t, err) +} + +func TestAnnotationHintsDoNotAliasNestedHistory(t *testing.T) { + numa := uint32(3) + history := AcceptedMembers{testNodeUID: {Node: testNodeUID, RDMANICs: []wire.RDMANIC{{Device: "mlx5_0", Port: 1, NUMANode: &numa}}}} + hints := cloneAccepted(history) + *hints[testNodeUID].RDMANICs[0].NUMANode = 99 + require.EqualValues(t, 3, *history[testNodeUID].RDMANICs[0].NUMANode) +} + +func TestAuthorityConstructionFreezesAuthenticationAndIssuance(t *testing.T) { + f := newServingFixture(t) + cfg := f.a.Topology.Config + cfg.CertificateLifetime = 2 * time.Minute + a := New(cfg, Dependencies{Reader: f.a.Topology.APIReader, Writer: f.a.Topology.Client}) + cfg.Cluster = "" + cfg.CertificateLifetime = time.Second + + require.Equal(t, f.a.Topology.Config.Cluster, a.bootstrap.Config.Cluster) + require.Equal(t, 2*time.Minute, a.bootstrap.Issuer.Config.CertificateLifetime) + identity := NodeIdentity{owner: a, bearer: true, cluster: a.config.Cluster, node: testNodeUID, expires: time.Now().Add(time.Hour)} + encoded, err := a.Issue(t.Context(), identity, f.request) + require.NoError(t, err) + response, err := wire.DecodeBootstrapResponse(bytes.NewReader(encoded)) + require.NoError(t, err) + leaf, err := x509.ParseCertificate(response.CertificateChain[0]) + require.NoError(t, err) + require.Equal(t, 3*time.Minute, leaf.NotAfter.Sub(leaf.NotBefore)) +} + +func TestCatalogGateCanceledAcquisition(t *testing.T) { + for _, held := range []bool{false, true} { + t.Run(map[bool]string{false: "available", true: "held"}[held], func(t *testing.T) { + gate := newCatalogGate() + if held { + if err := gate.Acquire(t.Context()); err != nil { + t.Fatal(err) + } + } + + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + for range 100 { + if err := gate.Acquire(ctx); !errors.Is(err, context.Canceled) { + t.Fatalf("already canceled acquisition: %v", err) + } + } + + if held { + gate.Release() + } + + live, stop := context.WithTimeout(t.Context(), time.Second) + defer stop() + + if err := gate.Acquire(live); err != nil { + t.Fatalf("canceled acquisition consumed the gate: %v", err) + } + + gate.Release() + }) + } +} + +func TestCatalogGateSerializesCallers(t *testing.T) { + gate := newCatalogGate() + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + + var wg sync.WaitGroup + // Deliberately non-atomic: the gate must protect each read/modify/write. + count := 0 + + for range 16 { + wg.Go(func() { + for range 100 { + if err := gate.Acquire(ctx); err != nil { + t.Errorf("acquire: %v", err) + return + } + + count++ + + gate.Release() + } + }) + } + + wg.Wait() + + if count != 1600 { + t.Fatalf("lost serialized updates: %d", count) + } +} + +func capacityCaches(count int) []racerv1.ClusterCache { + caches := make([]racerv1.ClusterCache, count) + for i := range caches { + caches[i] = catalogCache(fmt.Sprintf("capacity-%d", i), types.UID(fmt.Sprintf("%08x-0000-0000-0000-000000000000", i))) + } + + return caches +} + +func TestCatalogCapacityBoundary(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + _, b, _, _ := keyState(t, r) + + capacity, err := catalogCapacity(r.Config, b) + if err != nil { + t.Fatal(err) + } + + t.Logf("default admitted maximum: %d", capacity) + // Check the conservative byte envelope independently: max generation, four + // roots, two generations of both purposes, and the longest state. + for _, count := range []int{capacity, capacity + 1} { + keys := []map[string]any{} + + for _, cache := range capacityCaches(count) { + for range 2 { + for _, purpose := range []wire.KeyPurpose{wire.PageKey, wire.OriginCredentialsKey} { + keys = append(keys, map[string]any{"cache": cache.UID, "id": make([]byte, 16), "purpose": purpose, "state": wire.PreparedKey, "material": make([]byte, 32)}) + } + } + } + + var out bytes.Buffer + + err := json.NewEncoder(&out).Encode(map[string]any{"schema_version": wire.SchemaVersion, "cluster": r.Config.Cluster, "generation": fmt.Sprint(uint64(math.MaxUint64)), "peer_trust_roots": [][]byte{make([]byte, reservedRootBytes), make([]byte, reservedRootBytes), make([]byte, reservedRootBytes), make([]byte, reservedRootBytes)}, "cache_keys": keys}) + if err != nil || (out.Len() <= wire.MaxBundleBytes) != (count == capacity) { + t.Fatalf("capacity=%d count=%d bytes=%d: %v", capacity, count, out.Len(), err) + } + } + // Even an empty catalog cannot make an unbounded number of roots fit. Reject + // pathological policy before consuming the one-way initialization claim. + r.Config.Rotation.Interval, r.Config.Rotation.PrepareFor = time.Nanosecond, time.Nanosecond + if _, err := catalogCapacity(r.Config, b); !errors.Is(err, wire.TooLarge) { + t.Fatalf("unbounded trust reserve: %v", err) + } +} + +func TestCatalogAdmissionRotationCycles(t *testing.T) { + for _, policy := range []RotationPolicy{ + {Interval: 24 * time.Hour, PrepareFor: time.Hour, RetainFor: 48 * time.Hour}, + {Interval: 6 * time.Hour, PrepareFor: time.Hour, RetainFor: 48 * time.Hour}, + {Interval: 12 * time.Hour, PrepareFor: 12 * time.Hour, RetainFor: 48 * time.Hour}, + {Interval: 7 * 24 * time.Hour, PrepareFor: time.Hour, RetainFor: 24 * time.Hour}, + } { + t.Run(fmt.Sprint(policy), func(t *testing.T) { + r, now := testKeyring(t) + r.Config.Rotation = policy + runKeys(t, r) + shared, b, _, _ := keyState(t, r) + + capacity, err := catalogCapacity(r.Config, b) + if err != nil { + t.Fatal(err) + } + + for _, cache := range capacityCaches(capacity - 1) { + if err := r.Create(t.Context(), &cache); err != nil { + t.Fatal(err) + } + } + // Cross a decimal-width boundary immediately and continue through + // enough cycles to reach the steady-state retirement high watermark. + b.Generation = 99 + + shared.Data["bundle.json"], _ = wire.EncodeBundle(b) + if err := r.Update(t.Context(), shared); err != nil { + t.Fatal(err) + } + + runKeys(t, r) + + maxRoots := exerciseCapacityRotations(t, r, now, capacity) + + if policy.Interval == 6*time.Hour && maxRoots < 8 { + t.Fatalf("did not exercise multiple retiring generations: %d", maxRoots) + } + }) + } +} + +func exerciseCapacityRotations(t *testing.T, r *credentialsFixture, now *time.Time, capacity int) int { + t.Helper() + + maxRoots := 0 + + for range 30 { + _, before, previous, _ := keyState(t, r) + *now = previous.nextTransition() + + runKeys(t, r) + _, after, state, _ := keyState(t, r) + + maxRoots = max(maxRoots, len(after.PeerTrustRoots)) + if len(keyedCaches(after)) != capacity || !trustReady(r.Trust) || after.Generation <= before.Generation { + t.Fatal("rotation at capacity lost admission, readiness, or progress") + } + + for id, deadline := range previous.Retiring { + if now.Before(deadline) && !state.Retiring[id].Equal(deadline) { + t.Fatal("capacity shortened retirement") + } + } + } + + return maxRoots +} + +func TestCatalogAdmissionGrowthRemovalAndRestart(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + a := Assemble(r.Config, r.Client, r.APIReader) + first := reconcileTopology(t, a.Topology, t.Context()) + _, initial, _, _ := keyState(t, r) + + capacity, err := catalogCapacity(r.Config, initial) + if err != nil { + t.Fatal(err) + } + + caches := capacityCaches(capacity + 1) + for _, cache := range caches { + require.NoError(t, r.Create(t.Context(), &cache)) + } + + if got := reconcileTopology(t, a.Topology, t.Context()); got != first { + t.Fatal("published growth before its keys were committed") + } + + var logs strings.Builder + + logger := funcr.New(func(_, msg string) { logs.WriteString(msg) }, funcr.Options{}) + + ctx := ctrl.LoggerInto(t.Context(), logger) + if _, err := r.operations().ReconcileCredentials(ctx); err != nil || !trustReady(r.Trust) { + t.Fatalf("growth disabled healthy service: %v", err) + } + + if !strings.Contains(logs.String(), "rotation_capacity") || !strings.Contains(logs.String(), caches[capacity].Name) { + t.Fatal("capacity rejection was not observable") + } + + _, admitted, _, _ := keyState(t, r) + + ids := keyedCaches(admitted) + if !ids[wire.CacheID(testNodeUID)] || len(ids) != capacity || ids[wire.CacheID(caches[capacity-1].UID)] { + t.Fatal("growth displaced established UID or ignored sorted free-slot order") + } + + assertPublishedKeys(t, a.Topology, capacity) + // Removing a rejected candidate must neither change keys nor consume a + // publication sequence. A restart retains the admitted set from the Secret. + before := reconcileTopology(t, a.Topology, t.Context()) + require.NoError(t, r.Delete(t.Context(), &caches[capacity])) + + r = Assemble(r.Config, r.Client, r.APIReader).Keyring + runKeys(t, r) + + if after := reconcileTopology(t, a.Topology, t.Context()); after != before { + t.Fatal("rejected deletion changed publication") + } + + if err := r.Delete(t.Context(), &caches[0]); err != nil { + t.Fatal(err) + } + + assertPublishedKeys(t, a.Topology, capacity-1) + runKeys(t, r) + assertPublishedKeys(t, a.Topology, capacity) + + _, replaced, _, _ := keyState(t, r) + if keyedCaches(replaced)[wire.CacheID(caches[0].UID)] || !keyedCaches(replaced)[wire.CacheID(caches[capacity-1].UID)] { + t.Fatal("deletion did not admit next waiting UID") + } + + assertMissingCredentialsSuspendTopology(t, r, a.Topology) +} + +func assertMissingCredentialsSuspendTopology(t *testing.T, r *credentialsFixture, topology *topologyFixture) { + t.Helper() + // Missing/corrupt durable credentials remain fail-closed in both controllers. + shared := &corev1.Secret{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: r.Config.CredentialsSecretName}} + if err := r.Delete(t.Context(), shared); err != nil { + t.Fatal(err) + } + + if _, err := topology.operations().PublishTopology(t.Context(), topology.observeTopology); err == nil { + t.Fatal("missing credentials accepted by topology") + } + + if _, err := topology.Publications.Current(); !errors.Is(err, wire.Unavailable) { + t.Fatalf("missing credentials did not suspend publication: %v", err) + } +} + +func assertPublishedKeys(t *testing.T, r *topologyFixture, count int) { + t.Helper() + p := reconcileTopology(t, r, t.Context()) + + v, err := wire.DecodePublication(strings.NewReader(p.encoded)) + if err != nil || len(v.Caches) != count { + t.Fatalf("published caches=%d, want %d: %v", len(v.Caches), count, err) + } + + _, b, _, _ := keyState(t, Assemble(r.Config, r.Client, r.APIReader).Keyring) + for _, cache := range v.Caches { + if !keyedCaches(b)[cache.ID] { + t.Fatal("published cache without both active keys") + } + } +} + +func TestCatalogAdmissionDeterministicColdStart(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + _, b, _, _ := keyState(t, r) + b.CacheKeys = nil + + capacity, err := catalogCapacity(r.Config, b) + if err != nil { + t.Fatal(err) + } + + input := capacityCaches(capacity + 2) + slices.Reverse(input) + + catalog, err := BuildCatalog(input) + if err != nil { + t.Fatal(err) + } + + got, err := admitCatalog(context.Background(), r.Config, catalog, b) + if err != nil || !slices.Equal(got, catalog[:capacity]) { + t.Fatalf("cold admission is not a sorted UID prefix: %v", err) + } +} + +func TestCatalogCapacityRejectsBeforeInitializationClaim(t *testing.T) { + r, _ := testKeyring(t) + + r.Config.Rotation.Interval, r.Config.Rotation.PrepareFor = time.Nanosecond, time.Nanosecond + if _, err := r.operations().ReconcileCredentials(t.Context()); !errors.Is(err, wire.TooLarge) { + t.Fatalf("impossible root reserve: %v", err) + } + + cm, _, err := readVersion(t.Context(), r.APIReader, r.Config) + if err != nil || cm.Annotations[credentialClaim] != "" { + t.Fatalf("impossible policy consumed credential claim: %v", err) + } +} + +func TestCatalogAdmissionLegacyOvercommitDoesNotEvict(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + shared, b, state, _ := keyState(t, r) + + capacity, err := catalogCapacity(r.Config, b) + if err != nil { + t.Fatal(err) + } + + caches := capacityCaches(capacity) + for _, cache := range caches { + if err := r.Create(t.Context(), &cache); err != nil { + t.Fatal(err) + } + } + + caches = append(caches, catalogCache("cache", testNodeUID)) + + catalog, err := BuildCatalog(caches) + if err != nil { + t.Fatal(err) + } + // Model the older controller's active-only admission without removing the + // planner's independent final wire-size check. + b, state, _, err = planRotation(r.Config.Rotation, b, state, catalog, *now, nextGeneration(b.Generation)) + if err != nil { + t.Fatal(err) + } + + b.Generation++ + + shared.Data["bundle.json"], err = wire.EncodeBundle(b) + if err != nil { + t.Fatal(err) + } + + shared.Data["rotation.json"], _ = json.Marshal(state) + if err := r.Update(t.Context(), shared); err != nil { + t.Fatal(err) + } + + if _, err := r.operations().ReconcileCredentials(t.Context()); !errors.Is(err, wire.TooLarge) || trustReady(r.Trust) { + t.Fatalf("legacy overcommit silently accepted: %v", err) + } + + after, preserved, _, _ := keyState(t, r) + if after.ResourceVersion != shared.ResourceVersion || len(keyedCaches(preserved)) != capacity+1 { + t.Fatal("legacy overcommit evicted durable credentials") + } +} + +func TestCatalogAdmissionSerializesPublicationAndPruning(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + a := Assemble(r.Config, r.Client, r.APIReader) + read := make(chan struct{}) + proceed := make(chan struct{}) + a.Topology.APIReader = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + if key.Name == r.Config.CredentialsSecretName { + close(read) + <-proceed + } + + return err + }, + }) + done := make(chan error, 1) + + go func() { + _, err := a.Topology.operations().PublishTopology(t.Context(), a.Topology.observeTopology) + done <- err + }() + + <-read + // The keyring must be excluded for the entire read/commit/install window. + ctx, cancel := context.WithTimeout(t.Context(), 20*time.Millisecond) + defer cancel() + + err := a.authority.gate.Acquire(ctx) + if err == nil { + a.authority.gate.Release() + close(proceed) + <-done + t.Fatal("keyring can prune an in-progress topology candidate") + } + + close(proceed) + + if err := <-done; err != nil { + t.Fatal(err) + } + + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("unexpected gate wait error: %v", err) + } + + ctx, cancel = context.WithTimeout(t.Context(), time.Second) + defer cancel() + + if err := a.authority.gate.Acquire(ctx); err != nil { + t.Fatalf("publication did not release admission gate: %v", err) + } + + a.authority.gate.Release() +} + +func integrationCatalogCapacity(t *testing.T, c client.Client) { + a := integrationInstallation(t, c, "catalog-capacity") + // Use a root-heavy valid policy to exercise real admission with a small + // catalog. Unit tests above run the default maximum through repeated cycles. + // Reserve generations by activation interval, not interval plus preparation. + // 381 roots leave room for a small catalog with two symmetric generations; + // 382 roots would leave no capacity for a complete active/prepared key pair. + a.Keyring.Config.Rotation = RotationPolicy{Interval: time.Hour, PrepareFor: time.Hour, RetainFor: 379 * time.Hour} + + a.Topology.Config = a.Keyring.Config + if err := a.Recover(t.Context(), a.Topology.Client); err != nil { + t.Fatal(err) + } + + runKeys(t, a.Keyring) + _, b, _, _ := keyState(t, a.Keyring) + + capacity, err := catalogCapacity(a.Keyring.Config, b) + if err != nil || capacity < 1 || capacity > 5 { + t.Fatalf("integration capacity: %d %v", capacity, err) + } + + for i := range capacity + 1 { + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: fmt.Sprintf("capacity-real-%d", i)}} + if err := c.Create(t.Context(), cache); err != nil { + t.Fatal(err) + } + + t.Cleanup(func() { + if err := c.Delete(context.Background(), cache); err != nil { + t.Error(err) + } + }) + } + + assertPublishedKeys(t, a.Topology, 0) + runKeys(t, a.Keyring) + assertPublishedKeys(t, a.Topology, capacity) + + if !trustReady(a.Server.Trust) { + t.Fatal("API-backed capacity rejection withdrew readiness") + } +} + +func testIssuer(r *credentialsFixture) *Issuer { + return &Issuer{APIReader: r.APIReader, Config: r.Config.effective(), Trust: r.Trust, CatalogGate: r.authority.gate, Now: r.Now} +} + +// TrustRoots is a test adapter for authoritative signing observations. Production +// serving uses local Trust; only issuance and reconciliation read durable roots. +func (i *Issuer) TrustRoots(ctx context.Context) (*x509.CertPool, error) { + state, err := i.loadSigning(ctx, credentialTime(i.Now)) + if err != nil { + return nil, err + } + + return state.roots, nil +} + +func issuanceRequest(t *testing.T, r *credentialsFixture) (NodeIdentity, wire.BootstrapRequest, ed25519.PublicKey) { + t.Helper() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + csr, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{Subject: pkix.Name{CommonName: "untrusted"}, DNSNames: []string{"attacker"}, URIs: []*url.URL{{Scheme: "spiffe", Host: "attacker", Path: "/node/attacker"}}}, key) + require.NoError(t, err) + + return NodeIdentity{cluster: r.Config.Cluster, node: wire.NodeID(testNodeUID), expires: r.now().Add(time.Hour)}, wire.BootstrapRequest{SchemaVersion: wire.SchemaVersion, Cluster: r.Config.Cluster, Enrollment: wire.EnrollmentID(testOtherUID), CSRDER: csr, Shares: wire.DefaultShares}, pub +} + +func decodeIssuedResponse(t *testing.T, encoded []byte) wire.BootstrapResponse { + t.Helper() + require.NotEmpty(t, encoded) + require.LessOrEqual(t, len(encoded), wire.MaxBootstrapBytes) + response, err := wire.DecodeBootstrapResponse(bytes.NewReader(encoded)) + require.NoError(t, err) + + return response +} + +func TestIssuerCertificateContractAndTrustRotation(t *testing.T) { + r, now := testKeyring(t) + issuer := testIssuer(r) + runKeys(t, r) + identity, request, pub := issuanceRequest(t, r) + encoded, err := issuer.Issue(context.Background(), identity, request) + require.NoError(t, err) + response := decodeIssuedResponse(t, encoded) + cert, err := x509.ParseCertificate(response.CertificateChain[0]) + require.NoError(t, err) + require.Equal(t, identity.Node(), response.Node) + require.Equal(t, request.Cluster, response.Cluster) + require.Equal(t, request.Enrollment, response.Enrollment) + require.Len(t, response.CertificateChain, 2) + require.False(t, cert.IsCA) + require.Empty(t, cert.Subject.CommonName) + require.Empty(t, cert.DNSNames) + require.Len(t, cert.URIs, 1) + require.Equal(t, "spiffe://"+string(identity.Cluster())+"/node/"+string(identity.Node()), cert.URIs[0].String()) + require.Equal(t, x509.KeyUsageDigitalSignature, cert.KeyUsage) + require.True(t, cert.NotAfter.Equal(now.Add(wire.CertificateLifetime))) + require.True(t, cert.NotBefore.Equal(now.Add(-certificateClockSkew))) + require.Equal(t, pub, cert.PublicKey.(ed25519.PublicKey)) + + roots, err := issuer.TrustRoots(context.Background()) + require.NoError(t, err) + _, err = cert.Verify(x509.VerifyOptions{Roots: roots, CurrentTime: *now, KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}}) + require.NoError(t, err) + _, err = cert.Verify(x509.VerifyOptions{Roots: roots, CurrentTime: *now, KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}}) + require.Error(t, err, "node can act as HTTPS server") + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + + identity.expires = now.Add(time.Hour) + encoded, err = issuer.Issue(context.Background(), identity, request) + require.NoError(t, err) + staged := decodeIssuedResponse(t, encoded) + require.Equal(t, response.CertificateChain[1], staged.CertificateChain[1], "prepared issuer signed early") + _, _, preparation, _ := keyState(t, r) + *now = preparation.ActivateAt + + runKeys(t, r) + + identity.expires = now.Add(time.Hour) + encoded, err = issuer.Issue(context.Background(), identity, request) + require.NoError(t, err) + active := decodeIssuedResponse(t, encoded) + require.NotEqual(t, response.CertificateChain[1], active.CertificateChain[1], "new issuer not activated") + + roots, err = issuer.TrustRoots(context.Background()) + require.NoError(t, err) + oldLeaf, err := x509.ParseCertificate(staged.CertificateChain[0]) + require.NoError(t, err) + _, err = oldLeaf.Verify(x509.VerifyOptions{Roots: roots, CurrentTime: *now, KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}}) + require.NoError(t, err, "old leaf lost overlap") + + *now = now.Add(r.Config.Rotation.RetainFor) + runKeys(t, r) + + roots, err = issuer.TrustRoots(context.Background()) + require.NoError(t, err) + // Use a time at which the old leaf was valid to isolate root removal. + _, err = oldLeaf.Verify(x509.VerifyOptions{Roots: roots, CurrentTime: oldLeaf.NotBefore.Add(time.Minute), KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}}) + require.Error(t, err, "retired root remains trusted") +} + +func TestIssuerRejectsUntrustedRequests(t *testing.T) { + r, _ := testKeyring(t) + issuer := testIssuer(r) + runKeys(t, r) + + identity, request, _ := issuanceRequest(t, r) + for _, scenario := range []string{"zero identity", "expired identity", "wrong cluster", "unsupported version", "bad enrollment", "malformed csr", "bad proof", "wrong algorithm", "oversized", "canceled"} { + t.Run(scenario, func(t *testing.T) { + id, req := identity, request + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + switch scenario { + case "zero identity": + id = NodeIdentity{} + case "expired identity": + id.expires = r.now() + case "wrong cluster": + req.Cluster = wire.ClusterID(testNodeUID) + case "unsupported version": + req.SchemaVersion++ + case "bad enrollment": + req.Enrollment = "bad" + case "malformed csr": + req.CSRDER = []byte("invalid DER") + case "bad proof": + req.CSRDER = bytes.Clone(req.CSRDER) + req.CSRDER[len(req.CSRDER)-1] ^= 1 + case "wrong algorithm": + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + req.CSRDER, err = x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{}, key) + require.NoError(t, err) + case "oversized": + req.CSRDER = make([]byte, wire.MaxBootstrapBytes+1) + case "canceled": + cancel() + } + + response, err := issuer.Issue(ctx, id, req) + require.Error(t, err, "untrusted issuance accepted") + require.Nil(t, response) + }) + } +} + +func TestIssuerShortLifetimeAndRetirement(t *testing.T) { + r, now := testKeyring(t) + r.Config.CertificateLifetime = 2 * time.Minute + r.Config.Rotation = RotationPolicy{5 * time.Minute, 20 * time.Second, 2 * time.Minute} + issuer := testIssuer(r) + runKeys(t, r) + identity, request, _ := issuanceRequest(t, r) + encoded, err := issuer.Issue(context.Background(), identity, request) + require.NoError(t, err) + response := decodeIssuedResponse(t, encoded) + leaf, err := x509.ParseCertificate(response.CertificateChain[0]) + require.NoError(t, err) + require.True(t, leaf.NotAfter.Equal(now.Add(2*time.Minute))) + require.True(t, leaf.NotBefore.Equal(now.Add(-certificateClockSkew))) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + _, _, prepared, _ := keyState(t, r) + *now = prepared.ActivateAt + + runKeys(t, r) + + encoded, err = issuer.Issue(context.Background(), identity, request) + require.NoError(t, err) + renewed := decodeIssuedResponse(t, encoded) + require.NotEqual(t, response.CertificateChain[1], renewed.CertificateChain[1], "short rotation issuer activation") + + *now = now.Add(2 * time.Minute) + + runKeys(t, r) + _, bundle, _, material := keyState(t, r) + require.False(t, containsRoot(bundle, initial.ActiveIssuer)) + require.Len(t, material.Keys, 1, "short rotation did not retire old public/private issuer") +} + +func TestIssuerFullEncodedRequestBound(t *testing.T) { + r, _ := testKeyring(t) + issuer := testIssuer(r) + runKeys(t, r) + identity, request, _ := issuanceRequest(t, r) + _, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + + for _, size := range []int{47 * 1024, 49 * 1024} { + request.CSRDER, err = x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{Subject: pkix.Name{CommonName: strings.Repeat("x", size)}}, key) + require.NoError(t, err) + require.Less(t, len(request.CSRDER), wire.MaxBootstrapBytes, "fixture must fit the raw DER bound") + + encoded, err := issuer.Issue(context.Background(), identity, request) + if size == 49*1024 { + require.ErrorIs(t, err, wire.TooLarge) + require.Nil(t, encoded) + + continue + } + + require.NoError(t, err) + response := decodeIssuedResponse(t, encoded) + require.Equal(t, request.Enrollment, response.Enrollment) + require.Equal(t, identity.Node(), response.Node, "large valid request lost correlation") + } +} + +func TestIssuerConcurrentIssuanceAndReconciliation(t *testing.T) { + r, _ := testKeyring(t) + issuer := testIssuer(r) + runKeys(t, r) + identity, request, _ := issuanceRequest(t, r) + + var wg sync.WaitGroup + for range 16 { + wg.Go(func() { + for range 4 { + if _, err := issuer.Issue(context.Background(), identity, request); err != nil { + t.Error(err) + } + + if _, err := issuer.TrustRoots(context.Background()); err != nil { + t.Error(err) + } + } + }) + } + + for range 4 { + runKeys(t, r) + } + + wg.Wait() +} + +func writeSigningCredentials(t *testing.T, r *credentialsFixture, b wire.KeyringBundle, s RotationState, m issuerMaterial) { + t.Helper() + + bundle, err := wire.EncodeBundle(b) + require.NoError(t, err) + rotation, err := json.Marshal(s) + require.NoError(t, err) + material, err := json.Marshal(m) + require.NoError(t, err) + + secret := &corev1.Secret{} + require.NoError(t, r.APIReader.Get(t.Context(), client.ObjectKey{Namespace: r.Config.Namespace, Name: r.Config.CredentialsSecretName}, secret)) + secret.Data = map[string][]byte{"issuer.json": material, "bundle.json": bundle, "rotation.json": rotation} + require.NoError(t, r.Update(t.Context(), secret)) +} + +func editSigningCertificate(t *testing.T, m signingMaterial, edit func(*x509.Certificate)) signingMaterial { + t.Helper() + + cert, key, err := parseSigning(m) + require.NoError(t, err) + edit(cert) + m.Certificate, err = x509.CreateCertificate(rand.Reader, cert, cert, key.Public(), key) + require.NoError(t, err) + + return m +} + +func TestSigningRejectsCorruptPrivateEntries(t *testing.T) { + for _, role := range []string{"active", "prepared", "retiring", "extra", "pending"} { + for _, corruption := range []string{"missing", "root binding", "certificate", "trailing certificate bytes", "private key", "key mismatch", "not CA", "constraints", "key usage", "self signature"} { + t.Run(role+"/"+corruption, func(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, b, s, m := keyState(t, r) + id := signingRole(t, r, role, &b, &s, m) + writeSigningCredentials(t, r, b, s, m) + _, err := loadSigning(t.Context(), r.APIReader, r.Config, *now) + require.NoError(t, err, "valid %s rejected", role) + bad := corruptSigning(t, r, corruption, m.Keys[id]) + delete(m.Keys, id) + + if corruption == "root binding" { + m.Keys["wrong fingerprint"] = bad + } else if corruption != "missing" { + rebindSigning(&b, &s, m, id, bad, corruption != "certificate" && corruption != "trailing certificate bytes") + } + + writeSigningCredentials(t, r, b, s, m) + state, err := loadSigning(t.Context(), r.APIReader, r.Config, *now) + require.ErrorIs(t, err, wire.Unavailable) + require.Nil(t, state.certificate) + require.Nil(t, state.key) + require.Nil(t, state.roots) + _, err = testIssuer(r).TrustRoots(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + _, err = r.Trust.pool() + require.Error(t, err, "observed corruption retained trust") + _, err = r.operations().ReconcileCredentials(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + }) + } + } +} + +func signingRole(t *testing.T, r *credentialsFixture, role string, b *wire.KeyringBundle, s *RotationState, m issuerMaterial) string { + t.Helper() + + if role == "active" { + return s.ActiveIssuer + } + + cert, key, err := generateIssuer(r.now(), r.Config) + require.NoError(t, err) + + id := rootID(cert) + + m.Keys[id] = signingMaterial{Certificate: cert, PrivateKey: key} + if role == "extra" || role == "pending" { + writeSigningCredentials(t, r, *b, *s, m) + _, err := loadSigning(t.Context(), r.APIReader, r.Config, r.now()) + require.ErrorIs(t, err, wire.Unavailable, "unpublished private material accepted") + } + + b.PeerTrustRoots = append(b.PeerTrustRoots, cert) + + if role == "prepared" { + s.PreparedIssuer = id + s.ActivateAt = s.NextRotation.Add(r.Config.Rotation.PrepareFor) + } else { + s.Retiring[id] = s.NextRotation.Add(time.Hour) + } + + return id +} + +func corruptSigning(t *testing.T, r *credentialsFixture, corruption string, bad signingMaterial) signingMaterial { + t.Helper() + + switch corruption { + case "certificate": + bad.Certificate = []byte("invalid DER") + case "trailing certificate bytes": + bad.Certificate = append(bytes.Clone(bad.Certificate), 0) + case "private key": + bad.PrivateKey = []byte("invalid PKCS8") + case "key mismatch": + _, key, err := generateIssuer(r.now(), r.Config) + require.NoError(t, err) + + bad.PrivateKey = key + case "not CA": + bad = editSigningCertificate(t, bad, func(c *x509.Certificate) { + c.IsCA, c.MaxPathLenZero, c.MaxPathLen = false, false, -1 + }) + case "constraints": + bad = editSigningCertificate(t, bad, func(c *x509.Certificate) { c.BasicConstraintsValid = false }) + case "key usage": + bad = editSigningCertificate(t, bad, func(c *x509.Certificate) { c.KeyUsage = x509.KeyUsageDigitalSignature }) + case "self signature": + bad.Certificate = bytes.Clone(bad.Certificate) + bad.Certificate[len(bad.Certificate)-1] ^= 1 + } + + return bad +} + +func rebindSigning(b *wire.KeyringBundle, s *RotationState, m issuerMaterial, oldID string, material signingMaterial, replaceRoot bool) { + id := rootID(material.Certificate) + + delete(m.Keys, oldID) + m.Keys[id] = material + + for i, root := range b.PeerTrustRoots { + if replaceRoot && rootID(root) == oldID { + b.PeerTrustRoots[i] = material.Certificate + } + } + + if s.ActiveIssuer == oldID { + s.ActiveIssuer = id + } + + if s.PreparedIssuer == oldID { + s.PreparedIssuer = id + } + + if at, ok := s.Retiring[oldID]; ok { + delete(s.Retiring, oldID) + s.Retiring[id] = at + } +} + +func TestSigningActiveLifetimeBoundary(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + _, _, s, m := keyState(t, r) + cert, _, err := parseSigning(m.Keys[s.ActiveIssuer]) + require.NoError(t, err) + + for _, tc := range []struct { + name string + at time.Time + valid bool + }{ + {"not yet valid", cert.NotBefore.Add(-time.Second), false}, + {"starts now", cert.NotBefore, true}, + {"full leaf lifetime", cert.NotAfter.Add(-r.Config.CertificateLifetime), true}, + {"short by one second", cert.NotAfter.Add(-r.Config.CertificateLifetime).Add(time.Second), false}, + {"expired", cert.NotAfter, false}, + } { + t.Run(tc.name, func(t *testing.T) { + state, err := loadSigning(t.Context(), r.APIReader, r.Config, tc.at) + if tc.valid { + require.NoError(t, err) + require.Equal(t, cert.Raw, state.certificate.Raw) + require.True(t, state.certificate.PublicKey.(ed25519.PublicKey).Equal(state.key.Public())) + } else { + require.ErrorIs(t, err, wire.Unavailable, "invalid active lifetime accepted") + } + }) + } +} + +func TestSigningPoolOnlyIncludesTimeValidPublishedRoots(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, b, s, m := keyState(t, r) + active, _, err := parseSigning(m.Keys[s.ActiveIssuer]) + require.NoError(t, err) + + want := x509.NewCertPool() + want.AddCert(active) + + for _, role := range []string{"starts now", "expires now", "future", "extra", "pending"} { + cert, key, err := generateIssuer(*now, r.Config) + require.NoError(t, err) + material := editSigningCertificate(t, signingMaterial{Certificate: cert, PrivateKey: key}, func(c *x509.Certificate) { + switch role { + case "starts now": + c.NotBefore = *now + case "expires now": + c.NotAfter = *now + case "future": + c.NotBefore = now.Add(time.Second) + } + }) + id := rootID(material.Certificate) + + m.Keys[id] = material + if role == "extra" || role == "pending" { + writeSigningCredentials(t, r, b, s, m) + _, err := loadSigning(t.Context(), r.APIReader, r.Config, *now) + require.ErrorIs(t, err, wire.Unavailable, "unpublished private material accepted") + delete(m.Keys, id) + + continue + } + + b.PeerTrustRoots = append(b.PeerTrustRoots, material.Certificate) + s.Retiring[id] = s.NextRotation.Add(time.Hour) + + if role == "starts now" { + root, _, err := parseSigning(material) + require.NoError(t, err) + want.AddCert(root) + } + } + + writeSigningCredentials(t, r, b, s, m) + state, err := loadSigning(t.Context(), r.APIReader, r.Config, *now) + require.NoError(t, err) + require.True(t, state.roots.Equal(want)) + require.Equal(t, active.Raw, state.certificate.Raw) + + for id := range s.Retiring { + s.Retiring[id] = time.Time{} + break + } + + writeSigningCredentials(t, r, b, s, m) + _, err = loadSigning(t.Context(), r.APIReader, r.Config, *now) + require.ErrorIs(t, err, wire.Unavailable, "zero retirement deadline accepted") +} + +func TestLeafClockSkewPreservesExpirationAndUsage(t *testing.T) { + r, now := testKeyring(t) + // Simulate an issuing leader ahead of a follower's clock. + *now = now.Add(30 * time.Second) + + runKeys(t, r) + identity, request, _ := issuanceRequest(t, r) + encoded, err := testIssuer(r).Issue(t.Context(), identity, request) + require.NoError(t, err) + response := decodeIssuedResponse(t, encoded) + leaf, err := x509.ParseCertificate(response.CertificateChain[0]) + require.NoError(t, err) + root, err := x509.ParseCertificate(response.CertificateChain[1]) + require.NoError(t, err) + require.True(t, leaf.NotBefore.Equal(now.Add(-time.Minute))) + require.True(t, leaf.NotAfter.Equal(now.Add(r.Config.CertificateLifetime))) + + roots := x509.NewCertPool() + roots.AddCert(root) + + for _, tc := range []struct { + name string + at time.Time + usage x509.ExtKeyUsage + valid bool + }{ + {"follower behind", now.Add(-30 * time.Second), x509.ExtKeyUsageClientAuth, true}, + {"skew boundary", now.Add(-time.Minute), x509.ExtKeyUsageClientAuth, true}, + {"excess skew", now.Add(-time.Minute - time.Second), x509.ExtKeyUsageClientAuth, false}, + {"before expiry", leaf.NotAfter.Add(-time.Second), x509.ExtKeyUsageClientAuth, true}, + {"expired", leaf.NotAfter.Add(time.Second), x509.ExtKeyUsageClientAuth, false}, + {"server usage", *now, x509.ExtKeyUsageServerAuth, false}, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := leaf.Verify(x509.VerifyOptions{Roots: roots, CurrentTime: tc.at, KeyUsages: []x509.ExtKeyUsage{tc.usage}}) + require.Equal(t, tc.valid, err == nil, "verification: %v", err) + }) + } + + state := &tls.ConnectionState{HandshakeComplete: true, PeerCertificates: []*x509.Certificate{leaf, root}, VerifiedChains: [][]*x509.Certificate{{leaf, root}}} + _, err = AuthenticateCertificate(t.Context(), r.Trust, r.Config, state) + require.NoError(t, err, "follower rejected skewed leader's certificate") + // Expired leaves still fail each authorization, even on a verified connection. + leaf.NotAfter = time.Now().Add(-time.Second) + _, err = AuthenticateCertificate(t.Context(), r.Trust, r.Config, state) + require.ErrorIs(t, err, wire.Unauthenticated, "expired certificate accepted") + // Advancing past signing capacity still fails closed, rather than clipping expiry. + *now = root.NotAfter.Add(-r.Config.CertificateLifetime + time.Second) + identity.expires = now.Add(time.Hour) + _, err = testIssuer(r).Issue(t.Context(), identity, request) + require.ErrorIs(t, err, wire.Unavailable, "insufficient issuer lifetime accepted") +} + +func authenticatedBootstrapFixture(t *testing.T) (*servingFixture, authv1.TokenReviewStatus, string) { + t.Helper() + f := newServingFixture(t) + cfg, c := f.a.authority.config, f.a.Topology.Client + + var pods corev1.PodList + require.NoError(t, c.List(t.Context(), &pods)) + require.Len(t, pods.Items, 1) + pod := &pods.Items[0] + pod.Spec.ServiceAccountName = cfg.DataplaneServiceAccount + require.NoError(t, c.Update(t.Context(), pod)) + + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.DataplaneServiceAccount, UID: "service-account"}} + require.NoError(t, c.Create(t.Context(), sa)) + status := authv1.TokenReviewStatus{Authenticated: true, Audiences: []string{wire.TokenAudience}, User: authv1.UserInfo{ + Username: "system:serviceaccount:" + cfg.Namespace + ":" + cfg.DataplaneServiceAccount, UID: string(sa.UID), + Extra: map[string]authv1.ExtraValue{ + "authentication.kubernetes.io/pod-name": {pod.Name}, + "authentication.kubernetes.io/pod-uid": {string(pod.UID)}, + "authentication.kubernetes.io/node-name": {pod.Spec.NodeName}, + "authentication.kubernetes.io/node-uid": {testNodeUID}, + }, + }} + payload := fmt.Sprintf(`{"exp":%d}`, time.Now().Add(time.Hour).Unix()) + + return f, status, "header." + base64.RawURLEncoding.EncodeToString([]byte(payload)) + ".signature" +} + +func TestBootstrapTokenBindingAndEnrollment(t *testing.T) { + for _, tc := range []struct { + name string + edit func(*authv1.TokenReviewStatus, *string) + want error + }{ + {"valid", func(*authv1.TokenReviewStatus, *string) {}, nil}, + {"unauthenticated", func(s *authv1.TokenReviewStatus, _ *string) { s.Authenticated = false }, wire.Unauthenticated}, + {"review error", func(s *authv1.TokenReviewStatus, _ *string) { s.Error = "denied" }, wire.Unauthenticated}, + {"audience", func(s *authv1.TokenReviewStatus, _ *string) { s.Audiences = nil }, wire.Unauthenticated}, + {"account", func(s *authv1.TokenReviewStatus, _ *string) { s.User.Username = "other" }, wire.Forbidden}, + {"account uid", func(s *authv1.TokenReviewStatus, _ *string) { s.User.UID = "other" }, wire.Forbidden}, + {"missing uid", func(s *authv1.TokenReviewStatus, _ *string) { s.User.UID = "" }, wire.Unauthenticated}, + {"pod uid", func(s *authv1.TokenReviewStatus, _ *string) { + s.User.Extra["authentication.kubernetes.io/pod-uid"] = authv1.ExtraValue{"other"} + }, wire.Forbidden}, + {"ambiguous pod", func(s *authv1.TokenReviewStatus, _ *string) { + s.User.Extra["authentication.kubernetes.io/pod-name"] = authv1.ExtraValue{"one", "two"} + }, wire.Unauthenticated}, + {"missing pod", func(s *authv1.TokenReviewStatus, _ *string) { + s.User.Extra["authentication.kubernetes.io/pod-name"] = authv1.ExtraValue{"missing"} + }, wire.Forbidden}, + {"node uid", func(s *authv1.TokenReviewStatus, _ *string) { + s.User.Extra["authentication.kubernetes.io/node-uid"] = authv1.ExtraValue{"other"} + }, wire.Forbidden}, + {"malformed token", func(_ *authv1.TokenReviewStatus, token *string) { *token = "invalid" }, wire.Unauthenticated}, + {"invalid base64", func(_ *authv1.TokenReviewStatus, token *string) { *token = "header.!.signature" }, wire.Unauthenticated}, + {"missing expiration", func(_ *authv1.TokenReviewStatus, token *string) { *token = "header.e30.signature" }, wire.Unauthenticated}, + {"expired", func(_ *authv1.TokenReviewStatus, token *string) { + *token = "header." + base64.RawURLEncoding.EncodeToString([]byte(`{"exp":1}`)) + ".signature" + }, wire.Unauthenticated}, + } { + t.Run(tc.name, func(t *testing.T) { + f, status, token := authenticatedBootstrapFixture(t) + tc.edit(&status, &token) + + a := f.a.authority + a.bootstrap.Client = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Create: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.CreateOption) error { + review := obj.(*authv1.TokenReview) + require.Equal(t, token, review.Spec.Token) + require.Equal(t, []string{wire.TokenAudience}, review.Spec.Audiences) + review.Status = status + + return nil + }}) + request := httptest.NewRequest("POST", "/bootstrap", nil) + request.Header.Set("Authorization", "Bearer "+token) + + identity, err := a.Authenticate(t.Context(), request) + if tc.want != nil { + require.ErrorIs(t, err, tc.want) + require.Equal(t, NodeIdentity{}, identity) + + return + } + + require.NoError(t, err) + require.Equal(t, wire.NodeID(testNodeUID), identity.Node()) + require.Equal(t, a.config.Cluster, identity.Cluster()) + require.True(t, identity.Expires().After(time.Now())) + encoded, hint, err := a.EnrollWithHint(t.Context(), request, f.request) + require.NoError(t, err) + require.Equal(t, identity.Node(), decodeIssuedResponse(t, encoded).Node) + require.Equal(t, types.UID(testNodeUID), hint.Node.UID) + require.Equal(t, f.request.Shares, hint.Shares) + require.Equal(t, identity.Expires(), hint.Expires) + }) + } +} + +func TestBootstrapBlockDevices(t *testing.T) { + for _, tc := range []struct { + name string + pattern string + want string + warning bool + }{ + {name: "absent"}, + {name: "empty"}, + {name: "valid", pattern: `^nvme-eui\.[0-9a-f]+$`, want: `^nvme-eui\.[0-9a-f]+$`}, + {name: "invalid", pattern: "[", warning: true}, + {name: "oversized", pattern: strings.Repeat("a", 1025), warning: true}, + {name: "byte limit", pattern: strings.Repeat("é", 513), warning: true}, + {name: "at limit", pattern: strings.Repeat("a", 1024), want: strings.Repeat("a", 1024)}, + } { + t.Run(tc.name, func(t *testing.T) { + f, status, token := authenticatedBootstrapFixture(t) + a := f.a.authority + a.bootstrap.Client = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Create: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.CreateOption) error { + obj.(*authv1.TokenReview).Status = status + return nil + }, + }) + nodeReads := 0 + a.bootstrap.APIReader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if err := c.Get(ctx, key, obj, opts...); err != nil { + return err + } + + if node, ok := obj.(*corev1.Node); ok { + nodeReads++ + if nodeReads == 1 { + node.Annotations = map[string]string{wire.BlockDevicesAnnotation: "stale"} + } else if tc.name == "absent" { + node.Annotations = nil + } else { + node.Annotations = map[string]string{wire.BlockDevicesAnnotation: tc.pattern} + } + } + + return nil + }, + }) + + var logs strings.Builder + + logger := funcr.New(func(_, msg string) { logs.WriteString(msg) }, funcr.Options{}) + ctx := ctrl.LoggerInto(t.Context(), logger) + request := httptest.NewRequest(http.MethodPost, wire.BootstrapPath, nil) + request.Header.Set("Authorization", "Bearer "+token) + encoded, hint, err := a.EnrollWithHint(ctx, request, f.request) + require.NoError(t, err) + require.Equal(t, 2, nodeReads, "configuration must use the post-issuance live Node") + response := decodeIssuedResponse(t, encoded) + require.Equal(t, tc.want, response.BlockDevices) + require.Equal(t, wire.NodeID(testNodeUID), response.Node) + require.Equal(t, types.UID(testNodeUID), hint.Node.UID) + require.Equal(t, f.request.Shares, hint.Shares) + + var fields map[string]json.RawMessage + require.NoError(t, json.Unmarshal(encoded, &fields)) + _, present := fields["block_devices"] + require.Equal(t, tc.want != "", present) + require.Equal(t, tc.warning, strings.Contains(logs.String(), "warning: ignoring block device annotation")) + + if tc.warning { + require.Contains(t, logs.String(), wire.BlockDevicesAnnotation) + require.Contains(t, logs.String(), "file-backed storage") + } + }) + } +} diff --git a/internal/racer/authority/credentials.go b/internal/racer/authority/credentials.go new file mode 100644 index 000000000..684d90ac4 --- /dev/null +++ b/internal/racer/authority/credentials.go @@ -0,0 +1,1444 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/sha256" + "crypto/x509" + "encoding/base64" + "encoding/binary" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "reflect" + "strconv" + "strings" + "sync" + "time" + + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/types" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +type credentials struct { + client.Writer + APIReader client.Reader + Config Config + Trust *trustStore + Now func() time.Time +} + +// ReconcileCredentials completes authoritative post-write signing validation +// before releasing admission. Scheduling and conflict retries belong to root. +func (a *Authority) ReconcileCredentials(ctx context.Context) (time.Duration, error) { + r := a.credentials + if err := a.gate.Acquire(ctx); err != nil { + return 0, err + } + defer a.gate.Release() + + result, err := r.reconcileKeys(ctx) + + // Refresh committed trust while the catalog gate is still held. Failed gate + // admission must not reach this completion path or withdraw accepted trust. + // Install only after authoritative validation of the committed credentials, + // including installation binding, rotation consistency, and signing lifetime. + if err == nil { + var state signingState + + state, err = loadSigning(ctx, r.APIReader, r.Config, credentialTime(r.Now)) + if err == nil { + err = r.Trust.install(ctx, state.roots, state.bundle) + } + } + + // Cancellation after admission overrides even an unavailable authority read. + if ctx.Err() != nil { + err = ctx.Err() + } + + if shouldInvalidateTrust(err) { + r.Trust.invalidate() + } + + return result.RequeueAfter, err +} + +// Credential timestamps use the same precision as X.509 validity times. +func credentialTime(now func() time.Time) time.Time { + if now != nil { + return now().UTC().Truncate(time.Second) + } + + return time.Now().UTC().Truncate(time.Second) +} + +func credentialSecret(cfg Config, name, claim string) *corev1.Secret { + return &corev1.Secret{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: name, Annotations: map[string]string{credentialClaim: claim}}, Type: corev1.SecretTypeOpaque, Data: map[string][]byte{}} +} + +func (r *credentials) reconcileKeys(ctx context.Context) (ctrl.Result, error) { + cfg := r.Config + + if err := ctx.Err(); err != nil { + return ctrl.Result{}, err + } + + if err := cfg.Validate(); err != nil { + return ctrl.Result{}, err + } + + version, _, err := readVersion(ctx, r.APIReader, cfg) + if err != nil { + return ctrl.Result{}, err + } + + var caches racerv1.ClusterCacheList + if err := r.APIReader.List(ctx, &caches); err != nil { + return ctrl.Result{}, authorityReadFailure(err) + } + + catalog, err := members.BuildCatalog(caches.Items) + if err != nil { + return ctrl.Result{}, err + } + + claim := version.Annotations[credentialClaim] + if claim == "" { + return r.initializeKeys(ctx, version, catalog) + } + + if !validCredentialClaim(cfg, claim) { + return ctrl.Result{}, wire.Unavailable + } + + credentials, err := readBoundCredentials(ctx, r.APIReader, cfg, claim, version) + if err != nil { + return ctrl.Result{}, err + } + + if err := r.Trust.validateReplay(credentials.bundle); err != nil { + return ctrl.Result{}, err + } + + return r.rotateKeys(ctx, cfg, credentials, catalog) +} + +func (r *credentials) rotateKeys(ctx context.Context, cfg Config, credentials credentialState, catalog []wire.CacheDefinition) (ctrl.Result, error) { + var err error + + catalog, err = admitCatalog(ctx, cfg, catalog, credentials.bundle) + if err != nil { + return ctrl.Result{}, err + } + + now := credentialTime(r.Now) + credentials.discardStalePreparation(cfg, now) + + if err := credentials.prepareIssuer(cfg, now); err != nil { + return ctrl.Result{}, err + } + + var encoded []byte + + credentials.bundle, credentials.rotation, encoded, err = planRotation(cfg.Rotation, credentials.bundle, credentials.rotation, catalog, now, nextGeneration(credentials.bundle.Generation)) + if err != nil { + return ctrl.Result{}, err + } + + bundleChanged, err := credentials.encodeRotation(encoded) + if err != nil { + return ctrl.Result{}, err + } + + if bundleChanged { + if err := ctx.Err(); err != nil { + return ctrl.Result{}, err + } + + if err := r.Update(ctx, credentials.secret); err != nil { + return ctrl.Result{}, err + } + } + + return ctrl.Result{RequeueAfter: max(time.Second, credentials.rotation.nextTransition().Sub(now))}, nil +} + +func (c *credentialState) discardStalePreparation(cfg Config, now time.Time) { + b, s := &c.bundle, &c.rotation + // Downtime may exhaust a staged root's useful lifetime. Cancel that unused + // preparation and stage a fresh replacement with a full new preparation delay. + if s.PreparedIssuer != "" { + cert := c.signing[s.PreparedIssuer].certificate + + activation := now + if s.ActivateAt.After(activation) { + activation = s.ActivateAt + } + + // This issuer must sign until its replacement activates one full interval + // after actual activation, and cover the last leaf's entire lifetime. + if activation.Add(cfg.Rotation.Interval + cfg.CertificateLifetime).After(cert.NotAfter) { + roots := b.PeerTrustRoots[:0] + for _, root := range b.PeerTrustRoots { + if rootID(root) != s.PreparedIssuer { + roots = append(roots, root) + } + } + + b.PeerTrustRoots = roots + + keys := b.CacheKeys[:0] + for _, key := range b.CacheKeys { + if key.State != wire.PreparedKey { + keys = append(keys, key) + } + } + + b.CacheKeys = keys + s.PreparedIssuer, s.ActivateAt, s.NextRotation = "", time.Time{}, now + } + } +} + +func (c *credentialState) prepareIssuer(cfg Config, now time.Time) error { + b, s, material := &c.bundle, &c.rotation, &c.material + + if s.ActivateAt.IsZero() && !now.Before(s.NextRotation) { + cert, key, err := generateIssuer(now, cfg) + if err != nil { + return err + } + + id := rootID(cert) + material.Keys[id] = signingMaterial{Certificate: cert, PrivateKey: key} + b.PeerTrustRoots = append(b.PeerTrustRoots, cert) + s.PreparedIssuer = id + } + + return nil +} + +// encodeRotation validates the complete candidate, including its publication +// generation, before the single Secret CAS. Reuse the planner's encoded bundle. +func (c *credentialState) encodeRotation(encoded []byte) (bool, error) { + clean := issuerMaterial{Keys: map[string]signingMaterial{}} + + for _, root := range c.bundle.PeerTrustRoots { + id := rootID(root) + clean.Keys[id] = c.material.Keys[id] + } + + c.material = clean + if err := c.validateRotation(); err != nil { + return false, err + } + + if c.bundle.Generation == c.generation { + return false, nil + } + + stateBytes, err := json.Marshal(c.rotation) + if err != nil { + return false, err + } + + materialBytes, err := json.Marshal(c.material) + if err != nil { + return false, err + } + + c.secret.Data["bundle.json"] = encoded + c.secret.Data["rotation.json"] = stateBytes + c.secret.Data["issuer.json"] = materialBytes + + return true, nil +} + +func (r *credentials) initializeKeys(ctx context.Context, version *corev1.ConfigMap, catalog []wire.CacheDefinition) (ctrl.Result, error) { + cfg := r.Config + + existing := &corev1.Secret{} + + err := r.APIReader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.CredentialsSecretName}, existing) + if !apierrors.IsNotFound(err) { + if err != nil { + return ctrl.Result{}, err + } + + return r.commitStagedCredentials(ctx, version, existing) + } + + now := credentialTime(r.Now) + + cert, key, err := generateIssuer(now, cfg) + if err != nil { + return ctrl.Result{}, err + } + + id := rootID(cert) + b := wire.KeyringBundle{SchemaVersion: wire.SchemaVersion, Cluster: cfg.Cluster, Generation: 1, PeerTrustRoots: [][]byte{cert}} + s := rotationState{ActiveIssuer: id, NextRotation: now.Add(cfg.Rotation.Interval), Retiring: map[string]time.Time{}} + + catalog, err = admitCatalog(ctx, cfg, catalog, b) + if err != nil { + return ctrl.Result{}, err + } + + var encoded []byte + + b, s, encoded, err = planRotation(cfg.Rotation, b, s, catalog, now, 1) + if err != nil { + return ctrl.Result{}, err + } + + s.NextRotation = now.Add(cfg.Rotation.Interval - cfg.Rotation.PrepareFor) + + // The permanent claim is on the already-required version object. Topology + // preserves annotations with its CAS. Missing Secrets after this claim never + // authorize Create on recovery, even when the first create response was lost. + claim := fmt.Sprintf("%s/%s", cfg.CredentialsSecretName, id) + + secret := credentialSecret(cfg, cfg.CredentialsSecretName, claim) + material := issuerMaterial{Keys: map[string]signingMaterial{id: {Certificate: cert, PrivateKey: key}}} + + secret.Data["issuer.json"], err = json.Marshal(material) + if err != nil { + return ctrl.Result{}, err + } + + secret.Data["bundle.json"] = encoded + + secret.Data["rotation.json"], err = json.Marshal(s) + if err != nil { + return ctrl.Result{}, err + } + + candidate := credentialState{bundle: b, rotation: s, material: material} + if err := candidate.validateRotation(); err != nil { + return ctrl.Result{}, err + } + + if err := ctx.Err(); err != nil { + return ctrl.Result{}, err + } + + return r.createInitialCredentials(ctx, version, secret) +} + +func (r *credentials) createInitialCredentials(ctx context.Context, version *corev1.ConfigMap, secret *corev1.Secret) (ctrl.Result, error) { + // Commit the claim only after the complete Secret exists. Until then readers + // cannot use this candidate. Recovery must bind this exact Kubernetes UID. + secret.Annotations[installationUIDAnnotation] = version.Annotations[installationUIDAnnotation] + + secret.Annotations[initializationProtocol] = stagedInitialization + if err := r.Create(ctx, secret); err != nil { + return ctrl.Result{}, err + } + + return r.commitStagedCredentials(ctx, version, secret) +} + +type RotationPolicy struct { + Interval time.Duration + PrepareFor time.Duration + RetainFor time.Duration +} + +// RotationState is controller-only metadata beside bundle.json in the credentials +// Secret. Primary timestamps suffice to derive scheduling deadlines after restart. +type rotationState struct { + NextRotation time.Time `json:"next_rotation"` + ActivateAt time.Time `json:"activate_at"` + ActiveIssuer string `json:"active_issuer"` + PreparedIssuer string `json:"prepared_issuer"` + Retiring map[string]time.Time `json:"retiring"` // Root fingerprints only. +} + +func rootID(der []byte) string { sum := sha256.Sum256(der); return hex.EncodeToString(sum[:]) } +func keyID(k wire.CacheKey) string { + return string(k.Key.Cache) + "/" + string(k.Key.Purpose) + "/" + hex.EncodeToString(k.Key.ID) +} + +func keyScope(k wire.CacheKey) string { return string(k.Key.Cache) + "/" + string(k.Key.Purpose) } + +func (s rotationState) nextTransition() time.Time { + deadline := s.NextRotation + if s.PreparedIssuer != "" { + deadline = s.ActivateAt + } + + for _, at := range s.Retiring { + if at.Before(deadline) { + deadline = at + } + } + + return deadline +} + +func newCacheKey(cache wire.CacheID, purpose wire.KeyPurpose, state wire.KeyState, generation wire.Generation) (wire.CacheKey, error) { + if generation == 0 { + return wire.CacheKey{}, wire.Unavailable + } + + var material [32]byte + if _, err := rand.Read(material[:]); err != nil { + return wire.CacheKey{}, err + } + + id := make([]byte, 16) + if _, err := rand.Read(id); err != nil { + return wire.CacheKey{}, err + } + // Reserve a versioned namespace in the otherwise opaque wire ID. A node can + // reject reintroduced epochs using its bundle high-water mark, without keeping + // every retired ID. The suffix distinguishes keys minted in competing CAS attempts. + copy(id, "RKG1") + binary.BigEndian.PutUint64(id[4:12], uint64(generation)) + + return wire.NewCacheKey(wire.CacheKeyRef{Cache: cache, Purpose: purpose, ID: id}, state, material) +} + +func nextGeneration(g wire.Generation) wire.Generation { + if g == math.MaxUint64 { + return 0 + } + + return g + 1 +} + +// planRotation owns its output. Deadlines start at actual transitions, not missed +// intervals. reconcileKeys stages issuers before planning and commits atomically. +// creationGeneration is the publication that will first contain new keys. Zero +// forbids key creation when generations are exhausted, while allowing idle plans. +func planRotation(policy RotationPolicy, b wire.KeyringBundle, s rotationState, catalog []wire.CacheDefinition, now time.Time, creationGeneration wire.Generation) (wire.KeyringBundle, rotationState, []byte, error) { + // Keep wire validation and the encoded size bound at the input boundary. + // Ownership does not require decoding the just-validated representation. + if _, err := wire.EncodeBundle(b); err != nil { + return wire.KeyringBundle{}, rotationState{}, nil, err + } + + original, originalState := b, s + b, s = cloneRotation(b, s) + + wanted, err := rotationCatalog(catalog) + if err != nil { + return b, s, nil, err + } + + pruneRotation(&b, &s, wanted, now) + + if err := addMissingCacheKeys(&b, catalog, creationGeneration); err != nil { + return b, s, nil, err + } + + if !s.ActivateAt.IsZero() && !now.Before(s.ActivateAt) { + activateRotation(&b, &s, policy, now) + } else if s.ActivateAt.IsZero() && !now.Before(s.NextRotation) { + if err := prepareCacheKeys(&b, &s, policy, now, creationGeneration); err != nil { + return b, s, nil, err + } + } + // Generation counts changed publications, not rotation cycles. + if len(original.CacheKeys) == 0 { + original.CacheKeys = []wire.CacheKey{} + } + + if originalState.Retiring == nil { + originalState.Retiring = map[string]time.Time{} + } + + if !reflect.DeepEqual(original, b) || !reflect.DeepEqual(originalState, s) { + if creationGeneration == 0 { + return b, s, nil, wire.Unavailable + } + + b.Generation = creationGeneration + } + + encoded, err := wire.EncodeBundle(b) + + return b, s, encoded, err +} + +func rotationCatalog(catalog []wire.CacheDefinition) (map[wire.CacheID]bool, error) { + wanted := make(map[wire.CacheID]bool, len(catalog)) + for _, cache := range catalog { + if !wire.ValidUUID(string(cache.ID)) || wanted[cache.ID] { + return nil, wire.InvalidRequest + } + + wanted[cache.ID] = true + } + + return wanted, nil +} + +func cloneRotation(b wire.KeyringBundle, s rotationState) (wire.KeyringBundle, rotationState) { + rootsCopy := make([][]byte, len(b.PeerTrustRoots)) + for i, root := range b.PeerTrustRoots { + rootsCopy[i] = bytes.Clone(root) + } + + b.PeerTrustRoots = rootsCopy + // Like DecodeBundle, normalize an empty key collection to a non-nil slice. + keysCopy := make([]wire.CacheKey, len(b.CacheKeys)) + for i, key := range b.CacheKeys { + keysCopy[i] = key // Includes the value-owned [32]byte material. + keysCopy[i].Key.ID = bytes.Clone(key.Key.ID) + } + + b.CacheKeys = keysCopy + + retiring := make(map[string]time.Time, len(s.Retiring)) + for id, deadline := range s.Retiring { + retiring[id] = deadline + } + + s.Retiring = retiring + + return b, s +} + +func pruneRotation(b *wire.KeyringBundle, s *rotationState, wanted map[wire.CacheID]bool, now time.Time) { + keys := b.CacheKeys[:0] + for _, k := range b.CacheKeys { + if !wanted[k.Key.Cache] { + continue + } + + keys = append(keys, k) + } + + b.CacheKeys = keys + + roots := b.PeerTrustRoots[:0] + for _, root := range b.PeerTrustRoots { + id := rootID(root) + if deadline, ok := s.Retiring[id]; ok && !now.Before(deadline) { + delete(s.Retiring, id) + continue + } + + roots = append(roots, root) + } + + b.PeerTrustRoots = roots +} + +func addMissingCacheKeys(b *wire.KeyringBundle, catalog []wire.CacheDefinition, generation wire.Generation) error { + present := map[string]bool{} + for _, key := range b.CacheKeys { + present[keyScope(key)] = true + } + + for _, cache := range catalog { + for _, purpose := range []wire.KeyPurpose{wire.PageKey, wire.OriginCredentialsKey} { + if !present[string(cache.ID)+"/"+string(purpose)] { + k, err := newCacheKey(cache.ID, purpose, wire.ActiveKey, generation) + if err != nil { + return err + } + + b.CacheKeys = append(b.CacheKeys, k) + } + } + } + + return nil +} + +func activateRotation(b *wire.KeyringBundle, s *rotationState, policy RotationPolicy, now time.Time) { + prepared := map[string]bool{} + + for _, key := range b.CacheKeys { + if key.State == wire.PreparedKey { + prepared[keyScope(key)] = true + } + } + + keys := b.CacheKeys[:0] + for _, k := range b.CacheKeys { + // A cache added during preparation can have only its initial active key. + if k.State == wire.ActiveKey && prepared[keyScope(k)] { + continue + } + + if k.State == wire.PreparedKey { + k.State = wire.ActiveKey + } + + keys = append(keys, k) + } + + b.CacheKeys = keys + + s.Retiring[s.ActiveIssuer] = now.Add(policy.RetainFor) + s.ActiveIssuer, s.PreparedIssuer = s.PreparedIssuer, "" + s.ActivateAt = time.Time{} + s.NextRotation = now.Add(policy.Interval - policy.PrepareFor) +} + +func prepareCacheKeys(b *wire.KeyringBundle, s *rotationState, policy RotationPolicy, now time.Time, generation wire.Generation) error { + if s.PreparedIssuer == "" { + return wire.Unavailable + } + + var prepared []wire.CacheKey + + for _, k := range b.CacheKeys { + if k.State != wire.ActiveKey { + continue + } + + next, err := newCacheKey(k.Key.Cache, k.Key.Purpose, wire.PreparedKey, generation) + if err != nil { + return err + } + + prepared = append(prepared, next) + } + + b.CacheKeys = append(b.CacheKeys, prepared...) + s.ActivateAt = now.Add(policy.PrepareFor) + + return nil +} + +func containsRoot(b wire.KeyringBundle, id string) bool { + for _, root := range b.PeerTrustRoots { + if rootID(root) == id { + return true + } + } + + return false +} + +// Reserve a conservative DER ceiling for generated Ed25519 roots, including +// serial-number and ASN.1 time length variation. generateIssuer enforces it. +const reservedRootBytes = 1024 + +// catalogCapacity reserves active + prepared key generations and +// active + prepared + ceil(retention / cycle) roots. Actual activations are at +// least Interval apart. The extra prepared slot is reserved even when the oldest +// retiree expires before preparation. This favors a stable limit over phase-dependent fit. +func catalogCapacity(cfg Config, b wire.KeyringBundle) (int, error) { + cycle := cfg.Rotation.Interval + + retiring := cfg.Rotation.RetainFor / cycle + if cfg.Rotation.RetainFor%cycle != 0 { + retiring++ + } + + generations := retiring + 2 + + rootBytes := reservedRootBytes + for _, root := range b.PeerTrustRoots { + rootBytes = max(rootBytes, len(root)) + } + + rootCost := base64.StdEncoding.EncodedLen(rootBytes) + 3 // quotes and comma + if int64(generations) > int64(wire.MaxBundleBytes/rootCost) { + return 0, fmt.Errorf("rotation trust reserve: %w", wire.TooLarge) + } + + // Measure the wire envelope and fixed-width key pair through the real codec. + // Reserve all 20 generation digits and the longest key state spelling. + probe := wire.KeyringBundle{SchemaVersion: wire.SchemaVersion, Cluster: cfg.Cluster, Generation: math.MaxUint64, PeerTrustRoots: b.PeerTrustRoots[:1]} + + empty, err := wire.EncodeBundle(probe) + if err != nil { + return 0, err + } + + for i, purpose := range []wire.KeyPurpose{wire.PageKey, wire.OriginCredentialsKey} { + id := make([]byte, 16) + copy(id, "RKG1") + binary.BigEndian.PutUint64(id[4:12], 1) + + key, err := wire.NewCacheKey(wire.CacheKeyRef{Cache: wire.CacheID(cfg.Cluster), Purpose: purpose, ID: id}, wire.ActiveKey, [32]byte{byte(i)}) + if err != nil { + return 0, err + } + + probe.CacheKeys = append(probe.CacheKeys, key) + } + + withKeys, err := wire.EncodeBundle(probe) + if err != nil { + return 0, err + } + + pairCost := len(withKeys) - len(empty) + 1 + 2*(len(wire.PreparedKey)-len(wire.ActiveKey)) + envelope := len(empty) - base64.StdEncoding.EncodedLen(len(probe.PeerTrustRoots[0])) - 2 + + available := wire.MaxBundleBytes - envelope - int(generations)*rootCost + if available < 0 { + return 0, fmt.Errorf("rotation trust reserve: %w", wire.TooLarge) + } + + return available / (2 * pairCost), nil +} + +// keyedCaches is the durable admission record: both active purposes must exist. +// No process-local admission history is needed across leader changes. +func keyedCaches(b wire.KeyringBundle) map[wire.CacheID]bool { + purposes := map[wire.CacheID]int{} + + for _, key := range b.CacheKeys { + if key.State == wire.ActiveKey { + purposes[key.Key.Cache]++ + } + } + + ids := make(map[wire.CacheID]bool, len(purposes)) + for id, count := range purposes { + ids[id] = count == 2 + } + + return ids +} + +// admitCatalog retains existing UIDs before filling free slots in BuildCatalog's +// UID order. New low UIDs cannot evict working caches. Deletion frees a slot; +// recreation is a new identity. Rejections are input diagnostics, not key errors. +func admitCatalog(ctx context.Context, cfg Config, catalog []wire.CacheDefinition, b wire.KeyringBundle) ([]wire.CacheDefinition, error) { + capacity, err := catalogCapacity(cfg, b) + if err != nil { + return nil, err + } + + admitted := keyedCaches(b) + existing := 0 + + for _, cache := range catalog { + if admitted[cache.ID] { + existing++ + } + } + + if existing > capacity { + // An older controller or changed policy can have overcommitted durable + // state. Never silently evict its keys or shorten retirement to make room. + return nil, fmt.Errorf("admitted catalog exceeds rotation capacity %d: %w", capacity, wire.TooLarge) + } + + slots := capacity - existing + + accepted := make([]wire.CacheDefinition, 0, min(len(catalog), capacity)) + for _, cache := range catalog { + if !admitted[cache.ID] { + if slots == 0 { + ctrl.LoggerFrom(ctx).Info("cache catalog admission rejected", "cache", cache.Name, "uid", cache.ID, "reason", "rotation_capacity", "capacity", capacity) + continue + } + + slots-- + } + + accepted = append(accepted, cache) + } + + return accepted, nil +} + +// Trust atomically holds controller-validated public roots and the matching +// delivery bundle. Requests never refresh this state or fall back to Kubernetes. +// A failed observation cannot restore withdrawn trust. +type trustStore struct { + mu sync.RWMutex + roots *x509.CertPool + bundle *acceptedKeyring + changed chan struct{} + confirmed time.Time + maxAge time.Duration + // Retain only non-secret replay protection when serving state is withdrawn. + // Otherwise a rejected rollback could be accepted on the next reconcile. + highWater wire.Generation + digest [sha256.Size]byte + authority context.Context + revoke context.CancelFunc + process context.Context +} + +// acceptedKeyring owns an immutable, bounded wire encoding, never issuer material. +// Polls share it without copying secret bytes per waiting request. +type acceptedKeyring struct { + generation wire.Generation + encoded string +} + +func (*acceptedKeyring) String() string { return "" } +func (*acceptedKeyring) GoString() string { return "" } + +func (t *trustStore) validateReplay(bundle wire.KeyringBundle) error { + encoded, err := wire.EncodeBundle(bundle) + if err != nil { + return err + } + + t.mu.RLock() + defer t.mu.RUnlock() + + return t.validateReplayLocked(bundle.Generation, sha256.Sum256(encoded)) +} + +func (t *trustStore) validateReplayLocked(generation wire.Generation, digest [sha256.Size]byte) error { + if generation < t.highWater || generation == t.highWater && digest != t.digest { + return wire.Conflict + } + + return nil +} + +func (t *trustStore) install(ctx context.Context, roots *x509.CertPool, bundle wire.KeyringBundle) error { + encoded, err := wire.EncodeBundle(bundle) + if err != nil { + return err + } + + accepted := &acceptedKeyring{generation: bundle.Generation, encoded: string(encoded)} + digest := sha256.Sum256(encoded) + + t.mu.Lock() + defer t.mu.Unlock() + + if err := ctx.Err(); err != nil { + return err + } + + if roots == nil { + return wire.Unavailable + } + + if err := t.validateReplayLocked(accepted.generation, digest); err != nil { + return err + } + + if t.bundle != nil && accepted.generation == t.highWater { + accepted = t.bundle + } + + t.highWater, t.digest = accepted.generation, digest + t.roots = roots + + t.confirmed = time.Now() + if t.authority == nil || t.authority.Err() != nil { + t.authority, t.revoke = context.WithCancel(context.Background()) + } + + t.bundle = accepted + t.notifyLocked() + + return nil +} + +func (t *trustStore) notifyLocked() { + if t.changed != nil { + close(t.changed) + } + + t.changed = make(chan struct{}) +} + +// Caller holds mu so admission and accepted bundle are captured atomically. +// Invalidation revokes admission across recovery. Rotation and reconfirmation +// preserve it without extending its captured freshness deadline. +func (t *trustStore) admitLocked(parent context.Context) (*Admission, context.CancelFunc, error) { + if t.process != nil && t.process.Err() != nil { + return nil, nil, wire.Unavailable + } + + if t.roots == nil || t.authority == nil || t.authority.Err() != nil || t.maxAge > 0 && time.Since(t.confirmed) >= t.maxAge { + return nil, nil, wire.Unavailable + } + + expiry := time.Unix(1<<62, 0) + if t.maxAge > 0 { + expiry = t.confirmed.Add(t.maxAge) + } + + guard, cancel := newAdmission(parent, t.authority, expiry) + guard.trust, guard.bundle = true, t.bundle + + guard.process = t.process + if t.process != nil { + stop := context.AfterFunc(t.process, cancel) + return guard, func() { stop(); cancel() }, nil + } + + return guard, cancel, nil +} + +func (t *trustStore) invalidate() { + if t != nil { + t.mu.Lock() + defer t.mu.Unlock() + + t.roots, t.bundle = nil, nil + if t.revoke != nil { + t.revoke() + } + + t.notifyLocked() + } +} + +func (t *trustStore) keyring() (*acceptedKeyring, <-chan struct{}, error) { + if t == nil { + return nil, nil, wire.Unavailable + } + + t.mu.RLock() + defer t.mu.RUnlock() + + if t.roots == nil || t.bundle == nil || t.maxAge > 0 && time.Since(t.confirmed) >= t.maxAge { + return nil, nil, wire.Unavailable + } + + return t.bundle, t.changed, nil +} + +func (t *trustStore) waitKeyring(ctx context.Context, after *wire.Generation) (*acceptedKeyring, error) { + timer := time.NewTimer(wire.PollWait) + defer timer.Stop() + + expired := false + + for { + if err := ctx.Err(); err != nil { + return nil, err + } + + current, changed, err := t.keyring() + if err != nil { + return nil, err + } + + if after == nil || *after < current.generation && *after != 0 { + return current, nil + } + + if *after == 0 || *after > current.generation { + return nil, wire.Conflict + } + + if expired { + return nil, nil + } + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-changed: + case <-timer.C: + expired = true + } + } +} + +// pool is immutable after installation, including when shared with TLS configs. +func (t *trustStore) pool() (*x509.CertPool, error) { + if t == nil { + return nil, wire.Unavailable + } + + t.mu.RLock() + defer t.mu.RUnlock() + + if t.roots == nil || t.maxAge > 0 && time.Since(t.confirmed) >= t.maxAge { + return nil, wire.Unavailable + } + + return t.roots, nil +} + +// An unsuccessful read supplies no new authority facts. NotFound is an observed +// deletion, unlike an unavailable API. Validation errors are never wrapped here. +type authorityReadError struct{ error } + +func (e authorityReadError) Unwrap() error { return e.error } + +func authorityReadFailure(err error) error { + if apierrors.IsNotFound(err) { + return err + } + + return authorityReadError{err} +} + +// shouldInvalidateTrust is a fail-closed policy, not proof of invalid authority. +// Only an authorityReadError preserves accepted trust; every other non-nil error +// invalidates it, including NotFound, validation, write, and unclassified failures. +// Unwrapped cancellation also invalidates; callers that fail gate admission return +// before applying this policy because they have not started observing authority. +func shouldInvalidateTrust(err error) bool { + var unread authorityReadError + return err != nil && !errors.As(err, &unread) +} + +const credentialClaim = "racer.unbounded-cloud.io/credentials" + +func validCredentialClaim(cfg Config, claim string) bool { + name, fingerprint, ok := strings.Cut(claim, "/") + decoded, err := hex.DecodeString(fingerprint) + + return ok && name == cfg.CredentialsSecretName && err == nil && len(decoded) == sha256.Size && fingerprint == hex.EncodeToString(decoded) +} + +type credentialState struct { + secret *corev1.Secret + bundle wire.KeyringBundle + rotation rotationState + material issuerMaterial + generation wire.Generation + // Parsed once per authoritative read, never used to install candidate trust. + signing map[string]parsedSigning +} + +func readBoundCredentials(ctx context.Context, reader client.Reader, cfg Config, claim string, version *corev1.ConfigMap) (credentialState, error) { + var secret corev1.Secret + + if err := ctx.Err(); err != nil { + return credentialState{}, err + } + + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.CredentialsSecretName}, &secret); err != nil { + return credentialState{}, authorityReadFailure(err) + } + + if !validCredentialClaim(cfg, claim) || secret.Annotations[credentialClaim] != claim || secret.DeletionTimestamp != nil || secret.ResourceVersion == "" { + return credentialState{}, wire.Unavailable + } + + if version.Annotations[credentialClaim] != claim || secret.UID == "" || version.Annotations[credentialUID] != string(secret.UID) || secret.Annotations[initializationProtocol] != stagedInitialization || secret.Annotations[installationUIDAnnotation] != version.Annotations[installationUIDAnnotation] { + return credentialState{}, wire.Unavailable + } + + return decodeCredentials(cfg, &secret) +} + +func decodeCredentials(cfg Config, secret *corev1.Secret) (credentialState, error) { + b, err := wire.DecodeBundle(bytes.NewReader(secret.Data["bundle.json"])) + if err != nil { + return credentialState{}, wire.Unavailable + } + + var ( + s rotationState + material issuerMaterial + ) + + if b.Cluster != cfg.Cluster || decodeCredentialMetadata(secret.Data["rotation.json"], &s) != nil || decodeCredentialMetadata(secret.Data["issuer.json"], &material) != nil { + return credentialState{}, wire.Unavailable + } + + credentials := credentialState{secret: secret, bundle: b, rotation: s, material: material, generation: b.Generation} + if err := credentials.validateRotation(); err != nil { + return credentialState{}, err + } + + return credentials, nil +} + +func decodeCredentialMetadata(data []byte, out any) error { + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.DisallowUnknownFields() + + if err := decoder.Decode(out); err != nil { + return err + } + + if err := decoder.Decode(new(any)); err != io.EOF { + return wire.Unavailable + } + + return nil +} + +func (c *credentialState) validateRotation() error { + b, s, m := c.bundle, c.rotation, c.material + if len(m.Keys) != len(b.PeerTrustRoots) || s.NextRotation.IsZero() || s.Retiring == nil || !containsRoot(b, s.ActiveIssuer) || (s.PreparedIssuer == "") != s.ActivateAt.IsZero() { + return wire.Unavailable + } + + c.signing = make(map[string]parsedSigning, len(m.Keys)) + for id, material := range m.Keys { + if rootID(material.Certificate) != id { + return wire.Unavailable + } + + cert, key, err := parseSigning(material) + if err != nil { + return err + } + + c.signing[id] = parsedSigning{certificate: cert, key: key} + } + + if !s.ActivateAt.IsZero() { + if !containsRoot(b, s.PreparedIssuer) || s.PreparedIssuer == s.ActiveIssuer || !s.ActivateAt.After(s.NextRotation) { + return wire.Unavailable + } + } + + if err := c.validateRetiringRoots(); err != nil { + return err + } + + return validatePreparedKeys(b.CacheKeys, s.ActivateAt) +} + +func (c *credentialState) validateRetiringRoots() error { + b, s := c.bundle, c.rotation + required := map[string]struct{}{} + + for _, root := range b.PeerTrustRoots { + id := rootID(root) + + key, ok := c.signing[id] + if !ok || !bytes.Equal(key.certificate.Raw, root) { + return wire.Unavailable + } + + if id != s.ActiveIssuer && id != s.PreparedIssuer { + required[id] = struct{}{} + } + } + + if len(required) != len(s.Retiring) { + return wire.Unavailable + } + + for id := range required { + if at, ok := s.Retiring[id]; !ok || at.IsZero() { + return wire.Unavailable + } + } + + return nil +} + +func validatePreparedKeys(keys []wire.CacheKey, activateAt time.Time) error { + prepared := map[string]bool{} + + for _, key := range keys { + if key.State != wire.PreparedKey { + continue + } + + scope := keyScope(key) + if activateAt.IsZero() || prepared[scope] { + return wire.Unavailable + } + + prepared[scope] = true + } + + return nil +} + +// CatalogGate serializes authoritative catalog and credential operations while +// allowing callers to abandon admission when their context is canceled. +type catalogGate struct { + token chan struct{} +} + +func newCatalogGate() *catalogGate { + g := &catalogGate{token: make(chan struct{}, 1)} + g.token <- struct{}{} + + return g +} + +// Acquire returns ownership only for a live context. A failed acquisition must +// not be released and does not constitute an observation of invalid authority. +func (g *catalogGate) Acquire(ctx context.Context) error { + if err := ctx.Err(); err != nil { + return err + } + + select { + case <-ctx.Done(): + return ctx.Err() + case <-g.token: + if err := ctx.Err(); err != nil { + g.Release() + return err + } + + return nil + } +} + +// Release ends a successfully acquired critical section. +func (g *catalogGate) Release() { + g.token <- struct{}{} +} + +const installationUIDAnnotation = "racer.unbounded-cloud.io/installation-uid" + +func readInstallation(ctx context.Context, reader client.Reader, cfg Config, fresh bool) (*corev1.ConfigMap, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + + cm := &corev1.ConfigMap{} + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName}, cm); err != nil { + return nil, authorityReadFailure(err) + } + + return cm, validateMarker(cm, cfg, fresh) +} + +func validateMarker(cm *corev1.ConfigMap, cfg Config, fresh bool) error { + state := "consumed" + if fresh { + state = "fresh" + } + + immutable := cm.Immutable != nil && *cm.Immutable + if cm.UID == "" || cm.ResourceVersion == "" || cm.DeletionTimestamp != nil || cm.Data[markerInitializationProtocol] != stagedInitialization || cm.Data["cluster"] != string(cfg.Cluster) || cm.Data["version_configmap"] != cfg.VersionConfigMapName || cm.Data["state"] != state || immutable == fresh { + return fmt.Errorf("installation marker invalid: %w", wire.Unavailable) + } + + return nil +} + +// ensureInstalled recovers only staged-v1 installations with UID-bound authority. +func ensureInstalled(ctx context.Context, writer client.Writer, reader client.Reader, cfg Config) error { + if !wire.ValidUUID(string(cfg.Cluster)) { + return wire.InvalidRequest + } + + if err := ctx.Err(); err != nil { + return err + } + + marker := &corev1.ConfigMap{} + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName}, marker); err != nil { + return err + } + + return ensureStagedInstallation(ctx, writer, reader, cfg, marker) +} + +func versionData(v versionRecord) map[string]string { + return map[string]string{"cluster": string(v.Cluster), "sequence": strconv.FormatUint(uint64(v.Sequence), 10), "membership_version": strconv.FormatUint(uint64(v.MembershipVersion), 10), "content_hash": v.ContentHash, "membership_hash": v.MembershipHash} +} + +func validHash(s string) bool { + b, err := hex.DecodeString(s) + return err == nil && len(b) == 32 && hex.EncodeToString(b) == s +} + +func (v versionRecord) valid() bool { + return wire.ValidUUID(string(v.Cluster)) && v.Sequence > 0 && v.MembershipVersion > 0 && uint64(v.MembershipVersion) <= uint64(v.Sequence) && validHash(v.ContentHash) && validHash(v.MembershipHash) +} + +func parseVersion(cm *corev1.ConfigMap, cluster wire.ClusterID, markerUID types.UID) (versionRecord, error) { + sequence, e1 := strconv.ParseUint(cm.Data["sequence"], 10, 64) + membership, e2 := strconv.ParseUint(cm.Data["membership_version"], 10, 64) + + v := versionRecord{Cluster: wire.ClusterID(cm.Data["cluster"]), Sequence: wire.Sequence(sequence), MembershipVersion: wire.MembershipVersion(membership), ContentHash: cm.Data["content_hash"], MembershipHash: cm.Data["membership_hash"]} + if e1 != nil || e2 != nil || !v.valid() || v.Cluster != cluster || cm.ResourceVersion == "" || cm.DeletionTimestamp != nil || cm.Annotations[installationUIDAnnotation] != string(markerUID) || strconv.FormatUint(sequence, 10) != cm.Data["sequence"] || strconv.FormatUint(membership, 10) != cm.Data["membership_version"] { + return versionRecord{}, fmt.Errorf("durable version state invalid; explicit new-cluster rebootstrap required: %w", wire.Unavailable) + } + + return v, nil +} + +func readVersion(ctx context.Context, reader client.Reader, cfg Config) (*corev1.ConfigMap, versionRecord, error) { + marker, err := readInstallation(ctx, reader, cfg, false) + if err != nil { + return nil, versionRecord{}, err + } + + cm := &corev1.ConfigMap{} + + if err := ctx.Err(); err != nil { + return nil, versionRecord{}, err + } + + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.VersionConfigMapName}, cm); err != nil { + return nil, versionRecord{}, authorityReadFailure(err) + } + + v, err := parseVersion(cm, cfg.Cluster, marker.UID) + if marker.Data[versionUID] == "" || marker.Data[versionUID] != string(cm.UID) || cm.Annotations[initializationProtocol] != stagedInitialization { + err = wire.Unavailable + } + + return cm, v, err +} + +const ( + initializationProtocol = "racer.unbounded-cloud.io/initialization" + markerInitializationProtocol = "initialization_protocol" + stagedInitialization = "staged-v1" + versionUID = "version_uid" + credentialUID = "racer.unbounded-cloud.io/credentials-uid" +) + +// A candidate is not authority: the +// permanent parent CAS binds its Kubernetes UID before any reader can use it. +// Retrying Create cannot restore deleted authority because the UID changes. +func ensureStagedInstallation(ctx context.Context, writer client.Writer, reader client.Reader, cfg Config, marker *corev1.ConfigMap) error { + for { + err := stageInstallation(ctx, writer, reader, cfg, marker) + if !apierrors.IsConflict(err) && !apierrors.IsAlreadyExists(err) { + return err + } + + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(50 * time.Millisecond): + } + + marker = &corev1.ConfigMap{} + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName}, marker); err != nil { + return err + } + } +} + +func stageInstallation(ctx context.Context, writer client.Writer, reader client.Reader, cfg Config, marker *corev1.ConfigMap) error { + if marker.Data[markerInitializationProtocol] != stagedInitialization { + return wire.Unavailable + } + + if marker.Data["state"] == "consumed" { + _, _, err := readVersion(ctx, reader, cfg) + return err + } + + if err := validateMarker(marker, cfg, true); err != nil { + return err + } + + if marker.Data[versionUID] != "" { + return wire.Unavailable + } + + content, membership, err := wire.ContentHashes(wire.Publication{SchemaVersion: wire.SchemaVersion, Cluster: cfg.Cluster}) + if err != nil { + return err + } + + data := versionData(versionRecord{Cluster: cfg.Cluster, Sequence: 1, MembershipVersion: 1, ContentHash: content, MembershipHash: membership}) + key := client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.VersionConfigMapName} + + candidate := &corev1.ConfigMap{} + if err := reader.Get(ctx, key, candidate); apierrors.IsNotFound(err) { + candidate = &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: key.Namespace, Name: key.Name, Annotations: map[string]string{installationUIDAnnotation: string(marker.UID), initializationProtocol: stagedInitialization}}, Data: data} + + if err := ctx.Err(); err != nil { + return err + } + + if err := writer.Create(ctx, candidate); err != nil { + return err + } + } else if err != nil { + return err + } + // A concurrent winner may already have committed and advanced the candidate. + if _, _, err := readVersion(ctx, reader, cfg); err == nil { + return nil + } + + if !validStagedVersion(candidate, marker.UID, data) { + return wire.Unavailable + } + + marker.Data[versionUID] = string(candidate.UID) + marker.Data["state"] = "consumed" + immutable := true + marker.Immutable = &immutable + + if err := ctx.Err(); err != nil { + return err + } + + if err := writer.Update(ctx, marker); err != nil { + return err + } + + _, _, err = readVersion(ctx, reader, cfg) + + return err +} + +func validStagedVersion(candidate *corev1.ConfigMap, markerUID types.UID, data map[string]string) bool { + return candidate.UID != "" && candidate.ResourceVersion != "" && candidate.DeletionTimestamp == nil && + (candidate.Immutable == nil || !*candidate.Immutable) && candidate.Annotations[installationUIDAnnotation] == string(markerUID) && + candidate.Annotations[initializationProtocol] == stagedInitialization && candidate.Annotations[credentialClaim] == "" && reflect.DeepEqual(candidate.Data, data) +} + +// Commit only a complete generation-one candidate from this installation. The +// material stays exclusively in the ordinary credentials Secret. No pending +// private-key copy, second Secret, or new RBAC permission is necessary. +func (r *credentials) commitStagedCredentials(ctx context.Context, version *corev1.ConfigMap, secret *corev1.Secret) (ctrl.Result, error) { + cfg := r.Config + + claim := secret.Annotations[credentialClaim] + if !validStagedCredentials(cfg, version, secret) { + return ctrl.Result{}, wire.Unavailable + } + + bundle, err := wire.DecodeBundle(bytes.NewReader(secret.Data["bundle.json"])) + if err != nil { + return ctrl.Result{}, wire.Unavailable + } + + candidate := credentialState{bundle: bundle} + if bundle.Cluster != cfg.Cluster || bundle.Generation != 1 || decodeCredentialMetadata(secret.Data["rotation.json"], &candidate.rotation) != nil || decodeCredentialMetadata(secret.Data["issuer.json"], &candidate.material) != nil { + return ctrl.Result{}, wire.Unavailable + } + + if err := candidate.validateRotation(); err != nil { + return ctrl.Result{}, err + } + + if claim != cfg.CredentialsSecretName+"/"+candidate.rotation.ActiveIssuer || candidate.rotation.PreparedIssuer != "" { + return ctrl.Result{}, wire.Unavailable + } + + version.Annotations[credentialClaim] = claim + version.Annotations[credentialUID] = string(secret.UID) + + if err := ctx.Err(); err != nil { + return ctrl.Result{}, err + } + + if err := r.Update(ctx, version); err != nil { + return ctrl.Result{}, err + } + + return ctrl.Result{RequeueAfter: max(time.Second, candidate.rotation.nextTransition().Sub(credentialTime(r.Now)))}, nil +} + +func validStagedCredentials(cfg Config, version *corev1.ConfigMap, secret *corev1.Secret) bool { + return version.Annotations[credentialClaim] == "" && version.Annotations[credentialUID] == "" && + secret.UID != "" && secret.ResourceVersion != "" && secret.DeletionTimestamp == nil && (secret.Immutable == nil || !*secret.Immutable) && + secret.Annotations[initializationProtocol] == stagedInitialization && secret.Annotations[installationUIDAnnotation] == version.Annotations[installationUIDAnnotation] && + validCredentialClaim(cfg, secret.Annotations[credentialClaim]) +} diff --git a/internal/racer/authority/credentials_test.go b/internal/racer/authority/credentials_test.go new file mode 100644 index 000000000..25db2b4ab --- /dev/null +++ b/internal/racer/authority/credentials_test.go @@ -0,0 +1,1836 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "bytes" + "context" + "crypto/sha256" + "crypto/x509" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "math" + "net/http/httptest" + "reflect" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/types" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestCredentialsSingleCreateAndPermanentClaim(t *testing.T) { + r, _ := testKeyring(t) + base := r.Client.(client.WithWatch) + creates := 0 + r.Client = interceptor.NewClient(base, interceptor.Funcs{Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + creates++ + + secret, ok := obj.(*corev1.Secret) + if !ok || secret.Name != r.Config.CredentialsSecretName || len(secret.Data) != 3 { + t.Fatal("initialization did not create one complete credentials Secret") + } + + version, _, err := readVersion(ctx, base, r.Config) + if err != nil { + t.Fatal(err) + } + + var material issuerMaterial + if err := json.Unmarshal(secret.Data["issuer.json"], &material); err != nil { + t.Fatal(err) + } + + for id := range material.Keys { + want := secret.Name + "/" + id + if version.Annotations[credentialClaim] != "" || version.Annotations[credentialUID] != "" || secret.Annotations[credentialClaim] != want { + t.Fatal("candidate claim missing or parent committed before Create") + } + } + + return c.Create(ctx, obj, opts...) + }}) + runKeys(t, r) + runKeys(t, r) + version, _, err := readVersion(t.Context(), base, r.Config) + require.NoError(t, err) + + secret := &corev1.Secret{} + require.NoError(t, base.Get(t.Context(), client.ObjectKey{Namespace: r.Config.Namespace, Name: r.Config.CredentialsSecretName}, secret)) + require.NotEmpty(t, secret.UID) + require.Equal(t, string(secret.UID), version.Annotations[credentialUID]) + require.Equal(t, secret.Annotations[credentialClaim], version.Annotations[credentialClaim]) + + if creates != 1 { + t.Fatalf("credential Creates = %d, want 1", creates) + } +} + +func TestCredentialsClaimValidation(t *testing.T) { + r, _ := testKeyring(t) + name := r.Config.CredentialsSecretName + + fingerprint := strings.Repeat("ab", 32) + for _, claim := range []string{"", name + "/", name + "/" + fingerprint + "/extra", "other/" + fingerprint, name + "/" + strings.Repeat("x", 64), name + "/" + strings.ToUpper(fingerprint)} { + t.Run(claim, func(t *testing.T) { + if validCredentialClaim(r.Config, claim) { + t.Fatal("invalid claim accepted") + } + }) + } + + if !validCredentialClaim(r.Config, name+"/"+fingerprint) { + t.Fatal("valid claim rejected") + } +} + +func TestCredentialsMissingOrLegacyEntriesNeverRegenerate(t *testing.T) { + for _, corruption := range []string{"issuer.json", "bundle.json", "rotation.json", "pending", "trailing metadata", "symmetric retirement", "duplicate material"} { + t.Run(corruption, func(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + secret, bundle, state, material := keyState(t, r) + + switch corruption { + case "duplicate material": + var document map[string]any + if err := json.Unmarshal(secret.Data["bundle.json"], &document); err != nil { + t.Fatal(err) + } + + keys := document["cache_keys"].([]any) + keys[1].(map[string]any)["material"] = keys[0].(map[string]any)["material"] + document["generation"] = fmt.Sprint(uint64(bundle.Generation + 1)) + secret.Data["bundle.json"], _ = json.Marshal(document) + case "pending": + secret.Data["issuer.json"], _ = json.Marshal(map[string]any{"keys": material.Keys, "pending": state.ActiveIssuer}) + case "trailing metadata": + secret.Data["rotation.json"] = append(secret.Data["rotation.json"], []byte(" {}")...) + case "symmetric retirement": + state.Retiring[keyID(bundle.CacheKeys[0])] = state.NextRotation + secret.Data["rotation.json"], _ = json.Marshal(state) + default: + delete(secret.Data, corruption) + } + + if err := r.Update(t.Context(), secret); err != nil { + t.Fatal(err) + } + + base := r.Client.(client.WithWatch) + + r.Client = rejectWrites(t, base) + if _, err := r.operations().ReconcileCredentials(t.Context()); err == nil || trustReady(r.Trust) { + t.Fatal("corrupt atomic version accepted") + } else if corruption == "duplicate material" && !errors.Is(err, wire.Unavailable) { + t.Fatalf("committed duplicate material should be unavailable: %v", err) + } + }) + } +} + +func TestCredentialsIdlePreservesEncoding(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + secret, _, _, _ := keyState(t, r) + + for _, entry := range []string{"issuer.json", "rotation.json"} { + var indented bytes.Buffer + if err := json.Indent(&indented, secret.Data[entry], "", " "); err != nil { + t.Fatal(err) + } + + secret.Data[entry] = indented.Bytes() + } + + if err := r.Update(t.Context(), secret); err != nil { + t.Fatal(err) + } + + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + t.Fatal("idle plan normalized metadata with a write") + return nil + }}) + runKeys(t, r) + + after, _, _, _ := keyState(t, r) + if after.ResourceVersion != secret.ResourceVersion || !reflect.DeepEqual(after.Data, secret.Data) { + t.Fatal("idle credentials changed") + } +} + +func TestCredentialsCandidateIsOneCoherentCAS(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + base := r.Client.(client.WithWatch) + writes := 0 + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + writes++ + secret := obj.(*corev1.Secret) + assertPreparedCredentials(t, secret) + + return c.Update(ctx, obj, opts...) + }}) + runKeys(t, r) + + if writes != 1 { + t.Fatalf("rotation CAS count = %d", writes) + } + + accepted, _, err := r.Trust.keyring() + if err != nil || accepted.generation != 2 { + t.Fatal("committed publication not installed") + } + // The claim never follows rotating issuer identities. + version, _, err := readVersion(t.Context(), base, r.Config) + if err != nil || version.Annotations[credentialClaim] != r.Config.CredentialsSecretName+"/"+initial.ActiveIssuer { + t.Fatal("rotation changed the permanent creation claim") + } + + if _, err := loadSigning(t.Context(), base, r.Config, *now); err != nil { + t.Fatal(err) + } +} + +func assertPreparedCredentials(t *testing.T, secret *corev1.Secret) { + t.Helper() + + bundle, err := wire.DecodeBundle(bytes.NewReader(secret.Data["bundle.json"])) + require.NoError(t, err) + + var ( + rotation RotationState + material issuerMaterial + ) + + require.NoError(t, json.Unmarshal(secret.Data["rotation.json"], &rotation), "unreadable candidate rotation") + require.NoError(t, json.Unmarshal(secret.Data["issuer.json"], &material), "unreadable candidate issuers") + candidate := credentialState{bundle: bundle, rotation: rotation, material: material} + require.NoError(t, candidate.validateRotation(), "incoherent candidate") + require.EqualValues(t, 2, bundle.Generation) + require.NotEmpty(t, rotation.PreparedIssuer) + require.Len(t, material.Keys, 2) + + for _, key := range bundle.CacheKeys { + require.LessOrEqual(t, binary.BigEndian.Uint64(key.Key.ID[4:12]), uint64(bundle.Generation), "creation generation exceeds publication") + } +} + +func TestCredentialsStalePreparationReplacementIsAtomic(t *testing.T) { + for _, fail := range []bool{false, true} { + t.Run(fmt.Sprint(fail), func(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + _, bundle, state, material := keyState(t, r) + oldID := state.PreparedIssuer + short := editSigningCertificate(t, material.Keys[oldID], func(cert *x509.Certificate) { + cert.NotAfter = state.ActivateAt.Add(r.Config.Rotation.Interval + r.Config.CertificateLifetime - time.Second) + }) + shortID := rootID(short.Certificate) + + rebindSigning(&bundle, &state, material, oldID, short, true) + + bundle.Generation++ + writeSigningCredentials(t, r, bundle, state, material) + before, _, _, _ := keyState(t, r) + + base := r.Client.(client.WithWatch) + if fail { + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + return wire.Unavailable + }}) + _, err := r.operations().ReconcileCredentials(t.Context()) + require.Error(t, err, "failed replacement accepted") + + after, _, _, _ := keyState(t, r) + require.Equal(t, before.Data, after.Data, "failed stale replacement partially changed credentials") + + r.Client = base + } + + runKeys(t, r) + + _, after, replacement, keys := keyState(t, r) + if after.Generation != bundle.Generation+1 || replacement.PreparedIssuer == shortID || replacement.ActiveIssuer != initial.ActiveIssuer || !replacement.ActivateAt.Equal(now.Add(r.Config.Rotation.PrepareFor)) { + t.Fatal("stale preparation replacement reset active issuer or publication version") + } + + require.False(t, containsRoot(after, shortID), "stale root retained") + require.Len(t, keys.Keys, len(after.PeerTrustRoots), "private material not replaced atomically") + require.NotContains(t, keys.Keys, shortID, "stale private key retained") + + for _, key := range after.CacheKeys { + if key.State == wire.PreparedKey && binary.BigEndian.Uint64(key.Key.ID[4:12]) != uint64(after.Generation) { + t.Fatal("replacement prepared key has wrong creation generation") + } + } + }) + } +} + +func TestKeyringReplayProtectionSurvivesInvalidation(t *testing.T) { + for _, scenario := range []string{"rollback", "conflicting generation"} { + for _, restore := range []string{"accepted generation", "new generation"} { + t.Run(scenario+"/"+restore, func(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + shared, bundle, _, _ := keyState(t, r) + writeBundle := func(candidate wire.KeyringBundle) { + t.Helper() + + encoded, err := wire.EncodeBundle(candidate) + require.NoError(t, err) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(shared), shared)) + shared.Data["bundle.json"] = encoded + require.NoError(t, r.Update(t.Context(), shared)) + } + bundle.Generation = 2 + writeBundle(bundle) + runKeys(t, r) + + accepted, _, err := r.Trust.keyring() + require.NoError(t, err) + require.EqualValues(t, 2, accepted.generation) + + digest := sha256.Sum256([]byte(accepted.encoded)) + + candidate, err := wire.DecodeBundle(bytes.NewBufferString(accepted.encoded)) + require.NoError(t, err) + + if scenario == "rollback" { + candidate.Generation-- + } else { + key := candidate.CacheKeys[0] + + candidate.CacheKeys[0], err = wire.NewCacheKey(key.Key, key.State, [32]byte{1}) + require.NoError(t, err) + } + + writeBundle(candidate) + + for range 3 { + _, err := r.operations().ReconcileCredentials(t.Context()) + require.ErrorIs(t, err, wire.Conflict, "replay was not rejected") + _, _, err = r.Trust.keyring() + require.ErrorIs(t, err, wire.Unavailable, "replay restored delivery") + _, err = r.Trust.pool() + require.ErrorIs(t, err, wire.Unavailable, "replay restored trust") + require.EqualValues(t, 2, r.Trust.highWater) + require.Equal(t, digest, r.Trust.digest, "invalidation lost accepted replay protection") + } + + if restore == "new generation" { + bundle = candidate + bundle.Generation = 3 + } + + writeBundle(bundle) + runKeys(t, r) + + current, _, err := r.Trust.keyring() + require.NoError(t, err) + require.Equal(t, bundle.Generation, current.generation) + + expected, err := wire.EncodeBundle(bundle) + require.NoError(t, err) + require.Equal(t, string(expected), current.encoded, "restored wrong content") + + _, err = r.Trust.pool() + require.NoError(t, err, "restoration did not restore trust") + require.Equal(t, bundle.Generation, r.Trust.highWater) + require.Equal(t, sha256.Sum256(expected), r.Trust.digest, "restoration did not advance replay protection") + }) + } + } +} + +func TestAcceptedKeyringImmutableAndValidated(t *testing.T) { + f := newServingFixture(t) + trust := f.a.Server.Trust + + initial, _, err := trust.keyring() + if err != nil { + t.Fatal(err) + } + + roots, err := trust.pool() + if err != nil { + t.Fatal(err) + } + + _, bundle, _, _ := keyState(t, f.a.Keyring) + + bundle.Generation++ + if err := trust.install(f.ctx, roots, bundle); err != nil { + t.Fatal(err) + } + + accepted, _, err := trust.keyring() + if err != nil { + t.Fatal(err) + } + + expected := accepted.encoded + bundle.PeerTrustRoots[0][0] ^= 0xff + + if accepted.encoded != expected || initial.generation != 1 { + t.Fatal("caller mutated immutable delivery") + } + + for _, format := range []string{"%v", "%+v", "%#v"} { + require.Equal(t, "", fmt.Sprintf(format, accepted), "diagnostic exposed keyring") + } + + assertRejectedKeyringInstalls(t, f, roots, accepted) + + older := wire.Generation(1) + + current, err := trust.waitKeyring(f.ctx, &older) + if err != nil || current != accepted { + t.Fatalf("older cursor did not return newest: %v", err) + } +} + +func assertRejectedKeyringInstalls(t *testing.T, f *servingFixture, roots *x509.CertPool, accepted *acceptedKeyring) { + t.Helper() + + trust := f.a.Server.Trust + + for _, scenario := range []string{"invalid", "rollback", "same generation changed", "canceled", "nil roots"} { + t.Run(scenario, func(t *testing.T) { + _, candidate, _, _ := keyState(t, f.a.Keyring) + candidate.Generation = 3 + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + pool := roots + + switch scenario { + case "invalid": + candidate.SchemaVersion = 0 + case "rollback": + candidate.Generation = 1 + case "same generation changed": + candidate.Generation = 2 + candidate.PeerTrustRoots[0][0] ^= 0xff + case "canceled": + cancel() + case "nil roots": + pool = nil + } + + if err := trust.install(ctx, pool, candidate); err == nil { + t.Fatal("invalid installation succeeded") + } + + current, _, err := trust.keyring() + if err != nil || current != accepted { + t.Fatal("failed installation changed delivery") + } + }) + } +} + +func TestIssuerObservationInvalidatesKeyringDelivery(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + + accepted, _, err := f.a.Server.Trust.keyring() + if err != nil { + t.Fatal(err) + } + + issuer := f.a.Server.Bootstrap.Issuer + live := issuer.APIReader + + issuer.APIReader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + return errors.New("offline") + }}) + if _, err := issuer.TrustRoots(f.ctx); err == nil { + t.Fatal("outage hidden") + } + + if current, _, err := f.a.Server.Trust.keyring(); err != nil || current != accepted { + t.Fatal("issuer read outage withdrew accepted delivery") + } + + issuer.APIReader = live + done := make(chan error, 1) + + go func() { _, err := f.a.Server.Trust.waitKeyring(f.ctx, &accepted.generation); done <- err }() + + synctest.Wait() + + shared, _, _, _ := keyState(t, f.a.Keyring) + valid := bytes.Clone(shared.Data["bundle.json"]) + + shared.Data["bundle.json"] = []byte(`{}`) + if err := f.a.Topology.Update(f.ctx, shared); err != nil { + t.Fatal(err) + } + + if _, err := issuer.TrustRoots(f.ctx); err == nil { + t.Fatal("invalid authority accepted") + } + + select { + case err := <-done: + if !errors.Is(err, wire.Unavailable) { + t.Fatalf("invalidated delivery: %v", err) + } + case <-time.After(time.Second): + t.Fatal("issuer invalidation did not wake poll") + } + + shared.Data["bundle.json"] = valid + if err := f.a.Topology.Update(f.ctx, shared); err != nil { + t.Fatal(err) + } + + if _, err := issuer.TrustRoots(f.ctx); err != nil { + t.Fatal(err) + } + + if _, _, err := f.a.Server.Trust.keyring(); err == nil { + t.Fatal("issuance restored delivery without reconciliation") + } + + runKeys(t, f.a.Keyring) + + if _, _, err := f.a.Server.Trust.keyring(); err != nil { + t.Fatal(err) + } + }) +} + +func TestCredentialTime(t *testing.T) { + t.Run("injected", func(t *testing.T) { + now := time.Date(2026, time.October, 7, 12, 30, 5, 123456789, time.FixedZone("offset", 2*60*60)) + got := credentialTime(func() time.Time { return now }) + require.Equal(t, now.UTC().Truncate(time.Second), got) + require.Same(t, time.UTC, got.Location()) + }) + t.Run("default", func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + time.Sleep(123 * time.Millisecond) + require.Equal(t, time.Now().UTC().Truncate(time.Second), credentialTime(nil)) + }) + }) +} + +func testKeyring(t *testing.T) (*credentialsFixture, *time.Time) { + t.Helper() + + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache", UID: types.UID(testNodeUID)}} + topology := initializedTopology(t, cache) + a := Assemble(topology.Config, topology.Client, topology.APIReader) + now := time.Now().UTC().Truncate(time.Second) + a.Keyring.Now = func() time.Time { return now } + + return a.Keyring, &now +} + +func runKeys(t *testing.T, r *credentialsFixture) time.Duration { + t.Helper() + + result, err := r.operations().ReconcileCredentials(t.Context()) + if err != nil || result <= 0 { + t.Fatalf("reconcile: %v, %v", result, err) + } + + return result +} + +func keyState(t *testing.T, r *credentialsFixture) (*corev1.Secret, wire.KeyringBundle, RotationState, issuerMaterial) { + t.Helper() + + version, _, err := readVersion(context.Background(), r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + credentials, err := readBoundCredentials(context.Background(), r.APIReader, r.Config, version.Annotations[credentialClaim], version) + if err != nil { + t.Fatal(err) + } + + return credentials.secret, credentials.bundle, credentials.rotation, credentials.material +} + +func TestKeyringRotationLifecycle(t *testing.T) { + r, now := testKeyring(t) + started := *now + result := runKeys(t, r) + + shared, initial, state, _ := keyState(t, r) + require.EqualValues(t, 1, initial.Generation) + require.Len(t, initial.CacheKeys, 2) + require.Len(t, initial.PeerTrustRoots, 1) + require.Equal(t, r.Config.Rotation.Interval-r.Config.Rotation.PrepareFor, result) + require.True(t, trustReady(r.Trust), "initial readiness") + + assertCacheKeyGenerations(t, initial.CacheKeys, 1, 1) + // Repeated reconciliation must neither write nor consume a generation. + runKeys(t, r) + + unchanged, _, _, _ := keyState(t, r) + if shared.ResourceVersion != unchanged.ResourceVersion { + t.Fatal("unchanged credentials written") + } + + *now = state.NextRotation + result = runKeys(t, r) + + _, staged, prepared, _ := keyState(t, r) + if !prepared.ActivateAt.Equal(started.Add(r.Config.Rotation.Interval)) { + t.Fatal("activation cadence must include the preparation interval") + } + + require.EqualValues(t, 2, staged.Generation) + require.Len(t, staged.CacheKeys, 4) + require.Len(t, staged.PeerTrustRoots, 2) + require.Equal(t, state.ActiveIssuer, prepared.ActiveIssuer) + require.NotEmpty(t, prepared.PreparedIssuer) + require.Equal(t, r.Config.Rotation.PrepareFor, result) + + assertCacheKeyGenerations(t, staged.CacheKeys, 1, 2) + + if !reflect.DeepEqual(staged.CacheKeys[:len(initial.CacheKeys)], initial.CacheKeys) { + t.Fatal("staging changed retained keys") + } + // A fresh process resumes the persisted deadline, not a new delay. + restarted := Assemble(r.Config, r.Client, r.APIReader).Keyring + restarted.Now = r.Now + + *now = prepared.ActivateAt.Add(-time.Second) + if got := runKeys(t, restarted); got != time.Second { + t.Fatalf("lost activation deadline: %v", got) + } + + *now = prepared.ActivateAt + + runKeys(t, restarted) + + _, activated, active, _ := keyState(t, r) + require.EqualValues(t, 3, activated.Generation) + require.Equal(t, prepared.PreparedIssuer, active.ActiveIssuer) + require.Empty(t, active.PreparedIssuer) + require.Len(t, active.Retiring, 1) + require.Len(t, activated.CacheKeys, 2) + + assertActivatedCacheKeys(t, activated.CacheKeys, staged.CacheKeys[2:]) + // Further cycles overlap without evicting an earlier retirement prematurely. + *now = active.NextRotation + + runKeys(t, restarted) + _, _, nextStage, _ := keyState(t, r) + *now = nextStage.ActivateAt + + runKeys(t, restarted) + + _, overlap, overlapping, _ := keyState(t, r) + if len(overlap.PeerTrustRoots) != 3 || len(overlap.CacheKeys) != 2 { + t.Fatal("multiple retiring generations lost") + } + + *now = active.Retiring[state.ActiveIssuer] + + runKeys(t, restarted) + + _, pruned, prunedState, material := keyState(t, r) + if containsRoot(pruned, state.ActiveIssuer) || len(pruned.CacheKeys) != 4 || len(material.Keys) != 3 || !prunedState.Retiring[active.ActiveIssuer].Equal(overlapping.Retiring[active.ActiveIssuer]) { + t.Fatal("retirement pruning/reset") + } + // Topology CAS preserves the one-way initialization claim. + topology := Assemble(r.Config, r.Client, r.APIReader).Topology + reconcileTopology(t, topology, context.Background()) + keyState(t, r) +} + +func assertCacheKeyGenerations(t *testing.T, keys []wire.CacheKey, active, prepared uint64) { + t.Helper() + + for _, k := range keys { + want := active + if active == prepared { + require.Equal(t, wire.ActiveKey, k.State, "initial key not active") + } + + if k.State == wire.PreparedKey { + want = prepared + } + + require.Equal(t, "RKG1", string(k.Key.ID[:4])) + require.Equal(t, want, binary.BigEndian.Uint64(k.Key.ID[4:12]), "key creation generation changed") + } +} + +func assertActivatedCacheKeys(t *testing.T, activated, staged []wire.CacheKey) { + t.Helper() + + for i, k := range activated { + require.NotEqual(t, wire.PreparedKey, k.State, "prepared key not activated") + require.Equal(t, staged[i].Key, k.Key, "activation changed key identity") + require.True(t, k.EqualMaterial(staged[i]), "activation changed key material") + } +} + +func TestKeyringDerivedDeadlines(t *testing.T) { + for _, tc := range []struct { + name string + interval time.Duration + retain time.Duration + steps []struct{ at, next time.Duration } + }{ + { + name: "overlapping retirements", + interval: 24 * time.Hour, + retain: 48 * time.Hour, + steps: []struct{ at, next time.Duration }{ + {0, 23 * time.Hour}, + {23 * time.Hour, 24 * time.Hour}, + {24 * time.Hour, 47 * time.Hour}, + {47 * time.Hour, 48 * time.Hour}, + {48 * time.Hour, 71 * time.Hour}, + {71 * time.Hour, 72 * time.Hour}, + }, + }, + { + name: "retirement during preparation", + interval: 24 * time.Hour, + retain: 24*time.Hour + 30*time.Minute, + steps: []struct{ at, next time.Duration }{ + {0, 23 * time.Hour}, + {23 * time.Hour, 24 * time.Hour}, + {24 * time.Hour, 47 * time.Hour}, + {47 * time.Hour, 48 * time.Hour}, + {48 * time.Hour, 48*time.Hour + 30*time.Minute}, + {48*time.Hour + 30*time.Minute, 71 * time.Hour}, + }, + }, + { + name: "retirement before next rotation", + interval: 7 * 24 * time.Hour, + retain: 24 * time.Hour, + steps: []struct{ at, next time.Duration }{ + {0, 167 * time.Hour}, + {167 * time.Hour, 168 * time.Hour}, + {168 * time.Hour, 192 * time.Hour}, + {192 * time.Hour, 335 * time.Hour}, + }, + }, + } { + t.Run(tc.name, func(t *testing.T) { + r, now := testKeyring(t) + r.Config.Rotation = RotationPolicy{Interval: tc.interval, PrepareFor: time.Hour, RetainFor: tc.retain} + start := *now + + for i, step := range tc.steps { + *now = start.Add(step.at) + result := runKeys(t, r) + + shared, bundle, state, _ := keyState(t, r) + if want := start.Add(step.next); !state.nextTransition().Equal(want) || result != want.Sub(*now) { + t.Fatalf("step %d: deadline=%v delay=%v, want %v", i, state.nextTransition(), result, want) + } + + if bundle.Generation != wire.Generation(i+1) { + t.Fatalf("step %d did not publish exactly one transition: generation=%d", i, bundle.Generation) + } + + var persisted map[string]json.RawMessage + require.NoError(t, json.Unmarshal(shared.Data["rotation.json"], &persisted)) + require.NotContains(t, persisted, "next_transition", "derived deadline persisted") + + // Every phase must resume from only primary state without rewriting + // credentials or extending a deadline when the process restarts. + restarted := Assemble(r.Config, r.Client, r.APIReader).Keyring + restarted.Now = r.Now + + *now = now.Add(time.Second) + require.Equal(t, step.next-step.at-time.Second, runKeys(t, restarted), "step %d: restart delay", i) + + unchanged, _, _, _ := keyState(t, restarted) + require.Equal(t, shared.ResourceVersion, unchanged.ResourceVersion, "step %d: restart rewrote credentials", i) + require.Equal(t, shared.Data, unchanged.Data, "step %d: restart rewrote credentials", i) + + r = restarted + } + }) + } +} + +func TestGenerationBoundCacheKey(t *testing.T) { + if _, err := newCacheKey(wire.CacheID(testNodeUID), wire.PageKey, wire.ActiveKey, 0); err == nil { + t.Fatal("zero or wrapped generation accepted") + } + + for _, generation := range []wire.Generation{1, 256, math.MaxUint64} { + key, err := newCacheKey(wire.CacheID(testNodeUID), wire.PageKey, wire.PreparedKey, generation) + if err != nil { + t.Fatal(err) + } + + if len(key.Key.ID) != 16 || string(key.Key.ID[:4]) != "RKG1" || binary.BigEndian.Uint64(key.Key.ID[4:12]) != uint64(generation) { + t.Fatalf("creation generation not preserved: %x", key.Key.ID) + } + } +} + +func TestPlanRotationExhaustedKeyCreation(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, b, state, _ := keyState(t, r) + b.Generation = math.MaxUint64 + catalog := []wire.CacheDefinition{{ID: wire.CacheID(testNodeUID)}} + + next, nextState, _, err := planRotation(r.Config.Rotation, b, state, catalog, *now, nextGeneration(b.Generation)) + if err != nil || !reflect.DeepEqual(next, b) || !reflect.DeepEqual(nextState, state) { + t.Fatalf("exhausted idle plan: %v", err) + } + + catalog = append(catalog, wire.CacheDefinition{ID: wire.CacheID(testOtherUID)}) + if _, _, _, err := planRotation(r.Config.Rotation, b, state, catalog, *now, nextGeneration(b.Generation)); !errors.Is(err, wire.Unavailable) { + t.Fatalf("exhausted admission plan: %v", err) + } + + state.PreparedIssuer = state.ActiveIssuer + if _, _, _, err := planRotation(r.Config.Rotation, b, state, catalog[:1], state.NextRotation, nextGeneration(b.Generation)); !errors.Is(err, wire.Unavailable) { + t.Fatalf("exhausted staging plan: %v", err) + } +} + +func TestKeyringCatalogAndBounds(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + shared, original, s, _ := keyState(t, r) + // Active-only fits, but staging the same catalog must account for overlap. + var catalog []wire.CacheDefinition + for n := range 800 { + catalog = append(catalog, wire.CacheDefinition{ID: wire.CacheID(fmt.Sprintf("%08x-0000-0000-0000-000000000000", n+1))}) + } + + b, state, _, err := planRotation(r.Config.Rotation, original, s, catalog, *now, nextGeneration(original.Generation)) + if err != nil { + t.Fatal(err) + } + + state.PreparedIssuer = state.ActiveIssuer + if _, _, _, err := planRotation(r.Config.Rotation, b, state, catalog, state.NextRotation, nextGeneration(b.Generation)); !errors.Is(err, wire.TooLarge) { + t.Fatalf("overlap bound not enforced: %v", err) + } + + encoded, err := wire.EncodeBundle(original) + if err != nil { + t.Fatal(err) + } + + if !bytes.Equal(encoded, shared.Data["bundle.json"]) { + t.Fatal("planner mutated input") + } + + cache := &racerv1.ClusterCache{} + if err := r.APIReader.Get(context.Background(), client.ObjectKey{Name: "cache"}, cache); err != nil { + t.Fatal(err) + } + + if err := r.Delete(context.Background(), cache); err != nil { + t.Fatal(err) + } + + runKeys(t, r) + + _, removed, _, _ := keyState(t, r) + if len(removed.CacheKeys) != 0 || removed.Generation != 2 { + t.Fatal("removed cache keys retained") + } + + cache.ResourceVersion = "" + + cache.UID = types.UID(testOtherUID) + if err := r.Create(context.Background(), cache); err != nil { + t.Fatal(err) + } + + runKeys(t, r) + + _, recreated, _, _ := keyState(t, r) + for _, key := range recreated.CacheKeys { + if key.Key.Cache != wire.CacheID(testOtherUID) || key.EqualMaterial(original.CacheKeys[0]) { + t.Fatal("recreated cache inherited keys") + } + } +} + +func TestPlanRotationOwnsOutput(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + _, _, prepared, _ := keyState(t, r) + *now = prepared.ActivateAt + + runKeys(t, r) + _, original, state, _ := keyState(t, r) + catalog := []wire.CacheDefinition{{ID: wire.CacheID(testNodeUID)}} + // The old codec round trip preserves order, even when it is not sorted. + original.PeerTrustRoots[0], original.PeerTrustRoots[1] = original.PeerTrustRoots[1], original.PeerTrustRoots[0] + original.CacheKeys[0], original.CacheKeys[1] = original.CacheKeys[1], original.CacheKeys[0] + + before, err := wire.EncodeBundle(original) + if err != nil { + t.Fatal(err) + } + + beforeState, err := json.Marshal(state) + if err != nil { + t.Fatal(err) + } + + want, err := wire.DecodeBundle(bytes.NewReader(before)) + if err != nil { + t.Fatal(err) + } + + next, nextState, _, err := planRotation(r.Config.Rotation, original, state, catalog, *now, nextGeneration(original.Generation)) + if err != nil { + t.Fatal(err) + } + + if !reflect.DeepEqual(next, want) || !reflect.DeepEqual(nextState, state) { + t.Fatal("idle planner changed bundle representation or rotation state") + } + + // Exercise both outer collections and nested bytes, plus the retirement map. + next.PeerTrustRoots[0][0] ^= 0xff + next.PeerTrustRoots[1] = nil + next.CacheKeys[0].Key.ID[0] ^= 0xff + next.CacheKeys[1].State = wire.PreparedKey + + for id := range nextState.Retiring { + delete(nextState.Retiring, id) + } + + after, err := wire.EncodeBundle(original) + if err != nil || !bytes.Equal(before, after) { + t.Fatalf("output aliases input bundle: %v", err) + } + + afterState, err := json.Marshal(state) + if err != nil || !bytes.Equal(beforeState, afterState) { + t.Fatalf("output aliases input retirement map: %v", err) + } +} + +func TestPlanRotationInputValidation(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + + _, original, state, _ := keyState(t, r) + for _, tc := range []struct { + name string + edit func(*wire.KeyringBundle) + want error + }{ + {"schema", func(b *wire.KeyringBundle) { b.SchemaVersion++ }, wire.UnsupportedVersion}, + {"generation", func(b *wire.KeyringBundle) { b.Generation = 0 }, wire.InvalidRequest}, + {"duplicate root", func(b *wire.KeyringBundle) { b.PeerTrustRoots = append(b.PeerTrustRoots, b.PeerTrustRoots[0]) }, wire.InvalidRequest}, + {"invalid key", func(b *wire.KeyringBundle) { + b.CacheKeys = append([]wire.CacheKey(nil), b.CacheKeys...) + b.CacheKeys[0].Key.ID = nil + }, wire.InvalidRequest}, + {"encoded size", func(b *wire.KeyringBundle) { + key := b.CacheKeys[0] + + b.CacheKeys = make([]wire.CacheKey, 4000) + for i := range b.CacheKeys { + ref := key.Key + ref.Cache = wire.CacheID(fmt.Sprintf("%08x-0000-0000-0000-000000000000", i+1)) + + var material [32]byte + binary.BigEndian.PutUint64(material[:8], uint64(i+1)) + + var err error + + b.CacheKeys[i], err = wire.NewCacheKey(ref, key.State, material) + if err != nil { + t.Fatal(err) + } + } + }, wire.TooLarge}, + } { + t.Run(tc.name, func(t *testing.T) { + b := original + tc.edit(&b) + // An empty catalog would discard all keys; invalid/oversized input + // must still fail before the planner can shrink it into a valid output. + if _, _, _, err := planRotation(r.Config.Rotation, b, state, nil, *now, nextGeneration(b.Generation)); !errors.Is(err, tc.want) { + t.Fatalf("input validation: got %v, want %v", err, tc.want) + } + }) + } + + original.CacheKeys = nil + state.Retiring = nil + + next, nextState, _, err := planRotation(r.Config.Rotation, original, state, nil, *now, nextGeneration(original.Generation)) + if err != nil || next.CacheKeys == nil || nextState.Retiring == nil { + t.Fatalf("empty collection normalization: %v", err) + } +} + +func TestKeyringCorruptionAndGenerationExhaustion(t *testing.T) { + for _, corrupt := range []string{"timestamp", "transition mismatch", "missing activation", "missing prepared issuer", "zero root retirement", "zero key retirement", "missing root retirement", "missing key retirement", "replaced retirement", "nil retirement", "unknown retirement", "active issuer", "bundle", "private key", "generation", "binding"} { + t.Run(corrupt, func(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + + prepareCorruptionPhase(t, r, now, corrupt) + + shared, b, s, _ := keyState(t, r) + corruptCredentials(corrupt, shared, b, s, now) + assertCorruptCredentialsUnchanged(t, r, shared) + }) + } +} + +func corruptCredentials(corrupt string, shared *corev1.Secret, b wire.KeyringBundle, s RotationState, now *time.Time) { + switch corrupt { + case "timestamp": + s.NextRotation = time.Time{} + shared.Data["rotation.json"], _ = json.Marshal(s) + case "transition mismatch": + // Preparation must activate strictly after the rotation timestamp. + s.ActivateAt = s.NextRotation + shared.Data["rotation.json"], _ = json.Marshal(s) + case "missing activation": + s.ActivateAt = time.Time{} + shared.Data["rotation.json"], _ = json.Marshal(s) + case "missing prepared issuer": + s.PreparedIssuer = "" + shared.Data["rotation.json"], _ = json.Marshal(s) + case "zero root retirement", "missing root retirement", "replaced retirement": + corruptRootRetirement(corrupt, b, &s) + + shared.Data["rotation.json"], _ = json.Marshal(s) + case "zero key retirement", "missing key retirement": + // Symmetric retirement metadata is never valid, with or without a deadline. + s.Retiring[keyID(b.CacheKeys[0])] = time.Time{} + if corrupt == "missing key retirement" { + s.Retiring[keyID(b.CacheKeys[0])] = s.NextRotation + } + + shared.Data["rotation.json"], _ = json.Marshal(s) + case "nil retirement": + s.Retiring = nil + shared.Data["rotation.json"], _ = json.Marshal(s) + case "unknown retirement": + s.Retiring["unknown"] = s.NextRotation.Add(time.Hour) + shared.Data["rotation.json"], _ = json.Marshal(s) + case "active issuer": + s.ActiveIssuer = "missing" + shared.Data["rotation.json"], _ = json.Marshal(s) + case "bundle": + shared.Data["bundle.json"] = []byte("{}") + case "generation": + b.Generation = math.MaxUint64 + shared.Data["bundle.json"], _ = wire.EncodeBundle(b) + *now = s.NextRotation + case "binding": + shared.Annotations[credentialClaim] = "foreign" + case "private key": + shared.Data["issuer.json"] = []byte("{}") + } +} + +func prepareCorruptionPhase(t *testing.T, r *credentialsFixture, now *time.Time, corrupt string) { + t.Helper() + + switch corrupt { + case "transition mismatch", "missing activation", "missing prepared issuer", "zero root retirement", "zero key retirement", "missing root retirement", "missing key retirement", "replaced retirement", "unknown retirement": + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + } + + switch corrupt { + case "zero root retirement", "zero key retirement", "missing root retirement", "missing key retirement", "replaced retirement", "unknown retirement": + _, _, prepared, _ := keyState(t, r) + *now = prepared.ActivateAt + + runKeys(t, r) + } +} + +func corruptRootRetirement(corrupt string, b wire.KeyringBundle, s *RotationState) { + for _, root := range b.PeerTrustRoots { + id := rootID(root) + if id == s.ActiveIssuer { + continue + } + + switch corrupt { + case "zero root retirement": + s.Retiring[id] = time.Time{} + case "missing root retirement": + delete(s.Retiring, id) + case "replaced retirement": + s.Retiring[s.ActiveIssuer] = s.Retiring[id] + delete(s.Retiring, id) + } + } +} + +func assertCorruptCredentialsUnchanged(t *testing.T, r *credentialsFixture, shared *corev1.Secret) { + t.Helper() + require.NoError(t, r.Update(t.Context(), shared)) + before := shared.DeepCopy() + _, err := r.operations().ReconcileCredentials(t.Context()) + require.Error(t, err, "corrupt state accepted") + require.False(t, trustReady(r.Trust), "corrupt state trusted") + + after := &corev1.Secret{} + require.NoError(t, r.APIReader.Get(t.Context(), client.ObjectKeyFromObject(shared), after)) + require.Equal(t, before.Data, after.Data, "corrupt state rewritten") +} + +func TestKeyringExhaustedGenerationTransitions(t *testing.T) { + for _, transition := range []string{"idle", "admission", "stage", "empty stage", "activate", "remove", "prune"} { + t.Run(transition, func(t *testing.T) { + r, now := testKeyring(t) + + r.Config.Rotation.Interval = 7 * 24 * time.Hour + if transition == "empty stage" { + require.NoError(t, r.Delete(t.Context(), &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache"}})) + } + + runKeys(t, r) + + _, _, initial, _ := keyState(t, r) + if transition == "activate" || transition == "prune" { + *now = initial.NextRotation + + runKeys(t, r) + _, _, staged, _ := keyState(t, r) + *now = staged.ActivateAt + + if transition == "prune" { + runKeys(t, r) + _, _, active, _ := keyState(t, r) + *now = active.Retiring[initial.ActiveIssuer] + } + } + + shared, b, _, _ := keyState(t, r) + b.Generation = math.MaxUint64 + + var err error + + shared.Data["bundle.json"], err = wire.EncodeBundle(b) + require.NoError(t, err) + require.NoError(t, r.Update(t.Context(), shared)) + + switch transition { + case "admission": + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "added", UID: types.UID(testOtherUID)}} + require.NoError(t, r.Create(t.Context(), cache)) + case "stage", "empty stage": + *now = initial.NextRotation + case "remove": + require.NoError(t, r.Delete(t.Context(), &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache"}})) + } + + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + t.Fatal("exhausted generation wrote durable state") + return nil + }}) + if transition == "idle" { + runKeys(t, r) + } else if _, err := r.operations().ReconcileCredentials(t.Context()); !errors.Is(err, wire.Unavailable) { + t.Fatalf("exhausted generation transition: %v", err) + } + + after, preserved, _, _ := keyState(t, r) + require.Equal(t, shared.ResourceVersion, after.ResourceVersion, "exhausted generation wrote credentials") + require.Equal(t, b, preserved, "exhausted generation changed the published bundle") + }) + } +} + +func TestKeyringAdmissionAtLastGeneration(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + shared, b, _, _ := keyState(t, r) + b.Generation = math.MaxUint64 - 1 + + var err error + + shared.Data["bundle.json"], err = wire.EncodeBundle(b) + if err != nil { + t.Fatal(err) + } + + if err := r.Update(t.Context(), shared); err != nil { + t.Fatal(err) + } + + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "added", UID: types.UID(testOtherUID)}} + if err := r.Create(t.Context(), cache); err != nil { + t.Fatal(err) + } + + runKeys(t, r) + + _, admitted, _, _ := keyState(t, r) + if admitted.Generation != math.MaxUint64 || len(admitted.CacheKeys) != 4 || !reflect.DeepEqual(admitted.CacheKeys[:2], b.CacheKeys) { + t.Fatal("last generation admission lost existing keys or new scopes") + } + + for _, key := range admitted.CacheKeys[2:] { + if key.Key.Cache != wire.CacheID(testOtherUID) || key.State != wire.ActiveKey || binary.BigEndian.Uint64(key.Key.ID[4:12]) != math.MaxUint64 { + t.Fatal("admitted key did not bind the last publication generation") + } + } + + runKeys(t, r) +} + +func TestKeyringOversizedOverlapDoesNotWrite(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, initial, _, _ := keyState(t, r) + + for n := range 800 { + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: fmt.Sprintf("cache-%d", n), UID: types.UID(fmt.Sprintf("%08x-0000-0000-0000-000000000000", n+1))}} + if err := r.Create(context.Background(), cache); err != nil { + t.Fatal(err) + } + } + + runKeys(t, r) + shared, admitted, state, _ := keyState(t, r) + + capacity, err := catalogCapacity(r.Config, admitted) + if err != nil || len(admitted.CacheKeys) != 2*capacity || capacity >= 801 || !trustReady(r.Trust) { + t.Fatalf("rotation capacity not enforced: %d, %v", capacity, err) + } + + for _, key := range initial.CacheKeys { + found := false + for _, accepted := range admitted.CacheKeys { + found = found || key.EqualMaterial(accepted) + } + + if !found { + t.Fatal("growth evicted an existing cache key") + } + } + + base := r.Client + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + t.Fatal("rejected growth wrote durable state") + return nil + }}) + runKeys(t, r) + + unchanged, _, _, _ := keyState(t, r) + if unchanged.ResourceVersion != shared.ResourceVersion { + t.Fatal("rejected growth consumed generation") + } + + r.Client = base + *now = state.NextRotation + + runKeys(t, r) + + _, staged, _, _ := keyState(t, r) + if len(staged.CacheKeys) != 4*capacity || !trustReady(r.Trust) { + t.Fatal("admitted catalog could not rotate") + } +} + +func TestKeyringEmptyCatalogAndCacheAddedDuringPreparation(t *testing.T) { + r, now := testKeyring(t) + + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache"}} + if err := r.Delete(context.Background(), cache); err != nil { + t.Fatal(err) + } + + runKeys(t, r) + + _, b, s, _ := keyState(t, r) + if b.Generation != 1 || len(b.CacheKeys) != 0 { + t.Fatal("empty catalog initialization has wrong generation or keys") + } + + *now = s.NextRotation + + runKeys(t, r) + _, _, staged, _ := keyState(t, r) + + cache.UID = types.UID(testNodeUID) + if err := r.Create(context.Background(), cache); err != nil { + t.Fatal(err) + } + + runKeys(t, r) + + _, added, _, _ := keyState(t, r) + if added.Generation != 3 || len(added.CacheKeys) != 2 { + t.Fatal("new cache missing initial keys") + } + + for _, key := range added.CacheKeys { + if binary.BigEndian.Uint64(key.Key.ID[4:12]) != uint64(added.Generation) { + t.Fatal("new cache key did not bind its admission generation") + } + } + + *now = staged.ActivateAt + + runKeys(t, r) + + _, activated, _, _ := keyState(t, r) + for n, key := range activated.CacheKeys { + if key.State != wire.ActiveKey || !reflect.DeepEqual(key.Key, added.CacheKeys[n].Key) || !key.EqualMaterial(added.CacheKeys[n]) { + t.Fatal("new cache key retired without replacement") + } + } +} + +func TestKeyringRotationCrashRecovery(t *testing.T) { + for _, phase := range []string{"stage", "activate", "prune"} { + for _, failure := range []string{"before issuer", "after issuer", "before bundle", "after bundle"} { + t.Run(phase+"/"+failure, func(t *testing.T) { + exerciseRotationCrash(t, phase, failure) + }) + } + } +} + +func exerciseRotationCrash(t *testing.T, phase, failure string) { + t.Helper() + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + if phase != "stage" { + runKeys(t, r) + _, _, staged, _ := keyState(t, r) + *now = staged.ActivateAt + + if phase == "prune" { + runKeys(t, r) + *now = now.Add(r.Config.Rotation.RetainFor) + } + } + + base := r.Client.(client.WithWatch) + boom := errors.New("lost response") + failed := false + original, _, _, _ := keyState(t, r) + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + require.Equal(t, r.Config.CredentialsSecretName, obj.GetName(), "rotation wrote outside the credentials CAS") + + if strings.HasPrefix(failure, "before ") && !failed { + failed = true + return boom + } + + if err := c.Update(ctx, obj, opts...); err != nil { + return err + } + + if strings.HasPrefix(failure, "after ") && !failed { + failed = true + return boom + } + + return nil + }}) + + _, err := r.operations().ReconcileCredentials(t.Context()) + + require.True(t, failed, "write failure not injected") + require.ErrorIs(t, err, boom, "write failure accepted") + require.False(t, trustReady(r.Trust), "failed write trusted") + // Former two-Secret boundaries now fail one coherent atomic version. + committed, before, beforeState, _ := keyState(t, r) + if strings.HasPrefix(failure, "before ") && !reflect.DeepEqual(original.Data, committed.Data) { + t.Fatal("failed CAS partially changed credentials") + } + + recovered := Assemble(r.Config, base, base).Keyring + recovered.Now = r.Now + runKeys(t, recovered) + + _, after, afterState, _ := keyState(t, recovered) + require.GreaterOrEqual(t, after.Generation, before.Generation, "generation reset") + + if strings.HasPrefix(failure, "after ") { + require.Equal(t, before.Generation, after.Generation, "committed atomic generation replaced on recovery") + require.Equal(t, beforeState, afterState, "committed atomic state replaced on recovery") + } + + if !beforeState.ActivateAt.IsZero() && !afterState.ActivateAt.IsZero() && !beforeState.ActivateAt.Equal(afterState.ActivateAt) { + t.Fatal("committed preparation restarted") + } +} + +func TestKeyringPrivatePruneRecovery(t *testing.T) { + for _, afterWrite := range []bool{false, true} { + t.Run(fmt.Sprintf("response-lost=%t", afterWrite), func(t *testing.T) { + r, now := testKeyring(t) + r.Config.Rotation.Interval = 7 * 24 * time.Hour + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + _, _, staged, _ := keyState(t, r) + *now = staged.ActivateAt + + runKeys(t, r) + _, _, active, _ := keyState(t, r) + *now = active.Retiring[initial.ActiveIssuer] + base := r.Client.(client.WithWatch) + boom := errors.New("private prune interrupted") + + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if obj.GetName() != r.Config.CredentialsSecretName { + return c.Update(ctx, obj, opts...) + } + + _, b, _, _ := keyState(t, r) + require.True(t, containsRoot(b, initial.ActiveIssuer), "root changed before atomic pruning") + + if afterWrite { + if err := c.Update(ctx, obj, opts...); err != nil { + return err + } + } + + return boom + }}) + _, err := r.operations().ReconcileCredentials(t.Context()) + require.ErrorIs(t, err, boom, "prune interruption") + + _, before, _, beforeMaterial := keyState(t, r) + if containsRoot(before, initial.ActiveIssuer) == afterWrite || (len(beforeMaterial.Keys) == 2) == afterWrite { + t.Fatal("root and private key cleanup were not atomic") + } + + r.Client = base + runKeys(t, r) + + _, after, _, material := keyState(t, r) + + wantGeneration := before.Generation + if !afterWrite { + wantGeneration++ + } + + if wantGeneration != after.Generation || len(material.Keys) != 1 || containsRoot(after, initial.ActiveIssuer) { + t.Fatal("prune recovery changed publication or retained private material") + } + }) + } +} + +func TestKeyringBundleConflictKeepsPendingIssuer(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + base := r.Client.(client.WithWatch) + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if obj.GetName() == r.Config.CredentialsSecretName { + return apierrors.NewConflict(corev1.Resource("secrets"), obj.GetName(), wire.Conflict) + } + + return c.Update(ctx, obj, opts...) + }}) + + result, err := r.operations().ReconcileCredentials(t.Context()) + if !apierrors.IsConflict(err) || result != 0 { + t.Fatalf("bundle conflict: %v, %v", result, err) + } + + _, b, _, material := keyState(t, r) + if b.Generation != 1 || len(material.Keys) != 1 || !containsRoot(b, initial.ActiveIssuer) { + t.Fatal("conflicting bundle became authoritative") + } + + r.Client = base + runKeys(t, r) + + _, b, state, _ := keyState(t, r) + if state.PreparedIssuer == "" || state.PreparedIssuer == initial.ActiveIssuer || b.Generation != 2 { + t.Fatal("conflict recovery did not publish a coherent preparation") + } +} + +func TestKeyringInitializationNeverResurrects(t *testing.T) { + for _, failure := range []string{"before claim", "after claim", "before issuer", "after issuer", "before bundle", "after bundle"} { + t.Run(failure, func(t *testing.T) { + r, _ := testKeyring(t) + base := r.Client.(client.WithWatch) + boundary := strings.NewReplacer("claim", "commit", "issuer", "create", "bundle", "create").Replace(failure) + r.Client = interruptInitialization(base, boundary, func() {}) + _, err := r.operations().ReconcileCredentials(t.Context()) + require.ErrorIs(t, err, errInitializationInterrupted, "crash not injected") + + recovered := Assemble(r.Config, base, base).Keyring + recovered.Now = r.Now + + _, err = recovered.operations().ReconcileCredentials(t.Context()) + require.NoError(t, err, "staged initialization must recover an uncommitted candidate") + version, _, err := readVersion(t.Context(), base, r.Config) + require.NoError(t, err) + + secret := &corev1.Secret{} + require.NoError(t, base.Get(t.Context(), client.ObjectKey{Namespace: r.Config.Namespace, Name: r.Config.CredentialsSecretName}, secret)) + require.NotEmpty(t, secret.UID) + require.Equal(t, string(secret.UID), version.Annotations[credentialUID]) + }) + } + + for _, lost := range []string{"issuer", "bundle", "both", "version", "marker"} { + t.Run("lost "+lost, func(t *testing.T) { + r, _ := testKeyring(t) + issuer := testIssuer(r) + runKeys(t, r) + + if lost == "both" || lost == "issuer" || lost == "bundle" { + require.NoError(t, r.Delete(context.Background(), &corev1.Secret{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: r.Config.CredentialsSecretName}})) + } + + if lost == "version" || lost == "marker" { + name := r.Config.VersionConfigMapName + if lost == "marker" { + name = r.Config.InstallationConfigMapName + } + + require.NoError(t, r.Delete(context.Background(), &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: name}})) + } + + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + t.Fatal("recreated established state") + return nil + }}) + _, err := r.operations().ReconcileCredentials(t.Context()) + require.Error(t, err, "lost state accepted") + require.False(t, trustReady(r.Trust)) + + _, err = issuer.TrustRoots(context.Background()) + require.Error(t, err, "lost state still trusted") + }) + } +} + +func TestKeyringConflictCancellationAndAuthoritativeReads(t *testing.T) { + for _, cancelAt := range []string{"none", "before", "issuer", "bundle"} { + t.Run(cancelAt, func(t *testing.T) { + exerciseKeyringCancellation(t, cancelAt) + }) + } +} + +func exerciseKeyringCancellation(t *testing.T, cancelAt string) { + t.Helper() + r, now := testKeyring(t) + runKeys(t, r) + _, _, s, _ := keyState(t, r) + *now = s.NextRotation + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + base := r.Client.(client.WithWatch) + writes := 0 + r.Client = interceptor.NewClient(base, interceptor.Funcs{ + Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + t.Fatal("cached credential read") + return nil + }, + List: func(context.Context, client.WithWatch, client.ObjectList, ...client.ListOption) error { + t.Fatal("cached catalog read") + return nil + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + writes++ + + if (cancelAt == "issuer" || cancelAt == "bundle") && obj.GetName() == r.Config.CredentialsSecretName { + cancel() + } + + if cancelAt == "none" { + return apierrors.NewConflict(corev1.Resource("secrets"), obj.GetName(), wire.Conflict) + } + + return c.Update(ctx, obj, opts...) + }, + }) + + if cancelAt == "before" { + cancel() + } + + result, err := r.operations().ReconcileCredentials(ctx) + if cancelAt == "none" { + require.True(t, apierrors.IsConflict(err), "conflict must reach the controller queue") + require.Zero(t, result, "conflict supplied a fixed retry delay") + } else { + require.ErrorIs(t, err, context.Canceled) + + if cancelAt == "before" { + require.Zero(t, result, "failed admission must not plan rotation") + } else { + require.Equal(t, r.Config.Rotation.PrepareFor, result, "committed preparation retains its deadline; root discards it on cancellation") + } + } + + if cancelAt == "before" && writes != 0 || cancelAt == "issuer" && writes != 1 { + t.Fatal("write after cancellation") + } + + // Cancellation before admission observes no authority failure. + if trustReady(r.Trust) != (cancelAt == "before") { + t.Fatal("readiness did not reflect whether admission observed a failure") + } + + r.Client = base + runKeys(t, r) +} + +func TestKeyringExpiredPreparationRecovery(t *testing.T) { + for _, prepared := range []bool{false, true} { + t.Run(fmtBool(prepared), func(t *testing.T) { + r, now := testKeyring(t) + issuer := testIssuer(r) + runKeys(t, r) + + _, _, initial, _ := keyState(t, r) + if prepared { + *now = initial.NextRotation + + runKeys(t, r) + } + + *now = now.Add(30 * 24 * time.Hour) + // An expired active issuer cannot sign, but rotation can recover by + // staging fresh trust and waiting the complete preparation interval. + if _, err := issuer.TrustRoots(context.Background()); !errors.Is(err, wire.Unavailable) { + t.Fatalf("expired issuer accepted: %v", err) + } + + if _, err := r.operations().ReconcileCredentials(t.Context()); !errors.Is(err, wire.Unavailable) || trustReady(r.Trust) { + t.Fatalf("expired active readiness: %v", err) + } + + _, _, staged, _ := keyState(t, r) + if !staged.ActivateAt.Equal(now.Add(r.Config.Rotation.PrepareFor)) || staged.ActiveIssuer != initial.ActiveIssuer { + t.Fatal("recovery skipped preparation") + } + + *now = staged.ActivateAt + + runKeys(t, r) + + if _, err := issuer.TrustRoots(context.Background()); err != nil { + t.Fatal(err) + } + }) + } +} + +func fmtBool(v bool) string { + if v { + return "prepared" + } + + return "active" +} + +func TestReplicaAuthenticationBindings(t *testing.T) { + for _, scenario := range []string{"valid", "missing bearer", "review unavailable", "wrong account", "missing binding", "missing pod", "wrong pod uid", "failed pod", "missing account", "wrong account uid", "invalid expiration"} { + t.Run(scenario, func(t *testing.T) { + f, status, token := authenticatedBootstrapFixture(t) + a := f.a.authority + base := f.a.Topology.Client.(client.WithWatch) + a.config.ControllerServiceAccount = a.config.DataplaneServiceAccount + status.Audiences = []string{ReplicationAudience} + want := mutateReplicaBinding(t, a, base, scenario, &status, &token) + a.client = interceptor.NewClient(base, interceptor.Funcs{Create: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.CreateOption) error { + if scenario == "review unavailable" { + return wire.Unavailable + } + + review := obj.(*authv1.TokenReview) + require.Equal(t, []string{ReplicationAudience}, review.Spec.Audiences) + review.Status = status + + return nil + }}) + + request := httptest.NewRequest("GET", "/replica", nil) + if scenario != "missing bearer" { + request.Header.Set("Authorization", "Bearer "+token) + } + + identity, err := a.AuthenticateReplica(t.Context(), request) + if want != nil { + require.ErrorIs(t, err, want) + require.Equal(t, ReplicaIdentity{}, identity) + + return + } + + require.NoError(t, err) + require.Equal(t, "pod-uid", identity.UID()) + require.True(t, identity.Expires().After(time.Now())) + require.Same(t, a, identity.owner) + }) + } +} + +func mutateReplicaBinding(t *testing.T, a *Authority, c client.Client, scenario string, status *authv1.TokenReviewStatus, token *string) error { + t.Helper() + + switch scenario { + case "valid": + return nil + case "missing bearer", "missing binding", "invalid expiration": + if scenario == "missing binding" { + status.User.UID = "" + } + + if scenario == "invalid expiration" { + *token = "invalid" + } + + return wire.Unauthenticated + case "review unavailable": + return wire.Unavailable + case "wrong account": + status.User.Username = "other" + case "missing pod": + status.User.Extra["authentication.kubernetes.io/pod-name"] = authv1.ExtraValue{"missing"} + case "wrong pod uid": + status.User.Extra["authentication.kubernetes.io/pod-uid"] = authv1.ExtraValue{"other"} + case "failed pod": + var pod corev1.Pod + require.NoError(t, c.Get(t.Context(), client.ObjectKey{Namespace: a.config.Namespace, Name: singleExtra(status.User, "pod-name")}, &pod)) + pod.Status.Phase = corev1.PodFailed + require.NoError(t, c.Status().Update(t.Context(), &pod)) + case "missing account": + require.NoError(t, c.Delete(t.Context(), &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: a.config.Namespace, Name: a.config.ControllerServiceAccount}})) + case "wrong account uid": + status.User.UID = "other" + } + + return wire.Forbidden +} + +func TestPreparedIssuerCoversReplacementActivation(t *testing.T) { + for _, early := range []bool{false, true} { + for _, margin := range []time.Duration{-time.Second, 0, time.Second} { + t.Run(fmt.Sprintf("early=%v/margin=%s", early, margin), func(t *testing.T) { + exercisePreparedIssuerHorizon(t, early, margin) + }) + } + } +} + +func exercisePreparedIssuerHorizon(t *testing.T, early bool, margin time.Duration) { + t.Helper() + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + runKeys(t, r) + _, bundle, state, material := keyState(t, r) + oldID := state.PreparedIssuer + + activation := state.ActivateAt + if !early { + // Late activation must schedule the next cycle from actual time. + *now = state.ActivateAt.Add(2 * time.Hour) + activation = *now + } + + short := editSigningCertificate(t, material.Keys[oldID], func(cert *x509.Certificate) { + cert.NotAfter = activation.Add(r.Config.Rotation.Interval + r.Config.CertificateLifetime + margin) + }) + shortID := rootID(short.Certificate) + + rebindSigning(&bundle, &state, material, oldID, short, true) + + bundle.Generation++ + writeSigningCredentials(t, r, bundle, state, material) + runKeys(t, r) + + _, after, next, keys := keyState(t, r) + if margin < 0 { + if containsRoot(after, shortID) || next.PreparedIssuer == shortID || next.ActiveIssuer != initial.ActiveIssuer || !next.ActivateAt.Equal(now.Add(r.Config.Rotation.PrepareFor)) { + t.Fatal("insufficient signing horizon did not restart preparation") + } + + require.NotContains(t, keys.Keys, shortID, "stale private material retained") + } else if early { + if next.PreparedIssuer != shortID || !next.ActivateAt.Equal(state.ActivateAt) || after.Generation != bundle.Generation { + t.Fatal("usable preparation changed before activation") + } + } else if next.ActiveIssuer != shortID || next.PreparedIssuer != "" { + t.Fatal("usable prepared issuer did not activate") + } + + if next.PreparedIssuer != "" { + *now = next.ActivateAt + + runKeys(t, r) + } + + _, _, active, _ := keyState(t, r) + *now = active.NextRotation + + runKeys(t, r) + _, _, replacement, _ := keyState(t, r) + *now = replacement.ActivateAt.Add(-time.Second) + + identity, request, _ := issuanceRequest(t, r) + if _, err := testIssuer(r).Issue(t.Context(), identity, request); err != nil { + t.Fatalf("issuer failed before replacement activation: %v", err) + } + + *now = replacement.ActivateAt + + runKeys(t, r) +} diff --git a/internal/racer/authority/initialization_test.go b/internal/racer/authority/initialization_test.go new file mode 100644 index 000000000..c150c5603 --- /dev/null +++ b/internal/racer/authority/initialization_test.go @@ -0,0 +1,83 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + corev1 "k8s.io/api/core/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestInitializationRequiresStagedProtocol(t *testing.T) { + for _, state := range []string{"fresh", "consumed"} { + for _, protocol := range []string{"", "staged-v2"} { + t.Run(state+"/"+protocol, func(t *testing.T) { + r := testTopology(t) + if state == "consumed" { + require.NoError(t, r.authority.Recover(t.Context(), r.Client)) + } + + reader := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + if err == nil && key.Name == r.Config.InstallationConfigMapName { + obj.(*corev1.ConfigMap).Data[markerInitializationProtocol] = protocol + } + + return err + }}) + writer := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + t.Fatal("unsupported protocol created authority") + return nil + }, + Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + t.Fatal("unsupported protocol changed authority") + return nil + }, + }) + a := New(r.Config, Dependencies{Reader: reader, Writer: writer}) + require.ErrorIs(t, a.Recover(t.Context(), writer), wire.Unavailable) + _, _, err := readVersion(t.Context(), reader, r.Config) + require.ErrorIs(t, err, wire.Unavailable) + }) + } + } +} + +func TestCommittedCredentialsRequireUIDAndInstallationBinding(t *testing.T) { + for _, corruption := range []string{"missing uid", "wrong uid", "missing protocol", "wrong installation"} { + t.Run(corruption, func(t *testing.T) { + r, _ := testKeyring(t) + runKeys(t, r) + reader := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + if err == nil && key.Name == r.Config.CredentialsSecretName { + switch corruption { + case "missing uid": + obj.SetUID("") + case "wrong uid": + obj.SetUID("replacement") + case "missing protocol": + delete(obj.GetAnnotations(), initializationProtocol) + case "wrong installation": + obj.GetAnnotations()[installationUIDAnnotation] = "replacement" + } + } + + return err + }}) + issuer := testIssuer(r) + issuer.APIReader = reader + _, err := issuer.TrustRoots(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + require.False(t, trustReady(r.Trust)) + }) + } +} diff --git a/internal/racer/authority/issuance_test.go b/internal/racer/authority/issuance_test.go new file mode 100644 index 000000000..25d4e4c7e --- /dev/null +++ b/internal/racer/authority/issuance_test.go @@ -0,0 +1,181 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "context" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + corev1 "k8s.io/api/core/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestIssuanceReadCancellationPreservesTrust(t *testing.T) { + for _, deadline := range []bool{false, true} { + for _, boundary := range []string{"marker", "version", "credentials"} { + for _, successfulRead := range []bool{false, true} { + t.Run(boundary+"/"+map[bool]string{false: "cancel", true: "deadline"}[deadline]+"/"+map[bool]string{false: "failed read", true: "successful read"}[successfulRead], func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + identity := pollIdentity(a.config, testNodeUID) + identity.owner, identity.bearer = a, true + confirmed, epoch, bundle := a.trust.confirmed, a.trust.authority, a.trust.bundle + publicationConfirmed := a.publications.confirmed + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + want := context.Canceled + + if deadline { + var stop context.CancelFunc + + ctx, stop = context.WithTimeout(ctx, time.Second) + defer stop() + + want = context.DeadlineExceeded + } + + target := map[string]string{"marker": a.config.InstallationConfigMapName, "version": a.config.VersionConfigMapName, "credentials": a.config.CredentialsSecretName}[boundary] + reads := 0 + a.bootstrap.Issuer.APIReader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + if key.Name != target { + return err + } + + reads++ + + if deadline { + time.Sleep(time.Second) + } else { + cancel() + } + + if !successfulRead { + return ctx.Err() + } + + return err + }, + }) + encoded, err := a.Issue(ctx, identity, f.request) + require.ErrorIs(t, err, want) + require.Empty(t, encoded) + require.Equal(t, 1, reads) + require.NoError(t, a.TrustReady()) + require.NoError(t, a.PublicationReady()) + require.Equal(t, confirmed, a.trust.confirmed, "issuance must not refresh trust") + require.Equal(t, publicationConfirmed, a.publications.confirmed) + require.Same(t, bundle, a.trust.bundle) + require.Equal(t, epoch, a.trust.authority) + require.NoError(t, epoch.Err()) + + gateCtx, stop := context.WithTimeout(t.Context(), time.Second) + defer stop() + + require.NoError(t, a.gate.Acquire(gateCtx), "issuance must release its gate") + a.gate.Release() + }) + }) + } + } + } +} + +func TestIssuanceInvalidEvidenceWithCanceledRequestRevokesTrust(t *testing.T) { + for _, boundary := range []string{"marker", "version", "credentials"} { + t.Run(boundary, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + identity := pollIdentity(a.config, testNodeUID) + identity.owner, identity.bearer = a, true + epoch := a.trust.authority + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + target := map[string]string{"marker": a.config.InstallationConfigMapName, "version": a.config.VersionConfigMapName, "credentials": a.config.CredentialsSecretName}[boundary] + a.bootstrap.Issuer.APIReader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + if key.Name == target { + switch obj := obj.(type) { + case *corev1.ConfigMap: + obj.Data = nil + case *corev1.Secret: + obj.Data = nil + } + + cancel() + } + + return err + }, + }) + encoded, err := a.Issue(ctx, identity, f.request) + require.Error(t, err) + require.NotErrorIs(t, err, context.Canceled, "invalid evidence must not be replaced by request cancellation") + require.Empty(t, encoded) + require.ErrorIs(t, a.TrustReady(), wire.Unavailable) + require.ErrorIs(t, epoch.Err(), context.Canceled) + }) + } +} + +func TestIssuanceStormRetainsAuthoritativeReads(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + identity := pollIdentity(a.config, testNodeUID) + identity.owner, identity.bearer = a, true + confirmed := a.trust.confirmed + + var reads atomic.Int64 + + a.bootstrap.Issuer.APIReader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + reads.Add(1) + return c.Get(ctx, key, obj, opts...) + }, + }) + + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + + const workers, requests = 8, 8 + + errors := make(chan error, workers*requests) + start := time.Now() + + var wg sync.WaitGroup + for range workers { + wg.Go(func() { + for range requests { + _, err := a.Issue(ctx, identity, f.request) + errors <- err + } + }) + } + + wg.Wait() + close(errors) + + for err := range errors { + require.NoError(t, err) + } + + require.EqualValues(t, workers*requests*3, reads.Load(), "every issuance must read marker, version, and credentials") + require.Equal(t, confirmed, a.trust.confirmed) + t.Logf("%d issuances, %d workers, %d authoritative reads, elapsed %s (in-memory API fixture)", workers*requests, workers, reads.Load(), time.Since(start)) +} diff --git a/internal/racer/authority/publications_test.go b/internal/racer/authority/publications_test.go new file mode 100644 index 000000000..00cdbcefd --- /dev/null +++ b/internal/racer/authority/publications_test.go @@ -0,0 +1,1723 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "reflect" + "strings" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/google/uuid" + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func pollIdentity(cfg Config, node wire.NodeID) NodeIdentity { + return NodeIdentity{cluster: cfg.Cluster, node: node, expires: time.Now().Add(time.Hour)} +} + +func TestPublicationDeltaSelectionAndDisconnectedFallback(t *testing.T) { + r := initializedTopology(t) + ctx := context.Background() + reconcileTopology(t, r, ctx) + + members := AcceptedMembers{} + + for i := range 100 { + id := wire.NodeID(fmt.Sprintf("22222222-2222-4222-8222-%012d", i)) + members[id] = wire.Member{Node: id, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}} + } + + install := func() *CommittedPublication { + t.Helper() + + cm, previous, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + prepared, err := r.Publications.Prepare(previous, cm.ResourceVersion, members, nil) + if err != nil { + t.Fatal(err) + } + + committed, err := r.CommitVersion(ctx, prepared) + if err != nil { + t.Fatal(err) + } + + if err := r.Publications.Install(committed); err != nil { + t.Fatal(err) + } + + return committed + } + base := install() + id := wire.NodeID("22222222-2222-4222-8222-000000000000") + m := members[id] + m.Shares = 9 + members[id] = m + next := install() + + delta := next.ForBase(base.record.Sequence, base.record.ContentHash) + if len(delta.encoded) >= len(next.encoded) { + t.Fatal("delta was not selected") + } + + decoded, err := wire.DecodePublication(strings.NewReader(base.encoded)) + if err != nil { + t.Fatal(err) + } + + applied, err := wire.ApplyDelta(decoded, strings.NewReader(delta.encoded)) + if err != nil { + t.Fatal(err) + } + + hash, _, err := wire.ContentHashes(applied) + if err != nil || hash != next.record.ContentHash { + t.Fatalf("target mismatch: %v", err) + } + + if next.ForBase(base.record.Sequence, "missing").encoded != next.encoded || next.ForBase(0, "").encoded != next.encoded { + t.Fatal("missing-base fallback failed") + } + + for _, sequence := range []wire.Sequence{0, base.record.Sequence - 1, base.record.Sequence + 1} { + require.Equal(t, next.encoded, next.ForBase(sequence, base.record.ContentHash).encoded, "matching hash cannot authorize a delta from another sequence") + } +} + +func TestPublicationCurrentAndSubscribe(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + p := r.Publications + + leader, cancel := context.WithCancel(t.Context()) + defer cancel() + + current, changed, err := p.CurrentAndSubscribe() + if current != nil || !errors.Is(err, wire.Unavailable) { + t.Fatalf("empty publication: %p, %v", current, err) + } + + installed := reconcileTopology(t, r, leader) + // Install between the snapshot and waiting must close the captured channel. + <-changed + + current, changed, err = p.CurrentAndSubscribe() + if current != installed || err != nil { + t.Fatalf("installed publication: %p, %v", current, err) + } + + go p.Suspend() + + <-changed + + current, changed, err = p.CurrentAndSubscribe() + if current != nil || !errors.Is(err, wire.Unavailable) { + t.Fatalf("suspended publication: %p, %v", current, err) + } + + replayed := make(chan error, 1) + + go func() { replayed <- p.Install(installed) }() + + <-changed + + if err := <-replayed; err != nil { + t.Fatal(err) + } + + current, changed, err = p.CurrentAndSubscribe() + if err != nil || current.encoded != installed.encoded || current.authority == installed.authority || installed.authority.Err() == nil { + t.Fatalf("resumed publication: %p, %v", current, err) + } + + select { + case <-changed: + t.Fatal("subscription returned an already-closed channel for unchanged state") + default: + } + + cancel() + + current, _, err = p.CurrentAndSubscribe() + if current != nil || !errors.Is(err, context.Canceled) { + t.Fatalf("lost leadership: %p, %v", current, err) + } + }) +} + +func TestPublicationBoundsOverflowAndInstallProof(t *testing.T) { + r := initializedTopology(t) + + p := reconcileTopology(t, r, context.Background()) + for _, invalid := range []*CommittedPublication{nil, {}, {owner: r.Publications}} { + if err := r.Publications.Install(invalid); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("forged install: %v", err) + } + } + + if err := NewPublications().Install(p); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("foreign install: %v", err) + } + + cache, err := BuildCatalog([]racerv1.ClusterCache{catalogCache("cache", testNodeUID)}) + if err != nil { + t.Fatal(err) + } + + previous := p.record + + previous.Sequence = ^wire.Sequence(0) + if _, err := r.Publications.Prepare(previous, "rv", nil, cache); !errors.Is(err, wire.Unavailable) { + t.Fatalf("overflow: %v", err) + } + + if _, err := r.Publications.Prepare(previous, "rv", nil, nil); err != nil { + t.Fatalf("unchanged maximum counter rejected: %v", err) + } + + previous.MembershipVersion = ^wire.MembershipVersion(0) + + members := AcceptedMembers{testNodeUID: {Node: testNodeUID, Shares: 4, PeerEndpoint: "192.0.2.1:8082"}} + if _, err := r.Publications.Prepare(previous, "rv", members, nil); !errors.Is(err, wire.Unavailable) { + t.Fatalf("membership overflow: %v", err) + } + + if _, err := r.Publications.Prepare(p.record, "rv", AcceptedMembers{testNodeUID: {Node: testOtherUID}}, nil); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("map identity: %v", err) + } + + oversized := members[testNodeUID] + + oversized.RDMANICs = []wire.RDMANIC{{Device: strings.Repeat("a", wire.MaxPublicationBytes), Port: 1}} + if _, err := r.Publications.Prepare(p.record, "rv", AcceptedMembers{testNodeUID: oversized}, nil); !errors.Is(err, wire.TooLarge) { + t.Fatalf("oversized candidate: %v", err) + } +} + +func TestPublicationDeepIsolationAndReplay(t *testing.T) { + r := initializedTopology(t) + ctx := context.Background() + old := reconcileTopology(t, r, ctx) + + cm, previous, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + numa := uint32(1) + members := AcceptedMembers{testNodeUID: {Node: testNodeUID, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{{Device: "mlx5_0", Port: 1, NUMANode: &numa}}}} + + prepared, err := r.Publications.Prepare(previous, cm.ResourceVersion, members, nil) + if err != nil { + t.Fatal(err) + } + + numa = 9 + + delete(members, testNodeUID) + + committed, err := r.CommitVersion(ctx, prepared) + if err != nil { + t.Fatal(err) + } + + if err := r.Publications.Install(committed); err != nil { + t.Fatal(err) + } + + decoded, err := wire.DecodePublication(strings.NewReader(committed.encoded)) + if err != nil || len(decoded.Members) != 1 || *decoded.Members[0].RDMANICs[0].NUMANode != 1 { + t.Fatalf("mutable alias: %+v, %v", decoded, err) + } + + if err := r.Publications.Install(old); !errors.Is(err, wire.Conflict) { + t.Fatalf("rollback: %v", err) + } + + conflicting := *committed + + conflicting.encoded = "different" + if err := r.Publications.Install(&conflicting); !errors.Is(err, wire.Conflict) { + t.Fatalf("conflicting replay: %v", err) + } + + if err := r.Publications.Install(committed); err != nil { + t.Fatalf("idempotent replay: %v", err) + } +} + +func TestPollValidationAndCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + + leader, loseLeadership := context.WithCancel(context.Background()) + defer loseLeadership() + + current := reconcileTopology(t, r, leader) + identity := pollIdentity(r.Config, testNodeUID) + sequence := current.record.Sequence + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + done := make(chan error, 1) + + go func() { _, err := r.Publications.Wait(ctx, identity, &sequence); done <- err }() + + synctest.Wait() + + cancel() + + if err := <-done; !errors.Is(err, context.Canceled) { + t.Fatalf("wait cancellation: %v", err) + } + + if got, err := r.Publications.Wait(ctx, identity, nil); got != nil || !errors.Is(err, context.Canceled) { + t.Fatalf("canceled immediate poll: %p, %v", got, err) + } + + if got, err := r.Publications.Wait(context.Background(), identity, nil); err != nil || got != current { + t.Fatalf("immediate shared snapshot: %p, %v", got, err) + } + + assertPollValidation(t, r, identity, sequence) + + go func() { _, err := r.Publications.Wait(context.Background(), identity, &sequence); done <- err }() + + synctest.Wait() + loseLeadership() + + if err := <-done; !errors.Is(err, context.Canceled) { + t.Fatalf("leadership did not cancel wait: %v", err) + } + + if _, err := r.Publications.Current(); !errors.Is(err, context.Canceled) { + t.Fatalf("old leadership still serves: %v", err) + } + + if got, err := r.Publications.Wait(context.Background(), identity, nil); got != nil || !errors.Is(err, context.Canceled) { + t.Fatalf("immediate poll after leadership loss: %p, %v", got, err) + } + + if _, _, err := current.admit(t.Context()); !errors.Is(err, context.Canceled) { + t.Fatalf("write after leadership: %v", err) + } + }) +} + +func assertPollValidation(t *testing.T, r *topologyFixture, identity NodeIdentity, sequence wire.Sequence) { + t.Helper() + + zero, future := wire.Sequence(0), sequence+1 + wrong, expired, invalid := identity, identity, identity + wrong.cluster = testNodeUID + expired.expires = time.Now().Add(-time.Second) + + invalid.node = "not-a-uuid" + for _, tc := range []struct { + name string + identity NodeIdentity + cursor *wire.Sequence + want error + }{ + {"zero cursor", identity, &zero, wire.Conflict}, + {"future cursor", identity, &future, wire.Unavailable}, + {"wrong cluster", wrong, nil, wire.Forbidden}, + {"expired identity", expired, nil, wire.Unauthenticated}, + {"invalid identity", invalid, nil, wire.Unauthenticated}, + } { + _, err := r.Publications.Wait(t.Context(), tc.identity, tc.cursor) + require.ErrorIs(t, err, tc.want, tc.name) + } +} + +func TestPollImmediateReturnsWithoutAllocations(t *testing.T) { + r := initializedTopology(t) + ctx := context.Background() + previous := reconcileTopology(t, r, ctx).record.Sequence + cache := catalogCache("cache", testNodeUID) + + if err := r.Create(ctx, &cache); err != nil { + t.Fatal(err) + } + + runKeys(t, Assemble(r.Config, r.Client, r.APIReader).Keyring) + + current := reconcileTopology(t, r, ctx) + identity := pollIdentity(r.Config, testNodeUID) + + for _, tc := range []struct { + name string + after *wire.Sequence + }{ + {name: "absent cursor"}, + {name: "older cursor", after: &previous}, + } { + t.Run(tc.name, func(t *testing.T) { + var ( + got *CommittedPublication + err error + ) + + allocations := testing.AllocsPerRun(100, func() { + got, err = r.Publications.Wait(ctx, identity, tc.after) + }) + + if err != nil || got != current { + t.Fatalf("immediate shared publication: %p, %v", got, err) + } + + if allocations != 0 { + t.Fatalf("immediate poll allocated: %v allocations per call", allocations) + } + }) + } +} + +func TestPollCertificateExpiration(t *testing.T) { + for _, name := range []string{"with context deadline", "without context deadline"} { + t.Run(name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + p := reconcileTopology(t, r, context.Background()) + identity := pollIdentity(r.Config, testNodeUID) + identity.expires = time.Now().Add(time.Second) + sequence := p.record.Sequence + ctx := context.Background() + + if name == "with context deadline" { + var cancel context.CancelFunc + + ctx, cancel = context.WithTimeout(ctx, 5*time.Second) + defer cancel() + } + + got, err := r.Publications.Wait(ctx, identity, &sequence) + if got != nil || !errors.Is(err, wire.Unauthenticated) || !time.Now().Equal(identity.expires) { + t.Fatalf("expiration: %p, %v, time %v, expires %v", got, err, time.Now(), identity.expires) + } + }) + }) + } +} + +func TestPollNormalTimeout(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + r.Publications.maxAge = 2 * wire.PollWait + p := reconcileTopology(t, r, context.Background()) + identity := pollIdentity(r.Config, testNodeUID) + sequence := p.record.Sequence + start := time.Now() + + got, err := r.Publications.Wait(context.Background(), identity, &sequence) + if err != nil || got != nil || time.Since(start) != wire.PollWait { + t.Fatalf("normal timeout: %p, %v, %v", got, err, time.Since(start)) + } + }) +} + +func TestDurableLossWithdrawsPublication(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + p := reconcileTopology(t, r, ctx) + sequence := p.record.Sequence + done := make(chan error, 1) + + go func() { + _, err := r.Publications.Wait(ctx, pollIdentity(r.Config, testNodeUID), &sequence) + done <- err + }() + + synctest.Wait() + + cm, _, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + if err := r.Delete(ctx, cm); err != nil { + t.Fatal(err) + } + + if _, err := r.operations().PublishTopology(ctx, r.observeTopology); err == nil { + t.Fatal("missing counter accepted") + } + + if err := <-done; !errors.Is(err, wire.Unavailable) { + t.Fatalf("waiting poll did not fail closed: %v", err) + } + + if _, err := r.Publications.Current(); !errors.Is(err, wire.Unavailable) { + t.Fatalf("durable loss still serves: %v", err) + } + }) +} + +func TestPollFanoutSharesOnePublication(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + current := reconcileTopology(t, r, ctx) + sequence := current.record.Sequence + + const count = 256 + + results := make(chan *CommittedPublication, count) + errors := make(chan error, count) + + var wg sync.WaitGroup + + // Direct waiters share a node deliberately: only the HTTP server owns admission. + for range count { + identity := pollIdentity(r.Config, testNodeUID) + + wg.Go(func() { p, err := r.Publications.Wait(ctx, identity, &sequence); results <- p; errors <- err }) + } + + synctest.Wait() + + cache := catalogCache("cache", testNodeUID) + if err := r.Create(ctx, &cache); err != nil { + t.Fatal(err) + } + + runKeys(t, Assemble(r.Config, r.Client, r.APIReader).Keyring) + + next := reconcileTopology(t, r, ctx) + + wg.Wait() + + for range count { + if err := <-errors; err != nil { + t.Fatal(err) + } + + if p := <-results; p != next { + t.Fatalf("waiter copied or missed publication: %p != %p", p, next) + } + } + }) +} + +type boundedWriter struct { + largest int + calls int + cancel context.CancelFunc +} + +func (w *boundedWriter) Write(p []byte) (int, error) { + w.largest = max(w.largest, len(p)) + + w.calls++ + if w.cancel != nil { + w.cancel() + } + + return len(p), nil +} + +func TestPublicationWriteScratchAndCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + p := publicationResponse{encoded: strings.Repeat("x", 100_000)} + if n, err := p.writeTo(ctx, shortPublicationWriter{}); !errors.Is(err, io.ErrShortWrite) || n != 1 { + t.Fatalf("short write: %d, %v", n, err) + } + + w := &boundedWriter{} + if n, err := p.writeTo(ctx, w); err != nil || n != 100_000 || w.largest > 32*1024 || w.calls != 4 { + t.Fatalf("unbounded write scratch: %d, %v, %+v", n, err, w) + } + + w = &boundedWriter{cancel: cancel} + if n, err := p.writeTo(ctx, w); !errors.Is(err, context.Canceled) || n != 32*1024 || w.calls != 1 { + t.Fatalf("write ignored cancellation: %d, %v", n, err) + } +} + +type shortPublicationWriter struct{} + +func (shortPublicationWriter) Write([]byte) (int, error) { return 1, nil } + +func TestPublicationAdmissionRequiresOwner(t *testing.T) { + for _, image := range []*CommittedPublication{nil, {}, {leadership: t.Context()}} { + if _, _, err := image.admit(t.Context()); !errors.Is(err, wire.Unavailable) { + t.Fatalf("ownerless write admitted: %v", err) + } + } +} + +func TestPublicationInstalledStateNeverExceedsHighWater(t *testing.T) { + for _, mutation := range []string{"cluster", "sequence", "membership", "same-sequence hash"} { + t.Run(mutation, func(t *testing.T) { + r := initializedTopology(t) + + image := reconcileTopology(t, r, t.Context()) + if r.Publications.observed != image.record { + t.Fatal("install did not observe record") + } + + newer := image.record + newer.Sequence += 2 + + newer.MembershipVersion++ + if err := r.Publications.confirm(newer); err != nil { + t.Fatal(err) + } + + next := *image + next.record = newer + + switch mutation { + case "cluster": + next.record.Cluster = testNodeUID + case "sequence": + next.record.Sequence-- + case "membership": + next.record.MembershipVersion-- + case "same-sequence hash": + next.record.ContentHash = strings.Repeat("a", 64) + } + + if err := r.Publications.Install(&next); !errors.Is(err, wire.Conflict) { + t.Fatalf("high-water bypass: %v", err) + } + + if r.Publications.current != image || r.Publications.observed != newer { + t.Fatal("rejected install mutated state") + } + }) + } +} + +func TestPrepareCanonicalEquivalence(t *testing.T) { + numa := uint32(3) + members := AcceptedMembers{ + testNodeUID: {Node: testNodeUID, Shares: 4, PeerEndpoint: "192.0.2.1:7443", RDMANICs: []wire.RDMANIC{{Rail: 2, Device: "β<&>", Port: 1, NUMANode: &numa}, {Rail: 1, Device: "mlx5_0", Port: 1}}}, + testOtherUID: {Node: testOtherUID, Shares: 1, PeerEndpoint: "[2001:db8::1]:7443"}, + } + caches := []wire.CacheDefinition{{ID: testNodeUID, Name: "cache", ClientSocket: "/run/racer/cache/client/socket", OriginSocket: "/run/racer/cache/origin/socket"}} + v := wire.Publication{SchemaVersion: wire.SchemaVersion, Cluster: testNodeUID} + + content, membership, err := wire.ContentHashes(v) + if err != nil { + t.Fatal(err) + } + + previous := VersionRecord{Cluster: v.Cluster, Sequence: 1, MembershipVersion: 1, ContentHash: content, MembershipHash: membership} + p := NewPublications() + + for _, tc := range []struct { + name string + members AcceptedMembers + caches []wire.CacheDefinition + sequence wire.Sequence + membership wire.MembershipVersion + }{ + {"unchanged empty", nil, nil, 1, 1}, + {"cache only", nil, caches, 2, 1}, + {"membership", members, caches, 3, 2}, + {"unchanged populated", members, caches, 3, 2}, + {"remove cache", members, nil, 4, 2}, + } { + t.Run(tc.name, func(t *testing.T) { + prepared, err := p.Prepare(previous, "rv", tc.members, tc.caches) + if err != nil { + t.Fatal(err) + } + + v.Sequence, v.MembershipVersion, v.Caches = tc.sequence, tc.membership, tc.caches + + v.Members = nil + for _, member := range tc.members { + v.Members = append(v.Members, member) + } + + want, err := wire.EncodePublication(v) + if err != nil { + t.Fatal(err) + } + + content, membership, err := wire.ContentHashes(v) + if err != nil { + t.Fatal(err) + } + + wantRecord := VersionRecord{Cluster: v.Cluster, Sequence: tc.sequence, MembershipVersion: tc.membership, ContentHash: content, MembershipHash: membership} + if prepared.record != wantRecord || prepared.encoded != string(want) || prepared.previous != previous || prepared.resourceVersion != "rv" || prepared.owner != p { + t.Fatal("prepared bytes, hashes, counters, or commit metadata differ") + } + + previous = prepared.record + }) + } +} + +func TestPrepareRejectsFinalCounterGrowth(t *testing.T) { + v := wire.Publication{ + SchemaVersion: wire.SchemaVersion, Cluster: testNodeUID, Sequence: 9, MembershipVersion: 9, + Members: []wire.Member{{Node: testNodeUID, Shares: 1, PeerEndpoint: "192.0.2.1:1", RDMANICs: []wire.RDMANIC{{Device: "x", Port: 1}}}}, + } + + b, err := wire.EncodePublication(v) + if err != nil { + t.Fatal(err) + } + + v.Members[0].RDMANICs[0].Device += strings.Repeat("x", wire.MaxPublicationBytes-len(b)) + + content, membership, err := wire.ContentHashes(v) + if err != nil { + t.Fatal(err) + } + + previous := VersionRecord{Cluster: v.Cluster, Sequence: 9, MembershipVersion: 9, ContentHash: content, MembershipHash: membership} + members := AcceptedMembers{testNodeUID: v.Members[0]} + p := NewPublications() + + prepared, err := p.Prepare(previous, "rv", members, nil) + if err != nil || len(prepared.encoded) != wire.MaxPublicationBytes { + t.Fatalf("exact final bound: %v", err) + } + + // Same-width input change grows both assigned counters from 9 to 10. The + // counter-free hash still fits, but the final publication must be rejected. + m := members[testNodeUID] + m.Shares++ + + members[testNodeUID] = m + if prepared, err := p.Prepare(previous, "rv", members, nil); !errors.Is(err, wire.TooLarge) || prepared != nil { + t.Fatalf("oversized final encoding accepted: %v", err) + } +} + +func stagedTopology(t *testing.T) *topologyFixture { + t.Helper() + return testTopology(t) +} + +func stagedFakeClient(base client.WithWatch) client.WithWatch { + // The fake client does not assign server UIDs. Model the API's identity and + // immutable ConfigMap data rules, including metadata updates remaining legal. + return interceptor.NewClient(base, interceptor.Funcs{ + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + if obj.GetUID() == "" { + obj.SetUID(types.UID(uuid.NewString())) + } + + return c.Create(ctx, obj, opts...) + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if cm, ok := obj.(*corev1.ConfigMap); ok { + old := &corev1.ConfigMap{} + if err := c.Get(ctx, client.ObjectKeyFromObject(cm), old); err != nil { + return err + } + + if old.Immutable != nil && *old.Immutable && (!reflect.DeepEqual(old.Data, cm.Data) || cm.Immutable == nil || !*cm.Immutable) { + return errors.New("immutable data changed") + } + } + + return c.Update(ctx, obj, opts...) + }, + }) +} + +// Each write boundary is tested both before persistence and with an uncertain +// successful response, including cancellation immediately after persistence. +var errInitializationInterrupted = errors.New("interrupted initialization") + +func interruptInitialization(base client.WithWatch, boundary string, cancel context.CancelFunc) client.WithWatch { + boom := errInitializationInterrupted + + return interceptor.NewClient(base, interceptor.Funcs{ + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + if boundary == "before create" { + return boom + } + + if err := c.Create(ctx, obj, opts...); err != nil { + return err + } + + if boundary == "after create" { + return boom + } + + if boundary == "cancel create" { + cancel() + } + + return nil + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if boundary == "before commit" { + return boom + } + + if err := c.Update(ctx, obj, opts...); err != nil { + return err + } + + if boundary == "after commit" { + return boom + } + + if boundary == "cancel commit" { + cancel() + } + + return nil + }, + }) +} + +func TestStagedInitializationEveryBoundary(t *testing.T) { + for _, credentials := range []bool{false, true} { + for _, boundary := range []string{"before create", "after create", "cancel create", "before commit", "after commit", "cancel commit"} { + t.Run(fmt.Sprintf("credentials=%t/%s", credentials, boundary), func(t *testing.T) { + f := newStagedFixture(t, credentials) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + writer := interruptInitialization(f.base, boundary, cancel) + require.Error(t, f.run(ctx, writer), "interruption not observed") + before := f.object() + err := f.base.Get(t.Context(), client.ObjectKeyFromObject(before), before) + require.True(t, err == nil || apierrors.IsNotFound(err), "candidate read: %v", err) + + if boundary != "after commit" && boundary != "cancel commit" { + require.Error(t, f.read(t.Context()), "uncommitted authority usable") + } + + require.NoError(t, f.run(t.Context(), f.base), "restart") + + after := before.DeepCopyObject().(client.Object) + require.NoError(t, f.base.Get(t.Context(), client.ObjectKeyFromObject(before), after)) + + if before.GetUID() != "" { + require.Equal(t, before.GetUID(), after.GetUID(), "recovery replaced staged material") + } + + if secret, ok := before.(*corev1.Secret); ok && secret.UID != "" { + require.Equal(t, secret.Data, after.(*corev1.Secret).Data, "recovery changed exact secret material") + } + }) + } + } +} + +type stagedFixture struct { + config Config + base client.WithWatch + credentials bool +} + +func newStagedFixture(t *testing.T, credentials bool) stagedFixture { + t.Helper() + r := stagedTopology(t) + + f := stagedFixture{config: r.Config, base: r.Client.(client.WithWatch), credentials: credentials} + if credentials { + require.NoError(t, ensureInstalled(t.Context(), f.base, f.base, f.config)) + } + + return f +} + +func (f stagedFixture) object() client.Object { + metadata := metav1.ObjectMeta{Namespace: f.config.Namespace, Name: f.config.VersionConfigMapName} + if f.credentials { + metadata.Name = f.config.CredentialsSecretName + return &corev1.Secret{ObjectMeta: metadata} + } + + return &corev1.ConfigMap{ObjectMeta: metadata} +} + +func (f stagedFixture) run(ctx context.Context, writer client.WithWatch) error { + if f.credentials { + _, err := Assemble(f.config, writer, f.base).authority.ReconcileCredentials(ctx) + return err + } + + return ensureInstalled(ctx, writer, f.base, f.config) +} + +func (f stagedFixture) read(ctx context.Context) error { + if f.credentials { + _, err := loadSigning(ctx, f.base, f.config, time.Now().UTC().Truncate(time.Second)) + return err + } + + _, _, err := readVersion(ctx, f.base, f.config) + + return err +} + +func TestStagedCommittedDeletionAndReplacementFailClosed(t *testing.T) { + for _, credentials := range []bool{false, true} { + for _, replace := range []bool{false, true} { + t.Run(fmt.Sprintf("credentials=%t/replace=%t", credentials, replace), func(t *testing.T) { + f := newStagedFixture(t, credentials) + require.NoError(t, ensureInstalled(t.Context(), f.base, f.base, f.config)) + runKeys(t, Assemble(f.config, f.base, f.base).Keyring) + obj := f.object() + require.NoError(t, f.base.Get(t.Context(), client.ObjectKeyFromObject(obj), obj)) + require.NoError(t, f.base.Delete(t.Context(), obj)) + + if replace { + obj.SetResourceVersion("") + obj.SetUID("") + + require.NoError(t, f.base.Create(t.Context(), obj)) + } + + require.Error(t, f.run(t.Context(), rejectWrites(t, f.base)), "lost authority accepted") + }) + } + } +} + +func TestStagedCompetingInstallers(t *testing.T) { + r := stagedTopology(t) + base := r.Client.(client.WithWatch) + + var wg sync.WaitGroup + for range 8 { + wg.Go(func() { + if err := ensureInstalled(t.Context(), base, base, r.Config); err != nil { + t.Error(err) + } + }) + } + + wg.Wait() + + for range 8 { + wg.Go(func() { + delay, err := Assemble(r.Config, base, base).authority.ReconcileCredentials(t.Context()) + if apierrors.IsConflict(err) || apierrors.IsAlreadyExists(err) { + if delay != 0 { + t.Errorf("losing initializer supplied retry delay %v", delay) + } + } else if err != nil { + t.Error(err) + } + }) + } + + wg.Wait() + runKeys(t, Assemble(r.Config, base, base).Keyring) + version, _, err := readVersion(t.Context(), base, r.Config) + require.NoError(t, err) + shared, bundle, _, _ := keyState(t, Assemble(r.Config, base, base).Keyring) + require.Equal(t, string(shared.UID), version.Annotations[credentialUID]) + require.EqualValues(t, 1, bundle.Generation, "competitors must converge on the same initial generation") +} + +func TestStagedDelayedCreateCannotResurrectAuthority(t *testing.T) { + for _, credentials := range []bool{false, true} { + t.Run(fmt.Sprint(credentials), func(t *testing.T) { + r := stagedTopology(t) + + base := r.Client.(client.WithWatch) + if credentials { + require.NoError(t, ensureInstalled(t.Context(), base, base, r.Config)) + } + + writer := interceptor.NewClient(base, interceptor.Funcs{Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + // Let a competitor fully commit, then lose its resource, while this + // installer is paused after its NotFound but before its Create. + if credentials { + runKeys(t, Assemble(r.Config, base, base).Keyring) + } else { + require.NoError(t, ensureInstalled(ctx, base, base, r.Config)) + } + + require.NoError(t, base.Delete(ctx, obj)) + + return c.Create(ctx, obj, opts...) + }}) + if credentials { + _, _ = Assemble(r.Config, writer, base).authority.ReconcileCredentials(t.Context()) + _, err := Assemble(r.Config, base, base).authority.ReconcileCredentials(t.Context()) + require.Error(t, err, "delayed Create restored credentials authority") + } else { + require.Error(t, ensureInstalled(t.Context(), writer, base, r.Config), "delayed Create restored version authority") + _, _, err := readVersion(t.Context(), base, r.Config) + require.Error(t, err, "orphan version accepted") + } + }) + } +} + +func TestStagedRejectsUnboundOrCorruptCandidates(t *testing.T) { + for _, credentials := range []bool{false, true} { + for _, corruption := range []string{"binding", "protocol", "immutable", "data"} { + t.Run(fmt.Sprintf("credentials=%t/%s", credentials, corruption), func(t *testing.T) { + f := newStagedFixture(t, credentials) + writer := interruptInitialization(f.base, "after create", func() {}) + _ = f.run(t.Context(), writer) + obj := f.object() + require.NoError(t, f.base.Get(t.Context(), client.ObjectKeyFromObject(obj), obj)) + + switch corruption { + case "binding": + obj.GetAnnotations()[installationUIDAnnotation] = "foreign" + case "protocol": + delete(obj.GetAnnotations(), initializationProtocol) + case "immutable": + immutable := true + if credentials { + obj.(*corev1.Secret).Immutable = &immutable + } else { + obj.(*corev1.ConfigMap).Immutable = &immutable + } + case "data": + if credentials { + delete(obj.(*corev1.Secret).Data, "issuer.json") + } else { + obj.(*corev1.ConfigMap).Data["sequence"] = "2" + } + } + + require.NoError(t, f.base.Update(t.Context(), obj)) + require.Error(t, f.run(t.Context(), f.base), "invalid candidate committed") + + if credentials { + version, _, err := readVersion(t.Context(), f.base, f.config) + require.NoError(t, err) + require.Empty(t, version.Annotations[credentialClaim], "invalid candidate consumed claim") + } else { + _, err := readInstallation(t.Context(), f.base, f.config, true) + require.NoError(t, err, "invalid candidate consumed marker") + } + }) + } + } +} + +func integrationStagedInitialization(t *testing.T, c client.WithWatch) { + t.Helper() + + for _, credentials := range []bool{false, true} { + for i, boundary := range []string{"before create", "after create", "cancel create", "before commit", "after commit", "cancel commit"} { + t.Run(fmt.Sprintf("credentials=%t/%s", credentials, boundary), func(t *testing.T) { + a := integrationInstallation(t, c, fmt.Sprintf("staged-%t-%d", credentials, i)) + cfg := a.Topology.Config + + marker, err := readInstallation(t.Context(), c, cfg, true) + require.NoError(t, err) + + marker.Data[markerInitializationProtocol] = stagedInitialization + require.NoError(t, c.Update(t.Context(), marker)) + + if credentials { + require.NoError(t, ensureInstalled(t.Context(), c, c, cfg)) + } + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + writer := interruptInitialization(c, boundary, cancel) + if credentials { + _, err := Assemble(cfg, writer, c).authority.ReconcileCredentials(ctx) + require.Error(t, err, "interruption not injected") + + runKeys(t, Assemble(cfg, c, c).Keyring) + } else { + require.Error(t, ensureInstalled(ctx, writer, c, cfg), "interruption not injected") + require.NoError(t, ensureInstalled(t.Context(), c, c, cfg)) + } + + marker, err = readInstallation(t.Context(), c, cfg, false) + require.NoError(t, err) + + marker.Data[versionUID] = "replacement" + require.Error(t, c.Update(t.Context(), marker), "API allowed rewriting immutable binding") + }) + } + } +} + +func testConfig(t *testing.T) Config { + t.Helper() + t.Setenv("RACER_CLUSTER_ID", testOtherUID) + t.Setenv("POD_NAMESPACE", "racer") + + cfg, err := LoadConfig() + if err != nil { + t.Fatal(err) + } + + return cfg +} + +func testTopology(t *testing.T, objects ...client.Object) *topologyFixture { + t.Helper() + cfg := testConfig(t) + + scheme := runtime.NewScheme() + for _, add := range []func(*runtime.Scheme) error{corev1.AddToScheme, appsv1.AddToScheme, racerv1.AddToScheme} { + if err := add(scheme); err != nil { + t.Fatal(err) + } + } + + marker := &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName, UID: "installation-uid"}, Data: map[string]string{"cluster": string(cfg.Cluster), "version_configmap": cfg.VersionConfigMapName, "state": "fresh", markerInitializationProtocol: stagedInitialization}} + objects = append(objects, marker) + c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).WithIndex(&corev1.Pod{}, podNodeIndex, podNodeKeys).Build() + + api := stagedFakeClient(c) + + return Assemble(cfg, api, api).Topology +} + +func initializedTopology(t *testing.T, objects ...client.Object) *topologyFixture { + t.Helper() + + r := testTopology(t, objects...) + if err := ensureInstalled(context.Background(), r.Client, r.APIReader, r.Config); err != nil { + t.Fatal(err) + } + + return r +} + +func reconcileTopology(t *testing.T, r *topologyFixture, ctx context.Context) *CommittedPublication { + t.Helper() + + _, err := r.operations().PublishTopology(ctx, r.observeTopology) + if err != nil { + t.Fatalf("publish: %v", err) + } + + p, err := r.Publications.Current() + if err != nil { + t.Fatal(err) + } + + return p +} + +func TestInitializeCrashOrdering(t *testing.T) { + for _, stage := range []string{"before marker", "marker response lost", "after marker", "create response lost", "success"} { + t.Run(stage, func(t *testing.T) { + exerciseInitializationCrash(t, stage) + }) + } +} + +func exerciseInitializationCrash(t *testing.T, stage string) { + t.Helper() + r := testTopology(t) + base := r.Client.(client.WithWatch) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + writes := []string{} + boom := errors.New("simulated crash") + r.Client = interceptor.NewClient(base, interceptor.Funcs{ + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + writes = append(writes, "consume") + + if stage == "before marker" { + return boom + } + + if err := c.Update(ctx, obj, opts...); err != nil { + return err + } + + if stage == "marker response lost" { + return boom + } + + if stage == "after marker" { + cancel() + } + + return nil + }, + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + writes = append(writes, "create") + + marker, err := readInstallation(ctx, r.APIReader, r.Config, true) + require.NoError(t, err, "candidate must precede marker freeze") + require.Empty(t, marker.Data[versionUID], "uncreated candidate was committed") + + if err := c.Create(ctx, obj, opts...); err != nil { + return err + } + + if stage == "create response lost" { + return boom + } + + return nil + }, + }) + + err := ensureInstalled(ctx, r.Client, r.APIReader, r.Config) + require.Equal(t, stage == "success", err == nil, "initialize: %v", err) + + wantWrites := []string{"create"} + if stage != "create response lost" { + wantWrites = append(wantWrites, "consume") + } + + require.Equal(t, wantWrites, writes) + + writes = nil + + var candidate corev1.ConfigMap + require.NoError(t, base.Get(t.Context(), client.ObjectKey{Namespace: r.Config.Namespace, Name: r.Config.VersionConfigMapName}, &candidate)) + r.Client = interceptor.NewClient(base, interceptor.Funcs{Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + t.Fatal("restart recreated staged candidate") + return nil + }}) + require.NoError(t, ensureInstalled(t.Context(), r.Client, r.APIReader, r.Config)) + + _, _, err = readVersion(context.Background(), r.APIReader, r.Config) + + require.NoError(t, err) + marker, err := readInstallation(t.Context(), base, r.Config, false) + require.NoError(t, err) + require.Equal(t, string(candidate.UID), marker.Data[versionUID], "recovery did not bind staged UID") +} + +func TestInitializeConflictAndExistingState(t *testing.T) { + r := testTopology(t) + base := r.Client.(client.WithWatch) + creates := 0 + + r.Client = interceptor.NewClient(base, interceptor.Funcs{ + Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + return apierrors.NewConflict(corev1.Resource("configmaps"), "marker", errors.New("concurrent initializer")) + }, + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + creates++ + return c.Create(ctx, obj, opts...) + }, + }) + + ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) + defer cancel() + + if err := ensureInstalled(ctx, r.Client, r.APIReader, r.Config); !errors.Is(err, context.DeadlineExceeded) || creates != 1 { + t.Fatalf("marker conflict: %v, creates=%d", err, creates) + } + + r.Client = base + if err := base.Delete(t.Context(), &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: r.Config.VersionConfigMapName}}); err != nil { + t.Fatal(err) + } + + if err := base.Create(context.Background(), &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: r.Config.VersionConfigMapName}}); err != nil { + t.Fatal(err) + } + + if err := ensureInstalled(context.Background(), r.Client, r.APIReader, r.Config); !errors.Is(err, wire.Unavailable) { + t.Fatalf("existing counters accepted: %v", err) + } + + if _, err := readInstallation(context.Background(), r.APIReader, r.Config, true); err != nil { + t.Fatalf("marker consumed despite existing counters: %v", err) + } +} + +func TestInitializeCanceledBeforeMarker(t *testing.T) { + r := testTopology(t) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if err := ensureInstalled(ctx, r.Client, r.APIReader, r.Config); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled initialize: %v", err) + } + + if _, err := readInstallation(context.Background(), r.APIReader, r.Config, true); err != nil { + t.Fatalf("canceled initialize consumed marker: %v", err) + } +} + +func TestStalePreparedPublicationCannotCommit(t *testing.T) { + r := initializedTopology(t) + ctx := context.Background() + + cm, previous, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + p, err := r.Publications.Prepare(previous, cm.ResourceVersion, nil, nil) + if err != nil { + t.Fatal(err) + } + // Another writer changes only metadata, but even equal counters require the + // exact read resource version. No install token can escape a stale candidate. + cm.Labels = map[string]string{"changed": "true"} + if err := r.Update(ctx, cm); err != nil { + t.Fatal(err) + } + + if committed, err := r.CommitVersion(ctx, p); !apierrors.IsConflict(err) || committed != nil { + t.Fatalf("stale candidate committed: %p, %v", committed, err) + } +} + +func TestTopologyNamespaceOwnershipAndMissingDaemonSet(t *testing.T) { + node := memberNode() + pod := memberPod("pod", 1, "192.0.2.1") + pod.Namespace = "unrelated" + r := initializedTopology(t, &node, &pod) + ctx := context.Background() + + first := reconcileTopology(t, r, ctx) + if len(r.authority.accepted) != 0 { + t.Fatal("pod in unrelated namespace admitted") + } + + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: r.Config.DaemonSetName, UID: testDaemonSetUID}} + if err := r.Create(ctx, ds); err != nil { + t.Fatal(err) + } + + if next := reconcileTopology(t, r, ctx); next != first { + t.Fatal("foreign pod admitted by matching owner UID") + } + + pod.Namespace = r.Config.Namespace + + pod.ResourceVersion = "" + if err := r.Create(ctx, &pod); err != nil { + t.Fatal(err) + } + + member := reconcileTopology(t, r, ctx) + if len(r.authority.accepted) != 1 { + t.Fatal("managed endpoint not admitted") + } + + if err := r.Delete(ctx, ds); err != nil { + t.Fatal(err) + } + + if gap := reconcileTopology(t, r, ctx); gap != member { + t.Fatal("missing workload discarded warm endpoint") + } + + // Supply the persisted hint as input; root tests cover writing and removing it. + encoded, err := json.Marshal(r.authority.accepted[testNodeUID]) + require.NoError(t, err) + require.NoError(t, r.Get(ctx, client.ObjectKeyFromObject(&node), &node)) + node.Annotations = map[string]string{admittedMemberAnnotation: string(encoded)} + require.NoError(t, r.Update(ctx, &node)) + + r = Assemble(r.Config, r.Client, r.APIReader).Topology + if restarted := reconcileTopology(t, r, ctx); restarted.record != member.record || restarted.encoded != member.encoded { + t.Fatal("cold restart changed admitted membership during workload gap") + } + + if len(r.authority.accepted) != 1 { + t.Fatal("cold restart lost persisted admitted identity") + } +} + +func TestRecoveryNeverRecreatesCounters(t *testing.T) { + for _, mutation := range []string{"missing version", "missing marker", "corrupt", "wrong cluster", "wrong marker uid", "mutable marker", "fresh marker"} { + t.Run(mutation, func(t *testing.T) { + r := initializedTopology(t) + ctx := context.Background() + + cm, _, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + marker, err := readInstallation(ctx, r.APIReader, r.Config, false) + if err != nil { + t.Fatal(err) + } + + switch mutation { + case "missing version": + err = r.Delete(ctx, cm) + case "missing marker": + err = r.Delete(ctx, marker) + case "corrupt": + cm.Data["sequence"] = "01" + err = r.Update(ctx, cm) + case "wrong cluster": + cm.Data["cluster"] = testNodeUID + err = r.Update(ctx, cm) + case "wrong marker uid": + cm.Annotations[installationUIDAnnotation] = "replacement" + err = r.Update(ctx, cm) + case "mutable marker": + marker.Immutable = nil + case "fresh marker": + marker.Data["state"] = "fresh" + } + + if mutation == "mutable marker" || mutation == "fresh marker" { + // Inject invalid observed state without weakening the fake API's + // immutable data enforcement for normal operations. + r.APIReader = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if key.Name == marker.Name { + marker.DeepCopyInto(obj.(*corev1.ConfigMap)) + return nil + } + + return c.Get(ctx, key, obj, opts...) + }}) + } + + if err != nil { + t.Fatal(err) + } + + writes := 0 + + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + writes++ + return nil + }, + Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + writes++ + return nil + }, + }) + if _, err := r.operations().PublishTopology(ctx, r.observeTopology); err == nil || writes != 0 { + t.Fatalf("unsafe recovery: %v, writes=%d", err, writes) + } + + if _, err := r.Publications.Current(); err == nil { + t.Fatal("served invalid recovery") + } + }) + } +} + +func TestVersionCountersAndCrashAfterCommit(t *testing.T) { + r := initializedTopology(t) + ctx := context.Background() + + empty := reconcileTopology(t, r, ctx) + if v := empty.record; v.Sequence != 1 || v.MembershipVersion != 1 { + t.Fatalf("initial counters: %+v", v) + } + + if same := reconcileTopology(t, r, ctx); same != empty { + t.Fatal("unchanged install replaced shared allocation") + } + + cache := catalogCache("cache-a", testNodeUID) + if err := r.Create(ctx, &cache); err != nil { + t.Fatal(err) + } + + runKeys(t, Assemble(r.Config, r.Client, r.APIReader).Keyring) + + catalog := reconcileTopology(t, r, ctx) + if v := catalog.record; v.Sequence != 2 || v.MembershipVersion != 1 { + t.Fatalf("catalog counters: %+v", v) + } + + node, pod := memberNode(), memberPod("pod", 1, "192.0.2.1") + + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Name: r.Config.DaemonSetName, Namespace: r.Config.Namespace, UID: testDaemonSetUID}} + for _, obj := range []client.Object{&node, &pod, ds} { + if err := r.Create(ctx, obj); err != nil { + t.Fatal(err) + } + } + + member := reconcileTopology(t, r, ctx) + if v := member.record; v.Sequence != 3 || v.MembershipVersion != 2 { + t.Fatalf("member counters: %+v", v) + } + + assertCrashAfterVersionCommit(t, r, member) +} + +func assertCrashAfterVersionCommit(t *testing.T, r *topologyFixture, member *CommittedPublication) { + t.Helper() + ctx := t.Context() + + cm, previous, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + prepared, err := r.Publications.Prepare(previous, cm.ResourceVersion, nil, nil) + if err != nil { + t.Fatal(err) + } + + if _, err := r.CommitVersion(ctx, prepared); err != nil { + t.Fatal(err) + } + // Simulated crash before install: no candidate bytes or history were exposed. + if current, _ := r.Publications.Current(); current != member { + t.Fatal("commit installed prematurely") + } + + if len(r.authority.accepted) != 1 { + t.Fatal("commit replaced history") + } + + r = Assemble(r.Config, r.Client, r.APIReader).Topology + + recovered := reconcileTopology(t, r, ctx) + if v := recovered.record; v.Sequence != 5 || v.MembershipVersion != 4 { + t.Fatalf("recovery reused an unserved counter: %+v", v) + } +} + +func TestCASConflictRetriesFreshInputsAndKeepsHistory(t *testing.T) { + node, pod := memberNode(), memberPod("pod", 1, "192.0.2.1") + r := initializedTopology(t, &node, &pod, &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Name: "racer-dataplane", Namespace: "racer", UID: testDaemonSetUID}}) + ctx := context.Background() + initial := reconcileTopology(t, r, ctx) + base := r.Client.(client.WithWatch) + updates := 0 + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + updates++ + if updates == 1 { + n := &corev1.Node{} + if err := c.Get(ctx, client.ObjectKey{Name: node.Name}, n); err != nil { + return err + } + + n.Annotations = map[string]string{wire.SharesAnnotation: "9"} + if err := c.Update(ctx, n); err != nil { + return err + } + + return apierrors.NewConflict(corev1.Resource("configmaps"), obj.GetName(), wire.Conflict) + } + + return c.Update(ctx, obj, opts...) + }}) + + _, err := r.operations().PublishTopology(ctx, r.observeTopology) + if !apierrors.IsConflict(err) || updates != 1 { + t.Fatalf("conflict: %v, writes=%d", err, updates) + } + + if current, _ := r.Publications.Current(); current != initial || r.authority.accepted[testNodeUID].Shares != 4 { + t.Fatal("failed commit changed publication/history") + } + + next := reconcileTopology(t, r, ctx) + if r.authority.accepted[testNodeUID].Shares != 9 || next.record.Sequence != initial.record.Sequence+1 { + t.Fatal("retry reused stale inputs") + } +} + +func TestCancellationBeforeWritesAndInstall(t *testing.T) { + for _, stage := range []string{"before reconcile", "read", "conflict", "after commit"} { + t.Run(stage, func(t *testing.T) { + r := initializedTopology(t) + base := r.Client.(client.WithWatch) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + writes := 0 + r.Client = interceptor.NewClient(base, interceptor.Funcs{Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + writes++ + + if stage == "conflict" { + cancel() + return apierrors.NewConflict(corev1.Resource("configmaps"), obj.GetName(), wire.Conflict) + } + + err := c.Update(ctx, obj, opts...) + + cancel() + + return err + }}) + + if stage == "before reconcile" { + cancel() + } + + if stage == "read" { + r.APIReader = interceptor.NewClient(base, interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + + cancel() + + return err + }}) + } + + _, err := r.operations().PublishTopology(ctx, r.observeTopology) + if stage == "conflict" { + require.True(t, apierrors.IsConflict(err), "operation must preserve the write error") + } else { + require.ErrorIs(t, err, context.Canceled) + } + + if (stage == "before reconcile" || stage == "read") && writes != 0 { + t.Fatal("write after cancellation") + } + + if _, err := r.Publications.Current(); err == nil { + t.Fatal("installed after cancellation") + } + }) + } +} + +func TestCancellationBetweenCommitAndInstall(t *testing.T) { + // Also cover cancellation between returning a committed token and Install. + r := initializedTopology(t) + ctx, cancel := context.WithCancel(context.Background()) + + cm, previous, err := readVersion(ctx, r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + p, err := r.Publications.Prepare(previous, cm.ResourceVersion, nil, nil) + if err != nil { + t.Fatal(err) + } + + committed, err := r.CommitVersion(ctx, p) + if err != nil { + t.Fatal(err) + } + + cancel() + + if err := r.Publications.Install(committed); !errors.Is(err, context.Canceled) { + t.Fatalf("late install: %v", err) + } +} + +func TestReplicaAcceptanceAndPublicHandles(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + p, err := a.Current() + require.NoError(t, err) + image, err := wire.DecodePublication(strings.NewReader(p.image.encoded)) + require.NoError(t, err) + + process, stop := context.WithCancel(t.Context()) + defer stop() + + replica := New(a.config, Dependencies{Reader: a.reader, Writer: a.client}) + replica.BindProcess(process) + require.NoError(t, replica.AcceptReplica(t.Context(), process, image)) + current, changed, err := replica.CurrentAndSubscribe() + require.NoError(t, err) + require.NotNil(t, changed) + require.Equal(t, image.Sequence, current.Sequence()) + require.Zero(t, (*PublicationHandle)(nil).Sequence()) + require.Zero(t, (&PublicationHandle{}).Sequence()) + + identity := pollIdentity(a.config, testNodeUID) + identity.owner = replica + waited, err := replica.Wait(t.Context(), identity, nil) + require.NoError(t, err) + require.Equal(t, current.Sequence(), waited.Sequence()) + require.NoError(t, replica.PublicationReady()) + stop() + + _, err = replica.Wait(t.Context(), identity, nil) + require.ErrorIs(t, err, context.Canceled) +} + +func TestReplicaRejectsInvalidAndUncommittedImages(t *testing.T) { + for _, scenario := range []string{"invalid", "uncommitted", "canceled", "missing version", "rollback"} { + t.Run(scenario, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + p, err := a.Current() + require.NoError(t, err) + image, err := wire.DecodePublication(strings.NewReader(p.image.encoded)) + require.NoError(t, err) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + switch scenario { + case "invalid": + image.SchemaVersion = 0 + case "uncommitted": + image.Sequence++ + case "canceled": + cancel() + case "missing version": + require.NoError(t, f.a.Topology.Delete(t.Context(), &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: a.config.Namespace, Name: a.config.VersionConfigMapName}})) + case "rollback": + newer := p.image.record + newer.Sequence++ + require.NoError(t, a.publications.confirm(newer)) + } + + require.Error(t, a.AcceptReplica(ctx, t.Context(), image)) + }) + } +} + +func TestPublicKeyringWaitAndUnavailableTrust(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + keyring, err := a.WaitKeyring(t.Context(), nil) + require.NoError(t, err) + require.EqualValues(t, 1, keyring.Generation()) + require.Zero(t, (KeyringHandle{}).Generation()) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + _, err = a.WaitKeyring(ctx, nil) + require.ErrorIs(t, err, context.Canceled) + a.trust.invalidate() + _, err = a.TrustPool() + require.ErrorIs(t, err, wire.Unavailable) + _, _, err = a.AdmitTrust(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) +} diff --git a/internal/racer/authority/security_test.go b/internal/racer/authority/security_test.go new file mode 100644 index 000000000..e3706e40b --- /dev/null +++ b/internal/racer/authority/security_test.go @@ -0,0 +1,199 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package authority + +import ( + "context" + "crypto/x509" + "errors" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + corev1 "k8s.io/api/core/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestSecurityPublicationReplayBeforeObservation(t *testing.T) { + for _, scenario := range []string{"rollback", "replacement"} { + t.Run(scenario, func(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + image, err := a.Current() + require.NoError(t, err) + guard, stop, err := image.Admit(t.Context()) + require.NoError(t, err) + + defer stop() + + cm, record, err := readVersion(t.Context(), a.reader, a.config) + require.NoError(t, err) + + if scenario == "rollback" { + record.Sequence-- + record.MembershipVersion-- + } else { + record.ContentHash = strings.Repeat("a", 64) + } + + cm.Data = versionData(record) + require.NoError(t, f.a.Topology.Update(t.Context(), cm)) + + called := false + _, err = a.PublishTopology(t.Context(), func(context.Context) (TopologyObservation, error) { + called = true + return TopologyObservation{}, errors.New("discovery unavailable") + }) + require.ErrorIs(t, err, wire.Conflict) + require.False(t, called, "invalid authority reached discovery") + require.ErrorIs(t, guard.Check(t.Context()), context.Canceled) + }) + } +} + +func TestSecurityPublicationReplayBeforeCAS(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + cm, previous, err := readVersion(t.Context(), a.reader, a.config) + require.NoError(t, err) + prepared, err := a.publications.Prepare(previous, cm.ResourceVersion, nil, nil) + require.NoError(t, err) + + newer := previous + newer.Sequence += 2 + require.NoError(t, a.publications.confirm(newer)) + _, err = a.publisher.CommitVersion(t.Context(), prepared) + after, record, readErr := readVersion(t.Context(), a.reader, a.config) + require.NoError(t, readErr) + require.Equal(t, cm.ResourceVersion, after.ResourceVersion, "stale counter reached external CAS") + require.Equal(t, previous, record) + require.ErrorIs(t, err, wire.Conflict) + require.Equal(t, newer, a.publications.observed) +} + +func TestSecurityPublicationReplacementDuringObservation(t *testing.T) { + f := newServingFixture(t) + a := f.a.authority + image, err := a.Current() + require.NoError(t, err) + guard, stop, err := image.Admit(t.Context()) + require.NoError(t, err) + + defer stop() + + _, err = a.PublishTopology(t.Context(), func(ctx context.Context) (TopologyObservation, error) { + cm, record, err := readVersion(ctx, a.reader, a.config) + require.NoError(t, err) + + record.ContentHash = strings.Repeat("b", 64) + cm.Data = versionData(record) + require.NoError(t, f.a.Topology.Update(ctx, cm)) + + return f.a.Topology.observeTopology(ctx) + }) + require.ErrorIs(t, err, wire.Conflict) + require.ErrorIs(t, guard.Check(t.Context()), context.Canceled, "CAS mismatch concealed invalid authority") +} + +func TestSecurityCredentialReplayBeforeUse(t *testing.T) { + for _, scenario := range []string{"rollback", "replacement"} { + for _, operation := range []string{"issue", "catalog", "rotation"} { + t.Run(scenario+"/"+operation, func(t *testing.T) { + f := newServingFixture(t) + a, r := f.a.authority, f.a.Keyring + cache := catalogCache("cache", testNodeUID) + require.NoError(t, r.Create(t.Context(), &cache)) + runKeys(t, r) + _, oldBundle, oldRotation, oldMaterial := keyState(t, r) + _, bundle, rotation, material := keyState(t, r) + replacement := editSigningCertificate(t, material.Keys[rotation.ActiveIssuer], func(cert *x509.Certificate) { + cert.NotBefore = cert.NotBefore.Add(-time.Second) + }) + rebindSigning(&bundle, &rotation, material, rotation.ActiveIssuer, replacement, true) + bundle.Generation++ + writeSigningCredentials(t, r, bundle, rotation, material) + runKeys(t, r) + + confirmed, highWater, digest := a.trust.confirmed, a.trust.highWater, a.trust.digest + guard, stop, err := a.AdmitTrust(t.Context()) + require.NoError(t, err) + + defer stop() + + if scenario == "replacement" { + oldBundle.Generation = bundle.Generation + } + + writeSigningCredentials(t, r, oldBundle, oldRotation, oldMaterial) + before, _, _, _ := keyState(t, r) + version, _, err := readVersion(t.Context(), a.reader, a.config) + require.NoError(t, err) + + switch operation { + case "issue": + identity := pollIdentity(a.config, testNodeUID) + identity.owner, identity.bearer = a, true + encoded, issueErr := a.Issue(t.Context(), identity, f.request) + err = issueErr + + require.Zero(t, len(encoded), "replayed issuer signed a fresh credential") + case "catalog": + _, err = a.PublishTopology(t.Context(), f.a.Topology.observeTopology) + case "rotation": + a.credentials.Now = func() time.Time { return oldRotation.NextRotation } + _, err = a.ReconcileCredentials(t.Context()) + } + + after := &corev1.Secret{} + require.NoError(t, a.reader.Get(t.Context(), client.ObjectKeyFromObject(before), after)) + require.Equal(t, before.ResourceVersion, after.ResourceVersion, "replayed credentials reached rotation CAS") + afterVersion, _, readErr := readVersion(t.Context(), a.reader, a.config) + require.NoError(t, readErr) + require.Equal(t, version.ResourceVersion, afterVersion.ResourceVersion, "replayed catalog reached publication CAS") + require.ErrorIs(t, err, wire.Conflict) + require.ErrorIs(t, guard.Check(t.Context()), context.Canceled) + require.Equal(t, highWater, a.trust.highWater) + require.Equal(t, digest, a.trust.digest) + require.Equal(t, confirmed, a.trust.confirmed) + }) + } + } +} + +func TestSecurityCredentialValidationDoesNotInstallOrRefresh(t *testing.T) { + for _, newer := range []bool{false, true} { + f := newServingFixture(t) + a, r := f.a.authority, f.a.Keyring + + _, bundle, rotation, material := keyState(t, r) + if newer { + bundle.Generation++ + writeSigningCredentials(t, r, bundle, rotation, material) + } + + confirmed, highWater, digest := a.trust.confirmed, a.trust.highWater, a.trust.digest + accepted, epoch := a.trust.bundle, a.trust.authority + identity := pollIdentity(a.config, testNodeUID) + identity.owner, identity.bearer = a, true + encoded, err := a.Issue(t.Context(), identity, f.request) + require.NoError(t, err) + require.NotZero(t, len(encoded)) + _, err = a.PublishTopology(t.Context(), f.a.Topology.observeTopology) + require.NoError(t, err) + require.Equal(t, confirmed, a.trust.confirmed) + require.Equal(t, highWater, a.trust.highWater) + require.Equal(t, digest, a.trust.digest) + require.Same(t, accepted, a.trust.bundle) + require.Equal(t, epoch, a.trust.authority) + a.trust.invalidate() + require.NoError(t, a.trust.validateReplay(bundle)) + require.ErrorIs(t, a.TrustReady(), wire.Unavailable, "validation restored withdrawn trust") + require.Equal(t, confirmed, a.trust.confirmed) + require.Equal(t, highWater, a.trust.highWater) + require.Equal(t, digest, a.trust.digest) + } +} diff --git a/internal/racer/ci_test.go b/internal/racer/ci_test.go new file mode 100644 index 000000000..1055d9beb --- /dev/null +++ b/internal/racer/ci_test.go @@ -0,0 +1,68 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package racer + +import ( + "os" + "strings" + "testing" + + "github.com/stretchr/testify/require" + "sigs.k8s.io/yaml" +) + +func TestRacerEnvtestCIContract(t *testing.T) { + makefile, err := os.ReadFile("../../Makefile") + require.NoError(t, err) + + target := func(name string) string { + t.Helper() + + _, body, found := strings.Cut(string(makefile), "\n"+name+":") + require.True(t, found, "missing target %s", name) + + body, _, _ = strings.Cut(body, "\n\n") + + return body + } + + run := target("racer-envtest") + for _, pkg := range []string{"./internal/racer", "./internal/racer/authority"} { + require.Contains(t, strings.Fields(run), pkg) + } + + // A prefix includes new envtest cases without an allowlist. + require.Contains(t, run, "-run '^TestEnvtest'") + require.Contains(t, run, "$(GOTEST) -race") + require.Contains(t, run, "-count=1") + require.Contains(t, run, "-timeout=5m") + require.Contains(t, run, "timeout --signal=TERM --kill-after=10s 300s") + require.Contains(t, run, `test -n "$(KUBEBUILDER_ASSETS)" ||`) + require.Contains(t, run, `KUBEBUILDER_ASSETS="$(KUBEBUILDER_ASSETS)"`) + + provision := target("racer-envtest-ci") + require.Contains(t, provision, "$(SETUP_ENVTEST)") + require.Contains(t, provision, "use $(ENVTEST_K8S_VERSION)") + require.Contains(t, provision, `$(MAKE) racer-envtest KUBEBUILDER_ASSETS="$$assets"`) + + workflow, err := os.ReadFile("../../.github/workflows/ci.yaml") + require.NoError(t, err) + + var ci struct { + Jobs map[string]struct { + Steps []struct { + Run string `json:"run"` + } `json:"steps"` + } `json:"jobs"` + } + + require.NoError(t, yaml.Unmarshal(workflow, &ci)) + + var commands []string + for _, step := range ci.Jobs["racer-envtest"].Steps { + commands = append(commands, step.Run) + } + + require.Contains(t, commands, "timeout --signal=TERM --kill-after=10s 300s make racer-envtest-ci") +} diff --git a/internal/racer/controller_test.go b/internal/racer/controller_test.go new file mode 100644 index 000000000..e4a20ca66 --- /dev/null +++ b/internal/racer/controller_test.go @@ -0,0 +1,2023 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package racer + +import ( + "bytes" + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/http/httptrace" + "os" + "path/filepath" + "reflect" + "slices" + "strconv" + "strings" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + coordv1 "k8s.io/api/coordination/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/labels" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "k8s.io/client-go/kubernetes" + clientgoscheme "k8s.io/client-go/kubernetes/scheme" + "k8s.io/client-go/rest" + "k8s.io/client-go/tools/leaderelection/resourcelock" + "k8s.io/client-go/util/workqueue" + "k8s.io/utils/ptr" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + "sigs.k8s.io/controller-runtime/pkg/envtest" + "sigs.k8s.io/controller-runtime/pkg/event" + "sigs.k8s.io/controller-runtime/pkg/reconcile" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/testutil" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestMixedNetworkConfiguration(t *testing.T) { + values := map[string]string{ + "RACER_CLUSTER_ID": "11111111-1111-1111-1111-111111111111", + "RACER_CONTROL_URL": "https://controller:8443", "RACER_DATAPLANE_IMAGE": "racer:5e1", + "RACER_HOST_NETWORK": "true", + } + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + + for _, input := range []string{`[]`, `["node-b","node-a"]`} { + values["RACER_POD_NETWORK_NODES"] = input + _, err := testutil.ConfigFromLookup(lookup) + require.NoError(t, err) + } + + for _, input := range []string{"", "null", `{}`, `"node-a"`, `[1]`, `[null]`, `[""]`, `["Node-A"]`, `["node-a","node-a"]`, `["node-a"] trailing`} { + values["RACER_POD_NETWORK_NODES"] = input + _, err := testutil.ConfigFromLookup(lookup) + require.ErrorIs(t, err, wire.InvalidRequest, input) + } + + values["RACER_POD_NETWORK_NODES"] = `["node-a"]` + values["RACER_HOST_NETWORK"] = "false" + _, err := testutil.ConfigFromLookup(lookup) + require.ErrorIs(t, err, wire.InvalidRequest) +} + +func TestMixedNetworkBuilders(t *testing.T) { + cfg := workloadConfig(t) + legacy, err := testutil.DesiredDaemonSet(cfg) + require.NoError(t, err) + sets, err := testutil.DesiredDaemonSets(cfg) + require.NoError(t, err) + require.Equal(t, []*appsv1.DaemonSet{legacy}, sets) + + cfg.HostNetwork = true + cfg.PeerPort, cfg.DiagnosticsPort = 18082, 19090 + cfg.PodNetworkNodes = []string{"node-b", "node-a"} + _, err = testutil.DesiredDaemonSet(cfg) + require.ErrorIs(t, err, wire.InvalidRequest, "legacy planner must fail closed") + sets, err = testutil.DesiredDaemonSets(cfg) + require.NoError(t, err) + require.Len(t, sets, 2) + host, pod := sets[0], sets[1] + require.Equal(t, DataplaneDaemonSetName, host.Name) + require.Equal(t, PodNetworkDaemonSetName, pod.Name) + require.Equal(t, legacy.Spec.Selector, host.Spec.Selector) + require.True(t, host.Spec.Template.Spec.HostNetwork) + require.False(t, pod.Spec.Template.Spec.HostNetwork) + require.Equal(t, corev1.DNSClusterFirst, pod.Spec.Template.Spec.DNSPolicy) + require.Equal(t, host.Spec.Template.Spec.Volumes, pod.Spec.Template.Spec.Volumes) + require.Equal(t, host.Spec.Template.Spec.ServiceAccountName, pod.Spec.Template.Spec.ServiceAccountName) + require.Equal(t, host.Spec.Template.Spec.Containers, pod.Spec.Template.Spec.Containers) + + for i, ds := range sets { + selector, err := metav1.LabelSelectorAsSelector(ds.Spec.Selector) + require.NoError(t, err) + require.True(t, selector.Matches(labels.Set(ds.Spec.Template.Labels))) + require.False(t, selector.Matches(labels.Set(sets[1-i].Spec.Template.Labels))) + } + + hostTerms := host.Spec.Template.Spec.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms + podTerms := pod.Spec.Template.Spec.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms + + require.Len(t, hostTerms, 1) + require.Len(t, hostTerms[0].MatchFields, 2) + require.Len(t, podTerms, 2) + + for i, node := range []string{"node-a", "node-b"} { + require.Equal(t, corev1.NodeSelectorRequirement{Key: "metadata.name", Operator: corev1.NodeSelectorOpNotIn, Values: []string{node}}, hostTerms[0].MatchFields[i]) + require.Equal(t, []corev1.NodeSelectorRequirement{{Key: "metadata.name", Operator: corev1.NodeSelectorOpIn, Values: []string{node}}}, podTerms[i].MatchFields) + require.Equal(t, hostTerms[0].MatchExpressions, podTerms[i].MatchExpressions) + } + + require.Equal(t, []string{"node-b", "node-a"}, cfg.PodNetworkNodes) +} + +func TestMixedNetworkIdentities(t *testing.T) { + scheme := runtime.NewScheme() + require.NoError(t, appsv1.AddToScheme(scheme)) + + host := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: DataplaneDaemonSetName, UID: "host-current"}} + podnet := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: PodNetworkDaemonSetName, UID: "pod-current"}} + reader := fake.NewClientBuilder().WithScheme(scheme).WithObjects(host, podnet).Build() + ids, err := readManagedWorkloadIdentities(t.Context(), reader, Config{Namespace: "racer", DaemonSetName: DataplaneDaemonSetName}) + require.NoError(t, err) + + for _, ds := range []*appsv1.DaemonSet{host, podnet} { + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", OwnerReferences: []metav1.OwnerReference{*metav1.NewControllerRef(ds, appsv1.SchemeGroupVersion.WithKind("DaemonSet"))}}} + require.Equal(t, ds.Name == host.Name, ids.Owns(pod)) + + for _, mutate := range []func(*corev1.Pod){ + func(p *corev1.Pod) { p.Namespace = "other" }, + func(p *corev1.Pod) { p.OwnerReferences[0].UID = "stale" }, + func(p *corev1.Pod) { p.OwnerReferences[0].UID = "" }, + func(p *corev1.Pod) { p.OwnerReferences[0].Name = "arbitrary" }, + func(p *corev1.Pod) { p.OwnerReferences[0].Controller = ptr.To(false) }, + func(p *corev1.Pod) { p.OwnerReferences[0].APIVersion = "apps/v2" }, + func(p *corev1.Pod) { p.OwnerReferences = nil; p.Labels = ds.Labels }, + } { + bad := pod.DeepCopy() + mutate(bad) + require.False(t, ids.Owns(bad)) + } + + require.False(t, (DataplaneWorkloadIdentities{}).Owns(pod)) + } + + require.False(t, ids.Owns(nil)) + require.NoError(t, reader.Delete(t.Context(), podnet)) + ids, err = readManagedWorkloadIdentities(t.Context(), reader, Config{Namespace: "racer", DaemonSetName: DataplaneDaemonSetName}) + require.NoError(t, err) + require.Equal(t, host.Name, ids.Name) + require.Equal(t, host.UID, ids.UID) +} + +func TestAssemble(t *testing.T) { + // Nil Kubernetes dependencies make unintended constructor API calls fail. + a := Assemble(Config{}, nil, nil) + require.NotNil(t, a.authority, "bootstrap must have an issuer") + require.Same(t, a.authority, a.Topology.authority) + require.Same(t, a.authority, a.Keyring.authority, "controllers, issuance, and serving must share trust and catalog gate") + require.Same(t, a.Lifecycle, a.Server.Lifecycle) + require.Same(t, a.Replication, a.Server.Leader) + require.NotNil(t, a.Server) + + for _, cfg := range []Config{a.Topology.config, a.Keyring.config, a.Replication.config} { + require.Equal(t, wire.CertificateLifetime, cfg.CertificateLifetime) + require.Equal(t, 30*time.Second, cfg.SnapshotMaxAge) + } + + require.Equal(t, a.Topology.config.SnapshotMaxAge, a.Replication.config.SnapshotMaxAge) + require.Equal(t, a.Keyring.config.SnapshotMaxAge, a.Replication.config.SnapshotMaxAge) + require.Error(t, a.authority.PublicationReady()) + require.False(t, a.Server.NeedLeaderElection()) + require.False(t, a.Replication.NeedLeaderElection()) + require.ErrorIs(t, a.Server.Ready(nil), wire.Unavailable) + + // Composition freezes observer inputs before exposing any server entry point. + want := a.Replication.config + + cfg := Config{SnapshotMaxAge: time.Hour} + require.NotEqual(t, want, cfg) + require.Equal(t, want, a.Replication.config, "replication settings must be captured before serving composition") +} + +func TestFailClosedEntryPoints(t *testing.T) { + a := Assemble(Config{}, nil, nil) + ctx := context.Background() + + operations := map[string]func() error{ + "topology": func() error { _, err := a.Topology.Reconcile(ctx, ctrl.Request{}); return err }, + "keyring": func() error { _, err := a.Keyring.Reconcile(ctx, ctrl.Request{}); return err }, + "workload": func() error { _, err := testutil.DesiredDaemonSet(testutil.Config{}); return err }, + "server": func() error { return a.Server.Start(ctx) }, + "run": func() error { return Run(ctx, Config{}) }, + } + for name, operation := range operations { + t.Run(name, func(t *testing.T) { + if err := operation(); err == nil { + t.Fatalf("entry point did not fail closed: %v", err) + } + }) + } +} + +func TestScaffoldRoutesCannotAuthenticate(t *testing.T) { + handler := Assemble(Config{}, nil, nil).Server.Handler() + + for _, tc := range []struct{ method, path string }{ + {http.MethodPost, wire.BootstrapPath}, + {http.MethodGet, wire.SnapshotPath}, + } { + t.Run(tc.path, func(t *testing.T) { + response := httptest.NewRecorder() + handler.ServeHTTP(response, httptest.NewRequest(tc.method, tc.path, nil)) + + if response.Code != http.StatusServiceUnavailable || response.Body.String() != `{"code":"unavailable"}` { + t.Fatalf("unexpected scaffold response: %d %s", response.Code, response.Body.String()) + } + }) + } +} + +func TestSingletonCoalescesObjects(t *testing.T) { + requests := singleton(context.Background(), nil) + if len(requests) != 1 || requests[0].Name != "racer" || requests[0].Namespace != "" { + t.Fatalf("unexpected singleton requests: %v", requests) + } +} + +func TestInitialEnqueueEmptyInputsAndCoalescing(t *testing.T) { + q := workqueue.NewTypedRateLimitingQueue(workqueue.DefaultTypedControllerRateLimiter[reconcile.Request]()) + defer q.ShutDown() + + for range 10 { + if err := initialEnqueue().Start(context.Background(), q); err != nil { + t.Fatal(err) + } + } + + if q.Len() != 1 { + t.Fatalf("startup events not coalesced: %d", q.Len()) + } + + request, stopped := q.Get() + if stopped || request != singleton(context.Background(), nil)[0] { + t.Fatalf("initial request: %+v", request) + } + + q.Done(request) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if err := initialEnqueue().Start(ctx, q); !errors.Is(err, context.Canceled) || q.Len() != 0 { + t.Fatalf("canceled startup queued: %v", err) + } +} + +func TestTopologyWatchFiltering(t *testing.T) { + cfg := testConfig(t) + node := memberNode() + updated := node.DeepCopy() + + updated.Status.Conditions = []corev1.NodeCondition{{Type: corev1.NodeReady, Status: corev1.ConditionTrue}} + if nodeChanges().Update(event.UpdateEvent{ObjectOld: &node, ObjectNew: updated}) { + t.Fatal("node readiness enqueued topology") + } + + updated.Annotations = map[string]string{wire.SharesAnnotation: ""} + if !nodeChanges().Update(event.UpdateEvent{ObjectOld: &node, ObjectNew: updated}) { + t.Fatal("absent -> invalid empty annotation lost") + } + + updated.Annotations = nil + + updated.Labels = map[string]string{wire.ExclusionLabel: "false"} + if !nodeChanges().Update(event.UpdateEvent{ObjectOld: &node, ObjectNew: updated}) { + t.Fatal("exclusion presence lost") + } + + pod := memberPod("pod", 1, "192.0.2.1") + pod.OwnerReferences[0].Name = cfg.DaemonSetName + + pred := managedPodChanges(cfg) + if !pred.Create(event.CreateEvent{Object: &pod}) { + t.Fatal("managed pod ignored") + } + + changed := pod.DeepCopy() + + changed.Status.Conditions = []corev1.PodCondition{{Type: corev1.PodReady, Status: corev1.ConditionTrue}} + if pred.Update(event.UpdateEvent{ObjectOld: &pod, ObjectNew: changed}) { + t.Fatal("pod readiness enqueued topology") + } + + changed.OwnerReferences = nil + if pred.Create(event.CreateEvent{Object: changed}) || !pred.Update(event.UpdateEvent{ObjectOld: &pod, ObjectNew: changed}) { + t.Fatal("ownership loss filtering") + } + + changed = pod.DeepCopy() + + changed.Namespace = "unrelated" + if pred.Create(event.CreateEvent{Object: changed}) { + t.Fatal("foreign namespace pod admitted") + } + + if keys := podNodeKeys(&pod); len(keys) != 1 || keys[0] != pod.Spec.NodeName { + t.Fatalf("node index: %v", keys) + } + + pod.Spec.NodeName = "" + if len(podNodeKeys(&pod)) != 0 { + t.Fatal("unassigned pod indexed") + } + + cm := &corev1.ConfigMap{} + cm.Name, cm.Namespace, cm.ResourceVersion = cfg.VersionConfigMapName, cfg.Namespace, "1" + newCM := cm.DeepCopy() + + newCM.ResourceVersion = "2" + if versionChanges(cfg).Update(event.UpdateEvent{ObjectOld: cm, ObjectNew: newCM}) { + t.Fatal("CAS-only write caused reconcile loop") + } + + newCM.Data = map[string]string{"sequence": "2"} + if !versionChanges(cfg).Update(event.UpdateEvent{ObjectOld: cm, ObjectNew: newCM}) { + t.Fatal("version change ignored") + } +} + +// Opt-in, but never silently skip when assets were explicitly supplied. envtest +// runs real etcd/apiserver processes; it has no kubelet or workload controllers. +func TestEnvtestServer(t *testing.T) { + assets := os.Getenv("KUBEBUILDER_ASSETS") + if assets == "" { + t.Skip("set KUBEBUILDER_ASSETS to run the real API-server integration suite") + } + + scheme := runtime.NewScheme() + for _, add := range []func(*runtime.Scheme) error{clientgoscheme.AddToScheme, racerv1.AddToScheme} { + if err := add(scheme); err != nil { + t.Fatal(err) + } + } + + environment := &envtest.Environment{BinaryAssetsDirectory: assets, CRDDirectoryPaths: []string{"../../api/racer/v1alpha1/crd"}, ErrorIfCRDPathMissing: true} + + rc, err := environment.Start() + if err != nil { + t.Fatal(err) + } + + t.Cleanup(func() { + if err := environment.Stop(); err != nil { + t.Error(err) + } + }) + + c, err := client.NewWithWatch(rc, client.Options{Scheme: scheme}) + if err != nil { + t.Fatal(err) + } + + t.Run("initialization-and-CAS", func(t *testing.T) { integrationInitialization(t, c) }) + t.Run("cache-name-admission", func(t *testing.T) { integrationCacheNameAdmission(t, c) }) + t.Run("rotation-crash-recovery", func(t *testing.T) { integrationRotation(t, c) }) + t.Run("manager-election-HTTPS-failover", func(t *testing.T) { integrationManagers(t, rc, scheme, c) }) +} + +func integrationInstallation(t *testing.T, c client.Client, namespace string) *Application { + t.Helper() + cfg := testConfig(t) + cfg.Namespace = namespace + + cfg.MetricsAddress, cfg.ProbeAddress = "0", "0" + for _, obj := range []client.Object{ + &corev1.Namespace{ObjectMeta: metav1.ObjectMeta{Name: namespace}}, + &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: namespace, Name: cfg.InstallationConfigMapName}, Data: map[string]string{"cluster": string(cfg.Cluster), "version_configmap": cfg.VersionConfigMapName, "state": "fresh", "initialization_protocol": "staged-v1"}}, + } { + if err := c.Create(t.Context(), obj); err != nil { + t.Fatal(err) + } + } + + return assembleFixture(cfg, c, c) +} + +type interruptedClient struct { + client.Client + update func(context.Context, client.Object, ...client.UpdateOption) error + create func(context.Context, client.Object, ...client.CreateOption) error +} + +func (c interruptedClient) Update(ctx context.Context, obj client.Object, opts ...client.UpdateOption) error { + if c.update != nil { + return c.update(ctx, obj, opts...) + } + + return c.Client.Update(ctx, obj, opts...) +} + +func (c interruptedClient) Create(ctx context.Context, obj client.Object, opts ...client.CreateOption) error { + if c.create != nil { + return c.create(ctx, obj, opts...) + } + + return c.Client.Create(ctx, obj, opts...) +} + +func integrationInitialization(t *testing.T, c client.Client) { + ctx := t.Context() + concurrent := integrationInstallation(t, c, "init-concurrent") + arrived := make(chan struct{}, 2) + proceed := make(chan struct{}) + results := make(chan error, 2) + + for range 2 { + r := Assemble(concurrent.Topology.config, interruptedClient{Client: c, update: func(ctx context.Context, obj client.Object, opts ...client.UpdateOption) error { + arrived <- struct{}{} + + select { + case <-ctx.Done(): + return ctx.Err() + case <-proceed: + } + + return c.Update(ctx, obj, opts...) + }}, c).Topology + + go func() { results <- r.authority.Recover(ctx, r.Client) }() + } + + for range 2 { + select { + case <-arrived: + case <-time.After(5 * time.Second): + close(proceed) + t.Fatal("concurrent initializers did not reach marker CAS") + } + } + + close(proceed) + + winners := 0 + + for range 2 { + if err := <-results; err == nil { + winners++ + } else if !apierrors.IsConflict(err) { + t.Fatalf("marker CAS loser: %v", err) + } + } + + if winners != 2 { + t.Fatalf("successful concurrent startups: %d", winners) + } + + integrationInitializationCAS(t, c) + integrationInitializationCrash(t, c) +} + +func integrationInitializationCAS(t *testing.T, c client.Client) { + t.Helper() + ctx := t.Context() + a := integrationInstallation(t, c, "init-cas") + + r := a.Topology + if err := a.Recover(ctx, r.Client); err != nil { + t.Fatal(err) + } + + marker, err := readInstallation(ctx, r.APIReader, r.config, false) + if err != nil { + t.Fatal(err) + } + + for _, mutation := range []func(*corev1.ConfigMap){ + func(cm *corev1.ConfigMap) { cm.Data["state"] = "fresh" }, + func(cm *corev1.ConfigMap) { cm.Immutable = ptr.To(false) }, + } { + copy := marker.DeepCopy() + mutation(copy) + + if err := c.Update(ctx, copy); !apierrors.IsInvalid(err) { + t.Fatalf("API server allowed immutable marker rollback: %v", err) + } + } + + if err := a.Recover(ctx, r.Client); err != nil { + t.Fatalf("installed startup rejected: %v", err) + } + + cm, _, err := readVersion(ctx, r.APIReader, r.config) + if err != nil { + t.Fatal(err) + } + + // Race after CommitVersion's authoritative read, so the API server, rather + // than our preliminary resourceVersion comparison, must reject the write. + writer := interruptedClient{Client: c, update: func(ctx context.Context, obj client.Object, opts ...client.UpdateOption) error { + other := cm.DeepCopy() + + other.Labels = map[string]string{"concurrent": "writer"} + if err := c.Update(ctx, other); err != nil { + return err + } + + return c.Update(ctx, obj, opts...) + }} + + owner := authority.New(r.config.authorityConfig(), authority.Dependencies{Reader: c, Writer: writer}) + if _, err := owner.PublishTopology(ctx, r.observeTopology); !apierrors.IsConflict(err) { + t.Fatalf("real CAS failed: err=%v", err) + } + + r.Client = c + + cm, _, err = readVersion(ctx, r.APIReader, r.config) + if err != nil { + t.Fatal(err) + } + + leader, cancel := context.WithCancel(ctx) + + owner = authority.New(r.config.authorityConfig(), authority.Dependencies{Reader: c, Writer: interruptedClient{Client: c, update: func(ctx context.Context, obj client.Object, opts ...client.UpdateOption) error { + err := c.Update(ctx, obj, opts...) + + cancel() + + return err + }}}) + if _, err := owner.PublishTopology(leader, r.observeTopology); !errors.Is(err, context.Canceled) { + t.Fatalf("late install: %v", err) + } +} + +func integrationInitializationCrash(t *testing.T, c client.Client) { + t.Helper() + + ctx := t.Context() + for _, afterCreate := range []bool{false, true} { + a := integrationInstallation(t, c, fmt.Sprintf("init-crash-%t", afterCreate)) + boom := errors.New("ambiguous initialization response") + + a.Topology.Client = interruptedClient{Client: c, create: func(ctx context.Context, obj client.Object, opts ...client.CreateOption) error { + if afterCreate { + if err := c.Create(ctx, obj, opts...); err != nil { + return err + } + } + + return boom + }} + if err := a.Recover(ctx, a.Topology.Client); !errors.Is(err, boom) { + t.Fatal(err) + } + + restarted := Assemble(a.Topology.config, c, c) + if err := restarted.Recover(ctx, c); err != nil { + t.Fatalf("ambiguous initialization recovery: %v", err) + } + + if _, _, err := readVersion(ctx, restarted.Topology.APIReader, restarted.Topology.config); err != nil { + t.Fatalf("crash recovery afterCreate=%t: %v", afterCreate, err) + } + } +} + +func integrationRotation(t *testing.T, c client.Client) { + a := integrationInstallation(t, c, "rotation") + cfg := a.Keyring.config + cfg.Rotation.Interval = 7 * 24 * time.Hour + + a = assembleFixture(cfg, c, c) + require.NoError(t, a.Recover(t.Context(), a.Topology.Client)) + + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "rotation-cache"}} + if err := c.Create(t.Context(), cache); err != nil { + t.Fatal(err) + } + + t.Cleanup(func() { + if err := c.Delete(context.Background(), cache); err != nil { + t.Error(err) + } + }) + + r := a.Keyring + now := time.Now().UTC().Truncate(time.Second) + fixtureDependencies[a.authority].now = func() time.Time { return now } + + runKeys(t, r) + _, initial, state, _ := keyState(t, r) + now = state.NextRotation + oldIssuer := state.ActiveIssuer + // Interrupt each actual write boundary, including a committed response lost + // during activation. Every recovery uses a fresh application and real reads. + for _, step := range []struct { + name string + secret string + after bool + }{{"stage-private", r.config.CredentialsSecretName, true}, {"activate-bundle", r.config.CredentialsSecretName, true}, {"prune-private", r.config.CredentialsSecretName, false}} { + boom := errors.New(step.name) + failed := false + + fixtureDependencies[r.authority].Client = interruptedClient{Client: c, update: func(ctx context.Context, obj client.Object, opts ...client.UpdateOption) error { + if obj.GetName() != step.secret { + return c.Update(ctx, obj, opts...) + } + + failed = true + + if step.after { + if err := c.Update(ctx, obj, opts...); err != nil { + return err + } + } + + return boom + }} + _, err := r.Reconcile(t.Context(), ctrl.Request{}) + require.True(t, failed) + require.ErrorIs(t, err, boom) + require.Error(t, r.authority.TrustReady()) + + _, before, beforeState, private := keyState(t, r) + recoveredApp := assembleFixture(r.config, c, c) + recovered := recoveredApp.Keyring + fixtureDependencies[recovered.authority].now = func() time.Time { return now } + runKeys(t, recovered) + _, after, next, material := keyState(t, recovered) + + switch step.name { + case "stage-private": + require.Equal(t, initial.Generation+1, before.Generation) + require.Len(t, private.Keys, 2) + require.Equal(t, beforeState.PreparedIssuer, next.PreparedIssuer) + require.Len(t, after.CacheKeys, 4) + + now = next.ActivateAt + case "activate-bundle": + require.Equal(t, before.Generation, after.Generation) + require.Equal(t, beforeState.ActiveIssuer, next.ActiveIssuer) + require.Len(t, next.Retiring, 1) + require.Len(t, after.CacheKeys, 2) + + now = next.Retiring[oldIssuer] + case "prune-private": + require.True(t, containsRoot(before, oldIssuer)) + require.Len(t, private.Keys, 2) + require.Len(t, material.Keys, 1) + require.Equal(t, before.Generation+1, after.Generation) + require.Len(t, after.CacheKeys, 2) + require.False(t, containsRoot(after, oldIssuer)) + } + + r = recovered + } +} + +type transportFunc func(*http.Request) (*http.Response, error) + +func (f transportFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +func eventually(t *testing.T, description string, f func() bool) { + t.Helper() + + deadline := time.Now().Add(20 * time.Second) + for time.Now().Before(deadline) { + if f() { + return + } + + time.Sleep(20 * time.Millisecond) + } + + t.Fatal("timed out: " + description) +} + +func integrationManagers(t *testing.T, rc *rest.Config, scheme *runtime.Scheme, c client.Client) { + a := integrationInstallation(t, c, "managers") + if err := a.Recover(t.Context(), a.Topology.Client); err != nil { + t.Fatal(err) + } + + cfg := a.Topology.config + cfg.ControllerServiceAccount = "racer-controller" + cfg.ReplicationServerName = "racer-controller.managers.svc" + roots := integrationTLS(t, &cfg) + cfg.ReplicationTrustFile = cfg.TLSCertificateFile + + _, port, err := net.SplitHostPort(unusedAddress(t)) + if err != nil { + t.Fatal(err) + } + + number, err := strconv.ParseUint(port, 10, 16) + if err != nil { + t.Fatal(err) + } + + cfg.ReplicationPort = uint16(number) + + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.ControllerServiceAccount}} + if err := c.Create(t.Context(), sa); err != nil { + t.Fatal(err) + } + + kube, err := kubernetes.NewForConfig(rc) + if err != nil { + t.Fatal(err) + } + + var ( + apps [2]*Application + cancels [2]context.CancelFunc + done [2]chan error + denyRenewal [2]atomic.Bool + topologyCommitted [2]atomic.Bool + ) + + for i := range apps { + apps[i], cancels[i], done[i] = startIntegrationManager(t, rc, scheme, c, cfg, kube, sa, i, port, &denyRenewal[i], &topologyCommitted[i]) + } + + leader := -1 + + eventually(t, "elected manager becomes ready", func() bool { + for i, app := range apps { + if app.Replication.isLeader() && app.Server.Ready(nil) == nil { + leader = i + return true + } + } + + return false + }) + + follower := 1 - leader + + eventually(t, "follower installs replicated snapshot and serves", func() bool { return apps[follower].Server.Ready(nil) == nil }) + integrationManagerFailover(t, rc, c, cfg, roots, apps, cancels, done, &denyRenewal[leader], &topologyCommitted[follower], leader) +} + +func startIntegrationManager(t *testing.T, rc *rest.Config, scheme *runtime.Scheme, c client.Client, cfg Config, kube kubernetes.Interface, sa *corev1.ServiceAccount, i int, port string, denyRenewal, topologyCommitted *atomic.Bool) (*Application, context.CancelFunc, chan error) { + t.Helper() + + ip := fmt.Sprintf("127.0.0.%d", i+2) + cfg.ControlAddress, cfg.ProbeAddress = net.JoinHostPort(ip, port), unusedAddress(t) + + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: fmt.Sprintf("controller-%d", i)}, Spec: corev1.PodSpec{ServiceAccountName: sa.Name, Containers: []corev1.Container{{Name: "controller", Image: "example.invalid/controller:test"}}}} + if err := c.Create(t.Context(), pod); err != nil { + t.Fatal(err) + } + + pod.Status.PodIP = ip + if err := c.Status().Update(t.Context(), pod); err != nil { + t.Fatal(err) + } + + cfg.PodName, cfg.PodUID = pod.Name, string(pod.UID) + + token := boundPodToken(t, kube, pod, sa.Name, ReplicationAudience) + + cfg.ReplicationTokenFile = filepath.Join(t.TempDir(), "token") + if err := os.WriteFile(cfg.ReplicationTokenFile, []byte(token), 0o600); err != nil { + t.Fatal(err) + } + + options := managerOptions(cfg, scheme) + options.LeaseDuration, options.RenewDeadline, options.RetryPeriod = ptr.To(4*time.Second), ptr.To(2*time.Second), ptr.To(500*time.Millisecond) + options.Controller.SkipNameValidation = ptr.To(true) // Two real managers in one test process. + connection := rest.CopyConfig(rc) + connection.WrapTransport = func(base http.RoundTripper) http.RoundTripper { + return transportFunc(func(req *http.Request) (*http.Response, error) { + if denyRenewal.Load() && req.Method == http.MethodPut && strings.Contains(req.URL.Path, "/leases/") { + return nil, errors.New("injected Lease renewal partition") + } + + return base.RoundTrip(req) + }) + } + + lockClient, err := kubernetes.NewForConfig(connection) + if err != nil { + t.Fatal(err) + } + + options.LeaderElectionResourceLockInterface = &resourcelock.LeaseLock{LeaseMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: "racer-controller"}, Client: lockClient.CoordinationV1(), LockConfig: resourcelock.ResourceLockConfig{Identity: cfg.PodName + "/" + cfg.PodUID}} + + mgr, err := ctrl.NewManager(connection, options) + if err != nil { + t.Fatal(err) + } + + writer := interruptedClient{Client: mgr.GetClient(), update: func(ctx context.Context, obj client.Object, opts ...client.UpdateOption) error { + if err := mgr.GetClient().Update(ctx, obj, opts...); err != nil { + return err + } + + if _, ok := obj.(*corev1.ConfigMap); ok && obj.GetName() == cfg.VersionConfigMapName { + topologyCommitted.Store(true) + } + + return nil + }} + + app := Assemble(cfg, writer, mgr.GetAPIReader()) + if err := app.SetupWithManager(mgr); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(t.Context()) + done := make(chan error, 1) + + go func() { done <- mgr.Start(ctx) }() + + t.Cleanup(func() { + cancel() + + select { + case <-done: + case <-time.After(15 * time.Second): + t.Error("manager failed to stop") + } + }) + + return app, cancel, done +} + +func integrationManagerFailover(t *testing.T, rc *rest.Config, c client.Client, cfg Config, roots *x509.CertPool, apps [2]*Application, cancels [2]context.CancelFunc, done [2]chan error, denyRenewal, topologyCommitted *atomic.Bool, leader int) { + t.Helper() + + follower := 1 - leader + ds := integrationWorkloadDrift(t, c, cfg, apps) + peer := integrationEnrollment(t, rc, c, apps[follower], ds, roots) + integrationHTTPSFailover(t, rc, c, cfg, apps, cancels, done, denyRenewal, topologyCommitted, leader, peer) +} + +func integrationWorkloadDrift(t *testing.T, c client.Client, cfg Config, apps [2]*Application) *appsv1.DaemonSet { + t.Helper() + + for i, app := range apps { + response, err := http.Get("http://" + app.Topology.config.ProbeAddress + "/readyz") + + want := 200 + + if err != nil { + t.Fatal(err) + } + + response.Body.Close() + + if response.StatusCode != want { + t.Fatalf("manager %d readiness: %d", i, response.StatusCode) + } + } + // No Nodes/Pods/DaemonSets existed at startup. The ready Racer manager must + // not provision workloads; simulate the operator's independent installation. + ds := &appsv1.DaemonSet{} + if err := c.Get(t.Context(), client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.DaemonSetName}, ds); !apierrors.IsNotFound(err) { + t.Fatalf("Racer manager created a workload: %v", err) + } + + workload, err := testutil.DesiredDaemonSet(testutil.Config{ + Cluster: cfg.Cluster, Namespace: cfg.Namespace, + ControlURL: "https://127.0.0.1:8443", DataplaneImage: "example.invalid/racer:test", + BootstrapTrustConfigMap: "racer-bootstrap-trust", PeerPort: cfg.PeerPort, + DataplaneServiceAccount: cfg.DataplaneServiceAccount, DaemonSetName: cfg.DaemonSetName, + }) + if err != nil { + t.Fatal(err) + } + + if err := c.Create(t.Context(), workload); err != nil { + t.Fatal(err) + } + + if err := c.Get(t.Context(), client.ObjectKeyFromObject(workload), ds); err != nil { + t.Fatal(err) + } + + if !ptr.Deref(ds.Spec.Template.Spec.Containers[0].SecurityContext.ReadOnlyRootFilesystem, false) { + t.Fatal("operator-owned workload must initially have a read-only root filesystem") + } + + ds.Spec.Template.Spec.Containers[0].SecurityContext.ReadOnlyRootFilesystem = ptr.To(false) + if err := c.Update(t.Context(), ds); err != nil { + t.Fatal(err) + } + + rv := ds.ResourceVersion + + time.Sleep(300 * time.Millisecond) + + if err := c.Get(t.Context(), client.ObjectKeyFromObject(ds), ds); err != nil || ds.ResourceVersion != rv { + t.Fatalf("Racer manager mutated an operator-owned workload: %v", err) + } + + if ptr.Deref(ds.Spec.Template.Spec.Containers[0].SecurityContext.ReadOnlyRootFilesystem, true) { + t.Fatal("Racer manager reverted operator-owned security drift") + } + + return ds +} + +func integrationHTTPSFailover(t *testing.T, rc *rest.Config, c client.Client, cfg Config, apps [2]*Application, cancels [2]context.CancelFunc, done [2]chan error, denyRenewal, topologyCommitted *atomic.Bool, leader int, peer *http.Client) { + t.Helper() + + follower := 1 - leader + endpoint := "https://" + apps[leader].Topology.config.ControlAddress + response, err := peer.Get(endpoint + wire.SnapshotPath) + + publication, err := wire.DecodePublication(bytes.NewReader(responseBody(t, response, err, 200))) + if err != nil { + t.Fatal(err) + } + // Establish a real authenticated pending HTTPS request before loss of Lease. + eventually(t, "follower receives current image before failover", func() bool { + p, err := apps[follower].authority.Current() + return err == nil && p.Sequence() == publication.Sequence + }) + + followerResponse, followerErr := peer.Get("https://" + apps[follower].Topology.config.ControlAddress + wire.SnapshotPath) + responseBody(t, followerResponse, followerErr, http.StatusOK) + + pollDone := integrationPendingPoll(peer, endpoint, publication.Sequence) + + // A duplicate authenticated poll observes admission through the public route. + eventually(t, "leader parks authenticated poll", func() bool { + ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("%s%s?after=%d", endpoint, wire.SnapshotPath, publication.Sequence), nil) + if err != nil { + t.Fatal(err) + } + + duplicate, err := peer.Do(req) + if err != nil { + return false + } + defer duplicate.Body.Close() + + return duplicate.StatusCode == http.StatusTooManyRequests + }) + + lease := &coordv1.Lease{} + if err := c.Get(t.Context(), client.ObjectKey{Namespace: cfg.Namespace, Name: "racer-controller"}, lease); err != nil { + t.Fatal(err) + } + + oldHolder := *lease.Spec.HolderIdentity + start := time.Now() + + denyRenewal.Store(true) + eventually(t, "Lease renewal failure withdraws readiness", func() bool { return apps[leader].Server.Ready(nil) != nil }) + + select { + case err := <-pollDone: + require.Error(t, err, "old leader poll completed instead of closing") + case <-time.After(5 * time.Second): + t.Fatal("old leader HTTPS poll survived cancellation") + } + + select { + case err := <-done[leader]: + require.ErrorContains(t, err, "leader election lost") + + done[leader] <- err // Cleanup still joins this manager. + case <-time.After(5 * time.Second): + t.Fatal("lost leader did not stop") + } + + if conn, err := net.DialTimeout("tcp", apps[leader].Topology.config.ControlAddress, time.Second); err == nil { + conn.Close() + t.Fatal("old leader listener still accepts after manager exit") + } + + eventually(t, "follower takes expired Lease, commits topology, and serves", func() bool { + return apps[follower].Replication.isLeader() && topologyCommitted.Load() && apps[follower].Server.Ready(nil) == nil + }) + + if err := c.Get(t.Context(), client.ObjectKeyFromObject(lease), lease); err != nil || *lease.Spec.HolderIdentity == oldHolder { + t.Fatalf("Lease did not change holder: %v", err) + } + + _, committedVersion, err := readVersion(t.Context(), c, apps[follower].Topology.config) + require.NoError(t, err) + require.Equal(t, publication.Sequence, committedVersion.Sequence) + require.Equal(t, publication.MembershipVersion, committedVersion.MembershipVersion) + + response, err = peer.Get("https://" + apps[follower].Topology.config.ControlAddress + wire.SnapshotPath) + + recovered, err := wire.DecodePublication(bytes.NewReader(responseBody(t, response, err, 200))) + require.NoError(t, err) + require.Equal(t, publication.Sequence, recovered.Sequence) + require.Equal(t, publication.MembershipVersion, recovered.MembershipVersion) + + t.Logf("actual Lease failover and authenticated HTTPS recovery: %s; sequence=%d membership=%d", time.Since(start), recovered.Sequence, recovered.MembershipVersion) + integrationAuthorizationLoad(t, rc, c, apps[follower], peer) + + integrationManagerStop(t, apps[follower], cancels[follower], done[follower]) +} + +func integrationPendingPoll(peer *http.Client, endpoint string, sequence wire.Sequence) <-chan error { + done := make(chan error, 1) + + go func() { + path := fmt.Sprintf("%s%s?after=%d", endpoint, wire.SnapshotPath, sequence) + response, err := peer.Get(path) + // A duplicate probe can win admission first. Wait for its short request. + for attempts := 0; response != nil && response.StatusCode == http.StatusTooManyRequests && attempts < 20; attempts++ { + response.Body.Close() + time.Sleep(20 * time.Millisecond) + + response, err = peer.Get(path) + } + + if response != nil { + response.Body.Close() + + if response.StatusCode == http.StatusServiceUnavailable { + err = wire.Unavailable + } + } + + done <- err + }() + + return done +} + +func integrationManagerStop(t *testing.T, app *Application, cancel context.CancelFunc, done chan error) { + t.Helper() + + if err := app.Server.Ready(nil); err != nil { + t.Fatalf("new leader not ready before normal cancellation: %v", err) + } + + cancel() + eventually(t, "manager cancellation withdraws readiness", func() bool { return app.Server.Ready(nil) != nil }) + + select { + case err := <-done: + if err != nil { + t.Fatalf("normal manager cancellation: %v", err) + } + + done <- err + case <-time.After(5 * time.Second): + t.Fatal("canceled manager did not stop") + } + + if _, err := app.authority.Current(); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled manager still publishes: %v", err) + } +} + +type countedBody struct { + io.ReadCloser + bytes *atomic.Int64 +} + +func (b countedBody) Read(p []byte) (int, error) { + n, err := b.ReadCloser.Read(p) + b.bytes.Add(int64(n)) + + return n, err +} + +func integrationAuthorizationLoad(t *testing.T, rc *rest.Config, c client.Client, a *Application, peer *http.Client) { + t.Helper() + // A separate owner uses the instrumented API dependency for every operation. + // Validate replicated state through public operations before measuring serving. + var requests, nodeLists, podLists, received atomic.Int64 + + connection := rest.CopyConfig(rc) + connection.WrapTransport = func(base http.RoundTripper) http.RoundTripper { + return transportFunc(func(req *http.Request) (*http.Response, error) { + requests.Add(1) + + if req.URL.Path == "/api/v1/nodes" { + nodeLists.Add(1) + } + + if strings.HasSuffix(req.URL.Path, "/pods") { + podLists.Add(1) + } + + response, err := base.RoundTrip(req) + if err == nil { + response.Body = countedBody{ReadCloser: response.Body, bytes: &received} + } + + return response, err + }) + } + + reader, err := client.New(connection, client.Options{Scheme: c.Scheme()}) + if err != nil { + t.Fatal(err) + } + + measuredApp := Assemble(a.Topology.config, reader, reader) + + measured := measuredApp.Server + if err := measuredApp.authority.Observe(t.Context()); err != nil { + t.Fatal(err) + } + + image, err := wire.DecodePublication(strings.NewReader(capturePublication(t, a.authority).encoded)) + if err != nil { + t.Fatal(err) + } + + if err := measuredApp.authority.AcceptReplica(t.Context(), t.Context(), image); err != nil { + t.Fatal(err) + } + + startFixtureLifecycle(t, measured.Lifecycle, t.Context()) + // Positive control on this same owner proves the measurement is connected. + requests.Store(0) + received.Store(0) + + if err := measuredApp.authority.Observe(t.Context()); err != nil { + t.Fatal(err) + } + + require.NotZero(t, requests.Load()) + require.NotZero(t, received.Load()) + + config, err := measured.TLSConfig(t.Context()) + if err != nil { + t.Fatal(err) + } + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + + server := &http.Server{Handler: measured.Handler(), TLSConfig: config} + + go func() { server.Serve(tls.NewListener(listener, config)) }() + + t.Cleanup(func() { server.Close() }) + + endpoint := "https://" + listener.Addr().String() + wire.SnapshotPath + // The first request includes a real TLS handshake. Neither path may read API state. + if err := measuredApp.authority.Observe(t.Context()); err != nil { + t.Fatal(err) + } + + requests.Store(0) + received.Store(0) + + response, err := peer.Get(endpoint) + responseBody(t, response, err, 200) + + if requests.Load() != 0 || received.Load() != 0 { + t.Fatalf("TLS handshake/snapshot used API: requests=%d bytes=%d", requests.Load(), received.Load()) + } + + integrationAuthorizationScale(t, rc, c, a, measuredApp, peer, endpoint, &requests, &nodeLists, &podLists, &received) +} + +func integrationAuthorizationScale(t *testing.T, rc *rest.Config, c client.Client, a, measuredApp *Application, peer *http.Client, endpoint string, requests, nodeLists, podLists, received *atomic.Int64) { + t.Helper() + + for _, count := range []int{1, 1001} { + if count > 1 { + for i := range 1000 { + node := &corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: fmt.Sprintf("auth-scale-%04d", i), Labels: map[string]string{wire.ExclusionLabel: ""}}} + if err := c.Create(t.Context(), node); err != nil { + t.Fatal(err) + } + } + } + // Include all added Nodes in the real informer. Controller watch traffic + // is deliberately outside the request budget. + eventually(t, "authorization discovery cache convergence", func() bool { + var nodes corev1.NodeList + return a.Topology.List(t.Context(), &nodes) == nil && len(nodes.Items) == count + }) + + // This standalone owner has no observer loop. Setup may outlast snapshot + // freshness, so refresh outside the measured request window. + if err := measuredApp.authority.Observe(t.Context()); err != nil { + t.Fatal(err) + } + + requests.Store(0) + nodeLists.Store(0) + podLists.Store(0) + received.Store(0) + + start := time.Now() + + for range 10 { + response, err := peer.Get(endpoint) + responseBody(t, response, err, 200) + } + + if requests.Load() != 0 || nodeLists.Load() != 0 || podLists.Load() != 0 { + t.Fatalf("authorization API budget drift: requests=%d node_lists=%d pod_lists=%d", requests.Load(), nodeLists.Load(), podLists.Load()) + } + + if received.Load() != 0 { + t.Fatalf("snapshot read API bytes: %d", received.Load()) + } + + t.Logf("real HTTPS authorization: live_nodes=%d snapshots=10 elapsed=%s API_requests=%d Node_lists=%d Pod_lists=%d API_response_bytes=%d (warm TLS; envtest QPS=%g burst=%d)", count, time.Since(start), requests.Load(), nodeLists.Load(), podLists.Load(), received.Load(), rc.QPS, rc.Burst) + } + // Exclusion changes routing membership, not authorization of issued identities. + node := &corev1.Node{} + if err := c.Get(t.Context(), client.ObjectKey{Name: "server-node"}, node); err != nil { + t.Fatal(err) + } + + node.Labels = map[string]string{wire.ExclusionLabel: ""} + if err := c.Update(t.Context(), node); err != nil { + t.Fatal(err) + } + + response, err := peer.Get(endpoint) + responseBody(t, response, err, 200) +} + +func integrationEnrollment(t *testing.T, rc *rest.Config, c client.Client, a *Application, ds *appsv1.DaemonSet, roots *x509.CertPool) *http.Client { + t.Helper() + + cfg := a.Topology.config + node := &corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "server-node"}} + + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.DataplaneServiceAccount}} + for _, obj := range []client.Object{node, sa} { + require.NoError(t, c.Create(t.Context(), obj)) + } + + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Name: "server-pod", Namespace: cfg.Namespace, OwnerReferences: []metav1.OwnerReference{*metav1.NewControllerRef(ds, appsv1.SchemeGroupVersion.WithKind("DaemonSet"))}}, Spec: *ds.Spec.Template.Spec.DeepCopy()} + + pod.Spec.NodeName = node.Name + if err := c.Create(t.Context(), pod); err != nil { + t.Fatal(err) + } + + pod.Status.PodIP = "192.0.2.1" + if err := c.Status().Update(t.Context(), pod); err != nil { + t.Fatal(err) + } + + eventually(t, "managed Pod published through informer", func() bool { + _, err := a.authority.Current() + return err == nil && strings.Contains(capturePublication(t, a.authority).encoded, string(node.UID)) + }) + + kube, err := kubernetes.NewForConfig(rc) + if err != nil { + t.Fatal(err) + } + + requestToken := func(audience string) string { + return boundPodToken(t, kube, pod, sa.Name, audience) + } + + _, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + csr, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{DNSNames: []string{"untrusted"}}, key) + if err != nil { + t.Fatal(err) + } + + body, err := wire.EncodeBootstrapRequest(wire.BootstrapRequest{SchemaVersion: 1, Cluster: cfg.Cluster, Enrollment: wire.EnrollmentID(testOtherUID), CSRDER: csr, Shares: wire.DefaultShares}) + if err != nil { + t.Fatal(err) + } + + transport := &http.Transport{TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS13, RootCAs: roots}} + t.Cleanup(transport.CloseIdleConnections) + anonymous := &http.Client{Transport: transport, Timeout: 10 * time.Second} + + var enrollment wire.BootstrapResponse + + for _, audience := range []string{"wrong-audience", wire.TokenAudience} { + req, err := http.NewRequestWithContext(t.Context(), "POST", "https://"+cfg.ControlAddress+wire.BootstrapPath, bytes.NewReader(body)) + if err != nil { + t.Fatal(err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+requestToken(audience)) + response, err := anonymous.Do(req) + + want := 401 + if audience == wire.TokenAudience { + want = 200 + } + + encoded := responseBody(t, response, err, want) + if want == 200 { + enrollment, err = wire.DecodeBootstrapResponse(bytes.NewReader(encoded)) + require.NoError(t, err) + require.Equal(t, wire.NodeID(node.UID), enrollment.Node) + } + } + + peerTransport := &http.Transport{TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS13, RootCAs: roots, Certificates: []tls.Certificate{{Certificate: enrollment.CertificateChain, PrivateKey: key}}}} + t.Cleanup(peerTransport.CloseIdleConnections) + + return &http.Client{Transport: peerTransport, Timeout: 15 * time.Second} +} + +func TestConcurrentStartupSingleCreate(t *testing.T) { + for _, failure := range []string{"none", "create denied", "create response lost", "marker response lost"} { + t.Run(failure, func(t *testing.T) { concurrentStartup(t, failure) }) + } +} + +func concurrentStartup(t *testing.T, failure string) { + t.Helper() + r := testTopology(t) + base := r.Client.(client.WithWatch) + + const replicas = 8 + + arrived := make(chan struct{}, replicas) + proceed := make(chan struct{}) + results := make(chan error, replicas) + + var creates, persisted, consumed atomic.Int32 + + boom := errors.New(failure) + writer := interceptor.NewClient(base, interceptor.Funcs{ + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if err := c.Update(ctx, obj, opts...); err != nil { + return err + } + + consumed.Add(1) + + if failure == "marker response lost" { + return boom + } + + return nil + }, + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + creates.Add(1) + + arrived <- struct{}{} + + select { + case <-ctx.Done(): + return ctx.Err() + case <-proceed: + } + + if failure == "create denied" { + return boom + } + + if err := c.Create(ctx, obj, opts...); err != nil { + return err + } + + persisted.Add(1) + + if failure == "create response lost" { + return boom + } + + return nil + }, + }) + + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + + for range replicas { + go func() { results <- Assemble(r.config, writer, base).Recover(ctx, writer) }() + } + + for range replicas { + select { + case <-arrived: + case <-ctx.Done(): + close(proceed) + t.Fatal("startups did not reach candidate Create") + } + } + + close(proceed) + + successes := successfulStartups(results, replicas) + + wantPersisted, wantConsumed, wantSuccess := int32(1), int32(1), replicas-1 + + switch failure { + case "none": + wantSuccess = replicas + case "create denied": + wantPersisted, wantConsumed, wantSuccess = 0, 0, 0 + } + + require.Equal(t, wantConsumed, consumed.Load()) + require.EqualValues(t, replicas, creates.Load()) + require.Equal(t, wantPersisted, persisted.Load(), "concurrent Create attempts must persist only one candidate") + require.Equal(t, wantSuccess, successes) + require.NoError(t, Assemble(r.config, base, base).Recover(t.Context(), base), "fresh/candidate/committed recovery must converge") +} + +func successfulStartups(results <-chan error, replicas int) int { + successes := 0 + + for range replicas { + if <-results == nil { + successes++ + } + } + + return successes +} + +func TestStartupRecoveryNeverWrites(t *testing.T) { + for _, state := range []string{"valid", "missing", "corrupt", "wrong binding", "mutable marker", "read denied"} { + t.Run(state, func(t *testing.T) { + r := initializedTopology(t) + + cm, _, err := readVersion(t.Context(), r.APIReader, r.config) + if err != nil { + t.Fatal(err) + } + + switch state { + case "missing": + err = r.Delete(t.Context(), cm) + case "corrupt": + cm.Data["sequence"] = "0" + err = r.Update(t.Context(), cm) + case "wrong binding": + cm.Annotations[installationUIDAnnotation] = "foreign" + err = r.Update(t.Context(), cm) + case "mutable marker": + marker, getErr := readInstallation(t.Context(), r.APIReader, r.config, false) + if getErr != nil { + t.Fatal(getErr) + } + + marker.Immutable = nil + err = r.Update(t.Context(), marker) + } + + if err != nil { + t.Fatal(err) + } + + writes := 0 + c := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + writes++ + return errors.New("unexpected Update") + }, + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + writes++ + return errors.New("unexpected Create") + }, + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if state == "read denied" { + return apierrors.NewForbidden(corev1.Resource("configmaps"), key.Name, errors.New("denied")) + } + + return c.Get(ctx, key, obj, opts...) + }, + }) + + ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) + defer cancel() + + err = Assemble(r.config, c, c).Recover(ctx, c) + if (err == nil) != (state == "valid") || writes != 0 { + t.Fatalf("recovery err=%v writes=%d", err, writes) + } + }) + } +} + +func TestStartupWaitsForWinnerGap(t *testing.T) { + r := testTopology(t) + base := r.Client.(client.WithWatch) + + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + + creating := make(chan struct{}) + observedGap := make(chan struct{}, 1) + results := make(chan error, 2) + + var creates atomic.Int32 + + winner := interceptor.NewClient(base, interceptor.Funcs{ + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + close(creating) + + select { + case <-ctx.Done(): + return ctx.Err() + case <-observedGap: + } + + if err := c.Create(ctx, obj, opts...); err != nil { + return err + } + + creates.Add(1) + + return nil + }, + }) + + go func() { results <- Assemble(r.config, winner, base).Recover(ctx, winner) }() + + select { + case <-ctx.Done(): + t.Fatal(ctx.Err()) + case <-creating: + } + + follower := interceptor.NewClient(base, interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + if key.Name == r.config.VersionConfigMapName && apierrors.IsNotFound(err) { + select { + case observedGap <- struct{}{}: + default: + } + } + + return err + }, + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + if err := c.Create(ctx, obj, opts...); err != nil { + return err + } + + creates.Add(1) + + return nil + }, + }) + + go func() { results <- Assemble(r.config, follower, follower).Recover(ctx, follower) }() + + for range 2 { + if err := <-results; err != nil { + t.Fatal(err) + } + } + + if creates.Load() != 1 { + t.Fatalf("Create attempts: %d", creates.Load()) + } +} + +func TestHTTPSCertificateIndependentOfWorkloadChanges(t *testing.T) { + for _, scenario := range []string{"node deleted", "pod recreated", "pod deleted", "pod owner revoked", "ds recreated"} { + t.Run(scenario, func(t *testing.T) { certificateAfterWorkloadChange(t, scenario) }) + } +} + +func certificateAfterWorkloadChange(t *testing.T, scenario string) { + t.Helper() + f := newServingFixture(t) + endpoint := f.start(t) + peer := f.client(t, &f.certificate) + response, err := peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusOK) + + var obj client.Object = &corev1.Pod{} + + key := client.ObjectKey{Namespace: "racer", Name: "worker-pod"} + + switch scenario { + case "node deleted": + obj, key = &corev1.Node{}, client.ObjectKey{Name: "worker"} + case "ds recreated": + obj, key = &appsv1.DaemonSet{}, client.ObjectKey{Namespace: "racer", Name: "racer-dataplane"} + } + + if err := f.a.Topology.Get(t.Context(), key, obj); err != nil { + t.Fatal(err) + } + + if scenario == "pod owner revoked" { + obj.(*corev1.Pod).OwnerReferences[0].UID = "revoked-owner" + if err := f.a.Topology.Update(t.Context(), obj); err != nil { + t.Fatal(err) + } + } else { + if err := f.a.Topology.Delete(t.Context(), obj); err != nil { + t.Fatal(err) + } + + if scenario == "pod recreated" || scenario == "ds recreated" { + obj.SetUID("replacement") + obj.SetResourceVersion("") + + if err := f.a.Topology.Create(t.Context(), obj); err != nil { + t.Fatal(err) + } + } + } + + reused := false + ctx := httptrace.WithClientTrace(t.Context(), &httptrace.ClientTrace{GotConn: func(info httptrace.GotConnInfo) { reused = info.Reused }}) + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint+wire.SnapshotPath, nil) + if err != nil { + t.Fatal(err) + } + + response, err = peer.Do(req) + responseBody(t, response, err, http.StatusOK) + + if !reused { + t.Fatal("workload change check did not reuse TLS connection") + } +} + +func TestHTTPSSnapshotDoesNotReadKubernetes(t *testing.T) { + f := newServingFixture(t) + + var reads atomic.Int64 + + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + reads.Add(1) + return errors.New("unexpected Kubernetes GET") + }, + List: func(context.Context, client.WithWatch, client.ObjectList, ...client.ListOption) error { + reads.Add(1) + return errors.New("unexpected Kubernetes LIST") + }, + }) + endpoint := f.start(t) + peer := f.client(t, &f.certificate) + + for range 2 { + response, err := peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusOK) + } + + if reads.Load() != 0 { + t.Fatalf("snapshot read Kubernetes: %d", reads.Load()) + } +} + +func catalogCache(name string, uid types.UID) racerv1.ClusterCache { + return racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: name, UID: uid}} +} + +// CanonicalSocketPaths keeps these tests on the production wire validation boundary. +func CanonicalSocketPaths(name string) (clientPath, originPath string, err error) { + return wire.CanonicalSocketPaths(name) +} + +func TestCanonicalSocketPaths(t *testing.T) { + for _, name := range []string{"a", "cache-a", "cache.a", strings.Repeat("a", 63) + "." + strings.Repeat("b", 18)} { + client, origin, err := CanonicalSocketPaths(name) + if err != nil || client != "/run/racer/"+name+"/client/socket" || origin != "/run/racer/"+name+"/origin/socket" || len(client) > 107 || len(origin) > 107 { + t.Fatalf("name %q: %q, %q, %v", name, client, origin, err) + } + } + + for _, name := range []string{"", ".", "..", "../cache", "cache/child", "Cache", "cache_a", "a..b", "-a", "a-", "a.-b", "a.b-", "a\x00", "caché", strings.Repeat("a", 64), strings.Repeat("a", 63) + "." + strings.Repeat("b", 19)} { + client, origin, err := CanonicalSocketPaths(name) + if !errors.Is(err, wire.InvalidRequest) || client != "" || origin != "" { + t.Fatalf("invalid name %q: %q, %q, %v", name, client, origin, err) + } + } +} + +func TestBuildCatalog(t *testing.T) { + caches := []racerv1.ClusterCache{catalogCache("cache-b", testOtherUID), catalogCache("cache-a", testNodeUID), catalogCache("cache-c", testDaemonSetUID)} + + original := make([]racerv1.ClusterCache, len(caches)) + for i := range caches { + original[i] = *caches[i].DeepCopy() + } + + got, err := BuildCatalog(caches) + + want := []wire.CacheDefinition{ + {ID: testNodeUID, Name: "cache-a", ClientSocket: "/run/racer/cache-a/client/socket", OriginSocket: "/run/racer/cache-a/origin/socket"}, + {ID: testOtherUID, Name: "cache-b", ClientSocket: "/run/racer/cache-b/client/socket", OriginSocket: "/run/racer/cache-b/origin/socket"}, + {ID: wire.CacheID(testDaemonSetUID), Name: "cache-c", ClientSocket: "/run/racer/cache-c/client/socket", OriginSocket: "/run/racer/cache-c/origin/socket"}, + } + if err != nil || !reflect.DeepEqual(got, want) || !reflect.DeepEqual(caches, original) { + t.Fatalf("catalog: %#v, %v; inputs: %#v", got, err, caches) + } + + slices.Reverse(caches) + + again, err := BuildCatalog(caches) + if err != nil || !reflect.DeepEqual(again, want) { + t.Fatalf("order changed catalog: %#v, %v", again, err) + } + + empty, err := BuildCatalog(nil) + if err != nil || empty == nil || len(empty) != 0 { + t.Fatalf("empty catalog: %#v, %v", empty, err) + } + // A terminating object still exists; removal follows its absence from inputs. + cache := catalogCache("cache-a", testNodeUID) + cache.DeletionTimestamp = &metav1.Time{} + + got, err = BuildCatalog([]racerv1.ClusterCache{cache}) + if err != nil || len(got) != 1 { + t.Fatalf("terminating cache: %#v, %v", got, err) + } + + cache.UID = testOtherUID + + recreated, err := BuildCatalog([]racerv1.ClusterCache{cache}) + if err != nil || recreated[0].ID == got[0].ID || recreated[0].ClientSocket != got[0].ClientSocket { + t.Fatalf("recreation: %#v, %v", recreated, err) + } +} + +func TestBuildCatalogRejectsWholeInvalidCandidate(t *testing.T) { + valid := catalogCache("cache-a", testNodeUID) + for name, invalid := range map[string]racerv1.ClusterCache{ + "missing uid": catalogCache("cache-b", ""), + "invalid uid": catalogCache("cache-b", "invalid"), + "uppercase uid": catalogCache("cache-b", "AAAAAAAA-AAAA-AAAA-AAAA-AAAAAAAAAAAA"), + "duplicate uid": catalogCache("cache-b", testNodeUID), + "duplicate name": catalogCache("cache-a", testOtherUID), + "unsafe name": catalogCache("../cache", testOtherUID), + "long path": catalogCache(strings.Repeat("a", 63)+"."+strings.Repeat("b", 19), testOtherUID), + } { + t.Run(name, func(t *testing.T) { + got, err := BuildCatalog([]racerv1.ClusterCache{valid, invalid}) + if !errors.Is(err, wire.InvalidRequest) || got != nil { + t.Fatalf("partial catalog escaped: %#v, %v", got, err) + } + }) + } +} + +// Run against the generated CRD in TestEnvtestServer, including Kubernetes' +// built-in metadata validation rather than a fake client or a CEL-only evaluator. +func integrationCacheNameAdmission(t *testing.T, c client.Client) { + t.Helper() + + for _, tt := range []struct { + name string + cacheName string + wantMessage string + }{ + {name: "single-character", cacheName: "a"}, + {name: "digits-and-hyphens", cacheName: "0.cache-1.2"}, + {name: "63-character-label", cacheName: strings.Repeat("a", 63)}, + {name: "63-character-hyphenated-label", cacheName: "0" + strings.Repeat("-", 61) + "9"}, + {name: "63-character-middle-label", cacheName: "a." + strings.Repeat("b", 63) + ".c"}, + {name: "64-total-multiple-labels", cacheName: strings.Repeat("a", 62) + ".b"}, + {name: "82-total-first-label-boundary", cacheName: strings.Repeat("a", 63) + "." + strings.Repeat("b", 18)}, + {name: "82-total-last-label-boundary", cacheName: strings.Repeat("a", 18) + "." + strings.Repeat("b", 63)}, + {name: "82-total-many-labels", cacheName: strings.Repeat("a.", 40) + "bb"}, + {name: "64-character-label", cacheName: strings.Repeat("a", 64), wantMessage: "each name label must be at most 63 characters"}, + {name: "64-character-first-label", cacheName: strings.Repeat("a", 64) + ".b", wantMessage: "each name label must be at most 63 characters"}, + {name: "64-character-middle-label", cacheName: "a." + strings.Repeat("b", 64) + ".c", wantMessage: "each name label must be at most 63 characters"}, + {name: "64-character-last-label", cacheName: "a." + strings.Repeat("b", 64), wantMessage: "each name label must be at most 63 characters"}, + {name: "82-character-single-label", cacheName: strings.Repeat("a", 82), wantMessage: "each name label must be at most 63 characters"}, + {name: "83-total-valid-labels", cacheName: strings.Repeat("a", 63) + "." + strings.Repeat("b", 19), wantMessage: "name must fit the canonical Unix socket path"}, + {name: "83-total-many-labels", cacheName: strings.Repeat("a.", 41) + "b", wantMessage: "name must fit the canonical Unix socket path"}, + {name: "empty", cacheName: "", wantMessage: "metadata.name"}, + {name: "uppercase", cacheName: "Cache", wantMessage: "metadata.name"}, + {name: "underscore", cacheName: "cache_a", wantMessage: "metadata.name"}, + {name: "non-ASCII", cacheName: "caché", wantMessage: "metadata.name"}, + {name: "slash", cacheName: "cache/child", wantMessage: "metadata.name"}, + {name: "leading-dot", cacheName: ".cache", wantMessage: "metadata.name"}, + {name: "trailing-dot", cacheName: "cache.", wantMessage: "metadata.name"}, + {name: "empty-label", cacheName: "cache..a", wantMessage: "metadata.name"}, + {name: "leading-hyphen", cacheName: "-cache", wantMessage: "metadata.name"}, + {name: "trailing-hyphen", cacheName: "cache-", wantMessage: "metadata.name"}, + {name: "label-leading-hyphen", cacheName: "cache.-a", wantMessage: "metadata.name"}, + {name: "label-trailing-hyphen", cacheName: "cache.a-", wantMessage: "metadata.name"}, + } { + t.Run(tt.name, func(t *testing.T) { + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: tt.cacheName}} + + err := c.Create(t.Context(), cache) + if err == nil { + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + if err := c.Delete(ctx, cache); err != nil { + t.Error(err) + } + }) + } + + if tt.wantMessage != "" { + if !apierrors.IsInvalid(err) || !strings.Contains(err.Error(), tt.wantMessage) { + t.Fatalf("create %q: want Invalid containing %q, got %v", tt.cacheName, tt.wantMessage, err) + } + + return + } + + if err != nil { + t.Fatalf("create %q: %v", tt.cacheName, err) + } + + catalog, err := BuildCatalog([]racerv1.ClusterCache{*cache}) + if err != nil || len(catalog) != 1 { + t.Fatalf("admitted cache cannot enter wire catalog: %v, %v", catalog, err) + } + }) + } +} + +func TestStartupDeadlineCancelsAPI(t *testing.T) { + for _, operation := range []string{"initial marker", "initial version", "update", "create", "final marker", "final version"} { + for _, boundary := range []string{"application", "earlier caller", "caller cancellation"} { + t.Run(operation+"/"+boundary, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { startupDeadline(t, operation, boundary) }) + }) + } + } +} + +func startupDeadline(t *testing.T, operation, boundary string) { + t.Helper() + r := testTopology(t) + + parent, cancel := context.WithCancel(t.Context()) + defer cancel() + + wantDuration := 30 * time.Second + wantErr := context.DeadlineExceeded + + switch boundary { + case "earlier caller": + var stop context.CancelFunc + + parent, stop = context.WithTimeout(parent, time.Second) + defer stop() + + wantDuration = time.Second + case "caller cancellation": + time.AfterFunc(time.Second, cancel) + wantDuration = time.Second + wantErr = context.Canceled + } + + blocked := false + block := func(ctx context.Context) error { + blocked = true + + <-ctx.Done() + + return ctx.Err() + } + created := false + c := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + stage := "initial " + if created { + stage = "final " + } + + if key.Name == r.config.InstallationConfigMapName { + stage += "marker" + } else { + stage += "version" + } + + if operation == stage { + return block(ctx) + } + // Spending budget before later calls catches per-call resets. + time.Sleep(100 * time.Millisecond) + + return c.Get(ctx, key, obj, opts...) + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if operation == "update" { + return block(ctx) + } + + return c.Update(ctx, obj, opts...) + }, + Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + if operation == "create" { + return block(ctx) + } + + created = true + + return c.Create(ctx, obj, opts...) + }, + }) + + start := time.Now() + + err := Assemble(r.config, c, c).Recover(parent, c) + if !blocked || !errors.Is(err, wantErr) || time.Since(start) != wantDuration { + t.Fatalf("blocked=%v error=%v elapsed=%v; want %v after %v", blocked, err, time.Since(start), wantErr, wantDuration) + } + + if boundary == "application" && parent.Err() != nil { + t.Fatalf("recovery canceled parent: %v", parent.Err()) + } +} + +func TestStartupDeadlineSuccess(t *testing.T) { + for _, installed := range []bool{false, true} { + t.Run(map[bool]string{false: "fresh", true: "installed"}[installed], func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { startupDeadlineSuccess(t, installed) }) + }) + } +} + +func startupDeadlineSuccess(t *testing.T, installed bool) { + t.Helper() + + r := testTopology(t) + if installed { + if err := ensureInstalled(t.Context(), r.Client, r.APIReader, r.config); err != nil { + t.Fatal(err) + } + } + + var recovery []context.Context + + reader := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if len(recovery) == 0 { + recovery = append(recovery, ctx) + + deadline, ok := ctx.Deadline() + if !ok || time.Until(deadline) != 30*time.Second { + t.Fatalf("startup deadline=%v present=%v", deadline, ok) + } + } + + return c.Get(ctx, key, obj, opts...) + }, + }) + + a := Assemble(r.config, r.Client, reader) + if err := a.Recover(t.Context(), r.Client); err != nil { + t.Fatal(err) + } + + if len(recovery) == 0 || !errors.Is(recovery[0].Err(), context.Canceled) { + t.Fatal("recovery context not released on success") + } + + if t.Context().Err() != nil || a.Server.Ready(nil) == nil { + t.Fatal("recovery canceled caller or granted serving authority") + } +} + +func TestStartupDeadlineAllowsCompetingInstallerWait(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := testTopology(t) + + c := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + return errors.New("winner stopped before version creation") + }, + }) + if err := ensureInstalled(t.Context(), c, c, r.config); err == nil { + t.Fatal("expected incomplete installation") + } + + first := true + reader := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if first { + first = false + + time.Sleep(2 * time.Second) + } + + return c.Get(ctx, key, obj, opts...) + }, + }) + start := time.Now() + + err := Assemble(r.config, r.Client, reader).Recover(t.Context(), r.Client) + if err != nil || time.Since(start) != 2*time.Second { + t.Fatalf("competing installer wait: error=%v elapsed=%v", err, time.Since(start)) + } + }) +} diff --git a/internal/racer/hints_test.go b/internal/racer/hints_test.go new file mode 100644 index 000000000..214666cc5 --- /dev/null +++ b/internal/racer/hints_test.go @@ -0,0 +1,427 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package racer + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "testing" + "time" + + "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" + "k8s.io/apimachinery/pkg/fields" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "k8s.io/client-go/util/workqueue" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + "sigs.k8s.io/controller-runtime/pkg/reconcile" + + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestRecoveryHintsUnchangedLargeSnapshot(t *testing.T) { + r := testTopology(t) + gets := 0 + r.APIReader = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + gets++ + return errors.New("unexpected authoritative read") + }, + }) + queued := 0 + r.enqueueHint = func(reconcile.Request) { queued++ } + update := authority.TopologyHints{Members: make(members.History)} + + for i := range 100_000 { + id := wire.NodeID(fmt.Sprintf("00000000-0000-4000-8000-%012d", i)) + member := wire.Member{Node: id, Shares: wire.DefaultShares, PeerEndpoint: "192.0.2.1:8082"} + encoded, err := json.Marshal(member) + require.NoError(t, err) + + update.Members[id] = member + update.Nodes.Items = append(update.Nodes.Items, corev1.Node{ObjectMeta: metav1.ObjectMeta{ + Name: string(id), UID: types.UID(id), Annotations: map[string]string{admittedMemberAnnotation: string(encoded)}, + }}) + } + + update.Nodes.Items = append(update.Nodes.Items, + corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "excluded", Labels: map[string]string{wire.ExclusionLabel: ""}}}, + corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "unadmitted"}}, + ) + for range 2 { + require.NoError(t, r.queueHints(t.Context(), update)) + require.Empty(t, r.hints) + } + + require.Zero(t, gets) + require.Zero(t, queued) +} + +func TestRecoveryHintsQueueLatestSnapshot(t *testing.T) { + for _, change := range []string{"desired", "satisfied", "absent", "unadmitted", "replacement", "excluded", "stale cache"} { + t.Run(change, func(t *testing.T) { + f := newServingFixture(t) + r := f.a.Topology + base := r.Client.(client.WithWatch) + + var node corev1.Node + require.NoError(t, base.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + member := acceptedMembers(t, r)[testNodeUID] + member.Shares = 7 + update := authority.TopologyHints{Nodes: corev1.NodeList{Items: []corev1.Node{*node.DeepCopy()}}, Members: members.History{testNodeUID: member}} + + queue := workqueue.NewTypedRateLimitingQueue(workqueue.DefaultTypedControllerRateLimiter[reconcile.Request]()) + defer queue.ShutDown() + + adds, gets, patches := 0, 0, 0 + r.enqueueHint = func(request reconcile.Request) { adds++; queue.Add(request) } + r.APIReader = interceptor.NewClient(base, interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + gets++ + return c.Get(ctx, key, obj, opts...) + }, + }) + r.Client = interceptor.NewClient(base, interceptor.Funcs{ + Patch: func(ctx context.Context, c client.WithWatch, obj client.Object, patch client.Patch, opts ...client.PatchOption) error { + patches++ + return c.Patch(ctx, obj, patch, opts...) + }, + }) + require.NoError(t, r.queueHints(f.ctx, update)) + require.Equal(t, 1, queue.Len()) + require.Zero(t, gets) + require.Zero(t, patches) + + switch change { + case "desired": + member.Shares = 9 + update.Members[testNodeUID] = member + case "satisfied": + update.Members = acceptedMembers(t, r) + case "absent": + update.Nodes.Items = nil + case "unadmitted": + update.Members = nil + case "replacement": + require.NoError(t, base.Delete(f.ctx, &node)) + node.UID = types.UID(testOtherUID) + node.ResourceVersion = "" + node.Annotations = nil + require.NoError(t, base.Create(f.ctx, &node)) + + member.Node = testOtherUID + update.Nodes.Items[0] = *node.DeepCopy() + update.Members = members.History{testOtherUID: member} + case "excluded": + node.Labels = map[string]string{wire.ExclusionLabel: ""} + require.NoError(t, base.Update(f.ctx, &node)) + update.Nodes.Items[0] = *node.DeepCopy() + update.Members = nil + case "stale cache": + encoded, err := json.Marshal(member) + require.NoError(t, err) + + node.Annotations[admittedMemberAnnotation] = string(encoded) + require.NoError(t, base.Update(f.ctx, &node)) + } + + require.NoError(t, r.queueHints(f.ctx, update)) + require.Equal(t, 1, adds, "pending work must retain its retry delay") + require.Equal(t, 1, queue.Len()) + require.Zero(t, gets) + + request, shutdown := queue.Get() + require.False(t, shutdown) + + result, err := r.Reconcile(f.ctx, request) + queue.Done(request) + queue.Forget(request) + require.NoError(t, err) + require.Equal(t, ctrl.Result{}, result) + require.Empty(t, r.hints) + require.Zero(t, queue.Len()) + + if change == "satisfied" || change == "absent" || change == "unadmitted" { + require.Zero(t, gets, "obsolete requests must not read Nodes") + require.Zero(t, patches) + + return + } + + require.Equal(t, 1, gets) + + if change == "stale cache" { + require.Zero(t, patches, "fresh read must suppress a redundant patch") + } else { + require.Equal(t, 1, patches) + } + + require.NoError(t, base.Get(f.ctx, client.ObjectKeyFromObject(&node), &node)) + + if change == "excluded" { + require.Empty(t, node.Annotations[admittedMemberAnnotation]) + return + } + + var saved wire.Member + require.NoError(t, json.Unmarshal([]byte(node.Annotations[admittedMemberAnnotation]), &saved)) + require.Equal(t, member, saved) + }) + } +} + +func TestRecoveryHintRetriesWithoutPublication(t *testing.T) { + for _, change := range []string{"unrelated", "replacement", "deleted", "excluded", "inputs", "spoof"} { + t.Run(change, func(t *testing.T) { + f := newServingFixture(t) + r := f.a.Topology + + var node corev1.Node + require.NoError(t, r.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + node.Annotations[wire.SharesAnnotation] = "7" + require.NoError(t, r.Update(f.ctx, &node)) + base := r.Client.(client.WithWatch) + fail := true + r.Client = interceptor.NewClient(base, interceptor.Funcs{Patch: func(ctx context.Context, c client.WithWatch, obj client.Object, patch client.Patch, opts ...client.PatchOption) error { + if fail { + return apierrors.NewConflict(corev1.Resource("nodes"), obj.GetName(), errors.New("concurrent update")) + } + + return c.Patch(ctx, obj, patch, opts...) + }}) + + var queued reconcile.Request + + r.enqueueHint = func(request reconcile.Request) { queued = request } + result, err := r.Reconcile(f.ctx, ctrl.Request{}) + require.NoError(t, err) + require.Equal(t, ctrl.Result{}, result) + require.Equal(t, "worker", queued.Name) + require.Equal(t, "hints", queued.Namespace) + published := capturePublication(t, r.authority) + result, err = r.Reconcile(f.ctx, queued) + require.True(t, apierrors.IsConflict(err)) + require.Equal(t, ctrl.Result{}, result) + require.NotEmpty(t, r.hints) + require.NoError(t, base.Get(f.ctx, client.ObjectKeyFromObject(&node), &node)) + + switch change { + case "unrelated": + node.Annotations["other"] = "preserved" + case "replacement": + require.NoError(t, base.Delete(f.ctx, &node)) + node.UID = types.UID(testOtherUID) + node.ResourceVersion = "" + node.Annotations = nil + require.NoError(t, base.Create(f.ctx, &node)) + case "deleted": + require.NoError(t, base.Delete(f.ctx, &node)) + case "excluded": + node.Labels = map[string]string{wire.ExclusionLabel: ""} + case "inputs": + node.Annotations[wire.SharesAnnotation] = "9" + case "spoof": + node.Annotations[admittedMemberAnnotation] = "spoof" + } + + if change != "replacement" && change != "deleted" { + require.NoError(t, base.Update(f.ctx, &node)) + } + + fail = false + // Any topology or durable publication call here fails the test. + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + List: func(context.Context, client.WithWatch, client.ObjectList, ...client.ListOption) error { + t.Fatal("hint retry listed topology") + return nil + }, + Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { + t.Fatal("hint retry republished") + return nil + }, + }) + result, err = r.Reconcile(f.ctx, queued) + require.NoError(t, err) + require.Equal(t, ctrl.Result{}, result) + require.Empty(t, r.hints) + require.Equal(t, published.encoded, capturePublication(t, r.authority).encoded) + + if change == "deleted" { + return + } + + require.NoError(t, base.Get(f.ctx, client.ObjectKeyFromObject(&node), &node)) + + if change == "excluded" || change == "replacement" { + require.Empty(t, node.Annotations[admittedMemberAnnotation]) + return + } + + member, err := wire.DecodeAdmittedMember(strings.NewReader(node.Annotations[admittedMemberAnnotation])) + require.NoError(t, err) + + if change == "inputs" { + require.NotEqualValues(t, 7, member.Shares) + } else { + require.EqualValues(t, 7, member.Shares) + } + + if change == "unrelated" { + require.Equal(t, "preserved", node.Annotations["other"]) + } + }) + } +} + +func TestRecoveryHintCancellationOverridesConflict(t *testing.T) { + for _, stage := range []string{"before", "read", "patch"} { + t.Run(stage, func(t *testing.T) { + f := newServingFixture(t) + r := f.a.Topology + + var node corev1.Node + require.NoError(t, r.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + node.Annotations[wire.SharesAnnotation] = "7" + require.NoError(t, r.Update(f.ctx, &node)) + + var queued reconcile.Request + + r.enqueueHint = func(request reconcile.Request) { queued = request } + _, err := r.Reconcile(f.ctx, ctrl.Request{}) + require.NoError(t, err) + require.Equal(t, "hints", queued.Namespace) + published := capturePublication(t, r.authority) + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + conflict := apierrors.NewConflict(corev1.Resource("nodes"), node.Name, errors.New("concurrent update")) + gets, patches := 0, 0 + r.APIReader = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + gets++ + + if stage == "read" { + cancel() + return conflict + } + + return c.Get(ctx, key, obj, opts...) + }, + }) + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Patch: func(context.Context, client.WithWatch, client.Object, client.Patch, ...client.PatchOption) error { + patches++ + + cancel() + + return conflict + }, + }) + + if stage == "before" { + cancel() + } + + result, err := r.Reconcile(ctx, queued) + require.ErrorIs(t, err, context.Canceled) + require.ErrorIs(t, err, reconcile.TerminalError(nil)) + require.Equal(t, ctrl.Result{}, result) + require.Len(t, r.hints, 1) + require.Equal(t, published.encoded, capturePublication(t, r.authority).encoded) + + if stage == "before" { + require.Zero(t, gets) + } else { + require.Equal(t, 1, gets) + } + + if stage == "patch" { + require.Equal(t, 1, patches) + } else { + require.Zero(t, patches) + } + }) + } +} + +func TestHandshakeTimeoutConfiguration(t *testing.T) { + values := map[string]string{"RACER_CLUSTER_ID": testOtherUID} + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + cfg, err := ConfigFromLookup(lookup) + require.NoError(t, err) + require.Equal(t, 5*time.Second, cfg.Limits.HandshakeTimeout) + + values["RACER_HANDSHAKE_TIMEOUT"] = "2s" + cfg, err = ConfigFromLookup(lookup) + require.NoError(t, err) + require.Equal(t, 2*time.Second, cfg.serverConfig().Limits.HandshakeTimeout) + require.Equal(t, 30*time.Second, cfg.Limits.WriteTimeout) + + for _, value := range []string{"", "bad", "0s", "-1s", "500ms"} { + values["RACER_HANDSHAKE_TIMEOUT"] = value + _, err := ConfigFromLookup(lookup) + require.ErrorContains(t, err, "RACER_HANDSHAKE_TIMEOUT") + require.NotErrorIs(t, err, wire.InvalidRequest) + } +} + +func TestNamedConfigMapCache(t *testing.T) { + cfg := testConfig(t) + options := managerOptions(cfg, runtime.NewScheme()) + require.Empty(t, options.LeaderElectionID) + require.Empty(t, options.LeaderElectionNamespace) + + for obj, config := range options.Cache.ByObject { + if _, ok := obj.(*corev1.ConfigMap); !ok { + continue + } + + require.Len(t, config.Namespaces, 1) + require.Contains(t, config.Namespaces, cfg.Namespace) + require.True(t, config.Field.Matches(fields.Set{"metadata.name": cfg.VersionConfigMapName})) + require.False(t, config.Field.Matches(fields.Set{"metadata.name": "unrelated"})) + require.False(t, config.Field.Matches(fields.Set{"metadata.name": cfg.InstallationConfigMapName}), "installation uses a separate exact-name source") + + return + } + + t.Fatal("ConfigMap cache missing") +} + +func TestControllerConflictUsesErrorBackoff(t *testing.T) { + for _, component := range []string{"topology", "keyring"} { + t.Run(component, func(t *testing.T) { + r, now := testKeyring(t) + runKeys(t, r) + _, _, state, _ := keyState(t, r) + *now = state.NextRotation + d := fixtureDependencies[r.authority] + conflict := apierrors.NewConflict(corev1.Resource("configmaps"), "version", errors.New("concurrent update")) + wrapped := interceptor.NewClient(d.Client.(client.WithWatch), interceptor.Funcs{Update: func(context.Context, client.WithWatch, client.Object, ...client.UpdateOption) error { return conflict }}) + + var target reconcile.Reconciler = r + if component == "topology" { + target = Assemble(r.config, wrapped, d.reader).Topology + } else { + d.Client = wrapped + } + + result, err := target.Reconcile(t.Context(), ctrl.Request{}) + require.ErrorIs(t, err, conflict) + require.Equal(t, ctrl.Result{}, result) + require.NotErrorIs(t, err, reconcile.TerminalError(nil)) + }) + } +} diff --git a/internal/racer/members/members.go b/internal/racer/members/members.go new file mode 100644 index 000000000..96354a74c --- /dev/null +++ b/internal/racer/members/members.go @@ -0,0 +1,328 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// Package members derives deterministic membership candidates from observed +// Kubernetes values. Discovery uses caller-provided readers; this package owns +// no clients, publication state, workload builders, or annotation writes. +package members + +import ( + "cmp" + "context" + "fmt" + "net/netip" + "slices" + "strconv" + "strings" + + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/types" + "sigs.k8s.io/controller-runtime/pkg/client" + + machinav1 "github.com/Azure/unbounded/api/machina/v1alpha3" + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/wire" +) + +// ReadWorkloadIdentities snapshots current DaemonSet ownership for one discovery +// or authorization pass. Missing and terminating workloads grant no ownership. +func ReadWorkloadIdentities(ctx context.Context, reader client.Reader, namespace, daemonSetName string) (WorkloadIdentities, error) { + ids := WorkloadIdentities{Namespace: namespace, Name: daemonSetName} + + var ds appsv1.DaemonSet + if err := reader.Get(ctx, client.ObjectKey{Namespace: namespace, Name: daemonSetName}, &ds); err != nil { + if apierrors.IsNotFound(err) { + return ids, nil + } + + return WorkloadIdentities{}, err + } + + if ds.DeletionTimestamp == nil { + ids.UID = ds.UID + } + + return ids, nil +} + +// ControllerPod checks the live namespace and service-account identity shared by +// replication discovery and authorization. Readiness is not an identity signal. +func ControllerPod(pod *corev1.Pod, namespace, serviceAccount string) bool { + return pod.Namespace == namespace && pod.UID != "" && pod.DeletionTimestamp == nil && pod.Spec.ServiceAccountName == serviceAccount && pod.Status.Phase != corev1.PodFailed && pod.Status.Phase != corev1.PodSucceeded +} + +// Observed workload ownership and endpoint selection. + +// WorkloadIdentities is the configured live DaemonSet identity, not a label selector. +// Owns performs no I/O; callers refresh the snapshot with ReadWorkloadIdentities +// on every topology or authorization pass so ownership never relies on stale UIDs. +type WorkloadIdentities struct { + Namespace string + Name string + UID types.UID +} + +// Owns checks ownership only. Callers retain their Pod, Node, service-account, +// token and readiness-independent membership checks. +func (ids WorkloadIdentities) Owns(pod *corev1.Pod) bool { + if pod == nil { + return false + } + + owner := metav1.GetControllerOf(pod) + if owner == nil || owner.APIVersion != "apps/v1" || owner.Kind != "DaemonSet" || owner.UID == "" { + return false + } + + if pod.Namespace != ids.Namespace { + return false + } + + return owner.Name == ids.Name && owner.UID == ids.UID +} + +// SelectEndpoint verifies workload ownership, ignores terminal/terminating/IP-less +// Pods, and chooses the newest creation time, breaking ties by UID. Readiness is ignored. +func SelectEndpoint(pods []corev1.Pod, ownership WorkloadIdentities, nodeName string, port uint16) (string, error) { + if nodeName == "" || port == 0 { + return "", wire.InvalidRequest + } + + var ( + selected *corev1.Pod + address netip.Addr + ) + + for i := range pods { + pod := &pods[i] + + ip, eligible := endpointAddress(pod, ownership, nodeName) + if !eligible { + continue + } + + if selected == nil || pod.CreationTimestamp.After(selected.CreationTimestamp.Time) || + pod.CreationTimestamp.Equal(&selected.CreationTimestamp) && pod.UID > selected.UID { + selected, address = pod, ip + } + } + + if selected == nil { + return "", wire.Unavailable + } + + return netip.AddrPortFrom(address, port).String(), nil +} + +func endpointAddress(pod *corev1.Pod, ownership WorkloadIdentities, nodeName string) (netip.Addr, bool) { + if pod.Spec.NodeName != nodeName || pod.DeletionTimestamp != nil || pod.UID == "" { + return netip.Addr{}, false + } + + if pod.Status.Phase == corev1.PodFailed || pod.Status.Phase == corev1.PodSucceeded || !ownership.Owns(pod) { + return netip.Addr{}, false + } + + ip, err := netip.ParseAddr(pod.Status.PodIP) + + return ip, err == nil && ip.Zone() == "" +} + +// Membership candidates and recovery. + +const ( + EnrolledSharesAnnotation = "racer.unbounded-cloud.io/enrolled-shares" + EnrolledRDMANICsAnnotation = "racer.unbounded-cloud.io/enrolled-rdma-nics" + AdmittedMemberAnnotation = "racer.unbounded-cloud.io/last-admitted-member" +) + +// History contains only previously published members. Node annotations provide +// UID-bound recovery hints when an entry is absent. Callers must not advance this +// history until publication succeeds. +type History map[wire.NodeID]wire.Member + +// Input is one observed topology snapshot. Pods are grouped by assigned Node +// name; each group is still checked for assignment and workload ownership. +// Reconcile borrows these values without mutating or retaining them. +type Input struct { + Nodes []corev1.Node + PodsByNode map[string][]corev1.Pod + Ownership WorkloadIdentities + PeerPort uint16 +} + +// Result owns its members (including nested NICs) and diagnostics. Members is a +// candidate, not accepted state, until the caller successfully publishes it. +type Result struct { + Members History + Diagnostics []Diagnostic +} + +type MemberAttributes struct { + Shares uint32 + RDMANICs []wire.RDMANIC +} + +type Diagnostic struct { + Object string + Field string + Reason string +} + +// ParseAnnotations distinguishes absent defaults from malformed proposed updates. +func ParseAnnotations(node *corev1.Node) (MemberAttributes, error) { + if node == nil { + return MemberAttributes{}, wire.InvalidRequest + } + + attributes := MemberAttributes{Shares: wire.DefaultShares, RDMANICs: []wire.RDMANIC{}} + + if _, explicit := node.Annotations[wire.SharesAnnotation]; !explicit { + if value := node.Annotations[EnrolledSharesAnnotation]; value != "" { + shares, err := strconv.ParseUint(value, 10, 32) + if err != nil || shares == 0 { + return MemberAttributes{}, wire.InvalidRequest + } + + attributes.Shares = uint32(shares) + } + } + + if value, present := node.Annotations[wire.SharesAnnotation]; present { + shares, err := strconv.ParseUint(value, 10, 32) + if err != nil || shares == 0 || strings.HasPrefix(value, "+") { + return MemberAttributes{}, fmt.Errorf("%s: %w", wire.SharesAnnotation, wire.InvalidRequest) + } + + attributes.Shares = uint32(shares) + } + + field := wire.RDMANICsAnnotation + + value, present := node.Annotations[field] + if !present { + field = EnrolledRDMANICsAnnotation + value, present = node.Annotations[field] + } + + if present { + nics, err := wire.DecodeRDMANICs(strings.NewReader(value)) + if err != nil { + return MemberAttributes{}, fmt.Errorf("%s: %w", field, err) + } + + attributes.RDMANICs = nics + } + + return attributes, nil +} + +// Reconcile preserves admitted values across gaps using UID-bound Node +// annotations on restart. Never-admitted nodes with unavailable or malformed +// required inputs are omitted. Deletion and exclusion remove membership. +// Annotations are accepted as one unit, independently of the endpoint. Site is +// always derived from current labels, never from admitted history, so a +// malformed annotation cannot retain a removed or changed RDMA boundary. +// The caller installs returned history only after the candidate publication commits. +// Inputs and nested accepted state are never mutated or aliased by the result. +func Reconcile(input Input, accepted History) (Result, error) { + if input.PeerPort == 0 { + return Result{}, wire.InvalidRequest + } + + nodes := slices.Clone(input.Nodes) + slices.SortFunc(nodes, func(a, b corev1.Node) int { return cmp.Compare(a.UID, b.UID) }) + + result := Result{Members: make(History), Diagnostics: []Diagnostic{}} + ids, names := map[types.UID]bool{}, map[string]bool{} + + for i := range nodes { + node := &nodes[i] + if !wire.ValidUUID(string(node.UID)) || node.Name == "" || ids[node.UID] || names[node.Name] { + return Result{}, fmt.Errorf("node identity: %w", wire.InvalidRequest) + } + + ids[node.UID], names[node.Name] = true, true + if _, excluded := node.Labels[wire.ExclusionLabel]; excluded { + continue + } + + result.reconcileNode(node, input, accepted) + } + + if len(result.Members) > wire.MaxMembers { + return Result{}, wire.TooLarge + } + + return result, nil +} + +func (result *Result) reconcileNode(node *corev1.Node, input Input, accepted History) { + id := wire.NodeID(node.UID) + + previous, known := accepted[id] + if !known { + if saved, err := wire.DecodeAdmittedMember(strings.NewReader(node.Annotations[AdmittedMemberAnnotation])); err == nil && saved.Node == id { + previous, known = saved, true + } + } + + attributes, annotationErr := ParseAnnotations(node) + if annotationErr != nil { + result.Diagnostics = append(result.Diagnostics, Diagnostic{Object: node.Name, Field: "annotations", Reason: annotationErr.Error()}) + + if known { + attributes = MemberAttributes{Shares: previous.Shares, RDMANICs: previous.RDMANICs} + } + } + // Node identity and peer port were validated by Reconcile, so the only + // possible endpoint error here is Unavailable. + endpoint, endpointErr := SelectEndpoint(input.PodsByNode[node.Name], input.Ownership, node.Name, input.PeerPort) + if endpointErr != nil { + result.Diagnostics = append(result.Diagnostics, Diagnostic{Object: node.Name, Field: "peer_endpoint", Reason: "no eligible managed Pod endpoint"}) + + if known { + endpoint = previous.PeerEndpoint + } + } + + if !known && (annotationErr != nil || endpointErr != nil) { + return + } + + result.Members[id] = wire.Member{ + Node: id, Shares: attributes.Shares, RDMANICs: wire.CanonicalRDMANICs(attributes.RDMANICs), PeerEndpoint: endpoint, + // Current labels revoke a stale RDMA boundary even if annotations are invalid. + Site: node.Labels[machinav1.MachineSiteLabelKey], + } +} + +// BuildCatalog derives cache identities from UIDs and paths from names, sorted by UID. +// An invalid catalog never partially replaces the currently served publication. +func BuildCatalog(caches []racerv1.ClusterCache) ([]wire.CacheDefinition, error) { + catalog := make([]wire.CacheDefinition, 0, len(caches)) + ids := make(map[wire.CacheID]bool, len(caches)) + + names := make(map[string]bool, len(caches)) + for _, cache := range caches { + id := wire.CacheID(cache.UID) + if !wire.ValidUUID(string(id)) || ids[id] || names[cache.Name] { + return nil, fmt.Errorf("cache identity: %w", wire.InvalidRequest) + } + + client, origin, err := wire.CanonicalSocketPaths(cache.Name) + if err != nil { + return nil, fmt.Errorf("cache socket paths: %w", err) + } + + ids[id], names[cache.Name] = true, true + catalog = append(catalog, wire.CacheDefinition{ID: id, Name: cache.Name, ClientSocket: client, OriginSocket: origin}) + } + + slices.SortFunc(catalog, func(a, b wire.CacheDefinition) int { return cmp.Compare(a.ID, b.ID) }) + + return catalog, nil +} diff --git a/internal/racer/members/members_test.go b/internal/racer/members/members_test.go new file mode 100644 index 000000000..087db106c --- /dev/null +++ b/internal/racer/members/members_test.go @@ -0,0 +1,1241 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package members + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "path" + "reflect" + "slices" + "strconv" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "k8s.io/apimachinery/pkg/util/intstr" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + machinav1 "github.com/Azure/unbounded/api/machina/v1alpha3" + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/wire" +) + +// Workload configuration, projections, and placement. + +func TestReadWorkloadIdentities(t *testing.T) { + scheme := runtime.NewScheme() + require.NoError(t, appsv1.AddToScheme(scheme)) + + for _, name := range []string{DataplaneDaemonSetName, "custom"} { + t.Run(name, func(t *testing.T) { + live := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: name, UID: "current"}} + terminating := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: PodNetworkDaemonSetName, UID: "terminating", Finalizers: []string{"test"}, DeletionTimestamp: ptr.To(metav1.Now())}} + c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(live, terminating).Build() + ids, err := ReadWorkloadIdentities(t.Context(), c, "racer", name) + require.NoError(t, err) + require.Equal(t, "racer", ids.Namespace) + require.Equal(t, name, ids.Name) + require.Equal(t, live.UID, ids.UID) + + require.NoError(t, c.Delete(t.Context(), live)) + ids, err = ReadWorkloadIdentities(t.Context(), c, "racer", name) + require.NoError(t, err) + require.Empty(t, ids.UID) + require.Equal(t, name, ids.Name) + }) + } + + boom := errors.New("read denied") + c := fake.NewClientBuilder().WithScheme(scheme).WithInterceptorFuncs(interceptor.Funcs{Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + return boom + }}).Build() + ids, err := ReadWorkloadIdentities(t.Context(), c, "racer", DataplaneDaemonSetName) + require.ErrorIs(t, err, boom) + require.Zero(t, ids) +} + +func TestControllerPodIdentity(t *testing.T) { + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", UID: "pod"}, Spec: corev1.PodSpec{ServiceAccountName: "controller"}} + require.True(t, ControllerPod(pod, "racer", "controller")) + + for name, mutate := range map[string]func(*corev1.Pod){ + "namespace": func(p *corev1.Pod) { p.Namespace = "other" }, + "UID": func(p *corev1.Pod) { p.UID = "" }, + "terminating": func(p *corev1.Pod) { p.DeletionTimestamp = ptr.To(metav1.Now()) }, + "service account": func(p *corev1.Pod) { p.Spec.ServiceAccountName = "other" }, + "failed": func(p *corev1.Pod) { p.Status.Phase = corev1.PodFailed }, + "succeeded": func(p *corev1.Pod) { p.Status.Phase = corev1.PodSucceeded }, + } { + t.Run(name, func(t *testing.T) { + invalid := pod.DeepCopy() + mutate(invalid) + require.False(t, ControllerPod(invalid, "racer", "controller")) + }) + } +} + +func workloadConfig(t *testing.T) Config { + t.Helper() + + return Config{ + Cluster: "11111111-1111-1111-1111-111111111111", Namespace: "racer", + ControlURL: "https://racer-controller.racer.svc:8443", DataplaneImage: "racer:test", + BootstrapTrustConfigMap: "racer-bootstrap-trust", PeerPort: 8082, + DataplaneServiceAccount: "racer-dataplane", DaemonSetName: "racer-dataplane", + } +} + +func TestWorkloadProjectionAndStorage(t *testing.T) { + cfg := workloadConfig(t) + + ds, err := DesiredDaemonSet(cfg) + if err != nil { + t.Fatal(err) + } + + wantSelector := map[string]string{"app.kubernetes.io/name": "racer-dataplane", "app.kubernetes.io/instance": cfg.DaemonSetName} + if !reflect.DeepEqual(ds.Spec.Selector.MatchLabels, wantSelector) || !reflect.DeepEqual(ds.Spec.Template.Labels, wantSelector) { + t.Fatalf("unexpected fresh workload selector: %v", ds.Spec.Selector) + } + + pod := ds.Spec.Template.Spec + if pod.AutomountServiceAccountToken == nil || *pod.AutomountServiceAccountToken || pod.ServiceAccountName != cfg.DataplaneServiceAccount { + t.Fatal("automatic API token or wrong service account") + } + + assertWorkloadVolumes(t, pod, cfg) + assertWorkloadMounts(t, pod.Containers[0]) + + requirements := pod.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms[0].MatchExpressions + if requirements[0].Key != wire.ExclusionLabel || requirements[0].Operator != corev1.NodeSelectorOpDoesNotExist { + t.Fatal("exclusion must test presence") + } + + for _, env := range pod.Containers[0].Env { + if env.ValueFrom != nil && (env.Name != "RACER_POD_IP" || !reflect.DeepEqual(env.ValueFrom, &corev1.EnvVarSource{FieldRef: &corev1.ObjectFieldSelector{APIVersion: "v1", FieldPath: "status.podIP"}})) { + t.Fatal("node identity must come from verified enrollment") + } + } + + if len(pod.Containers[0].Args) != 0 || len(pod.Containers[0].Command) != 0 { + t.Fatal("workload must use the image entrypoint") + } +} + +func assertWorkloadVolumes(t *testing.T, pod corev1.PodSpec, cfg Config) { + t.Helper() + + volumes := map[string]corev1.Volume{} + + for _, v := range pod.Volumes { + assertNoSecretVolume(t, v) + volumes[v.Name] = v + } + + token := volumes["token"].Projected.Sources[0].ServiceAccountToken + if token.Audience != wire.TokenAudience || *token.ExpirationSeconds != 3600 { + t.Fatal("wrong token projection") + } + + if _, exists := volumes["keyring"]; exists { + t.Fatal("shared keyring volume must not exist") + } + + if volumes["bootstrap"].ConfigMap.Name != cfg.BootstrapTrustConfigMap { + t.Fatal("bootstrap trust not independent") + } + + if !reflect.DeepEqual(volumes["bootstrap"].ConfigMap.Items, []corev1.KeyToPath{{Key: "ca.crt", Path: "ca.crt"}}) { + t.Fatal("only controller CA trust may be projected as cryptographic material") + } + + for _, name := range []string{"identity", "slabs", "sockets", "infiniband"} { + if volumes[name].HostPath == nil || *volumes[name].HostPath.Type != corev1.HostPathDirectoryOrCreate { + t.Fatalf("missing persistent host mount %s", name) + } + } + + if volumes["identity"].HostPath.Path == volumes["slabs"].HostPath.Path { + t.Fatal("private keys share disposable storage") + } + + wantDevices := &corev1.HostPathVolumeSource{Path: "/dev", Type: ptr.To(corev1.HostPathDirectory)} + if !reflect.DeepEqual(volumes["devices"].HostPath, wantDevices) { + t.Fatal("block device discovery requires the existing host /dev directory") + } +} + +func assertWorkloadMounts(t *testing.T, container corev1.Container) { + t.Helper() + + mounts := map[string]corev1.VolumeMount{} + for _, mount := range container.VolumeMounts { + mounts[mount.Name] = mount + if mount.SubPath != "" { + t.Fatal("subPath prevents projection rotation") + } + + if (mount.Name == "token" || mount.Name == "bootstrap") && !mount.ReadOnly { + t.Fatal("writable credential projection") + } + } + + wantMounts := map[string]corev1.VolumeMount{ + "token": {Name: "token", MountPath: "/var/run/racer-token", ReadOnly: true}, + "bootstrap": {Name: "bootstrap", MountPath: "/etc/racer/bootstrap", ReadOnly: true}, + "identity": {Name: "identity", MountPath: "/var/lib/racer/identity"}, + "slabs": {Name: "slabs", MountPath: "/var/lib/racer/slabs"}, + "sockets": {Name: "sockets", MountPath: "/run/racer"}, + "infiniband": {Name: "infiniband", MountPath: "/dev/infiniband", ReadOnly: true}, + "devices": {Name: "devices", MountPath: "/host/dev", ReadOnly: true}, + } + if !reflect.DeepEqual(mounts, wantMounts) { + t.Fatal("mounts must retain token, controller trust, private identity, slabs, sockets, RDMA devices and host devices only") + } +} + +func assertNoSecretVolume(t *testing.T, volume corev1.Volume) { + t.Helper() + + if volume.Secret != nil { + t.Fatal("dataplane must fetch shared keys over control HTTPS, not mount Secrets") + } + + if volume.Projected == nil { + return + } + + for _, source := range volume.Projected.Sources { + if source.Secret != nil { + t.Fatal("dataplane must not project Secrets") + } + } +} + +func TestWorkloadNativeRDMAAccess(t *testing.T) { + for _, hostNetwork := range []bool{false, true} { + t.Run(strconv.FormatBool(hostNetwork), func(t *testing.T) { + cfg := workloadConfig(t) + cfg.HostNetwork = hostNetwork + + ds, err := DesiredDaemonSet(cfg) + if err != nil { + t.Fatal(err) + } + + pod := ds.Spec.Template.Spec + if pod.HostNetwork != hostNetwork || pod.HostPID || pod.HostIPC { + t.Fatal("RDMA access must not implicitly enable host namespaces") + } + + wantSecurity := &corev1.SecurityContext{ + Privileged: ptr.To(true), AllowPrivilegeEscalation: ptr.To(true), ReadOnlyRootFilesystem: ptr.To(true), + } + if !reflect.DeepEqual(pod.Containers[0].SecurityContext, wantSecurity) || !reflect.DeepEqual(pod.SecurityContext, &corev1.PodSecurityContext{RunAsUser: ptr.To(int64(0))}) { + t.Fatal("native dataplane must explicitly run privileged as root with a read-only rootfs") + } + + wantVolume := corev1.Volume{Name: "infiniband", VolumeSource: corev1.VolumeSource{HostPath: &corev1.HostPathVolumeSource{ + Path: "/dev/infiniband", Type: ptr.To(corev1.HostPathDirectoryOrCreate), + }}} + + if len(pod.Volumes) != 7 { + t.Fatalf("unexpected volumes: %v", pod.Volumes) + } + + found := false + + for _, volume := range pod.Volumes { + if volume.Name == "infiniband" { + found = reflect.DeepEqual(volume, wantVolume) + } + } + + if !found { + t.Fatal("RDMA hostPath must tolerate an absent directory on HTTP-only nodes") + } + }) + } +} + +func TestWorkloadDataplaneEnvironment(t *testing.T) { + for _, port := range []uint16{8082, 7443, 9090, 9091, 65535} { + t.Run(strconv.Itoa(int(port)), func(t *testing.T) { + cfg := workloadConfig(t) + cfg.PeerPort = port + + ds, err := DesiredDaemonSet(cfg) + if err != nil { + t.Fatal(err) + } + + container := ds.Spec.Template.Spec.Containers[0] + + env := environmentValues(t, container) + + diagnosticsPort := "9090" + if port == 9090 { + diagnosticsPort = "9091" + } + // Rust Config::from_lookup consumes these settings after kubelet expands + // the Pod IP helper into both listener addresses. + expected := map[string]string{ + "RACER_CLUSTER_ID": string(cfg.Cluster), + "RACER_CONTROL_ENDPOINT": cfg.ControlURL, + "RACER_PEER_LISTEN": "[$(RACER_POD_IP)]:" + strconv.Itoa(int(port)), + "RACER_POD_IP": "", + "RACER_DIAGNOSTICS_LISTEN": "[$(RACER_POD_IP)]:" + diagnosticsPort, + "RACER_TRUST_BUNDLE": "/etc/racer/bootstrap/ca.crt", + "RACER_SERVICE_ACCOUNT_TOKEN": "/var/run/racer-token/token", + "RACER_IDENTITY_DIRECTORY": "/var/lib/racer/identity/private", + "RACER_SLAB_DIRECTORY": "/var/lib/racer/slabs", + "RACER_DEVICE_DIRECTORY": "/host/dev", + } + if len(env) != len(expected) { + t.Fatalf("unexpected configuration: %v", env) + } + + for name, value := range expected { + if env[name] != value { + t.Errorf("%s = %q, want %q", name, env[name], value) + } + } + + if container.Ports[0].ContainerPort != int32(port) { + t.Fatal("listener disagrees with advertised peer port") + } + + assertWorkloadReadiness(t, ds) + + assertEnvironmentMountPaths(t, container, env) + }) + } +} + +func assertWorkloadReadiness(t *testing.T, ds *appsv1.DaemonSet) { + t.Helper() + + rolling := ds.Spec.UpdateStrategy.RollingUpdate + + wantRolling := &appsv1.RollingUpdateDaemonSet{MaxUnavailable: ptr.To(intstr.FromInt32(1)), MaxSurge: ptr.To(intstr.FromInt32(0))} + if ds.Spec.UpdateStrategy.Type != appsv1.RollingUpdateDaemonSetStrategyType || !reflect.DeepEqual(rolling, wantRolling) || ds.Spec.MinReadySeconds != 10 { + t.Fatal("rollout must wait for sustained readiness with at most one unavailable Pod") + } + + container := ds.Spec.Template.Spec.Containers[0] + + expected := &corev1.Probe{ + ProbeHandler: corev1.ProbeHandler{HTTPGet: &corev1.HTTPGetAction{Path: "/readyz", Port: intstr.FromString("diagnostics"), Scheme: corev1.URISchemeHTTP}}, + PeriodSeconds: 5, TimeoutSeconds: 2, SuccessThreshold: 1, FailureThreshold: 1, + } + if !reflect.DeepEqual(container.ReadinessProbe, expected) || container.LivenessProbe != nil || container.StartupProbe != nil { + t.Fatal("must probe actual Pod-IP readiness without dependency-driven restarts") + } + + ports := map[string]int32{} + + for _, port := range container.Ports { + if port.Protocol != corev1.ProtocolTCP || port.HostPort != 0 || port.HostIP != "" { + t.Fatal("listeners must use Pod TCP ports") + } + + ports[port.Name] = port.ContainerPort + } + + if ports["diagnostics"] == 0 || ports["diagnostics"] == ports["peer"] { + t.Fatal("diagnostics missing or collides with peer listener") + } + + assertDiagnosticsEnvironment(t, container, ports["diagnostics"]) +} + +func assertDiagnosticsEnvironment(t *testing.T, container corev1.Container, diagnosticsPort int32) { + t.Helper() + + podIPSeen, diagnosticsSeen := false, false + + for _, env := range container.Env { + switch env.Name { + case "RACER_POD_IP": + if podIPSeen || env.Value != "" || !reflect.DeepEqual(env.ValueFrom, &corev1.EnvVarSource{FieldRef: &corev1.ObjectFieldSelector{APIVersion: "v1", FieldPath: "status.podIP"}}) { + t.Fatal("bind address must come from the downward API Pod IP") + } + + podIPSeen = true + case "RACER_DIAGNOSTICS_LISTEN": + if !podIPSeen || diagnosticsSeen || env.ValueFrom != nil || env.Value != "[$(RACER_POD_IP)]:"+strconv.Itoa(int(diagnosticsPort)) { + t.Fatal("diagnostics must expand the preceding Pod IP and match the probe port") + } + + diagnosticsSeen = true + default: + if env.ValueFrom != nil { + t.Fatal("unexpected indirect configuration") + } + } + } + + if !podIPSeen || !diagnosticsSeen { + t.Fatal("missing diagnostics bind configuration") + } +} + +func environmentValues(t *testing.T, container corev1.Container) map[string]string { + t.Helper() + + env := map[string]string{} + for _, value := range container.Env { + if _, exists := env[value.Name]; exists { + t.Fatalf("duplicate configuration: %s", value.Name) + } + + env[value.Name] = value.Value + } + + return env +} + +func assertEnvironmentMountPaths(t *testing.T, container corev1.Container, env map[string]string) { + t.Helper() + + mounts := map[string]string{} + for _, mount := range container.VolumeMounts { + mounts[mount.Name] = mount.MountPath + } + + for name, location := range map[string]string{ + "RACER_TRUST_BUNDLE": path.Join(mounts["bootstrap"], "ca.crt"), + "RACER_SERVICE_ACCOUNT_TOKEN": path.Join(mounts["token"], "token"), + "RACER_IDENTITY_DIRECTORY": path.Join(mounts["identity"], "private"), + "RACER_SLAB_DIRECTORY": mounts["slabs"], + "RACER_DEVICE_DIRECTORY": mounts["devices"], + } { + if env[name] != location { + t.Errorf("%s does not match its mounted projection or storage", name) + } + } +} + +func TestDesiredDaemonSetRejectsInvalidEndpoint(t *testing.T) { + for _, endpoint := range []string{"", "http://host", "https://host:0", "https://host:65536", "https://user@host", "https://host/path", "https://host?", "https://host/#fragment"} { + cfg := workloadConfig(t) + + cfg.ControlURL = endpoint + if ds, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) || ds != nil { + t.Fatalf("endpoint %q: %v", endpoint, err) + } + } +} + +func TestDesiredDaemonSetRejectsMissingImage(t *testing.T) { + for _, image := range []string{"", " \t"} { + cfg := workloadConfig(t) + + cfg.DataplaneImage = image + if ds, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) || ds != nil { + t.Fatalf("image %q: %v", image, err) + } + } +} + +func TestWorkloadConfigIdentityAndNames(t *testing.T) { + for name, mutate := range map[string]func(*Config){ + "cluster": func(c *Config) { c.Cluster = "invalid" }, + "missing cluster": func(c *Config) { c.Cluster = "" }, + "namespace": func(c *Config) { c.Namespace = "invalid.namespace" }, + "missing namespace": func(c *Config) { c.Namespace = "" }, + "zero port": func(c *Config) { c.PeerPort = 0 }, + } { + t.Run(name, func(t *testing.T) { + cfg := workloadConfig(t) + mutate(&cfg) + + if ds, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) || ds != nil { + t.Fatalf("invalid workload config accepted: %v", err) + } + }) + } + + for _, name := range []string{"daemonset", "trust", "serviceaccount"} { + t.Run(name, func(t *testing.T) { + for _, value := range []string{"", "../name", "Uppercase", strings.Repeat("a", 254)} { + cfg := workloadConfig(t) + fields := map[string]*string{ + "daemonset": &cfg.DaemonSetName, + "trust": &cfg.BootstrapTrustConfigMap, "serviceaccount": &cfg.DataplaneServiceAccount, + } + + *fields[name] = value + if ds, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) || ds != nil { + t.Fatalf("invalid resource name %q accepted: %v", value, err) + } + } + }) + } +} + +func TestWorkloadConfigLookupDefaultsAndOverrides(t *testing.T) { + values := map[string]string{ + "RACER_CLUSTER_ID": "11111111-1111-1111-1111-111111111111", + "RACER_CONTROL_URL": "https://controller:8443", "RACER_DATAPLANE_IMAGE": "racer:test", + } + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + + cfg, err := ConfigFromLookup(lookup) + if err != nil { + t.Fatal(err) + } + + if cfg.Namespace != "unbounded-system" || cfg.PeerPort != 8082 || cfg.DaemonSetName != "racer-dataplane" || cfg.DataplaneServiceAccount != "racer-dataplane" || cfg.BootstrapTrustConfigMap != "racer-bootstrap-trust" { + t.Fatalf("unexpected workload defaults: %+v", cfg) + } + + for key, value := range map[string]string{ + "POD_NAMESPACE": "custom", "RACER_PEER_PORT": "65535", "RACER_DAEMONSET_NAME": "custom.dataplane", + "RACER_DATAPLANE_SERVICE_ACCOUNT": "custom.account", + "RACER_BOOTSTRAP_TRUST_CONFIGMAP": "custom.trust", + } { + values[key] = value + } + + for key := range values { + t.Setenv(key, "invalid-process-value") + } + + cfg, err = ConfigFromLookup(lookup) + + want := Config{ + Cluster: "11111111-1111-1111-1111-111111111111", Namespace: "custom", PeerPort: 65535, + ControlURL: "https://controller:8443", DataplaneImage: "racer:test", DaemonSetName: "custom.dataplane", + DataplaneServiceAccount: "custom.account", BootstrapTrustConfigMap: "custom.trust", + } + if err != nil || !reflect.DeepEqual(cfg, want) { + t.Fatalf("custom lookup: %+v, %v", cfg, err) + } + + assertInvalidLookupValues(t, values) + + for _, port := range []string{"0", "65536", "-1"} { + values["RACER_PEER_PORT"] = port + if _, err := ConfigFromLookup(lookup); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("port %q accepted: %v", port, err) + } + } +} + +func assertInvalidLookupValues(t *testing.T, values map[string]string) { + t.Helper() + + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + for key := range values { + previous := values[key] + for _, invalid := range []string{"", "invalid value"} { + values[key] = invalid + // Image syntax remains the container runtime's responsibility. + if key == "RACER_DATAPLANE_IMAGE" && invalid != "" { + continue + } + + if _, err := ConfigFromLookup(lookup); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("%s=%q accepted: %v", key, invalid, err) + } + } + + values[key] = previous + } +} + +func TestWorkloadConfigIgnoresControllerRuntime(t *testing.T) { + want := workloadConfig(t) + values := map[string]string{ + "RACER_CLUSTER_ID": string(want.Cluster), "POD_NAMESPACE": want.Namespace, + "RACER_CONTROL_URL": want.ControlURL, "RACER_DATAPLANE_IMAGE": want.DataplaneImage, + } + + cfg, err := ConfigFromLookup(func(key string) (string, bool) { + if value, ok := values[key]; ok { + return value, true + } + + switch key { + case "RACER_PEER_PORT", "RACER_HOST_NETWORK", "RACER_POD_NETWORK_NODES", "RACER_DIAGNOSTICS_PORT", "RACER_DATAPLANE_SERVICE_ACCOUNT", "RACER_DAEMONSET_NAME", "RACER_BOOTSTRAP_TRUST_CONFIGMAP": + return "", false + default: + t.Errorf("workload parser requested runtime setting %s", key) + return "invalid", true + } + }) + if err != nil || !reflect.DeepEqual(cfg, want) { + t.Fatalf("workload needs runtime configuration: %+v, %v", cfg, err) + } + + if _, err := DesiredDaemonSet(cfg); err != nil { + t.Fatal(err) + } +} + +func TestWorkloadIgnoresLegacyKeyringSecret(t *testing.T) { + want := workloadConfig(t) + values := map[string]string{ + "RACER_CLUSTER_ID": string(want.Cluster), "POD_NAMESPACE": want.Namespace, + "RACER_CONTROL_URL": want.ControlURL, "RACER_DATAPLANE_IMAGE": want.DataplaneImage, + "RACER_CREDENTIALS_SECRET_NAME": "../obsolete", + } + + cfg, err := ConfigFromLookup(func(key string) (string, bool) { + value, ok := values[key] + return value, ok + }) + if err != nil || !reflect.DeepEqual(cfg, want) { + t.Fatalf("legacy environment changed workload configuration: %+v, %v", cfg, err) + } + + ds, err := DesiredDaemonSet(cfg) + if err != nil { + t.Fatal(err) + } + + for _, volume := range ds.Spec.Template.Spec.Volumes { + if volume.Secret != nil || volume.Name == "keyring" { + t.Fatal("legacy environment must not restore shared key mounts") + } + } +} + +func TestWorkloadNetworkPortBounds(t *testing.T) { + values := map[string]string{ + "RACER_CLUSTER_ID": "11111111-1111-1111-1111-111111111111", + "RACER_CONTROL_URL": "https://controller:8443", "RACER_DATAPLANE_IMAGE": "racer:test", + "RACER_HOST_NETWORK": "true", "RACER_PEER_PORT": "1024", "RACER_DIAGNOSTICS_PORT": "65535", + } + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + + for range 2 { + cfg, err := ConfigFromLookup(lookup) + if err != nil { + t.Fatal(err) + } + + if _, err := DesiredDaemonSet(cfg); err != nil { + t.Fatal(err) + } + + cfg.DiagnosticsPort = cfg.PeerPort + if _, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) { + t.Fatal("direct builder must reject colliding ports") + } + + cfg.DiagnosticsPort = 1023 + if _, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) { + t.Fatal("direct builder must reject privileged diagnostics") + } + + values["RACER_PEER_PORT"], values["RACER_DIAGNOSTICS_PORT"] = values["RACER_DIAGNOSTICS_PORT"], values["RACER_PEER_PORT"] + } +} + +func TestManagedNames(t *testing.T) { + for _, tt := range []struct { + name string + want []string + }{ + {DataplaneDaemonSetName, []string{DataplaneDaemonSetName}}, + {"custom-racer", []string{"custom-racer"}}, + {PodNetworkDaemonSetName, []string{PodNetworkDaemonSetName}}, + {"", []string{""}}, + } { + t.Run(tt.name, func(t *testing.T) { + got := ManagedNames(tt.name) + if !reflect.DeepEqual(got, tt.want) { + t.Fatalf("ManagedNames(%q) = %v, want %v", tt.name, got, tt.want) + } + + got[0] = "mutated" + + if !reflect.DeepEqual(ManagedNames(tt.name), tt.want) { + t.Fatal("caller mutation changed managed names") + } + }) + } +} + +// Observed membership, recovery, and catalog candidates. + +const ( + nodeID = "11111111-1111-4111-8111-111111111111" + otherID = "22222222-2222-4222-8222-222222222222" +) + +func observedInput() Input { + return Input{ + Nodes: []corev1.Node{{ObjectMeta: metav1.ObjectMeta{Name: "node", UID: nodeID}}}, + PodsByNode: map[string][]corev1.Pod{"node": {{ + ObjectMeta: metav1.ObjectMeta{ + Namespace: "racer", Name: "pod", UID: "pod-uid", + OwnerReferences: []metav1.OwnerReference{{APIVersion: "apps/v1", Kind: "DaemonSet", Name: "dataplane", UID: "workload-uid", Controller: ptr.To(true)}}, + }, + Spec: corev1.PodSpec{NodeName: "node"}, + Status: corev1.PodStatus{PodIP: "192.0.2.1"}, + }}}, + Ownership: WorkloadIdentities{Namespace: "racer", Name: "dataplane", UID: "workload-uid"}, + PeerPort: 7443, + } +} + +func TestReconcileCandidateRequiresExplicitHistoryAdvance(t *testing.T) { + input := observedInput() + input.Nodes[0].Annotations = map[string]string{wire.SharesAnnotation: "8", wire.RDMANICsAnnotation: `[{"device":"nic","port":1,"rail":0,"numa_node":1}]`} + accepted := History{} + + candidate, err := Reconcile(input, accepted) + if err != nil || len(candidate.Diagnostics) != 0 || len(candidate.Members) != 1 || len(accepted) != 0 { + t.Fatalf("candidate changed accepted history: %+v, %v, %v", candidate, accepted, err) + } + + input.PodsByNode = nil + input.Nodes[0].Annotations[wire.SharesAnnotation] = "invalid" + + unpublished, err := Reconcile(input, accepted) + if err != nil || len(unpublished.Members) != 0 || len(unpublished.Diagnostics) != 2 { + t.Fatalf("unpublished candidate survived a gap: %+v, %v", unpublished, err) + } + + accepted = candidate.Members // Simulate a successful publication. + + retained, err := Reconcile(input, accepted) + if err != nil || !reflect.DeepEqual(retained.Members, accepted) { + t.Fatalf("accepted member did not survive gap: %+v, %v", retained, err) + } + + retained.Members[nodeID].RDMANICs[0].Device = "changed" + *retained.Members[nodeID].RDMANICs[0].NUMANode = 9 + delete(retained.Members, nodeID) + + if accepted[nodeID].RDMANICs[0].Device != "nic" || *accepted[nodeID].RDMANICs[0].NUMANode != 1 { + t.Fatal("result aliases nested accepted history") + } +} + +func TestReconcileRecoveryIsUIDBoundAndSiteIsCurrent(t *testing.T) { + input := observedInput() + input.Nodes[0].Labels = map[string]string{machinav1.MachineSiteLabelKey: "old-site"} + + initial, err := Reconcile(input, nil) + if err != nil { + t.Fatal(err) + } + + saved, err := json.Marshal(initial.Members[nodeID]) + if err != nil { + t.Fatal(err) + } + + input.PodsByNode = nil + input.Nodes[0].Annotations = map[string]string{AdmittedMemberAnnotation: string(saved), wire.SharesAnnotation: "invalid"} + input.Nodes[0].Labels = nil + + recovered, err := Reconcile(input, nil) + if err != nil || len(recovered.Members) != 1 || recovered.Members[nodeID].Site != "" || recovered.Members[nodeID].PeerEndpoint != "192.0.2.1:7443" { + t.Fatalf("recovery retained stale site or lost endpoint: %+v, %v", recovered, err) + } + + input.Nodes[0].UID = otherID + + recreated, err := Reconcile(input, nil) + if err != nil || len(recreated.Members) != 0 { + t.Fatalf("recreated Node inherited history: %+v, %v", recreated, err) + } + + input.Nodes[0].UID = nodeID + input.Nodes[0].Annotations[AdmittedMemberAnnotation] = "malformed" + + invalid, err := Reconcile(input, nil) + if err != nil || len(invalid.Members) != 0 { + t.Fatalf("malformed recovery admitted: %+v, %v", invalid, err) + } +} + +func TestReconcileRejectsWholeInputAndOrdersDiagnostics(t *testing.T) { + input := observedInput() + input.Nodes = append(input.Nodes, input.Nodes[0]) + + result, err := Reconcile(input, nil) + if !errors.Is(err, wire.InvalidRequest) || result.Members != nil || result.Diagnostics != nil { + t.Fatalf("partial result escaped: %+v, %v", result, err) + } + + input.Nodes[1] = corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "other", UID: otherID}} + input.PodsByNode = nil + + first, err := Reconcile(input, nil) + if err != nil || len(first.Diagnostics) != 2 || first.Diagnostics[0].Object != "node" || first.Diagnostics[1].Object != "other" { + t.Fatalf("unexpected diagnostics: %+v, %v", first, err) + } + + slices.Reverse(input.Nodes) + + second, err := Reconcile(input, nil) + if err != nil || !reflect.DeepEqual(first, second) { + t.Fatalf("order changed result: %+v, %v", second, err) + } + + input.PeerPort = 0 + if _, err := Reconcile(input, nil); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("zero port admitted: %v", err) + } +} + +func TestObservedOwnershipAndEndpoint(t *testing.T) { + input := observedInput() + + pod := input.PodsByNode["node"][0] + if !input.Ownership.Owns(&pod) || input.Ownership.Owns(nil) { + t.Fatal("observed ownership mismatch") + } + + for _, identity := range []WorkloadIdentities{{Name: "dataplane"}, {Name: "other", UID: "workload-uid"}, {Name: "dataplane", UID: "recreated"}} { + ownership := identity + + ownership.Namespace = input.Ownership.Namespace + if ownership.Owns(&pod) { + t.Fatalf("unobserved identity admitted: %+v", identity) + } + } + // A newer Pod from another workload is not eligible. + newer := pod.DeepCopy() + newer.UID = "z" + newer.OwnerReferences[0].Name, newer.OwnerReferences[0].UID = "podnet", "podnet-uid" + newer.Status.PodIP = "2001:db8::1" + + endpoint, err := SelectEndpoint([]corev1.Pod{*newer, pod}, input.Ownership, "node", 7443) + if err != nil || endpoint != "192.0.2.1:7443" { + t.Fatalf("mixed workload endpoint: %q, %v", endpoint, err) + } + + newer.Status.Phase = corev1.PodSucceeded + + endpoint, err = SelectEndpoint([]corev1.Pod{*newer, pod}, input.Ownership, "node", 7443) + if err != nil || endpoint != "192.0.2.1:7443" { + t.Fatalf("terminal endpoint admitted: %q, %v", endpoint, err) + } +} + +func TestCatalogOrderingAndWholeCandidateValidation(t *testing.T) { + caches := []racerv1.ClusterCache{ + {ObjectMeta: metav1.ObjectMeta{Name: "cache-b", UID: types.UID(otherID)}}, + {ObjectMeta: metav1.ObjectMeta{Name: "cache-a", UID: types.UID(nodeID)}}, + } + + catalog, err := BuildCatalog(caches) + if err != nil || len(catalog) != 2 || catalog[0].ID != nodeID || catalog[0].ClientSocket != "/run/racer/cache-a/client/socket" || caches[0].Name != "cache-b" { + t.Fatalf("catalog order or input mutation: %+v, %v", catalog, err) + } + + caches[1].Name = "../invalid" + + catalog, err = BuildCatalog(caches) + if !errors.Is(err, wire.InvalidRequest) || catalog != nil { + t.Fatalf("partial catalog escaped: %+v, %v", catalog, err) + } +} + +func TestParseAnnotations(t *testing.T) { + nicJSON := `[{"device":"nic","port":1,"rail":0}]` + + nics := []wire.RDMANIC{{Device: "nic", Port: 1, Rail: 0}} + for _, tt := range []struct { + name string + annotations map[string]string + shares uint32 + nics []wire.RDMANIC + errorField string + }{ + {name: "defaults", shares: wire.DefaultShares}, + {name: "empty enrolled shares", annotations: map[string]string{EnrolledSharesAnnotation: ""}, shares: wire.DefaultShares}, + {name: "enrolled shares", annotations: map[string]string{EnrolledSharesAnnotation: "8"}, shares: 8}, + {name: "explicit overrides invalid enrolled", annotations: map[string]string{EnrolledSharesAnnotation: "invalid", wire.SharesAnnotation: "9"}, shares: 9}, + {name: "maximum shares", annotations: map[string]string{wire.SharesAnnotation: "4294967295"}, shares: 4294967295}, + {name: "zero enrolled", annotations: map[string]string{EnrolledSharesAnnotation: "0"}, errorField: "bare"}, + {name: "invalid enrolled", annotations: map[string]string{EnrolledSharesAnnotation: "invalid"}, errorField: "bare"}, + {name: "zero explicit", annotations: map[string]string{wire.SharesAnnotation: "0"}, errorField: wire.SharesAnnotation}, + {name: "empty explicit", annotations: map[string]string{wire.SharesAnnotation: ""}, errorField: wire.SharesAnnotation}, + {name: "plus explicit", annotations: map[string]string{wire.SharesAnnotation: "+1"}, errorField: wire.SharesAnnotation}, + {name: "negative explicit", annotations: map[string]string{wire.SharesAnnotation: "-1"}, errorField: wire.SharesAnnotation}, + {name: "overflow explicit", annotations: map[string]string{wire.SharesAnnotation: "4294967296"}, errorField: wire.SharesAnnotation}, + {name: "enrolled NICs", annotations: map[string]string{EnrolledRDMANICsAnnotation: nicJSON}, shares: wire.DefaultShares, nics: nics}, + {name: "explicit NICs", annotations: map[string]string{wire.RDMANICsAnnotation: nicJSON}, shares: wire.DefaultShares, nics: nics}, + {name: "explicit clears enrolled NICs", annotations: map[string]string{wire.RDMANICsAnnotation: "[]", EnrolledRDMANICsAnnotation: nicJSON}, shares: wire.DefaultShares}, + {name: "invalid enrolled NICs", annotations: map[string]string{EnrolledRDMANICsAnnotation: "invalid"}, errorField: EnrolledRDMANICsAnnotation}, + {name: "invalid explicit NICs", annotations: map[string]string{wire.RDMANICsAnnotation: "", EnrolledRDMANICsAnnotation: nicJSON}, errorField: wire.RDMANICsAnnotation}, + } { + t.Run(tt.name, func(t *testing.T) { + node := &corev1.Node{ObjectMeta: metav1.ObjectMeta{Annotations: tt.annotations}} + before := node.DeepCopy() + + got, err := ParseAnnotations(node) + if !reflect.DeepEqual(node, before) { + t.Fatal("annotations mutated") + } + + if tt.errorField != "" { + assertAnnotationError(t, got, err, tt.errorField) + return + } + + if err != nil || got.Shares != tt.shares || !slices.EqualFunc(got.RDMANICs, tt.nics, func(a, b wire.RDMANIC) bool { return reflect.DeepEqual(a, b) }) { + t.Fatalf("attributes = %+v, %v", got, err) + } + }) + } + + got, err := ParseAnnotations(nil) + assertAnnotationError(t, got, err, "bare") +} + +func assertAnnotationError(t *testing.T, got MemberAttributes, err error, field string) { + t.Helper() + + if !errors.Is(err, wire.InvalidRequest) || !reflect.DeepEqual(got, MemberAttributes{}) { + t.Fatalf("partial attributes or wrong error: %+v, %v", got, err) + } + + if field == "bare" { + if err != wire.InvalidRequest { + t.Fatalf("wrapped enrolled error: %v", err) + } + } else if !strings.HasPrefix(err.Error(), field+": ") { + t.Fatalf("missing field %s: %v", field, err) + } +} + +func TestEndpointEligibility(t *testing.T) { + for name, mutate := range map[string]func(*corev1.Pod){ + "wrong node": func(p *corev1.Pod) { p.Spec.NodeName = "other" }, + "terminating": func(p *corev1.Pod) { p.DeletionTimestamp = ptr.To(metav1.Now()) }, + "missing UID": func(p *corev1.Pod) { p.UID = "" }, + "failed": func(p *corev1.Pod) { p.Status.Phase = corev1.PodFailed }, + "succeeded": func(p *corev1.Pod) { p.Status.Phase = corev1.PodSucceeded }, + "missing IP": func(p *corev1.Pod) { p.Status.PodIP = "" }, + "invalid IP": func(p *corev1.Pod) { p.Status.PodIP = "hostname" }, + "zoned IP": func(p *corev1.Pod) { p.Status.PodIP = "fe80::1%eth0" }, + "wrong namespace": func(p *corev1.Pod) { p.Namespace = "other" }, + "no owner": func(p *corev1.Pod) { p.OwnerReferences = nil }, + "not controller": func(p *corev1.Pod) { p.OwnerReferences[0].Controller = ptr.To(false) }, + "wrong API": func(p *corev1.Pod) { p.OwnerReferences[0].APIVersion = "apps/v2" }, + "wrong kind": func(p *corev1.Pod) { p.OwnerReferences[0].Kind = "Deployment" }, + "empty owner UID": func(p *corev1.Pod) { p.OwnerReferences[0].UID = "" }, + } { + t.Run(name, func(t *testing.T) { + input := observedInput() + pod := input.PodsByNode["node"][0] + mutate(&pod) + + endpoint, err := SelectEndpoint([]corev1.Pod{pod}, input.Ownership, "node", 7443) + if !errors.Is(err, wire.Unavailable) || endpoint != "" { + t.Fatalf("ineligible endpoint: %q, %v", endpoint, err) + } + }) + } +} + +func TestEndpointSelectionOrderAndArguments(t *testing.T) { + input := observedInput() + older := input.PodsByNode["node"][0] + newer := *older.DeepCopy() + newer.UID = "a" + newer.CreationTimestamp = metav1.NewTime(time.Unix(100, 0)) + newer.Status.PodIP = "192.0.2.2" + + newer.Status.Conditions = []corev1.PodCondition{{Type: corev1.PodReady, Status: corev1.ConditionFalse}} + for _, pods := range [][]corev1.Pod{{older, newer}, {newer, older}} { + endpoint, err := SelectEndpoint(pods, input.Ownership, "node", 7443) + if err != nil || endpoint != "192.0.2.2:7443" { + t.Fatalf("creation time or readiness selection: %q, %v", endpoint, err) + } + } + + for _, tt := range []struct { + node string + port uint16 + }{{"", 7443}, {"node", 0}} { + endpoint, err := SelectEndpoint(nil, input.Ownership, tt.node, tt.port) + if !errors.Is(err, wire.InvalidRequest) || endpoint != "" { + t.Fatalf("invalid arguments: %q, %v", endpoint, err) + } + } +} + +func TestReconcileIdentityValidation(t *testing.T) { + for name, mutate := range map[string]func(*Input){ + "invalid UID": func(i *Input) { i.Nodes[0].UID = "invalid" }, + "empty name": func(i *Input) { i.Nodes[0].Name = "" }, + "duplicate name": func(i *Input) { n := *i.Nodes[0].DeepCopy(); n.UID = otherID; i.Nodes = append(i.Nodes, n) }, + "excluded invalid identity": func(i *Input) { + i.Nodes[0].UID = "invalid" + i.Nodes[0].Labels = map[string]string{wire.ExclusionLabel: ""} + }, + } { + t.Run(name, func(t *testing.T) { + input := observedInput() + mutate(&input) + + got, err := Reconcile(input, nil) + if !errors.Is(err, wire.InvalidRequest) || !reflect.DeepEqual(got, Result{}) { + t.Fatalf("invalid identity: %+v, %v", got, err) + } + }) + } +} + +func TestReconcileRemovalAndLegacyDiagnostics(t *testing.T) { + input := observedInput() + + initial, err := Reconcile(input, nil) + if err != nil { + t.Fatal(err) + } + + input.Nodes[0].Annotations = map[string]string{"racer.unbounded-cloud.io/rails": "ignored", "racer.unbounded-cloud.io/aligned-rails": "ignored"} + result, err := Reconcile(input, initial.Members) + + wantDiagnostics := []Diagnostic{} + if err != nil || !reflect.DeepEqual(result.Diagnostics, wantDiagnostics) || !reflect.DeepEqual(result.Members, initial.Members) { + t.Fatalf("legacy diagnostics: %+v, %v", result, err) + } + + for _, excluded := range []bool{true, false} { + if excluded { + input.Nodes[0].Labels = map[string]string{wire.ExclusionLabel: "false"} + } else { + input.Nodes = nil + } + + result, err := Reconcile(input, initial.Members) + if err != nil || len(result.Members) != 0 || len(result.Diagnostics) != 0 { + t.Fatalf("removed node retained: %+v, %v", result, err) + } + } +} + +func TestReconcileMemberLimit(t *testing.T) { + input := Input{PeerPort: 7443} + accepted := History{} + + for i := range wire.MaxMembers + 1 { + id := fmt.Sprintf("%08x-1111-4111-8111-111111111111", i) + input.Nodes = append(input.Nodes, corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: fmt.Sprintf("node-%d", i), UID: types.UID(id)}}) + accepted[wire.NodeID(id)] = wire.Member{Node: wire.NodeID(id), Shares: 1, PeerEndpoint: "192.0.2.1:7443"} + } + + got, err := Reconcile(input, accepted) + if !errors.Is(err, wire.TooLarge) || !reflect.DeepEqual(got, Result{}) { + t.Fatalf("oversized membership: %d, %v", len(got.Members), err) + } + + input.Nodes = input.Nodes[:wire.MaxMembers] + + got, err = Reconcile(input, accepted) + if err != nil || len(got.Members) != wire.MaxMembers { + t.Fatalf("boundary membership: %d, %v", len(got.Members), err) + } +} + +func TestCatalogRejectsInvalidIdentities(t *testing.T) { + valid := racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache", UID: nodeID}} + for _, identity := range []metav1.ObjectMeta{{Name: "other", UID: "invalid"}, {Name: "other", UID: nodeID}, {Name: "cache", UID: otherID}} { + invalid := valid + invalid.ObjectMeta = identity + + got, err := BuildCatalog([]racerv1.ClusterCache{valid, invalid}) + if !errors.Is(err, wire.InvalidRequest) || got != nil { + t.Fatalf("invalid catalog identity: %+v, %v", got, err) + } + } +} + +func TestConfigValidationErrors(t *testing.T) { + for _, tt := range []struct { + name string + mutate func(*Config) + message string + }{ + {"exceptions without host networking", func(c *Config) { c.PodNetworkNodes = []string{"node"} }, "pod network exceptions require host networking and the fixed dataplane name"}, + {"exceptions with custom name", func(c *Config) { + c.HostNetwork = true + c.DaemonSetName = "custom" + c.PodNetworkNodes = []string{"node"} + }, "pod network exceptions require host networking and the fixed dataplane name"}, + {"duplicate nodes", func(c *Config) { c.HostNetwork = true; c.PodNetworkNodes = []string{"node", "node"} }, "pod network nodes must be unique valid node names"}, + {"invalid node", func(c *Config) { c.HostNetwork = true; c.PodNetworkNodes = []string{"Uppercase"} }, "pod network nodes must be unique valid node names"}, + {"privileged peer port", func(c *Config) { c.PeerPort = 1023 }, "cluster, namespace, or peer port"}, + {"long selector label", func(c *Config) { c.DaemonSetName = strings.Repeat("a", 64) }, "DaemonSet name must fit a label value"}, + {"escaped path", func(c *Config) { c.ControlURL = "https://host/%2f" }, "workload endpoint or image"}, + {"query", func(c *Config) { c.ControlURL = "https://host/?x=y" }, "workload endpoint or image"}, + {"malformed URL", func(c *Config) { c.ControlURL = "https://host:%" }, "workload endpoint or image"}, + {"missing host", func(c *Config) { c.ControlURL = "https:///" }, "workload endpoint or image"}, + {"zero endpoint port", func(c *Config) { c.ControlURL = "https://host:0" }, "workload endpoint port"}, + {"overflow endpoint port", func(c *Config) { c.ControlURL = "https://host:65536" }, "workload endpoint port"}, + } { + t.Run(tt.name, func(t *testing.T) { + cfg := workloadConfig(t) + tt.mutate(&cfg) + err := cfg.Validate() + + want := tt.message + ": " + wire.InvalidRequest.Error() + if !errors.Is(err, wire.InvalidRequest) || err.Error() != want { + t.Fatalf("error = %v, want %s", err, want) + } + + sets, err := DesiredDaemonSets(cfg) + if !errors.Is(err, wire.InvalidRequest) || sets != nil { + t.Fatalf("invalid planner config: %+v, %v", sets, err) + } + }) + } +} + +func TestConfigLookupNetworkSettings(t *testing.T) { + for _, tt := range []struct { + key, value string + valid bool + }{ + {"RACER_HOST_NETWORK", "true", true}, + {"RACER_HOST_NETWORK", "false", true}, + {"RACER_HOST_NETWORK", "TRUE", false}, + {"RACER_HOST_NETWORK", "", false}, + {"RACER_DIAGNOSTICS_PORT", "1024", true}, + {"RACER_DIAGNOSTICS_PORT", "65535", true}, + {"RACER_DIAGNOSTICS_PORT", "1023", false}, + {"RACER_DIAGNOSTICS_PORT", "65536", false}, + {"RACER_DIAGNOSTICS_PORT", "", false}, + {"RACER_POD_NETWORK_NODES", "[]", true}, + {"RACER_POD_NETWORK_NODES", `["node-b","node-a"]`, true}, + {"RACER_POD_NETWORK_NODES", "null", false}, + {"RACER_POD_NETWORK_NODES", "{}", false}, + {"RACER_POD_NETWORK_NODES", "[1]", false}, + {"RACER_POD_NETWORK_NODES", "", false}, + } { + t.Run(tt.key+"="+tt.value, func(t *testing.T) { + values := map[string]string{"RACER_CLUSTER_ID": nodeID, "RACER_CONTROL_URL": "https://controller/", "RACER_DATAPLANE_IMAGE": "racer:test", "RACER_HOST_NETWORK": "true"} + values[tt.key] = tt.value + + cfg, err := ConfigFromLookup(func(key string) (string, bool) { value, ok := values[key]; return value, ok }) + if tt.valid { + if err != nil { + t.Fatal(err) + } + + if _, err := DesiredDaemonSets(cfg); err != nil { + t.Fatal(err) + } + } else if !errors.Is(err, wire.InvalidRequest) || !reflect.DeepEqual(cfg, Config{}) { + t.Fatalf("parse failure = %+v, %v", cfg, err) + } + }) + } +} + +func TestDesiredDaemonSets(t *testing.T) { + for _, hostNetwork := range []bool{false, true} { + cfg := workloadConfig(t) + cfg.HostNetwork = hostNetwork + + one, err := DesiredDaemonSet(cfg) + if err != nil { + t.Fatal(err) + } + + sets, err := DesiredDaemonSets(cfg) + if err != nil || !reflect.DeepEqual(sets, []*appsv1.DaemonSet{one}) { + t.Fatalf("single-workload plan differs: %+v, %v", sets, err) + } + } + + cfg := workloadConfig(t) + cfg.HostNetwork = true + + cfg.PodNetworkNodes = []string{"node-b", "node-a"} + if ds, err := DesiredDaemonSet(cfg); !errors.Is(err, wire.InvalidRequest) || ds != nil { + t.Fatalf("single builder accepted mixed mode: %+v, %v", ds, err) + } + + sets, err := DesiredDaemonSets(cfg) + if err != nil || len(sets) != 2 { + t.Fatalf("mixed plan: %+v, %v", sets, err) + } + + if !slices.Equal(cfg.PodNetworkNodes, []string{"node-b", "node-a"}) { + t.Fatal("planner reordered caller nodes") + } + + assertMixedPlacement(t, sets) + + sets[1].Spec.Template.Spec.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms[0].MatchExpressions[1].Values[0] = "mutated" + if sets[1].Spec.Template.Spec.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms[1].MatchExpressions[1].Values[0] != "linux" { + t.Fatal("placement terms alias each other") + } +} + +func assertMixedPlacement(t *testing.T, sets []*appsv1.DaemonSet) { + t.Helper() + + for i, ds := range sets { + name := []string{DataplaneDaemonSetName, PodNetworkDaemonSetName}[i] + + pod := ds.Spec.Template.Spec + if ds.Name != name || pod.HostNetwork != (i == 0) { + t.Fatalf("workload identity or network: %+v", ds) + } + + wantDNS := []corev1.DNSPolicy{corev1.DNSClusterFirstWithHostNet, corev1.DNSClusterFirst}[i] + if pod.DNSPolicy != wantDNS { + t.Fatalf("DNS = %s, want %s", pod.DNSPolicy, wantDNS) + } + + labels := map[string]string{"app.kubernetes.io/name": name, "app.kubernetes.io/instance": name} + if !reflect.DeepEqual(ds.Labels, labels) || !reflect.DeepEqual(ds.Spec.Selector.MatchLabels, labels) || !reflect.DeepEqual(ds.Spec.Template.Labels, labels) { + t.Fatalf("overlapping selectors: %+v", ds.Spec.Selector) + } + + terms := pod.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms + assertPlacementTerms(t, terms, i == 0) + } +} + +func assertPlacementTerms(t *testing.T, terms []corev1.NodeSelectorTerm, host bool) { + t.Helper() + + base := []corev1.NodeSelectorRequirement{ + {Key: wire.ExclusionLabel, Operator: corev1.NodeSelectorOpDoesNotExist}, + {Key: "kubernetes.io/os", Operator: corev1.NodeSelectorOpIn, Values: []string{"linux"}}, + } + + want := []corev1.NodeSelectorTerm{} + if host { + want = append(want, corev1.NodeSelectorTerm{MatchExpressions: base, MatchFields: []corev1.NodeSelectorRequirement{ + {Key: "metadata.name", Operator: corev1.NodeSelectorOpNotIn, Values: []string{"node-a"}}, + {Key: "metadata.name", Operator: corev1.NodeSelectorOpNotIn, Values: []string{"node-b"}}, + }}) + } else { + for _, node := range []string{"node-a", "node-b"} { + want = append(want, corev1.NodeSelectorTerm{MatchExpressions: base, MatchFields: []corev1.NodeSelectorRequirement{{Key: "metadata.name", Operator: corev1.NodeSelectorOpIn, Values: []string{node}}}}) + } + } + + if !reflect.DeepEqual(terms, want) { + t.Fatalf("placement = %+v, want %+v", terms, want) + } +} diff --git a/internal/racer/members/workload_fixture_test.go b/internal/racer/members/workload_fixture_test.go new file mode 100644 index 000000000..4c15202a2 --- /dev/null +++ b/internal/racer/members/workload_fixture_test.go @@ -0,0 +1,21 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package members + +import "github.com/Azure/unbounded/internal/racer/testutil" + +// Keep external-workload security assertions without a production builder API. +type Config = testutil.Config + +const ( + DataplaneDaemonSetName = testutil.DataplaneDaemonSetName + PodNetworkDaemonSetName = testutil.PodNetworkDaemonSetName +) + +var ( + ConfigFromLookup = testutil.ConfigFromLookup + DesiredDaemonSet = testutil.DesiredDaemonSet + DesiredDaemonSets = testutil.DesiredDaemonSets + ManagedNames = testutil.ManagedNames +) diff --git a/internal/racer/racer.go b/internal/racer/racer.go new file mode 100644 index 000000000..1ab4b734e --- /dev/null +++ b/internal/racer/racer.go @@ -0,0 +1,1048 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// Package racer composes the Racer controllers and replica observer. +// Constructors compose only; serving requires locally validated replicated state. +// Authority owns accepted state, server owns HTTPS serving, members owns workload +// discovery and builders, and wire owns protocol encoding. +package racer + +import ( + "context" + "crypto/tls" + "crypto/x509" + "encoding/json" + "fmt" + "math/rand/v2" + "net" + "net/http" + "os" + "reflect" + "strconv" + "strings" + "sync" + "time" + + appsv1 "k8s.io/api/apps/v1" + coordv1 "k8s.io/api/coordination/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/fields" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "k8s.io/apimachinery/pkg/util/validation" + "k8s.io/client-go/kubernetes" + clientgoscheme "k8s.io/client-go/kubernetes/scheme" + "k8s.io/client-go/tools/leaderelection/resourcelock" + "k8s.io/client-go/util/workqueue" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/builder" + "sigs.k8s.io/controller-runtime/pkg/cache" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/controller" + "sigs.k8s.io/controller-runtime/pkg/event" + "sigs.k8s.io/controller-runtime/pkg/handler" + "sigs.k8s.io/controller-runtime/pkg/healthz" + metricsserver "sigs.k8s.io/controller-runtime/pkg/metrics/server" + "sigs.k8s.io/controller-runtime/pkg/predicate" + "sigs.k8s.io/controller-runtime/pkg/reconcile" + "sigs.k8s.io/controller-runtime/pkg/source" + + machinav1 "github.com/Azure/unbounded/api/machina/v1alpha3" + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/server" + "github.com/Azure/unbounded/internal/racer/wire" +) + +type Application struct { + authority *authority.Authority + Topology *TopologyReconciler + Keyring *KeyringReconciler + Server *server.Server + Lifecycle *server.Lifecycle + Replication *Replication +} + +// Assemble performs no Kubernetes calls, I/O, cryptography, or goroutine startup. +// c supplies indexed discovery (configured by SetupWithManager in production); +// reader must bypass the cache for authorization and durable-state validation. +func Assemble(cfg Config, c client.Client, reader client.Reader) *Application { + cfg = cfg.effective() + a := authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: c, Reader: reader}) + lifecycle := server.NewLifecycle(a) + replication := &Replication{config: cfg, Client: c, APIReader: reader, authority: a} + + return &Application{ + authority: a, + Topology: &TopologyReconciler{Client: c, APIReader: reader, config: cfg, authority: a}, + Keyring: &KeyringReconciler{config: cfg, authority: a}, + Server: server.New(cfg.serverConfig(), c, a, lifecycle, replication), + Lifecycle: lifecycle, + Replication: replication, + } +} + +func (a *Application) SetupWithManager(mgr ctrl.Manager) error { + a.Lifecycle.SetCacheSync(mgr.GetCache().WaitForCacheSync) + + if err := mgr.Add(a.Lifecycle); err != nil { + return fmt.Errorf("register serving lifecycle: %w", err) + } + + if err := mgr.Add(a.Replication); err != nil { + return fmt.Errorf("register replica observer: %w", err) + } + + if err := mgr.Add(&publisherLifetime{replication: a.Replication}); err != nil { + return fmt.Errorf("register publisher lifetime: %w", err) + } + + if err := a.Topology.SetupWithManager(mgr); err != nil { + return fmt.Errorf("register topology: %w", err) + } + + if err := a.Keyring.SetupWithManager(mgr); err != nil { + return fmt.Errorf("register keyring: %w", err) + } + + if err := mgr.Add(a.Server); err != nil { + return fmt.Errorf("register HTTPS server: %w", err) + } + + if err := mgr.AddHealthzCheck("healthz", healthz.Ping); err != nil { + return fmt.Errorf("register liveness: %w", err) + } + + return mgr.AddReadyzCheck("racer", a.Server.Ready) +} + +// Run uses standard manager leader election. On leadership loss the process +// exits; Lease release-on-cancel stays disabled to avoid overlapping writers. +func Run(ctx context.Context, cfg Config) error { + if err := cfg.Validate(); err != nil { + return err + } + + scheme := runtime.NewScheme() + if err := clientgoscheme.AddToScheme(scheme); err != nil { + return err + } + + if err := racerv1.AddToScheme(scheme); err != nil { + return err + } + + // controller-runtime disables default client-side QPS throttling here; + // serving admission bounds concurrency and API priority/fairness governs load. + restConfig, err := ctrl.GetConfig() + if err != nil { + return err + } + + if err := cfg.validateReplication(); err != nil { + return err + } + + kube, err := kubernetes.NewForConfig(restConfig) + if err != nil { + return err + } + + options := managerOptions(cfg, scheme) + options.LeaderElectionResourceLockInterface = &resourcelock.LeaseLock{ + LeaseMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: "racer-controller"}, + Client: kube.CoordinationV1(), + LockConfig: resourcelock.ResourceLockConfig{Identity: cfg.PodName + "/" + cfg.PodUID}, + } + + mgr, err := ctrl.NewManager(restConfig, options) + if err != nil { + return err + } + + app := Assemble(cfg, mgr.GetClient(), mgr.GetAPIReader()) + // Recovery runs before election so every replica validates durable state. + // First-install initialization uses resource-version CAS and is safe when + // replicas race. Only the elected leader publishes subsequent changes. + if err := app.Recover(ctx, mgr.GetClient()); err != nil { + return err + } + + if err := app.SetupWithManager(mgr); err != nil { + return err + } + + return mgr.Start(ctx) +} + +// Recover runs the production startup guard before any manager runnable starts. +// The writer need not have a running cache; all reads use the authoritative reader +// supplied to Assemble. Constructors and recovery never grant serving authority. +func (a *Application) Recover(ctx context.Context, writer client.Writer) error { + return a.authority.Recover(ctx, writer) +} + +func managerOptions(cfg Config, scheme *runtime.Scheme) ctrl.Options { + return ctrl.Options{ + Scheme: scheme, + LeaderElection: true, + LeaderElectionReleaseOnCancel: false, + Metrics: metricsserver.Options{BindAddress: cfg.MetricsAddress}, + HealthProbeBindAddress: cfg.ProbeAddress, + Cache: cache.Options{ByObject: map[client.Object]cache.ByObject{ + &corev1.Pod{}: {Namespaces: map[string]cache.Config{cfg.Namespace: {}}}, + // Named Secret RBAC requires this selector on both LIST and WATCH. + &corev1.Secret{}: {Namespaces: map[string]cache.Config{cfg.Namespace: {}}, Field: fields.OneTermEqualSelector("metadata.name", cfg.CredentialsSecretName)}, + &corev1.ConfigMap{}: {Namespaces: map[string]cache.Config{cfg.Namespace: {}}, Field: fields.OneTermEqualSelector("metadata.name", cfg.VersionConfigMapName)}, + &appsv1.DaemonSet{}: {Namespaces: map[string]cache.Config{cfg.Namespace: {}}}, + }}, + } +} + +// Config contains controller runtime settings, not operator workload inputs. +type Config struct { + Cluster wire.ClusterID + Namespace string + ControlAddress string + MetricsAddress string + ProbeAddress string + TLSCertificateFile string + TLSPrivateKeyFile string + PeerPort uint16 + DataplaneServiceAccount string + DaemonSetName string + CredentialsSecretName string + VersionConfigMapName string + InstallationConfigMapName string + Limits server.Limits + Rotation authority.RotationPolicy + CertificateLifetime time.Duration + PodName string + PodUID string + ControllerServiceAccount string + ReplicationTokenFile string + ReplicationTrustFile string + ReplicationServerName string + ReplicationPort uint16 + SnapshotMaxAge time.Duration +} + +// LoadConfig reads deployment configuration. Initialization state is deliberately +// not an environment setting: it is read authoritatively on every recovery. +func LoadConfig() (Config, error) { + return ConfigFromLookup(os.LookupEnv) +} + +// ConfigFromLookup parses the controller configuration without process-global state. +func ConfigFromLookup(lookup func(string) (string, bool)) (Config, error) { + env := func(key, fallback string) string { + if value, ok := lookup(key); ok { + return value + } + + return fallback + } + + port, err := strconv.ParseUint(env("RACER_PEER_PORT", "8082"), 10, 16) + if err != nil { + return Config{}, fmt.Errorf("RACER_PEER_PORT must be an integer from 1 to 65535: %w", err) + } + + cfg := Config{ + Cluster: wire.ClusterID(env("RACER_CLUSTER_ID", "")), + Namespace: env("POD_NAMESPACE", "unbounded-system"), + ControlAddress: env("RACER_CONTROL_ADDRESS", ":8443"), + MetricsAddress: env("RACER_METRICS_ADDRESS", ":8080"), + ProbeAddress: env("RACER_PROBE_ADDRESS", ":8081"), + TLSCertificateFile: env("RACER_TLS_CERTIFICATE_FILE", "/etc/racer/tls/tls.crt"), + TLSPrivateKeyFile: env("RACER_TLS_PRIVATE_KEY_FILE", "/etc/racer/tls/tls.key"), + PeerPort: uint16(port), + DataplaneServiceAccount: env("RACER_DATAPLANE_SERVICE_ACCOUNT", "racer-dataplane"), + DaemonSetName: env("RACER_DAEMONSET_NAME", "racer-dataplane"), + CredentialsSecretName: env("RACER_CREDENTIALS_SECRET_NAME", "racer-credentials"), + VersionConfigMapName: env("RACER_VERSION_CONFIGMAP_NAME", "racer-version"), + InstallationConfigMapName: env("RACER_INSTALLATION_CONFIGMAP_NAME", "racer-installation"), + PodName: env("POD_NAME", ""), + PodUID: env("POD_UID", ""), + ControllerServiceAccount: env("RACER_CONTROLLER_SERVICE_ACCOUNT", "racer-controller"), + ReplicationTokenFile: env("RACER_REPLICATION_TOKEN_FILE", "/var/run/secrets/racer-controller/token"), + ReplicationTrustFile: env("RACER_REPLICATION_TRUST_FILE", "/etc/racer/tls/ca.crt"), + ReplicationServerName: env("RACER_REPLICATION_SERVER_NAME", "racer-controller."+env("POD_NAMESPACE", "unbounded-system")+".svc"), + SnapshotMaxAge: 30 * time.Second, + Limits: server.Limits{ + MaxConnections: 2*wire.MaxMembers + 128, + MaxConcurrentHandshakes: 32, + MaxPolls: wire.MaxMembers, + MaxConcurrentWrites: 128, + MaxConcurrentBootstrap: 32, + HeaderBytes: 16 * 1024, + WriteTimeout: 30 * time.Second, + HandshakeTimeout: 5 * time.Second, + ShutdownTimeout: 10 * time.Second, + }, + Rotation: authority.RotationPolicy{ + Interval: 24 * time.Hour, + PrepareFor: time.Hour, + RetainFor: 48 * time.Hour, + }, + } + + replicationPort, err := strconv.ParseUint(env("RACER_REPLICATION_PORT", "8443"), 10, 16) + if err != nil || replicationPort == 0 { + return Config{}, fmt.Errorf("RACER_REPLICATION_PORT must be an integer from 1 to 65535") + } + + cfg.ReplicationPort = uint16(replicationPort) + + for _, setting := range []struct { + name string + value *time.Duration + }{ + {"RACER_CERTIFICATE_LIFETIME", &cfg.CertificateLifetime}, + {"RACER_SNAPSHOT_MAX_AGE", &cfg.SnapshotMaxAge}, + {"RACER_HANDSHAKE_TIMEOUT", &cfg.Limits.HandshakeTimeout}, + {"RACER_ROTATION_INTERVAL", &cfg.Rotation.Interval}, + {"RACER_ROTATION_PREPARE_FOR", &cfg.Rotation.PrepareFor}, + {"RACER_ROTATION_RETAIN_FOR", &cfg.Rotation.RetainFor}, + } { + if value, ok := lookup(setting.name); ok { + duration, err := time.ParseDuration(value) + if err != nil || duration <= 0 || duration%time.Second != 0 { + return Config{}, fmt.Errorf("%s must be a positive duration in whole seconds", setting.name) + } + + *setting.value = duration + } + } + + cfg = cfg.effective() + + return cfg, cfg.Validate() +} + +// effective resolves optional lifetimes without I/O or validation side effects. +// Zero retains the defaults accepted by programmatic callers. +func (c Config) effective() Config { + if c.CertificateLifetime == 0 { + c.CertificateLifetime = wire.CertificateLifetime + } + + if c.SnapshotMaxAge == 0 { + c.SnapshotMaxAge = 30 * time.Second + } + + return c +} + +func (c Config) Validate() error { + c = c.effective() + if c.SnapshotMaxAge < 0 || c.SnapshotMaxAge > 0 && c.SnapshotMaxAge < time.Second { + return fmt.Errorf("SnapshotMaxAge must be at least one second") + } + + if !wire.ValidUUID(string(c.Cluster)) { + return fmt.Errorf("cluster must be a UUID") + } + + if len(validation.IsDNS1123Label(c.Namespace)) != 0 { + return fmt.Errorf("namespace must be a DNS label") + } + + if c.PeerPort == 0 { + return fmt.Errorf("PeerPort must be from 1 to 65535") + } + + for field, name := range map[string]string{"VersionConfigMapName": c.VersionConfigMapName, "InstallationConfigMapName": c.InstallationConfigMapName, "DaemonSetName": c.DaemonSetName, "CredentialsSecretName": c.CredentialsSecretName, "DataplaneServiceAccount": c.DataplaneServiceAccount} { + if len(validation.IsDNS1123Subdomain(name)) != 0 { + return fmt.Errorf("%s must be a DNS subdomain", field) + } + } + + if c.VersionConfigMapName == c.InstallationConfigMapName { + return fmt.Errorf("VersionConfigMapName and InstallationConfigMapName must differ") + } + + if err := c.serverConfig().Validate(); err != nil { + return fmt.Errorf("server configuration: %v", err) + } + + return c.validateRotation() +} + +func (c Config) validateRotation() error { + // Two minutes leaves a full poll turn between renewal at two-thirds of the + // lifetime and expiry. X.509 and rotation deadlines have second precision. + lifetime := c.CertificateLifetime + if lifetime < 2*time.Minute || lifetime > wire.CertificateLifetime || lifetime%time.Second != 0 { + return fmt.Errorf("CertificateLifetime must be from 2m to %s in whole seconds", wire.CertificateLifetime) + } + + if c.Rotation.PrepareFor <= 0 || c.Rotation.Interval < c.Rotation.PrepareFor || c.Rotation.RetainFor < lifetime || + c.Rotation.Interval > 365*24*time.Hour || c.Rotation.RetainFor > 365*24*time.Hour { + return fmt.Errorf("rotation requires positive PrepareFor <= Interval <= 8760h and CertificateLifetime <= RetainFor <= 8760h") + } + + return nil +} + +func (c Config) validateReplication() error { + if len(validation.IsDNS1123Subdomain(c.PodName)) != 0 || c.PodUID == "" || + len(validation.IsDNS1123Subdomain(c.ControllerServiceAccount)) != 0 || + len(validation.IsDNS1123Subdomain(c.ReplicationServerName)) != 0 || c.ReplicationPort == 0 || + c.ReplicationTokenFile == "" || c.ReplicationTrustFile == "" { + return fmt.Errorf("replication requires valid POD_NAME, POD_UID, controller identity, TLS trust, token, server name, and nonzero port") + } + + return nil +} + +const podNodeIndex = "spec.nodeName" + +// Controllers start sources only under leadership and wait for cache sync before +// workers run. Queueing directly guarantees a first reconcile even for empty lists. +func initialEnqueue() source.Source { + return source.Func(func(ctx context.Context, q workqueue.TypedRateLimitingInterface[reconcile.Request]) error { + if err := ctx.Err(); err != nil { + return err + } + + q.Add(singleton(ctx, nil)[0]) + + return nil + }) +} + +func podNodeKeys(obj client.Object) []string { + pod, ok := obj.(*corev1.Pod) + if !ok || pod.Spec.NodeName == "" { + return nil + } + + return []string{pod.Spec.NodeName} +} + +func changes(relevant func(client.Object) bool, equal func(client.Object, client.Object) bool) predicate.Predicate { + return predicate.Funcs{ + CreateFunc: func(e event.CreateEvent) bool { return relevant(e.Object) }, + DeleteFunc: func(e event.DeleteEvent) bool { return relevant(e.Object) }, + GenericFunc: func(e event.GenericEvent) bool { return relevant(e.Object) }, + UpdateFunc: func(e event.UpdateEvent) bool { + return (relevant(e.ObjectOld) || relevant(e.ObjectNew)) && !equal(e.ObjectOld, e.ObjectNew) + }, + } +} + +func nodeChanges() predicate.Predicate { + return changes(func(client.Object) bool { return true }, func(a, b client.Object) bool { + _, excludedA := a.GetLabels()[wire.ExclusionLabel] + + _, excludedB := b.GetLabels()[wire.ExclusionLabel] + if a.GetUID() != b.GetUID() || a.GetName() != b.GetName() || excludedA != excludedB { + return false + } + + if a.GetLabels()[machinav1.MachineSiteLabelKey] != b.GetLabels()[machinav1.MachineSiteLabelKey] { + return false + } + + for _, key := range []string{wire.SharesAnnotation, wire.RDMANICsAnnotation, enrolledRDMANICsAnnotation, enrolledSharesAnnotation} { + av, ap := a.GetAnnotations()[key] + + bv, bp := b.GetAnnotations()[key] + if ap != bp || av != bv { + return false + } + } + + return true + }) +} + +func managedPodChanges(cfg Config) predicate.Predicate { + return changes(func(obj client.Object) bool { + if obj.GetNamespace() != cfg.Namespace { + return false + } + + owner := metav1.GetControllerOf(obj) + + return owner != nil && owner.APIVersion == "apps/v1" && owner.Kind == "DaemonSet" && owner.Name == cfg.DaemonSetName + }, func(a, b client.Object) bool { + x, xok := a.(*corev1.Pod) + + y, yok := b.(*corev1.Pod) + if !xok || !yok { + return false + } + + return x.UID == y.UID && x.Spec.NodeName == y.Spec.NodeName && x.Status.PodIP == y.Status.PodIP && x.Status.Phase == y.Status.Phase && x.CreationTimestamp.Equal(&y.CreationTimestamp) && reflect.DeepEqual(x.DeletionTimestamp, y.DeletionTimestamp) && reflect.DeepEqual(x.OwnerReferences, y.OwnerReferences) + }) +} + +func cacheChanges() predicate.Predicate { + return changes(func(client.Object) bool { return true }, func(a, b client.Object) bool { + x, xok := a.(*racerv1.ClusterCache) + + y, yok := b.(*racerv1.ClusterCache) + if !xok || !yok { + return false + } + + return x.UID == y.UID && x.Name == y.Name + }) +} + +func namedObjects(namespace string, names ...string) func(client.Object) bool { + return func(obj client.Object) bool { + if obj.GetNamespace() != namespace { + return false + } + + for _, name := range names { + if obj.GetName() == name { + return true + } + } + + return false + } +} + +func namedChanges(namespace string, names ...string) predicate.Predicate { + return changes(namedObjects(namespace, names...), func(a, b client.Object) bool { return a.GetResourceVersion() == b.GetResourceVersion() }) +} + +// Ignore our own resource-version-only CAS writes to avoid an endless hot loop. +func versionChanges(cfg Config) predicate.Predicate { + return changes(namedObjects(cfg.Namespace, cfg.VersionConfigMapName, cfg.InstallationConfigMapName), func(a, b client.Object) bool { + x, xok := a.(*corev1.ConfigMap) + + y, yok := b.(*corev1.ConfigMap) + if !xok || !yok { + return false + } + + return x.UID == y.UID && reflect.DeepEqual(x.Data, y.Data) && reflect.DeepEqual(x.Annotations, y.Annotations) && reflect.DeepEqual(x.Immutable, y.Immutable) && reflect.DeepEqual(x.DeletionTimestamp, y.DeletionTimestamp) + }) +} + +func (c Config) serverConfig() server.Config { + return server.Config{ + ControlAddress: c.ControlAddress, + TLSCertificateFile: c.TLSCertificateFile, + TLSPrivateKeyFile: c.TLSPrivateKeyFile, + ReplicationServerName: c.ReplicationServerName, + Limits: c.Limits, + } +} + +func (c Config) authorityConfig() authority.Config { + return authority.Config{ + Cluster: c.Cluster, Namespace: c.Namespace, + DataplaneServiceAccount: c.DataplaneServiceAccount, ControllerServiceAccount: c.ControllerServiceAccount, + DaemonSetName: c.DaemonSetName, CredentialsSecretName: c.CredentialsSecretName, + VersionConfigMapName: c.VersionConfigMapName, InstallationConfigMapName: c.InstallationConfigMapName, + Rotation: c.Rotation, CertificateLifetime: c.CertificateLifetime, SnapshotMaxAge: c.SnapshotMaxAge, + MaxTokenBytes: c.Limits.HeaderBytes, + } +} + +type TopologyReconciler struct { + authority *authority.Authority + client.Client + APIReader client.Reader + config Config + hints map[string]nodeHint + enqueueHint func(reconcile.Request) +} + +type nodeHint struct { + node *corev1.Node + member string +} + +// Reconcile builds from the synchronized cache, reads the version ConfigMap +// authoritatively, commits counters/hashes with CAS, then installs the result. +// Conflicts requeue from fresh inputs; missing established counters fail closed. +func (r *TopologyReconciler) Reconcile(ctx context.Context, request ctrl.Request) (ctrl.Result, error) { + if err := ctx.Err(); err != nil { + return ctrl.Result{}, reconcile.TerminalError(err) + } + + if err := r.config.Validate(); err != nil { + return ctrl.Result{}, reconcile.TerminalError(err) + } + + var err error + if request.Namespace == "hints" { + err = r.reconcileHint(ctx, request.Name) + } else { + err = r.reconcile(ctx) + } + // Cancellation takes precedence even if a transport concurrently reports Conflict. + if ctx.Err() != nil { + return ctrl.Result{}, reconcile.TerminalError(ctx.Err()) + } + + return ctrl.Result{}, err +} + +func (r *TopologyReconciler) reconcile(ctx context.Context) error { + update, err := r.authority.PublishTopology(ctx, r.observeTopology) + if err != nil { + return err + } + + // Hint failures have their own keyed retries and never republish topology. + return r.queueHints(ctx, update) +} + +func (r *TopologyReconciler) observeTopology(ctx context.Context) (authority.TopologyObservation, error) { + cfg := r.config + + var nodes corev1.NodeList + if err := r.List(ctx, &nodes); err != nil { + return authority.TopologyObservation{}, err + } + + var caches racerv1.ClusterCacheList + if err := r.APIReader.List(ctx, &caches); err != nil { + return authority.TopologyObservation{}, err + } + + catalog, err := members.BuildCatalog(caches.Items) + if err != nil { + return authority.TopologyObservation{}, err + } + + ownership, err := members.ReadWorkloadIdentities(ctx, r.APIReader, cfg.Namespace, cfg.DaemonSetName) + if err != nil { + return authority.TopologyObservation{}, err + } + // Indexed namespace-scoped queries avoid scanning unrelated Pods for each + // Node. Ownership is still verified against the current DaemonSet UID. + podsByNode := make(map[string][]corev1.Pod, len(nodes.Items)) + + for _, node := range nodes.Items { + if err := ctx.Err(); err != nil { + return authority.TopologyObservation{}, err + } + + var list corev1.PodList + if err := r.List(ctx, &list, client.InNamespace(cfg.Namespace), client.MatchingFields{podNodeIndex: node.Name}); err != nil { + return authority.TopologyObservation{}, err + } + + podsByNode[node.Name] = list.Items + } + + return authority.TopologyObservation{Nodes: nodes, Catalog: catalog, Input: members.Input{ + Nodes: nodes.Items, PodsByNode: podsByNode, Ownership: ownership, PeerPort: cfg.PeerPort, + }}, nil +} + +func (r *TopologyReconciler) queueHints(ctx context.Context, update authority.TopologyHints) error { + hints := make(map[string]nodeHint) + + for i := range update.Nodes.Items { + if err := ctx.Err(); err != nil { + return err + } + + node := &update.Nodes.Items[i] + hint := nodeHint{} + + member, ok := update.Members[wire.NodeID(node.UID)] + if !ok { + if _, excluded := node.Labels[wire.ExclusionLabel]; !excluded { + continue + } + } else { + encoded, err := json.Marshal(member) + if err != nil { + return err + } + + if len(encoded) > 64*1024 { + return wire.TooLarge + } + + hint.member = string(encoded) + } + + // The cache can avoid work, but never authorize a write. Queued hints + // still require a fresh read and UID/input checks in reconcileHint. + if node.Annotations[admittedMemberAnnotation] == hint.member { + continue + } + + hint.node = node.DeepCopy() + hints[node.Name] = hint + } + + previous := r.hints + + r.hints = hints + for name := range hints { + // Pending keys already have a queue entry or a rate-limited retry. + // Replace their desired state without bypassing the retry delay. + if _, pending := previous[name]; !pending && r.enqueueHint != nil { + r.enqueueHint(reconcile.Request{NamespacedName: types.NamespacedName{Namespace: "hints", Name: name}}) + } + } + + return nil +} + +func (r *TopologyReconciler) reconcileHint(ctx context.Context, name string) error { + hint, ok := r.hints[name] + if !ok { + return nil + } + + var current corev1.Node + if err := r.APIReader.Get(ctx, client.ObjectKey{Name: name}, ¤t); err != nil { + if apierrors.IsNotFound(err) { + delete(r.hints, name) + return nil + } + + return err + } + + if current.UID != hint.node.UID { + delete(r.hints, name) + return nil + } + + _, excluded := current.Labels[wire.ExclusionLabel] + if excluded { + hint.member = "" + } else if nodeChanges().Update(event.UpdateEvent{ObjectOld: hint.node, ObjectNew: ¤t}) { + // The topology watch will publish the new inputs before supplying a hint. + delete(r.hints, name) + return nil + } + + if current.Annotations[admittedMemberAnnotation] != hint.member { + before := current.DeepCopy() + if hint.member == "" { + delete(current.Annotations, admittedMemberAnnotation) + } else { + if current.Annotations == nil { + current.Annotations = map[string]string{} + } + + current.Annotations[admittedMemberAnnotation] = hint.member + } + + if err := r.Patch(ctx, ¤t, client.MergeFromWithOptions(before, client.MergeFromWithOptimisticLock{})); err != nil { + return err + } + } + + delete(r.hints, name) + + return nil +} + +func (r *TopologyReconciler) SetupWithManager(mgr ctrl.Manager) error { + cfg := r.config + + installation, err := installationSource(mgr, cfg) + if err != nil { + return err + } + + if err := mgr.GetFieldIndexer().IndexField(context.Background(), &corev1.Pod{}, podNodeIndex, podNodeKeys); err != nil { + return err + } + + return ctrl.NewControllerManagedBy(mgr). + Named("racer-topology"). + WatchesRawSource(installation). + WatchesRawSource(source.Func(func(_ context.Context, q workqueue.TypedRateLimitingInterface[reconcile.Request]) error { + r.enqueueHint = q.Add + return nil + })). + WatchesRawSource(initialEnqueue()). + Watches(&corev1.Node{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(nodeChanges())). + Watches(&corev1.Pod{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(managedPodChanges(cfg))). + Watches(&appsv1.DaemonSet{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(namedChanges(cfg.Namespace, cfg.DaemonSetName))). + Watches(&racerv1.ClusterCache{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(cacheChanges())). + Watches(&corev1.Secret{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(namedChanges(cfg.Namespace, cfg.CredentialsSecretName))). + Watches(&corev1.ConfigMap{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(versionChanges(cfg))). + WithOptions(controller.Options{MaxConcurrentReconciles: 1}). + Complete(r) +} + +// singleton coalesces input changes without introducing a singleton CR. +func singleton(_ context.Context, _ client.Object) []reconcile.Request { + return []reconcile.Request{{NamespacedName: types.NamespacedName{Name: "racer"}}} +} + +const ( + enrolledSharesAnnotation = members.EnrolledSharesAnnotation + enrolledRDMANICsAnnotation = members.EnrolledRDMANICsAnnotation + admittedMemberAnnotation = members.AdmittedMemberAnnotation +) + +type KeyringReconciler struct { + config Config + authority *authority.Authority +} + +func (r *KeyringReconciler) Reconcile(ctx context.Context, _ ctrl.Request) (ctrl.Result, error) { + delay, err := r.authority.ReconcileCredentials(ctx) + if ctx.Err() != nil { + return ctrl.Result{}, reconcile.TerminalError(ctx.Err()) + } + + if err != nil { + return ctrl.Result{}, err + } + + return ctrl.Result{RequeueAfter: delay}, nil +} + +func (r *KeyringReconciler) SetupWithManager(mgr ctrl.Manager) error { + cfg := r.config + + installation, err := installationSource(mgr, cfg) + if err != nil { + return err + } + + return ctrl.NewControllerManagedBy(mgr). + Named("racer-credentials"). + WatchesRawSource(installation). + WatchesRawSource(initialEnqueue()). + Watches(&racerv1.ClusterCache{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(cacheChanges())). + Watches(&corev1.Secret{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(namedChanges(cfg.Namespace, cfg.CredentialsSecretName))). + Watches(&corev1.ConfigMap{}, handler.EnqueueRequestsFromMapFunc(singleton), builder.WithPredicates(versionChanges(cfg))). + WithOptions(controller.Options{MaxConcurrentReconciles: 1}).Complete(r) +} + +// Kubernetes field selectors cannot OR two names. Each controller watches the +// installation marker through a separate named cache; the manager caches only +// the version object. Authorization always uses the uncached API reader. +func installationSource(mgr ctrl.Manager, cfg Config) (source.SyncingSource, error) { + c, err := cache.New(mgr.GetConfig(), cache.Options{ + Scheme: mgr.GetScheme(), Mapper: mgr.GetRESTMapper(), + DefaultNamespaces: map[string]cache.Config{cfg.Namespace: {}}, + DefaultFieldSelector: fields.OneTermEqualSelector("metadata.name", cfg.InstallationConfigMapName), + }) + if err != nil { + return nil, err + } + + if err := mgr.Add(c); err != nil { + return nil, err + } + + return source.Kind[client.Object](c, &corev1.ConfigMap{}, handler.EnqueueRequestsFromMapFunc(singleton), versionChanges(cfg)), nil +} + +const ReplicationAudience = authority.ReplicationAudience + +// Replication observes durable authority on every replica. No request from a +// dataplane performs these reads. Only the elected publisher supplies image bytes. +type Replication struct { + authority *authority.Authority + config Config + Client client.Client + APIReader client.Reader + mu sync.Mutex + leader context.Context +} + +type publisherLifetime struct{ replication *Replication } + +func (*publisherLifetime) NeedLeaderElection() bool { return true } +func (p *publisherLifetime) Start(ctx context.Context) error { + p.replication.mu.Lock() + p.replication.leader = ctx + p.replication.mu.Unlock() + <-ctx.Done() + + return nil +} + +func (r *Replication) isLeader() bool { + _, leader := r.LeaderContext() + return leader +} + +func (*Replication) NeedLeaderElection() bool { return false } + +func (r *Replication) interval() time.Duration { + return min(5*time.Second, r.config.SnapshotMaxAge/3) +} + +func (r *Replication) Start(ctx context.Context) error { + // Credential observations continue while the follower poll is blocked. + done := make(chan struct{}) + + go func() { + defer close(done) + + for ctx.Err() == nil { + observation, cancel := context.WithTimeout(ctx, r.interval()) + r.observe(observation) + cancel() + + if !replicationSleep(ctx, r.interval()) { + return + } + } + }() + + defer func() { <-done }() + + for ctx.Err() == nil { + if !r.isLeader() { + poll, cancel := context.WithTimeout(ctx, 2*r.interval()+r.config.Limits.WriteTimeout) + if err := r.poll(poll, ctx); err != nil && ctx.Err() == nil { + ctrl.LoggerFrom(ctx).V(1).Info("snapshot replication retry", "error", err) + } + + cancel() + } + + if !replicationSleep(ctx, r.interval()/2+time.Duration(rand.Int64N(int64(r.interval()/2)+1))) { + break + } + } + + return nil +} + +func replicationSleep(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func (r *Replication) observe(ctx context.Context) { + if err := r.authority.Observe(ctx); err != nil && ctx.Err() == nil { + ctrl.LoggerFrom(ctx).V(1).Info("authority observation retry", "error", err) + } +} + +func (r *Replication) leaderAddress(ctx context.Context) (string, error) { + cfg := r.config + + var lease coordv1.Lease + if err := r.APIReader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: "racer-controller"}, &lease); err != nil { + return "", err + } + + if lease.Spec.HolderIdentity == nil || lease.Spec.RenewTime == nil || lease.Spec.LeaseDurationSeconds == nil || *lease.Spec.LeaseDurationSeconds <= 0 || time.Since(lease.Spec.RenewTime.Time) >= time.Duration(*lease.Spec.LeaseDurationSeconds)*time.Second { + return "", wire.Unavailable + } + + name, uid, ok := strings.Cut(*lease.Spec.HolderIdentity, "/") + if !ok || name == "" || uid == "" { + return "", wire.Unavailable + } + + var pod corev1.Pod + if err := r.APIReader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: name}, &pod); err != nil { + return "", err + } + + if string(pod.UID) != uid || !members.ControllerPod(&pod, cfg.Namespace, cfg.ControllerServiceAccount) || net.ParseIP(pod.Status.PodIP) == nil { + return "", wire.Unavailable + } + + return net.JoinHostPort(pod.Status.PodIP, strconv.Itoa(int(cfg.ReplicationPort))), nil +} + +func (r *Replication) poll(ctx, process context.Context) error { + address, err := r.leaderAddress(ctx) + if err != nil { + return err + } + + pem, err := os.ReadFile(r.config.ReplicationTrustFile) + if err != nil { + return err + } + + roots := x509.NewCertPool() + if !roots.AppendCertsFromPEM(pem) { + return wire.Unavailable + } + + token, err := os.ReadFile(r.config.ReplicationTokenFile) + if err != nil { + return err + } + + transport := &http.Transport{TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS13, RootCAs: roots, ServerName: r.config.ReplicationServerName}, TLSHandshakeTimeout: r.interval(), DisableKeepAlives: true} + defer transport.CloseIdleConnections() + + httpClient := &http.Client{Transport: transport, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }} + + path := "https://" + address + server.ReplicationPath + if current, err := r.authority.Current(); err == nil { + path += "?after=" + strconv.FormatUint(uint64(current.Sequence()), 10) + } + + request, err := http.NewRequestWithContext(ctx, http.MethodGet, path, nil) + if err != nil { + return err + } + + request.Header.Set("Authorization", "Bearer "+strings.TrimSpace(string(token))) + + response, err := httpClient.Do(request) + if err != nil { + return err + } + + defer func() { + if err := response.Body.Close(); err != nil { + ctrl.LoggerFrom(ctx).V(1).Info("close replication response", "error", err) + } + }() + + if response.StatusCode == http.StatusNoContent { + return nil + } // Not a freshness confirmation. + + if response.StatusCode != http.StatusOK { + return fmt.Errorf("replication HTTP status %d", response.StatusCode) + } + + image, err := wire.DecodePublication(response.Body) + if err != nil { + return err + } + + return r.authority.AcceptReplica(ctx, process, image) +} + +// LeaderContext snapshots publisher lifetime and leadership under the same lock. +func (r *Replication) LeaderContext() (context.Context, bool) { + r.mu.Lock() + defer r.mu.Unlock() + + return r.leader, r.leader != nil && r.leader.Err() == nil +} + +func (r *Replication) PollInterval() time.Duration { return r.interval() } + +func (r *Replication) AuthenticateReplica(ctx context.Context, request *http.Request) (string, time.Time, error) { + identity, err := r.authority.AuthenticateReplica(ctx, request) + return identity.UID(), identity.Expires(), err +} diff --git a/internal/racer/racer_test.go b/internal/racer/racer_test.go new file mode 100644 index 000000000..e5ca78830 --- /dev/null +++ b/internal/racer/racer_test.go @@ -0,0 +1,1560 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package racer + +import ( + "bytes" + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/sha256" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/base64" + "encoding/hex" + "encoding/json" + "encoding/pem" + "errors" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/fields" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "k8s.io/apimachinery/pkg/util/validation" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + "sigs.k8s.io/controller-runtime/pkg/event" + "sigs.k8s.io/controller-runtime/pkg/reconcile" + + machinav1 "github.com/Azure/unbounded/api/machina/v1alpha3" + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/server" + "github.com/Azure/unbounded/internal/racer/testutil" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestFrozenConfigUsedByRuntimeOperations(t *testing.T) { + f := newServingFixture(t) + a := f.a + request := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + request.Header.Set("Authorization", "Bearer "+f.token) + + identity, err := a.authority.Authenticate(f.ctx, request) + if err != nil { + t.Fatal(err) + } + + a.Replication.observe(f.ctx) + handler := a.Server.Handler() + // Runtime components retain only private constructor inputs. + input := a.Topology.config + require.NotZero(t, input) + input = Config{} + require.Zero(t, input) + + runKeys(t, a.Keyring) + reconcileTopology(t, a.Topology, f.ctx) + a.Replication.observe(f.ctx) + + if _, err := a.authority.Authenticate(f.ctx, request); err != nil { + t.Fatal("bootstrap reread config", err) + } + + encoded, err := a.authority.Issue(f.ctx, identity, f.request) + if err != nil { + t.Fatal("issuer reread config", err) + } + + if response, err := wire.DecodeBootstrapResponse(bytes.NewReader(encoded)); err != nil || response.Cluster != f.request.Cluster { + t.Fatal("issued identity changed", err) + } + + request.TLS = f.requestState(t) + request.Header.Del("Authorization") + + w := httptest.NewRecorder() + handler.ServeHTTP(w, request) + + if w.Code != http.StatusOK { + t.Fatal("serving/replication reread config", w.Code) + } +} + +func TestComponentConfigFreezesDefaultsAtFirstUse(t *testing.T) { + for _, name := range []string{"topology", "keyring", "replication", "bootstrap", "issuer", "server"} { + t.Run(name, func(t *testing.T) { componentConfigFreezes(t, name) }) + } +} + +func componentConfigFreezes(t *testing.T, name string) { + t.Helper() + // Authority policy now belongs to root construction; the server only + // freezes transport settings, exercised through its public handler. + if name == "bootstrap" || name == "issuer" || name == "server" { + cfg := testConfig(t) + cfg.CertificateLifetime, cfg.SnapshotMaxAge = 0, 0 + + cfg.PeerPort = 9443 + if err := cfg.Validate(); err != nil { + t.Fatal(err) + } + + a := Assemble(cfg, nil, nil) + if a.Topology.config != cfg.effective() { + t.Fatal("construction ignored pre-use inputs/defaults") + } + + handler := a.Server.Handler() + cfg.Limits.HeaderBytes = 0 + + var readers sync.WaitGroup + for range 8 { + readers.Go(func() { + r := httptest.NewRequest(http.MethodPost, wire.BootstrapPath, nil) + r.Header.Set("X-Large", strings.Repeat("x", 16*1024)) + + w := httptest.NewRecorder() + handler.ServeHTTP(w, r) + + if w.Code != http.StatusRequestEntityTooLarge { + t.Error("runtime reread mutated construction inputs") + } + }) + } + + readers.Wait() + + return + } + + input := testConfig(t) + input.CertificateLifetime = 0 + input.SnapshotMaxAge = 0 + input.PeerPort = 9443 + + want := input.effective() + if err := input.Validate(); err != nil { + t.Fatal("zero optional lifetimes rejected", err) + } + + a := Assemble(input, nil, nil) + + get := map[string]func() Config{ + "topology": func() Config { return a.Topology.config }, + "keyring": func() Config { return a.Keyring.config }, + "replication": func() Config { return a.Replication.config }, + }[name] + if got := get(); got != want || got.CertificateLifetime != wire.CertificateLifetime || got.SnapshotMaxAge != 30*time.Second { + t.Fatal("construction ignored inputs/defaults") + } + + input = Config{} + require.Zero(t, input) + + var readers sync.WaitGroup + for range 8 { + readers.Go(func() { + if get() != want { + t.Error("runtime reread mutated construction inputs") + } + }) + } + + readers.Wait() +} + +func TestDirectComponentConfigDefaults(t *testing.T) { + for name, get := range map[string]func() Config{ + "topology": func() Config { return Assemble(Config{}, nil, nil).Topology.config }, + "keyring": func() Config { return Assemble(Config{}, nil, nil).Keyring.config }, + "bootstrap": func() Config { return Assemble(Config{}, nil, nil).Topology.config }, + "issuer": func() Config { return Assemble(Config{}, nil, nil).Keyring.config }, + "replication": func() Config { return Assemble(Config{}, nil, nil).Replication.config }, + } { + t.Run(name, func(t *testing.T) { + cfg := get() + if cfg.CertificateLifetime != wire.CertificateLifetime || cfg.SnapshotMaxAge != 30*time.Second { + t.Fatal("direct zero-default semantics lost") + } + }) + } +} + +func TestConfigDeploymentIdentityAndBounds(t *testing.T) { + cfg := testConfig(t) + for name, mutate := range map[string]func(*Config){ + "cluster": func(c *Config) { c.Cluster = "" }, + "namespace": func(c *Config) { c.Namespace = "../namespace" }, + "missing marker name": func(c *Config) { c.InstallationConfigMapName = "" }, + "aliased durable objects": func(c *Config) { c.InstallationConfigMapName = c.VersionConfigMapName }, + "aliased credential secrets": func(c *Config) { c.CredentialsSecretName = "" }, + "no preparation": func(c *Config) { c.Rotation.PrepareFor = 0 }, + "short overlap": func(c *Config) { c.Rotation.RetainFor = wire.CertificateLifetime - 1 }, + "short interval": func(c *Config) { c.Rotation.Interval = c.Rotation.PrepareFor - 1 }, + "zero port": func(c *Config) { c.PeerPort = 0 }, + "unbounded polls": func(c *Config) { c.Limits.MaxPolls = 0 }, + "unbounded writes": func(c *Config) { c.Limits.MaxConcurrentWrites = 0 }, + "unbounded bootstrap": func(c *Config) { c.Limits.MaxConcurrentBootstrap = 0 }, + "unbounded headers": func(c *Config) { c.Limits.HeaderBytes = 0 }, + "unbounded write duration": func(c *Config) { c.Limits.WriteTimeout = 0 }, + "unbounded shutdown": func(c *Config) { c.Limits.ShutdownTimeout = 0 }, + } { + t.Run(name, func(t *testing.T) { + invalid := cfg + mutate(&invalid) + + if err := invalid.Validate(); err == nil { + t.Fatalf("invalid config accepted: %v", err) + } + }) + } + + for _, port := range []string{"0", "65536", "-1", "invalid"} { + t.Setenv("RACER_PEER_PORT", port) + + if _, err := LoadConfig(); err == nil { + t.Fatalf("port %q: %v", port, err) + } + } + + t.Setenv("RACER_PEER_PORT", "65535") + t.Setenv("RACER_INSTALLATION_CONFIGMAP_NAME", "permanent-installation") + t.Setenv("RACER_CREDENTIALS_SECRET_NAME", "custom-credentials") + + loaded, err := LoadConfig() + if err != nil || loaded.PeerPort != 65535 || loaded.InstallationConfigMapName != "permanent-installation" || loaded.CredentialsSecretName != "custom-credentials" { + t.Fatalf("deployment overrides: %+v, %v", loaded, err) + } +} + +func TestRuntimeConfigDoesNotReadWorkloadOnlySettings(t *testing.T) { + _, err := ConfigFromLookup(func(key string) (string, bool) { + switch key { + case "RACER_CLUSTER_ID": + return "11111111-1111-1111-1111-111111111111", true + case "RACER_CONTROL_URL", "RACER_DATAPLANE_IMAGE", "RACER_BOOTSTRAP_TRUST_CONFIGMAP": + t.Errorf("runtime requested workload-only setting %s", key) + return "invalid", true + default: + return "", false + } + }) + if err != nil { + t.Fatal(err) + } +} + +func TestConfigShortRotationDurations(t *testing.T) { + testConfig(t) + + cfg, err := LoadConfig() + if err != nil || cfg.CertificateLifetime != wire.CertificateLifetime { + t.Fatalf("default lifetime: %v", err) + } + + for name, value := range map[string]string{ + "RACER_CERTIFICATE_LIFETIME": "2m", + "RACER_ROTATION_INTERVAL": "5m", + "RACER_ROTATION_PREPARE_FOR": "20s", + "RACER_ROTATION_RETAIN_FOR": "2m", + } { + t.Setenv(name, value) + } + + cfg, err = LoadConfig() + if err != nil || cfg.CertificateLifetime != 2*time.Minute || cfg.Rotation != (RotationPolicy{Interval: 5 * time.Minute, PrepareFor: 20 * time.Second, RetainFor: 2 * time.Minute}) { + t.Fatalf("short rotation config: %v", err) + } + + for name, value := range map[string]string{ + "RACER_CERTIFICATE_LIFETIME": "119s", + "RACER_ROTATION_INTERVAL": "19s", + "RACER_ROTATION_PREPARE_FOR": "500ms", + "RACER_ROTATION_RETAIN_FOR": "119s", + } { + t.Run(name, func(t *testing.T) { + for _, invalid := range []string{value, "", "nonsense", "0", "-1s", "8761h", "120.5s"} { + t.Setenv(name, invalid) + + if _, err := LoadConfig(); err == nil { + t.Fatalf("%s=%q accepted: %v", name, invalid, err) + } + } + }) + } +} + +func TestConfigDurationsUseProvidedLookup(t *testing.T) { + values := map[string]string{ + "RACER_CLUSTER_ID": string(testConfig(t).Cluster), + "RACER_CERTIFICATE_LIFETIME": "2m", + "RACER_ROTATION_INTERVAL": "5m", + "RACER_ROTATION_PREPARE_FOR": "1m", + "RACER_ROTATION_RETAIN_FOR": "2m", + } + for name := range values { + t.Setenv(name, "invalid-process-value") + } + + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + + cfg, err := ConfigFromLookup(lookup) + if err != nil || cfg.CertificateLifetime != 2*time.Minute || cfg.Rotation != (RotationPolicy{Interval: 5 * time.Minute, PrepareFor: time.Minute, RetainFor: 2 * time.Minute}) { + t.Fatalf("custom lookup ignored: %v", err) + } + + for _, name := range []string{"RACER_CERTIFICATE_LIFETIME", "RACER_ROTATION_INTERVAL", "RACER_ROTATION_PREPARE_FOR", "RACER_ROTATION_RETAIN_FOR"} { + previous := values[name] + + values[name] = "invalid-lookup-value" + if _, err := ConfigFromLookup(lookup); err == nil { + t.Fatalf("invalid custom %s accepted: %v", name, err) + } + + values[name] = previous + } + + delete(values, "RACER_CERTIFICATE_LIFETIME") + delete(values, "RACER_ROTATION_RETAIN_FOR") + + cfg, err = ConfigFromLookup(lookup) + if err != nil || cfg.CertificateLifetime != wire.CertificateLifetime || cfg.Rotation.RetainFor != 48*time.Hour { + t.Fatalf("absent custom values did not use defaults: %v", err) + } +} + +func TestReplicationConfigDefaultsAndOverrides(t *testing.T) { + values := map[string]string{"RACER_CLUSTER_ID": string(testConfig(t).Cluster), "POD_NAMESPACE": "controllers"} + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + + cfg, err := ConfigFromLookup(lookup) + if err != nil { + t.Fatal(err) + } + + require.Equal(t, 30*time.Second, cfg.SnapshotMaxAge) + require.EqualValues(t, 8443, cfg.ReplicationPort) + require.Equal(t, "racer-controller.controllers.svc", cfg.ReplicationServerName) + require.Equal(t, "/var/run/secrets/racer-controller/token", cfg.ReplicationTokenFile) + require.Equal(t, "/etc/racer/tls/ca.crt", cfg.ReplicationTrustFile) + require.Equal(t, "racer-controller", cfg.ControllerServiceAccount) + + values["RACER_SNAPSHOT_MAX_AGE"] = "45s" + values["RACER_REPLICATION_PORT"] = "9443" + values["POD_NAME"] = "controller-0" + values["POD_UID"] = "pod-uid" + + cfg, err = ConfigFromLookup(lookup) + if err != nil || cfg.SnapshotMaxAge != 45*time.Second || cfg.ReplicationPort != 9443 || cfg.PodName != "controller-0" || cfg.PodUID != "pod-uid" { + t.Fatal("replication overrides", err) + } + + for name, invalid := range map[string][]string{"RACER_REPLICATION_PORT": {"0", "65536", "-1", "bad"}, "RACER_SNAPSHOT_MAX_AGE": {"0s", "-1s", "500ms", "bad"}} { + previous := values[name] + for _, value := range invalid { + values[name] = value + if _, err := ConfigFromLookup(lookup); err == nil { + t.Fatalf("accepted %s=%s", name, value) + } + } + + values[name] = previous + } +} + +func TestCredentialsCacheSelector(t *testing.T) { + for _, name := range []string{"racer-credentials", "custom-credentials"} { + t.Run(name, func(t *testing.T) { + options := managerOptions(Config{Namespace: "custom-system", CredentialsSecretName: name}, runtime.NewScheme()) + for obj, config := range options.Cache.ByObject { + if _, ok := obj.(*corev1.Secret); !ok { + continue + } + + require.Len(t, config.Namespaces, 1) + require.Contains(t, config.Namespaces, "custom-system") + require.NotNil(t, config.Field) + require.Equal(t, "metadata.name="+name, config.Field.String()) + require.True(t, config.Field.Matches(fields.Set{"metadata.name": name})) + + for _, excluded := range []string{"racer-controller-tls", "unrelated", ""} { + require.False(t, config.Field.Matches(fields.Set{"metadata.name": excluded})) + } + + return + } + + t.Fatal("Secret cache configuration missing") + }) + } +} + +func (f *servingFixture) requestState(t *testing.T) *tls.ConnectionState { + t.Helper() + + certs := make([]*x509.Certificate, len(f.certificate.Certificate)) + for i, der := range f.certificate.Certificate { + var err error + + certs[i], err = x509.ParseCertificate(der) + if err != nil { + t.Fatal(err) + } + } + + return &tls.ConnectionState{HandshakeComplete: true, PeerCertificates: certs, VerifiedChains: [][]*x509.Certificate{certs}} +} + +func authFixture(t *testing.T) (*Application, authv1.TokenReviewStatus, string) { + t.Helper() + + controller := true + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "racer-dataplane", UID: "ds-uid"}} + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "racer-dataplane", UID: "sa-uid"}} + node := &corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "worker", UID: types.UID(testNodeUID)}} + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "worker-pod", UID: "pod-uid", OwnerReferences: []metav1.OwnerReference{{APIVersion: "apps/v1", Kind: "DaemonSet", Name: ds.Name, UID: ds.UID, Controller: &controller}}}, Spec: corev1.PodSpec{NodeName: node.Name, ServiceAccountName: sa.Name}, Status: corev1.PodStatus{PodIP: "192.0.2.1"}} + r := initializedTopology(t, ds, sa, node, pod) + a := assembleFixture(r.config, r.Client, r.APIReader) + status := authv1.TokenReviewStatus{Authenticated: true, Audiences: []string{wire.TokenAudience}, User: authv1.UserInfo{Username: "system:serviceaccount:racer:racer-dataplane", UID: string(sa.UID), Extra: map[string]authv1.ExtraValue{"authentication.kubernetes.io/pod-name": {pod.Name}, "authentication.kubernetes.io/pod-uid": {string(pod.UID)}, "authentication.kubernetes.io/node-name": {node.Name}, "authentication.kubernetes.io/node-uid": {string(node.UID)}}}} + token := "header." + base64.RawURLEncoding.EncodeToString(fmt.Appendf(nil, `{"exp":%d}`, time.Now().Add(time.Hour).Unix())) + ".signature" + + return a, status, token +} + +func installReview(t *testing.T, a *Application, status authv1.TokenReviewStatus, token string) { + t.Helper() + + fixtureDependencies[a.authority].Client = interceptor.NewClient(a.Topology.Client.(client.WithWatch), interceptor.Funcs{Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + review, ok := obj.(*authv1.TokenReview) + if !ok { + return c.Create(ctx, obj, opts...) + } + + if review.Spec.Token != token || len(review.Spec.Audiences) != 1 || review.Spec.Audiences[0] != wire.TokenAudience { + t.Error("TokenReview did not bind token/audience") + } + + review.Status = *status.DeepCopy() + + return ctx.Err() + }}) +} + +func TestWorkloadNameLabelBounds(t *testing.T) { + for _, name := range []string{"racer", "racer.custom", strings.Repeat("a", 63), strings.Repeat("a", 64), strings.Repeat("a", 63) + ".b", "", "Invalid"} { + t.Run(name, func(t *testing.T) { + cfg := workloadConfig(t) + cfg.DaemonSetName = name + valid := len(validation.IsDNS1123Subdomain(name)) == 0 && len(validation.IsValidLabelValue(name)) == 0 + + ds, err := testutil.DesiredDaemonSet(cfg) + if !valid { + if !errors.Is(err, wire.InvalidRequest) || ds != nil { + t.Fatalf("invalid name accepted: %v", err) + } + + return + } + + if err != nil || ds.Name != name || ds.Spec.Selector.MatchLabels["app.kubernetes.io/instance"] != name || ds.Spec.Template.Labels["app.kubernetes.io/instance"] != name { + t.Fatalf("valid name not preserved: %v", err) + } + }) + } + + // Other resource names are not instance labels and retain DNS subdomain bounds. + cfg := workloadConfig(t) + cfg.BootstrapTrustConfigMap = strings.Repeat("a", 63) + ".trust" + + cfg.DataplaneServiceAccount = strings.Repeat("a", 63) + ".account" + if _, err := testutil.DesiredDaemonSet(cfg); err != nil { + t.Fatal(err) + } +} + +func TestRDMANICAnnotationWatches(t *testing.T) { + for _, field := range []string{wire.RDMANICsAnnotation, enrolledRDMANICsAnnotation} { + t.Run(field, func(t *testing.T) { + before := memberNode() + after := before.DeepCopy() + after.Annotations = map[string]string{field: "[]"} + require.True(t, nodeChanges().Update(event.UpdateEvent{ObjectOld: &before, ObjectNew: after})) + require.True(t, nodeChanges().Update(event.UpdateEvent{ObjectOld: after, ObjectNew: &before})) + require.False(t, nodeChanges().Update(event.UpdateEvent{ObjectOld: after, ObjectNew: after.DeepCopy()})) + }) + } +} + +func TestMixedControllerBootstrapBindings(t *testing.T) { + a, status, token := authFixture(t) + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: PodNetworkDaemonSetName, UID: "podnet"}} + require.NoError(t, a.Topology.Create(t.Context(), ds)) + + pod := &corev1.Pod{} + require.NoError(t, a.Topology.Get(t.Context(), client.ObjectKey{Namespace: "racer", Name: "worker-pod"}, pod)) + pod.OwnerReferences = []metav1.OwnerReference{*metav1.NewControllerRef(ds, appsv1.SchemeGroupVersion.WithKind("DaemonSet"))} + require.NoError(t, a.Topology.Update(t.Context(), pod)) + + for _, scenario := range []string{"success", "sa", "pod", "node"} { + bound := *status.DeepCopy() + + switch scenario { + case "sa": + bound.User.UID = "old-sa" + case "pod": + bound.User.Extra["authentication.kubernetes.io/pod-uid"] = authv1.ExtraValue{"old-pod"} + case "node": + bound.User.Extra["authentication.kubernetes.io/node-uid"] = authv1.ExtraValue{"old-node"} + } + + installReview(t, a, bound, token) + + req := httptest.NewRequest("POST", "https://racer/bootstrap", nil) + req.Header.Set("Authorization", "Bearer "+token) + + identity, err := a.authority.Authenticate(t.Context(), req) + require.ErrorIs(t, err, wire.Forbidden, scenario) + require.Empty(t, identity.Node(), "unconfigured workload must not authenticate") + } +} + +func TestMixedControllerEvents(t *testing.T) { + cfg := Config{Namespace: "racer", DaemonSetName: DataplaneDaemonSetName} + p := memberPod("pod", 1, "192.0.2.1") + p.OwnerReferences[0].Name = PodNetworkDaemonSetName + pred := managedPodChanges(cfg) + require.False(t, pred.Create(event.CreateEvent{Object: &p})) + require.False(t, pred.Delete(event.DeleteEvent{Object: &p})) + changed := p.DeepCopy() + changed.Status.PodIP = "192.0.2.2" + require.False(t, pred.Update(event.UpdateEvent{ObjectOld: &p, ObjectNew: changed})) + changed = p.DeepCopy() + changed.Status.Conditions = []corev1.PodCondition{{Type: corev1.PodReady, Status: corev1.ConditionTrue}} + require.False(t, pred.Update(event.UpdateEvent{ObjectOld: &p, ObjectNew: changed})) + p.OwnerReferences[0].Name = "arbitrary" + require.False(t, pred.Create(event.CreateEvent{Object: &p})) + + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: PodNetworkDaemonSetName}} + require.False(t, namedChanges(cfg.Namespace, cfg.DaemonSetName).Create(event.CreateEvent{Object: ds})) +} + +func TestLocalSnapshotsDuringAPIOutage(t *testing.T) { + f := newServingFixture(t) + + var calls atomic.Int64 + + unavailable := interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + calls.Add(1) + return errors.New("API offline") + }, + Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + calls.Add(1) + return errors.New("API offline") + }, + List: func(context.Context, client.WithWatch, client.ObjectList, ...client.ListOption) error { + calls.Add(1) + return errors.New("API offline") + }, + }) + f.a.Topology.APIReader = unavailable + fixtureDependencies[f.a.authority].reader = unavailable + + fixtureDependencies[f.a.authority].Client = unavailable + if _, err := f.a.Keyring.Reconcile(f.ctx, ctrl.Request{}); err == nil { + t.Fatal("reconciliation hid API failure") + } + + if _, err := f.a.Topology.Reconcile(f.ctx, ctrl.Request{}); err == nil { + t.Fatal("topology hid API failure") + } + + calls.Store(0) + + endpoint := f.start(t) + peer := f.client(t, &f.certificate) + + peer.Timeout = wire.PollWait + 5*time.Second + for range 2 { + response, err := peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusOK) + } + + current, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + response, err := peer.Get(fmt.Sprintf("%s%s?after=%d", endpoint, wire.SnapshotPath, current.Sequence())) + responseBody(t, response, err, http.StatusServiceUnavailable) + + if calls.Load() != 0 { + t.Fatalf("handshake/warm/204 used API: %d", calls.Load()) + } + // Issuance still needs live authorization during the same outage. + body, err := wire.EncodeBootstrapRequest(f.request) + if err != nil { + t.Fatal(err) + } + + req, err := http.NewRequestWithContext(f.ctx, http.MethodPost, endpoint+wire.BootstrapPath, bytes.NewReader(body)) + if err != nil { + t.Fatal(err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+f.token) + response, err = peer.Do(req) + responseBody(t, response, err, http.StatusServiceUnavailable) + + if calls.Load() != 0 { + t.Fatal("stale replica attempted enrollment authorization") + } + + f.cancel() + + response, err = peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusServiceUnavailable) +} + +func TestObservedInvalidTrustCannotRecoverFromReadFailure(t *testing.T) { + for _, observer := range []string{"keyring", "topology", "issuance"} { + t.Run(observer, func(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + peer := f.client(t, &f.certificate) + response, err := peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusOK) + + shared := &corev1.Secret{} + + key := client.ObjectKey{Namespace: f.a.Keyring.config.Namespace, Name: f.a.Keyring.config.CredentialsSecretName} + if err := f.a.Topology.Get(f.ctx, key, shared); err != nil { + t.Fatal(err) + } + + valid := bytes.Clone(shared.Data["bundle.json"]) + + shared.Data["bundle.json"] = []byte(`{}`) + if err := f.a.Topology.Update(f.ctx, shared); err != nil { + t.Fatal(err) + } + + switch observer { + case "keyring": + _, err = f.a.Keyring.Reconcile(f.ctx, ctrl.Request{}) + case "issuance": + _, err = f.a.authority.Issue(f.ctx, fixtureIdentity(t, f), f.request) + default: + _, err = f.a.Topology.Reconcile(f.ctx, ctrl.Request{}) + } + + if err == nil { + t.Fatal("invalid trust accepted") + } + + live := fixtureDependencies[f.a.authority].reader + + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + return errors.New("API offline after invalid observation") + }}) + if _, err := f.a.Keyring.Reconcile(f.ctx, ctrl.Request{}); err == nil { + t.Fatal("API failure hidden") + } + + response, err = peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusServiceUnavailable) + + if err := f.a.authority.TrustReady(); err == nil { + t.Fatal("read failure restored invalidated roots") + } + + fresh := f.client(t, &f.certificate) + + response, err = fresh.Get(endpoint + wire.SnapshotPath) + if err == nil { + response.Body.Close() + t.Fatal("handshake accepted missing local trust") + } + + fixtureDependencies[f.a.authority].reader = live + + shared.Data["bundle.json"] = valid + if err := f.a.Topology.Update(f.ctx, shared); err != nil { + t.Fatal(err) + } + + runKeys(t, f.a.Keyring) + reconcileTopology(t, f.a.Topology, f.ctx) + + response, err = peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusOK) + }) + } +} + +func TestTrustReadOutageAtEachAuthorityRead(t *testing.T) { + for _, resource := range []string{"racer-installation", "racer-version", "issuer.json", "bundle.json"} { + t.Run(resource, func(t *testing.T) { + if resource == "issuer.json" || resource == "bundle.json" { + resource = "racer-credentials" + } + + f := newServingFixture(t) + reads := 0 + + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if key.Name == resource { + reads++ + return errors.New("API offline") + } + + return c.Get(ctx, key, obj, opts...) + }}) + if _, err := f.a.Keyring.Reconcile(f.ctx, ctrl.Request{}); err == nil || reads != 1 { + t.Fatalf("expected read outage at %s: %v, reads=%d", resource, err, reads) + } + + if err := f.a.Server.Ready(nil); err != nil { + t.Fatalf("read outage withdrew local state: %v", err) + } + }) + } +} + +func TestIssuanceTrustObservationLockHonorsDeadline(t *testing.T) { + f := newServingFixture(t) + + release := holdFixtureGate(t, f) + defer release() + + ctx, cancel := context.WithTimeout(f.ctx, 20*time.Millisecond) + defer cancel() + + done := make(chan error, 1) + + go func() { + _, err := f.a.authority.Issue(ctx, fixtureIdentity(t, f), f.request) + done <- err + }() + + select { + case err := <-done: + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("gate wait ignored deadline: %v", err) + } + case <-time.After(time.Second): + t.Fatal("gate wait held enrollment admission past deadline") + } + + if err := f.a.authority.TrustReady(); err != nil { + t.Fatalf("canceled gate wait invalidated accepted trust: %v", err) + } +} + +func TestCatalogGateCancellationPreservesAcceptedState(t *testing.T) { + for _, operation := range []string{"topology", "keyring", "issuance"} { + for _, held := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/held=%t", operation, held), func(t *testing.T) { catalogGateCancellation(t, operation, held) }) + } + } +} + +func catalogGateCancellation(t *testing.T, operation string, held bool) { + t.Helper() + f := newServingFixture(t) + + identity := fixtureIdentity(t, f) + + roots, err := f.a.authority.TrustPool() + if err != nil { + t.Fatal(err) + } + + publication, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + var reads atomic.Int64 + + reader := interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + reads.Add(1) + return wire.Unavailable + }, + }) + if held { + release := holdFixtureGate(t, f) + defer release() + } + + fixtureDependencies[f.a.authority].reader = reader + + ctx, cancel := context.WithTimeout(f.ctx, 20*time.Millisecond) + defer cancel() + + want := context.DeadlineExceeded + + if !held { + cancel() + + want = context.Canceled + } + + done := make(chan error, 1) + + go func() { + var err error + + switch operation { + case "topology": + _, err = f.a.Topology.Reconcile(ctx, ctrl.Request{}) + case "keyring": + _, err = f.a.Keyring.Reconcile(ctx, ctrl.Request{}) + case "issuance": + _, err = f.a.authority.Issue(ctx, identity, f.request) + } + + done <- err + }() + + select { + case err := <-done: + if !errors.Is(err, want) { + t.Fatalf("gate wait cancellation: %v", err) + } + + if operation != "issuance" && !errors.Is(err, reconcile.TerminalError(nil)) { + t.Fatalf("canceled reconcile can retry: %v", err) + } + case <-time.After(time.Second): + t.Fatal("gate wait ignored cancellation") + } + + if reads.Load() != 0 { + t.Fatalf("canceled admission read authority: %d", reads.Load()) + } + + currentRoots, err := f.a.authority.TrustPool() + require.NoError(t, err) + require.True(t, currentRoots.Equal(roots), "canceled admission changed accepted trust") + + current, err := f.a.authority.Current() + require.NoError(t, err) + require.Equal(t, publication.Sequence(), current.Sequence(), "canceled admission changed publication") + + if err := f.a.Server.Ready(nil); err != nil { + t.Fatalf("canceled admission withdrew readiness: %v", err) + } +} + +func TestNodeChangesSiteLabels(t *testing.T) { + for _, key := range []string{machinav1.MachineSiteLabelKey, "net.unbounded-cloud.io/site"} { + for _, values := range [][2]string{{"", "site-a"}, {"site-a", "site-b"}, {"site-a", ""}} { + old, next := memberNode(), memberNode() + if values[0] != "" { + old.Labels = map[string]string{key: values[0]} + } + + if values[1] != "" { + next.Labels = map[string]string{key: values[1]} + } + + require.Equal(t, key == machinav1.MachineSiteLabelKey, nodeChanges().Update(event.UpdateEvent{ObjectOld: &old, ObjectNew: &next}), "%s %v", key, values) + } + } + + old, next := memberNode(), memberNode() + next.Labels = map[string]string{"unrelated": "site-a"} + require.False(t, nodeChanges().Update(event.UpdateEvent{ObjectOld: &old, ObjectNew: &next})) +} + +func TestWorkloadDefaultsAgreeWithController(t *testing.T) { + values := map[string]string{ + "RACER_CLUSTER_ID": "11111111-1111-1111-1111-111111111111", + "RACER_CONTROL_URL": "https://controller:8443", "RACER_DATAPLANE_IMAGE": "racer:test", + } + lookup := func(key string) (string, bool) { value, ok := values[key]; return value, ok } + + cfg, err := testutil.ConfigFromLookup(lookup) + if err != nil { + t.Fatal(err) + } + + runtime, err := ConfigFromLookup(lookup) + if err != nil { + t.Fatal(err) + } + + if cfg.Namespace != runtime.Namespace || cfg.PeerPort != runtime.PeerPort || cfg.DaemonSetName != runtime.DaemonSetName || cfg.DataplaneServiceAccount != runtime.DataplaneServiceAccount || cfg.BootstrapTrustConfigMap != "racer-bootstrap-trust" { + t.Fatalf("workload defaults disagree with runtime: %+v", cfg) + } +} + +type servingFixture struct { + a *Application + token string + key ed25519.PrivateKey + request wire.BootstrapRequest + certificate tls.Certificate + serverCertificate tls.Certificate + roots *x509.CertPool + ctx context.Context + cancel context.CancelFunc +} + +func newServingFixture(t *testing.T) *servingFixture { + t.Helper() + a, status, token := authFixture(t) + installReview(t, a, status, token) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + runKeys(t, a.Keyring) + reconcileTopology(t, a.Topology, ctx) + startFixtureLifecycle(t, a.Lifecycle, ctx) + + pub, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + csr, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{DNSNames: []string{"attacker"}, Subject: pkix.Name{CommonName: "attacker"}}, key) + if err != nil { + t.Fatal(err) + } + + request := wire.BootstrapRequest{SchemaVersion: 1, Cluster: a.Topology.config.Cluster, Enrollment: wire.EnrollmentID(testOtherUID), CSRDER: csr, Shares: wire.DefaultShares} + authRequest := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + authRequest.Header.Set("Authorization", "Bearer "+token) + + identity, err := a.authority.Authenticate(ctx, authRequest) + if err != nil { + t.Fatal(err) + } + + encoded, err := a.authority.Issue(ctx, identity, request) + if err != nil { + t.Fatal(err) + } + + response := decodeIssuedResponse(t, encoded) + + template := &x509.Certificate{SerialNumber: big.NewInt(1), NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), DNSNames: []string{a.Topology.config.ReplicationServerName}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}} + + der, err := x509.CreateCertificate(rand.Reader, template, template, pub, key) + if err != nil { + t.Fatal(err) + } + + cert, err := x509.ParseCertificate(der) + if err != nil { + t.Fatal(err) + } + + roots := x509.NewCertPool() + roots.AddCert(cert) + fixtureTLS(t, a, ctx, tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}) + + return &servingFixture{a: a, token: token, key: key, request: request, certificate: tls.Certificate{Certificate: response.CertificateChain, PrivateKey: key}, serverCertificate: tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}, roots: roots, ctx: ctx, cancel: cancel} +} + +func (f *servingFixture) client(t *testing.T, cert *tls.Certificate) *http.Client { + t.Helper() + + config := &tls.Config{RootCAs: f.roots, MinVersion: tls.VersionTLS13, ClientSessionCache: tls.NewLRUClientSessionCache(4)} + if cert != nil { + config.GetClientCertificate = func(*tls.CertificateRequestInfo) (*tls.Certificate, error) { return cert, nil } + } + + transport := &http.Transport{TLSClientConfig: config} + t.Cleanup(transport.CloseIdleConnections) + + return &http.Client{Transport: transport, Timeout: 5 * time.Second} +} + +func (f *servingFixture) start(t *testing.T) string { + t.Helper() + + s := httptest.NewUnstartedServer(f.a.Server.Handler()) + s.TLS = fixtureTLS(t, f.a, f.ctx, f.serverCertificate) + s.StartTLS() + t.Cleanup(s.Close) + + return s.URL +} + +func responseBody(t *testing.T, response *http.Response, err error, status int) []byte { + t.Helper() + + if err != nil { + t.Fatal(err) + } + + defer response.Body.Close() + + b, err := io.ReadAll(response.Body) + if err != nil { + t.Fatal(err) + } + + if response.StatusCode != status { + t.Fatalf("status %d, want %d: %s", response.StatusCode, status, b) + } + + if status >= 400 { + if _, err := wire.DecodeError(bytes.NewReader(b)); err != nil { + t.Fatalf("non-protocol error %q", b) + } + + if status == 429 || status == 503 { + if response.Header.Get("Retry-After") != "1" { + t.Fatal("missing retry bound") + } + } + } + + return b +} + +func startFixtureLifecycle(t *testing.T, l *server.Lifecycle, ctx context.Context) { + t.Helper() + + ctx, cancel := context.WithCancel(ctx) + done := make(chan error, 1) + + t.Cleanup(func() { + cancel() + + select { + case err := <-done: + if err != nil { + t.Error(err) + } + case <-time.After(5 * time.Second): + t.Error("fixture lifecycle did not stop") + } + }) + + synced := make(chan struct{}) + + l.SetCacheSync(func(context.Context) bool { close(synced); return true }) + + go func() { done <- l.Start(ctx) }() + + <-synced + l.SetServingReady(true) + + deadline := time.Now().Add(5 * time.Second) + for l.Ready(nil) != nil { + if ctx.Err() != nil || time.Now().After(deadline) { + t.Fatal("fixture lifecycle did not become ready") + } + + time.Sleep(time.Millisecond) + } +} + +func fixtureTLS(t *testing.T, a *Application, ctx context.Context, certificate tls.Certificate) *tls.Config { + t.Helper() + + var chain []byte + for _, der := range certificate.Certificate { + chain = append(chain, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})...) + } + + key, err := x509.MarshalPKCS8PrivateKey(certificate.PrivateKey) + if err != nil { + t.Fatal(err) + } + + if err := os.WriteFile(a.Topology.config.TLSCertificateFile, chain, 0o600); err != nil { + t.Fatal(err) + } + + if err := os.WriteFile(a.Topology.config.TLSPrivateKeyFile, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: key}), 0o600); err != nil { + t.Fatal(err) + } + + config, err := a.Server.TLSConfig(ctx) + if err != nil { + t.Fatal(err) + } + + return config +} + +func testConfig(t *testing.T) Config { + t.Helper() + t.Setenv("RACER_CLUSTER_ID", testOtherUID) + t.Setenv("POD_NAMESPACE", "racer") + + cfg, err := LoadConfig() + if err != nil { + t.Fatal(err) + } + + dir := t.TempDir() + cfg.TLSCertificateFile = filepath.Join(dir, "tls.crt") + cfg.TLSPrivateKeyFile = filepath.Join(dir, "tls.key") + + return cfg +} + +func testTopology(t *testing.T, objects ...client.Object) *TopologyReconciler { + t.Helper() + cfg := testConfig(t) + + scheme := runtime.NewScheme() + for _, add := range []func(*runtime.Scheme) error{corev1.AddToScheme, appsv1.AddToScheme, racerv1.AddToScheme} { + if err := add(scheme); err != nil { + t.Fatal(err) + } + } + + objects = append(objects, &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName, UID: "installation-uid"}, Data: map[string]string{"cluster": string(cfg.Cluster), "version_configmap": cfg.VersionConfigMapName, "state": "fresh", "initialization_protocol": "staged-v1"}}) + c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).WithIndex(&corev1.Pod{}, podNodeIndex, podNodeKeys).Build() + c = interceptor.NewClient(c, interceptor.Funcs{Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + if obj.GetUID() == "" { + obj.SetUID(types.UID(fmt.Sprintf("fake-%s-%d", obj.GetName(), time.Now().UnixNano()))) + } + + return c.Create(ctx, obj, opts...) + }}) + + return Assemble(cfg, c, c).Topology +} + +func initializedTopology(t *testing.T, objects ...client.Object) *TopologyReconciler { + t.Helper() + + r := testTopology(t, objects...) + if err := r.authority.Recover(t.Context(), r.Client); err != nil { + t.Fatal(err) + } + + return r +} + +// Captured publication data is decoded through the public response operation, +// never an installation proof. It belongs exclusively to the integration test. +type ( + CommittedPublication struct { + handle *authority.PublicationHandle + encoded string + record VersionRecord + leadership context.Context + } + VersionRecord struct { + Cluster wire.ClusterID + Sequence wire.Sequence + MembershipVersion wire.MembershipVersion + ContentHash string + MembershipHash string + } +) + +func capturePublication(t *testing.T, a *authority.Authority) *CommittedPublication { + t.Helper() + + h, err := a.Current() + if err != nil { + t.Fatal(err) + } + + return captureHandle(t, h) +} + +func captureHandle(t *testing.T, h *authority.PublicationHandle) *CommittedPublication { + t.Helper() + + var b bytes.Buffer + + guard, cancel, err := h.Admit(t.Context()) + if err != nil { + t.Fatal(err) + } + defer cancel() + + ctx := guard.Context() + if _, err := h.ForBase(0, "").WriteTo(ctx, guard, &b); err != nil { + t.Fatal(err) + } + + p, err := wire.DecodePublication(bytes.NewReader(b.Bytes())) + if err != nil { + t.Fatal(err) + } + + content, members, err := wire.ContentHashes(p) + if err != nil { + t.Fatal(err) + } + + return &CommittedPublication{handle: h, encoded: b.String(), record: VersionRecord{Cluster: p.Cluster, Sequence: p.Sequence, MembershipVersion: p.MembershipVersion, ContentHash: content, MembershipHash: members}, leadership: ctx} +} + +func (p *CommittedPublication) admit(ctx context.Context) (*authority.Admission, context.CancelFunc, error) { + return p.handle.Admit(ctx) +} + +func reconcileTopology(t *testing.T, r *TopologyReconciler, ctx context.Context) *CommittedPublication { + t.Helper() + + result, err := r.Reconcile(ctx, ctrl.Request{}) + if err != nil || result.RequeueAfter != 0 { + t.Fatalf("reconcile: %v, %v", result, err) + } + + for name := range r.hints { + result, err := r.Reconcile(ctx, ctrl.Request{NamespacedName: types.NamespacedName{Namespace: "hints", Name: name}}) + if err != nil || result != (ctrl.Result{}) { + t.Fatalf("hint reconcile: %v, %v", result, err) + } + } + + return capturePublication(t, r.authority) +} + +func acceptedMembers(t *testing.T, r *TopologyReconciler) AcceptedMembers { + t.Helper() + + p, err := r.authority.Current() + if err != nil { + return nil + } + + captured := captureHandle(t, p) + + image, err := wire.DecodePublication(bytes.NewBufferString(captured.encoded)) + if err != nil { + t.Fatal(err) + } + + members := make(AcceptedMembers, len(image.Members)) + for _, m := range image.Members { + members[m.Node] = m + } + + return members +} + +func runKeys(t *testing.T, r *KeyringReconciler) ctrl.Result { + t.Helper() + + result, err := r.Reconcile(t.Context(), ctrl.Request{}) + if err != nil || result.RequeueAfter <= 0 { + t.Fatalf("reconcile: %v, %v", result, err) + } + + return result +} + +type ( + RotationState struct { + NextRotation time.Time `json:"next_rotation"` + ActivateAt time.Time `json:"activate_at"` + ActiveIssuer string `json:"active_issuer"` + PreparedIssuer string `json:"prepared_issuer"` + Retiring map[string]time.Time `json:"retiring"` + } + signingMaterial struct { + PrivateKey []byte `json:"private_key"` + Certificate []byte `json:"certificate"` + } + issuerMaterial struct { + Keys map[string]signingMaterial `json:"keys"` + } +) + +func keyState(t *testing.T, r *KeyringReconciler) (*corev1.Secret, wire.KeyringBundle, RotationState, issuerMaterial) { + t.Helper() + + var secret corev1.Secret + + deps := fixtureDependencies[r.authority] + if err := deps.reader.Get(t.Context(), client.ObjectKey{Namespace: r.config.Namespace, Name: r.config.CredentialsSecretName}, &secret); err != nil { + t.Fatal(err) + } + + b, err := wire.DecodeBundle(bytes.NewReader(secret.Data["bundle.json"])) + if err != nil { + t.Fatal(err) + } + + var state RotationState + if err := json.Unmarshal(secret.Data["rotation.json"], &state); err != nil { + t.Fatal(err) + } + + var material issuerMaterial + if err := json.Unmarshal(secret.Data["issuer.json"], &material); err != nil { + t.Fatal(err) + } + + return &secret, b, state, material +} + +// Fault injection is a Kubernetes dependency supplied at construction, not an +// authority mutation hook. Tests may change the transport under that dependency. +type fixtureDependency struct { + client.Client + reader client.Reader + now func() time.Time +} + +func (d *fixtureDependency) Get(ctx context.Context, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + return d.reader.Get(ctx, key, obj, opts...) +} + +func (d *fixtureDependency) List(ctx context.Context, list client.ObjectList, opts ...client.ListOption) error { + return d.reader.List(ctx, list, opts...) +} + +var fixtureDependencies = map[*authority.Authority]*fixtureDependency{} + +func assembleFixture(cfg Config, c client.Client, reader client.Reader) *Application { + d := &fixtureDependency{Client: c, reader: reader, now: time.Now} + a := Assemble(cfg, d, d) + // Supply a clock through construction; no setter is exposed by authority. + owner := authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: d, Reader: d, Now: func() time.Time { return d.now() }}) + a.authority = owner + a.Topology.authority = owner + a.Keyring.authority = owner + a.Lifecycle = server.NewLifecycle(owner) + a.Server = server.New(cfg.serverConfig(), d, owner, a.Lifecycle, a.Replication) + a.Replication.authority = owner + a.Topology.Client = c + a.Topology.APIReader = reader + a.Replication.Client = c + a.Replication.APIReader = reader + fixtureDependencies[owner] = d + + return a +} + +func decodeIssuedResponse(t *testing.T, encoded []byte) wire.BootstrapResponse { + t.Helper() + + response, err := wire.DecodeBootstrapResponse(bytes.NewReader(encoded)) + if err != nil { + t.Fatal(err) + } + + return response +} + +func fixtureIdentity(t *testing.T, f *servingFixture) NodeIdentity { + t.Helper() + + r := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + r.Header.Set("Authorization", "Bearer "+f.token) + + identity, err := f.a.authority.Authenticate(t.Context(), r) + if err != nil { + t.Fatal(err) + } + + return identity +} + +type podAuthorizationReader struct { + client.Reader + pod *corev1.Pod +} + +func (r podAuthorizationReader) Get(ctx context.Context, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + switch value := obj.(type) { + case *corev1.Pod: + *value = *r.pod.DeepCopy() + return nil + case *corev1.Node: + *value = corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: r.pod.Spec.NodeName, UID: testNodeUID}} + return nil + default: + return r.Reader.Get(ctx, key, obj, opts...) + } +} + +type reviewWriter struct { + client.Writer + status authv1.TokenReviewStatus +} + +func (w reviewWriter) Create(ctx context.Context, obj client.Object, opts ...client.CreateOption) error { + if review, ok := obj.(*authv1.TokenReview); ok { + review.Status = w.status + return ctx.Err() + } + + return w.Writer.Create(ctx, obj, opts...) +} + +// Preserve authorization-only scenarios through public token authentication. +func authorizePod(ctx context.Context, reader client.Reader, cfg Config, pod *corev1.Pod, saUID string) error { + status := authv1.TokenReviewStatus{Authenticated: true, Audiences: []string{wire.TokenAudience}, User: authv1.UserInfo{Username: "system:serviceaccount:" + cfg.Namespace + ":" + cfg.DataplaneServiceAccount, UID: saUID, Extra: map[string]authv1.ExtraValue{"authentication.kubernetes.io/pod-name": {pod.Name}, "authentication.kubernetes.io/pod-uid": {string(pod.UID)}, "authentication.kubernetes.io/node-name": {pod.Spec.NodeName}, "authentication.kubernetes.io/node-uid": {testNodeUID}}}} + if pod.UID == "" { + status.User.Extra["authentication.kubernetes.io/pod-uid"] = authv1.ExtraValue{"missing"} + } + + a := authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: reviewWriter{status: status}, Reader: podAuthorizationReader{Reader: reader, pod: pod}}) + token := "header." + base64.RawURLEncoding.EncodeToString(fmt.Appendf(nil, `{"exp":%d}`, time.Now().Add(time.Hour).Unix())) + ".signature" + r := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + r.Header.Set("Authorization", "Bearer "+token) + _, err := a.Authenticate(ctx, r) + + return err +} + +const ( + installationUIDAnnotation = "racer.unbounded-cloud.io/installation-uid" + credentialClaim = "racer.unbounded-cloud.io/credentials" +) + +func readInstallation(ctx context.Context, reader client.Reader, cfg Config, fresh bool) (*corev1.ConfigMap, error) { + var cm corev1.ConfigMap + + err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName}, &cm) + + return &cm, err +} + +func readVersion(ctx context.Context, reader client.Reader, cfg Config) (*corev1.ConfigMap, VersionRecord, error) { + if err := authority.New(cfg.authorityConfig(), authority.Dependencies{Reader: reader}).Recover(ctx, nil); err != nil { + return nil, VersionRecord{}, err + } + + var cm corev1.ConfigMap + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.VersionConfigMapName}, &cm); err != nil { + return nil, VersionRecord{}, err + } + + seq, err := strconv.ParseUint(cm.Data["sequence"], 10, 64) + if err != nil { + return nil, VersionRecord{}, err + } + + members, err := strconv.ParseUint(cm.Data["membership_version"], 10, 64) + if err != nil { + return nil, VersionRecord{}, err + } + + return &cm, VersionRecord{Cluster: wire.ClusterID(cm.Data["cluster"]), Sequence: wire.Sequence(seq), MembershipVersion: wire.MembershipVersion(members), ContentHash: cm.Data["content_hash"], MembershipHash: cm.Data["membership_hash"]}, nil +} + +func ensureInstalled(ctx context.Context, writer client.Writer, reader client.Reader, cfg Config) error { + return authority.New(cfg.authorityConfig(), authority.Dependencies{Reader: reader, Writer: writer}).Recover(ctx, writer) +} + +func containsRoot(b wire.KeyringBundle, id string) bool { + for _, root := range b.PeerTrustRoots { + sum := sha256.Sum256(root) + if hex.EncodeToString(sum[:]) == id { + return true + } + } + + return false +} + +func withdrawPublication(t *testing.T, r *TopologyReconciler) func() { + t.Helper() + + cm, _, err := readVersion(t.Context(), r.APIReader, r.config) + if err != nil { + t.Fatal(err) + } + + saved := cm.DeepCopy() + + cm.Data["sequence"] = "0" + if err := r.Update(t.Context(), cm); err != nil { + t.Fatal(err) + } + + if _, err := r.authority.PublishTopology(t.Context(), r.observeTopology); err == nil { + t.Fatal("invalid publication accepted") + } + + return func() { + if err := r.Get(t.Context(), client.ObjectKeyFromObject(cm), cm); err != nil { + t.Fatal(err) + } + + cm.Data = saved.Data + if err := r.Update(t.Context(), cm); err != nil { + t.Fatal(err) + } + } +} + +func versionData(v VersionRecord) map[string]string { + return map[string]string{"cluster": string(v.Cluster), "sequence": strconv.FormatUint(uint64(v.Sequence), 10), "membership_version": strconv.FormatUint(uint64(v.MembershipVersion), 10), "content_hash": v.ContentHash, "membership_hash": v.MembershipHash} +} + +func advanceFixturePublication(t *testing.T, r *TopologyReconciler) { + t.Helper() + + members := AcceptedMembers{testNodeUID: {Node: testNodeUID, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}}} + replicationSmokePublish(t, t.Context(), r, members) +} + +func testKeyring(t *testing.T) (*KeyringReconciler, *time.Time) { + t.Helper() + + cache := &racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: "cache", UID: testNodeUID}} + r := initializedTopology(t, cache) + a := assembleFixture(r.config, r.Client, r.APIReader) + now := time.Now().UTC().Truncate(time.Second) + fixtureDependencies[a.authority].now = func() time.Time { return now } + + return a.Keyring, &now +} + +func holdFixtureGate(t *testing.T, f *servingFixture) func() { + t.Helper() + + entered, release, done := make(chan struct{}), make(chan struct{}), make(chan struct{}) + + go func() { + defer close(done) + + _, err := f.a.authority.PublishTopology(t.Context(), func(context.Context) (TopologyObservation, error) { + close(entered) + <-release + + return TopologyObservation{}, wire.Unavailable + }) + if err == nil { + t.Error("failed discovery succeeded") + } + }() + + <-entered + + return func() { close(release); <-done } +} diff --git a/internal/racer/reconcilers_test.go b/internal/racer/reconcilers_test.go new file mode 100644 index 000000000..f2a6ab7d2 --- /dev/null +++ b/internal/racer/reconcilers_test.go @@ -0,0 +1,2138 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package racer + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "reflect" + "slices" + "strconv" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + coordv1 "k8s.io/api/coordination/v1" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + "k8s.io/utils/ptr" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + "sigs.k8s.io/controller-runtime/pkg/event" + "sigs.k8s.io/controller-runtime/pkg/reconcile" + + machinav1 "github.com/Azure/unbounded/api/machina/v1alpha3" + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/server" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestTerminalPodsCannotReplaceEndpoints(t *testing.T) { + ownership := memberOwnership(t, testDaemonSetUID) + for _, phase := range []corev1.PodPhase{corev1.PodPending, corev1.PodRunning, corev1.PodUnknown, corev1.PodFailed, corev1.PodSucceeded} { + t.Run(string(phase), func(t *testing.T) { terminalPodEndpoint(t, ownership, phase) }) + } +} + +func terminalPodEndpoint(t *testing.T, ownership members.WorkloadIdentities, phase corev1.PodPhase) { + t.Helper() + + old := memberPod("old", 1, "192.0.2.1") + newest := memberPod("new", 2, "192.0.2.2") + newest.Status.Phase = phase + terminal := phase == corev1.PodFailed || phase == corev1.PodSucceeded + + want := "192.0.2.2:8082" + if terminal { + want = "192.0.2.1:8082" + } + + endpoint, err := selectEndpoint([]corev1.Pod{old, newest}, ownership, old.Spec.NodeName, 8082) + require.NoError(t, err) + require.Equal(t, want, endpoint) + + endpoint, err = selectEndpoint([]corev1.Pod{newest}, ownership, old.Spec.NodeName, 8082) + if terminal { + require.Empty(t, endpoint) + require.ErrorIs(t, err, wire.Unavailable) + + node := memberNode() + groups := map[string][]corev1.Pod{node.Name: {newest}} + + candidate, _, err := reconcileMembers([]corev1.Node{node}, groups, ownership, nil, 8082) + require.NoError(t, err) + require.Empty(t, candidate, "terminal-only Pod admitted a new node") + + previous := wire.Member{Node: testNodeUID, Shares: wire.DefaultShares, PeerEndpoint: want} + + candidate, _, err = reconcileMembers([]corev1.Node{node}, groups, ownership, AcceptedMembers{testNodeUID: previous}, 8082) + require.NoError(t, err) + require.Equal(t, previous.PeerEndpoint, candidate[testNodeUID].PeerEndpoint) + } else { + require.NoError(t, err) + require.Equal(t, want, endpoint) + } + + pred := managedPodChanges(Config{Namespace: "racer", DaemonSetName: DataplaneDaemonSetName}) + before := newest.DeepCopy() + + before.Status.Phase = corev1.PodRunning + if pred.Update(event.UpdateEvent{ObjectOld: before, ObjectNew: &newest}) != (phase != corev1.PodRunning) { + t.Fatal("Pod phase transition predicate mismatch") + } + + before = newest.DeepCopy() + + newest.Status.Conditions = []corev1.PodCondition{{Type: corev1.PodReady, Status: corev1.ConditionTrue}} + if pred.Update(event.UpdateEvent{ObjectOld: before, ObjectNew: &newest}) { + t.Fatal("readiness-only change triggered topology") + } +} + +func TestReconcilerDependencyCancellationRetries(t *testing.T) { + for _, controller := range []string{"topology", "keyring"} { + for _, stage := range []string{"read", "write", "completion read"} { + if controller == "topology" && stage == "completion read" { + continue + } + + for _, dependencyErr := range []error{context.DeadlineExceeded, context.Canceled} { + for _, cancelParent := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/%s/%v/parent=%v", controller, stage, dependencyErr, cancelParent), func(t *testing.T) { + dependencyCancellation(t, controller, stage, dependencyErr, cancelParent) + }) + } + } + } + } +} + +func dependencyCancellation(t *testing.T, controller, stage string, dependencyErr error, cancelParent bool) { + t.Helper() + topology := initializedTopology(t) + app := assembleFixture(topology.config, topology.Client, topology.APIReader) + + var target reconcile.Reconciler = app.Topology + if controller == "keyring" { + runKeys(t, app.Keyring) + _, _, state, _ := keyState(t, app.Keyring) + fixtureDependencies[app.authority].now = func() time.Time { return state.NextRotation } + target = app.Keyring + } + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + injected, reads := false, 0 + fail := func() error { + injected = true + + if cancelParent { + cancel() + } + + return fmt.Errorf("dependency: %w", dependencyErr) + } + base := topology.Client.(client.WithWatch) + wrapped := interceptor.NewClient(base, interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if _, ok := obj.(*corev1.ConfigMap); ok && key.Name == topology.config.VersionConfigMapName { + reads++ + if stage == "read" || stage == "completion read" && reads == 2 { + return fail() + } + } + + return c.Get(ctx, key, obj, opts...) + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + if stage == "write" { + return fail() + } + + return c.Update(ctx, obj, opts...) + }, + }) + app.Topology.Client, app.Topology.APIReader = wrapped, wrapped + fixtureDependencies[app.authority].Client, fixtureDependencies[app.authority].reader = wrapped, wrapped + + _, err := target.Reconcile(ctx, ctrl.Request{}) + + want := dependencyErr + if cancelParent { + want = context.Canceled + } + + if !injected || !errors.Is(err, want) || errors.Is(err, reconcile.TerminalError(nil)) != cancelParent { + t.Fatalf("injected=%v, error=%v, parent canceled=%v", injected, err, cancelParent) + } + + // A fresh reconcile succeeds without waiting for another watch event. + app.Topology.Client, app.Topology.APIReader = base, base + fixtureDependencies[app.authority].Client, fixtureDependencies[app.authority].reader = base, base + + if _, err := target.Reconcile(t.Context(), ctrl.Request{}); err != nil { + t.Fatalf("retry failed: %v", err) + } +} + +// These are real HTTPS protocol clients, not Rust processes or in-memory Wait +// calls. Kubernetes authority is fake. Never interpret this as API capacity. +func replicationSmokePublish(t *testing.T, ctx context.Context, r *TopologyReconciler, accepted AcceptedMembers) *CommittedPublication { + t.Helper() + + _, err := r.authority.PublishTopology(ctx, func(context.Context) (TopologyObservation, error) { + nodes := corev1.NodeList{} + + for id, member := range accepted { + encoded, err := json.Marshal(member) + if err != nil { + return TopologyObservation{}, err + } + + nodes.Items = append(nodes.Items, corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: string(id), UID: types.UID(id), Annotations: map[string]string{admittedMemberAnnotation: string(encoded), wire.SharesAnnotation: strconv.FormatUint(uint64(member.Shares), 10)}}}) + } + + return TopologyObservation{Nodes: nodes, Input: members.Input{Nodes: nodes.Items, PeerPort: r.config.PeerPort}}, nil + }) + if err != nil { + t.Fatal(err) + } + + return capturePublication(t, r.authority) +} + +const ( + testNodeUID = "11111111-1111-4111-8111-111111111111" + testOtherUID = "22222222-2222-4222-8222-222222222222" + testDaemonSetUID types.UID = "33333333-3333-4333-8333-333333333333" +) + +func memberNode() corev1.Node { + return corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "node-a", UID: testNodeUID}} +} + +func memberOwnership(t *testing.T, uid types.UID) DataplaneWorkloadIdentities { + t.Helper() + r := initializedTopology(t, &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Name: DataplaneDaemonSetName, Namespace: "racer", UID: uid}}) + + ids, err := readManagedWorkloadIdentities(t.Context(), r.APIReader, Config{Namespace: "racer", DaemonSetName: DataplaneDaemonSetName}) + if err != nil { + t.Fatal(err) + } + + return ids +} + +func memberPod(uid types.UID, created int64, ip string) corev1.Pod { + controller := true + + return corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "racer-" + string(uid), UID: uid, Namespace: "racer", + CreationTimestamp: metav1.NewTime(time.Unix(created, 0)), + OwnerReferences: []metav1.OwnerReference{{APIVersion: "apps/v1", Kind: "DaemonSet", Name: DataplaneDaemonSetName, UID: testDaemonSetUID, Controller: &controller}}, + }, + Spec: corev1.PodSpec{NodeName: "node-a"}, + Status: corev1.PodStatus{PodIP: ip}, + } +} + +func TestParseAnnotations(t *testing.T) { + zero, maxNUMA := uint32(0), ^uint32(0) + for _, tc := range []struct { + name string + annotations map[string]string + want MemberAttributes + wantError bool + }{ + {"defaults", nil, MemberAttributes{Shares: 4, RDMANICs: []wire.RDMANIC{}}, false}, + {"explicit empty NICs", map[string]string{wire.RDMANICsAnnotation: "[]", enrolledRDMANICsAnnotation: `[{"device":"a","port":1,"rail":0}]`}, MemberAttributes{Shares: 4, RDMANICs: []wire.RDMANIC{}}, false}, + {"enrolled fallback", map[string]string{enrolledRDMANICsAnnotation: `[{"device":"a","port":1,"rail":0}]`}, MemberAttributes{Shares: 4, RDMANICs: []wire.RDMANIC{{Device: "a", Port: 1}}}, false}, + {"legacy ignored", map[string]string{"racer.unbounded-cloud.io/rails": "malformed", "racer.unbounded-cloud.io/aligned-rails": "false"}, MemberAttributes{Shares: 4, RDMANICs: []wire.RDMANIC{}}, false}, + {"bounds and sorting", map[string]string{ + wire.SharesAnnotation: "4294967295", + wire.RDMANICsAnnotation: `[{"rail":65535,"device":"β<&>","port":255,"numa_node":4294967295},{"rail":0,"device":"a","port":1,"numa_node":0}]`, + }, MemberAttributes{Shares: ^uint32(0), RDMANICs: []wire.RDMANIC{{Rail: 0, Device: "a", Port: 1, NUMANode: &zero}, {Rail: 65535, Device: "β<&>", Port: 255, NUMANode: &maxNUMA}}}, false}, + { + "identical duplicates", + map[string]string{wire.RDMANICsAnnotation: `[{"rail":2,"device":"b","port":1},{"rail":2,"device":"b","port":1}]`}, + MemberAttributes{}, + true, + }, + { + "same rail different physical ports", + map[string]string{wire.RDMANICsAnnotation: `[{"rail":1,"device":"b","port":1},{"rail":1,"device":"a","port":2},{"rail":1,"device":"a","port":1}]`}, + MemberAttributes{Shares: 4, RDMANICs: []wire.RDMANIC{{Rail: 1, Device: "a", Port: 1}, {Rail: 1, Device: "a", Port: 2}, {Rail: 1, Device: "b", Port: 1}}}, + false, + }, + { + "unknown fields rejected", + map[string]string{wire.RDMANICsAnnotation: `[{"rail":0,"device":"a","port":1,"Rail":1,"future":{"x":true}}]`}, + MemberAttributes{}, + true, + }, + {"decimal shares", map[string]string{wire.SharesAnnotation: "0008"}, MemberAttributes{Shares: 8, RDMANICs: []wire.RDMANIC{}}, false}, + } { + t.Run(tc.name, func(t *testing.T) { + node := memberNode() + node.Annotations = tc.annotations + + got, err := ParseAnnotations(&node) + if tc.wantError { + if err == nil { + t.Fatal("unknown wire fields accepted") + } + + return + } + + if err != nil || !reflect.DeepEqual(got, tc.want) { + t.Fatalf("got %#v, %v; want %#v", got, err, tc.want) + } + }) + } +} + +func TestParseAnnotationsRejectsInvalidUpdates(t *testing.T) { + for field, values := range map[string][]string{ + wire.SharesAnnotation: {"", "0", "-1", "+1", "4294967296", "1.0", "1e2", "0x10", " 4", "4 ", "Ù¤"}, + wire.RDMANICsAnnotation: { + "", "null", "{}", "[null]", "[{}]", `[{"rail":0}]`, `[{"device":"a","port":1}]`, + `[{"rail":65536,"device":"a","port":1}]`, `[{"rail":-1,"device":"a","port":1}]`, `[{"rail":1.0,"device":"a","port":1}]`, + `[{"rail":"0","device":"a","port":1}]`, `[{"rail":0,"device":"","port":1}]`, `[{"rail":0,"device":"a\n","port":1}]`, + `[{"rail":0,"device":"a\r","port":1}]`, `[{"rail":0,"device":"a\u0000","port":1}]`, + `[{"rail":0,"device":"a","port":1,"numa_node":null}]`, `[{"rail":0,"device":"a","port":1,"numa_node":4294967296}]`, + `[{"rail":0,"device":"a","port":1,"numa_node":-1}]`, `[{"rail":0,"device":"a","port":1,"numa_node":1e0}]`, + `[{"rail":0,"device":"a","port":1,"rail":1}]`, `[{"rail":0,"device":"a","port":1,"\u0072ail":1}]`, + `[{"rail":0,"device":"a","port":1,"future":{"x":1,"x":2}}]`, + `[{"Rail":0,"device":"a","port":1}]`, `[{"rail":0,"device":"\ud800","port":1}]`, + "[{\"rail\":0,\"device\":\"\xff\",\"port\":1}]", "[] []", + `[{"rail":0,"device":"a","port":1},{"rail":1,"device":"a","port":1}]`, + `[{"rail":0,"device":"a","port":1},{"rail":0,"device":"a","port":1,"numa_node":0}]`, + `[{"rail":0,"device":"a","port":1,"numa_node":1},{"rail":0,"device":"a","port":1,"numa_node":2}]`, + `[{"rail":0,"device":"a","port":1,"future":` + strings.Repeat("[", 64) + "0" + strings.Repeat("]", 64) + "}]", + }, + } { + for _, value := range values { + t.Run(field+"/"+value, func(t *testing.T) { + node := memberNode() + node.Annotations = map[string]string{field: value} + + got, err := ParseAnnotations(&node) + if !errors.Is(err, wire.InvalidRequest) || !reflect.DeepEqual(got, MemberAttributes{}) { + t.Fatalf("invalid update produced %#v, %v", got, err) + } + }) + } + } + + if _, err := ParseAnnotations(nil); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("nil node: %v", err) + } + + node := memberNode() + + node.Annotations = map[string]string{wire.RDMANICsAnnotation: strings.Repeat(" ", 256*1024+1)} + if _, err := ParseAnnotations(&node); !errors.Is(err, wire.TooLarge) { + t.Fatalf("oversized rails: %v", err) + } +} + +func TestSelectEndpoint(t *testing.T) { + pods := []corev1.Pod{ + memberPod("a", 1, "192.0.2.1"), memberPod("z", 2, "2001:db8::1"), memberPod("b", 2, "192.0.2.2"), + } + // The newest Pod need not be Ready; a ready older Pod is not preferred. + pods[0].Status.Conditions = []corev1.PodCondition{{Type: corev1.PodReady, Status: corev1.ConditionTrue}} + for range 2 { + got, err := selectEndpoint(pods, memberOwnership(t, testDaemonSetUID), "node-a", 7443) + if err != nil || got != "[2001:db8::1]:7443" { + t.Fatalf("got %q, %v", got, err) + } + + slices.Reverse(pods) + } + + for name, mutate := range map[string]func(*corev1.Pod){ + "terminating": func(p *corev1.Pod) { p.DeletionTimestamp = &metav1.Time{} }, + "other node": func(p *corev1.Pod) { p.Spec.NodeName = "node-b" }, + "unassigned": func(p *corev1.Pod) { p.Spec.NodeName = "" }, + "missing uid": func(p *corev1.Pod) { p.UID = "" }, + "no owner": func(p *corev1.Pod) { p.OwnerReferences = nil }, + "other owner": func(p *corev1.Pod) { p.OwnerReferences[0].UID = "other" }, + "not controller": func(p *corev1.Pod) { p.OwnerReferences[0].Controller = nil }, + "false controller": func(p *corev1.Pod) { *p.OwnerReferences[0].Controller = false }, + "wrong kind": func(p *corev1.Pod) { p.OwnerReferences[0].Kind = "ReplicaSet" }, + "wrong api": func(p *corev1.Pod) { p.OwnerReferences[0].APIVersion = "other/v1" }, + "wrong namespace": func(p *corev1.Pod) { p.Namespace = "other" }, + "wrong owner name": func(p *corev1.Pod) { p.OwnerReferences[0].Name = "other" }, + "no ip": func(p *corev1.Pod) { p.Status.PodIP = "" }, + "hostname": func(p *corev1.Pod) { p.Status.PodIP = "example.com" }, + "ip with port": func(p *corev1.Pod) { p.Status.PodIP = "192.0.2.9:7443" }, + "ip with zone": func(p *corev1.Pod) { p.Status.PodIP = "fe80::1%eth0" }, + } { + t.Run(name, func(t *testing.T) { + ineligible := memberPod("new", 10, "192.0.2.9") + mutate(&ineligible) + + if _, err := selectEndpoint([]corev1.Pod{ineligible}, memberOwnership(t, testDaemonSetUID), "node-a", 7443); !errors.Is(err, wire.Unavailable) { + t.Fatalf("ineligible Pod admitted: %v", err) + } + + got, err := selectEndpoint([]corev1.Pod{ineligible, memberPod("old", 1, "192.0.2.1")}, memberOwnership(t, testDaemonSetUID), "node-a", 7443) + if err != nil || got != "192.0.2.1:7443" { + t.Fatalf("eligible older Pod lost: %q, %v", got, err) + } + }) + } + + for _, tc := range []struct { + uid types.UID + node string + port uint16 + }{ + {testDaemonSetUID, "", 7443}, {testDaemonSetUID, "node-a", 0}, + } { + if _, err := selectEndpoint(pods, memberOwnership(t, tc.uid), tc.node, tc.port); !errors.Is(err, wire.InvalidRequest) { + t.Fatalf("invalid endpoint configuration: %v", err) + } + } + + if _, err := selectEndpoint(pods, DataplaneWorkloadIdentities{}, "node-a", 7443); !errors.Is(err, wire.Unavailable) { + t.Fatalf("missing workload must have no eligible endpoint: %v", err) + } +} + +func TestReconcileMembersColdStartAndRetention(t *testing.T) { + node := memberNode() + pod := memberPod("a", 1, "192.0.2.1") + initial, diagnostics, err := reconcileMembers([]corev1.Node{node}, map[string][]corev1.Pod{node.Name: {pod}}, memberOwnership(t, testDaemonSetUID), nil, 7443) + + want := wire.Member{Node: testNodeUID, Shares: 4, RDMANICs: []wire.RDMANIC{}, PeerEndpoint: "192.0.2.1:7443"} + + require.NoError(t, err) + require.Empty(t, diagnostics) + require.Equal(t, want, initial[testNodeUID]) + + for _, tc := range []struct { + name string + annotations map[string]string + pods []corev1.Pod + shares uint32 + endpoint string + diagnostics int + }{ + {"pod gap", map[string]string{wire.SharesAnnotation: "8"}, nil, 8, "192.0.2.1:7443", 1}, + {"invalid annotations", map[string]string{wire.SharesAnnotation: "0"}, []corev1.Pod{memberPod("b", 2, "192.0.2.2")}, 4, "192.0.2.2:7443", 1}, + {"both missing", map[string]string{wire.SharesAnnotation: "0"}, nil, 4, "192.0.2.1:7443", 2}, + {"annotation unit", map[string]string{wire.SharesAnnotation: "8", wire.RDMANICsAnnotation: "invalid"}, []corev1.Pod{pod}, 4, "192.0.2.1:7443", 1}, + } { + t.Run(tc.name, func(t *testing.T) { + node := node.DeepCopy() + node.Annotations = tc.annotations + + got, diagnostics, err := reconcileMembers([]corev1.Node{*node}, map[string][]corev1.Pod{node.Name: tc.pods}, memberOwnership(t, testDaemonSetUID), initial, 7443) + require.NoError(t, err) + require.Len(t, diagnostics, tc.diagnostics) + require.Equal(t, tc.shares, got[testNodeUID].Shares) + require.Equal(t, tc.endpoint, got[testNodeUID].PeerEndpoint) + + require.Equal(t, want, initial[testNodeUID], "mutated accepted input") + + cold, diagnostics, err := reconcileMembers([]corev1.Node{*node}, map[string][]corev1.Pod{node.Name: tc.pods}, memberOwnership(t, testDaemonSetUID), nil, 7443) + require.NoError(t, err) + require.Empty(t, cold) + require.Len(t, diagnostics, tc.diagnostics) + + for _, d := range diagnostics { + require.Equal(t, node.Name, d.Object) + require.NotEmpty(t, d.Field) + require.NotEmpty(t, d.Reason) + } + }) + } +} + +func TestReconcileMembersIdentityAndRemoval(t *testing.T) { + node := memberNode() + pod := memberPod("a", 1, "192.0.2.1") + + pods := map[string][]corev1.Pod{node.Name: {pod}} + + accepted, _, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + if err != nil { + t.Fatal(err) + } + + for _, label := range []string{"", "false", "true"} { + node.Labels = map[string]string{wire.ExclusionLabel: label} + + got, diagnostics, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + require.Empty(t, got) + require.Empty(t, diagnostics) + + node.Labels = nil + + got, _, err = reconcileMembers([]corev1.Node{node}, nil, memberOwnership(t, testDaemonSetUID), got, 7443) + require.NoError(t, err) + require.Empty(t, got, "exclusion history survived") + } + + got, _, err := reconcileMembers(nil, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + require.NotNil(t, got) + require.Empty(t, got) + + node.UID = testOtherUID + + got, _, err = reconcileMembers([]corev1.Node{node}, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + require.Empty(t, got, "same-name recreation inherited history") + + got, _, err = reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + require.Len(t, got, 1) + require.EqualValues(t, testOtherUID, got[testOtherUID].Node) + // Node readiness and deletion timestamps do not change ownership while the + // Node remains in the Kubernetes input and is not explicitly excluded. + node.DeletionTimestamp = &metav1.Time{} + node.Status.Conditions = []corev1.NodeCondition{{Type: corev1.NodeReady, Status: corev1.ConditionFalse}} + + got, _, err = reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), got, 7443) + require.NoError(t, err) + require.Len(t, got, 1, "readiness removed ownership") +} + +func TestReconcileMembersDefaultsAndIsolation(t *testing.T) { + node := memberNode() + node.Annotations = map[string]string{wire.SharesAnnotation: "8", wire.RDMANICsAnnotation: `[{"rail":0,"device":"a","port":1,"numa_node":1}]`} + pod := memberPod("a", 1, "192.0.2.1") + + pods := map[string][]corev1.Pod{node.Name: {pod}} + + accepted, _, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + if err != nil { + t.Fatal(err) + } + + pods[node.Name][0].Status.PodIP = "192.0.2.9" + *pods[node.Name][0].OwnerReferences[0].Controller = false + delete(pods, node.Name) + node.Annotations[wire.SharesAnnotation] = "invalid" + + if accepted[testNodeUID].PeerEndpoint != "192.0.2.1:7443" || accepted[testNodeUID].Shares != 8 { + t.Fatal("candidate aliases Kubernetes inputs") + } + + got, _, err := reconcileMembers([]corev1.Node{node}, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + if err != nil || !reflect.DeepEqual(got, accepted) { + t.Fatalf("retention: %v, %v", got, err) + } + + got[testNodeUID].RDMANICs[0].Device = "changed" + *got[testNodeUID].RDMANICs[0].NUMANode = 9 + delete(got, testNodeUID) + + if accepted[testNodeUID].RDMANICs[0].Device != "a" || *accepted[testNodeUID].RDMANICs[0].NUMANode != 1 { + t.Fatal("retention aliases accepted state") + } + + got, _, err = reconcileMembers([]corev1.Node{node}, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + if err != nil { + t.Fatal(err) + } + + accepted[testNodeUID].RDMANICs[0].Device = "changed input" + + *accepted[testNodeUID].RDMANICs[0].NUMANode = 7 + if got[testNodeUID].RDMANICs[0].Device != "a" || *got[testNodeUID].RDMANICs[0].NUMANode != 1 { + t.Fatal("accepted input mutation changed retained output") + } + + node.Annotations = nil + + got, _, err = reconcileMembers([]corev1.Node{node}, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + if err != nil || got[testNodeUID].Shares != 4 || len(got[testNodeUID].RDMANICs) != 0 { + t.Fatalf("removed annotations did not default: %v, %v", got, err) + } +} + +func TestReconcileMembersRejectsInvalidInput(t *testing.T) { + for _, tc := range []struct { + name string + nodes []corev1.Node + uid types.UID + port uint16 + }{ + {"missing port", nil, testDaemonSetUID, 0}, + {"duplicate node", []corev1.Node{memberNode(), memberNode()}, testDaemonSetUID, 7443}, + {"missing uid", []corev1.Node{{ObjectMeta: metav1.ObjectMeta{Name: "a"}}}, testDaemonSetUID, 7443}, + {"malformed uid", []corev1.Node{{ObjectMeta: metav1.ObjectMeta{Name: "a", UID: "invalid"}}}, testDaemonSetUID, 7443}, + {"missing name", []corev1.Node{{ObjectMeta: metav1.ObjectMeta{UID: testNodeUID}}}, testDaemonSetUID, 7443}, + {"duplicate name", []corev1.Node{memberNode(), {ObjectMeta: metav1.ObjectMeta{Name: "node-a", UID: testOtherUID}}}, testDaemonSetUID, 7443}, + } { + t.Run(tc.name, func(t *testing.T) { + got, _, err := reconcileMembers(tc.nodes, nil, memberOwnership(t, tc.uid), nil, tc.port) + if !errors.Is(err, wire.InvalidRequest) || got != nil { + t.Fatalf("invalid inputs accepted: %v, %v", got, err) + } + }) + } +} + +func TestReconcileMembersGroupedPods(t *testing.T) { + nodeA, nodeB := memberNode(), memberNode() + nodeB.Name, nodeB.UID = "node-b", testOtherUID + nodeB.Annotations = map[string]string{wire.SharesAnnotation: "invalid"} + podA := memberPod("a", 1, "192.0.2.1") + wrongNode := memberPod("wrong-node", 3, "192.0.2.3") + wrongNode.Spec.NodeName = nodeB.Name + wrongOwner := memberPod("wrong-owner", 4, "192.0.2.4") + wrongOwner.OwnerReferences[0].UID = "old-daemonset" + nodes := []corev1.Node{nodeB, nodeA} + pods := map[string][]corev1.Pod{ + nodeA.Name: {wrongNode, podA, wrongOwner}, + nodeB.Name: {podA}, + "absent": {wrongNode}, + } + + podsBefore := make(map[string][]corev1.Pod, len(pods)) + for name, group := range pods { + for _, pod := range group { + podsBefore[name] = append(podsBefore[name], *pod.DeepCopy()) + } + } + + var firstDiagnostics []Diagnostic + + for range 2 { + members, diagnostics, err := reconcileMembers(nodes, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + require.NoError(t, err) + require.Len(t, members, 1) + require.Equal(t, "192.0.2.1:7443", members[testNodeUID].PeerEndpoint) + + require.Len(t, diagnostics, 2) + require.Equal(t, nodeB.Name, diagnostics[0].Object) + require.Equal(t, "annotations", diagnostics[0].Field) + require.Equal(t, nodeB.Name, diagnostics[1].Object) + require.Equal(t, "peer_endpoint", diagnostics[1].Field) + + if firstDiagnostics != nil && !reflect.DeepEqual(diagnostics, firstDiagnostics) { + t.Fatalf("input order changed diagnostics: %v, %v", diagnostics, firstDiagnostics) + } + + firstDiagnostics = diagnostics + + if !reflect.DeepEqual(pods, podsBefore) { + t.Fatal("mutated grouped Pod inputs") + } + + slices.Reverse(nodes) + slices.Reverse(pods[nodeA.Name]) + slices.Reverse(podsBefore[nodeA.Name]) + } + + // Multiple rejected nodes must report in UID order, not input or map order. + nodeA.Annotations = nodeB.Annotations + for _, nodes := range [][]corev1.Node{{nodeB, nodeA}, {nodeA, nodeB}} { + _, diagnostics, err := reconcileMembers(nodes, nil, memberOwnership(t, testDaemonSetUID), nil, 7443) + require.NoError(t, err) + require.Len(t, diagnostics, 4) + + for i, want := range []string{nodeA.Name, nodeA.Name, nodeB.Name, nodeB.Name} { + require.Equal(t, want, diagnostics[i].Object) + } + } +} + +func TestReconcileCandidateHashesAndOrdering(t *testing.T) { + nodeA, nodeB := memberNode(), memberNode() + nodeB.Name, nodeB.UID = "node-b", testOtherUID + nodeA.Annotations = map[string]string{wire.RDMANICsAnnotation: `[{"rail":1,"device":"b","port":1},{"rail":1,"device":"a","port":1}]`} + podA, podB := memberPod("a", 1, "192.0.2.1"), memberPod("b", 1, "192.0.2.2") + podB.Spec.NodeName = "node-b" + nodes := []corev1.Node{nodeB, nodeA} + pods := map[string][]corev1.Pod{nodeA.Name: {podA}, nodeB.Name: {podB}} + nodesBefore := []corev1.Node{*nodeB.DeepCopy(), *nodeA.DeepCopy()} + podsBefore := map[string][]corev1.Pod{nodeA.Name: {*podA.DeepCopy()}, nodeB.Name: {*podB.DeepCopy()}} + + members, diagnostics, err := reconcileMembers(nodes, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + require.NoError(t, err) + require.Empty(t, diagnostics) + require.Equal(t, nodesBefore, nodes) + require.Equal(t, podsBefore, pods) + + candidate := wire.Publication{SchemaVersion: wire.SchemaVersion, Cluster: testNodeUID, Members: []wire.Member{members[testOtherUID], members[testNodeUID]}} + + content, membership, err := wire.ContentHashes(candidate) + if err != nil { + t.Fatal(err) + } + + slices.Reverse(nodes) + + pods = map[string][]corev1.Pod{nodeB.Name: {podB}, nodeA.Name: {podA}} + + nodes[0].Annotations[wire.RDMANICsAnnotation] = `[{"rail":1,"device":"a","port":1},{"rail":1,"device":"b","port":1}]` + + members, _, err = reconcileMembers(nodes, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + if err != nil { + t.Fatal(err) + } + + candidate.Members = []wire.Member{members[testNodeUID], members[testOtherUID]} + + contentAgain, membershipAgain, err := wire.ContentHashes(candidate) + require.NoError(t, err) + require.Equal(t, content, contentAgain) + require.Equal(t, membership, membershipAgain) + + candidate.Caches, err = BuildCatalog([]racerv1.ClusterCache{catalogCache("cache-a", testNodeUID)}) + if err != nil { + t.Fatal(err) + } + + contentAgain, membershipAgain, err = wire.ContentHashes(candidate) + if err != nil || contentAgain == content || membershipAgain != membership { + t.Fatalf("cache-only hash change: %v", err) + } + + candidate.Members[0].PeerEndpoint = "192.0.2.9:7443" + + _, membershipAgain, err = wire.ContentHashes(candidate) + if err != nil || membershipAgain == membership { + t.Fatalf("endpoint must change membership hash: %v", err) + } +} + +func TestReconcileMembersLimit(t *testing.T) { + nodes := make([]corev1.Node, wire.MaxMembers+1) + + accepted := make(AcceptedMembers, len(nodes)) + for i := range nodes { + id := fmt.Sprintf("%08x-0000-0000-0000-000000000000", i) + nodes[i] = corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: id, UID: types.UID(id)}} + accepted[wire.NodeID(id)] = wire.Member{Node: wire.NodeID(id), Shares: 4, PeerEndpoint: "192.0.2.1:7443", RDMANICs: []wire.RDMANIC{}} + } + + got, _, err := reconcileMembers(nodes, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + if !errors.Is(err, wire.TooLarge) || got != nil { + t.Fatalf("oversized membership: %d, %v", len(got), err) + } + // The bound is on admitted members, not all observed Nodes. + nodes[0].Labels = map[string]string{wire.ExclusionLabel: ""} + + got, _, err = reconcileMembers(nodes, nil, memberOwnership(t, testDaemonSetUID), accepted, 7443) + if err != nil || len(got) != wire.MaxMembers { + t.Fatalf("membership at bound: %d, %v", len(got), err) + } +} + +func TestReconcileMembersMissingWorkloadAndRecovery(t *testing.T) { + empty, diagnostics, err := reconcileMembers(nil, nil, DataplaneWorkloadIdentities{}, nil, 7443) + require.NoError(t, err) + require.NotNil(t, empty) + require.Empty(t, empty) + require.Empty(t, diagnostics) + + node := memberNode() + pods := map[string][]corev1.Pod{node.Name: {memberPod("a", 1, "192.0.2.1")}} + + accepted, _, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + if err != nil { + t.Fatal(err) + } + + var got AcceptedMembers + for _, uid := range []types.UID{"", "replacement-daemonset"} { + got, diagnostics, err = reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, uid), accepted, 7443) + require.NoError(t, err) + require.Equal(t, accepted, got) + require.Len(t, diagnostics, 1) + + got, diagnostics, err = reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, uid), nil, 7443) + require.NoError(t, err) + require.Empty(t, got) + require.Len(t, diagnostics, 1) + } + + node.Annotations = map[string]string{wire.SharesAnnotation: "invalid-sensitive-input"} + + cold, diagnostics, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + if err != nil || len(cold) != 0 || len(diagnostics) != 1 || strings.Contains(diagnostics[0].Reason, "invalid-sensitive-input") { + t.Fatalf("cold invalid annotations: %v, %v, %v", cold, diagnostics, err) + } + + node.Annotations[wire.SharesAnnotation] = "16" + + got, diagnostics, err = reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), cold, 7443) + if err != nil || got[testNodeUID].Shares != 16 || len(diagnostics) != 0 { + t.Fatalf("corrected inputs not admitted: %v, %v, %v", got, diagnostics, err) + } +} + +func TestTopologyIndexedPodGroups(t *testing.T) { + for _, stage := range []string{"success", "list error", "canceled list"} { + t.Run(stage, func(t *testing.T) { topologyIndexedPodGroups(t, stage) }) + } +} + +func topologyIndexedPodGroups(t *testing.T, stage string) { + t.Helper() + + nodeA, nodeB := memberNode(), memberNode() + nodeB.Name, nodeB.UID = "node-b", testOtherUID + podA, podB := memberPod("a", 1, "192.0.2.1"), memberPod("b", 1, "192.0.2.2") + podB.Spec.NodeName = nodeB.Name + r := initializedTopology(t, &nodeA, &nodeB, &podA, &podB, &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{ + Name: "racer-dataplane", Namespace: "racer", UID: testDaemonSetUID, + }}) + r.config.PeerPort = 7443 + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + boom := errors.New("pod list failed") + queries := map[string]int{} + writes := 0 + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + List: func(ctx context.Context, c client.WithWatch, list client.ObjectList, opts ...client.ListOption) error { + if _, ok := list.(*corev1.PodList); ok { + options := (&client.ListOptions{}).ApplyOptions(opts) + if options.Namespace != r.config.Namespace || options.FieldSelector == nil { + t.Fatalf("Pod query lacks namespace or node index: %+v", options) + } + + name, exact := options.FieldSelector.RequiresExactMatch(podNodeIndex) + require.True(t, exact) + require.Contains(t, []string{nodeA.Name, nodeB.Name}, name) + + queries[name]++ + + if stage == "list error" { + return boom + } + + if stage == "canceled list" { + defer cancel() + } + } + + return c.List(ctx, list, opts...) + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + writes++ + return c.Update(ctx, obj, opts...) + }, + }) + + result, err := r.Reconcile(ctx, ctrl.Request{}) + if stage == "success" { + require.NoError(t, err) + require.Zero(t, result.RequeueAfter) + require.Equal(t, 1, queries[nodeA.Name]) + require.Equal(t, 1, queries[nodeB.Name]) + accepted := acceptedMembers(t, r) + require.Len(t, accepted, 2) + require.Equal(t, "192.0.2.1:7443", accepted[testNodeUID].PeerEndpoint) + require.Equal(t, "192.0.2.2:7443", accepted[testOtherUID].PeerEndpoint) + + return + } + + wantErr := boom + if stage == "canceled list" { + wantErr = context.Canceled + + if !errors.Is(err, reconcile.TerminalError(nil)) { + t.Fatalf("canceled list is not terminal: %v", err) + } + } + + if !errors.Is(err, wantErr) || result.RequeueAfter != 0 || len(queries) != 1 || writes != 0 || len(acceptedMembers(t, r)) != 0 { + t.Fatalf("failed listing changed state or continued: queries=%v writes=%d members=%v result=%v err=%v", queries, writes, acceptedMembers(t, r), result, err) + } + + if _, err := r.authority.Current(); err == nil { + t.Fatal("installed publication after failed listing") + } +} + +// Exercise the same ownership snapshot through discovery, durable publication, +// key admission, and restart rather than injecting an endpoint callback. +func TestTopologyOwnershipHistoryAndCatalogRestart(t *testing.T) { + for _, workload := range []string{DataplaneDaemonSetName, PodNetworkDaemonSetName, "standalone-racer"} { + t.Run(workload, func(t *testing.T) { + node := memberNode() + node.Annotations = map[string]string{wire.SharesAnnotation: "8"} + pod := memberPod("current", 1, "192.0.2.1") + pod.OwnerReferences[0].Name = workload + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: workload, UID: testDaemonSetUID}} + + r := initializedTopology(t, &node, &pod, ds) + r.config.DaemonSetName = workload + + r.config.PeerPort = 7443 + a := Assemble(r.config, r.Client, r.APIReader) + r = a.Topology + runKeys(t, a.Keyring) + first := reconcileTopology(t, r, t.Context()) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + history := node.Annotations[admittedMemberAnnotation] + require.NotEmpty(t, history) + + cache := catalogCache("cache-a", testOtherUID) + require.NoError(t, r.Create(t.Context(), &cache)) + require.Equal(t, first.encoded, reconcileTopology(t, r, t.Context()).encoded, "cache waits for committed keys") + runKeys(t, a.Keyring) + withCache := reconcileTopology(t, r, t.Context()) + published, err := wire.DecodePublication(strings.NewReader(withCache.encoded)) + require.NoError(t, err) + require.Equal(t, []wire.CacheDefinition{{ID: testOtherUID, Name: "cache-a", ClientSocket: "/run/racer/cache-a/client/socket", OriginSocket: "/run/racer/cache-a/origin/socket"}}, published.Caches) + require.Len(t, published.Members, 1) + require.Equal(t, first.record.MembershipVersion, withCache.record.MembershipVersion) + + // The workload was recreated while the old Pod still exists. A fresh + // process must recover history, not admit that Pod under its stale UID. + require.NoError(t, r.Delete(t.Context(), ds)) + ds.UID, ds.ResourceVersion = "replacement", "" + require.NoError(t, r.Create(t.Context(), ds)) + + node.Annotations[wire.SharesAnnotation] = "malformed" + require.NoError(t, r.Update(t.Context(), &node)) + a = Assemble(r.config, r.Client, r.APIReader) + r = a.Topology + restarted := reconcileTopology(t, r, t.Context()) + require.Equal(t, withCache.encoded, restarted.encoded) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + require.Equal(t, history, node.Annotations[admittedMemberAnnotation]) + + // A new owned endpoint recovers independently of malformed attributes. + replacement := memberPod("replacement", 2, "2001:db8::2") + replacement.OwnerReferences[0] = *metav1.NewControllerRef(ds, appsv1.SchemeGroupVersion.WithKind("DaemonSet")) + require.NoError(t, r.Create(t.Context(), &replacement)) + recovered := reconcileTopology(t, r, t.Context()) + published, err = wire.DecodePublication(strings.NewReader(recovered.encoded)) + require.NoError(t, err) + require.Equal(t, "[2001:db8::2]:7443", published.Members[0].PeerEndpoint) + require.Equal(t, uint32(8), published.Members[0].Shares) + require.Len(t, published.Caches, 1) + + // Whole-candidate rejection leaves both publication and Node history + // unchanged even when there is a valid membership update to publish. + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + history = node.Annotations[admittedMemberAnnotation] + node.Annotations[wire.SharesAnnotation] = "16" + require.NoError(t, r.Update(t.Context(), &node)) + + invalid := catalogCache("cache-b", "not-a-uuid") + require.NoError(t, r.Create(t.Context(), &invalid)) + _, err = r.Reconcile(t.Context(), ctrl.Request{}) + require.ErrorIs(t, err, wire.InvalidRequest) + current, err := r.authority.Current() + require.NoError(t, err) + require.Equal(t, recovered.encoded, captureHandle(t, current).encoded) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + require.Equal(t, history, node.Annotations[admittedMemberAnnotation]) + + require.NoError(t, r.Delete(t.Context(), &invalid)) + require.NoError(t, r.Delete(t.Context(), &cache)) + + node.Labels = map[string]string{wire.ExclusionLabel: ""} + require.NoError(t, r.Update(t.Context(), &node)) + excluded := reconcileTopology(t, r, t.Context()) + published, err = wire.DecodePublication(strings.NewReader(excluded.encoded)) + require.NoError(t, err) + require.Empty(t, published.Members) + require.Empty(t, published.Caches) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + require.Empty(t, node.Annotations[admittedMemberAnnotation]) + }) + } +} + +func TestTopologyCustomWorkloadOwnership(t *testing.T) { + for _, scenario := range []string{"current", "wrong name", "stale UID", "wrong namespace", "wrong kind", "wrong API version", "not controller", "labels only", "missing workload", "deleting workload"} { + t.Run(scenario, func(t *testing.T) { + node := memberNode() + pod := memberPod("custom", 1, "192.0.2.1") + pod.OwnerReferences[0].Name = "standalone-racer" + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "standalone-racer", UID: testDaemonSetUID}} + + switch scenario { + case "wrong name": + pod.OwnerReferences[0].Name = DataplaneDaemonSetName + case "stale UID": + pod.OwnerReferences[0].UID = "stale" + case "wrong namespace": + pod.Namespace = "other" + case "wrong kind": + pod.OwnerReferences[0].Kind = "Deployment" + case "wrong API version": + pod.OwnerReferences[0].APIVersion = "apps/v2" + case "not controller": + pod.OwnerReferences[0].Controller = nil + case "labels only": + pod.OwnerReferences = nil + pod.Labels = map[string]string{"app.kubernetes.io/name": ds.Name} + case "deleting workload": + ds.Finalizers = []string{"test/hold"} + } + + r := initializedTopology(t, &node, &pod, ds) + + r.config.DaemonSetName = ds.Name + if scenario == "missing workload" || scenario == "deleting workload" { + require.NoError(t, r.Delete(t.Context(), ds)) + } + + committed := reconcileTopology(t, r, t.Context()) + published, err := wire.DecodePublication(strings.NewReader(committed.encoded)) + require.NoError(t, err) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + + if scenario == "current" { + require.Len(t, published.Members, 1) + require.NotEmpty(t, node.Annotations[admittedMemberAnnotation]) + } else { + require.Empty(t, published.Members) + require.Empty(t, node.Annotations[admittedMemberAnnotation]) + } + }) + } +} + +// Test-local vocabulary keeps the original membership scenarios readable without +// re-exporting the extracted packages through the production controller API. +type ( + AcceptedMembers = members.History + MemberAttributes = members.MemberAttributes + Diagnostic = members.Diagnostic + DataplaneWorkloadIdentities = members.WorkloadIdentities + TopologyObservation = authority.TopologyObservation + NodeIdentity = authority.NodeIdentity + RotationPolicy = authority.RotationPolicy +) + +const ( + DataplaneDaemonSetName = "racer-dataplane" + PodNetworkDaemonSetName = "racer-dataplane-podnet" +) + +func readManagedWorkloadIdentities(ctx context.Context, reader client.Reader, cfg Config) (members.WorkloadIdentities, error) { + return members.ReadWorkloadIdentities(ctx, reader, cfg.Namespace, cfg.DaemonSetName) +} + +func ParseAnnotations(node *corev1.Node) (MemberAttributes, error) { + return members.ParseAnnotations(node) +} + +func selectEndpoint(pods []corev1.Pod, ownership DataplaneWorkloadIdentities, nodeName string, port uint16) (string, error) { + return members.SelectEndpoint(pods, ownership, nodeName, port) +} + +func reconcileMembers(nodes []corev1.Node, podsByNode map[string][]corev1.Pod, ownership DataplaneWorkloadIdentities, accepted AcceptedMembers, port uint16) (AcceptedMembers, []Diagnostic, error) { + result, err := members.Reconcile(members.Input{ + Nodes: nodes, PodsByNode: podsByNode, Ownership: ownership, PeerPort: port, + }, accepted) + + return result.Members, result.Diagnostics, err +} + +func BuildCatalog(caches []racerv1.ClusterCache) ([]wire.CacheDefinition, error) { + return members.BuildCatalog(caches) +} + +func TestLegacyRDMADiagnosticsAndMalformedUnitRetention(t *testing.T) { + node := memberNode() + node.Annotations = map[string]string{"racer.unbounded-cloud.io/rails": "malformed", "racer.unbounded-cloud.io/aligned-rails": "false"} + pods := map[string][]corev1.Pod{node.Name: {memberPod("a", 1, "192.0.2.1")}} + accepted, diagnostics, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), nil, 7443) + require.NoError(t, err) + require.Empty(t, diagnostics) + require.Empty(t, accepted[testNodeUID].RDMANICs) + + node.Annotations = map[string]string{wire.SharesAnnotation: "8", enrolledRDMANICsAnnotation: `[{"device":"a","port":1,"rail":0}]`} + accepted, _, err = reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + + node.Annotations[wire.SharesAnnotation] = "9" + node.Annotations[wire.RDMANICsAnnotation] = "null" + retained, diagnostics, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + require.Len(t, diagnostics, 1) + require.Equal(t, accepted, retained) + + node.Annotations[wire.RDMANICsAnnotation] = "[]" + attributes, err := ParseAnnotations(&node) + require.NoError(t, err) + require.Equal(t, uint32(9), attributes.Shares) + require.Empty(t, attributes.RDMANICs) +} + +func TestMixedControllerTopologyLiveOwnership(t *testing.T) { + node := memberNode() + pod := memberPod("podnet", 1, "192.0.2.2") + pod.OwnerReferences[0].Name = PodNetworkDaemonSetName + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: PodNetworkDaemonSetName, UID: testDaemonSetUID}} + r := initializedTopology(t, &node, &pod, ds) + r.config.DaemonSetName = ds.Name + reconcileTopology(t, r, t.Context()) + require.Len(t, acceptedMembers(t, r), 1) + before := acceptedMembers(t, r)[testNodeUID] + r.APIReader = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if _, ok := obj.(*appsv1.DaemonSet); ok { + return errors.New("injected workload read failure") + } + + return c.Get(ctx, key, obj, opts...) + }, + }) + _, err := r.Reconcile(t.Context(), ctrl.Request{}) + require.Error(t, err) + require.Equal(t, before, acceptedMembers(t, r)[testNodeUID]) + r.APIReader = r.Client + require.NoError(t, r.Delete(t.Context(), ds)) + ds.UID, ds.ResourceVersion = "replacement", "" + require.NoError(t, r.Create(t.Context(), ds)) + r = Assemble(r.config, r.Client, r.APIReader).Topology + reconcileTopology(t, r, t.Context()) + require.Equal(t, before, acceptedMembers(t, r)[testNodeUID]) +} + +type mixedFailReader struct{ client.Reader } + +func (r mixedFailReader) Get(context.Context, client.ObjectKey, client.Object, ...client.GetOption) error { + return errors.New("injected read failure") +} + +func TestMixedControllerAuthorization(t *testing.T) { + cfg := Config{Namespace: "racer", DaemonSetName: DataplaneDaemonSetName, DataplaneServiceAccount: "racer"} + scheme := runtime.NewScheme() + require.NoError(t, appsv1.AddToScheme(scheme)) + require.NoError(t, corev1.AddToScheme(scheme)) + + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: "racer", UID: "sa-current"}} + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: PodNetworkDaemonSetName, UID: "pod-current"}} + c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(sa, ds).Build() + pod := memberPod("pod", 1, "192.0.2.1") + pod.Spec.ServiceAccountName = "racer" + pod.OwnerReferences = []metav1.OwnerReference{*metav1.NewControllerRef(ds, appsv1.SchemeGroupVersion.WithKind("DaemonSet"))} + require.ErrorIs(t, authorizePod(t.Context(), c, cfg, &pod, string(sa.UID)), wire.Forbidden) + cfg.DaemonSetName = ds.Name + require.NoError(t, authorizePod(t.Context(), c, cfg, &pod, string(sa.UID))) + require.ErrorIs(t, authorizePod(t.Context(), c, cfg, &pod, "old-sa"), wire.Forbidden) + require.ErrorIs(t, authorizePod(t.Context(), mixedFailReader{c}, cfg, &pod, string(sa.UID)), wire.Unavailable) + ids, err := readManagedWorkloadIdentities(t.Context(), mixedFailReader{c}, cfg) + require.Error(t, err) + require.False(t, ids.Owns(&pod)) + + for _, mutate := range []func(*corev1.Pod){ + func(p *corev1.Pod) { p.OwnerReferences[0].Name = "arbitrary" }, + func(p *corev1.Pod) { p.OwnerReferences[0].UID = "stale" }, + func(p *corev1.Pod) { + p.OwnerReferences = nil + p.Labels = map[string]string{"app": DataplaneDaemonSetName} + }, + func(p *corev1.Pod) { p.Status.Phase = corev1.PodFailed }, + func(p *corev1.Pod) { p.Spec.ServiceAccountName = "other" }, + func(p *corev1.Pod) { p.UID = "" }, + func(p *corev1.Pod) { p.DeletionTimestamp = &metav1.Time{} }, + } { + bad := pod.DeepCopy() + mutate(bad) + require.ErrorIs(t, authorizePod(t.Context(), c, cfg, bad, string(sa.UID)), wire.Forbidden) + } + + require.NoError(t, c.Delete(t.Context(), ds)) + require.ErrorIs(t, authorizePod(t.Context(), c, cfg, &pod, string(sa.UID)), wire.Forbidden) + ds.ResourceVersion, ds.UID = "", "recreated" + require.NoError(t, c.Create(t.Context(), ds)) + require.ErrorIs(t, authorizePod(t.Context(), c, cfg, &pod, string(sa.UID)), wire.Forbidden) + pod.OwnerReferences[0].UID = ds.UID + require.NoError(t, authorizePod(t.Context(), c, cfg, &pod, string(sa.UID))) + + ds.Finalizers = []string{"test/hold"} + require.NoError(t, c.Update(t.Context(), ds)) + require.NoError(t, c.Delete(t.Context(), ds)) + require.ErrorIs(t, authorizePod(t.Context(), c, cfg, &pod, string(sa.UID)), wire.Forbidden) +} + +func TestMixedControllerMembership(t *testing.T) { + node := memberNode() + node.Annotations = map[string]string{enrolledSharesAnnotation: "8"} + host := memberPod("host", 1, "192.0.2.1") + host.OwnerReferences[0].Name = DataplaneDaemonSetName + pod := memberPod("pod", 2, "192.0.2.2") + pod.OwnerReferences[0].Name, pod.OwnerReferences[0].UID = PodNetworkDaemonSetName, "podnet" + hostDS := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: DataplaneDaemonSetName, UID: testDaemonSetUID}} + podDS := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: PodNetworkDaemonSetName, UID: "podnet"}} + r := initializedTopology(t, &node, &host, hostDS, podDS) + r.config.PeerPort = 7443 + reconcileTopology(t, r, t.Context()) + require.NoError(t, r.Create(t.Context(), &pod)) + reconcileTopology(t, r, t.Context()) + after := acceptedMembers(t, r) + require.Len(t, after, 1) + require.Equal(t, uint32(8), after[testNodeUID].Shares) + require.Equal(t, "192.0.2.1:7443", after[testNodeUID].PeerEndpoint, "unconfigured workload cannot replace endpoint") + + require.NoError(t, r.Delete(t.Context(), &host)) + require.NoError(t, r.Delete(t.Context(), podDS)) + podDS.UID, podDS.ResourceVersion = "recreated", "" + require.NoError(t, r.Create(t.Context(), podDS)) + r = Assemble(r.config, r.Client, r.APIReader).Topology + reconcileTopology(t, r, t.Context()) + require.Equal(t, after, acceptedMembers(t, r), "restart retains UID-bound last admitted endpoint") + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + delete(node.Annotations, admittedMemberAnnotation) + require.NoError(t, r.Update(t.Context(), &node)) + r = Assemble(r.config, r.Client, r.APIReader).Topology + reconcileTopology(t, r, t.Context()) + require.Empty(t, acceptedMembers(t, r), "stale owner cannot admit a new member") + + pod.OwnerReferences[0].UID = hostDS.UID + pod.OwnerReferences[0].Name = "arbitrary" + require.NoError(t, r.Update(t.Context(), &pod)) + reconcileTopology(t, r, t.Context()) + require.Empty(t, acceptedMembers(t, r), "a live UID with the wrong owner name cannot admit a member") +} + +func TestTrustRequiresFreshPostReconcileCredentials(t *testing.T) { + for _, resource := range []string{"racer-installation", "racer-version", "issuer.json", "bundle.json"} { + for _, failure := range []string{"outage", "deleted", "malformed"} { + t.Run(resource+"/"+failure, func(t *testing.T) { postReconcileCredentials(t, resource, failure) }) + } + } +} + +func postReconcileCredentials(t *testing.T, resource, failure string) { + t.Helper() + + resourceName := resource + if resource == "issuer.json" || resource == "bundle.json" { + resourceName = "racer-credentials" + } + + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + accepted, err := r.authority.TrustPool() + if err != nil { + t.Fatal(err) + } + + acceptedBundle, err := r.authority.Keyring() + require.NoError(t, err) + require.EqualValues(t, 1, acceptedBundle.Generation()) + + reads := 0 + d := fixtureDependencies[r.authority] + d.reader = interceptor.NewClient(d.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if key.Name == resourceName { + reads++ + if reads == 2 { + return failedCredentialRead(ctx, c, key, obj, failure, opts...) + } + } + + return c.Get(ctx, key, obj, opts...) + }}) + + _, err = r.Reconcile(t.Context(), ctrl.Request{}) + require.Error(t, err) + require.Equal(t, 2, reads) + + current, err := r.authority.TrustPool() + if failure == "outage" { + require.NoError(t, err) + require.True(t, current.Equal(accepted)) + require.NoError(t, r.authority.TrustReady()) + } else { + require.Error(t, err) + require.Error(t, r.authority.TrustReady()) + } + + currentBundle, bundleErr := r.authority.Keyring() + if failure == "outage" { + require.NoError(t, bundleErr) + require.True(t, acceptedBundle == currentBundle) + } else { + require.Error(t, bundleErr) + } + + d.reader = d.Client + + _, staged, _, _ := keyState(t, r) + require.EqualValues(t, 2, staged.Generation) + require.Len(t, staged.PeerTrustRoots, 2) + + runKeys(t, r) + + current, err = r.authority.TrustPool() + require.NoError(t, err) + require.False(t, current.Equal(accepted)) + + currentBundle, bundleErr = r.authority.Keyring() + require.NoError(t, bundleErr) + require.EqualValues(t, 2, currentBundle.Generation()) +} + +func failedCredentialRead(ctx context.Context, c client.Reader, key client.ObjectKey, obj client.Object, failure string, opts ...client.GetOption) error { + switch failure { + case "outage": + return errors.New("post-reconcile API outage") + case "deleted": + return apierrors.NewNotFound(corev1.Resource("secrets"), key.Name) + default: + if err := c.Get(ctx, key, obj, opts...); err != nil { + return err + } + + switch value := obj.(type) { + case *corev1.Secret: + value.Data = nil + case *corev1.ConfigMap: + value.Data = nil + } + + return nil + } +} + +func TestKeyringCancellationOverridesPostReconcileReadFailure(t *testing.T) { + for _, failure := range []string{"outage", "conflict", "success"} { + t.Run(failure, func(t *testing.T) { keyringCompletionCancellation(t, failure) }) + } +} + +func keyringCompletionCancellation(t *testing.T, failure string) { + t.Helper() + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + reads := 0 + d := fixtureDependencies[r.authority] + d.reader = interceptor.NewClient(d.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if key.Name == r.config.CredentialsSecretName { + reads++ + if reads == 2 { + defer cancel() + + switch failure { + case "outage": + return errors.New("post-reconcile API outage") + case "conflict": + return apierrors.NewConflict(corev1.Resource("secrets"), key.Name, wire.Conflict) + } + } + } + + return c.Get(ctx, key, obj, opts...) + }}) + + result, err := r.Reconcile(ctx, ctrl.Request{}) + if reads != 2 || !errors.Is(err, context.Canceled) || !errors.Is(err, reconcile.TerminalError(nil)) || result != (ctrl.Result{}) { + t.Fatalf("post-reconcile cancellation: reads=%d result=%v err=%v", reads, result, err) + } + + if err := r.authority.TrustReady(); err == nil { + t.Fatal("cancellation after admission retained trust or issuer readiness") + } + + if _, err := r.authority.Keyring(); err == nil { + t.Fatal("cancellation after admission retained delivery") + } + + d.reader = d.Client + + _, staged, _, _ := keyState(t, r) + if staged.Generation != 2 || len(staged.PeerTrustRoots) != 2 { + t.Fatal("cancellation preceded successful rotation publication") + } +} + +func TestReconcilerAlreadyExistsHandling(t *testing.T) { + for _, operation := range []string{"keyring", "topology"} { + t.Run(operation, func(t *testing.T) { reconcilerAlreadyExists(t, operation) }) + } +} + +func TestProductionReconcilerConflictAndCancellation(t *testing.T) { + for _, operation := range []string{"topology", "keyring"} { + for _, stage := range []string{"before reconcile", "read", "conflict", "canceled conflict", "after commit"} { + t.Run(operation+"/"+stage, func(t *testing.T) { + productionConflictAndCancellation(t, operation, stage) + }) + } + } +} + +func productionConflictAndCancellation(t *testing.T, operation, stage string) { + t.Helper() + r := initializedTopology(t) + a := assembleFixture(r.config, r.Client, r.APIReader) + d := fixtureDependencies[a.authority] + base := d.Client.(client.WithWatch) + + var target reconcile.Reconciler = a.Topology + if operation == "keyring" { + runKeys(t, a.Keyring) + _, _, state, _ := keyState(t, a.Keyring) + d.now = func() time.Time { return state.NextRotation } + target = a.Keyring + } + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + writes := 0 + wrapped := interceptor.NewClient(base, interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + err := c.Get(ctx, key, obj, opts...) + + if stage == "read" { + cancel() + } + + return err + }, + Update: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.UpdateOption) error { + writes++ + + if stage == "conflict" || stage == "canceled conflict" { + if stage == "canceled conflict" { + cancel() + } + + return apierrors.NewConflict(corev1.Resource("configmaps"), obj.GetName(), wire.Conflict) + } + + err := c.Update(ctx, obj, opts...) + + cancel() + + return err + }, + }) + d.Client, d.reader = wrapped, wrapped + a.Topology.Client, a.Topology.APIReader = wrapped, wrapped + + if stage == "before reconcile" { + cancel() + } + + result, err := target.Reconcile(ctx, ctrl.Request{}) + require.Equal(t, ctrl.Result{}, result, "errors must not use a fixed retry delay") + + if stage == "conflict" { + require.True(t, apierrors.IsConflict(err)) + require.NotErrorIs(t, err, reconcile.TerminalError(nil), "queue must retry conflicts") + } else { + require.ErrorIs(t, err, context.Canceled) + require.ErrorIs(t, err, reconcile.TerminalError(nil), "canceled leader must not retry") + } + + if stage == "before reconcile" || stage == "read" { + require.Zero(t, writes) + } else { + require.Equal(t, 1, writes) + } + + if operation == "topology" { + require.ErrorIs(t, a.authority.PublicationReady(), wire.Unavailable) + } else if stage == "before reconcile" { + require.NoError(t, a.authority.TrustReady()) + } else { + require.ErrorIs(t, a.authority.TrustReady(), wire.Unavailable) + } + + d.Client, d.reader = base, base + a.Topology.Client, a.Topology.APIReader = base, base + result, err = target.Reconcile(t.Context(), ctrl.Request{}) + require.NoError(t, err, "fresh controller attempt must recover") + + if operation == "keyring" { + require.Positive(t, result.RequeueAfter, "successful rotation retains its deadline") + require.NoError(t, a.authority.TrustReady()) + } else { + require.Equal(t, ctrl.Result{}, result) + require.NoError(t, a.authority.PublicationReady()) + } +} + +func reconcilerAlreadyExists(t *testing.T, operation string) { + t.Helper() + r, now := testKeyring(t) + runKeys(t, r) + _, _, initial, _ := keyState(t, r) + *now = initial.NextRotation + + accepted, err := r.authority.TrustPool() + if err != nil { + t.Fatal(err) + } + + writes := 0 + d := fixtureDependencies[r.authority] + writer := interceptor.NewClient(d.Client.(client.WithWatch), interceptor.Funcs{Update: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.UpdateOption) error { + writes++ + return apierrors.NewAlreadyExists(corev1.Resource("secrets"), obj.GetName()) + }}) + + var result ctrl.Result + + if operation == "keyring" { + d.Client = writer + + result, err = r.Reconcile(t.Context(), ctrl.Request{}) + if !apierrors.IsAlreadyExists(err) || result != (ctrl.Result{}) { + t.Fatalf("keyring AlreadyExists not requeued: %v %v", result, err) + } + + if err := r.authority.TrustReady(); err == nil { + t.Fatal("keyring write failure retained trust or issuer readiness") + } + } else { + topology := Assemble(r.config, writer, d.reader).Topology + + result, err = topology.Reconcile(t.Context(), ctrl.Request{}) + if !apierrors.IsAlreadyExists(err) || result != (ctrl.Result{}) { + t.Fatalf("topology AlreadyExists treated as Conflict: %v %v", result, err) + } + + if current, err := r.authority.TrustPool(); err != nil || !current.Equal(accepted) { + t.Fatalf("topology publication write failure changed trust: %v", err) + } + } + + if writes != 1 { + t.Fatalf("expected one failed write, got %d", writes) + } +} + +func TestTopologyAnnotationDoesNotBlockObserver(t *testing.T) { + for _, outcome := range []string{"success", "failure", "cancellation"} { + t.Run(outcome, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { annotationDoesNotBlockObserver(t, outcome) }) + }) + } +} + +func annotationDoesNotBlockObserver(t *testing.T, outcome string) { + t.Helper() + f := newServingFixture(t) + r := f.a.Topology + + var node corev1.Node + require.NoError(t, r.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + node.Annotations["racer.unbounded-cloud.io/shares"] = "7" + require.NoError(t, r.Update(f.ctx, &node)) + + entered, release := make(chan struct{}), make(chan struct{}) + patchError := errors.New("patch failed") + r.Client = interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{ + Patch: func(ctx context.Context, c client.WithWatch, obj client.Object, patch client.Patch, opts ...client.PatchOption) error { + close(entered) + + select { + case <-ctx.Done(): + return ctx.Err() + case <-release: + } + + if outcome == "failure" { + return patchError + } + + return c.Patch(ctx, obj, patch, opts...) + }, + }) + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + done := make(chan error, 1) + + var queued reconcile.Request + + r.enqueueHint = func(request reconcile.Request) { queued = request } + _, err := r.Reconcile(ctx, ctrl.Request{}) + require.NoError(t, err) + require.Equal(t, "hints", queued.Namespace) + + go func() { _, err := r.Reconcile(ctx, queued); done <- err }() + + <-entered + require.EqualValues(t, 7, acceptedMembers(t, r)[testNodeUID].Shares) + // Cross the original freshness deadline while the actual Node patch + // remains blocked. The real observer must renew both accepted states. + for range 3 { + time.Sleep(20 * time.Second) + + observed := make(chan struct{}) + + go func() { f.a.Replication.observe(f.ctx); close(observed) }() + + synctest.Wait() + + select { + case <-observed: + default: + t.Fatal("observer blocked behind annotation patch") + } + + require.NoError(t, f.a.Server.Ready(nil)) + } + + if outcome == "cancellation" { + cancel() + } else { + close(release) + } + + err = <-done + + switch outcome { + case "success": + require.NoError(t, err) + case "failure": + require.ErrorIs(t, err, patchError) + require.NotEmpty(t, r.hints) + case "cancellation": + require.ErrorIs(t, err, context.Canceled) + } + + gateCtx, stop := context.WithTimeout(f.ctx, time.Second) + defer stop() + + require.NoError(t, f.a.authority.Observe(gateCtx)) +} + +func TestReplicaInstallationAndFreshness(t *testing.T) { + synctest.Test(t, replicaInstallationAndFreshness) +} + +func replicaInstallationAndFreshness(t *testing.T) { + t.Helper() + leader := initializedTopology(t) + publication := reconcileTopology(t, leader, t.Context()) + follower := Assemble(leader.config, leader.Client, leader.APIReader) + + process, cancel := context.WithCancel(t.Context()) + defer cancel() + + follower.authority.BindProcess(process) + + image, err := wire.DecodePublication(strings.NewReader(publication.encoded)) + if err != nil { + t.Fatal(err) + } + + bad := image + + bad.Sequence++ + if err := follower.authority.AcceptReplica(t.Context(), process, bad); err == nil { + t.Fatal("unconfirmed counters installed") + } + + if err := follower.authority.AcceptReplica(t.Context(), process, image); err != nil { + t.Fatal(err) + } + + _, err = follower.authority.Current() + if err != nil || capturePublication(t, follower.authority).encoded != publication.encoded { + t.Fatal("replica did not install canonical image", err) + } + + time.Sleep(20 * time.Second) + + if err := follower.authority.AcceptReplica(t.Context(), process, image); err != nil { + t.Fatal(err) + } + + time.Sleep(20 * time.Second) + + if follower.authority.PublicationReady() != nil { + t.Fatal("unchanged authoritative confirmation did not renew freshness") + } + + time.Sleep(11 * time.Second) + + if follower.authority.PublicationReady() == nil { + t.Fatal("expired image still serves") + } + + if err := follower.authority.AcceptReplica(t.Context(), process, image); err != nil { + t.Fatal(err) + } + + if capturePublication(t, follower.authority).encoded != publication.encoded { + t.Fatal("interruption discarded validated image") + } + + rollback := publication.record + + rollback.ContentHash = strings.Repeat("0", 64) + + cm, _, err := readVersion(t.Context(), leader.APIReader, leader.config) + if err != nil { + t.Fatal(err) + } + + cm.Data = versionData(rollback) + if err := leader.Update(t.Context(), cm); err != nil { + t.Fatal(err) + } + + if follower.authority.AcceptReplica(t.Context(), process, image) == nil { + t.Fatal("same-counter corruption accepted") + } + + cancel() + + if follower.authority.PublicationReady() == nil { + t.Fatal("process cancellation ignored") + } +} + +func TestReplicaServingSurvivesPublisherCancellation(t *testing.T) { + r := initializedTopology(t) + r.authority.BindProcess(t.Context()) + publisher, cancel := context.WithCancel(t.Context()) + publication := reconcileTopology(t, r, publisher) + + cancel() + + if _, err := r.authority.Current(); err != nil { + t.Fatal("publisher lifetime leaked into serving", err) + } + + if _, _, err := publication.admit(t.Context()); err != nil { + t.Fatal("image bound to publisher instead of process") + } +} + +func TestReplicaObservationsFailClosed(t *testing.T) { + f := newServingFixture(t) + r := f.a.Replication + r.observe(f.ctx) + + if f.a.Server.Ready(nil) != nil { + t.Fatal("valid observation withdrew readiness") + } + + r.APIReader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + return errors.New("offline") + }}) + fixtureDependencies[r.authority].reader = r.APIReader + r.observe(f.ctx) + + if f.a.Server.Ready(nil) != nil { + t.Fatal("transport interruption discarded recent state") + } + + r.APIReader = f.a.Topology.APIReader + fixtureDependencies[r.authority].reader = r.APIReader + + cm, _, err := readVersion(f.ctx, r.APIReader, r.config) + if err != nil { + t.Fatal(err) + } + + cm.Data["sequence"] = "0" + if err := r.Client.Update(f.ctx, cm); err != nil { + t.Fatal(err) + } + + r.observe(f.ctx) + + if f.a.Server.Ready(nil) == nil { + t.Fatal("observed invalid authority still serves") + } +} + +func TestReplicaLeaderDiscovery(t *testing.T) { + f := newServingFixture(t) + + r := f.a.Replication + if err := coordv1.AddToScheme(r.Client.Scheme()); err != nil { + t.Fatal(err) + } + + r.config.ControllerServiceAccount = "racer-controller" + r.config.ReplicationPort = 8443 + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: r.config.Namespace, Name: "controller", UID: "controller-uid"}, Spec: corev1.PodSpec{ServiceAccountName: "racer-controller"}, Status: corev1.PodStatus{PodIP: "192.0.2.10"}} + + lease := &coordv1.Lease{ObjectMeta: metav1.ObjectMeta{Namespace: r.config.Namespace, Name: "racer-controller"}, Spec: coordv1.LeaseSpec{HolderIdentity: ptr.To("controller/controller-uid"), RenewTime: ptr.To(metav1.NewMicroTime(time.Now())), LeaseDurationSeconds: ptr.To(int32(15))}} + for _, obj := range []client.Object{pod, lease} { + if err := r.Client.Create(f.ctx, obj); err != nil { + t.Fatal(err) + } + } + + if address, err := r.leaderAddress(f.ctx); err != nil || address != "192.0.2.10:8443" { + t.Fatal(address, err) + } + + lease.Spec.HolderIdentity = ptr.To("controller/replaced-uid") + if err := r.Client.Update(f.ctx, lease); err != nil { + t.Fatal(err) + } + + if _, err := r.leaderAddress(f.ctx); err == nil { + t.Fatal("replaced leader Pod accepted") + } +} + +func TestReplicaObservedHighWaterWithoutImage(t *testing.T) { + for _, initial := range []string{"installed10", "no image", "suspended"} { + for _, mutation := range []string{"rollback10", "conflicting11", "membership hash", "membership rollback", "membership jump"} { + t.Run(initial+"/"+mutation, func(t *testing.T) { observedHighWater(t, initial, mutation) }) + } + } +} + +func observedHighWater(t *testing.T, initial, mutation string) { + t.Helper() + f := newServingFixture(t) + r := f.a.Replication + base := *capturePublication(t, f.a.authority) + base.record.Sequence, base.record.MembershipVersion = 10, 5 + + newer := base.record + newer.Sequence = 11 + newer.ContentHash = strings.Repeat("a", 64) + setVersion := func(record VersionRecord) { + cm := &corev1.ConfigMap{} + + err := r.APIReader.Get(f.ctx, client.ObjectKey{Namespace: r.config.Namespace, Name: r.config.VersionConfigMapName}, cm) + if err != nil { + t.Fatal(err) + } + + cm.Data = versionData(record) + if err := r.Client.Update(f.ctx, cm); err != nil { + t.Fatal(err) + } + } + setVersion(base.record) + + if initial == "no image" { + d := fixtureDependencies[f.a.authority] + a := authority.New(r.config.authorityConfig(), authority.Dependencies{Reader: d, Writer: d}) + f.a.authority = a + r.authority = a + f.a.Lifecycle = server.NewLifecycle(a) + f.a.Server = server.New(r.config.serverConfig(), d, a, f.a.Lifecycle, r) + fixtureTLS(t, f.a, f.ctx, f.serverCertificate) + f.a.Lifecycle.SetCacheSync(func(context.Context) bool { return true }) + + go func() { _ = f.a.Lifecycle.Start(f.ctx) }() + + f.a.Lifecycle.SetServingReady(true) + f.a.Topology.authority = a + f.a.Keyring.authority = a + fixtureDependencies[a] = d + } else { + image, err := wire.DecodePublication(strings.NewReader(base.encoded)) + if err != nil { + t.Fatal(err) + } + + image.Sequence, image.MembershipVersion = 10, 5 + if err := r.authority.AcceptReplica(f.ctx, f.ctx, image); err != nil { + t.Fatal(err) + } + + if initial == "suspended" { + restore := withdrawPublication(t, f.a.Topology) + restore() + } + } + + setVersion(newer) + r.observe(f.ctx) + + setVersion(invalidHighWater(base.record, newer, mutation)) + r.observe(f.ctx) + + require.Error(t, f.a.authority.PublicationReady()) + require.Error(t, f.a.Server.Ready(nil)) + + require.Error(t, r.authority.TrustReady(), "invalid authority retained trust") + + image, decodeErr := wire.DecodePublication(strings.NewReader(base.encoded)) + if decodeErr != nil { + t.Fatal(decodeErr) + } + + image.Sequence, image.MembershipVersion = 10, 5 + if err := r.authority.AcceptReplica(f.ctx, f.ctx, image); err == nil { + t.Fatal("install bypassed observed high-water") + } + + // Exercise the public serving boundary, not only store readiness: + // rejected authority must not expose even the previously valid image. + request := httptest.NewRequest(http.MethodGet, wire.SnapshotPath, nil) + request.TLS = f.requestState(t) + response := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(response, request) + + if response.Code != http.StatusServiceUnavailable { + t.Fatalf("invalid authority still served snapshot: %d", response.Code) + } + + // Restoring the last valid authority permits a new reconcile/CAS, + // without forgetting the observed watermark or reusing revoked bytes. + setVersion(newer) + runKeys(t, f.a.Keyring) + + recovered := reconcileTopology(t, f.a.Topology, f.ctx) + require.Equal(t, newer.Sequence+1, recovered.record.Sequence) + require.Equal(t, newer.MembershipVersion, recovered.record.MembershipVersion) + + response = httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(response, request.Clone(f.ctx)) + + require.Equal(t, http.StatusOK, response.Code) + require.Equal(t, recovered.encoded, response.Body.String()) +} + +func invalidHighWater(base, newer VersionRecord, mutation string) VersionRecord { + bad := newer + + switch mutation { + case "rollback10": + bad = base + case "conflicting11": + bad.ContentHash = strings.Repeat("b", 64) + case "membership hash": + bad.Sequence++ + bad.MembershipHash = strings.Repeat("b", 64) + case "membership rollback": + bad.Sequence++ + bad.MembershipVersion-- + case "membership jump": + bad.Sequence++ + bad.MembershipVersion += 2 + } + + return bad +} + +type firstChunkWriter struct { + calls int + entered, unblock chan struct{} +} + +func (w *firstChunkWriter) Write(b []byte) (int, error) { + w.calls++ + if w.calls == 1 { + close(w.entered) + <-w.unblock + } + + return len(b), nil +} + +func TestPublicationWriteAuthorityRevocation(t *testing.T) { + for _, change := range []string{"superseded", "suspend recover", "confirmation after write admission"} { + t.Run(change, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + membership := AcceptedMembers{} + + for i := range 500 { + id := wire.NodeID(fmt.Sprintf("22222222-2222-4222-8222-%012d", i)) + membership[id] = wire.Member{Node: id, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}} + } + + image := replicationSmokePublish(t, t.Context(), r, membership) + copy := *image + require.Greater(t, len(copy.encoded), 32*1024) + + w := &firstChunkWriter{entered: make(chan struct{}), unblock: make(chan struct{})} + done := make(chan error, 1) + + guard, stopWrite, err := copy.admit(t.Context()) + if err != nil { + t.Fatal(err) + } + defer stopWrite() + + go func() { _, err := copy.handle.ForBase(0, "").WriteTo(t.Context(), guard, w); done <- err }() + + <-w.entered + + switch change { + case "superseded": + advanceFixturePublication(t, r) + case "suspend recover": + restore := withdrawPublication(t, r) + restore() + replicationSmokePublish(t, t.Context(), r, membership) + case "confirmation after write admission": + time.Sleep(20 * time.Second) + + replicationSmokePublish(t, t.Context(), r, membership) + + time.Sleep(11 * time.Second) + } + + if r.authority.PublicationReady() != nil { + t.Fatal("replacement or confirmed image unavailable") + } + + close(w.unblock) + + err = <-done + if change == "superseded" { + require.NoError(t, err, "ordinary advancement revoked admitted response") + require.Equal(t, (len(copy.encoded)+32*1024-1)/(32*1024), w.calls) + _, _, err = copy.admit(t.Context()) + require.ErrorIs(t, err, wire.Unavailable) + } else if err == nil || w.calls != 1 { + t.Fatalf("revoked response continued: calls=%d err=%v", w.calls, err) + } + }) + }) + } +} + +func TestMemberSiteLabelsOverrideHistory(t *testing.T) { + for _, tc := range []struct { + name string + labels map[string]string + want string + }{ + {"unlabeled", nil, ""}, + {"canonical", map[string]string{machinav1.MachineSiteLabelKey: "site-a"}, "site-a"}, + {"deprecated ignored", map[string]string{"net.unbounded-cloud.io/site": "site-b"}, ""}, + {"canonical only", map[string]string{machinav1.MachineSiteLabelKey: "site-a", "net.unbounded-cloud.io/site": "site-b"}, "site-a"}, + {"empty canonical no fallback", map[string]string{machinav1.MachineSiteLabelKey: "", "net.unbounded-cloud.io/site": "site-b"}, ""}, + {"empty labels", map[string]string{machinav1.MachineSiteLabelKey: ""}, ""}, + } { + t.Run(tc.name, func(t *testing.T) { + for _, mode := range []string{"cold", "memory", "restart"} { + t.Run(mode, func(t *testing.T) { + node := memberNode() + node.Labels = tc.labels + pods := map[string][]corev1.Pod{node.Name: {memberPod("a", 1, "192.0.2.1")}} + + var accepted AcceptedMembers + + previous := wire.Member{Node: testNodeUID, Shares: 8, PeerEndpoint: "192.0.2.1:7443", RDMANICs: []wire.RDMANIC{{Rail: 1, Device: "mlx5_0", Port: 1}}, Site: "stale-site"} + + if mode != "cold" { + node.Annotations = map[string]string{wire.RDMANICsAnnotation: "malformed"} + pods = nil + + if mode == "memory" { + accepted = AcceptedMembers{testNodeUID: previous} + } else { + encoded, err := json.Marshal(previous) + require.NoError(t, err) + + node.Annotations[admittedMemberAnnotation] = string(encoded) + } + } + + got, diagnostics, err := reconcileMembers([]corev1.Node{node}, pods, memberOwnership(t, testDaemonSetUID), accepted, 7443) + require.NoError(t, err) + require.Len(t, got, 1) + require.Equal(t, tc.want, got[testNodeUID].Site) + + if mode == "cold" { + require.Empty(t, diagnostics) + } else { + require.Len(t, diagnostics, 2) + + previous.Site = tc.want + require.Equal(t, previous, got[testNodeUID]) + } + }) + } + }) + } +} + +func TestTopologySiteChangesPersistAcrossRestart(t *testing.T) { + node := memberNode() + node.Labels = map[string]string{machinav1.MachineSiteLabelKey: "site-a"} + node.Annotations = map[string]string{wire.RDMANICsAnnotation: `[{"rail":0,"device":"mlx5_0","port":1}]`} + pod := memberPod("a", 1, "192.0.2.1") + r := initializedTopology(t, &node, &pod, &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Name: DataplaneDaemonSetName, Namespace: "racer", UID: testDaemonSetUID}}) + first := reconcileTopology(t, r, t.Context()) + base, err := wire.DecodePublication(strings.NewReader(first.encoded)) + require.NoError(t, err) + require.Equal(t, "site-a", base.Members[0].Site) + + for _, site := range []string{"site-b", ""} { + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + + node.Labels = nil + if site != "" { + node.Labels = map[string]string{machinav1.MachineSiteLabelKey: site} + } + + node.Annotations[wire.RDMANICsAnnotation] = "malformed" + require.NoError(t, r.Update(t.Context(), &node)) + // Drop all in-memory history before processing the changed boundary. + r = Assemble(r.config, r.Client, r.APIReader).Topology + committed := reconcileTopology(t, r, t.Context()) + next, err := wire.DecodePublication(strings.NewReader(committed.encoded)) + require.NoError(t, err) + require.Equal(t, site, next.Members[0].Site) + require.Equal(t, base.Members[0].RDMANICs, next.Members[0].RDMANICs) + require.Equal(t, base.Sequence+1, next.Sequence) + require.Equal(t, base.MembershipVersion+1, next.MembershipVersion) + oldContent, oldMembership, err := wire.ContentHashes(base) + require.NoError(t, err) + content, membership, err := wire.ContentHashes(next) + require.NoError(t, err) + require.NotEqual(t, oldContent, content) + require.NotEqual(t, oldMembership, membership) + + delta, err := wire.EncodeDelta(base, next) + require.NoError(t, err) + applied, err := wire.ApplyDelta(base, strings.NewReader(string(delta))) + require.NoError(t, err) + require.Equal(t, next, applied) + require.NoError(t, r.Get(t.Context(), client.ObjectKeyFromObject(&node), &node)) + + var saved wire.Member + require.NoError(t, json.Unmarshal([]byte(node.Annotations[admittedMemberAnnotation]), &saved)) + require.Equal(t, next.Members[0], saved) + + r = Assemble(r.config, r.Client, r.APIReader).Topology + require.Equal(t, committed.encoded, reconcileTopology(t, r, t.Context()).encoded) + + base = next + } +} diff --git a/internal/racer/replication_test.go b/internal/racer/replication_test.go new file mode 100644 index 000000000..2d17ccff3 --- /dev/null +++ b/internal/racer/replication_test.go @@ -0,0 +1,399 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package racer + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "errors" + "math/big" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "os" + "path/filepath" + "strconv" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + authv1 "k8s.io/api/authentication/v1" + coordv1 "k8s.io/api/coordination/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/client-go/kubernetes" + "k8s.io/utils/ptr" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + "sigs.k8s.io/controller-runtime/pkg/event" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/testutil" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func unusedAddress(t *testing.T) string { + t.Helper() + + l, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + address := l.Addr().String() + require.NoError(t, l.Close()) + + return address +} + +func integrationTLS(t *testing.T, cfg *Config) *x509.CertPool { + t.Helper() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{SerialNumber: big.NewInt(1), NotBefore: time.Now().Add(-time.Minute), NotAfter: time.Now().Add(time.Hour), DNSNames: []string{cfg.ReplicationServerName}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("127.0.0.2"), net.ParseIP("127.0.0.3")}, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}} + der, err := x509.CreateCertificate(rand.Reader, template, template, pub, key) + require.NoError(t, err) + private, err := x509.MarshalPKCS8PrivateKey(key) + require.NoError(t, err) + dir := t.TempDir() + + cfg.TLSCertificateFile, cfg.TLSPrivateKeyFile = filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key") + for path, block := range map[string]*pem.Block{cfg.TLSCertificateFile: {Type: "CERTIFICATE", Bytes: der}, cfg.TLSPrivateKeyFile: {Type: "PRIVATE KEY", Bytes: private}} { + require.NoError(t, os.WriteFile(path, pem.EncodeToMemory(block), 0o600)) + } + + roots := x509.NewCertPool() + roots.AppendCertsFromPEM(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})) + + return roots +} + +func boundPodToken(t *testing.T, kube kubernetes.Interface, pod *corev1.Pod, serviceAccount, audience string) string { + t.Helper() + token, err := kube.CoreV1().ServiceAccounts(pod.Namespace).CreateToken(t.Context(), serviceAccount, &authv1.TokenRequest{Spec: authv1.TokenRequestSpec{Audiences: []string{audience}, ExpirationSeconds: ptr.To(int64(3600)), BoundObjectRef: &authv1.BoundObjectReference{APIVersion: "v1", Kind: "Pod", Name: pod.Name, UID: pod.UID}}}, metav1.CreateOptions{}) + require.NoError(t, err) + + return token.Status.Token +} + +func workloadConfig(t *testing.T) testutil.Config { + t.Helper() + + return testutil.Config{ + Cluster: "11111111-1111-1111-1111-111111111111", Namespace: "racer", + ControlURL: "https://racer-controller.racer.svc:8443", DataplaneImage: "racer:test", + BootstrapTrustConfigMap: "racer-bootstrap-trust", PeerPort: 8082, + DataplaneServiceAccount: "racer-dataplane", DaemonSetName: "racer-dataplane", + } +} + +func TestWorkloadPeerMembership(t *testing.T) { + for _, port := range []uint16{8082, 7443, 9090, 9091, 65535} { + t.Run(strconv.Itoa(int(port)), func(t *testing.T) { + cfg := workloadConfig(t) + cfg.PeerPort = port + ds, err := testutil.DesiredDaemonSet(cfg) + require.NoError(t, err) + assertWorkloadPeerMembership(t, ds, port) + }) + } +} + +// Exercise ordered downward-API expansion and membership together. +func assertWorkloadPeerMembership(t *testing.T, ds *appsv1.DaemonSet, peerPort uint16) { + t.Helper() + + for _, ips := range [][]string{{"192.0.2.1"}, {"2001:db8::1"}, {"192.0.2.1", "2001:db8::1"}, {"2001:db8::1", "192.0.2.1"}} { + pod := memberPod("peer", 1, ips[0]) + for _, ip := range ips { + pod.Status.PodIPs = append(pod.Status.PodIPs, corev1.PodIP{IP: ip}) + } + + podIP, listen := "", "" + + for _, env := range ds.Spec.Template.Spec.Containers[0].Env { + switch env.Name { + case "RACER_POD_IP": + require.Empty(t, podIP) + require.Empty(t, env.Value) + require.Equal(t, &corev1.EnvVarSource{FieldRef: &corev1.ObjectFieldSelector{APIVersion: "v1", FieldPath: "status.podIP"}}, env.ValueFrom) + + podIP = pod.Status.PodIP + case "RACER_PEER_LISTEN": + if podIP == "" || listen != "" || env.ValueFrom != nil || env.Value != "[$(RACER_POD_IP)]:"+strconv.Itoa(int(peerPort)) { + t.Fatal("peer listener must expand the preceding Pod IP and configured peer port") + } + + listen = strings.ReplaceAll(env.Value, "$(RACER_POD_IP)", podIP) + } + } + + host, port, err := net.SplitHostPort(listen) + require.NoError(t, err) + ip, err := netip.ParseAddr(host) + require.NoError(t, err) + require.False(t, ip.IsUnspecified()) + require.Equal(t, strconv.Itoa(int(peerPort)), port) + candidate, diagnostics, err := reconcileMembers([]corev1.Node{memberNode()}, map[string][]corev1.Pod{pod.Spec.NodeName: {pod}}, memberOwnership(t, testDaemonSetUID), nil, peerPort) + require.NoError(t, err) + require.Empty(t, diagnostics) + require.Len(t, candidate, 1) + require.Equal(t, netip.AddrPortFrom(ip, peerPort).String(), candidate[testNodeUID].PeerEndpoint) + } +} + +func TestCacheChanges(t *testing.T) { + cache := catalogCache("cache", testNodeUID) + p := cacheChanges() + require.True(t, p.Create(event.CreateEvent{Object: &cache})) + require.True(t, p.Delete(event.DeleteEvent{Object: &cache})) + + for _, tt := range []struct { + name string + mutate func(*racerv1.ClusterCache) + want bool + }{ + {"unchanged", func(*racerv1.ClusterCache) {}, false}, + {"resource version", func(c *racerv1.ClusterCache) { c.ResourceVersion = "2" }, false}, + {"labels", func(c *racerv1.ClusterCache) { c.Labels = map[string]string{"test": "value"} }, false}, + {"uid", func(c *racerv1.ClusterCache) { c.UID = testOtherUID }, true}, + {"name", func(c *racerv1.ClusterCache) { c.Name = "other" }, true}, + } { + t.Run(tt.name, func(t *testing.T) { + updated := cache.DeepCopy() + tt.mutate(updated) + require.Equal(t, tt.want, p.Update(event.UpdateEvent{ObjectOld: &cache, ObjectNew: updated})) + }) + } + + require.True(t, p.Update(event.UpdateEvent{ObjectOld: &corev1.Node{}, ObjectNew: &cache})) +} + +func TestReplicationValidation(t *testing.T) { + cfg := testConfig(t) + cfg.PodName, cfg.PodUID = "controller", "controller-uid" + require.NoError(t, cfg.validateReplication()) + + for name, mutate := range map[string]func(*Config){ + "pod name": func(c *Config) { c.PodName = "Invalid" }, + "pod UID": func(c *Config) { c.PodUID = "" }, + "service account": func(c *Config) { c.ControllerServiceAccount = "" }, + "server name": func(c *Config) { c.ReplicationServerName = "" }, + "port": func(c *Config) { c.ReplicationPort = 0 }, + "token": func(c *Config) { c.ReplicationTokenFile = "" }, + "trust": func(c *Config) { c.ReplicationTrustFile = "" }, + } { + t.Run(name, func(t *testing.T) { + invalid := cfg + mutate(&invalid) + require.Error(t, invalid.validateReplication()) + }) + } +} + +func TestReplicationLifetime(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := &Replication{config: Config{SnapshotMaxAge: 3 * time.Second}.effective()} + p := &publisherLifetime{replication: r} + require.True(t, p.NeedLeaderElection()) + require.False(t, r.NeedLeaderElection()) + require.Equal(t, time.Second, r.PollInterval()) + require.Equal(t, 5*time.Second, Assemble(Config{}, nil, nil).Replication.PollInterval()) + + ctx, leader := r.LeaderContext() + require.Nil(t, ctx) + require.False(t, leader) + + ctx, cancel := context.WithCancel(t.Context()) + done := make(chan error, 1) + + go func() { done <- p.Start(ctx) }() + + synctest.Wait() + + observed, leader := r.LeaderContext() + require.Same(t, ctx, observed) + require.True(t, leader) + cancel() + require.NoError(t, <-done) + require.False(t, r.isLeader()) + require.False(t, replicationSleep(ctx, time.Hour)) + require.True(t, replicationSleep(t.Context(), time.Second)) + }) +} + +func TestReplicationObserverLoop(t *testing.T) { + for _, leader := range []bool{false, true} { + t.Run(strconv.FormatBool(leader), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + require.NoError(t, coordv1.AddToScheme(r.Scheme())) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + observer := Assemble(r.config, r.Client, r.APIReader).Replication + if leader { + observer.leader = ctx + } + + done := make(chan error, 1) + + go func() { done <- observer.Start(ctx) }() + + time.Sleep(6 * time.Second) + cancel() + require.NoError(t, <-done) + }) + }) + } +} + +func replicationPeer(t *testing.T, handler http.Handler) (*Replication, *httptest.Server) { + t.Helper() + r := initializedTopology(t) + cfg := r.config + integrationTLS(t, &cfg) + cfg.ReplicationTrustFile = cfg.TLSCertificateFile + cfg.ReplicationTokenFile = filepath.Join(t.TempDir(), "token") + require.NoError(t, os.WriteFile(cfg.ReplicationTokenFile, []byte(" token-value\n"), 0o600)) + certificate, err := tls.LoadX509KeyPair(cfg.TLSCertificateFile, cfg.TLSPrivateKeyFile) + require.NoError(t, err) + + peer := httptest.NewUnstartedServer(handler) + peer.TLS = &tls.Config{MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{certificate}} + peer.StartTLS() + t.Cleanup(peer.Close) + host, port, err := net.SplitHostPort(peer.Listener.Addr().String()) + require.NoError(t, err) + number, err := strconv.ParseUint(port, 10, 16) + require.NoError(t, err) + + cfg.ReplicationPort = uint16(number) + + require.NoError(t, coordv1.AddToScheme(r.Scheme())) + + for _, obj := range []client.Object{ + &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: "controller", UID: "controller-uid"}, Spec: corev1.PodSpec{ServiceAccountName: cfg.ControllerServiceAccount}, Status: corev1.PodStatus{PodIP: host}}, + &coordv1.Lease{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: "racer-controller"}, Spec: coordv1.LeaseSpec{HolderIdentity: ptr.To("controller/controller-uid"), RenewTime: ptr.To(metav1.NewMicroTime(time.Now())), LeaseDurationSeconds: ptr.To(int32(60))}}, + } { + require.NoError(t, r.Create(t.Context(), obj)) + } + + return Assemble(cfg, r.Client, r.APIReader).Replication, peer +} + +func TestReplicationPollResponses(t *testing.T) { + for _, status := range []int{http.StatusOK, http.StatusNoContent, http.StatusForbidden, http.StatusTemporaryRedirect} { + t.Run(strconv.Itoa(status), func(t *testing.T) { + var image string + + requests := make(chan *http.Request, 2) + r, _ := replicationPeer(t, http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + requests <- req + + w.Header().Set("Location", "/redirected") + w.WriteHeader(status) + + if status == http.StatusOK { + _, _ = w.Write([]byte(image)) + } + })) + topology := Assemble(r.config, r.Client, r.APIReader).Topology + image = reconcileTopology(t, topology, t.Context()).encoded + + err := r.poll(t.Context(), t.Context()) + if status == http.StatusOK || status == http.StatusNoContent { + require.NoError(t, err) + } else { + require.EqualError(t, err, "replication HTTP status "+strconv.Itoa(status)) + } + + req := <-requests + require.Equal(t, "Bearer token-value", req.Header.Get("Authorization")) + require.Empty(t, req.URL.RawQuery) + require.Empty(t, requests, "redirects must not forward credentials") + + if status == http.StatusOK { + require.NoError(t, r.poll(t.Context(), t.Context())) + require.Equal(t, "1", (<-requests).URL.Query().Get("after")) + } else { + require.Error(t, r.authority.PublicationReady(), "responses must not grant freshness") + } + }) + } +} + +func TestReplicationPollFailures(t *testing.T) { + for _, stage := range []string{"leader", "trust missing", "trust invalid", "token", "TLS", "decode", "unconfirmed"} { + t.Run(stage, func(t *testing.T) { + body := "invalid" + r, _ := replicationPeer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte(body)) })) + cfg := r.config + + switch stage { + case "leader": + r.APIReader = fake.NewClientBuilder().WithScheme(r.Client.Scheme()).Build() + case "trust missing": + cfg.ReplicationTrustFile += ".missing" + case "trust invalid": + require.NoError(t, os.WriteFile(cfg.ReplicationTrustFile, []byte("invalid"), 0o600)) + case "token": + cfg.ReplicationTokenFile += ".missing" + case "TLS": + cfg.ReplicationServerName = "wrong.invalid" + case "unconfirmed": + p := reconcileTopology(t, Assemble(cfg, r.Client, r.APIReader).Topology, t.Context()) + image, err := wire.DecodePublication(strings.NewReader(p.encoded)) + require.NoError(t, err) + + image.Sequence++ + encoded, err := wire.EncodePublication(image) + require.NoError(t, err) + + body = string(encoded) + } + + r = Assemble(cfg, r.Client, r.APIReader).Replication + require.Error(t, r.poll(t.Context(), t.Context())) + require.Error(t, r.authority.PublicationReady()) + }) + } +} + +func TestLeaderDiscoveryRejectsInvalidLease(t *testing.T) { + for _, holder := range []*string{nil, ptr.To(""), ptr.To("controller"), ptr.To("/uid"), ptr.To("controller/"), ptr.To("missing/uid")} { + r, _ := replicationPeer(t, http.NotFoundHandler()) + + var lease coordv1.Lease + + key := client.ObjectKey{Namespace: r.config.Namespace, Name: "racer-controller"} + require.NoError(t, r.Client.Get(t.Context(), key, &lease)) + lease.Spec.HolderIdentity = holder + require.NoError(t, r.Client.Update(t.Context(), &lease)) + _, err := r.leaderAddress(t.Context()) + require.Error(t, err) + } +} + +func TestReplicaAuthenticationDelegates(t *testing.T) { + r := initializedTopology(t) + boom := errors.New("review failed") + c := interceptor.NewClient(r.Client.(client.WithWatch), interceptor.Funcs{Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { return boom }}) + request := httptest.NewRequest(http.MethodGet, "/", nil) + request.Header.Set("Authorization", "Bearer token") + uid, expires, err := Assemble(r.config, c, c).Replication.AuthenticateReplica(t.Context(), request) + require.Error(t, err) + require.Empty(t, uid) + require.Zero(t, expires) +} diff --git a/internal/racer/server/handlers_test.go b/internal/racer/server/handlers_test.go new file mode 100644 index 000000000..e5f88407b --- /dev/null +++ b/internal/racer/server/handlers_test.go @@ -0,0 +1,2066 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "bytes" + "context" + "crypto/tls" + "crypto/x509" + "encoding/base64" + "encoding/json" + "encoding/pem" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/http/httptrace" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "syscall" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/types" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestHTTPSBootstrapSnapshotAndStrictRoutes(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + anonymous := f.client(t, nil) + + req := bootstrapTestRequest(t, f.ctx, endpoint, f.token, f.request) + response, err := anonymous.Do(req) + encoded := responseBody(t, response, err, 200) + + enrollment, err := wire.DecodeBootstrapResponse(bytes.NewReader(encoded)) + if err != nil { + t.Fatal(err) + } + + leaf, err := x509.ParseCertificate(enrollment.CertificateChain[0]) + if err != nil { + t.Fatal(err) + } + + if enrollment.Node != wire.NodeID(testNodeUID) || len(leaf.DNSNames) != 0 || leaf.Subject.CommonName != "" { + t.Fatal("CSR became authority") + } + + response, err = anonymous.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, 401) + peer := f.client(t, &f.certificate) + response, err = peer.Get(endpoint + wire.SnapshotPath) + + publication, err := wire.DecodePublication(bytes.NewReader(responseBody(t, response, err, 200))) + if err != nil || len(publication.Members) != 1 { + t.Fatalf("snapshot: %v", err) + } + + for _, tc := range []struct { + method, path string + status int + }{ + {"GET", "/v1/bootstrap", 400}, + {"POST", "/v1/snapshot", 400}, + {"GET", "/v1/snapshot/", 400}, + {"GET", "/v1//snapshot", 400}, + {"GET", "/v1/%73napshot", 400}, + {"GET", "/healthz", 400}, + {"GET", "/v1/snapshot?after=0", 409}, + {"GET", "/v1/snapshot?after=999", 503}, + {"GET", "/v1/snapshot?after=01", 400}, + {"GET", "/v1/snapshot?after=+1", 400}, + {"GET", "/v1/snapshot?after=%31", 400}, + {"GET", "/v1/snapshot?after=1&after=1", 400}, + {"GET", "/v1/snapshot?other=1", 400}, + {"GET", "/v1/snapshot?", 400}, + } { + t.Run(tc.path+tc.method, func(t *testing.T) { + r, err := http.NewRequestWithContext(f.ctx, tc.method, endpoint+tc.path, nil) + if err != nil { + t.Fatal(err) + } + + response, err := peer.Do(r) + responseBody(t, response, err, tc.status) + }) + } +} + +func TestAuthenticatedSharesProposalAndExplicitNodePrecedence(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + f.request.Shares = 9 + + req := bootstrapTestRequest(t, f.ctx, endpoint, f.token, f.request) + response, err := f.client(t, nil).Do(req) + responseBody(t, response, err, 200) + + node := &corev1.Node{} + if err := f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, node); err != nil { + t.Fatal(err) + } + + if node.Annotations[enrolledSharesAnnotation] != "9" { + t.Fatal("authenticated proposal not persisted") + } + + published := reconcileTopology(t, f.a.Topology, f.ctx) + + publication, err := wire.DecodePublication(strings.NewReader(published.encoded)) + if err != nil || publication.Members[0].Shares != 9 { + t.Fatalf("proposal not published: %+v %v", publication, err) + } + + if err := f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, node); err != nil { + t.Fatal(err) + } + + node.Annotations[wire.SharesAnnotation] = "12" + if err := f.a.Topology.Update(f.ctx, node); err != nil { + t.Fatal(err) + } + + published = reconcileTopology(t, f.a.Topology, f.ctx) + + publication, err = wire.DecodePublication(strings.NewReader(published.encoded)) + if err != nil || publication.Members[0].Shares != 12 { + t.Fatalf("explicit shares lost: %+v %v", publication, err) + } +} + +func TestHTTPSBootstrapBoundsAndErrors(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + c := f.client(t, nil) + + valid, err := wire.EncodeBootstrapRequest(f.request) + if err != nil { + t.Fatal(err) + } + + for _, tc := range []struct { + name, body, token, media string + status int + }{ + {"no token", string(valid), "", "application/json", 401}, + {"duplicate fields", `{"schema_version":1,"schema_version":1}`, f.token, "application/json", 400}, + {"unsupported", strings.Replace(string(valid), `"schema_version":1`, `"schema_version":2`, 1), f.token, "application/json", 426}, + {"large", strings.Repeat(" ", wire.MaxBootstrapBytes+1), f.token, "application/json", 413}, + {"media", string(valid), f.token, "text/plain", 400}, + {"cluster", strings.Replace(string(valid), string(f.request.Cluster), testNodeUID, 1), f.token, "application/json", 403}, + } { + t.Run(tc.name, func(t *testing.T) { + r, err := http.NewRequestWithContext(f.ctx, "POST", endpoint+wire.BootstrapPath, strings.NewReader(tc.body)) + if err != nil { + t.Fatal(err) + } + + if tc.token != "" { + r.Header.Set("Authorization", "Bearer "+tc.token) + } + + r.Header.Set("Content-Type", tc.media) + response, err := c.Do(r) + responseBody(t, response, err, tc.status) + }) + } +} + +func TestHTTPSCertificateRejectionAndRecovery(t *testing.T) { + f := newServingFixture(t) + + endpoint := f.start(t) + for name, tc := range map[string]struct { + mutate func(*x509.Certificate) + status int + allowTLSRejection bool + }{ + "expired": {func(c *x509.Certificate) { c.NotAfter = time.Now().Add(-time.Second) }, http.StatusUnauthorized, true}, + "future": {func(c *x509.Certificate) { c.NotBefore = time.Now().Add(time.Minute) }, http.StatusUnauthorized, true}, + "wrong usage": {func(c *x509.Certificate) { c.ExtKeyUsage = []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth} }, http.StatusUnauthorized, true}, + "wrong cluster": {func(c *x509.Certificate) { c.URIs[0].Host = testNodeUID }, http.StatusForbidden, false}, + // A signed, well-formed UID need not appear in routing membership. + "wrong uid": {func(c *x509.Certificate) { c.URIs[0].Path = "/node/" + testOtherUID }, http.StatusOK, false}, + "bad identity": {func(c *x509.Certificate) { c.URIs[0].Path = "/node/not-a-uuid" }, http.StatusUnauthorized, false}, + "ambiguous SAN": {func(c *x509.Certificate) { c.URIs = append(c.URIs, c.URIs[0]) }, http.StatusUnauthorized, false}, + "query SAN": {func(c *x509.Certificate) { c.URIs[0].RawQuery = "admin=true" }, http.StatusUnauthorized, false}, + "no signing usage": {func(c *x509.Certificate) { c.KeyUsage = x509.KeyUsageKeyEncipherment }, http.StatusUnauthorized, false}, + } { + t.Run(name, func(t *testing.T) { + cert := f.signLeaf(t, tc.mutate) + + response, err := f.client(t, &cert).Get(endpoint + wire.SnapshotPath) + if err != nil && tc.allowTLSRejection && strings.Contains(err.Error(), "remote error: tls:") { + return + } + + if tc.status == http.StatusOK { + responseBody(t, response, err, tc.status) + return + } + // Require only the expected wire error, with no snapshot bytes admitted. + want := map[int]string{ + http.StatusUnauthorized: `{"code":"unauthenticated"}`, + http.StatusForbidden: `{"code":"forbidden"}`, + http.StatusServiceUnavailable: `{"code":"unavailable"}`, + }[tc.status] + if body := responseBody(t, response, err, tc.status); string(body) != want { + t.Fatalf("rejection response: %s, want %s", body, want) + } + }) + } + // Expired identities can recover by omitting the certificate entirely. + req := bootstrapTestRequest(t, f.ctx, endpoint, f.token, f.request) + response, err := f.client(t, nil).Do(req) + + enrollment, err := wire.DecodeBootstrapResponse(bytes.NewReader(responseBody(t, response, err, http.StatusOK))) + if err != nil { + t.Fatal(err) + } + + if enrollment.Node != wire.NodeID(testNodeUID) { + t.Fatalf("recovered Node %s, want %s", enrollment.Node, testNodeUID) + } + + cert := tls.Certificate{Certificate: enrollment.CertificateChain, PrivateKey: f.key} + response, err = f.client(t, &cert).Get(endpoint + wire.SnapshotPath) + + publication, err := wire.DecodePublication(bytes.NewReader(responseBody(t, response, err, http.StatusOK))) + if err != nil || len(publication.Members) != 1 { + t.Fatalf("recovered snapshot: %v", err) + } +} + +func TestPooledTLSIgnoresWorkloadChangesButRejectsExpiry(t *testing.T) { + for _, scenario := range []string{"excluded", "recreated node", "pod gone", "expired"} { + t.Run(scenario, func(t *testing.T) { + f := newServingFixture(t) + + cert := f.certificate + if scenario == "expired" { + cert = f.signLeaf(t, func(c *x509.Certificate) { c.NotAfter = time.Now().Add(2 * time.Second).Truncate(time.Second) }) + } + + endpoint := f.start(t) + c := f.client(t, &cert) + response, err := c.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, 200) + + switch scenario { + case "excluded", "recreated node": + node := &corev1.Node{} + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, node)) + + if scenario == "excluded" { + node.Labels = map[string]string{wire.ExclusionLabel: ""} + } else { + node.UID = "replacement" + } + + require.NoError(t, f.a.Topology.Update(f.ctx, node)) + case "pod gone": + pod := &corev1.Pod{} + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Namespace: "racer", Name: "worker-pod"}, pod)) + require.NoError(t, f.a.Topology.Delete(f.ctx, pod)) + case "expired": + leaf, err := x509.ParseCertificate(cert.Certificate[0]) + require.NoError(t, err) + + time.Sleep(time.Until(leaf.NotAfter) + 10*time.Millisecond) + } + + reused := false + ctx := httptrace.WithClientTrace(f.ctx, &httptrace.ClientTrace{GotConn: func(info httptrace.GotConnInfo) { reused = info.Reused }}) + + req, err := http.NewRequestWithContext(ctx, "GET", endpoint+wire.SnapshotPath, nil) + require.NoError(t, err) + + response, err = c.Do(req) + + want := 200 + + switch scenario { + case "expired": + want = 401 + } + + responseBody(t, response, err, want) + + require.True(t, reused, "test failed to reuse TLS connection") + }) + } +} + +func TestPooledTLSRetiredTrustAndNoResumption(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + c := f.client(t, &f.certificate) + response, err := c.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, 200) + // Install a coherent post-retirement credential state while the original + // leaf is still valid. This isolates trust revocation from leaf expiration. + shared, bundle, state, material := keyState(t, f.a.Keyring) + + der, key, err := generateIssuer(time.Now().UTC().Truncate(time.Second), f.a.Keyring.Config) + if err != nil { + t.Fatal(err) + } + + id := rootID(der) + material.Keys = map[string]signingMaterial{id: {Certificate: der, PrivateKey: key}} + + shared.Data["issuer.json"], err = json.Marshal(material) + if err != nil { + t.Fatal(err) + } + + bundle.PeerTrustRoots = [][]byte{der} + bundle.Generation++ + state.ActiveIssuer = id + + shared.Data["bundle.json"], err = wire.EncodeBundle(bundle) + if err != nil { + t.Fatal(err) + } + + shared.Data["rotation.json"], err = json.Marshal(state) + if err != nil { + t.Fatal(err) + } + + if err := f.a.Topology.Update(f.ctx, shared); err != nil { + t.Fatal(err) + } + + runKeys(t, f.a.Keyring) + + reused := false + ctx := httptrace.WithClientTrace(f.ctx, &httptrace.ClientTrace{GotConn: func(info httptrace.GotConnInfo) { reused = info.Reused }}) + + req, err := http.NewRequestWithContext(ctx, "GET", endpoint+wire.SnapshotPath, nil) + if err != nil { + t.Fatal(err) + } + + response, err = c.Do(req) + responseBody(t, response, err, 401) + + if !reused { + t.Fatal("trust rotation did not use pooled connection") + } + + c.Transport.(*http.Transport).CloseIdleConnections() + + response, err = c.Get(endpoint + wire.SnapshotPath) + if err == nil { + response.Body.Close() + t.Fatal("new TLS connection accepted retired root") + } + // A fresh identity reconnects successfully but cannot resume an old session. + encoded, err := f.a.authority.Issue(f.ctx, fixtureIdentity(t, f), f.request) + if err != nil { + t.Fatal(err) + } + + responseChain := decodeIssuedResponse(t, encoded) + + fresh := tls.Certificate{Certificate: responseChain.CertificateChain, PrivateKey: f.key} + + freshClient := f.client(t, &fresh) + for range 2 { + response, err = freshClient.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, 200) + + if response.TLS.DidResume { + t.Fatal("TLS resumed authentication") + } + + freshClient.Transport.(*http.Transport).CloseIdleConnections() + } +} + +func keyringRequest(t *testing.T, f *servingFixture, bearer bool, query string) *http.Request { + t.Helper() + + r := httptest.NewRequestWithContext(f.ctx, http.MethodGet, wire.KeyringPath+query, nil) + if bearer { + r.TLS = &tls.ConnectionState{HandshakeComplete: true} + r.Header.Set("Authorization", "Bearer "+f.token) + } else { + r.TLS = f.requestState(t) + } + + return r +} + +func requireKeyringResponse(t *testing.T, w *httptest.ResponseRecorder, status int) []byte { + t.Helper() + + body := responseBody(t, w.Result(), nil, status) + if w.Header().Get("Cache-Control") != "no-store" || !w.Flushed { + t.Fatal("response not flushed or cacheable") + } + + if len(body) > wire.MaxBundleBytes { + t.Fatal("unbounded keyring response") + } + + return body +} + +func TestHTTPSKeyringAuthenticationAndEncoding(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + _, bundle, _, material := keyState(t, f.a.Keyring) + + expected, err := wire.EncodeBundle(bundle) + if err != nil { + t.Fatal(err) + } + + for _, bearer := range []bool{true, false} { + r, err := http.NewRequestWithContext(f.ctx, http.MethodGet, endpoint+wire.KeyringPath, nil) + if err != nil { + t.Fatal(err) + } + + peer := f.client(t, &f.certificate) + if bearer { + peer = f.client(t, nil) + r.Header.Set("Authorization", "Bearer "+f.token) + } + + response, err := peer.Do(r) + + body := responseBody(t, response, err, http.StatusOK) + if !bytes.Equal(body, expected) || response.Header.Get("Cache-Control") != "no-store" || response.Header.Get("Content-Type") != "application/json" { + t.Fatal("response was not the full committed wire bundle") + } + + for _, key := range material.Keys { + if bytes.Contains(body, []byte(base64.StdEncoding.EncodeToString(key.PrivateKey))) { + t.Fatal("issuer private key disclosed") + } + } + } +} + +func TestKeyringRequestValidation(t *testing.T) { + f := newServingFixture(t) + handler := f.a.Server.Handler() + + for _, tc := range []struct { + name string + mutate func(*http.Request) + status int + }{ + {"anonymous", func(r *http.Request) { r.TLS = &tls.ConnectionState{HandshakeComplete: true} }, 401}, + {"no TLS", func(r *http.Request) { r.TLS = nil }, 401}, + {"both credentials", func(r *http.Request) { r.Header.Set("Authorization", "Bearer "+f.token) }, 401}, + {"empty auth with certificate", func(r *http.Request) { r.Header["Authorization"] = []string{""} }, 401}, + {"duplicate bearer", func(r *http.Request) { + r.TLS = &tls.ConnectionState{HandshakeComplete: true} + r.Header["Authorization"] = []string{"Bearer " + f.token, "Bearer " + f.token} + }, 401}, + {"bad bearer", func(r *http.Request) { + r.TLS = &tls.ConnectionState{HandshakeComplete: true} + r.Header.Set("Authorization", "Basic abc") + }, 401}, + {"unverified certificate", func(r *http.Request) { r.TLS.VerifiedChains = nil }, 401}, + {"body", func(r *http.Request) { r.ContentLength = 1 }, 400}, + {"unknown body length", func(r *http.Request) { r.ContentLength = -1 }, 400}, + {"chunked", func(r *http.Request) { r.TransferEncoding = []string{"chunked"} }, 400}, + {"encoding", func(r *http.Request) { r.Header.Set("Content-Encoding", "gzip") }, 400}, + {"empty query", func(r *http.Request) { r.URL.ForceQuery = true }, 400}, + {"post", func(r *http.Request) { r.Method = http.MethodPost }, 400}, + {"head", func(r *http.Request) { r.Method = http.MethodHead }, 400}, + {"slash", func(r *http.Request) { r.URL.Path += "/" }, 400}, + {"escaped path", func(r *http.Request) { r.URL.RawPath = "/v1/%6beyring" }, 400}, + {"header limit", func(r *http.Request) { + r.Header.Set("X-Large", strings.Repeat("x", f.a.Server.config.Limits.HeaderBytes)) + }, 413}, + } { + t.Run(tc.name, func(t *testing.T) { + r := keyringRequest(t, f, false, "") + tc.mutate(r) + + w := httptest.NewRecorder() + handler.ServeHTTP(w, r) + requireKeyringResponse(t, w, tc.status) + }) + } + + for _, query := range []string{"after=0", "after=999", "after=01", "after=+1", "after=%31", "after=1&after=1", "after=1&x=2", "other=1", "after=", "after=-1", "after=18446744073709551616"} { + t.Run(query, func(t *testing.T) { + w := httptest.NewRecorder() + handler.ServeHTTP(w, keyringRequest(t, f, false, "?"+query)) + + status := 400 + if query == "after=0" || query == "after=999" { + status = 409 + } + + requireKeyringResponse(t, w, status) + }) + } +} + +func TestKeyringBearerLiveBindings(t *testing.T) { + for _, scenario := range []string{"token", "audience", "pod", "service account", "daemonset", "node", "API outage"} { + t.Run(scenario, func(t *testing.T) { + f := newServingFixture(t) + status := http.StatusForbidden + + var obj client.Object + + key := client.ObjectKey{Namespace: "racer", Name: "racer-dataplane"} + + switch scenario { + case "token", "audience": + status = http.StatusUnauthorized + fixtureDependencies[f.a.authority].Client = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Create: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.CreateOption) error { + review := obj.(*authv1.TokenReview) + review.Status.Authenticated = scenario == "audience" + review.Status.Audiences = []string{"wrong-audience"} + + return nil + }}) + case "pod": + obj, key.Name = &corev1.Pod{}, "worker-pod" + case "service account": + obj = &corev1.ServiceAccount{} + case "daemonset": + obj = &appsv1.DaemonSet{} + case "node": + obj, key = &corev1.Node{}, client.ObjectKey{Name: "worker"} + case "API outage": + status = http.StatusServiceUnavailable + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + return errors.New("offline") + }}) + } + + if obj != nil { + if err := f.a.Topology.Get(f.ctx, key, obj); err != nil { + t.Fatal(err) + } + + if scenario == "node" { + obj.SetLabels(map[string]string{wire.ExclusionLabel: ""}) + } else { + obj.SetUID("recreated") + } + + if err := f.a.Topology.Update(f.ctx, obj); err != nil { + t.Fatal(err) + } + } + + w := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(w, keyringRequest(t, f, true, "")) + requireKeyringResponse(t, w, status) + }) + } +} + +func TestKeyringMTLSNeverReadsAPI(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + + var calls atomic.Int64 + + unavailable := interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + calls.Add(1) + return errors.New("offline") + }, + Create: func(context.Context, client.WithWatch, client.Object, ...client.CreateOption) error { + calls.Add(1) + return errors.New("offline") + }, + }) + fixtureDependencies[f.a.authority].Client, fixtureDependencies[f.a.authority].reader = unavailable, unavailable + + if _, err := f.a.Keyring.Reconcile(f.ctx, ctrl.Request{}); err == nil { + t.Fatal("outage hidden") + } + + calls.Store(0) + + handler := f.a.Server.Handler() + + for _, query := range []string{"", "?after=1"} { + w := httptest.NewRecorder() + handler.ServeHTTP(w, keyringRequest(t, f, false, query)) + + want := 200 + if query != "" { + // Without background confirmation the replica expires during the poll. + want = 503 + } + + body := requireKeyringResponse(t, w, want) + if want == 204 && len(body) != 0 { + t.Fatal("204 body") + } + } + + if calls.Load() != 0 { + t.Fatal("mTLS keyring used API") + } + }) +} + +func TestKeyringPollWakeAndTermination(t *testing.T) { + for _, bearer := range []bool{false, true} { + for _, scenario := range []string{"timeout", "rotation", "invalidation", "leader canceled", "request canceled", "expired", "bearer revoked", "retired trust"} { + t.Run(fmt.Sprintf("bearer=%t/%s", bearer, scenario), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { testKeyringPollTermination(t, bearer, scenario) }) + }) + } + } +} + +func testKeyringPollTermination(t *testing.T, bearer bool, scenario string) { + t.Helper() + f := newServingFixture(t) + // Isolate poll termination from the default 30-second freshness gate. + configureFixtureAge(t, f, time.Minute) + + if scenario == "expired" { + f.expireIdentitySoon(t, bearer) + } + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + r := keyringRequest(t, f, bearer, "?after=1").WithContext(ctx) + handler := f.a.Server.Handler() + w := httptest.NewRecorder() + done := make(chan struct{}) + + go func() { defer close(done); handler.ServeHTTP(w, r) }() + + synctest.Wait() + + polls := f.a.Server.keyringPolls.count() + + require.Equal(t, 1, polls) + require.Empty(t, f.a.Server.writes) + require.Empty(t, f.a.Server.bootstrapSlots) + + want := terminateKeyringPoll(t, f, cancel, bearer, scenario) + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("poll did not wake") + } + + body := requireKeyringResponse(t, w, want) + if want == 200 { + bundle, err := wire.DecodeBundle(bytes.NewReader(body)) + require.NoError(t, err) + require.Equal(t, wire.Generation(2), bundle.Generation) + } + + require.Zero(t, f.a.Server.keyringPolls.count(), "admission leaked") +} + +func terminateKeyringPoll(t *testing.T, f *servingFixture, cancel context.CancelFunc, bearer bool, scenario string) int { + t.Helper() + + want := 503 + + switch scenario { + case "timeout": + want = 204 + // Repeated unchanged reconciliations must not extend the wait. + time.Sleep(20 * time.Second) + runKeys(t, f.a.Keyring) + time.Sleep(10 * time.Second) + case "expired": + want = 401 + + time.Sleep(time.Second) + case "rotation": + want = 200 + _, _, rotation, _ := keyState(t, f.a.Keyring) + fixtureDependencies[f.a.authority].now = func() time.Time { return rotation.NextRotation } + runKeys(t, f.a.Keyring) + case "invalidation": + invalidateFixtureTrust(t, f) + case "leader canceled": + f.cancel() + case "request canceled": + cancel() + case "bearer revoked": + want = 204 + if bearer { + want = 403 + } + + pod := &corev1.Pod{} + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Namespace: "racer", Name: "worker-pod"}, pod)) + require.NoError(t, f.a.Topology.Delete(f.ctx, pod)) + + time.Sleep(wire.PollWait) + case "retired trust": + want = 401 + if bearer { + want = 200 + } + + replaceFixtureCredentials(t, f) + } + + return want +} + +func (f *servingFixture) expireIdentitySoon(t *testing.T, bearer bool) { + t.Helper() + + if !bearer { + f.certificate = f.signLeaf(t, func(c *x509.Certificate) { c.NotAfter = time.Now().Add(time.Second) }) + return + } + + _, status, _ := authFixture(t) + f.token = "header." + base64.RawURLEncoding.EncodeToString(fmt.Appendf(nil, `{"exp":%d}`, time.Now().Add(time.Second).Unix())) + ".signature" + installReview(t, f.a, status, f.token) +} + +func TestKeyringAdmissionHeldThroughResponse(t *testing.T) { + for _, flush := range []bool{false, true} { + for _, status := range []int{200, 204, 409} { + t.Run(fmt.Sprintf("flush=%t/status=%d", flush, status), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxPolls = 1 }) + // Keep authority fresh through the 204 wait and blocked response. + configureFixtureAge(t, f, time.Minute) + handler := f.a.Server.Handler() + + query := "" + if status == 204 { + query = "?after=1" + } + + if status == 409 { + query = "?after=2" + } + + w := &blockingResponse{ResponseRecorder: httptest.NewRecorder(), entered: make(chan struct{}), unblock: make(chan struct{}), blockFlush: flush} + // A 204 has no Write, only Flush. + if status == 204 { + w.blockFlush = true + } + + unblock := sync.OnceFunc(func() { close(w.unblock) }) + defer unblock() + + done := make(chan struct{}) + r := keyringRequest(t, f, false, query) + + go func() { defer close(done); handler.ServeHTTP(w, r) }() + + <-w.entered + + other := *f + + other.certificate = f.signLeaf(t, func(c *x509.Certificate) { c.URIs[0].Path = "/node/" + testOtherUID }) + for _, fixture := range []*servingFixture{f, &other} { + second := httptest.NewRecorder() + handler.ServeHTTP(second, keyringRequest(t, fixture, false, "")) + requireKeyringResponse(t, second, 429) + } + // Independent snapshot admission is available while delivery blocks. + require.True(t, f.a.Server.polls.acquire(wire.NodeID(testNodeUID)), "keyring consumed snapshot admission") + + f.a.Server.polls.release(wire.NodeID(testNodeUID)) + unblock() + <-done + requireKeyringResponse(t, w.ResponseRecorder, status) + + requireNoAdmission(t, f.a.Server) + }) + }) + } + } +} + +func TestKeyringBearerAdmissionAndDeadline(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { + c.Limits.MaxConcurrentBootstrap = 1 + c.Limits.WriteTimeout = time.Second + }) + s := f.a.Server + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, _ client.WithWatch, _ client.ObjectKey, _ client.Object, _ ...client.GetOption) error { + <-ctx.Done() + return ctx.Err() + }}) + handler := s.Handler() + w := httptest.NewRecorder() + r := keyringRequest(t, f, true, "") + done := make(chan struct{}) + + go func() { defer close(done); handler.ServeHTTP(w, r) }() + + synctest.Wait() + + if len(s.bootstrapSlots) != 1 { + t.Fatal("bearer API work not admitted") + } + + second := httptest.NewRecorder() + handler.ServeHTTP(second, keyringRequest(t, f, true, "")) + requireKeyringResponse(t, second, 429) + + local := httptest.NewRecorder() + handler.ServeHTTP(local, keyringRequest(t, f, false, "")) + requireKeyringResponse(t, local, 200) + time.Sleep(time.Second) + <-done + requireKeyringResponse(t, w, 503) + + if len(s.bootstrapSlots) != 0 || s.keyringPolls.count() != 0 { + t.Fatal("API timeout leaked admission") + } + }) +} + +func TestKeyringPerNodeAdmissionIndependentOfSnapshots(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxPolls = 2 }) + handler := f.a.Server.Handler() + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + request := keyringRequest(t, f, false, "?after=1").WithContext(ctx) + done := make(chan struct{}) + + go func() { defer close(done); handler.ServeHTTP(httptest.NewRecorder(), request) }() + + synctest.Wait() + + for _, bearer := range []bool{true, false} { + w := httptest.NewRecorder() + handler.ServeHTTP(w, keyringRequest(t, f, bearer, "")) + requireKeyringResponse(t, w, 429) + } + + other := *f + other.certificate = f.signLeaf(t, func(c *x509.Certificate) { c.URIs[0].Path = "/node/" + testOtherUID }) + w := httptest.NewRecorder() + handler.ServeHTTP(w, keyringRequest(t, &other, false, "")) + requireKeyringResponse(t, w, 200) + w = httptest.NewRecorder() + snapshot := keyringRequest(t, f, false, "") + snapshot.URL.Path = wire.SnapshotPath + handler.ServeHTTP(w, snapshot) + responseBody(t, w.Result(), nil, 200) + cancel() + <-done + }) +} + +func TestKeyringBearerExpiresDuringReauthentication(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + f.expireIdentitySoon(t, true) + + reads := 0 + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if _, ok := obj.(*corev1.Pod); ok { + reads++ + if reads == 2 { + <-ctx.Done() + return ctx.Err() + } + } + + return c.Get(ctx, key, obj, opts...) + }}) + w := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(w, keyringRequest(t, f, true, "")) + requireKeyringResponse(t, w, 401) + + if reads != 2 { + t.Fatal("did not exercise post-wait expiration") + } + }) +} + +// TestRustKeyringInterop runs the production Rust transport against the real Go +// HTTPS handler, issuer, and keyring reconciler with the existing fake API fixture. +// Opt in with RACER_RUST_INTEROP=1; no Kubernetes cluster is contacted. +func TestRustKeyringInterop(t *testing.T) { + if os.Getenv("RACER_RUST_INTEROP") != "1" { + t.Skip("set RACER_RUST_INTEROP=1 to run the Rust client") + } + + root, err := filepath.Abs("../..") + require.NoError(t, err) + require.NoError(t, os.MkdirAll(filepath.Join(root, "tmp"), 0o700)) + directory, err := os.MkdirTemp(filepath.Join(root, "tmp"), "keyring-interop-") + require.NoError(t, err) + + t.Cleanup(func() { _ = os.RemoveAll(directory) }) + f := newServingFixture(t) + + cache := catalogCache("interop", testOtherUID) + require.NoError(t, f.a.Topology.Create(f.ctx, &cache)) + + runKeys(t, f.a.Keyring) + reconcileTopology(t, f.a.Topology, f.ctx) + endpoint := f.start(t) + + writeInteropConfig(t, f, directory, endpoint) + + ctx, cancel := context.WithTimeout(t.Context(), 240*time.Second) + defer cancel() + + cmd := exec.CommandContext(ctx, "timeout", "--signal=TERM", "--kill-after=10s", "230s", "cargo", "test", "--locked", "--manifest-path", filepath.Join(root, "cmd/racer-dataplane/Cargo.toml"), "--test", "keyring_interop", "--", "--ignored", "--nocapture") + + cmd.Env = append(os.Environ(), "RACER_KEYRING_INTEROP_DIR="+directory) + cmd.Stdout, cmd.Stderr = os.Stdout, os.Stderr + cmd.Cancel = func() error { return cmd.Process.Signal(syscall.SIGTERM) } + + cmd.WaitDelay = 10 * time.Second + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + + done := make(chan error, 1) + + go func() { done <- cmd.Wait() }() + + reaped := false + + defer func() { + if !reaped { + cancel() + <-done + } + }() + + ticker := time.NewTicker(10 * time.Millisecond) + defer ticker.Stop() + + // No manager runs in this fixture. Keep authority observations fresh through + // Cargo startup and the full no-change poll without relaxing freshness gates. + refresh := time.NewTicker(time.Second) + defer refresh.Stop() + + rotated := false + + for { + select { + case err := <-done: + reaped = true + + require.NoError(t, err, "Rust interoperability client") + require.True(t, rotated, "Rust client did not reach the rotation poll") + + return + case <-refresh.C: + runKeys(t, f.a.Keyring) + reconcileTopology(t, f.a.Topology, f.ctx) + case <-ticker.C: + if rotated { + continue + } + + if _, err := os.Stat(filepath.Join(directory, "rotate")); err != nil { + continue + } + + parked := f.a.Server.keyringPolls.count() == 1 + + if parked { + _, _, rotation, _ := keyState(t, f.a.Keyring) + fixtureDependencies[f.a.authority].now = func() time.Time { return rotation.NextRotation } + runKeys(t, f.a.Keyring) + + rotated = true + } + } + } +} + +func writeInteropConfig(t *testing.T, f *servingFixture, directory, endpoint string) { + t.Helper() + + trust := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: f.serverCertificate.Certificate[0]}) + require.NoError(t, os.WriteFile(filepath.Join(directory, "trust.pem"), trust, 0o600)) + + config, err := json.Marshal(map[string]string{ + "endpoint": endpoint, "cluster": string(f.a.Topology.Config.Cluster), + "node": testNodeUID, "token": f.token, + }) + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(directory, "config.json"), config, 0o600)) +} + +func TestAuthenticatedRDMANICProposalAndEmptyRemoval(t *testing.T) { + f := newServingFixture(t) + f.request.Shares = 9 + f.request.RDMANICs = []wire.RDMANIC{{Device: "mlx5_1", Port: 1, Rail: 0}, {Device: "mlx5_0", Port: 2, Rail: 0}} + req, err := http.NewRequestWithContext(f.ctx, http.MethodPost, wire.BootstrapPath, nil) + require.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+f.token) + _, err = f.a.Server.enroll(f.ctx, req, f.request) + require.NoError(t, err) + + var node corev1.Node + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + require.Equal(t, "9", node.Annotations[enrolledSharesAnnotation]) + nics, err := wire.DecodeRDMANICs(strings.NewReader(node.Annotations[enrolledRDMANICsAnnotation])) + require.NoError(t, err) + require.Equal(t, wire.CanonicalRDMANICs(f.request.RDMANICs), nics) + + node.Annotations[wire.RDMANICsAnnotation] = "[]" + require.NoError(t, f.a.Topology.Update(f.ctx, &node)) + attributes, err := members.ParseAnnotations(&node) + require.NoError(t, err) + require.Empty(t, attributes.RDMANICs) + f.request.RDMANICs = nil + f.request.Shares = 10 + _, err = f.a.Server.enroll(f.ctx, req, f.request) + require.NoError(t, err) + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + require.NotContains(t, node.Annotations, enrolledRDMANICsAnnotation) + require.Equal(t, "10", node.Annotations[enrolledSharesAnnotation]) + require.Equal(t, "[]", node.Annotations[wire.RDMANICsAnnotation]) +} + +func TestEnrollmentRDMANICAtomicPatchFailure(t *testing.T) { + f := newServingFixture(t) + f.request.Shares = 9 + f.request.RDMANICs = []wire.RDMANIC{{Device: "mlx5_0", Port: 1}} + patchFailure := errors.New("patch rejected") + patched := false + fixtureDependencies[f.a.authority].Client = interceptor.NewClient(fixtureDependencies[f.a.authority].Client.(client.WithWatch), interceptor.Funcs{ + Patch: func(_ context.Context, _ client.WithWatch, obj client.Object, patch client.Patch, _ ...client.PatchOption) error { + patched = true + data, err := patch.Data(obj) + require.NoError(t, err) + require.Contains(t, string(data), enrolledSharesAnnotation) + require.Contains(t, string(data), enrolledRDMANICsAnnotation) + require.Contains(t, string(data), "resourceVersion") + + return patchFailure + }, + }) + req, err := http.NewRequestWithContext(f.ctx, http.MethodPost, wire.BootstrapPath, nil) + require.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+f.token) + _, err = f.a.Server.enroll(f.ctx, req, f.request) + require.ErrorIs(t, err, patchFailure) + require.True(t, patched) + + var node corev1.Node + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + require.NotContains(t, node.Annotations, enrolledSharesAnnotation) + require.NotContains(t, node.Annotations, enrolledRDMANICsAnnotation) +} + +func TestEnrollmentRDMANICLiveNodeRecheck(t *testing.T) { + for _, scenario := range []string{"replacement UID", "excluded"} { + t.Run(scenario, func(t *testing.T) { + f := newServingFixture(t) + f.request.RDMANICs = []wire.RDMANIC{{Device: "mlx5_0", Port: 1}} + nodeReads := 0 + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if err := c.Get(ctx, key, obj, opts...); err != nil { + return err + } + + if node, ok := obj.(*corev1.Node); ok { + nodeReads++ + if nodeReads > 1 { + if scenario == "replacement UID" { + node.UID = testOtherUID + } else { + node.Labels = map[string]string{wire.ExclusionLabel: ""} + } + } + } + + return nil + }, + }) + req, err := http.NewRequestWithContext(f.ctx, http.MethodPost, wire.BootstrapPath, nil) + require.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+f.token) + _, err = f.a.Server.enroll(f.ctx, req, f.request) + require.ErrorIs(t, err, wire.Forbidden) + require.Equal(t, 2, nodeReads) + + var node corev1.Node + require.NoError(t, f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, &node)) + require.NotContains(t, node.Annotations, enrolledSharesAnnotation) + require.NotContains(t, node.Annotations, enrolledRDMANICsAnnotation) + }) + } +} + +func TestLocalTrustInvalidationDuringPoll(t *testing.T) { + f := newServingFixture(t) + + current, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("%s?after=%d", wire.SnapshotPath, current.Sequence()), nil) + req.TLS = f.requestState(t) + handler := f.a.Server.Handler() + done := make(chan *httptest.ResponseRecorder, 1) + + go func() { + w := httptest.NewRecorder() + handler.ServeHTTP(w, req) + + done <- w + }() + + deadline := time.After(5 * time.Second) + + for { + n := f.a.Server.polls.count() + + if n == 1 { + break + } + + select { + case <-deadline: + t.Fatal("poll did not park") + default: + time.Sleep(time.Millisecond) + } + } + + withdrawServerTrust(t, f.a.Server) + + node := &corev1.Node{} + if err := f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, node); err != nil { + t.Fatal(err) + } + + node.Labels = map[string]string{wire.ExclusionLabel: ""} + if err := f.a.Topology.Update(f.ctx, node); err != nil { + t.Fatal(err) + } + + _, _ = f.a.Topology.Reconcile(f.ctx, ctrl.Request{}) + + select { + case w := <-done: + if w.Code != http.StatusServiceUnavailable || w.Body.String() != `{"code":"unavailable"}` { + t.Fatalf("invalidated trust disclosed publication: %d %s", w.Code, w.Body.String()) + } + case <-deadline: + t.Fatal("poll did not recheck local trust") + } +} + +func TestKeyringReauthenticationCannotDiscloseWithdrawnBundle(t *testing.T) { + f := newServingFixture(t) + reads := 0 + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + if _, ok := obj.(*corev1.Pod); ok { + reads++ + if reads == 2 { + invalidateFixtureTrust(t, f) + } + } + + return c.Get(ctx, key, obj, opts...) + }}) + w := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(w, keyringRequest(t, f, true, "")) + requireKeyringResponse(t, w, 503) + + if reads != 2 { + t.Fatal("bearer not rechecked before response") + } +} + +type completionResponse struct { + *httptest.ResponseRecorder + onWrite func() + onFlush func() +} + +func (w *completionResponse) Write(p []byte) (int, error) { + n, err := w.ResponseRecorder.Write(p) + if w.onWrite != nil { + w.onWrite() + } + + return n, err +} + +func (w *completionResponse) Flush() { + w.ResponseRecorder.Flush() + + if w.onFlush != nil { + w.onFlush() + } +} + +func TestTrustAuthorityImmediateCompletion(t *testing.T) { + for _, route := range []string{"bootstrap", "keyring", "keyring empty", "snapshot"} { + for _, stage := range []string{"write", "flush"} { + if route == "keyring empty" && stage == "write" { + continue + } + + for _, change := range []string{"invalidate", "recover", "rotate"} { + t.Run(route+"/"+stage+"/"+change, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { testTrustCompletion(t, route, stage, change) }) + }) + } + } + } +} + +func testTrustCompletion(t *testing.T, route, stage, change string) { + t.Helper() + f := newServingFixture(t) + s := f.a.Server + + request := f.publicRequest(t, route) + called := false + changeAuthority := func() { + called = true + + if change != "rotate" { + withdrawServerTrust(t, s) + } + + if change == "recover" { + restoreServerTrust(t, s) + } + + if change == "rotate" { + rotateFixtureTrust(t, f) + } + // Deliberately do not yield or wait for cancellation callbacks. + } + + w := &completionResponse{ResponseRecorder: httptest.NewRecorder()} + if stage == "write" { + w.onWrite = changeAuthority + } else { + w.onFlush = changeAuthority + } + + aborted := serveRecover(s.Handler(), w, request) + require.True(t, called, "completion hook not reached: %d", w.Code) + + if change == "rotate" { + require.Nil(t, aborted, "ordinary rotation aborted admitted response") + } else { + require.Equal(t, http.ErrAbortHandler, aborted, "revoked response completed") + } + + requireNoAdmission(t, s) +} + +func TestTrustAuthoritySynchronousRevocation(t *testing.T) { + f := newServingFixture(t) + + guard, cancel, err := f.a.authority.AdmitTrust(t.Context()) + if err != nil { + t.Fatal(err) + } + defer cancel() + + withdrawServerTrust(t, f.a.Server) + + if guard.Check(t.Context()) != context.Canceled { + t.Fatal("revocation depends on callback scheduling") + } +} + +func TestPublicBlockedWriteTrustAuthority(t *testing.T) { + for _, route := range []string{"snapshot", "bootstrap", "keyring"} { + for _, change := range []string{"invalidate", "invalidate recover", "expire", "reconfirm"} { + t.Run(route+"/"+change, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + s := f.a.Server + configureFixtureAge(t, f, 5*time.Second) + + server, peer := net.Pipe() + defer server.Close() + defer peer.Close() + + request := f.publicRequest(t, route) + request = request.WithContext(connectionContext(f.ctx, server)) + w := &pipeResponse{ResponseRecorder: httptest.NewRecorder(), conn: server} + done := make(chan any, 1) + + go func() { defer func() { done <- recover() }(); s.Handler().ServeHTTP(w, request) }() + + synctest.Wait() + + require.Len(t, s.writes, 1, "response not blocked in write") + + start := time.Now() + + switch change { + case "invalidate": + withdrawServerTrust(t, s) + case "invalidate recover": + withdrawServerTrust(t, s) + restoreServerTrust(t, s) + case "expire": + time.Sleep(3 * time.Second) + reconcileTopology(t, f.a.Topology, f.ctx) + time.Sleep(2 * time.Second) + case "reconfirm": + time.Sleep(3 * time.Second) + + _, err := f.a.authority.ReconcileCredentials(t.Context()) + require.NoError(t, err) + + reconcileTopology(t, f.a.Topology, f.ctx) + + time.Sleep(2 * time.Second) + } + + require.Equal(t, http.ErrAbortHandler, <-done, "blocked response did not abort") + + want := time.Duration(0) + if change == "expire" || change == "reconfirm" { + want = 5 * time.Second + } + + require.Equal(t, want, time.Since(start), "trust cancellation delay") + require.NoError(t, f.a.authority.PublicationReady(), "test must retain fresh publication independently of trust") + require.Empty(t, s.writes) + require.Empty(t, s.bootstrapSlots) + + if change == "invalidate recover" || change == "reconfirm" { + require.NoError(t, f.a.authority.TrustReady(), "new requests should have usable trust") + } + }) + }) + } + } +} + +func TestTrustRotationKeepsAdmittedResponseBounded(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + configureFixtureAge(t, f, 5*time.Second) + + guard, cancel, err := f.a.authority.AdmitTrust(t.Context()) + if err != nil { + t.Fatal(err) + } + defer cancel() + + ctx := guard.Context() + deadline, _ := ctx.Deadline() + + time.Sleep(3 * time.Second) + + rotateFixtureTrust(t, f) + + if ctx.Err() != nil { + t.Fatal("normal rotation revoked admitted response") + } + + if got, _ := ctx.Deadline(); got != deadline { + t.Fatal("rotation extended admitted freshness") + } + + time.Sleep(2 * time.Second) + synctest.Wait() + + if ctx.Err() == nil { + t.Fatal("admitted trust outlived pinned freshness") + } + + if err := f.a.authority.TrustReady(); err != nil { + t.Fatal("rotation should admit fresh requests", err) + } + }) +} + +func eventually(t *testing.T, description string, ready func() bool) { + t.Helper() + + deadline := time.Now().Add(10 * time.Second) + for !ready() { + if time.Now().After(deadline) { + t.Fatal(description) + } + + time.Sleep(time.Millisecond) + } +} + +func TestReplicationRouteAuthorizationAndEarlyListener(t *testing.T) { + f := newServingFixture(t) + r := f.a.Replication + r.Config.ControllerServiceAccount = "racer-controller" + r.leader = f.ctx + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: "controller", UID: "controller-uid"}, Spec: corev1.PodSpec{ServiceAccountName: "racer-controller"}} + + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: "racer-controller", UID: "controller-sa"}} + for _, obj := range []client.Object{pod, sa} { + require.NoError(t, r.Client.Create(f.ctx, obj)) + } + + username := "system:serviceaccount:" + r.Config.Namespace + ":racer-controller" + audience := ReplicationAudience + r.Client = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Create: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.CreateOption) error { + review := obj.(*authv1.TokenReview) + require.Equal(t, []string{ReplicationAudience}, review.Spec.Audiences, "wrong review audience") + + review.Status = authv1.TokenReviewStatus{Authenticated: true, Audiences: []string{audience}, User: authv1.UserInfo{Username: username, UID: string(sa.UID), Extra: map[string]authv1.ExtraValue{"authentication.kubernetes.io/pod-name": {pod.Name}, "authentication.kubernetes.io/pod-uid": {string(pod.UID)}}}} + + return nil + }}) + fixtureDependencies[r.authority].Client = r.Client + + for _, unchanged := range []bool{false, true} { + for _, fail := range []bool{false, true} { + t.Run(fmt.Sprintf("blocked flush unchanged=%v failure=%v", unchanged, fail), func(t *testing.T) { + testReplicationFlush(t, f, unchanged, fail) + }) + } + } + + for _, tc := range []struct { + name string + code int + }{{"controller", 200}, {"duplicate poll", 429}, {"dataplane", 403}, {"wrong audience", 401}} { + t.Run(tc.name, func(t *testing.T) { + if tc.name == "duplicate poll" { + require.True(t, f.a.Server.replicationPolls.acquire(string(pod.UID)), "could not reserve replication poll") + defer f.a.Server.replicationPolls.release(string(pod.UID)) + } + + if tc.name == "dataplane" { + username = "system:serviceaccount:" + r.Config.Namespace + ":racer-dataplane" + } + + if tc.name == "wrong audience" { + audience = wire.TokenAudience + } + + request := httptest.NewRequest(http.MethodGet, ReplicationPath, nil) + request.TLS = f.requestState(t) + request.Header.Set("Authorization", "Bearer "+f.token) + + response := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(response, request) + + require.Equal(t, tc.code, response.Code, response.Body.String()) + }) + } + // TLS must be reachable before public readiness, including before trust is + // initialized. Public routes remain unavailable; internal auth is independent. + invalidateFixtureTrust(t, f) + endpoint := f.start(t) + peer := f.client(t, nil) + response, err := peer.Get(endpoint + ReplicationPath) + responseBody(t, response, err, http.StatusUnauthorized) +} + +func TestSnapshotBlockedWriteClosesAtPinnedFreshness(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + configureFixtureAge(t, f, 5*time.Second) + + server, peer := net.Pipe() + defer server.Close() + defer peer.Close() + + request := httptest.NewRequest(http.MethodGet, wire.SnapshotPath, nil) + request.TLS = f.requestState(t) + request = request.WithContext(connectionContext(f.ctx, server)) + w := &pipeResponse{ResponseRecorder: httptest.NewRecorder(), conn: server} + done := make(chan any, 1) + + go func() { defer func() { done <- recover() }(); f.a.Server.Handler().ServeHTTP(w, request) }() + + synctest.Wait() + time.Sleep(3 * time.Second) + + if err := f.a.authority.Observe(f.ctx); err != nil { + t.Fatal(err) + } + + time.Sleep(2 * time.Second) + + if aborted := <-done; aborted != http.ErrAbortHandler { + t.Fatalf("blocked write did not abort: %v", aborted) + } + + if len(f.a.Server.writes) != 0 { + t.Fatal("blocked write retained admission") + } + + if f.a.authority.PublicationReady() != nil { + t.Fatal("confirmation should allow a new request") + } + }) +} + +type pipeResponse struct { + *httptest.ResponseRecorder + conn net.Conn +} + +func (w *pipeResponse) SetWriteDeadline(deadline time.Time) error { + return w.conn.SetWriteDeadline(deadline) +} + +func TestRevokedDeltaCannotBorrowNewAuthority(t *testing.T) { + r := initializedTopology(t) + membership := AcceptedMembers{} + + for i := range 100 { + id := wire.NodeID(fmt.Sprintf("22222222-2222-4222-8222-%012d", i)) + membership[id] = wire.Member{Node: id, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}} + } + + base := replicationSmokePublish(t, t.Context(), r, membership) + id := wire.NodeID("22222222-2222-4222-8222-000000000000") + member := membership[id] + member.Shares++ + membership[id] = member + p := replicationSmokePublish(t, t.Context(), r, membership) + delta := p.handle.ForBase(base.record.Sequence, base.record.ContentHash) + + guard, cancel, err := p.admit(t.Context()) + if err != nil { + t.Fatal(err) + } + defer cancel() + + var encoded bytes.Buffer + + _, err = delta.WriteTo(t.Context(), guard, &encoded) + require.NoError(t, err) + require.Less(t, encoded.Len(), len(p.encoded), "test must exercise the real delta response") + baseImage, err := wire.DecodePublication(strings.NewReader(base.encoded)) + require.NoError(t, err) + _, err = wire.ApplyDelta(baseImage, &encoded) + require.NoError(t, err) + + restore := withdrawPublication(t, r) + restore() + reconcileTopology(t, r, t.Context()) + + <-guard.Context().Done() + + if _, err := delta.WriteTo(t.Context(), guard, io.Discard); !errors.Is(err, context.Canceled) { + t.Fatalf("revoked delta: %v", err) + } + + if _, _, err := p.admit(t.Context()); !errors.Is(err, wire.Unavailable) { + t.Fatalf("old image borrowed new authority: %v", err) + } +} + +func TestSnapshotAuthorityHeldThroughFlush(t *testing.T) { + for _, action := range []string{"suspend", "freshness", "advance"} { + t.Run(action, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + configureFixtureAge(t, f, 5*time.Second) + + w := &blockingResponse{ResponseRecorder: httptest.NewRecorder(), entered: make(chan struct{}), unblock: make(chan struct{}), blockFlush: true} + request := httptest.NewRequest(http.MethodGet, wire.SnapshotPath, nil) + request.TLS = f.requestState(t) + done := make(chan any, 1) + + go func() { defer func() { done <- recover() }(); f.a.Server.Handler().ServeHTTP(w, request) }() + + <-w.entered + + _, err := f.a.authority.Current() + require.NoError(t, err) + + switch action { + case "advance": + replicationSmokePublish(t, f.ctx, f.a.Topology, AcceptedMembers{testNodeUID: {Node: testNodeUID, Shares: 9, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}}}) + case "suspend": + restore := withdrawPublication(t, f.a.Topology) + restore() + reconcileTopology(t, f.a.Topology, f.ctx) + default: + time.Sleep(3 * time.Second) + + require.NoError(t, f.a.authority.Observe(f.ctx)) + + time.Sleep(2 * time.Second) + } + + require.Len(t, f.a.Server.writes, 1) + require.Equal(t, 1, f.a.Server.polls.count()) + + close(w.unblock) + + if action == "advance" { + require.Nil(t, <-done, "ordinary advancement aborted admitted flush") + } else { + require.Equal(t, http.ErrAbortHandler, <-done, "revoked flush completed") + } + + require.Empty(t, f.a.Server.writes) + require.Zero(t, f.a.Server.polls.count()) + }) + }) + } +} + +func TestSnapshotDeltaRequiresSequenceAndHash(t *testing.T) { + f := newServingFixture(t) + membership := AcceptedMembers{} + + for i := range 100 { + id := wire.NodeID(fmt.Sprintf("22222222-2222-4222-8222-%012d", i)) + membership[id] = wire.Member{Node: id, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}} + } + + base := replicationSmokePublish(t, f.ctx, f.a.Topology, membership) + id := wire.NodeID("22222222-2222-4222-8222-000000000000") + member := membership[id] + member.Shares++ + membership[id] = member + next := replicationSmokePublish(t, f.ctx, f.a.Topology, membership) + baseImage, err := wire.DecodePublication(strings.NewReader(base.encoded)) + require.NoError(t, err) + + for _, tc := range []struct { + name, query, hash string + delta bool + }{ + {"exact", fmt.Sprintf("?after=%d", base.record.Sequence), base.record.ContentHash, true}, + {"older sequence same hash", fmt.Sprintf("?after=%d", base.record.Sequence-1), base.record.ContentHash, false}, + {"no sequence same hash", "", base.record.ContentHash, false}, + {"wrong hash", fmt.Sprintf("?after=%d", base.record.Sequence), strings.Repeat("a", 64), false}, + } { + t.Run(tc.name, func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, wire.SnapshotPath+tc.query, nil) + r.TLS = f.requestState(t) + r.Header.Set(wire.DeltaHeader, tc.hash) + + w := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(w, r) + require.Equal(t, http.StatusOK, w.Code) + + if tc.delta { + applied, err := wire.ApplyDelta(baseImage, bytes.NewReader(w.Body.Bytes())) + require.NoError(t, err) + require.Equal(t, next.record.Sequence, applied.Sequence) + require.Less(t, w.Body.Len(), len(next.encoded)) + } else { + require.Equal(t, next.encoded, w.Body.String()) + } + + requireNoAdmission(t, f.a.Server) + }) + } +} + +func (w *pipeResponse) Write(b []byte) (int, error) { return w.conn.Write(b) } + +func authFixture(t *testing.T) (*Application, authv1.TokenReviewStatus, string) { + t.Helper() + + controller := true + ds := &appsv1.DaemonSet{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "racer-dataplane", UID: "ds-uid"}} + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "racer-dataplane", UID: "sa-uid"}} + node := &corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "worker", UID: types.UID(testNodeUID)}} + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: "racer", Name: "worker-pod", UID: "pod-uid", OwnerReferences: []metav1.OwnerReference{{APIVersion: "apps/v1", Kind: "DaemonSet", Name: ds.Name, UID: ds.UID, Controller: &controller}}}, Spec: corev1.PodSpec{NodeName: node.Name, ServiceAccountName: sa.Name}, Status: corev1.PodStatus{PodIP: "192.0.2.1"}} + r := initializedTopology(t, ds, sa, node, pod) + a := assembleFixture(r.Config, r.Client, r.APIReader) + status := authv1.TokenReviewStatus{Authenticated: true, Audiences: []string{wire.TokenAudience}, User: authv1.UserInfo{Username: "system:serviceaccount:racer:racer-dataplane", UID: string(sa.UID), Extra: map[string]authv1.ExtraValue{"authentication.kubernetes.io/pod-name": {pod.Name}, "authentication.kubernetes.io/pod-uid": {string(pod.UID)}, "authentication.kubernetes.io/node-name": {node.Name}, "authentication.kubernetes.io/node-uid": {string(node.UID)}}}} + token := "header." + base64.RawURLEncoding.EncodeToString(fmt.Appendf(nil, `{"exp":%d}`, time.Now().Add(time.Hour).Unix())) + ".signature" + + return a, status, token +} + +func installReview(t *testing.T, a *Application, status authv1.TokenReviewStatus, token string) { + t.Helper() + + fixtureDependencies[a.authority].Client = interceptor.NewClient(a.Topology.Client.(client.WithWatch), interceptor.Funcs{Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + review, ok := obj.(*authv1.TokenReview) + if !ok { + return c.Create(ctx, obj, opts...) + } + + if review.Spec.Token != token || len(review.Spec.Audiences) != 1 || review.Spec.Audiences[0] != wire.TokenAudience { + t.Error("TokenReview did not bind token/audience") + } + + review.Status = *status.DeepCopy() + + return ctx.Err() + }}) +} + +func TestBootstrapRequiresUnambiguousNodeBindings(t *testing.T) { + for _, key := range []string{"node-name", "node-uid"} { + for _, values := range []authv1.ExtraValue{nil, {}, {""}, {"wrong"}, {"worker", "worker"}, {testNodeUID, testNodeUID}} { + t.Run(fmt.Sprintf("%s/%v", key, values), func(t *testing.T) { + a, status, token := authFixture(t) + fullKey := "authentication.kubernetes.io/" + key + delete(status.User.Extra, fullKey) + + if values != nil { + status.User.Extra[fullKey] = values + } + + installReview(t, a, status, token) + + r := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + r.Header.Set("Authorization", "Bearer "+token) + + if _, err := a.authority.Authenticate(t.Context(), r); err == nil { + t.Fatal("accepted missing or ambiguous node binding") + } + }) + } + } +} + +func TestBootstrapAuthoritativeBindings(t *testing.T) { + for _, scenario := range []string{"success", "audience", "not authenticated", "review error", "username", "sa uid", "pod uid", "missing bound pod", "ambiguous bound pod", "node extra", "node extra uid", "recreated pod", "recreated sa", "recreated ds", "owner name", "owner kind", "owner not controller", "pod sa", "unscheduled", "terminal pod", "excluded node", "deleted node", "api failure", "canceled", "expired token", "duplicate bearer"} { + t.Run(scenario, func(t *testing.T) { + a, status, token := authFixture(t) + runKeys(t, a.Keyring) + reconcileTopology(t, a.Topology, t.Context()) + a.Lifecycle.process, a.Lifecycle.synced = t.Context(), true + a.Lifecycle.SetServingReady(true) + a.Server.tlsConfig(t.Context(), servingTestCertificate(t, 1, time.Now().Add(-time.Minute), time.Now().Add(time.Hour), nil, false)) + _, enrollment, _ := issuanceRequest(t, a.Keyring) + enrollment.RDMANICs = []wire.RDMANIC{{Device: "mlx5_0", Port: 1}} + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + pod := &corev1.Pod{} + require.NoError(t, a.Topology.Get(ctx, client.ObjectKey{Namespace: "racer", Name: "worker-pod"}, pod)) + + mutateBootstrapReview(scenario, &status) + mutateBootstrapPod(scenario, pod) + + switch scenario { + case "recreated sa", "recreated ds": + var obj client.Object = &corev1.ServiceAccount{} + if scenario == "recreated ds" { + obj = &appsv1.DaemonSet{} + } + + require.NoError(t, a.Topology.Get(ctx, client.ObjectKey{Namespace: "racer", Name: "racer-dataplane"}, obj)) + + obj.SetUID("replacement") + + require.NoError(t, a.Topology.Update(ctx, obj)) + case "terminal pod": + pod.Status.Phase = corev1.PodFailed + require.NoError(t, a.Topology.Client.Status().Update(ctx, pod)) + case "excluded node", "deleted node": + node := &corev1.Node{} + require.NoError(t, a.Topology.Get(ctx, client.ObjectKey{Name: "worker"}, node)) + + if scenario == "deleted node" { + require.NoError(t, a.Topology.Delete(ctx, node)) + } else { + node.Labels = map[string]string{wire.ExclusionLabel: ""} + require.NoError(t, a.Topology.Update(ctx, node)) + } + case "expired token": + token = "header." + base64.RawURLEncoding.EncodeToString([]byte(`{"exp":1}`)) + ".signature" + } + + require.NoError(t, a.Topology.Update(ctx, pod)) + + installReview(t, a, status, token) + + if scenario == "api failure" { + fixtureDependencies[a.authority].reader = interceptor.NewClient(a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(context.Context, client.WithWatch, client.ObjectKey, client.Object, ...client.GetOption) error { + return fmt.Errorf("private upstream failure") + }}) + } + + if scenario == "canceled" { + cancel() + } + + req := bootstrapTestRequest(t, ctx, "", token, enrollment) + req.TLS = &tls.ConnectionState{HandshakeComplete: true} + + if scenario == "duplicate bearer" { + req.Header.Add("Authorization", "Bearer "+token) + } + + w := httptest.NewRecorder() + a.Server.Handler().ServeHTTP(w, req) + + requireBootstrapBindingResult(t, a, scenario, w) + }) + } +} + +func requireBootstrapBindingResult(t *testing.T, a *Application, scenario string, w *httptest.ResponseRecorder) { + t.Helper() + + if scenario == "success" { + issued := decodeIssuedResponse(t, responseBody(t, w.Result(), nil, http.StatusOK)) + leaf, err := x509.ParseCertificate(issued.CertificateChain[0]) + require.NoError(t, err) + require.Equal(t, wire.NodeID(testNodeUID), issued.Node) + require.Equal(t, a.Topology.Config.Cluster, issued.Cluster) + require.True(t, leaf.NotAfter.After(time.Now())) + + return + } + + want := http.StatusForbidden + + switch scenario { + case "audience", "not authenticated", "review error", "missing bound pod", "ambiguous bound pod", "expired token", "duplicate bearer": + want = http.StatusUnauthorized + case "api failure", "canceled": + want = http.StatusServiceUnavailable + } + + responseBody(t, w.Result(), nil, want) + + if scenario != "deleted node" { + var node corev1.Node + require.NoError(t, a.Topology.Get(t.Context(), client.ObjectKey{Name: "worker"}, &node)) + require.Empty(t, node.Annotations[enrolledSharesAnnotation], "rejected enrollment persisted shares") + require.Empty(t, node.Annotations[enrolledRDMANICsAnnotation], "rejected enrollment persisted NICs") + } +} + +func mutateBootstrapReview(scenario string, status *authv1.TokenReviewStatus) { + switch scenario { + case "audience": + status.Audiences = []string{"api"} + case "not authenticated": + status.Authenticated = false + case "review error": + status.Error = "private upstream details" + case "username": + status.User.Username = "system:serviceaccount:other:racer-dataplane" + case "sa uid": + status.User.UID = "old-sa" + case "pod uid": + status.User.Extra["authentication.kubernetes.io/pod-uid"] = authv1.ExtraValue{"old-pod"} + case "missing bound pod": + delete(status.User.Extra, "authentication.kubernetes.io/pod-name") + case "ambiguous bound pod": + status.User.Extra["authentication.kubernetes.io/pod-name"] = authv1.ExtraValue{"worker-pod", "other"} + case "node extra": + status.User.Extra["authentication.kubernetes.io/node-name"] = authv1.ExtraValue{"other"} + case "node extra uid": + status.User.Extra["authentication.kubernetes.io/node-uid"] = authv1.ExtraValue{"old"} + } +} + +func mutateBootstrapPod(scenario string, pod *corev1.Pod) { + switch scenario { + case "recreated pod": + pod.UID = "replacement" + case "owner name": + pod.OwnerReferences[0].Name = "other" + case "owner kind": + pod.OwnerReferences[0].Kind = "ReplicaSet" + case "owner not controller": + pod.OwnerReferences[0].Controller = nil + case "pod sa": + pod.Spec.ServiceAccountName = "other" + case "unscheduled": + pod.Spec.NodeName = "" + } +} + +func TestEnrollmentSaturationPreservesLocalAuthentication(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxConcurrentBootstrap = 1 }) + s := f.a.Server + entered := make(chan struct{}) + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{ + Get: func(ctx context.Context, _ client.WithWatch, _ client.ObjectKey, _ client.Object, _ ...client.GetOption) error { + close(entered) + <-ctx.Done() + + return ctx.Err() + }, + }) + endpoint := f.start(t) + + body, err := wire.EncodeBootstrapRequest(f.request) + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + r := httptest.NewRequestWithContext(ctx, http.MethodPost, wire.BootstrapPath, bytes.NewReader(body)) + r.Header.Set("Content-Type", "application/json") + r.Header.Set("Authorization", "Bearer "+f.token) + r.TLS = &tls.ConnectionState{HandshakeComplete: true} + done := make(chan *httptest.ResponseRecorder, 1) + + go func() { + w := httptest.NewRecorder() + s.Handler().ServeHTTP(w, r) + + done <- w + }() + + select { + case <-entered: + case <-time.After(5 * time.Second): + t.Fatal("enrollment did not enter API wait") + } + + // A fresh TLS connection and snapshot both succeed while enrollment is full. + peer := f.client(t, &f.certificate) + response, err := peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusOK) + // Isolated admission must preserve certificate authentication. + anonymous := f.client(t, nil) + response, err = anonymous.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusUnauthorized) + // New enrollment reaches HTTP backpressure without another API operation. + response, err = anonymous.Post(endpoint+wire.BootstrapPath, "application/json", bytes.NewReader(body)) + responseBody(t, response, err, http.StatusTooManyRequests) + + if len(s.bootstrapSlots) != 1 || len(s.authSlots) != 0 { + t.Fatal("enrollment consumed local authentication capacity") + } + + // Saturation cannot bypass withdrawn trust, even on an established connection. + fixtureDependencies[f.a.authority].reader = f.a.Topology.Client + invalidateFixtureTrust(t, f) + + response, err = peer.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusServiceUnavailable) + fresh := f.client(t, &f.certificate) + + response, err = fresh.Get(endpoint + wire.SnapshotPath) + if err == nil { + response.Body.Close() + t.Fatal("fresh handshake accepted withdrawn trust") + } + + cancel() + + select { + case w := <-done: + responseBody(t, w.Result(), nil, http.StatusServiceUnavailable) + case <-time.After(5 * time.Second): + t.Fatal("canceled enrollment retained admission") + } + + if len(s.bootstrapSlots) != 0 || len(s.authSlots) != 0 { + t.Fatal("admission leaked after cancellation") + } +} + +func TestBootstrapIssuanceBeforeWriteAdmission(t *testing.T) { + for _, scenario := range []string{"success", "write saturation", "readiness lost", "canceled"} { + t.Run(scenario, func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxConcurrentWrites = 1 }) + s := f.a.Server + handler := s.Handler() + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + wantWrites, wantStatus := 0, http.StatusOK + + if scenario == "write saturation" { + take(s.writes) + defer release(s.writes) + + wantWrites, wantStatus = 1, http.StatusTooManyRequests + } else if scenario != "success" { + wantStatus = http.StatusServiceUnavailable + } + + reads := 0 + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, c client.WithWatch, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + reads++ + + if len(s.writes) != wantWrites || len(s.bootstrapSlots) != 1 || len(s.authSlots) != 0 { + t.Error("issuance changed write/auth admission timing") + } + + switch scenario { + case "readiness lost": + f.a.Lifecycle.SetServingReady(false) + case "canceled": + cancel() + } + + return c.Get(ctx, key, obj, opts...) + }}) + + r := bootstrapTestRequest(t, ctx, "", f.token, f.request) + r.TLS = &tls.ConnectionState{HandshakeComplete: true} + + w := httptest.NewRecorder() + handler.ServeHTTP(w, r) + + encoded := responseBody(t, w.Result(), nil, wantStatus) + require.NotZero(t, reads, "issuance skipped") + require.Empty(t, s.bootstrapSlots) + require.Empty(t, s.authSlots) + require.Len(t, s.writes, wantWrites) + + if scenario == "success" { + response := decodeIssuedResponse(t, encoded) + require.Equal(t, f.request.Enrollment, response.Enrollment) + require.Equal(t, wire.NodeID(testNodeUID), response.Node) + require.True(t, w.Flushed) + require.Equal(t, "no-store", w.Header().Get("Cache-Control")) + } + }) + } +} diff --git a/internal/racer/server/metrics.go b/internal/racer/server/metrics.go new file mode 100644 index 000000000..1b17f4b3b --- /dev/null +++ b/internal/racer/server/metrics.go @@ -0,0 +1,48 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "github.com/prometheus/client_golang/prometheus" + "sigs.k8s.io/controller-runtime/pkg/metrics" +) + +var servingTransportMetrics = newTransportMetrics(metrics.Registry) + +// Collectors aggregate all Racer listeners in this process. Tests use their own +// registry and collectors, without replacing or resetting the shared registry. +type transportMetrics struct { + connections prometheus.Gauge + handshakes prometheus.Gauge + connectionRejected prometheus.Counter + handshakeRejected prometheus.Counter + handshakeTimeouts prometheus.Counter +} + +func newTransportMetrics(reg prometheus.Registerer) *transportMetrics { + connections := prometheus.NewGauge(prometheus.GaugeOpts{ + Name: "racer_server_connections", + Help: "Admitted sockets from connection admission until raw close, including TLS handshakes, HTTP requests, and idle connections.", + }) + handshakes := prometheus.NewGauge(prometheus.GaugeOpts{ + Name: "racer_server_tls_handshakes", + Help: "Full TLS handshake slots held from admission before reading TLS bytes until HandshakeContext returns, not just TLS configuration callbacks.", + }) + rejected := prometheus.NewCounterVec(prometheus.CounterOpts{ + Name: "racer_server_admission_rejections_total", + Help: "Sockets rejected by the first exhausted admission stage: connection or full_handshake. Excludes shutdown and TLS errors.", + }, []string{"stage"}) + timeouts := prometheus.NewCounter(prometheus.CounterOpts{ + Name: "racer_server_tls_handshake_timeouts_total", + Help: "Failed full TLS handshakes returning a context deadline or transport timeout error, counted once. Excludes HTTP timeouts and cancellation.", + }) + reg.MustRegister(connections, handshakes, rejected, timeouts) + + return &transportMetrics{ + connections: connections, handshakes: handshakes, + connectionRejected: rejected.WithLabelValues("connection"), + handshakeRejected: rejected.WithLabelValues("full_handshake"), + handshakeTimeouts: timeouts, + } +} diff --git a/internal/racer/server/metrics_test.go b/internal/racer/server/metrics_test.go new file mode 100644 index 000000000..e6761efe9 --- /dev/null +++ b/internal/racer/server/metrics_test.go @@ -0,0 +1,107 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "context" + "crypto/tls" + "net" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/require" + "sigs.k8s.io/controller-runtime/pkg/metrics" +) + +func TestTransportMetricsRegistryIsolation(t *testing.T) { + first, second := prometheus.NewPedanticRegistry(), prometheus.NewPedanticRegistry() + m := newTransportMetrics(first) + newTransportMetrics(second) + m.connections.Inc() + m.handshakes.Inc() + m.connectionRejected.Inc() + m.handshakeRejected.Add(2) + m.handshakeTimeouts.Inc() + + require.NoError(t, testutil.GatherAndCompare(first, strings.NewReader(` +# HELP racer_server_admission_rejections_total Sockets rejected by the first exhausted admission stage: connection or full_handshake. Excludes shutdown and TLS errors. +# TYPE racer_server_admission_rejections_total counter +racer_server_admission_rejections_total{stage="connection"} 1 +racer_server_admission_rejections_total{stage="full_handshake"} 2 +# HELP racer_server_connections Admitted sockets from connection admission until raw close, including TLS handshakes, HTTP requests, and idle connections. +# TYPE racer_server_connections gauge +racer_server_connections 1 +# HELP racer_server_tls_handshake_timeouts_total Failed full TLS handshakes returning a context deadline or transport timeout error, counted once. Excludes HTTP timeouts and cancellation. +# TYPE racer_server_tls_handshake_timeouts_total counter +racer_server_tls_handshake_timeouts_total 1 +# HELP racer_server_tls_handshakes Full TLS handshake slots held from admission before reading TLS bytes until HandshakeContext returns, not just TLS configuration callbacks. +# TYPE racer_server_tls_handshakes gauge +racer_server_tls_handshakes 1 +`))) + // Both rejection stages are present at zero before any traffic. The other + // collectors have no labels; no peer, identity, path, or error text is used. + families, err := second.Gather() + require.NoError(t, err) + require.Len(t, families, 4) + + for _, family := range families { + for _, sample := range family.GetMetric() { + require.Zero(t, sample.GetGauge().GetValue()) + require.Zero(t, sample.GetCounter().GetValue()) + } + } + + // Production registers once with the controller-runtime metrics endpoint. + count, err := testutil.GatherAndCount(metrics.Registry, + "racer_server_connections", "racer_server_tls_handshakes", + "racer_server_admission_rejections_total", "racer_server_tls_handshake_timeouts_total") + require.NoError(t, err) + require.Equal(t, 5, count) +} + +func TestTransportMetricsShutdownAcceptance(t *testing.T) { + for _, capacity := range []int{0, 1} { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + accepted, peer := net.Pipe() + defer accepted.Close() + defer peer.Close() + + listener := teardownListener{close: func() error { return nil }, accept: func() (net.Conn, error) { + // Accept can return a socket while shutdown is already in progress. + cancel() + return accepted, nil + }} + + l := newTransportListenerWithMetrics(ctx, listener, &tls.Config{}, Limits{MaxConnections: capacity}, newTransportMetrics(prometheus.NewPedanticRegistry())) + defer l.Close() + + <-l.acceptDone + synctest.Wait() + + assertTransportMetrics(t, l, 0, 0, 0, 0, 0) + require.Empty(t, l.connections) + require.Empty(t, l.handshakes) + }) + } +} + +func assertTransportMetrics(t *testing.T, l *transportListener, connections, handshakes, connectionRejected, handshakeRejected, timeouts float64) { + t.Helper() + + // Slot changes and gauge changes are separate atomic operations. Wait for + // both before asserting the counters at this lifecycle boundary. + require.Eventually(t, func() bool { + return testutil.ToFloat64(l.metrics.connections) == connections && testutil.ToFloat64(l.metrics.handshakes) == handshakes + }, time.Second, time.Millisecond) + require.Equal(t, connectionRejected, testutil.ToFloat64(l.metrics.connectionRejected)) + require.Equal(t, handshakeRejected, testutil.ToFloat64(l.metrics.handshakeRejected)) + require.Equal(t, timeouts, testutil.ToFloat64(l.metrics.handshakeTimeouts)) +} diff --git a/internal/racer/server/replication_test.go b/internal/racer/server/replication_test.go new file mode 100644 index 000000000..c8f6780bd --- /dev/null +++ b/internal/racer/server/replication_test.go @@ -0,0 +1,1017 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/sha256" + "crypto/tls" + "crypto/x509" + "encoding/json" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "net/url" + "os" + "runtime" + "slices" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + authv1 "k8s.io/api/authentication/v1" + corev1 "k8s.io/api/core/v1" + "k8s.io/apimachinery/pkg/api/meta" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/fields" + "k8s.io/apimachinery/pkg/runtime/schema" + "k8s.io/apimachinery/pkg/types" + "k8s.io/client-go/rest" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/cache" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +type fixtureConfig struct { + authority.Config + ServerConfig Config + PeerPort uint16 +} + +func testReplicationFlush(t *testing.T, f *servingFixture, unchanged, fail bool) { + t.Helper() + + request := httptest.NewRequest(http.MethodGet, ReplicationPath, nil) + request.TLS = f.requestState(t) + request.Header.Set("Authorization", "Bearer "+f.token) + + want := http.StatusOK + + if unchanged { + p, err := f.a.authority.Current() + require.NoError(t, err) + + request.URL.RawQuery = fmt.Sprintf("after=%d", p.Sequence()) + want = http.StatusNoContent + } + + w := &blockingResponse{ResponseRecorder: httptest.NewRecorder(), entered: make(chan struct{}), unblock: make(chan struct{}), blockFlush: true, fail: fail} + + unblock := sync.OnceFunc(func() { close(w.unblock) }) + defer unblock() + + done := make(chan any, 1) + + go func() { done <- serveRecover(f.a.Server.Handler(), w, request) }() + + select { + case <-w.entered: + case <-time.After(8 * time.Second): + t.Fatal("explicit flush not reached") + } + + require.Len(t, f.a.Server.writes, 1, "write admission released before flush") + require.Equal(t, 1, f.a.Server.replicationPolls.count(), "poll admission released before flush") + + duplicate := httptest.NewRecorder() + f.a.Server.Handler().ServeHTTP(duplicate, request.Clone(f.ctx)) + require.Equal(t, http.StatusTooManyRequests, duplicate.Code, "duplicate during flush") + unblock() + + var wantAbort any + if fail { + wantAbort = http.ErrAbortHandler + } + + require.Equal(t, wantAbort, <-done, "flush result") + require.Equal(t, want, w.Code) + require.Empty(t, f.a.Server.writes, "write admission leaked after flush") + require.Zero(t, f.a.Server.replicationPolls.count(), "poll admission leaked after flush") +} + +func (r *TopologyReconciler) observeTopology(ctx context.Context) (TopologyObservation, error) { + cfg := r.Config + + var nodes corev1.NodeList + if err := r.List(ctx, &nodes); err != nil { + return TopologyObservation{}, err + } + + var caches racerv1.ClusterCacheList + if err := r.APIReader.List(ctx, &caches); err != nil { + return TopologyObservation{}, err + } + + catalog, err := members.BuildCatalog(caches.Items) + if err != nil { + return TopologyObservation{}, err + } + + ownership, err := members.ReadWorkloadIdentities(ctx, r.APIReader, cfg.Namespace, cfg.DaemonSetName) + if err != nil { + return TopologyObservation{}, err + } + // Indexed namespace-scoped queries avoid scanning unrelated Pods per Node. + podsByNode := make(map[string][]corev1.Pod, len(nodes.Items)) + for _, node := range nodes.Items { + if err := ctx.Err(); err != nil { + return TopologyObservation{}, err + } + + var list corev1.PodList + if err := r.List(ctx, &list, client.InNamespace(cfg.Namespace), client.MatchingFields{podNodeIndex: node.Name}); err != nil { + return TopologyObservation{}, err + } + + podsByNode[node.Name] = list.Items + } + + return TopologyObservation{Nodes: nodes, Catalog: catalog, Input: members.Input{ + Nodes: nodes.Items, PodsByNode: podsByNode, Ownership: ownership, PeerPort: cfg.PeerPort, + }}, nil +} + +func (c fixtureConfig) authorityConfig() authority.Config { return c.Config } + +type Application struct { + authority *authority.Authority + Topology *TopologyReconciler + Keyring *KeyringReconciler + Server *Server + Lifecycle *Lifecycle + Replication *fixtureLeader +} + +type TopologyReconciler struct { + client.Client + APIReader client.Reader + Config fixtureConfig + authority *authority.Authority +} + +func (r *TopologyReconciler) Reconcile(ctx context.Context, _ ctrl.Request) (ctrl.Result, error) { + _, err := r.authority.PublishTopology(ctx, r.observeTopology) + return ctrl.Result{}, err +} + +type KeyringReconciler struct { + Config fixtureConfig + authority *authority.Authority +} + +func (r *KeyringReconciler) Reconcile(ctx context.Context, _ ctrl.Request) (ctrl.Result, error) { + delay, err := r.authority.ReconcileCredentials(ctx) + return ctrl.Result{RequeueAfter: delay}, err +} + +type fixtureLeader struct { + Config fixtureConfig + Client client.Client + APIReader client.Reader + authority *authority.Authority + mu sync.Mutex + leader context.Context +} + +func (r *fixtureLeader) LeaderContext() (context.Context, bool) { + r.mu.Lock() + defer r.mu.Unlock() + + return r.leader, r.leader != nil && r.leader.Err() == nil +} + +func (r *fixtureLeader) PollInterval() time.Duration { + return min(5*time.Second, r.Config.SnapshotMaxAge/3) +} + +func (r *fixtureLeader) AuthenticateReplica(ctx context.Context, req *http.Request) (string, time.Time, error) { + i, err := r.authority.AuthenticateReplica(ctx, req) + return i.UID(), i.Expires(), err +} + +func (r *fixtureLeader) observe(ctx context.Context) { _ = r.authority.Observe(ctx) } + +func (r *fixtureLeader) installReplica(ctx, process context.Context, p wire.Publication) error { + return r.authority.AcceptReplica(ctx, process, p) +} + +func Assemble(cfg fixtureConfig, c client.Client, reader client.Reader) *Application { + a := authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: c, Reader: reader}) + l := NewLifecycle(a) + r := &fixtureLeader{Config: cfg, Client: c, APIReader: reader, authority: a} + + return &Application{authority: a, Topology: &TopologyReconciler{Client: c, APIReader: reader, Config: cfg, authority: a}, Keyring: &KeyringReconciler{Config: cfg, authority: a}, Server: New(cfg.ServerConfig, c, a, l, r), Lifecycle: l, Replication: r} +} + +type cachedScaleClient struct { + client.Client + reader client.Reader +} + +func (c cachedScaleClient) Get(ctx context.Context, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + return c.reader.Get(ctx, key, obj, opts...) +} + +func (c cachedScaleClient) List(ctx context.Context, list client.ObjectList, opts ...client.ListOption) error { + return c.reader.List(ctx, list, opts...) +} + +// The informer uses a synthetic list/watch HTTP source, but the cache, field +// index, deep copies, reconciler, canonical hashes, encoding and waiting are real. +// Only durable version CAS uses a fake client. This is deliberately not an HTTPS +// authentication or API-server capacity benchmark. +func TestServerScale(t *testing.T) { + if os.Getenv("RACER_SCALE") != "1" { + t.Skip("set RACER_SCALE=1 for 100,000-member reconciliation and waiter measurements") + } + + for _, count := range []int{1_000, 10_000, 100_000} { + t.Run(fmt.Sprint(count), func(t *testing.T) { + r := initializedTopology(t) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + reader := scaleCache(t, r, count) + r.Client = cachedScaleClient{Client: r.Client, reader: reader} + + runtime.GC() + + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + + start := time.Now() + first := reconcileTopology(t, r, ctx) + cold := time.Since(start) + + runtime.ReadMemStats(&after) + require.Len(t, acceptedMembers(t, r), count) + t.Logf("members=%d cold_reconcile=%s allocated_bytes=%d publication_bytes=%d", count, cold, after.TotalAlloc-before.TotalAlloc, len(first.encoded)) + + start = time.Now() + + if current := reconcileTopology(t, r, ctx); current != first { + t.Fatal("no-op reconcile replaced publication") + } + + t.Logf("members=%d unchanged_reconcile=%s", count, time.Since(start)) + + if count == 100_000 { + scaleFanout(t, r, ctx, count) + } + }) + } +} + +func scaleLists(t *testing.T, r *TopologyReconciler, count int) map[string]any { + t.Helper() + + nodes := &corev1.NodeList{TypeMeta: metav1.TypeMeta{APIVersion: "v1", Kind: "NodeList"}, ListMeta: metav1.ListMeta{ResourceVersion: "1"}} + pods := &corev1.PodList{TypeMeta: metav1.TypeMeta{APIVersion: "v1", Kind: "PodList"}, ListMeta: metav1.ListMeta{ResourceVersion: "1"}} + + for i := range count { + uid := types.UID(fmt.Sprintf("%08x-0000-4000-8000-000000000000", i)) + node := memberNode() + node.Name, node.UID, node.ResourceVersion = fmt.Sprintf("node-%d", i), uid, "1" + node.Annotations = map[string]string{wire.RDMANICsAnnotation: `[{"rail":0,"device":"mlx5_0","port":1,"numa_node":0},{"rail":1,"device":"mlx5_1","port":1,"numa_node":1}]`} + pod := memberPod(uid, 1, fmt.Sprintf("10.%d.%d.%d", i>>16, (i>>8)&255, i&255)) + pod.Spec.NodeName, pod.ResourceVersion = node.Name, "1" + pod.OwnerReferences[0].Name = r.Config.DaemonSetName + + nodes.Items, pods.Items = append(nodes.Items, node), append(pods.Items, pod) + } + + ds := &appsv1.DaemonSetList{TypeMeta: metav1.TypeMeta{APIVersion: "apps/v1", Kind: "DaemonSetList"}, ListMeta: metav1.ListMeta{ResourceVersion: "1"}, Items: []appsv1.DaemonSet{{ObjectMeta: metav1.ObjectMeta{Namespace: r.Config.Namespace, Name: r.Config.DaemonSetName, UID: testDaemonSetUID, ResourceVersion: "1"}}}} + + caches := &racerv1.ClusterCacheList{TypeMeta: metav1.TypeMeta{APIVersion: racerv1.GroupVersion.String(), Kind: "ClusterCacheList"}, ListMeta: metav1.ListMeta{ResourceVersion: "1"}} + for i := range 16 { + caches.Items = append(caches.Items, catalogCache(fmt.Sprintf("cache-%d", i), types.UID(fmt.Sprintf("%08x-1111-4000-8000-000000000000", i)))) + require.NoError(t, r.Create(t.Context(), &caches.Items[i])) + } + + runKeys(t, Assemble(r.Config, r.Client, r.APIReader).Keyring) + + return map[string]any{ + "/api/v1/nodes": nodes, + "/api/v1/namespaces/" + r.Config.Namespace + "/pods": pods, + "/apis/apps/v1/namespaces/" + r.Config.Namespace + "/daemonsets": ds, + "/apis/" + racerv1.GroupVersion.String() + "/clustercaches": caches, + } +} + +func scaleSource(lists map[string]any) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + w.Header().Set("Content-Type", "application/json") + + if req.URL.Query().Get("sendInitialEvents") == "true" { + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(metav1.Status{TypeMeta: metav1.TypeMeta{APIVersion: "v1", Kind: "Status"}, Status: "Failure", Reason: metav1.StatusReasonBadRequest, Code: 400, Message: "synthetic source supports ordinary list/watch"}) + + return + } + + if req.URL.Query().Get("watch") == "true" { + w.WriteHeader(http.StatusOK) + http.NewResponseController(w).Flush() + <-req.Context().Done() + + return + } + + list, ok := lists[req.URL.Path] + if !ok { + http.NotFound(w, req) + return + } + + json.NewEncoder(w).Encode(list) + }) +} + +func scaleCache(t *testing.T, r *TopologyReconciler, count int) cache.Cache { + t.Helper() + source := httptest.NewServer(scaleSource(scaleLists(t, r, count))) + t.Cleanup(source.Close) + + mapper := meta.NewDefaultRESTMapper([]schema.GroupVersion{corev1.SchemeGroupVersion, appsv1.SchemeGroupVersion, racerv1.GroupVersion}) + mapper.Add(corev1.SchemeGroupVersion.WithKind("Node"), meta.RESTScopeRoot) + mapper.Add(corev1.SchemeGroupVersion.WithKind("Pod"), meta.RESTScopeNamespace) + mapper.Add(corev1.SchemeGroupVersion.WithKind("PodList"), meta.RESTScopeNamespace) + mapper.Add(corev1.SchemeGroupVersion.WithKind("Secret"), meta.RESTScopeNamespace) + mapper.Add(corev1.SchemeGroupVersion.WithKind("ConfigMap"), meta.RESTScopeNamespace) + mapper.Add(appsv1.SchemeGroupVersion.WithKind("DaemonSet"), meta.RESTScopeNamespace) + mapper.Add(racerv1.GroupVersion.WithKind("ClusterCache"), meta.RESTScopeRoot) + + options := cache.Options{ByObject: map[client.Object]cache.ByObject{ + &corev1.Pod{}: {Namespaces: map[string]cache.Config{r.Config.Namespace: {}}}, + &corev1.Secret{}: {Namespaces: map[string]cache.Config{r.Config.Namespace: {}}, Field: fields.OneTermEqualSelector("metadata.name", r.Config.CredentialsSecretName)}, + &corev1.ConfigMap{}: {Namespaces: map[string]cache.Config{r.Config.Namespace: {}}}, + &appsv1.DaemonSet{}: {Namespaces: map[string]cache.Config{r.Config.Namespace: {}}}, + }} + options.Scheme, options.Mapper = r.Scheme(), mapper + reader, err := cache.New(&rest.Config{Host: source.URL, QPS: 1000, Burst: 1000}, options) + require.NoError(t, err) + require.NoError(t, reader.IndexField(t.Context(), &corev1.Pod{}, podNodeIndex, podNodeKeys)) + + for _, obj := range []client.Object{&corev1.Node{}, &appsv1.DaemonSet{}, &racerv1.ClusterCache{}} { + _, err := reader.GetInformer(t.Context(), obj) + require.NoError(t, err) + } + + ctx, cancel := context.WithCancel(t.Context()) + done := make(chan error, 1) + + go func() { done <- reader.Start(ctx) }() + + t.Cleanup(func() { + cancel() + + if err := <-done; err != nil { + t.Error(err) + } + }) + + syncCtx, stopSync := context.WithTimeout(ctx, time.Minute) + defer stopSync() + + require.True(t, reader.WaitForCacheSync(syncCtx), "scale informer did not synchronize") + + return reader +} + +func scaleFanout(t *testing.T, r *TopologyReconciler, ctx context.Context, count int) { + t.Helper() + + current, err := r.authority.Current() + require.NoError(t, err) + // Keep both realistic full-size encodings alive. Prepare before admission so + // fanout measures Install plus delivery, independent of canonical encoding. + members := make(AcceptedMembers, count) + + for id, member := range acceptedMembers(t, r) { + member.Shares++ + members[id] = member + } + + sequence := current.Sequence() + server := New(r.Config.ServerConfig, nil, nil, nil, nil) + + waiting, cancel := context.WithCancel(ctx) + defer cancel() + + results := make(chan *authority.PublicationHandle, count) + failures := make(chan error, count) + + var wg sync.WaitGroup + + runtime.GC() + runtime.GC() // Clear temporary encoding sync.Pools before the waiter baseline. + + var before, parked runtime.MemStats + runtime.ReadMemStats(&before) + + start := time.Now() + + for id := range members { + wg.Go(func() { + if !server.polls.acquire(id) { + results <- nil + + failures <- wire.Overloaded + + return + } + defer server.polls.release(id) + + p, err := waitFixturePublication(waiting, r.authority, sequence) + results <- p + + failures <- err + }) + } + + defer wg.Wait() + defer cancel() + + eventually(t, "100000 admitted waiters", func() bool { return server.polls.count() == count }) + + admit := time.Since(start) + + runtime.GC() + runtime.ReadMemStats(&parked) + require.False(t, server.polls.acquire(testOtherUID), "global bound failed") + + for id := range members { + require.False(t, server.polls.acquire(id), "duplicate bound failed") + break + } + + start = time.Now() + next := replicationSmokePublish(t, ctx, r, members) + install := time.Since(start) + + wg.Wait() + + fanout := time.Since(start) + + for range count { + require.NoError(t, <-failures) + + got := <-results + require.NotNil(t, got) + require.Equal(t, next.record.Sequence, got.Sequence(), "waiter missed or copied full publication") + } + + awaitServerPolls(t, server, 0) + t.Logf("waiters=%d GOMAXPROCS=%d admission=%s install=%s all_delivered=%s heap_delta=%d stack_delta=%d next_bytes=%d", count, runtime.GOMAXPROCS(0), admit, install, fanout, int64(parked.HeapAlloc)-int64(before.HeapAlloc), int64(parked.StackInuse)-int64(before.StackInuse), len(next.encoded)) +} + +// These are real HTTPS protocol clients, not Rust processes or in-memory Wait +// calls. Kubernetes authority is fake. Never interpret this as API capacity. +func TestReplicatedServingSmoke(t *testing.T) { replicatedServingSmoke(t, 12) } + +func TestReplicatedServingCapacity(t *testing.T) { + value := os.Getenv("RACER_REPLICATION_CLIENTS") + if value == "" { + t.Skip("set RACER_REPLICATION_CLIENTS to an exact count in [1,10000]") + } + + count, err := strconv.Atoi(value) + if err != nil || count < 1 || count > 10_000 { + t.Fatal("RACER_REPLICATION_CLIENTS must be in [1,10000]; 100k requires distributed validation, not this loopback harness") + } + + replicatedServingSmoke(t, count) +} + +type replicationSmokeListener struct { + net.Listener + accepted atomic.Int64 + live atomic.Int64 +} + +type replicationSmokeConn struct { + net.Conn + once sync.Once + owner *replicationSmokeListener +} + +func (l *replicationSmokeListener) Accept() (net.Conn, error) { + c, err := l.Listener.Accept() + if err != nil { + return nil, err + } + + l.accepted.Add(1) + l.live.Add(1) + + return &replicationSmokeConn{Conn: c, owner: l}, nil +} + +func (c *replicationSmokeConn) Close() error { + c.once.Do(func() { c.owner.live.Add(-1) }) + return c.Conn.Close() +} + +type replicationSmokeReplica struct { + a *Application + ctx context.Context + cancel context.CancelFunc + endpoint string + listener *replicationSmokeListener + done chan error +} + +type replicationSmokePeer struct { + client *http.Client + replica int +} + +type replicationSmokeResult struct { + index int + bytes int64 + digest [32]byte + elapsed time.Duration + err error +} + +type replicationSmoke struct { + ctx context.Context + fixture *servingFixture + peers []replicationSmokePeer + replicas []*replicationSmokeReplica + internal *http.Client + reviews atomic.Int64 + overloaded atomic.Int64 + reconnected atomic.Int64 + workers sync.WaitGroup +} + +func replicatedServingSmoke(t *testing.T, count int) { + t.Helper() + + ctx, cancel := context.WithTimeout(t.Context(), 240*time.Second) + defer cancel() + + setup := time.Now() + f := newServingFixture(t) + smoke := &replicationSmoke{ctx: ctx, fixture: f} + accepted := smoke.createPeers(t, count) + base := replicationSmokePublish(t, ctx, f.a.Topology, accepted) + smoke.startReplicas(t) + smoke.authorizeReplicas(t) + smoke.install(t, base) + smoke.internal = f.client(t, nil) + smoke.internal.Timeout = 30 * time.Second + smoke.replicate(t, base) + + leader := smoke.replicas[0] + for _, path := range []string{wire.SnapshotPath, ReplicationPath} { + request, err := http.NewRequestWithContext(ctx, http.MethodGet, leader.endpoint+path, nil) + require.NoError(t, err) + response, err := smoke.internal.Do(request) + responseBody(t, response, err, http.StatusUnauthorized) + } + + observed := make(chan struct{}) + + go func() { defer close(observed); smoke.observe() }() + + defer func() { cancel(); <-observed }() + defer func() { cancel(); smoke.workers.Wait() }() + + t.Logf("clients=%d replicas=3 GOMAXPROCS=%d setup=%s full_bytes=%d max_writes=%d max_auth=%d", count, runtime.GOMAXPROCS(0), time.Since(setup), len(base.encoded), leader.a.Server.config.Limits.MaxConcurrentWrites, leader.a.Server.config.Limits.MaxConcurrentBootstrap) + replicationSmokeStats(t, "baseline", smoke.replicas) + + start := time.Now() + smoke.collect(t, "cold", start, smoke.run(0, true), base) + + for phase := range 2 { + start = time.Now() + results := smoke.run(base.record.Sequence, false) + replicationSmokePark(t, ctx, smoke.replicas, count, results) + replicationSmokeStats(t, "parked", smoke.replicas) + + if phase == 1 { + smoke.replicas[2].cancel() + replicationSmokePark(t, ctx, smoke.replicas[:2], count, results) + t.Logf("failure_repark=%s", time.Since(start)) + replicationSmokeStats(t, "failure-parked", smoke.replicas) + } + + for id, member := range accepted { + member.Shares++ + accepted[id] = member + } + + start = time.Now() + base = replicationSmokePublish(t, ctx, f.a.Topology, accepted) + smoke.install(t, base) + smoke.replicate(t, base) + smoke.collect(t, []string{"update", "replica-failure-update"}[phase], start, results, base) + } + + require.Equal(t, int64(count/3), smoke.reconnected.Load(), "reconnected clients") +} + +func (s *replicationSmoke) createPeers(t *testing.T, count int) AcceptedMembers { + t.Helper() + + f := s.fixture + accepted := make(AcceptedMembers, count) + s.peers = make([]replicationSmokePeer, count) + handshakes := make(chan struct{}, 24) + _, _, rotation, material := keyState(t, f.a.Keyring) + ca, signingKey, err := parseSigning(material.Keys[rotation.ActiveIssuer]) + require.NoError(t, err) + + for i := range s.peers { + id := wire.NodeID(fmt.Sprintf("22222222-2222-4222-8222-%012d", i)) + accepted[id] = wire.Member{Node: id, Shares: 4, PeerEndpoint: fmt.Sprintf("10.%d.%d.%d:8082", i>>16, (i>>8)&255, i&255), RDMANICs: []wire.RDMANIC{}} + pub, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{SerialNumber: big.NewInt(int64(i + 1)), NotBefore: time.Now().Add(-time.Minute), NotAfter: time.Now().Add(time.Hour), KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, URIs: []*url.URL{{Scheme: "spiffe", Host: string(f.request.Cluster), Path: "/node/" + string(id)}}} + der, err := x509.CreateCertificate(rand.Reader, template, ca, pub, signingKey) + require.NoError(t, err) + + transport := &http.Transport{TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS13, RootCAs: f.roots, Certificates: []tls.Certificate{{Certificate: [][]byte{der, ca.Raw}, PrivateKey: key}}}, MaxConnsPerHost: 1, MaxIdleConns: 1, MaxIdleConnsPerHost: 1, IdleConnTimeout: time.Minute, TLSHandshakeTimeout: 10 * time.Second} + transport.DialTLSContext = func(ctx context.Context, network, address string) (net.Conn, error) { + select { + case handshakes <- struct{}{}: + case <-ctx.Done(): + return nil, ctx.Err() + } + + defer func() { <-handshakes }() + + bounded, stop := context.WithTimeout(ctx, 10*time.Second) + defer stop() + + return (&tls.Dialer{Config: transport.TLSClientConfig}).DialContext(bounded, network, address) + } + t.Cleanup(transport.CloseIdleConnections) + s.peers[i] = replicationSmokePeer{client: &http.Client{Transport: transport, Timeout: 45 * time.Second}, replica: i % 3} + } + + return accepted +} + +func (s *replicationSmoke) startReplicas(t *testing.T) { + t.Helper() + + f := s.fixture + for range 3 { + a := assembleFixture(f.a.Topology.Config, f.a.Topology.Client, f.a.Topology.APIReader) + process, stop := context.WithCancel(s.ctx) + a.authority.BindProcess(process) + replicationSmokeLifecycle(a.Lifecycle, process) + a.Replication.observe(process) + listener, err := (&net.ListenConfig{}).Listen(process, "tcp", "127.0.0.1:0") + require.NoError(t, err) + + r := &replicationSmokeReplica{a: a, ctx: process, cancel: stop, endpoint: "https://" + listener.Addr().String(), listener: &replicationSmokeListener{Listener: listener}, done: make(chan error, 1)} + s.replicas = append(s.replicas, r) + config := a.Server.tlsConfig(process, f.serverCertificate) + + go func() { r.done <- a.Server.serve(process, r.listener, config) }() + + t.Cleanup(func() { + stop() + + select { + case err := <-r.done: + if err != nil { + t.Error(err) + } + case <-time.After(12 * time.Second): + t.Error("replica shutdown exceeded bound") + } + }) + } +} + +func (s *replicationSmoke) authorizeReplicas(t *testing.T) { + t.Helper() + // Test-local leader selection and TokenReview, but the production TLS route, + // controller Pod/SA checks, bounded decoder and durable install are exercised. + f, leader := s.fixture, s.replicas[0] + leader.a.Replication.leader = leader.ctx + sa := &corev1.ServiceAccount{ObjectMeta: metav1.ObjectMeta{Namespace: f.a.Topology.Config.Namespace, Name: f.a.Topology.Config.ControllerServiceAccount, UID: "smoke-controller-sa"}} + require.NoError(t, f.a.Topology.Create(s.ctx, sa)) + + tokens := map[string]*corev1.Pod{} + + for i := 1; i < len(s.replicas); i++ { + pod := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: sa.Namespace, Name: fmt.Sprintf("controller-%d", i), UID: types.UID(fmt.Sprintf("controller-uid-%d", i))}, Spec: corev1.PodSpec{ServiceAccountName: sa.Name}} + require.NoError(t, f.a.Topology.Create(s.ctx, pod)) + tokens[f.token+strconv.Itoa(i)] = pod + } + + leader.a.Replication.Client = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Create: func(_ context.Context, _ client.WithWatch, obj client.Object, _ ...client.CreateOption) error { + review, ok := obj.(*authv1.TokenReview) + if !ok { + return fmt.Errorf("unexpected API create %T", obj) + } + + s.reviews.Add(1) + + pod := tokens[review.Spec.Token] + if pod == nil || !slices.Equal(review.Spec.Audiences, []string{ReplicationAudience}) { + return nil + } + + review.Status = authv1.TokenReviewStatus{Authenticated: true, Audiences: []string{ReplicationAudience}, User: authv1.UserInfo{Username: "system:serviceaccount:" + sa.Namespace + ":" + sa.Name, UID: string(sa.UID), Extra: map[string]authv1.ExtraValue{"authentication.kubernetes.io/pod-name": {pod.Name}, "authentication.kubernetes.io/pod-uid": {string(pod.UID)}}}} + + return nil + }}) + fixtureDependencies[leader.a.authority].Client = leader.a.Replication.Client +} + +func (s *replicationSmoke) install(t *testing.T, publication *CommittedPublication) { + t.Helper() + + image, err := wire.DecodePublication(strings.NewReader(publication.encoded)) + require.NoError(t, err) + + leader := s.replicas[0] + require.NoError(t, leader.a.Replication.installReplica(s.ctx, leader.ctx, image)) +} + +func (s *replicationSmoke) replicate(t *testing.T, publication *CommittedPublication) { + t.Helper() + + start := time.Now() + + for i, r := range s.replicas { + if i == 0 || r.ctx.Err() != nil { + continue + } + + request, err := http.NewRequestWithContext(s.ctx, http.MethodGet, s.replicas[0].endpoint+ReplicationPath, nil) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+s.fixture.token+strconv.Itoa(i)) + response, err := s.internal.Do(request) + require.NoError(t, err) + + image, decodeErr := wire.DecodePublication(response.Body) + closeErr := response.Body.Close() + require.Equal(t, http.StatusOK, response.StatusCode) + require.NoError(t, decodeErr) + require.NoError(t, closeErr) + require.NoError(t, r.a.Replication.installReplica(s.ctx, r.ctx, image)) + _, err = r.a.authority.Current() + require.NoError(t, err) + require.Equal(t, publication.encoded, capturePublication(t, r.a.authority).encoded, "replica diverged") + } + + t.Logf("replication sequence=%d elapsed=%s mock_token_reviews=%d", publication.record.Sequence, time.Since(start), s.reviews.Load()) +} + +func (s *replicationSmoke) observe() { + // Keep production freshness bounds; observations use fake Kubernetes reads. + ticker := time.NewTicker(5 * time.Second) + defer ticker.Stop() + + for { + select { + case <-s.ctx.Done(): + return + case <-ticker.C: + for _, r := range s.replicas { + if r.ctx.Err() == nil { + r.a.Replication.observe(r.ctx) + } + } + } + } +} + +func (s *replicationSmoke) run(after wire.Sequence, cold bool) <-chan replicationSmokeResult { + results := make(chan replicationSmokeResult, len(s.peers)) + slots := make(chan struct{}, 24) + + for i := range s.peers { + s.workers.Go(func() { + if cold { + select { + case slots <- struct{}{}: + case <-s.ctx.Done(): + results <- replicationSmokeResult{index: i, err: s.ctx.Err()} + return + } + + defer func() { <-slots }() + } + + start := time.Now() + result := s.peers[i].poll(s.ctx, s.replicas, i, after, &s.overloaded, &s.reconnected) + + result.elapsed = time.Since(start) + results <- result + }) + } + + return results +} + +func (s *replicationSmoke) collect(t *testing.T, phase string, started time.Time, results <-chan replicationSmokeResult, want *CommittedPublication) { + t.Helper() + + count := len(s.peers) + digest := sha256.Sum256([]byte(want.encoded)) + latencies := make([]time.Duration, 0, count) + + var total int64 + + for range s.peers { + select { + case result := <-results: + if result.err != nil || result.digest != digest || result.bytes != int64(len(want.encoded)) { + t.Fatalf("%s client=%d bytes=%d err=%v digest_match=%v", phase, result.index, result.bytes, result.err, result.digest == digest) + } + + total += result.bytes + latencies = append(latencies, result.elapsed) + case <-s.ctx.Done(): + t.Fatal(s.ctx.Err()) + } + } + + slices.Sort(latencies) + t.Logf("phase=%s clients=%d all_delivered=%s request_p50=%s request_p99=%s bytes=%d cumulative_429=%d reconnects=%d", phase, count, time.Since(started), latencies[len(latencies)/2], latencies[(len(latencies)-1)*99/100], total, s.overloaded.Load(), s.reconnected.Load()) + replicationSmokeStats(t, phase, s.replicas) +} + +func replicationSmokeLifecycle(l *Lifecycle, ctx context.Context) { l.process, l.synced = ctx, true } + +func (p *replicationSmokePeer) poll(ctx context.Context, replicas []*replicationSmokeReplica, index int, after wire.Sequence, overloaded, reconnected *atomic.Int64) replicationSmokeResult { + result := replicationSmokeResult{index: index} + + for attempt := range 6 { + r := replicas[p.replica] + + status, reconnect, err := p.readSnapshot(ctx, r, after, &result) + if reconnect { + p.replica = index % 2 + + reconnected.Add(1) + + continue + } + + if err != nil { + result.err = err + return result + } + + if status != http.StatusTooManyRequests { + return result + } + + overloaded.Add(1) + + if !replicationSleep(ctx, time.Second+time.Duration((index*31+attempt*97)%900)*time.Millisecond) { + break + } + } + + result.err = fmt.Errorf("six-attempt request budget exhausted") + + return result +} + +func (p *replicationSmokePeer) readSnapshot(ctx context.Context, replica *replicationSmokeReplica, after wire.Sequence, result *replicationSmokeResult) (int, bool, error) { + path := replica.endpoint + wire.SnapshotPath + if after != 0 { + path += fmt.Sprintf("?after=%d", after) + } + + request, err := http.NewRequestWithContext(ctx, http.MethodGet, path, nil) + if err != nil { + return 0, false, err + } + + response, err := p.client.Do(request) + if err != nil { + return 0, replica.ctx.Err() != nil && ctx.Err() == nil, err + } + + hash := sha256.New() + result.bytes, err = io.Copy(hash, io.LimitReader(response.Body, wire.MaxPublicationBytes+1)) + + closeErr := response.Body.Close() + if replica.ctx.Err() != nil && ctx.Err() == nil && (err != nil || response.StatusCode == http.StatusServiceUnavailable) { + return response.StatusCode, true, nil + } + + if err != nil || closeErr != nil { + return response.StatusCode, false, fmt.Errorf("body read=%v close=%v", err, closeErr) + } + + if response.StatusCode == http.StatusTooManyRequests { + return response.StatusCode, false, nil + } + + if response.StatusCode != http.StatusOK || response.TLS == nil || response.TLS.Version != tls.VersionTLS13 { + return response.StatusCode, false, fmt.Errorf("HTTP status %d or missing TLS 1.3", response.StatusCode) + } + + copy(result.digest[:], hash.Sum(nil)) + + return response.StatusCode, false, nil +} + +func replicationSmokePublish(t *testing.T, ctx context.Context, r *TopologyReconciler, accepted AcceptedMembers) *CommittedPublication { + t.Helper() + + _, err := r.authority.PublishTopology(ctx, func(context.Context) (TopologyObservation, error) { + nodes := corev1.NodeList{} + + for id, member := range accepted { + encoded, err := json.Marshal(member) + if err != nil { + return TopologyObservation{}, err + } + + nodes.Items = append(nodes.Items, corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: string(id), UID: types.UID(id), Annotations: map[string]string{admittedMemberAnnotation: string(encoded), wire.SharesAnnotation: strconv.FormatUint(uint64(member.Shares), 10)}}}) + } + + return TopologyObservation{Nodes: nodes, Input: members.Input{Nodes: nodes.Items, PeerPort: r.Config.PeerPort}}, nil + }) + require.NoError(t, err) + + return capturePublication(t, r.authority) +} + +func replicationSmokePark(t *testing.T, ctx context.Context, replicas []*replicationSmokeReplica, count int, results <-chan replicationSmokeResult) { + t.Helper() + + deadline, cancel := context.WithTimeout(ctx, 20*time.Second) + defer cancel() + + for { + select { + case result := <-results: + t.Fatalf("client %d completed before publication: %v", result.index, result.err) + default: + } + + total := 0 + for _, r := range replicas { + total += r.a.Server.polls.count() + } + + if total == count { + return + } + + if !replicationSleep(deadline, 10*time.Millisecond) { + t.Fatalf("parked %d/%d real HTTPS requests before deadline", total, count) + } + } +} + +func replicationSmokeStats(t *testing.T, phase string, replicas []*replicationSmokeReplica) { + t.Helper() + + var mem runtime.MemStats + runtime.ReadMemStats(&mem) + + var polls, live, accepted []int64 + for _, r := range replicas { + polls = append(polls, int64(r.a.Server.polls.count())) + live = append(live, r.listener.live.Load()) + accepted = append(accepted, r.listener.accepted.Load()) + } + + status, _ := os.ReadFile("/proc/self/status") + + var resource []string + + for _, line := range strings.Split(string(status), "\n") { + if strings.HasPrefix(line, "VmRSS:") || strings.HasPrefix(line, "VmHWM:") || strings.HasPrefix(line, "Threads:") { + resource = append(resource, strings.TrimSpace(line)) + } + } + + fds, _ := os.ReadDir("/proc/self/fd") + t.Logf("resources phase=%s polls=%v live_tcp=%v accepted_tcp=%v heap=%d stack=%d total_alloc=%d goroutines=%d fd=%d process=%v", phase, polls, live, accepted, mem.HeapAlloc, mem.StackInuse, mem.TotalAlloc, runtime.NumGoroutine(), len(fds), resource) +} diff --git a/internal/racer/server/server.go b/internal/racer/server/server.go new file mode 100644 index 000000000..8bf0126f3 --- /dev/null +++ b/internal/racer/server/server.go @@ -0,0 +1,1837 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// Package server serves Racer HTTPS endpoints from locally validated authority. +package server + +import ( + "bytes" + "context" + "crypto/tls" + "crypto/x509" + "encoding/json" + "encoding/pem" + "errors" + "fmt" + "io" + "mime" + "net" + "net/http" + "os" + "path/filepath" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/manager" + + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +type Server struct { + authority *authority.Authority + writer client.Writer + config Config + Lifecycle *Lifecycle + Leader Leader + polls *identityAdmission[wire.NodeID] + keyringPolls *identityAdmission[wire.NodeID] + replicationPolls *identityAdmission[string] + authSlots chan struct{} + bootstrapSlots chan struct{} + writes chan struct{} + servingCertificate atomic.Pointer[servingCertificateReloader] +} + +var ( + _ manager.Runnable = (*Server)(nil) + _ manager.LeaderElectionRunnable = (*Server)(nil) +) + +// Config contains only the HTTPS serving inputs, copied by New. +type Config struct { + ControlAddress string + TLSCertificateFile string + TLSPrivateKeyFile string + ReplicationServerName string + Limits Limits +} + +// Validate preserves controller validation: listener paths are checked by Start, +// and the replication client validates its server name before manager startup. +func (c Config) Validate() error { + if c.Limits.MaxConnections <= 0 || c.Limits.MaxConcurrentHandshakes <= 0 || + c.Limits.MaxPolls <= 0 || c.Limits.MaxConcurrentWrites <= 0 || c.Limits.MaxConcurrentBootstrap <= 0 || + c.Limits.HeaderBytes <= 0 || c.Limits.HandshakeTimeout <= 0 || c.Limits.WriteTimeout <= 0 || c.Limits.ShutdownTimeout <= 0 { + return fmt.Errorf("resource names or limits: %w", wire.InvalidRequest) + } + + return nil +} + +type Limits struct { + MaxConnections int + MaxConcurrentHandshakes int + MaxPolls int + MaxConcurrentWrites int + MaxConcurrentBootstrap int + HeaderBytes int + HandshakeTimeout time.Duration + WriteTimeout time.Duration + ShutdownTimeout time.Duration +} + +// Leader is the publisher view needed by the internal replication endpoint. +type Leader interface { + LeaderContext() (context.Context, bool) + PollInterval() time.Duration + AuthenticateReplica(context.Context, *http.Request) (string, time.Time, error) +} + +const ReplicationPath = "/internal/v1/snapshot" + +// New composes serving dependencies without starting work or granting readiness. +func New(cfg Config, writer client.Writer, auth *authority.Authority, lifecycle *Lifecycle, leader Leader) *Server { + return &Server{ + config: cfg, writer: writer, authority: auth, Lifecycle: lifecycle, Leader: leader, + polls: newIdentityAdmission[wire.NodeID](cfg.Limits.MaxPolls), + keyringPolls: newIdentityAdmission[wire.NodeID](cfg.Limits.MaxPolls), + replicationPolls: newIdentityAdmission[string](cfg.Limits.MaxConcurrentBootstrap), + authSlots: make(chan struct{}, max(0, cfg.Limits.MaxConcurrentBootstrap)), + // API-backed bearer authentication must not starve local TLS authentication. + // Enrollment and keyring bearer checks share this bounded API work pool. + bootstrapSlots: make(chan struct{}, max(0, cfg.Limits.MaxConcurrentBootstrap)), + writes: make(chan struct{}, max(0, cfg.Limits.MaxConcurrentWrites)), + } +} + +func (*Server) NeedLeaderElection() bool { return false } + +// TLSConfig must use VerifyClientCertIfGiven: bootstrap can omit the client +// certificate, while snapshot explicitly requires VerifiedChains. Resumption +// and pooled requests must not extend certificate validity or stale trust. +// The caller owns the reload lifetime and must supply and cancel a cancelable +// context, including when listener setup fails. Start owns this context itself. +func (s *Server) TLSConfig(ctx context.Context) (*tls.Config, error) { + if err := s.config.Validate(); err != nil { + return nil, err + } + + if s.authority == nil { + return nil, wire.Unavailable + } + + if ctx.Done() == nil { + return nil, wire.InvalidRequest + } + + if err := ctx.Err(); err != nil { + return nil, err + } + + reloader, err := newServingCertificateReloader(s.config.TLSCertificateFile, s.config.TLSPrivateKeyFile) + if err != nil { + return nil, wire.Unavailable + } + + go reloader.run(ctx, servingCertificatePollInterval) + + s.servingCertificate.Store(reloader) + + return s.tlsConfigWithCertificate(ctx, reloader.getCertificate), nil +} + +func (s *Server) tlsConfigWithCertificate(ctx context.Context, certificate func(*tls.ClientHelloInfo) (*tls.Certificate, error)) *tls.Config { + base := &tls.Config{MinVersion: tls.VersionTLS13, GetCertificate: certificate, ClientAuth: tls.VerifyClientCertIfGiven, SessionTicketsDisabled: true, NextProtos: []string{"http/1.1"}} + base.GetConfigForClient = func(_ *tls.ClientHelloInfo) (*tls.Config, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + + // Bound trust-pool work independently of the full handshake admission + // held by transportListener before any TLS bytes are read. + if !take(s.authSlots) { + return nil, wire.Overloaded + } + defer release(s.authSlots) + + roots, err := s.authority.TrustPool() + if err != nil { + // Replication uses a bearer token, not a dataplane certificate. Allow + // TLS startup before issuer trust exists to avoid bootstrap deadlock. + roots = x509.NewCertPool() + } + + cfg := base.Clone() + cfg.GetConfigForClient = nil + cfg.ClientCAs = roots + + return cfg, nil + } + + return base +} + +// Start opens TLS before public readiness so controller replication cannot +// deadlock on bootstrap. Process cancellation closes connections and polls. +func (s *Server) Start(ctx context.Context) error { + if err := s.config.Validate(); err != nil { + return err + } + + if s.config.ControlAddress == "" || s.config.TLSCertificateFile == "" || s.config.TLSPrivateKeyFile == "" { + return wire.InvalidRequest + } + + if s.Lifecycle == nil || s.authority == nil { + return wire.Unavailable + } + + serving, cancel := context.WithCancel(ctx) + defer cancel() + + config, err := s.TLSConfig(serving) + if err != nil { + return err + } + + listener, err := (&net.ListenConfig{}).Listen(serving, "tcp", s.config.ControlAddress) + if err != nil { + return err + } + + return s.serve(serving, listener, config) +} + +// serve owns the listener and every accepted connection. Close, rather than a +// grace period for active traffic, is required as soon as the process stops. +func (s *Server) serve(ctx context.Context, listener net.Listener, config *tls.Config) error { + ctx, cancel := context.WithCancel(ctx) + defer cancel() + + server := &http.Server{ + Handler: s.Handler(), + TLSConfig: config, + ReadHeaderTimeout: s.config.Limits.WriteTimeout, + ReadTimeout: s.config.Limits.WriteTimeout, + WriteTimeout: wire.PollWait + 3*s.config.Limits.WriteTimeout, + IdleTimeout: wire.PollWait, + MaxHeaderBytes: s.config.Limits.HeaderBytes, + BaseContext: func(net.Listener) context.Context { return ctx }, + } + + server.ConnContext = connectionContext + transport := newTransportListener(ctx, listener, config, s.config.Limits) + + done := make(chan error, 1) + + go func(done chan<- error) { done <- server.Serve(transport) }(done) + + s.Lifecycle.SetServingReady(true) + + var result error + + select { + case err := <-done: + result = err + done = nil + case <-ctx.Done(): + } + + s.Lifecycle.SetServingReady(false) + + servingCanceled := ctx.Err() != nil + + cancel() + + shutdown, stop := context.WithTimeout(context.Background(), s.config.Limits.ShutdownTimeout) + defer stop() + + closed := make(chan error, 1) + + go func(closed chan<- error) { + // Force-close TCP before net/http closes TLS connections: close-notify + // can otherwise block on a slow reader. No graceful drain is allowed. + closed <- errors.Join(transport.Close(), server.Close()) + }(closed) + + var closeErr error + + for done != nil || closed != nil { + select { + case result = <-done: + done = nil + case closeErr = <-closed: + closed = nil + case <-shutdown.Done(): + closeErr = errors.Join(closeErr, shutdown.Err()) + done, closed = nil, nil + } + } + + if errors.Is(result, http.ErrServerClosed) || servingCanceled && errors.Is(result, net.ErrClosed) { + result = nil + } + + return errors.Join(result, closeErr) +} + +func take(slots chan struct{}) bool { + select { + case slots <- struct{}{}: + return true + default: + return false + } +} + +func release(slots chan struct{}) { <-slots } + +func (s *Server) Ready(r *http.Request) error { + if s.Lifecycle == nil { + return wire.Unavailable + } + + reloader := s.servingCertificate.Load() + if reloader == nil { + return wire.Unavailable + } + + certificate, err := reloader.getCertificate(nil) + if err != nil { + return err + } + + // A valid chain for a different service cannot serve controller clients. + // The reloader has already parsed and validated this immutable leaf. + if certificate.Leaf == nil || certificate.Leaf.VerifyHostname(s.config.ReplicationServerName) != nil { + return wire.Unavailable + } + + if err := s.authority.TrustReady(); err != nil { + return err + } + + return s.Lifecycle.Ready(r) +} + +// Handler uses exact paths/methods without ServeMux redirects or implicit HEAD. +func (s *Server) Handler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + responseControl(http.NewResponseController(w).SetWriteDeadline(time.Now().Add(s.config.Limits.WriteTimeout))) + + if s.config.Limits.HeaderBytes > 0 && requestHeaderBytes(r) > s.config.Limits.HeaderBytes { + writeFailure(w, wire.TooLarge) + return + } + + if r.URL.EscapedPath() != r.URL.Path { + writeFailure(w, wire.InvalidRequest) + return + } + + handler, replication := s.route(r) + if handler == nil { + writeFailure(w, wire.InvalidRequest) + return + } + + if replication { + if r.TLS == nil || !r.TLS.HandshakeComplete { + writeFailure(w, wire.Unauthenticated) + return + } + + handler(w, r) + + return + } + + if s.Ready(r) != nil { + writeFailure(w, wire.Unavailable) + return + } + + if r.TLS == nil || !r.TLS.HandshakeComplete { + writeFailure(w, wire.Unauthenticated) + return + } + + ctx, cancel := s.Lifecycle.ProcessContext(r.Context()) + defer cancel() + + stop := context.AfterFunc(ctx, func() { + if conn, ok := ctx.Value(connectionKey{}).(net.Conn); ok { + closeTransport(conn) + } + }) + defer stop() + + handler(w, r.WithContext(ctx)) + }) +} + +func (s *Server) route(r *http.Request) (http.HandlerFunc, bool) { + switch { + case r.Method == http.MethodGet && r.URL.Path == ReplicationPath && s.Leader != nil: + return s.serveReplication, true + case r.Method == http.MethodPost && r.URL.Path == wire.BootstrapPath: + return s.serveBootstrap, false + case r.Method == http.MethodGet && r.URL.Path == wire.SnapshotPath: + return s.serveSnapshot, false + case r.Method == http.MethodGet && r.URL.Path == wire.KeyringPath: + return s.serveKeyring, false + default: + return nil, false + } +} + +// net/http permits parser slop above MaxHeaderBytes. Apply the application +// bound before authentication or API work. +func requestHeaderBytes(r *http.Request) int { + size := len(r.RequestURI) + len(r.Host) + for key, values := range r.Header { + for _, value := range values { + size += len(key) + len(value) + 4 + } + } + + return size +} + +func flushResponse(ctx context.Context, w http.ResponseWriter) { + if ctx.Err() != nil { + panic(http.ErrAbortHandler) + } + + responseControl(http.NewResponseController(w).Flush()) + + if ctx.Err() != nil { + panic(http.ErrAbortHandler) + } +} + +func flushAdmitted(ctx context.Context, guard *authority.Admission, w http.ResponseWriter) { + if guard.Check(ctx) != nil { + panic(http.ErrAbortHandler) + } + + flushResponse(ctx, w) + + if guard.Check(ctx) != nil { + panic(http.ErrAbortHandler) + } +} + +// The production HTTP/1 server supports these operations. In-memory handler +// recorders have no connection deadline; all actual connection errors abort it. +func responseControl(err error) { + if err != nil && !errors.Is(err, http.ErrNotSupported) { + panic(http.ErrAbortHandler) + } +} + +type connectionKey struct{} + +func connectionContext(ctx context.Context, conn net.Conn) context.Context { + return context.WithValue(ctx, connectionKey{}, conn) +} + +func closeTransport(conn net.Conn) { + if secured, ok := conn.(*tls.Conn); ok { + conn = secured.NetConn() + } + + if err := conn.Close(); err != nil { + return + } // Already closed is harmless. +} + +// TLS Close can attempt close-notify with its own deadline. Closing the raw +// transport at our deadline prevents that path from extending write admission. +func boundConnection(ctx context.Context, deadline time.Time) func() { + conn, ok := ctx.Value(connectionKey{}).(net.Conn) + if !ok { + return func() {} + } + + timer := time.AfterFunc(time.Until(deadline), func() { closeTransport(conn) }) + stop := context.AfterFunc(ctx, func() { closeTransport(conn) }) + + return func() { timer.Stop(); stop() } +} + +func minTime(a, b time.Time) time.Time { + if a.Before(b) { + return a + } + + return b +} + +func writeFailure(w http.ResponseWriter, err error) { + code := wire.Unavailable + + var protocol wire.ErrorCode + if errors.As(err, &protocol) { + code = protocol + } + + var status int + + switch code { + case wire.InvalidRequest: + status = http.StatusBadRequest + case wire.Unauthenticated: + status = http.StatusUnauthorized + case wire.Forbidden: + status = http.StatusForbidden + case wire.Conflict: + status = http.StatusConflict + case wire.TooLarge: + status = http.StatusRequestEntityTooLarge + case wire.UnsupportedVersion: + status = http.StatusUpgradeRequired + case wire.Overloaded: + status = http.StatusTooManyRequests + default: + code, status = wire.Unavailable, http.StatusServiceUnavailable + } + + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Cache-Control", "no-store") + + if code == wire.Unavailable || code == wire.Overloaded { + w.Header().Set("Retry-After", "1") + } + + w.WriteHeader(status) + + encoded, encodeErr := wire.EncodeError(wire.ErrorResponse{Code: code}) + if encodeErr != nil { + panic(http.ErrAbortHandler) + } + + if _, err := w.Write(encoded); err != nil { + return + } + + responseControl(http.NewResponseController(w).Flush()) +} + +// identityAdmission retains identities through response writes and flushes. +// Each endpoint owns an independent set and fixes its capacity at construction. +type identityAdmission[T comparable] struct { + mu sync.Mutex + limit int + held map[T]struct{} +} + +func newIdentityAdmission[T comparable](limit int) *identityAdmission[T] { + return &identityAdmission[T]{limit: limit, held: make(map[T]struct{})} +} + +func (a *identityAdmission[T]) acquire(id T) bool { + a.mu.Lock() + defer a.mu.Unlock() + + if _, exists := a.held[id]; exists || len(a.held) >= a.limit { + return false + } + + a.held[id] = struct{}{} + + return true +} + +func (a *identityAdmission[T]) release(id T) { + a.mu.Lock() + defer a.mu.Unlock() + + delete(a.held, id) +} + +func (a *identityAdmission[T]) count() int { + a.mu.Lock() + defer a.mu.Unlock() + + return len(a.held) +} + +// Lifecycle owns process serving, independently of the leader-owned publishers. +type Lifecycle struct { + authority *authority.Authority + mu sync.Mutex + process context.Context + synced bool + serving bool + waitForCacheSync func(context.Context) bool +} + +func NewLifecycle(a *authority.Authority) *Lifecycle { + return &Lifecycle{authority: a} +} + +func (*Lifecycle) NeedLeaderElection() bool { return false } + +// ProcessContext binds a request to the serving process lifetime. +// Missing or canceled process lifetime returns an already-canceled child. +func (l *Lifecycle) ProcessContext(parent context.Context) (context.Context, context.CancelFunc) { + ctx, cancel := context.WithCancel(parent) + if l == nil { + cancel() + return ctx, cancel + } + + l.mu.Lock() + process := l.process + l.mu.Unlock() + + if process == nil { + cancel() + return ctx, cancel + } + + stop := context.AfterFunc(process, cancel) + if process.Err() != nil { + cancel() + } + + return ctx, func() { stop(); cancel() } +} + +func (l *Lifecycle) Start(ctx context.Context) error { + l.mu.Lock() + if l.process != nil { + l.mu.Unlock() + return wire.Conflict + } + + l.process = ctx + l.authority.BindProcess(ctx) + l.mu.Unlock() + + defer func() { + l.mu.Lock() + l.synced, l.serving = false, false + l.mu.Unlock() + }() + + if l.waitForCacheSync == nil || !l.waitForCacheSync(ctx) { + if ctx.Err() != nil { + return nil + } + + return wire.Unavailable + } + + l.mu.Lock() + l.synced = ctx.Err() == nil + l.mu.Unlock() + <-ctx.Done() + + return nil +} + +// SetServingReady is set only after the authenticated listener is accepting. +func (l *Lifecycle) SetServingReady(ready bool) { + l.mu.Lock() + l.serving = ready + l.mu.Unlock() +} + +func (l *Lifecycle) Ready(_ *http.Request) error { + l.mu.Lock() + defer l.mu.Unlock() + + if l.process == nil || l.process.Err() != nil || !l.synced || !l.serving { + return wire.Unavailable + } + + return l.authority.PublicationReady() +} + +// SetCacheSync supplies the cache barrier before Start is called. +func (l *Lifecycle) SetCacheSync(wait func(context.Context) bool) { + l.waitForCacheSync = wait +} + +func (s *Server) serveBootstrap(w http.ResponseWriter, r *http.Request) { + if !take(s.bootstrapSlots) { + writeFailure(w, wire.Overloaded) + return + } + defer release(s.bootstrapSlots) + + ctx, cancel := context.WithTimeout(r.Context(), s.config.Limits.WriteTimeout) + defer cancel() + + deadline, _ := ctx.Deadline() + + stopWrite := boundConnection(ctx, deadline) + defer stopWrite() + + responseControl(http.NewResponseController(w).SetReadDeadline(deadline)) + responseControl(http.NewResponseController(w).SetWriteDeadline(deadline)) + + request, err := bootstrapRequest(r) + if err != nil { + writeFailure(w, err) + return + } + + trust, cancelTrust, err := s.authority.AdmitTrust(ctx) + if err != nil { + writeFailure(w, err) + return + } + defer cancelTrust() + + trustCtx := trust.Context() + + trustDeadline, _ := trustCtx.Deadline() + + stopTrustWrite := boundConnection(trustCtx, trustDeadline) + defer stopTrustWrite() + + responseControl(http.NewResponseController(w).SetWriteDeadline(trustDeadline)) + + encoded, err := s.enroll(trustCtx, r, request) + if err != nil { + writeFailure(w, err) + return + } + + if trust.Check(trustCtx) != nil || s.Ready(r) != nil { + writeFailure(w, wire.Unavailable) + return + } + + if !take(s.writes) { + writeFailure(w, wire.Overloaded) + return + } + defer release(s.writes) + + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Cache-Control", "no-store") + + if trust.Check(trustCtx) != nil { + panic(http.ErrAbortHandler) + } + + if _, err := w.Write(encoded); err != nil || trust.Check(trustCtx) != nil { + panic(http.ErrAbortHandler) + } + + flushAdmitted(trustCtx, trust, w) +} + +func bootstrapRequest(r *http.Request) (wire.BootstrapRequest, error) { + media, _, err := mime.ParseMediaType(r.Header.Get("Content-Type")) + if err != nil || media != "application/json" || r.URL.RawQuery != "" || r.URL.ForceQuery || r.Header.Get("Content-Encoding") != "" { + return wire.BootstrapRequest{}, wire.InvalidRequest + } + + if r.ContentLength > wire.MaxBootstrapBytes { + return wire.BootstrapRequest{}, wire.TooLarge + } + + return wire.DecodeBootstrap(r.Body) +} + +func (s *Server) enroll(ctx context.Context, r *http.Request, request wire.BootstrapRequest) ([]byte, error) { + response, hint, err := s.authority.EnrollWithHint(ctx, r, request) + if err != nil { + return nil, err + } + + ctx, cancel := context.WithDeadline(ctx, hint.Expires) + defer cancel() + + if err := annotateEnrollment(ctx, s.writer, hint); err != nil { + return nil, err + } + + return response, nil +} + +func annotateEnrollment(ctx context.Context, writer client.Writer, hint authority.EnrollmentHint) error { + node := hint.Node.DeepCopy() + value := strconv.FormatUint(uint64(hint.Shares), 10) + + nics, err := json.Marshal(hint.RDMANICs) + if err != nil { + return err + } + + nicValue := string(nics) + if len(hint.RDMANICs) == 0 { + nicValue = "" + } + + _, nicPresent := node.Annotations[members.EnrolledRDMANICsAnnotation] + if node.Annotations[members.EnrolledSharesAnnotation] == value && node.Annotations[members.EnrolledRDMANICsAnnotation] == nicValue && (nicValue != "" || !nicPresent) { + return nil + } + + before := node.DeepCopy() + if node.Annotations == nil { + node.Annotations = map[string]string{} + } + + node.Annotations[members.EnrolledSharesAnnotation] = value + if nicValue == "" { + delete(node.Annotations, members.EnrolledRDMANICsAnnotation) + } else { + node.Annotations[members.EnrolledRDMANICsAnnotation] = nicValue + } + + return writer.Patch(ctx, node, client.MergeFromWithOptions(before, client.MergeFromWithOptimisticLock{})) +} + +func (s *Server) authenticateSnapshot(ctx context.Context, state *tls.ConnectionState) (authority.NodeIdentity, error) { + if !take(s.authSlots) { + return authority.NodeIdentity{}, wire.Overloaded + } + defer release(s.authSlots) + + ctx, cancel := context.WithTimeout(ctx, s.config.Limits.WriteTimeout) + defer cancel() + + return s.authority.AuthenticateCertificate(ctx, state) +} + +func (s *Server) serveSnapshot(w http.ResponseWriter, r *http.Request) { + after, err := snapshotCursor(r) + if err != nil { + writeFailure(w, err) + return + } + + identity, err := s.authenticateSnapshot(r.Context(), r.TLS) + if err != nil { + writeFailure(w, err) + return + } + + if !s.polls.acquire(identity.Node()) { + writeFailure(w, wire.Overloaded) + return + } + defer s.polls.release(identity.Node()) + + ctx, cancel := context.WithDeadline(r.Context(), identity.Expires()) + defer cancel() + // A long poll does not consume a write slot. Give its eventual response a + // fresh bounded write window, capped by the verified chain's expiration. + responseControl(http.NewResponseController(w).SetWriteDeadline(minTime(identity.Expires(), time.Now().Add(wire.PollWait+s.config.Limits.WriteTimeout)))) + + publication, err := s.authority.Wait(ctx, identity, after) + if !time.Now().Before(identity.Expires()) { + err = wire.Unauthenticated + } + + if err != nil { + // Expiration forbids snapshot bytes, but a bounded error can still tell + // a pooled client to recover its expired identity through bootstrap. + responseControl(http.NewResponseController(w).SetWriteDeadline(time.Now().Add(s.config.Limits.WriteTimeout))) + writeFailure(w, err) + + return + } + + if ctx.Err() != nil || s.Ready(r) != nil { + writeFailure(w, wire.Unavailable) + return + } + + s.writeSnapshot(ctx, w, r, publication) +} + +func (s *Server) writeSnapshot(ctx context.Context, w http.ResponseWriter, r *http.Request, publication *authority.PublicationHandle) { + // Revalidate local trust after waiting: rotation or observed invalidity must + // also take effect on pooled connections before returning snapshot bytes. + trust, cancelTrust, err := s.authority.AdmitTrust(ctx) + if err != nil { + writeFailure(w, err) + return + } + defer cancelTrust() + + trustCtx := trust.Context() + + if _, err := s.authenticateSnapshot(trustCtx, r.TLS); err != nil { + writeFailure(w, err) + return + } + + if !take(s.writes) { + writeFailure(w, wire.Overloaded) + return + } + defer release(s.writes) + + image := publication + if image == nil { + image, err = s.authority.Current() + if err != nil { + writeFailure(w, err) + return + } + } + + boundedCtx, stopWindow := context.WithTimeout(trustCtx, s.config.Limits.WriteTimeout) + defer stopWindow() + + guard, cancelWrite, err := image.AdmitWithTrust(boundedCtx, trust) + if err != nil { + writeFailure(w, err) + return + } + defer cancelWrite() + + writeCtx := guard.Context() + + deadline, _ := writeCtx.Deadline() + + stopWrite := boundConnection(writeCtx, deadline) + defer stopWrite() + + responseControl(http.NewResponseController(w).SetWriteDeadline(deadline)) + w.Header().Set("Cache-Control", "no-store") + + if publication == nil { + w.WriteHeader(http.StatusNoContent) + flushAdmitted(writeCtx, guard, w) + + return + } + + w.Header().Set("Content-Type", "application/json") + + var sequence wire.Sequence + if after, err := snapshotCursor(r); err == nil && after != nil { + sequence = *after + } + + if _, err := publication.ForBase(sequence, r.Header.Get(wire.DeltaHeader)).WriteTo(writeCtx, guard, w); err != nil { + // A partial JSON response cannot be repaired with a protocol error. + panic(http.ErrAbortHandler) + } + + flushAdmitted(writeCtx, guard, w) +} + +func snapshotCursor(r *http.Request) (*wire.Sequence, error) { + if r.ContentLength != 0 || len(r.TransferEncoding) != 0 || r.URL.ForceQuery { + return nil, wire.InvalidRequest + } + + if r.URL.RawQuery == "" { + return nil, nil + } + + value, ok := strings.CutPrefix(r.URL.RawQuery, "after=") + if !ok { + return nil, wire.InvalidRequest + } + + n, err := strconv.ParseUint(value, 10, 64) + if err != nil || strconv.FormatUint(n, 10) != value { + return nil, wire.InvalidRequest + } + + if n == 0 { + return nil, wire.Conflict + } + + after := wire.Sequence(n) + + return &after, nil +} + +func keyringCursor(r *http.Request) (*wire.Generation, error) { + if r.Header.Get("Content-Encoding") != "" { + return nil, wire.InvalidRequest + } + // Both counters use the same canonical nonzero decimal query grammar. + after, err := snapshotCursor(r) + if err != nil || after == nil { + return nil, err + } + + generation := wire.Generation(*after) + + return &generation, nil +} + +func (s *Server) authenticateKeyring(r *http.Request) (authority.NodeIdentity, error) { + bearer := len(r.Header.Values("Authorization")) != 0 + + certificate := r.TLS != nil && (len(r.TLS.PeerCertificates) != 0 || len(r.TLS.VerifiedChains) != 0) + if bearer && certificate { + return authority.NodeIdentity{}, wire.Unauthenticated + } + + if !bearer { + return s.authenticateSnapshot(r.Context(), r.TLS) + } + + if s.authority == nil { + return authority.NodeIdentity{}, wire.Unavailable + } + + if !take(s.bootstrapSlots) { + return authority.NodeIdentity{}, wire.Overloaded + } + defer release(s.bootstrapSlots) + + ctx, cancel := context.WithTimeout(r.Context(), s.config.Limits.WriteTimeout) + defer cancel() + + return s.authority.Authenticate(ctx, r) +} + +func (s *Server) serveKeyring(w http.ResponseWriter, r *http.Request) { + after, err := keyringCursor(r) + if err != nil { + writeFailure(w, err) + return + } + + identity, err := s.authenticateKeyring(r) + if err != nil { + writeFailure(w, err) + return + } + + if !s.keyringPolls.acquire(identity.Node()) { + writeFailure(w, wire.Overloaded) + return + } + defer s.keyringPolls.release(identity.Node()) + + ctx, cancel := context.WithDeadline(r.Context(), identity.Expires()) + defer cancel() + + responseControl(http.NewResponseController(w).SetWriteDeadline(minTime(identity.Expires(), time.Now().Add(wire.PollWait+2*s.config.Limits.WriteTimeout)))) + + bundle, err := s.authority.WaitKeyring(ctx, after) + if !time.Now().Before(identity.Expires()) { + err = wire.Unauthenticated + } + + if err != nil { + responseControl(http.NewResponseController(w).SetWriteDeadline(time.Now().Add(s.config.Limits.WriteTimeout))) + writeFailure(w, err) + + return + } + + if ctx.Err() != nil || s.Ready(r) != nil { + writeFailure(w, wire.Unavailable) + return + } + + trust, cancelTrust, err := s.authority.AdmitTrust(ctx) + if err != nil { + writeFailure(w, err) + return + } + defer cancelTrust() + + s.writeKeyring(trust.Context(), trust, w, r, identity, after, bundle != nil) +} + +func (s *Server) reauthenticateKeyring(r *http.Request, identity authority.NodeIdentity) error { + verified, err := s.authenticateKeyring(r) + + if !time.Now().Before(identity.Expires()) { + return wire.Unauthenticated + } + + if err == nil && (verified.Node() != identity.Node() || verified.Cluster() != identity.Cluster()) { + return wire.Forbidden + } + + return err +} + +func (s *Server) writeKeyring(ctx context.Context, guard *authority.Admission, w http.ResponseWriter, r *http.Request, identity authority.NodeIdentity, after *wire.Generation, changed bool) { + accepted, err := s.authority.Keyring() + if err != nil { + writeFailure(w, err) + return + } + // Recheck local certificate trust or live bearer authorization after waiting. + // Never fall back from a rejected certificate to a bearer token. + if err := s.reauthenticateKeyring(r.WithContext(ctx), identity); err != nil { + responseControl(http.NewResponseController(w).SetWriteDeadline(time.Now().Add(s.config.Limits.WriteTimeout))) + writeFailure(w, err) + + return + } + // Authentication may have waited on the API. Do not deliver an old accepted + // encoding if reconciliation invalidated or replaced it in the meantime. + current, err := s.authority.Keyring() + if err != nil || current != accepted || guard.Check(ctx) != nil || s.Ready(r) != nil { + writeFailure(w, wire.Unavailable) + return + } + + changed = changed || after != nil && current.Generation() > *after + + if !take(s.writes) { + writeFailure(w, wire.Overloaded) + return + } + defer release(s.writes) + + deadline := minTime(identity.Expires(), time.Now().Add(s.config.Limits.WriteTimeout)) + if freshness, ok := ctx.Deadline(); ok { + deadline = minTime(deadline, freshness) + } + + ctx, cancel := context.WithDeadline(ctx, deadline) + defer cancel() + + stopWrite := boundConnection(ctx, deadline) + defer stopWrite() + + responseControl(http.NewResponseController(w).SetWriteDeadline(deadline)) + w.Header().Set("Cache-Control", "no-store") + + if !changed { + w.WriteHeader(http.StatusNoContent) + flushAdmitted(ctx, guard, w) + + return + } + + w.Header().Set("Content-Type", "application/json") + // Copy in bounded chunks so cancellation is observed between writes without + // allocating a bundle-sized byte slice for each request. + if _, err := current.Response().WriteTo(ctx, guard, w); err != nil { + panic(http.ErrAbortHandler) + } + + flushAdmitted(ctx, guard, w) +} + +func (s *Server) serveReplication(w http.ResponseWriter, request *http.Request) { + r := s.Leader + if _, ok := r.LeaderContext(); !ok { + writeFailure(w, wire.Unavailable) + return + } + + after, err := snapshotCursor(request) + if err != nil { + writeFailure(w, err) + return + } + + uid, expires, err := s.authenticateReplica(request) + if err != nil { + writeFailure(w, err) + return + } + + if !s.replicationPolls.acquire(uid) { + writeFailure(w, wire.Overloaded) + return + } + defer s.replicationPolls.release(uid) + + leader, _ := r.LeaderContext() + + ctx, cancel := context.WithDeadline(request.Context(), minTime(expires, time.Now().Add(r.PollInterval()))) + defer cancel() + + stop := context.AfterFunc(leader, cancel) + defer stop() + + responseControl(http.NewResponseController(w).SetWriteDeadline(time.Now().Add(r.PollInterval() + s.config.Limits.WriteTimeout))) + publication, err := s.waitReplica(ctx, after) + + if _, ok := r.LeaderContext(); !ok || !time.Now().Before(expires) { + writeFailure(w, wire.Unavailable) + return + } + + unchanged := errors.Is(err, context.DeadlineExceeded) + if unchanged { + publication, err = s.authority.Current() + } + + if err != nil { + writeFailure(w, err) + return + } + + s.writeReplica(w, request, leader, expires, publication, unchanged) +} + +func (s *Server) authenticateReplica(r *http.Request) (string, time.Time, error) { + if !take(s.bootstrapSlots) { + return "", time.Time{}, wire.Overloaded + } + defer release(s.bootstrapSlots) + + ctx, cancel := context.WithTimeout(r.Context(), s.config.Limits.WriteTimeout) + defer cancel() + + return s.Leader.AuthenticateReplica(ctx, r) +} + +func (s *Server) waitReplica(ctx context.Context, after *wire.Sequence) (*authority.PublicationHandle, error) { + for { + publication, changed, err := s.authority.CurrentAndSubscribe() + if err != nil || after == nil || publication.Sequence() > *after { + return publication, err + } + + select { + case <-ctx.Done(): + return publication, ctx.Err() + case <-changed: + } + } +} + +func (s *Server) writeReplica(w http.ResponseWriter, request *http.Request, leader context.Context, expires time.Time, publication *authority.PublicationHandle, unchanged bool) { + if !take(s.writes) { + writeFailure(w, wire.Overloaded) + return + } + defer release(s.writes) + + windowCtx, stopWrite := context.WithDeadline(request.Context(), minTime(expires, time.Now().Add(s.config.Limits.WriteTimeout))) + defer stopWrite() + + stopLeader := context.AfterFunc(leader, stopWrite) + defer stopLeader() + + guard, stopAuthority, err := publication.Admit(windowCtx) + if err != nil { + writeFailure(w, err) + return + } + defer stopAuthority() + + writeCtx := guard.Context() + + deadline, _ := writeCtx.Deadline() + + stopConnection := boundConnection(writeCtx, deadline) + defer stopConnection() + + responseControl(http.NewResponseController(w).SetWriteDeadline(deadline)) + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Cache-Control", "no-store") + + if unchanged { + w.WriteHeader(http.StatusNoContent) + flushAdmitted(writeCtx, guard, w) + + return + } + + if _, err := publication.ForBase(0, "").WriteTo(writeCtx, guard, w); err != nil { + panic(http.ErrAbortHandler) + } + + flushAdmitted(writeCtx, guard, w) +} + +// transportListener owns sockets from acceptance until close, including sockets +// not yet visible to net/http. Excess sockets are closed in the accept loop, +// without spawning a goroutine or queuing work. The kernel backlog is separate. +type transportListener struct { + net.Listener + config *tls.Config + timeout time.Duration + ctx context.Context + cancel context.CancelFunc + ready chan net.Conn + acceptDone chan struct{} + acceptErr error + connections chan struct{} + handshakes chan struct{} + metrics *transportMetrics + mu sync.Mutex + closed bool + live map[*transportConn]struct{} + closeOnce sync.Once + closeErr error +} + +func newTransportListener(ctx context.Context, listener net.Listener, config *tls.Config, limits Limits) *transportListener { + return newTransportListenerWithMetrics(ctx, listener, config, limits, servingTransportMetrics) +} + +func newTransportListenerWithMetrics(ctx context.Context, listener net.Listener, config *tls.Config, limits Limits, metrics *transportMetrics) *transportListener { + ctx, cancel := context.WithCancel(ctx) + + l := &transportListener{ + Listener: listener, config: config, timeout: limits.HandshakeTimeout, + ctx: ctx, cancel: cancel, ready: make(chan net.Conn), acceptDone: make(chan struct{}), + connections: make(chan struct{}, max(0, limits.MaxConnections)), + handshakes: make(chan struct{}, max(0, limits.MaxConcurrentHandshakes)), + live: make(map[*transportConn]struct{}), + metrics: metrics, + } + go l.run() + + return l +} + +func (l *transportListener) run() { + defer close(l.acceptDone) + + var retryDelay time.Duration + + for { + if l.ctx.Err() != nil { + l.acceptErr = net.ErrClosed + return + } + + conn, err := l.Listener.Accept() + if err != nil { + // Retry here, not in net/http: once this pump exits, Accept can + // only replay acceptErr and cannot resume accepting sockets. + var temporary net.Error + if errors.As(err, &temporary) && temporary.Temporary() { //nolint:staticcheck // Match net/http's accept-error retry contract, including EMFILE. + retryDelay = nextAcceptDelay(retryDelay) + if !waitAcceptRetry(l.ctx, retryDelay) { + l.acceptErr = net.ErrClosed + return + } + + continue + } + + l.acceptErr = err + + return + } + + retryDelay = 0 + + if !take(l.connections) { + if l.ctx.Err() == nil { + l.metrics.connectionRejected.Inc() + } + + closeTransport(conn) + + continue + } + + l.metrics.connections.Inc() + + c := &transportConn{Conn: conn, owner: l} + l.mu.Lock() + + closed := l.closed || l.ctx.Err() != nil + if !closed { + l.live[c] = struct{}{} + } + l.mu.Unlock() + + if closed { + closeTransport(c) + continue + } + + if !take(l.handshakes) { + if l.ctx.Err() == nil { + l.metrics.handshakeRejected.Inc() + } + + closeTransport(c) + + continue + } + + l.metrics.handshakes.Inc() + + go l.handshake(c) + } +} + +func nextAcceptDelay(previous time.Duration) time.Duration { + if previous == 0 { + return 5 * time.Millisecond + } + + return min(2*previous, time.Second) +} + +func waitAcceptRetry(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-timer.C: + return true + case <-ctx.Done(): + return false + } +} + +func (l *transportListener) handshake(raw *transportConn) { + // Only this goroutine releases handshake admission, after HandshakeContext + // actually returns, never from a TLS configuration/verification callback. + conn := tls.Server(raw, l.config) + ctx, cancel := context.WithTimeout(l.ctx, l.timeout) + + err := raw.SetDeadline(time.Now().Add(l.timeout)) + if err == nil { + err = conn.HandshakeContext(ctx) + } + + cancel() + + var timeout net.Error + if err != nil && (errors.Is(err, context.DeadlineExceeded) || errors.As(err, &timeout) && timeout.Timeout()) { + l.metrics.handshakeTimeouts.Inc() + } + + l.metrics.handshakes.Dec() + release(l.handshakes) + + if err != nil { + closeTransport(raw) + return + } + + if err := raw.SetDeadline(time.Time{}); err != nil { + closeTransport(raw) + return + } + // Return the concrete *tls.Conn, not a wrapper: net/http requires that type + // to populate Request.TLS and enforce client-certificate authentication. + select { + case l.ready <- conn: + case <-l.ctx.Done(): + closeTransport(raw) + } +} + +func (l *transportListener) Accept() (net.Conn, error) { + select { + case <-l.acceptDone: + return nil, l.acceptErr + case conn := <-l.ready: + return conn, nil + } +} + +func (l *transportListener) Close() error { + l.closeOnce.Do(func() { + l.cancel() + l.mu.Lock() + l.closed = true + + connections := make([]*transportConn, 0, len(l.live)) + for conn := range l.live { + connections = append(connections, conn) + } + l.mu.Unlock() + // Raw close bypasses TLS close-notify, including for handshakes and + // completed handshakes still waiting to be handed to net/http. + for _, conn := range connections { + closeTransport(conn) + } + + l.closeErr = l.Listener.Close() + }) + + return l.closeErr +} + +type transportConn struct { + net.Conn + owner *transportListener + once sync.Once + err error +} + +func (c *transportConn) Close() error { + c.once.Do(func() { + c.err = c.Conn.Close() + c.owner.mu.Lock() + delete(c.owner.live, c) + c.owner.mu.Unlock() + c.owner.metrics.connections.Dec() + release(c.owner.connections) + }) + + return c.err +} + +const servingCertificatePollInterval = time.Second + +type servingCertificate struct { + prefixes []servingCertificatePrefix +} + +type servingCertificatePrefix struct { + certificate tls.Certificate + notBefore, notAfter time.Time +} + +// Published certificates are immutable. Handshakes only load a pointer and check +// its validity window; all filesystem access and parsing happens in the poller. +type servingCertificateReloader struct { + certificateFile, keyFile string + current atomic.Pointer[servingCertificate] +} + +func newServingCertificateReloader(certificateFile, keyFile string) (*servingCertificateReloader, error) { + r := &servingCertificateReloader{certificateFile: certificateFile, keyFile: keyFile} + if err := r.reload(); err != nil { + return nil, err + } + + return r, nil +} + +func (r *servingCertificateReloader) getCertificate(*tls.ClientHelloInfo) (*tls.Certificate, error) { + return r.getCertificateAt(time.Now()) +} + +func (r *servingCertificateReloader) getCertificateAt(now time.Time) (*tls.Certificate, error) { + certificate := r.current.Load() + if certificate == nil { + return nil, wire.Unavailable + } + + return certificate.at(now) +} + +func (c *servingCertificate) at(now time.Time) (*tls.Certificate, error) { + // Longest first: retain compatibility until its suffix expires, then use + // an already validated prefix. No file reads, parsing, or verification here. + for i := range c.prefixes { + if prefix := &c.prefixes[i]; !now.Before(prefix.notBefore) && now.Before(prefix.notAfter) { + return &prefix.certificate, nil + } + } + + return nil, wire.Unavailable +} + +func (r *servingCertificateReloader) run(ctx context.Context, interval time.Duration) { + ticker := time.NewTicker(interval) + defer ticker.Stop() + + var nextLog time.Time + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if ctx.Err() != nil { + return + } + + if err := r.reload(); err != nil && !time.Now().Before(nextLog) { + // Do not log paths, PEM contents, or parser errors. Bound repeated + // failures even if a broken projection remains mounted indefinitely. + ctrl.LoggerFrom(ctx).Error(wire.Unavailable, "serving TLS certificate reload rejected; retaining last valid certificate") + + nextLog = time.Now().Add(time.Minute) + } + } + } +} + +func (r *servingCertificateReloader) reload() error { + return r.reloadAt(time.Now()) +} + +func (r *servingCertificateReloader) reloadAt(now time.Time) error { + certificatePEM, keyPEM, err := r.readPair() + if err != nil { + return err + } + + if err := validateCertificatePEM(certificatePEM); err != nil { + return err + } + + certificate, err := tls.X509KeyPair(certificatePEM, keyPEM) + if err != nil { + return wire.Unavailable + } + + validated, err := validateServingCertificate(certificate, now) + if err != nil { + return err + } + + r.current.Store(validated) + + return nil +} + +func (r *servingCertificateReloader) readPair() ([]byte, []byte, error) { + certificatePath, certificateGeneration, err := servingFilePath(r.certificateFile) + if err != nil { + return nil, nil, wire.Unavailable + } + + keyPath, keyGeneration, err := servingFilePath(r.keyFile) + if err != nil || certificateGeneration != keyGeneration { + return nil, nil, wire.Unavailable + } + + certificatePEM, err := readServingFile(certificatePath) + if err != nil { + return nil, nil, wire.Unavailable + } + + keyPEM, err := readServingFile(keyPath) + if err != nil { + return nil, nil, wire.Unavailable + } + // Recheck both names and contents. Projected paths must still name the same + // immutable generation; standalone files must be stable across both reads. + for _, file := range []struct { + name, path, generation string + contents []byte + }{ + {r.certificateFile, certificatePath, certificateGeneration, certificatePEM}, + {r.keyFile, keyPath, keyGeneration, keyPEM}, + } { + path, generation, err := servingFilePath(file.name) + if err != nil || path != file.path || generation != file.generation { + return nil, nil, wire.Unavailable + } + + contents, err := readServingFile(path) + if err != nil || !bytes.Equal(contents, file.contents) { + return nil, nil, wire.Unavailable + } + } + + return certificatePEM, keyPEM, nil +} + +func validateCertificatePEM(certificatePEM []byte) error { + // X509KeyPair tolerates trailing malformed PEM. Do not accidentally accept + // a truncated chain as a valid leaf-only deployment. + for remaining := bytes.TrimSpace(certificatePEM); len(remaining) > 0; { + if !bytes.HasPrefix(remaining, []byte("-----BEGIN CERTIFICATE-----")) { + return wire.Unavailable + } + + block, rest := pem.Decode(remaining) + if block == nil || block.Type != "CERTIFICATE" || len(block.Headers) != 0 { + return wire.Unavailable + } + + remaining = bytes.TrimSpace(rest) + } + + return nil +} + +// Kubernetes AtomicWriter projections have a ..data symlink at the volume root. +// Resolve each file and require both to belong to that same generation, including +// nested projected paths. Never follow a second generation while reading a pair. +func servingFilePath(name string) (string, string, error) { + abs, err := filepath.Abs(name) + if err != nil { + return "", "", err + } + + path, err := filepath.EvalSymlinks(abs) + if err != nil { + return "", "", err + } + + for dir := filepath.Dir(abs); ; dir = filepath.Dir(dir) { + data := filepath.Join(dir, "..data") + if _, err := os.Lstat(data); err == nil { + generation, err := filepath.EvalSymlinks(data) + if err != nil { + return "", "", err + } + + relative, err := filepath.Rel(generation, path) + if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) { + return "", "", wire.Unavailable + } + + return path, generation, nil + } else if !os.IsNotExist(err) { + return "", "", err + } + + if filepath.Dir(dir) == dir { + break + } + } + + return path, "", nil +} + +func readServingFile(path string) (data []byte, result error) { + // Reject special files before opening, and bound memory for malformed input. + info, err := os.Stat(path) + if err != nil || !info.Mode().IsRegular() { + return nil, wire.Unavailable + } + + f, err := os.Open(path) + if err != nil { + return nil, err + } + + defer func() { result = errors.Join(result, f.Close()) }() + + const maxBytes = 1024 * 1024 + + contents, err := io.ReadAll(io.LimitReader(f, maxBytes+1)) + if err != nil || len(contents) > maxBytes { + return nil, wire.Unavailable + } + + return contents, nil +} + +func validateServingCertificate(certificate tls.Certificate, now time.Time) (*servingCertificate, error) { + chain, err := parseServingChain(certificate.Certificate, now) + if err != nil { + return nil, err + } + + return validateServingChain(certificate, chain, now) +} + +func parseServingChain(certificates [][]byte, now time.Time) ([]*x509.Certificate, error) { + if len(certificates) == 0 { + return nil, wire.Unavailable + } + + chain := make([]*x509.Certificate, len(certificates)) + for i, der := range certificates { + cert, err := x509.ParseCertificate(der) + if err != nil || cert.NotBefore.After(now) || !cert.NotBefore.Before(cert.NotAfter) { + return nil, wire.Unavailable + } + + chain[i] = cert + if i > 0 && chain[i-1].CheckSignatureFrom(cert) != nil { + return nil, wire.Unavailable + } + } + + if !now.Before(chain[0].NotAfter) || chain[0].IsCA { + return nil, wire.Unavailable + } + + return chain, nil +} + +func validateServingChain(certificate tls.Certificate, chain []*x509.Certificate, now time.Time) (*servingCertificate, error) { + // Check the ENTIRE supplied chain's constraints before allowing any trim. + // Expired compatibility certificates need not overlap a newly issued leaf, + // so verify structure using copies with neutral validity windows. The actual + // windows above and in each cached prefix are enforced separately. Raw DER, + // signatures, EKU, name constraints, and path length constraints are unchanged. + structural := make([]*x509.Certificate, len(chain)) + for i, cert := range chain { + copy := *cert + copy.NotBefore, copy.NotAfter = now.Add(-time.Hour), now.Add(time.Hour) + structural[i] = © + } + + if err := verifyServingPrefix(structural, now); err != nil { + return nil, wire.Unavailable + } + + certificate.Leaf = chain[0] + validated := &servingCertificate{} + + for n := len(chain); n > 0; n-- { + // Only a cross-signed CA starts an optional compatibility suffix. Do + // not reinterpret an expired ordinary self-signed issuer as optional. + if n < len(chain) && (!chain[n].IsCA || chain[n].CheckSignatureFrom(chain[n]) == nil) { + continue + } + + if err := verifyServingPrefix(structural[:n], now); err != nil { + return nil, wire.Unavailable + } + + validated.prefixes = append(validated.prefixes, servingPrefix(certificate, chain, n)) + } + + if _, err := validated.at(now); err != nil { + return nil, err + } + + return validated, nil +} + +func servingPrefix(certificate tls.Certificate, chain []*x509.Certificate, n int) servingCertificatePrefix { + prefix := servingCertificatePrefix{certificate: certificate, notBefore: chain[0].NotBefore, notAfter: chain[0].NotAfter} + + prefix.certificate.Certificate = certificate.Certificate[:n:n] + for _, cert := range chain[:n] { + if cert.NotBefore.After(prefix.notBefore) { + prefix.notBefore = cert.NotBefore + } + + prefix.notAfter = minTime(prefix.notAfter, cert.NotAfter) + } + + if n < len(chain) { + // A shorter prefix is eligible only after its compatibility suffix + // expires, never to work around a not-yet-valid supplied certificate. + expires := chain[n].NotAfter + for _, cert := range chain[n:] { + expires = minTime(expires, cert.NotAfter) + if cert.NotBefore.After(prefix.notBefore) { + prefix.notBefore = cert.NotBefore + } + } + + if expires.After(prefix.notBefore) { + prefix.notBefore = expires + } + } + + return prefix +} + +// Serving trust is deployment-provided, distinct from peer trust. The last +// certificate in a prefix is the local validation anchor, not a client trust +// decision. Clients must still build a path to their own trusted current CA. +func verifyServingPrefix(chain []*x509.Certificate, now time.Time) error { + roots, intermediates := x509.NewCertPool(), x509.NewCertPool() + roots.AddCert(chain[len(chain)-1]) + + for _, cert := range chain[1:] { + intermediates.AddCert(cert) + } + + if _, err := chain[0].Verify(x509.VerifyOptions{Roots: roots, Intermediates: intermediates, CurrentTime: now, KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}}); err != nil { + return wire.Unavailable + } + + return nil +} diff --git a/internal/racer/server/server_test.go b/internal/racer/server/server_test.go new file mode 100644 index 000000000..4e6c936cd --- /dev/null +++ b/internal/racer/server/server_test.go @@ -0,0 +1,1965 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "bytes" + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/sha256" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "net/url" + "path/filepath" + "strconv" + "strings" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/types" + ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" + "sigs.k8s.io/controller-runtime/pkg/client/fake" + "sigs.k8s.io/controller-runtime/pkg/client/interceptor" + + racerv1 "github.com/Azure/unbounded/api/racer/v1alpha1" + "github.com/Azure/unbounded/internal/racer/authority" + "github.com/Azure/unbounded/internal/racer/members" + "github.com/Azure/unbounded/internal/racer/wire" +) + +var fixtureConfigs = map[*authority.Authority]fixtureConfig{} + +type ( + NodeIdentity = authority.NodeIdentity + TopologyObservation = authority.TopologyObservation + AcceptedMembers = members.History +) + +const ( + podNodeIndex = "spec.nodeName" + enrolledSharesAnnotation = members.EnrolledSharesAnnotation + enrolledRDMANICsAnnotation = members.EnrolledRDMANICsAnnotation + admittedMemberAnnotation = members.AdmittedMemberAnnotation + ReplicationAudience = authority.ReplicationAudience +) + +func podNodeKeys(obj client.Object) []string { + pod, ok := obj.(*corev1.Pod) + if !ok || pod.Spec.NodeName == "" { + return nil + } + + return []string{pod.Spec.NodeName} +} + +func replicationSleep(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func catalogCache(name string, uid types.UID) racerv1.ClusterCache { + return racerv1.ClusterCache{ObjectMeta: metav1.ObjectMeta{Name: name, UID: uid}} +} + +func TestIdentityAdmission(t *testing.T) { + a := newIdentityAdmission[string](1) + b := newIdentityAdmission[string](1) + + require.True(t, a.acquire("one")) + require.False(t, a.acquire("one")) + require.False(t, a.acquire("two")) + require.True(t, b.acquire("one"), "independent endpoint admission") + a.release("one") + require.True(t, a.acquire("two")) + require.False(t, newIdentityAdmission[string](0).acquire("one")) + require.False(t, newIdentityAdmission[string](-1).acquire("one")) +} + +func TestAdmissionLimitsFrozen(t *testing.T) { + cfg := testConfig(t).ServerConfig + cfg.Limits.MaxPolls = 1 + cfg.Limits.MaxConcurrentBootstrap = 1 + s := New(cfg, nil, nil, nil, nil) + cfg.Limits.MaxPolls = 2 + cfg.Limits.MaxConcurrentBootstrap = 2 + cfg.Limits.HeaderBytes = 1 + + require.Equal(t, 1, s.polls.limit) + require.Equal(t, 1, s.keyringPolls.limit) + require.Equal(t, cap(s.bootstrapSlots), s.replicationPolls.limit) + require.NotEqual(t, cfg.Limits.HeaderBytes, s.config.Limits.HeaderBytes) +} + +const ( + testNodeUID = "11111111-1111-4111-8111-111111111111" + testOtherUID = "22222222-2222-4222-8222-222222222222" + testDaemonSetUID types.UID = "33333333-3333-4333-8333-333333333333" +) + +func memberNode() corev1.Node { + return corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "node-a", UID: testNodeUID}} +} + +func memberPod(uid types.UID, created int64, ip string) corev1.Pod { + controller := true + + return corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "racer-" + string(uid), UID: uid, Namespace: "racer", + CreationTimestamp: metav1.NewTime(time.Unix(created, 0)), + OwnerReferences: []metav1.OwnerReference{{APIVersion: "apps/v1", Kind: "DaemonSet", Name: DataplaneDaemonSetName, UID: testDaemonSetUID, Controller: &controller}}, + }, + Spec: corev1.PodSpec{NodeName: "node-a"}, + Status: corev1.PodStatus{PodIP: ip}, + } +} + +const DataplaneDaemonSetName = "racer-dataplane" + +func TestHTTPRoutesBeforeReadinessAndTLS(t *testing.T) { + f := newServingFixture(t) + handler := f.a.Server.Handler() + + for _, ready := range []bool{false, true} { + t.Run(fmt.Sprintf("ready=%t", ready), func(t *testing.T) { + f.a.Lifecycle.SetServingReady(ready) + + for _, tc := range []struct { + method, path string + valid bool + }{ + {http.MethodPost, wire.BootstrapPath, true}, + {http.MethodGet, wire.SnapshotPath, true}, + {http.MethodGet, wire.BootstrapPath, false}, + {http.MethodPost, wire.SnapshotPath, false}, + {http.MethodHead, wire.BootstrapPath, false}, + {http.MethodHead, wire.SnapshotPath, false}, + {http.MethodOptions, wire.BootstrapPath, false}, + {http.MethodOptions, wire.SnapshotPath, false}, + {http.MethodPost, "/unknown", false}, + {http.MethodGet, "/unknown", false}, + {http.MethodPost, "/v1/%62ootstrap", false}, + {http.MethodGet, "/v1/%73napshot", false}, + } { + t.Run(tc.method+tc.path, func(t *testing.T) { + r := httptest.NewRequest(tc.method, tc.path, nil) + w := httptest.NewRecorder() + handler.ServeHTTP(w, r) + + wantStatus, wantCode := http.StatusBadRequest, wire.InvalidRequest + if tc.valid { + wantStatus, wantCode = http.StatusServiceUnavailable, wire.Unavailable + if ready { + wantStatus, wantCode = http.StatusUnauthorized, wire.Unauthenticated + } + } + + body := responseBody(t, w.Result(), nil, wantStatus) + if string(body) != fmt.Sprintf(`{"code":%q}`, wantCode) { + t.Fatalf("response %s, want code %s", body, wantCode) + } + }) + } + }) + } +} + +func TestHTTPFailureResponses(t *testing.T) { + for _, tc := range []struct { + name string + err error + status int + code wire.ErrorCode + retryAfter string + }{ + {"invalid request", wire.InvalidRequest, http.StatusBadRequest, wire.InvalidRequest, ""}, + {"unauthenticated", wire.Unauthenticated, http.StatusUnauthorized, wire.Unauthenticated, ""}, + {"forbidden", wire.Forbidden, http.StatusForbidden, wire.Forbidden, ""}, + {"conflict", wire.Conflict, http.StatusConflict, wire.Conflict, ""}, + {"too large", wire.TooLarge, http.StatusRequestEntityTooLarge, wire.TooLarge, ""}, + {"unsupported version", wire.UnsupportedVersion, http.StatusUpgradeRequired, wire.UnsupportedVersion, ""}, + {"overloaded", wire.Overloaded, http.StatusTooManyRequests, wire.Overloaded, "1"}, + {"unavailable", wire.Unavailable, http.StatusServiceUnavailable, wire.Unavailable, "1"}, + {"wrapped protocol error", fmt.Errorf("internal detail: %w", wire.Forbidden), http.StatusForbidden, wire.Forbidden, ""}, + {"unknown protocol error", wire.ErrorCode("unknown"), http.StatusServiceUnavailable, wire.Unavailable, "1"}, + {"nonprotocol error", io.ErrUnexpectedEOF, http.StatusServiceUnavailable, wire.Unavailable, "1"}, + {"nil error", nil, http.StatusServiceUnavailable, wire.Unavailable, "1"}, + } { + t.Run(tc.name, func(t *testing.T) { + w := httptest.NewRecorder() + writeFailure(w, tc.err) + + body := responseBody(t, w.Result(), nil, tc.status) + if string(body) != fmt.Sprintf(`{"code":%q}`, tc.code) { + t.Fatalf("response %s, want code %s", body, tc.code) + } + + if got := w.Header().Get("Retry-After"); got != tc.retryAfter { + t.Fatalf("Retry-After %q, want %q", got, tc.retryAfter) + } + + if w.Header().Get("Content-Type") != "application/json" || w.Header().Get("Cache-Control") != "no-store" || !w.Flushed { + t.Fatal("failure response headers or flush missing") + } + }) + } +} + +func (f *servingFixture) signLeaf(t *testing.T, mutate func(*x509.Certificate)) tls.Certificate { + t.Helper() + _, _, rotation, material := keyState(t, f.a.Keyring) + + ca, key, err := parseSigning(material.Keys[rotation.ActiveIssuer]) + if err != nil { + t.Fatal(err) + } + + template := &x509.Certificate{SerialNumber: big.NewInt(123), NotBefore: time.Now().Add(-time.Minute), NotAfter: time.Now().Add(time.Hour), KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, URIs: []*url.URL{{Scheme: "spiffe", Host: string(f.request.Cluster), Path: "/node/" + testNodeUID}}} + mutate(template) + + der, err := x509.CreateCertificate(rand.Reader, template, ca, f.key.Public(), key) + if err != nil { + t.Fatal(err) + } + + return tls.Certificate{Certificate: [][]byte{der, ca.Raw}, PrivateKey: f.key} +} + +type blockingResponse struct { + *httptest.ResponseRecorder + entered, unblock chan struct{} + once sync.Once + blockFlush bool + fail bool +} + +func (w *blockingResponse) Write(b []byte) (int, error) { + if !w.blockFlush { + w.once.Do(func() { close(w.entered); <-w.unblock }) + + if w.fail { + return 0, io.ErrClosedPipe + } + } + + return w.ResponseRecorder.Write(b) +} + +func (w *blockingResponse) FlushError() error { + if w.blockFlush { + w.once.Do(func() { close(w.entered); <-w.unblock }) + + if w.fail { + return io.ErrClosedPipe + } + } + + w.Flush() + + return nil +} + +func (w *blockingResponse) Unwrap() http.ResponseWriter { return w.ResponseRecorder } + +func awaitServerPolls(t *testing.T, s *Server, count int) { + t.Helper() + require.Eventually(t, func() bool { return s.polls.count() == count }, 5*time.Second, time.Millisecond, "poll admission did not reach %d", count) +} + +func bootstrapTestRequest(t *testing.T, ctx context.Context, endpoint, token string, request wire.BootstrapRequest) *http.Request { + t.Helper() + + body, err := wire.EncodeBootstrapRequest(request) + require.NoError(t, err) + r, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint+wire.BootstrapPath, bytes.NewReader(body)) + require.NoError(t, err) + r.Header.Set("Content-Type", "application/json") + r.Header.Set("Authorization", "Bearer "+token) + + return r +} + +func (f *servingFixture) publicRequest(t *testing.T, route string) *http.Request { + t.Helper() + + r := httptest.NewRequest(http.MethodGet, wire.SnapshotPath, nil) + if route == "keyring" { + r.URL.Path = wire.KeyringPath + } + + if route == "bootstrap" { + r = bootstrapTestRequest(t, f.ctx, "", f.token, f.request) + } + + if route == "keyring empty" { + r = httptest.NewRequest(http.MethodGet, wire.KeyringPath+"?after=1", nil) + + configureFixtureAge(t, f, 2*wire.PollWait) + } + + r.TLS = f.requestState(t) + + return r +} + +func requireNoAdmission(t *testing.T, s *Server) { + t.Helper() + require.Empty(t, s.writes, "write admission leaked") + require.Empty(t, s.bootstrapSlots, "bootstrap admission leaked") + require.Zero(t, s.keyringPolls.count(), "keyring admission leaked") + require.Zero(t, s.polls.count(), "snapshot admission leaked") +} + +func serveRecover(handler http.Handler, w http.ResponseWriter, r *http.Request) (aborted any) { + defer func() { aborted = recover() }() + + handler.ServeHTTP(w, r) + + return nil +} + +func (f *servingFixture) requestState(t *testing.T) *tls.ConnectionState { + t.Helper() + + certs := make([]*x509.Certificate, len(f.certificate.Certificate)) + for i, der := range f.certificate.Certificate { + var err error + + certs[i], err = x509.ParseCertificate(der) + if err != nil { + t.Fatal(err) + } + } + + return &tls.ConnectionState{HandshakeComplete: true, PeerCertificates: certs, VerifiedChains: [][]*x509.Certificate{certs}} +} + +func TestAdmissionHeldThroughWriteCompletion(t *testing.T) { + for _, tc := range []struct { + name string + flush, fail bool + status int + }{ + {"write", false, false, http.StatusOK}, + {"flush", true, false, http.StatusOK}, + {"failed write", false, true, http.StatusOK}, + {"failed flush", true, true, http.StatusOK}, + {"no content flush", true, false, http.StatusNoContent}, + {"error write", false, false, http.StatusServiceUnavailable}, + {"error flush", true, false, http.StatusServiceUnavailable}, + } { + t.Run(tc.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxPolls = 1 }) + configureFixtureAge(t, f, 2*wire.PollWait) + f.configureServer(func(c *Config) { c.Limits.MaxConcurrentWrites = 1 }) + other := *f + other.certificate = f.signLeaf(t, func(c *x509.Certificate) { c.URIs[0].Path = "/node/" + testOtherUID }) + otherState := other.requestState(t) + handler := f.a.Server.Handler() + r := httptest.NewRequest("GET", wire.SnapshotPath, nil) + + r.TLS = f.requestState(t) + r.URL.RawQuery = snapshotQueryForStatus(t, f, tc.status) + + w := &blockingResponse{ResponseRecorder: httptest.NewRecorder(), entered: make(chan struct{}), unblock: make(chan struct{}), blockFlush: tc.flush, fail: tc.fail} + + unblock := sync.OnceFunc(func() { close(w.unblock) }) + defer unblock() + + done := make(chan any, 1) + + go func() { defer func() { done <- recover() }(); handler.ServeHTTP(w, r) }() + + select { + case <-w.entered: + case <-time.After(wire.PollWait + time.Second): + t.Fatal("response did not start") + } + + // Use an immediate cursor so this also proves error responses retain admission. + r = r.Clone(f.ctx) + + r.URL.RawQuery = "" + for _, state := range []*tls.ConnectionState{r.TLS, otherState} { + second := httptest.NewRecorder() + request := r.Clone(f.ctx) + request.TLS = state + handler.ServeHTTP(second, request) + responseBody(t, second.Result(), nil, http.StatusTooManyRequests) + } + + wantWrites := 1 + if tc.status == http.StatusServiceUnavailable { + wantWrites = 0 + } + + require.Len(t, f.a.Server.writes, wantWrites, "write slots during response") + + unblock() + + aborted := <-done + if tc.fail { + require.Equal(t, http.ErrAbortHandler, aborted) + } else { + require.Nil(t, aborted) + require.Equal(t, tc.status, w.Code) + } + + require.Empty(t, f.a.Server.writes, "write slot leaked") + + third := httptest.NewRecorder() + handler.ServeHTTP(third, r.Clone(f.ctx)) + + require.Equal(t, http.StatusOK, third.Code, "admission leaked") + }) + }) + } +} + +func snapshotQueryForStatus(t *testing.T, f *servingFixture, status int) string { + t.Helper() + + if status == http.StatusOK { + return "" + } + + current, err := f.a.authority.Current() + require.NoError(t, err) + + cursor := current.Sequence() + if status == http.StatusServiceUnavailable { + cursor++ + } + + return fmt.Sprintf("after=%d", cursor) +} + +func TestHTTPPollAdmissionAndCancellation(t *testing.T) { + for _, limit := range []int{1, 2} { + t.Run(fmt.Sprint(limit), func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxPolls = limit }) + other := *f + other.certificate = f.signLeaf(t, func(c *x509.Certificate) { c.URIs[0].Path = "/node/" + testOtherUID }) + handler := f.a.Server.Handler() + + current, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + r := httptest.NewRequestWithContext(ctx, "GET", fmt.Sprintf("%s?after=%d", wire.SnapshotPath, current.Sequence()), nil) + r.TLS = f.requestState(t) + done := make(chan *httptest.ResponseRecorder, 1) + + go func() { w := httptest.NewRecorder(); handler.ServeHTTP(w, r); done <- w }() + + awaitServerPolls(t, f.a.Server, 1) + + if len(f.a.Server.writes) != 0 { + t.Fatal("long poll consumed write slot") + } + + for _, state := range []*tls.ConnectionState{r.TLS, other.requestState(t)} { + request := httptest.NewRequest("GET", wire.SnapshotPath, nil) + request.TLS = state + w := httptest.NewRecorder() + handler.ServeHTTP(w, request) + + want := http.StatusTooManyRequests + if state != r.TLS && limit == 2 { + want = http.StatusOK + } + + responseBody(t, w.Result(), nil, want) + } + + cancel() + + select { + case w := <-done: + responseBody(t, w.Result(), nil, http.StatusServiceUnavailable) + case <-time.After(5 * time.Second): + t.Fatal("canceled poll did not return") + } + + awaitServerPolls(t, f.a.Server, 0) + + for _, state := range []*tls.ConnectionState{r.TLS, other.requestState(t)} { + request := httptest.NewRequest("GET", wire.SnapshotPath, nil) + request.TLS = state + w := httptest.NewRecorder() + handler.ServeHTTP(w, request) + responseBody(t, w.Result(), nil, http.StatusOK) + } + }) + } +} + +func TestTLSAdmissionSaturationSendsInternalError(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.MaxConcurrentBootstrap = 1 }) + + endpoint := f.start(t) + if !take(f.a.Server.authSlots) { + t.Fatal("could not saturate admission") + } + defer release(f.a.Server.authSlots) + + conn, err := net.DialTimeout("tcp", strings.TrimPrefix(endpoint, "https://"), time.Second) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + + client := tls.Client(conn, &tls.Config{RootCAs: f.roots, ServerName: "example.com", MinVersion: tls.VersionTLS13}) + + err = client.HandshakeContext(f.ctx) + if err == nil || !strings.Contains(err.Error(), "internal error") { + t.Fatalf("saturated ClientHello should send TLS internal_error, got %v", err) + } +} + +func TestAdmissionAndAPIDeadlines(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { + c.Limits.MaxConcurrentBootstrap = 1 + c.Limits.WriteTimeout = 100 * time.Millisecond + }) + + entered := make(chan struct{}, 1) + fixtureDependencies[f.a.authority].reader = interceptor.NewClient(f.a.Topology.Client.(client.WithWatch), interceptor.Funcs{Get: func(ctx context.Context, _ client.WithWatch, _ client.ObjectKey, _ client.Object, _ ...client.GetOption) error { + entered <- struct{}{} + + <-ctx.Done() + + return ctx.Err() + }}) + handler := f.a.Server.Handler() + + body, err := wire.EncodeBootstrapRequest(f.request) + if err != nil { + t.Fatal(err) + } + + r := httptest.NewRequest("POST", wire.BootstrapPath, bytes.NewReader(body)) + r.Header.Set("Content-Type", "application/json") + r.Header.Set("Authorization", "Bearer "+f.token) + r.TLS = f.requestState(t) + done := make(chan *httptest.ResponseRecorder, 1) + + go func() { w := httptest.NewRecorder(); handler.ServeHTTP(w, r); done <- w }() + + <-entered + + w := httptest.NewRecorder() + handler.ServeHTTP(w, r.Clone(f.ctx)) + + if w.Code != 429 { + t.Fatalf("auth concurrency unbounded: %d", w.Code) + } + + select { + case w := <-done: + if w.Code != 503 { + t.Fatalf("deadline: %d", w.Code) + } + case <-time.After(time.Second): + t.Fatal("API deadline not enforced") + } +} + +func testConfig(t *testing.T) fixtureConfig { + t.Helper() + + return fixtureConfig{ + Config: authority.Config{ + Cluster: testOtherUID, Namespace: "racer", DataplaneServiceAccount: "racer-dataplane", + ControllerServiceAccount: "racer-controller", DaemonSetName: "racer-dataplane", + CredentialsSecretName: "racer-credentials", VersionConfigMapName: "racer-version", + InstallationConfigMapName: "racer-installation", CertificateLifetime: wire.CertificateLifetime, + SnapshotMaxAge: 30 * time.Second, MaxTokenBytes: 16 * 1024, + Rotation: authority.RotationPolicy{Interval: 24 * time.Hour, PrepareFor: time.Hour, RetainFor: 48 * time.Hour}, + }, + PeerPort: 8082, ServerConfig: Config{ + ControlAddress: ":8443", TLSCertificateFile: "/etc/racer/tls/tls.crt", TLSPrivateKeyFile: "/etc/racer/tls/tls.key", ReplicationServerName: "racer-controller.racer.svc", + Limits: Limits{MaxConnections: 2*wire.MaxMembers + 128, MaxConcurrentHandshakes: 32, MaxPolls: wire.MaxMembers, MaxConcurrentWrites: 128, MaxConcurrentBootstrap: 32, HeaderBytes: 16 * 1024, HandshakeTimeout: 5 * time.Second, WriteTimeout: 30 * time.Second, ShutdownTimeout: 10 * time.Second}, + }, + } +} + +func testTopology(t *testing.T, objects ...client.Object) *TopologyReconciler { + t.Helper() + cfg := testConfig(t) + + scheme := runtime.NewScheme() + for _, add := range []func(*runtime.Scheme) error{corev1.AddToScheme, appsv1.AddToScheme, racerv1.AddToScheme} { + if err := add(scheme); err != nil { + t.Fatal(err) + } + } + + objects = append(objects, &corev1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Namespace: cfg.Namespace, Name: cfg.InstallationConfigMapName, UID: "installation-uid"}, Data: map[string]string{"cluster": string(cfg.Cluster), "version_configmap": cfg.VersionConfigMapName, "state": "fresh", "initialization_protocol": "staged-v1"}}) + c := fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).WithIndex(&corev1.Pod{}, podNodeIndex, podNodeKeys).Build() + c = interceptor.NewClient(c, interceptor.Funcs{Create: func(ctx context.Context, c client.WithWatch, obj client.Object, opts ...client.CreateOption) error { + if obj.GetUID() == "" { + obj.SetUID(types.UID(fmt.Sprintf("fake-%s-%d", obj.GetName(), time.Now().UnixNano()))) + } + + return c.Create(ctx, obj, opts...) + }}) + + return Assemble(cfg, c, c).Topology +} + +func initializedTopology(t *testing.T, objects ...client.Object) *TopologyReconciler { + t.Helper() + + r := testTopology(t, objects...) + if err := r.authority.Recover(t.Context(), r.Client); err != nil { + t.Fatal(err) + } + + return r +} + +// Captured publication data is decoded through the public response operation, +// never an installation proof. It belongs exclusively to the integration test. +type ( + CommittedPublication struct { + handle *authority.PublicationHandle + encoded string + record VersionRecord + leadership context.Context + } + VersionRecord struct { + Cluster wire.ClusterID + Sequence wire.Sequence + MembershipVersion wire.MembershipVersion + ContentHash string + MembershipHash string + } +) + +func capturePublication(t *testing.T, a *authority.Authority) *CommittedPublication { + t.Helper() + + h, err := a.Current() + if err != nil { + t.Fatal(err) + } + + return captureHandle(t, h) +} + +func captureHandle(t *testing.T, h *authority.PublicationHandle) *CommittedPublication { + t.Helper() + + var b bytes.Buffer + + guard, cancel, err := h.Admit(t.Context()) + if err != nil { + t.Fatal(err) + } + defer cancel() + + ctx := guard.Context() + if _, err := h.ForBase(0, "").WriteTo(ctx, guard, &b); err != nil { + t.Fatal(err) + } + + p, err := wire.DecodePublication(bytes.NewReader(b.Bytes())) + if err != nil { + t.Fatal(err) + } + + content, members, err := wire.ContentHashes(p) + if err != nil { + t.Fatal(err) + } + + return &CommittedPublication{handle: h, encoded: b.String(), record: VersionRecord{Cluster: p.Cluster, Sequence: p.Sequence, MembershipVersion: p.MembershipVersion, ContentHash: content, MembershipHash: members}, leadership: ctx} +} + +func (p *CommittedPublication) admit(ctx context.Context) (*authority.Admission, context.CancelFunc, error) { + return p.handle.Admit(ctx) +} + +func reconcileTopology(t *testing.T, r *TopologyReconciler, ctx context.Context) *CommittedPublication { + t.Helper() + + result, err := r.Reconcile(ctx, ctrl.Request{}) + if err != nil || result.RequeueAfter != 0 { + t.Fatalf("reconcile: %v, %v", result, err) + } + + return capturePublication(t, r.authority) +} + +func acceptedMembers(t *testing.T, r *TopologyReconciler) AcceptedMembers { + t.Helper() + + p, err := r.authority.Current() + if err != nil { + return nil + } + + captured := captureHandle(t, p) + + image, err := wire.DecodePublication(bytes.NewBufferString(captured.encoded)) + if err != nil { + t.Fatal(err) + } + + members := make(AcceptedMembers, len(image.Members)) + for _, m := range image.Members { + members[m.Node] = m + } + + return members +} + +func runKeys(t *testing.T, r *KeyringReconciler) ctrl.Result { + t.Helper() + + result, err := r.Reconcile(t.Context(), ctrl.Request{}) + if err != nil || result.RequeueAfter <= 0 { + t.Fatalf("reconcile: %v, %v", result, err) + } + + return result +} + +type ( + RotationState struct { + NextRotation time.Time `json:"next_rotation"` + ActivateAt time.Time `json:"activate_at"` + ActiveIssuer string `json:"active_issuer"` + PreparedIssuer string `json:"prepared_issuer"` + Retiring map[string]time.Time `json:"retiring"` + } + signingMaterial struct { + PrivateKey []byte `json:"private_key"` + Certificate []byte `json:"certificate"` + } + issuerMaterial struct { + Keys map[string]signingMaterial `json:"keys"` + } +) + +func keyState(t *testing.T, r *KeyringReconciler) (*corev1.Secret, wire.KeyringBundle, RotationState, issuerMaterial) { + t.Helper() + + var secret corev1.Secret + + deps := fixtureDependencies[r.authority] + if err := deps.reader.Get(t.Context(), client.ObjectKey{Namespace: r.Config.Namespace, Name: r.Config.CredentialsSecretName}, &secret); err != nil { + t.Fatal(err) + } + + b, err := wire.DecodeBundle(bytes.NewReader(secret.Data["bundle.json"])) + if err != nil { + t.Fatal(err) + } + + var state RotationState + if err := json.Unmarshal(secret.Data["rotation.json"], &state); err != nil { + t.Fatal(err) + } + + var material issuerMaterial + if err := json.Unmarshal(secret.Data["issuer.json"], &material); err != nil { + t.Fatal(err) + } + + return &secret, b, state, material +} + +// Fault injection is a Kubernetes dependency supplied at construction, not an +// authority mutation hook. Tests may change the transport under that dependency. +type fixtureDependency struct { + client.Client + reader client.Reader + now func() time.Time +} + +func (d *fixtureDependency) Get(ctx context.Context, key client.ObjectKey, obj client.Object, opts ...client.GetOption) error { + return d.reader.Get(ctx, key, obj, opts...) +} + +func (d *fixtureDependency) List(ctx context.Context, list client.ObjectList, opts ...client.ListOption) error { + return d.reader.List(ctx, list, opts...) +} + +var fixtureDependencies = map[*authority.Authority]*fixtureDependency{} + +func assembleFixture(cfg fixtureConfig, c client.Client, reader client.Reader) *Application { + d := &fixtureDependency{Client: c, reader: reader, now: time.Now} + a := Assemble(cfg, d, d) + // Supply a clock through construction; no setter is exposed by authority. + owner := authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: d, Reader: d, Now: func() time.Time { return d.now() }}) + a.authority = owner + a.Topology.authority = owner + a.Keyring.authority = owner + a.Server.authority = owner + a.Replication.authority = owner + a.Lifecycle.authority = owner + a.Topology.Client = c + a.Topology.APIReader = reader + a.Replication.Client = c + a.Replication.APIReader = reader + fixtureDependencies[owner] = d + fixtureConfigs[owner] = cfg + + return a +} + +func decodeIssuedResponse(t *testing.T, encoded []byte) wire.BootstrapResponse { + t.Helper() + + response, err := wire.DecodeBootstrapResponse(bytes.NewReader(encoded)) + if err != nil { + t.Fatal(err) + } + + return response +} + +func issuanceRequest(t *testing.T, r *KeyringReconciler) (NodeIdentity, wire.BootstrapRequest, ed25519.PublicKey) { + t.Helper() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + csr, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{}, key) + if err != nil { + t.Fatal(err) + } + + return NodeIdentity{}, wire.BootstrapRequest{SchemaVersion: 1, Cluster: r.Config.Cluster, Enrollment: testOtherUID, CSRDER: csr, Shares: wire.DefaultShares}, pub +} + +func fixtureIdentity(t *testing.T, f *servingFixture) NodeIdentity { + t.Helper() + + r := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + r.Header.Set("Authorization", "Bearer "+f.token) + + identity, err := f.a.authority.Authenticate(t.Context(), r) + if err != nil { + t.Fatal(err) + } + + return identity +} + +func invalidateFixtureTrust(t *testing.T, f *servingFixture) { + t.Helper() + secret, _, _, _ := keyState(t, f.a.Keyring) + + secret.Data["bundle.json"] = []byte(`{}`) + if err := f.a.Topology.Update(t.Context(), secret); err != nil { + t.Fatal(err) + } + + if err := f.a.authority.Observe(t.Context()); err == nil { + t.Fatal("invalid credential observation succeeded") + } +} + +func readVersion(ctx context.Context, reader client.Reader, cfg fixtureConfig) (*corev1.ConfigMap, VersionRecord, error) { + if err := authority.New(cfg.authorityConfig(), authority.Dependencies{Reader: reader}).Recover(ctx, nil); err != nil { + return nil, VersionRecord{}, err + } + + var cm corev1.ConfigMap + if err := reader.Get(ctx, client.ObjectKey{Namespace: cfg.Namespace, Name: cfg.VersionConfigMapName}, &cm); err != nil { + return nil, VersionRecord{}, err + } + + seq, err := strconv.ParseUint(cm.Data["sequence"], 10, 64) + if err != nil { + return nil, VersionRecord{}, err + } + + members, err := strconv.ParseUint(cm.Data["membership_version"], 10, 64) + if err != nil { + return nil, VersionRecord{}, err + } + + return &cm, VersionRecord{Cluster: wire.ClusterID(cm.Data["cluster"]), Sequence: wire.Sequence(seq), MembershipVersion: wire.MembershipVersion(members), ContentHash: cm.Data["content_hash"], MembershipHash: cm.Data["membership_hash"]}, nil +} + +func configureFixtureAge(t *testing.T, f *servingFixture, age time.Duration) { + t.Helper() + + cfg := f.a.Topology.Config + cfg.SnapshotMaxAge = age + d := fixtureDependencies[f.a.authority] + a := authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: d, Reader: d, Now: func() time.Time { return d.now() }}) + f.a.authority = a + f.a.Topology.authority = a + f.a.Keyring.authority = a + f.a.Server.authority = a + f.a.Replication.authority = a + f.a.Lifecycle.authority = a + + fixtureDependencies[a] = d + + fixtureConfigs[a] = cfg + if err := a.Observe(t.Context()); err != nil { + t.Fatal(err) + } + + reconcileTopology(t, f.a.Topology, f.ctx) +} + +func replaceFixtureCredentials(t *testing.T, f *servingFixture) { + t.Helper() + other := newServingFixture(t) + candidate, _, _, _ := keyState(t, other.a.Keyring) + current, bundle, _, _ := keyState(t, f.a.Keyring) + + var ( + replacement wire.KeyringBundle + err error + ) + + replacement, err = wire.DecodeBundle(bytes.NewReader(candidate.Data["bundle.json"])) + if err != nil { + t.Fatal(err) + } + + replacement.Generation = bundle.Generation + 1 + + candidate.Data["bundle.json"], err = wire.EncodeBundle(replacement) + if err != nil { + t.Fatal(err) + } + + current.Data = candidate.Data + if err := f.a.Topology.Update(t.Context(), current); err != nil { + t.Fatal(err) + } + + if err := f.a.authority.Observe(t.Context()); err != nil { + t.Fatal(err) + } +} + +var withdrawnSecrets = map[*authority.Authority]*corev1.Secret{} + +func withdrawServerTrust(t *testing.T, s *Server) { + t.Helper() + + d := fixtureDependencies[s.authority] + + var secret corev1.Secret + if err := d.Client.Get(t.Context(), client.ObjectKey{Namespace: fixtureConfigs[s.authority].Namespace, Name: fixtureConfigs[s.authority].CredentialsSecretName}, &secret); err != nil { + t.Fatal(err) + } + + if withdrawnSecrets[s.authority] == nil { + withdrawnSecrets[s.authority] = secret.DeepCopy() + } + + secret.Data["bundle.json"] = []byte(`{}`) + if err := d.Update(t.Context(), &secret); err != nil { + t.Fatal(err) + } + + if _, err := s.authority.ReconcileCredentials(t.Context()); err == nil { + t.Fatal("invalid trust accepted") + } +} + +func restoreServerTrust(t *testing.T, s *Server) { + t.Helper() + + d := fixtureDependencies[s.authority] + + saved := withdrawnSecrets[s.authority] + if saved == nil { + return + } + + var secret corev1.Secret + if err := d.Client.Get(t.Context(), client.ObjectKeyFromObject(saved), &secret); err != nil { + t.Fatal(err) + } + + secret.Data = saved.DeepCopy().Data + if err := d.Update(t.Context(), &secret); err != nil { + t.Fatal(err) + } + + if _, err := s.authority.ReconcileCredentials(t.Context()); err != nil { + t.Fatal(err) + } +} + +func withdrawPublication(t *testing.T, r *TopologyReconciler) func() { + t.Helper() + + cm, _, err := readVersion(t.Context(), r.APIReader, r.Config) + if err != nil { + t.Fatal(err) + } + + saved := cm.DeepCopy() + + cm.Data["sequence"] = "0" + if err := r.Update(t.Context(), cm); err != nil { + t.Fatal(err) + } + + if _, err := r.authority.PublishTopology(t.Context(), r.observeTopology); err == nil { + t.Fatal("invalid publication accepted") + } + + return func() { + if err := r.Get(t.Context(), client.ObjectKeyFromObject(cm), cm); err != nil { + t.Fatal(err) + } + + cm.Data = saved.Data + if err := r.Update(t.Context(), cm); err != nil { + t.Fatal(err) + } + } +} + +func parseSigning(m signingMaterial) (*x509.Certificate, ed25519.PrivateKey, error) { + cert, err := x509.ParseCertificate(m.Certificate) + if err != nil { + return nil, nil, err + } + + key, err := x509.ParsePKCS8PrivateKey(m.PrivateKey) + if err != nil { + return nil, nil, err + } + + return cert, key.(ed25519.PrivateKey), nil +} + +func waitFixturePublication(ctx context.Context, a *authority.Authority, after wire.Sequence) (*authority.PublicationHandle, error) { + for { + p, changed, err := a.CurrentAndSubscribe() + if err != nil { + return nil, err + } + + if p.Sequence() > after { + return p, nil + } + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-changed: + } + } +} + +func rootID(der []byte) string { sum := sha256.Sum256(der); return hex.EncodeToString(sum[:]) } + +func generateIssuer(now time.Time, cfg fixtureConfig) ([]byte, []byte, error) { + pub, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, nil, err + } + + cert := &x509.Certificate{SerialNumber: big.NewInt(1), NotBefore: now.Add(-time.Minute), NotAfter: now.Add(cfg.Rotation.Interval + cfg.Rotation.PrepareFor + cfg.Rotation.RetainFor + 2*cfg.CertificateLifetime), IsCA: true, BasicConstraintsValid: true, MaxPathLenZero: true, KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign} + + der, err := x509.CreateCertificate(rand.Reader, cert, cert, pub, key) + if err != nil { + return nil, nil, err + } + + encoded, err := x509.MarshalPKCS8PrivateKey(key) + + return der, encoded, err +} + +func largeFixturePublication(t *testing.T, f *servingFixture) { + t.Helper() + + members := make(AcceptedMembers, 50000) + + for i := range 50000 { + id := wire.NodeID(fmt.Sprintf("33333333-3333-4333-8333-%012d", i)) + members[id] = wire.Member{Node: id, Shares: 4, PeerEndpoint: "192.0.2.1:8082", RDMANICs: []wire.RDMANIC{}} + } + + replicationSmokePublish(t, f.ctx, f.a.Topology, members) +} + +func rotateFixtureTrust(t *testing.T, f *servingFixture) { + t.Helper() + secret, bundle, _, _ := keyState(t, f.a.Keyring) + bundle.Generation++ + + encoded, err := wire.EncodeBundle(bundle) + if err != nil { + t.Fatal(err) + } + + secret.Data["bundle.json"] = encoded + if err := f.a.Topology.Update(t.Context(), secret); err != nil { + t.Fatal(err) + } + + if _, err := f.a.authority.ReconcileCredentials(t.Context()); err != nil { + t.Fatal(err) + } +} + +type teardownListener struct { + net.Listener + accept func() (net.Conn, error) + close func() error +} + +func (l teardownListener) Accept() (net.Conn, error) { return l.accept() } + +func (l teardownListener) Close() error { return l.close() } + +func TestServeTeardownErrorsAndLateAccept(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + s := New(Config{Limits: Limits{ShutdownTimeout: time.Second}}, nil, nil, NewLifecycle(nil), nil) + + accepted, peer := net.Pipe() + defer peer.Close() + defer accepted.Close() + + closing := make(chan struct{}) + acceptErr := errors.New("accept failed during cancellation") + closeErr := errors.New("listener close failed") + calls := 0 + listener := teardownListener{ + accept: func() (net.Conn, error) { + if accepted != nil { + // Return a connection only after the force-close sweep, while + // net/http is closing the listener and waiting for Serve. + <-closing + + conn := accepted + accepted = nil + + return conn, nil + } + + return nil, acceptErr + }, + close: func() error { + calls++ + + close(closing) + + return closeErr + }, + } + done := make(chan error, 1) + + go func() { done <- s.serve(ctx, listener, &tls.Config{}) }() + + synctest.Wait() + cancel() + synctest.Wait() + + // net/http reports ErrServerClosed once Close starts, even if Accept + // itself failed. The listener's close failure must still be returned. + if err := <-done; !errors.Is(err, closeErr) { + t.Fatalf("teardown error: %v", err) + } + + if calls != 1 || s.Lifecycle.serving { + t.Fatalf("close calls = %d, serving = %v", calls, s.Lifecycle.serving) + } + + if _, err := peer.Read(make([]byte, 1)); !errors.Is(err, io.EOF) { + t.Fatalf("late connection was not closed: %v", err) + } + }) +} + +func TestServeTeardownCompletionBound(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + s := New(Config{Limits: Limits{ShutdownTimeout: time.Second}}, nil, nil, NewLifecycle(nil), nil) + unblock := make(chan struct{}) + + release := sync.OnceFunc(func() { close(unblock) }) + defer release() + + listener := teardownListener{ + accept: func() (net.Conn, error) { <-unblock; return nil, net.ErrClosed }, + close: func() error { <-unblock; return nil }, + } + done := make(chan error, 1) + + go func() { done <- s.serve(ctx, listener, &tls.Config{}) }() + + synctest.Wait() + cancel() + + start := time.Now() + + if err := <-done; !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("blocked teardown: %v", err) + } + + if elapsed := time.Since(start); elapsed != s.config.Limits.ShutdownTimeout { + t.Fatalf("completion wait = %s", elapsed) + } + + if s.Lifecycle.serving { + t.Fatal("timed-out server still serving ready") + } + + release() + synctest.Wait() + }) +} + +func TestServeTeardownClosesSlowTLSWrite(t *testing.T) { + for _, cause := range []string{"cancellation", "accept failure"} { + t.Run(cause, func(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { + c.Limits.WriteTimeout = time.Minute + c.Limits.ShutdownTimeout = time.Second + }) + largeFixturePublication(t, f) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + t.Cleanup(func() { _ = listener.Close() }) + + done := make(chan error, 1) + + config := f.a.Server.tlsConfigWithCertificate(f.ctx, func(*tls.ClientHelloInfo) (*tls.Certificate, error) { + return &f.serverCertificate, nil + }) + + go func() { done <- f.a.Server.serve(f.ctx, listener, config) }() + + conn, err := tls.Dial("tcp", listener.Addr().String(), &tls.Config{RootCAs: f.roots, MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{f.certificate}}) + require.NoError(t, err) + + defer conn.Close() + + _, err = io.WriteString(conn, "GET /v1/snapshot HTTP/1.1\r\nHost: localhost\r\n\r\n") + require.NoError(t, err) + require.Eventually(t, func() bool { return len(f.a.Server.writes) != 0 }, 5*time.Second, time.Millisecond, "write not admitted") + // Keep the peer open without reading. Teardown must release the + // socket and admission before the much longer write deadline. + if cause == "cancellation" { + f.cancel() + } else if err := listener.Close(); err != nil { + t.Fatal(err) + } + + select { + case err := <-done: + if cause == "cancellation" && err != nil || cause == "accept failure" && !errors.Is(err, net.ErrClosed) { + t.Fatalf("serve result: %v", err) + } + case <-time.After(2 * time.Second): + t.Fatal("TLS teardown blocked") + } + + awaitServerPolls(t, f.a.Server, 0) + + require.Empty(t, f.a.Server.writes, "teardown retained admission") + require.Error(t, f.a.Server.Ready(nil), "teardown retained readiness") + }) + } +} + +func TestServingChainFreezesBeforeFirstRequest(t *testing.T) { + for _, boundary := range []string{"handler", "tls"} { + t.Run(boundary, func(t *testing.T) { + f := newServingFixture(t) + s := f.a.Server + // Pre-use overrides now belong to construction, not mutable engines. + cfg := f.a.Topology.Config + cfg.CertificateLifetime = 2 * time.Minute + cfg.SnapshotMaxAge = 30 * time.Second + d := fixtureDependencies[f.a.authority] + + s.authority = authority.New(cfg.authorityConfig(), authority.Dependencies{Writer: d, Reader: d}) + require.NoError(t, s.authority.Observe(f.ctx)) + + f.a.Replication.Config = cfg + + want := cfg + input := s.config + s = New(input, s.writer, s.authority, s.Lifecycle, s.Leader) + s.installTestServingCertificate(f.serverCertificate) + + wantServer := input + + var handler http.Handler + if boundary == "handler" { + handler = s.Handler() + } else { + s.tlsConfigWithCertificate(f.ctx, func(*tls.ClientHelloInfo) (*tls.Certificate, error) { return &f.serverCertificate, nil }) + } + // Mutate sequentially, before any request, without calling dependency + // getters first: those calls would accidentally hide lazy freezing. + cfg.DataplaneServiceAccount = "wrong-account" + cfg.Cluster = "" + cfg.CertificateLifetime = time.Second + input.Limits.HeaderBytes = 1 + + cfg.ControllerServiceAccount = "wrong-controller" + + if handler == nil { + handler = s.Handler() + } + + require.Equal(t, wantServer, s.config, "server did not freeze transport before exposure") + request := bootstrapTestRequest(t, f.ctx, "", f.token, f.request) + request.TLS = &tls.ConnectionState{HandshakeComplete: true} + w := httptest.NewRecorder() + handler.ServeHTTP(w, request) + + require.Equal(t, http.StatusOK, w.Code, "first request used post-exposure config") + + response := decodeIssuedResponse(t, w.Body.Bytes()) + + leaf, err := x509.ParseCertificate(response.CertificateChain[0]) + if err != nil { + t.Fatal(err) + } + + require.Equal(t, want.Cluster, response.Cluster) + require.Equal(t, want.CertificateLifetime+time.Minute, leaf.NotAfter.Sub(leaf.NotBefore)) + }) + } +} + +func TestLifecycleProcessContext(t *testing.T) { + for _, source := range []string{"parent", "process", "child"} { + t.Run(source, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + process, stopProcess := context.WithCancel(t.Context()) + defer stopProcess() + + parent, stopParent := context.WithTimeout(t.Context(), time.Minute) + defer stopParent() + + key := connectionKey{} + parent = context.WithValue(parent, key, source) + l := NewLifecycle(nil) + l.process = process + + child, cancel := l.ProcessContext(parent) + defer cancel() + + deadline, ok := child.Deadline() + + parentDeadline, _ := parent.Deadline() + + require.NoError(t, child.Err()) + require.Equal(t, source, child.Value(key)) + require.True(t, ok) + require.Equal(t, parentDeadline, deadline) + + switch source { + case "parent": + stopParent() + case "process": + stopProcess() + case "child": + cancel() + } + + synctest.Wait() + + if !errors.Is(child.Err(), context.Canceled) { + t.Fatalf("child ignored %s cancellation: %v", source, child.Err()) + } + + if source != "process" && process.Err() != nil || source != "parent" && parent.Err() != nil { + t.Fatal("child cancellation propagated to an independent parent") + } + }) + }) + } + + process, cancel := context.WithCancel(t.Context()) + cancel() + + for _, l := range []*Lifecycle{nil, NewLifecycle(nil), {process: process}} { + child, stop := l.ProcessContext(t.Context()) + if !errors.Is(child.Err(), context.Canceled) { + t.Fatal("absent or canceled process did not immediately cancel child") + } + + stop() + } +} + +func TestLifecycleHTTPReadinessTransitions(t *testing.T) { + for _, tc := range []struct { + name string + set func(*Server, bool) + }{ + {name: "issuer", set: func(s *Server, ready bool) { + if ready { + restoreServerTrust(t, s) + } else { + withdrawServerTrust(t, s) + } + }}, + {name: "serving", set: func(s *Server, ready bool) { s.Lifecycle.SetServingReady(ready) }}, + } { + t.Run(tc.name, func(t *testing.T) { + f := newServingFixture(t) + handler := f.a.Server.Handler() + + tc.set(f.a.Server, false) + + for _, step := range []struct { + name string + ready bool + }{ + {name: "initial false"}, + {name: "become ready", ready: true}, + {name: "remain ready", ready: true}, + {name: "withdraw readiness"}, + {name: "remain unready"}, + {name: "restore readiness", ready: true}, + } { + t.Run(step.name, func(t *testing.T) { + tc.set(f.a.Server, step.ready) + + r := httptest.NewRequest(http.MethodGet, wire.SnapshotPath, nil) + r.TLS = f.requestState(t) + w := httptest.NewRecorder() + handler.ServeHTTP(w, r) + + want := http.StatusServiceUnavailable + if step.ready { + want = http.StatusOK + } + + responseBody(t, w.Result(), nil, want) + }) + } + }) + } +} + +func TestLifecycleFollowerWithValidatedPublicationIsReady(t *testing.T) { + r := initializedTopology(t) + reconcileTopology(t, r, t.Context()) + + if err := r.authority.PublicationReady(); err != nil { + t.Fatal(err) + } + + l := NewLifecycle(r.authority) + l.process, l.synced = t.Context(), true + l.SetServingReady(true) + + if err := l.Ready(nil); err != nil { + t.Fatalf("follower with validated state must receive Service traffic: %v", err) + } +} + +func TestLifecycleGatesAndCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + + l := NewLifecycle(r.authority) + if l.NeedLeaderElection() || l.Ready(nil) == nil { + t.Fatal("follower ready") + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + syncCache := make(chan struct{}) + l.waitForCacheSync = func(ctx context.Context) bool { + select { + case <-ctx.Done(): + return false + case <-syncCache: + return true + } + } + started := make(chan error, 1) + + go func() { started <- l.Start(ctx) }() + + l.SetServingReady(true) + reconcileTopology(t, r, ctx) + + if l.Ready(nil) == nil { + t.Fatal("ready before synchronized inputs") + } + + close(syncCache) + synctest.Wait() + + if err := l.Ready(nil); err != nil { + t.Fatal(err) + } + + restore := withdrawPublication(t, r) + + if l.Ready(nil) == nil { + t.Fatal("ready without publication authority") + } + + restore() + reconcileTopology(t, r, ctx) + l.SetServingReady(false) + + if l.Ready(nil) == nil { + t.Fatal("ready without listener") + } + + cancel() + + if err := <-started; err != nil { + t.Fatal(err) + } + + l.SetServingReady(true) + + if l.Ready(nil) == nil { + t.Fatal("old process resurrected") + } + + request, stop := l.ProcessContext(t.Context()) + defer stop() + + if !errors.Is(request.Err(), context.Canceled) { + t.Fatalf("request resurrected: %v", request.Err()) + } + + if err := l.Start(context.Background()); !errors.Is(err, wire.Conflict) { + t.Fatalf("process restarted: %v", err) + } + }) +} + +func TestLifecycleRequiresPublicationAndHonorsRequestCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + r := initializedTopology(t) + l := NewLifecycle(r.authority) + l.waitForCacheSync = func(context.Context) bool { return true } + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + done := make(chan error, 1) + + go func() { done <- l.Start(ctx) }() + + l.SetServingReady(true) + synctest.Wait() + + requestCtx, stop := context.WithCancel(context.Background()) + stop() + + request, stopRequest := l.ProcessContext(requestCtx) + defer stopRequest() + + if !errors.Is(request.Err(), context.Canceled) { + t.Fatalf("request ignores cancellation: %v", request.Err()) + } + + if err := l.Ready(nil); !errors.Is(err, wire.Unavailable) { + t.Fatalf("ready without publication: %v", err) + } + + reconcileTopology(t, r, ctx) + + if err := l.Ready(nil); err != nil { + t.Fatal(err) + } + + cancel() + + if err := <-done; err != nil { + t.Fatal(err) + } + + require.ErrorIs(t, l.Ready(nil), wire.Unavailable) + + beforeStartup := NewLifecycle(r.authority) + beforeStartup.waitForCacheSync = func(ctx context.Context) bool { return ctx.Err() == nil } + require.NoError(t, beforeStartup.Start(ctx)) + beforeStartup.SetServingReady(true) + require.ErrorIs(t, beforeStartup.Ready(nil), wire.Unavailable, "canceled startup became ready") + }) +} + +func TestLifecycleHTTPStartupAdmission(t *testing.T) { + for _, scenario := range []string{"synchronized", "cache failed", "canceled during sync"} { + t.Run(scenario, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + l := NewLifecycle(f.a.authority) + f.a.Server.Lifecycle = l + l.SetServingReady(true) + + syncCache := make(chan struct{}) + l.waitForCacheSync = func(ctx context.Context) bool { + select { + case <-ctx.Done(): + return false + case <-syncCache: + return scenario == "synchronized" + } + } + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + done := make(chan error, 1) + + go func() { done <- l.Start(ctx) }() + + synctest.Wait() + + handler := f.a.Server.Handler() + request := httptest.NewRequest(http.MethodGet, wire.SnapshotPath, nil) + request.TLS = f.requestState(t) + before := httptest.NewRecorder() + handler.ServeHTTP(before, request) + responseBody(t, before.Result(), nil, http.StatusServiceUnavailable) + + if scenario == "canceled during sync" { + cancel() + } else { + close(syncCache) + } + + synctest.Wait() + + after := httptest.NewRecorder() + handler.ServeHTTP(after, request) + + want := http.StatusServiceUnavailable + if scenario == "synchronized" { + want = http.StatusOK + } + + responseBody(t, after.Result(), nil, want) + + cancel() + + err := <-done + if scenario == "cache failed" { + if !errors.Is(err, wire.Unavailable) { + t.Fatalf("cache failure: %v", err) + } + } else if err != nil { + t.Fatal(err) + } + + stopped := httptest.NewRecorder() + handler.ServeHTTP(stopped, request) + responseBody(t, stopped.Result(), nil, http.StatusServiceUnavailable) + }) + }) + } +} + +type servingFixture struct { + a *Application + token string + key ed25519.PrivateKey + request wire.BootstrapRequest + certificate tls.Certificate + serverCertificate tls.Certificate + roots *x509.CertPool + ctx context.Context + cancel context.CancelFunc +} + +// Fixed-certificate adapter for socket tests; production always uses the reloader. +func (s *Server) tlsConfig(ctx context.Context, certificate tls.Certificate) *tls.Config { + r := s.installTestServingCertificate(certificate) + return s.tlsConfigWithCertificate(ctx, r.getCertificate) +} + +func (s *Server) installTestServingCertificate(certificate tls.Certificate) *servingCertificateReloader { + validated, err := validateServingCertificate(certificate, time.Now()) + if err != nil { + panic(err) + } + + r := &servingCertificateReloader{} + r.current.Store(validated) + s.servingCertificate.Store(r) + + return r +} + +func newServingFixture(t *testing.T) *servingFixture { + t.Helper() + a, status, token := authFixture(t) + installReview(t, a, status, token) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + runKeys(t, a.Keyring) + reconcileTopology(t, a.Topology, ctx) + a.Lifecycle.mu.Lock() + a.Lifecycle.process, a.Lifecycle.synced, a.Lifecycle.serving = ctx, true, true + a.Lifecycle.mu.Unlock() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + csr, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{DNSNames: []string{"attacker"}, Subject: pkix.Name{CommonName: "attacker"}}, key) + if err != nil { + t.Fatal(err) + } + + request := wire.BootstrapRequest{SchemaVersion: 1, Cluster: a.Topology.Config.Cluster, Enrollment: wire.EnrollmentID(testOtherUID), CSRDER: csr, Shares: wire.DefaultShares} + authRequest := httptest.NewRequest(http.MethodGet, wire.KeyringPath, nil) + authRequest.Header.Set("Authorization", "Bearer "+token) + + identity, err := a.authority.Authenticate(ctx, authRequest) + if err != nil { + t.Fatal(err) + } + + encoded, err := a.authority.Issue(ctx, identity, request) + if err != nil { + t.Fatal(err) + } + + response := decodeIssuedResponse(t, encoded) + + template := &x509.Certificate{SerialNumber: big.NewInt(1), NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), DNSNames: []string{a.Server.config.ReplicationServerName}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}} + + der, err := x509.CreateCertificate(rand.Reader, template, template, pub, key) + if err != nil { + t.Fatal(err) + } + + cert, err := x509.ParseCertificate(der) + if err != nil { + t.Fatal(err) + } + + roots := x509.NewCertPool() + roots.AddCert(cert) + a.Server.installTestServingCertificate(tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}) + + return &servingFixture{a: a, token: token, key: key, request: request, certificate: tls.Certificate{Certificate: response.CertificateChain, PrivateKey: key}, serverCertificate: tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}, roots: roots, ctx: ctx, cancel: cancel} +} + +// configureServer replaces the unstarted fixture server with constructor inputs. +func (f *servingFixture) configureServer(change func(*Config)) { + s := f.a.Server + cfg := s.config + change(&cfg) + f.a.Server = New(cfg, s.writer, s.authority, s.Lifecycle, s.Leader) + f.a.Server.servingCertificate.Store(s.servingCertificate.Load()) +} + +func (f *servingFixture) client(t *testing.T, cert *tls.Certificate) *http.Client { + t.Helper() + + config := &tls.Config{RootCAs: f.roots, MinVersion: tls.VersionTLS13, ClientSessionCache: tls.NewLRUClientSessionCache(4)} + if cert != nil { + config.GetClientCertificate = func(*tls.CertificateRequestInfo) (*tls.Certificate, error) { return cert, nil } + } + + transport := &http.Transport{TLSClientConfig: config} + t.Cleanup(transport.CloseIdleConnections) + + return &http.Client{Transport: transport, Timeout: 5 * time.Second} +} + +func (f *servingFixture) start(t *testing.T) string { + t.Helper() + + s := httptest.NewUnstartedServer(f.a.Server.Handler()) + s.Config.ConnContext = connectionContext + s.TLS = f.a.Server.tlsConfig(f.ctx, f.serverCertificate) + s.StartTLS() + t.Cleanup(s.Close) + + return s.URL +} + +func responseBody(t *testing.T, response *http.Response, err error, status int) []byte { + t.Helper() + + if err != nil { + t.Fatal(err) + } + + defer response.Body.Close() + + b, err := io.ReadAll(response.Body) + if err != nil { + t.Fatal(err) + } + + if response.StatusCode != status { + t.Fatalf("status %d, want %d: %s", response.StatusCode, status, b) + } + + if status >= 400 { + if _, err := wire.DecodeError(bytes.NewReader(b)); err != nil { + t.Fatalf("non-protocol error %q", b) + } + + if status == 429 || status == 503 { + if response.Header.Get("Retry-After") != "1" { + t.Fatal("missing retry bound") + } + } + } + + return b +} + +func TestOperationalStartCancellationAndTLSFiles(t *testing.T) { + f := newServingFixture(t) + dir := t.TempDir() + + f.configureServer(func(c *Config) { + c.ControlAddress = "127.0.0.1:0" + c.TLSCertificateFile = filepath.Join(dir, "tls.crt") + c.TLSPrivateKeyFile = filepath.Join(dir, "tls.key") + }) + writeServingTestPair(t, dir, f.serverCertificate) + + f.a.Lifecycle.SetServingReady(false) + + done := make(chan error, 1) + + go func() { done <- f.a.Server.Start(f.ctx) }() + + deadline := time.After(5 * time.Second) + + for f.a.Server.Ready(nil) != nil { + select { + case err := <-done: + t.Fatalf("start: %v", err) + case <-deadline: + t.Fatal("not ready") + default: + time.Sleep(time.Millisecond) + } + } + + f.cancel() + + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("shutdown blocked") + } + + if f.a.Server.Ready(nil) == nil { + t.Fatal("canceled server ready") + } +} + +func TestHTTPWriteBootstrapAndGlobalAdmission(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { + c.Limits.MaxPolls = 1 + c.Limits.MaxConcurrentWrites = 1 + c.Limits.MaxConcurrentBootstrap = 1 + }) + handler := f.a.Server.Handler() + r := httptest.NewRequest("GET", wire.SnapshotPath, nil) + + r.TLS = f.requestState(t) + for _, resource := range []string{"write", "bootstrap", "headers"} { + t.Run(resource, func(t *testing.T) { + request := r.Clone(f.ctx) + + switch resource { + case "write": + take(f.a.Server.writes) + defer release(f.a.Server.writes) + case "bootstrap": + take(f.a.Server.bootstrapSlots) + defer release(f.a.Server.bootstrapSlots) + + request.Method = "POST" + request.URL.Path = wire.BootstrapPath + case "headers": + request.Header.Set("X-Large", strings.Repeat("a", f.a.Server.config.Limits.HeaderBytes)) + } + + w := httptest.NewRecorder() + handler.ServeHTTP(w, request) + + want := 429 + if resource == "headers" { + want = 413 + } + + if w.Code != want { + t.Fatalf("unbounded %s: %d", resource, w.Code) + } + }) + } +} diff --git a/internal/racer/server/transport_test.go b/internal/racer/server/transport_test.go new file mode 100644 index 000000000..dce05a33c --- /dev/null +++ b/internal/racer/server/transport_test.go @@ -0,0 +1,1722 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package server + +import ( + "bytes" + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "errors" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "net/http/httptrace" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/stretchr/testify/require" + corev1 "k8s.io/api/core/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + + "github.com/Azure/unbounded/internal/racer/wire" +) + +func TestStartupGuards(t *testing.T) { + f := newServingFixture(t) + canceled, cancel := context.WithCancel(t.Context()) + cancel() + + for _, tc := range []struct { + name string + mutate func(*Server) + want error + }{ + {"invalid limits", func(s *Server) { s.config.Limits.MaxPolls = 0 }, wire.InvalidRequest}, + {"missing address", func(s *Server) { s.config.ControlAddress = "" }, wire.InvalidRequest}, + {"missing lifecycle", func(s *Server) { s.Lifecycle = nil }, wire.Unavailable}, + {"missing authority", func(s *Server) { s.authority = nil }, wire.Unavailable}, + } { + t.Run(tc.name, func(t *testing.T) { + s := New(f.a.Server.config, f.a.Server.writer, f.a.authority, f.a.Lifecycle, f.a.Replication) + tc.mutate(s) + require.ErrorIs(t, s.Start(canceled), tc.want) + }) + } + + s := New(f.a.Server.config, nil, nil, nil, nil) + require.False(t, s.NeedLeaderElection()) + _, err := s.TLSConfig(canceled) + require.ErrorIs(t, err, wire.Unavailable) + + s = New(Config{}, nil, f.a.authority, nil, nil) + _, err = s.TLSConfig(canceled) + require.ErrorIs(t, err, wire.InvalidRequest) + + lifecycle := NewLifecycle(f.a.authority) + lifecycle.SetCacheSync(func(ctx context.Context) bool { return ctx.Err() == nil }) + require.NoError(t, lifecycle.Start(canceled)) + require.PanicsWithValue(t, http.ErrAbortHandler, func() { flushResponse(canceled, httptest.NewRecorder()) }) +} + +func servingTestCertificate(t *testing.T, serial int64, before, after time.Time, parent *tls.Certificate, ca bool) tls.Certificate { + t.Helper() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + template := &x509.Certificate{SerialNumber: big.NewInt(serial), Subject: pkix.Name{CommonName: fmt.Sprint(serial)}, NotBefore: before, NotAfter: after, DNSNames: []string{"racer-controller.racer.svc"}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, IsCA: ca, BasicConstraintsValid: true} + if ca { + template.KeyUsage |= x509.KeyUsageCertSign + } + + issuer, signer := template, key + if parent != nil { + issuer = parent.Leaf + signer = parent.PrivateKey.(ed25519.PrivateKey) + } + + der, err := x509.CreateCertificate(rand.Reader, template, issuer, pub, signer) + if err != nil { + t.Fatal(err) + } + + leaf, err := x509.ParseCertificate(der) + if err != nil { + t.Fatal(err) + } + + certificate := tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key, Leaf: leaf} + if parent != nil { + certificate.Certificate = append(certificate.Certificate, parent.Certificate...) + } + + return certificate +} + +func writeServingTestPair(t *testing.T, dir string, certificate tls.Certificate) { + t.Helper() + + var chain []byte + for _, der := range certificate.Certificate { + chain = append(chain, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})...) + } + + key, err := x509.MarshalPKCS8PrivateKey(certificate.PrivateKey) + if err != nil { + t.Fatal(err) + } + + for name, data := range map[string][]byte{"tls.crt": chain, "tls.key": pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: key})} { + if err := os.WriteFile(filepath.Join(dir, name), data, 0o600); err != nil { + t.Fatal(err) + } + } +} + +func TestServingCertificateValidation(t *testing.T) { + now := time.Now() + root := servingTestCertificate(t, 1, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + valid := servingTestCertificate(t, 2, now.Add(-time.Minute), now.Add(time.Minute), &root, false) + + other := servingTestCertificate(t, 3, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + for _, tc := range []struct { + name string + certificate tls.Certificate + valid bool + }{ + {"valid chain", valid, true}, + {"leaf only", servingTestCertificate(t, 4, now.Add(-time.Minute), now.Add(time.Minute), nil, false), true}, + {"expired", servingTestCertificate(t, 5, now.Add(-time.Hour), now.Add(-time.Second), &root, false), false}, + {"future", servingTestCertificate(t, 6, now.Add(time.Minute), now.Add(time.Hour), &root, false), false}, + {"wrong chain", tls.Certificate{Certificate: [][]byte{valid.Certificate[0], other.Certificate[0]}, PrivateKey: valid.PrivateKey}, false}, + {"malformed chain", tls.Certificate{Certificate: [][]byte{valid.Certificate[0], {1, 2, 3}}, PrivateKey: valid.PrivateKey}, false}, + {"CA leaf", root, false}, + {"expired issuer", servingTestCertificate(t, 7, now.Add(-time.Minute), now.Add(time.Minute), certificatePointer(servingTestCertificate(t, 8, now.Add(-time.Hour), now.Add(-time.Second), nil, true)), false), false}, + {"future issuer", servingTestCertificate(t, 9, now.Add(-time.Minute), now.Add(time.Minute), certificatePointer(servingTestCertificate(t, 10, now.Add(time.Minute), now.Add(time.Hour), nil, true)), false), false}, + } { + t.Run(tc.name, func(t *testing.T) { + dir := t.TempDir() + writeServingTestPair(t, dir, valid) + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + initial := r.current.Load() + + writeServingTestPair(t, dir, tc.certificate) + + _, err = newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if (err == nil) != tc.valid { + t.Fatalf("valid=%v error=%v", tc.valid, err) + } + + err = r.reload() + if (err == nil) != tc.valid || !tc.valid && r.current.Load() != initial { + t.Fatal("replacement validation or last-valid retention failed") + } + }) + } +} + +func certificatePointer(c tls.Certificate) *tls.Certificate { return &c } + +func TestServingCertificateChainOnlyRotation(t *testing.T) { + now := time.Now() + old := servingTestCertificate(t, 1, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + current := servingTestCertificate(t, 2, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + leaf := servingTestCertificate(t, 3, now.Add(-time.Minute), now.Add(time.Hour), ¤t, false) + bridge, err := x509.CreateCertificate(rand.Reader, current.Leaf, old.Leaf, current.Leaf.PublicKey, old.PrivateKey) + require.NoError(t, err) + + compatible := leaf + compatible.Certificate = [][]byte{leaf.Certificate[0], bridge, old.Certificate[0]} + dir := t.TempDir() + writeServingTestPair(t, dir, compatible) + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + require.NoError(t, err) + testCachedHandshake(t, r, now, old.Leaf, true, 3) + + initial := r.current.Load() + + writeServingTestPair(t, dir, leaf) + require.NoError(t, r.reload()) + require.NotSame(t, initial, r.current.Load()) + selected, err := r.getCertificate(nil) + require.NoError(t, err) + require.Equal(t, compatible.Certificate[0], selected.Certificate[0], "leaf must remain unchanged") + require.Equal(t, leaf.Certificate, selected.Certificate, "chain-only update was ignored") + testCachedHandshake(t, r, now, current.Leaf, true, 2) + testCachedHandshake(t, r, now, old.Leaf, false, 0) + + bad := compatible + bad.Certificate = [][]byte{leaf.Certificate[0], old.Certificate[0]} + last := r.current.Load() + + writeServingTestPair(t, dir, bad) + require.Error(t, r.reload(), "same-leaf invalid chain bypassed validation") + require.Same(t, last, r.current.Load()) + testCachedHandshake(t, r, now, current.Leaf, true, 2) +} + +func TestServerReadinessTracksCachedServingCertificate(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + s := f.a.Server + s.servingCertificate.Store(nil) + + if s.Ready(nil) == nil { + t.Fatal("ready without serving certificate initialization") + } + + r := &servingCertificateReloader{} + s.servingCertificate.Store(r) + + if s.Ready(nil) == nil { + t.Fatal("ready with empty certificate cache") + } + + dir := t.TempDir() + now := time.Now() + first := servingTestCertificate(t, 1, now.Add(-time.Minute), now.Add(2*time.Second), nil, false) + writeServingTestPair(t, dir, first) + + r.certificateFile, r.keyFile = filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key") + if err := r.reload(); err != nil { + t.Fatal(err) + } + + if err := s.Ready(nil); err != nil { + t.Fatal(err) + } + + if err := os.WriteFile(r.certificateFile, []byte("broken projection"), 0o600); err != nil { + t.Fatal(err) + } + + if r.reload() == nil { + t.Fatal("accepted invalid reload") + } + + if err := s.Ready(nil); err != nil { + t.Fatal("lost valid last-good certificate", err) + } + + time.Sleep(2 * time.Second) + + if s.Ready(nil) == nil { + t.Fatal("ready with expired cached certificate") + } + + second := servingTestCertificate(t, 2, now.Add(-time.Minute), now.Add(time.Hour), nil, false) + writeServingTestPair(t, dir, second) + + if err := r.reload(); err != nil { + t.Fatal(err) + } + + if err := s.Ready(nil); err != nil { + t.Fatal("reload did not restore readiness", err) + } + }) +} + +func TestServerReadinessUsesUsableChainPrefix(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + now := time.Now() + old := servingTestCertificate(t, 1, now.Add(-time.Hour), now.Add(time.Second), nil, true) + current := servingTestCertificate(t, 2, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + leaf := servingTestCertificate(t, 3, now.Add(-time.Minute), now.Add(time.Minute), ¤t, false) + bridge := *current.Leaf + bridge.NotAfter = old.Leaf.NotAfter + + der, err := x509.CreateCertificate(rand.Reader, &bridge, old.Leaf, current.Leaf.PublicKey, old.PrivateKey) + if err != nil { + t.Fatal(err) + } + + leaf.Certificate = [][]byte{leaf.Certificate[0], der, old.Certificate[0]} + + r := f.a.Server.installTestServingCertificate(leaf) + if err := f.a.Server.Ready(nil); err != nil { + t.Fatal(err) + } + + time.Sleep(time.Second) + + if err := f.a.Server.Ready(nil); err != nil { + t.Fatal("optional suffix expiration withdrew readiness", err) + } + + selected, err := r.getCertificate(nil) + if err != nil || len(selected.Certificate) != 1 { + t.Fatal("handshake did not select usable prefix", err) + } + }) +} + +func TestServingCertificateProjectionAndRecovery(t *testing.T) { + dir := t.TempDir() + now := time.Now() + first := servingTestCertificate(t, 1, now.Add(-time.Hour), now.Add(time.Hour), nil, false) + + second := servingTestCertificate(t, 2, now.Add(-time.Hour), now.Add(time.Hour), nil, false) + for name, certificate := range map[string]tls.Certificate{"one": first, "two": second} { + path := filepath.Join(dir, name) + require.NoError(t, os.Mkdir(path, 0o700)) + + writeServingTestPair(t, path, certificate) + } + + project := func(generation string) { + t.Helper() + + require.NoError(t, os.Symlink(generation, filepath.Join(dir, "..next"))) + require.NoError(t, os.Rename(filepath.Join(dir, "..next"), filepath.Join(dir, "..data"))) + } + project("one") + + for _, name := range []string{"tls.crt", "tls.key"} { + require.NoError(t, os.Symlink(filepath.Join("..data", name), filepath.Join(dir, name))) + } + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + project("two") + + if err := r.reload(); err != nil { + t.Fatal(err) + } + + if certificate, err := r.getCertificate(nil); err != nil || certificate.Leaf.SerialNumber.Int64() != 2 { + t.Fatal("atomic projection did not replace certificate") + } + + project("one") + + if err := r.reload(); err != nil { + t.Fatal(err) + } + + initial := r.current.Load() + // Deliberately pin the key to the old generation. Reject even when both + // generations reuse the same key and X509KeyPair alone would accept them. + second.PrivateKey = first.PrivateKey + second.Certificate = first.Certificate + writeServingTestPair(t, filepath.Join(dir, "two"), second) + + require.NoError(t, os.Remove(filepath.Join(dir, "tls.key"))) + + require.NoError(t, os.Symlink("one/tls.key", filepath.Join(dir, "tls.key"))) + + project("two") + + require.Error(t, r.reload(), "mixed generation accepted") + require.Same(t, initial, r.current.Load()) + + require.NoError(t, os.Remove(filepath.Join(dir, "tls.key"))) + + require.NoError(t, os.Symlink("..data/tls.key", filepath.Join(dir, "tls.key"))) + + if err := r.reload(); err != nil { + t.Fatal(err) + } + + require.NotSame(t, initial, r.current.Load(), "projection did not recover") + + last := r.current.Load() + + require.NoError(t, os.WriteFile(filepath.Join(dir, "two/tls.crt"), []byte("malformed"), 0o600)) + + if err := r.reload(); err == nil || r.current.Load() != last { + t.Fatal("malformed update replaced certificate") + } + + require.NoError(t, os.Remove(filepath.Join(dir, "two/tls.key"))) + + if err := r.reload(); err == nil || r.current.Load() != last { + t.Fatal("missing update replaced certificate") + } + + writeServingTestPair(t, filepath.Join(dir, "two"), first) + + if err := r.reload(); err != nil { + t.Fatal(err) + } +} + +func TestServingCertificateStandaloneTornPair(t *testing.T) { + dir := t.TempDir() + now := time.Now() + first := servingTestCertificate(t, 1, now.Add(-time.Hour), now.Add(time.Hour), nil, false) + second := servingTestCertificate(t, 2, now.Add(-time.Hour), now.Add(time.Hour), nil, false) + writeServingTestPair(t, dir, first) + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + initial := r.current.Load() + torn := second + torn.PrivateKey = first.PrivateKey + writeServingTestPair(t, dir, torn) + + if err := r.reload(); err == nil || r.current.Load() != initial { + t.Fatal("torn pair accepted") + } + + writeServingTestPair(t, dir, second) + + if err := r.reload(); err != nil { + t.Fatal(err) + } + + if certificate, err := r.getCertificate(nil); err != nil || certificate.Leaf.SerialNumber.Int64() != 2 { + t.Fatal("pair did not recover") + } + + last := r.current.Load() + + file, err := os.OpenFile(filepath.Join(dir, "tls.crt"), os.O_APPEND|os.O_WRONLY, 0o600) + if err != nil { + t.Fatal(err) + } + + _, writeErr := file.WriteString("-----BEGIN CERTIFICATE-----\ntruncated") + + closeErr := file.Close() + if writeErr != nil || closeErr != nil { + t.Fatalf("append: %v, close: %v", writeErr, closeErr) + } + + if err := r.reload(); err == nil || r.current.Load() != last { + t.Fatal("truncated chain replaced certificate") + } +} + +func TestServingCertificateExpirationAndCancellation(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + now := time.Now() + certificate := servingTestCertificate(t, 1, now.Add(-time.Minute), now.Add(2*time.Second), nil, false) + writeServingTestPair(t, dir, certificate) + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(t.Context()) + done := make(chan struct{}) + + go func() { defer close(done); r.run(ctx, time.Second) }() + // Bad files do not invalidate a still-valid cached certificate. + if err := os.Remove(filepath.Join(dir, "tls.crt")); err != nil { + t.Fatal(err) + } + + if _, err := r.getCertificate(nil); err != nil { + t.Fatal(err) + } + + time.Sleep(2 * time.Second) + + if _, err := r.getCertificate(nil); err == nil { + t.Fatal("expired cached certificate served") + } + + cancel() + <-done + + last := r.current.Load() + + writeServingTestPair(t, dir, servingTestCertificate(t, 2, now, now.Add(time.Hour), nil, false)) + time.Sleep(2 * time.Second) + + if r.current.Load() != last { + t.Fatal("reload continued after cancellation") + } + }) +} + +func TestServingTLSConfigInitialFailure(t *testing.T) { + f := newServingFixture(t) + dir := t.TempDir() + + f.configureServer(func(c *Config) { + c.TLSCertificateFile = filepath.Join(dir, "tls.crt") + c.TLSPrivateKeyFile = filepath.Join(dir, "tls.key") + }) + + if _, err := f.a.Server.TLSConfig(f.ctx); err == nil { + t.Fatal("missing initial pair accepted") + } + + writeServingTestPair(t, dir, servingTestCertificate(t, 2, time.Now().Add(time.Hour), time.Now().Add(2*time.Hour), nil, false)) + + if _, err := f.a.Server.TLSConfig(f.ctx); err == nil { + t.Fatal("future initial pair accepted") + } + + writeServingTestPair(t, dir, f.serverCertificate) + + if _, err := f.a.Server.TLSConfig(context.Background()); err == nil { + t.Fatal("unowned reload lifetime accepted") + } + + f.cancel() + + if _, err := f.a.Server.TLSConfig(f.ctx); err == nil { + t.Fatal("canceled reload lifetime accepted") + } +} + +func TestServingTLSLiveReloadPreservesConnectionsAndPolls(t *testing.T) { + f := newServingFixture(t) + dir := t.TempDir() + + f.configureServer(func(c *Config) { + c.TLSCertificateFile = filepath.Join(dir, "tls.crt") + c.TLSPrivateKeyFile = filepath.Join(dir, "tls.key") + }) + + now := time.Now() + root := servingTestCertificate(t, 10, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + first := servingTestCertificate(t, 11, now.Add(-time.Minute), now.Add(time.Hour), &root, false) + // The replacement uses a new CA cross-signed by the old root. A client + // retaining only the old root must still complete fresh TLS handshakes. + newRoot := servingTestCertificate(t, 20, now.Add(-time.Hour), now.Add(time.Hour), nil, true) + + crossDER, err := x509.CreateCertificate(rand.Reader, newRoot.Leaf, root.Leaf, newRoot.Leaf.PublicKey, root.PrivateKey) + require.NoError(t, err) + + second := servingTestCertificate(t, 12, now.Add(-time.Minute), now.Add(time.Hour), &newRoot, false) + second.Certificate = [][]byte{second.Certificate[0], crossDER, root.Certificate[0]} + f.roots = x509.NewCertPool() + f.roots.AddCert(root.Leaf) + writeServingTestPair(t, dir, first) + + config, err := f.a.Server.TLSConfig(f.ctx) + require.NoError(t, err) + + s := httptest.NewUnstartedServer(f.a.Server.Handler()) + s.Config.ConnContext = connectionContext + s.TLS = config + s.StartTLS() + t.Cleanup(s.Close) + peer := f.client(t, &f.certificate) + response, err := peer.Get(s.URL + wire.SnapshotPath) + body := responseBody(t, response, err, http.StatusOK) + + publication, err := wire.DecodePublication(bytes.NewReader(body)) + if err != nil { + t.Fatal(err) + } + + require.Equal(t, int64(11), response.TLS.PeerCertificates[0].SerialNumber.Int64(), "wrong initial certificate") + + type result struct { + response *http.Response + err error + } + + done := make(chan result, 1) + + go func() { + pollResponse, pollErr := peer.Get(fmt.Sprintf("%s%s?after=%d", s.URL, wire.SnapshotPath, publication.Sequence)) + done <- result{pollResponse, pollErr} + }() + + awaitServerPolls(t, f.a.Server, 1) + writeServingTestPair(t, dir, second) + awaitServingSerial(t, f, s.URL, 12) + + select { + case result := <-done: + if result.response != nil { + result.response.Body.Close() + } + + t.Fatalf("rotation interrupted long poll: %v", result.err) + default: + } + + node := &corev1.Node{} + if err := f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, node); err != nil { + t.Fatal(err) + } + + node.Labels = map[string]string{wire.ExclusionLabel: ""} + if err := f.a.Topology.Update(f.ctx, node); err != nil { + t.Fatal(err) + } + + reconcileTopology(t, f.a.Topology, f.ctx) + + select { + case result := <-done: + responseBody(t, result.response, result.err, http.StatusOK) + + require.Equal(t, int64(11), result.response.TLS.PeerCertificates[0].SerialNumber.Int64(), "long poll reconnected") + case <-time.After(5 * time.Second): + t.Fatal("long poll did not complete") + } + + reused := false + ctx := httptrace.WithClientTrace(f.ctx, &httptrace.ClientTrace{GotConn: func(info httptrace.GotConnInfo) { reused = info.Reused }}) + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, s.URL+wire.SnapshotPath, nil) + if err != nil { + t.Fatal(err) + } + + response, err = peer.Do(req) + responseBody(t, response, err, http.StatusOK) + + require.True(t, reused, "persistent connection replaced") + require.Equal(t, int64(11), response.TLS.PeerCertificates[0].SerialNumber.Int64()) +} + +func awaitServingSerial(t *testing.T, f *servingFixture, endpoint string, serial int64) { + t.Helper() + fresh := f.client(t, nil) + fresh.Transport.(*http.Transport).DisableKeepAlives = true + deadline := time.Now().Add(5 * time.Second) + + for { + response, err := fresh.Get(endpoint + wire.SnapshotPath) + responseBody(t, response, err, http.StatusUnauthorized) + + if response.TLS.PeerCertificates[0].SerialNumber.Int64() == serial { + return + } + + if time.Now().After(deadline) { + t.Fatal("new handshakes did not see replacement") + } + + time.Sleep(10 * time.Millisecond) + } +} + +func TestServingCertificateConcurrentReload(t *testing.T) { + dir := t.TempDir() + now := time.Now() + first := servingTestCertificate(t, 1, now.Add(-time.Hour), now.Add(time.Hour), nil, false) + second := servingTestCertificate(t, 2, now.Add(-time.Hour), now.Add(time.Hour), nil, false) + writeServingTestPair(t, dir, first) + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + var wg sync.WaitGroup + for range 8 { + wg.Go(func() { + for range 1000 { + certificate, err := r.getCertificate(nil) + if err != nil || certificate.Leaf.SerialNumber.Int64() < 1 || certificate.Leaf.SerialNumber.Int64() > 2 { + t.Error("invalid concurrent certificate") + } + } + }) + } + + for range 10 { + for _, certificate := range []tls.Certificate{second, first} { + writeServingTestPair(t, dir, certificate) + + if err := r.reload(); err != nil { + t.Fatal(err) + } + } + } + + wg.Wait() +} + +func TestServingTLSRejectsExpiredCachedCertificate(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + f := newServingFixture(t) + dir := t.TempDir() + now := time.Now() + certificate := servingTestCertificate(t, 1, now.Add(-time.Minute), now.Add(time.Second), nil, false) + writeServingTestPair(t, dir, certificate) + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + config := f.a.Server.tlsConfigWithCertificate(f.ctx, r.getCertificate) + // Advance past expiration without any reload. Even a client disabling + // verification cannot make the server disclose an expired cached chain. + time.Sleep(time.Second) + + serverSide, clientSide := net.Pipe() + defer serverSide.Close() + defer clientSide.Close() + + done := make(chan error, 1) + + go func() { done <- tls.Server(serverSide, config).HandshakeContext(f.ctx) }() + + peer := tls.Client(clientSide, &tls.Config{MinVersion: tls.VersionTLS13, InsecureSkipVerify: true}) //nolint:gosec // Verify rejection by the server, not the client. + if err := peer.HandshakeContext(f.ctx); err == nil { + t.Fatal("expired server certificate accepted") + } + + if err := <-done; err == nil { + t.Fatal("server completed expired handshake") + } + + if len(peer.ConnectionState().PeerCertificates) != 0 { + t.Fatal("expired chain sent to client") + } + }) +} + +func TestServingTLSCompatibilitySuffixExpiration(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + old := servingTestCertificate(t, 100, now.Add(-27*24*time.Hour), now.Add(24*time.Hour), nil, true) + current := servingTestCertificate(t, 101, now.Add(-time.Hour), now.Add(28*24*time.Hour), nil, true) + leaf := servingTestCertificate(t, 102, now.Add(-time.Minute), now.Add(7*24*time.Hour), ¤t, false) + bridgeTemplate := *current.Leaf + bridgeTemplate.NotAfter = old.Leaf.NotAfter + + bridgeDER, err := x509.CreateCertificate(rand.Reader, &bridgeTemplate, old.Leaf, current.Leaf.PublicKey, old.PrivateKey) + require.NoError(t, err) + + leaf.Certificate = [][]byte{leaf.Certificate[0], bridgeDER, old.Certificate[0]} + dir := t.TempDir() + writeServingTestPair(t, dir, leaf) + + r := &servingCertificateReloader{certificateFile: filepath.Join(dir, "tls.crt"), keyFile: filepath.Join(dir, "tls.key")} + require.NoError(t, r.reloadAt(now)) + + initial := r.current.Load() + // Remove source files: selection after the frozen expiration boundary must + // work using only immutable cached, prevalidated chain prefixes. + require.NoError(t, os.Remove(r.certificateFile)) + require.NoError(t, os.Remove(r.keyFile)) + + handshake := func(at time.Time, root *x509.Certificate, want bool, length int) { + t.Helper() + testCachedHandshake(t, r, at, root, want, length) + } + handshake(now, old.Leaf, true, 3) + handshake(now, current.Leaf, true, 3) + after := old.Leaf.NotAfter + handshake(after, current.Leaf, true, 1) + handshake(after, old.Leaf, false, 0) + handshake(leaf.Leaf.NotAfter, current.Leaf, false, 0) + + require.Same(t, initial, r.current.Load(), "handshake changed published certificate") + // The operator normally sends leaf + bridges, omitting the old root. + // Exercise that wire layout as well as the explicit-root layout above. + withoutRoot := leaf + withoutRoot.Certificate = leaf.Certificate[:2:2] + writeServingTestPair(t, dir, withoutRoot) + + require.NoError(t, r.reloadAt(now)) + + handshake(now, old.Leaf, true, 2) + handshake(after, current.Leaf, true, 1) + handshake(after, old.Leaf, false, 0) + + require.NoError(t, r.reloadAt(after), "expired bridge prevented initial load") + // Initial/repeated loading of an unpruned Secret must also accept the + // current path, including a leaf issued after the old bridge expired. + newLeaf := servingTestCertificate(t, 103, after.Add(time.Hour), after.Add(7*24*time.Hour), ¤t, false) + newLeaf.Certificate = [][]byte{newLeaf.Certificate[0], bridgeDER, old.Certificate[0]} + writeServingTestPair(t, dir, newLeaf) + + require.NoError(t, r.reloadAt(after.Add(2*time.Hour))) + + handshake(after.Add(2*time.Hour), current.Leaf, true, 1) + last := r.current.Load() + + for _, scenario := range []string{"future bridge", "wrong EKU", "path length", "malformed suffix", "expired leaf"} { + t.Run(scenario, func(t *testing.T) { + bad := newLeaf + bad.Certificate = append([][]byte(nil), newLeaf.Certificate...) + template := bridgeTemplate + + switch scenario { + case "future bridge": + template.NotBefore, template.NotAfter = after.Add(3*time.Hour), after.Add(4*time.Hour) + case "wrong EKU": + template.ExtKeyUsage = []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth} + case "path length": + rootTemplate := *old.Leaf + rootTemplate.MaxPathLen, rootTemplate.MaxPathLenZero = 0, true + + der, err := x509.CreateCertificate(rand.Reader, &rootTemplate, &rootTemplate, old.Leaf.PublicKey, old.PrivateKey) + require.NoError(t, err) + + bad.Certificate[2] = der + case "malformed suffix": + bad.Certificate[2] = []byte{1, 2, 3} + case "expired leaf": + bad = leaf + } + + if scenario == "future bridge" || scenario == "wrong EKU" { + der, err := x509.CreateCertificate(rand.Reader, &template, old.Leaf, current.Leaf.PublicKey, old.PrivateKey) + require.NoError(t, err) + + bad.Certificate[1] = der + } + + writeServingTestPair(t, dir, bad) + + at := after.Add(2 * time.Hour) + if scenario == "expired leaf" { + at = leaf.Leaf.NotAfter + } + + require.Error(t, r.reloadAt(at), "invalid replacement accepted") + require.Same(t, last, r.current.Load(), "last good lost") + }) + } +} + +func testCachedHandshake(t *testing.T, r *servingCertificateReloader, at time.Time, root *x509.Certificate, want bool, length int) { + t.Helper() + + roots := x509.NewCertPool() + roots.AddCert(root) + + serverSide, clientSide := net.Pipe() + defer serverSide.Close() + defer clientSide.Close() + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + + done := make(chan error, 1) + + go func() { + defer serverSide.Close() + + server := tls.Server(serverSide, &tls.Config{MinVersion: tls.VersionTLS13, GetCertificate: func(*tls.ClientHelloInfo) (*tls.Certificate, error) { return r.getCertificateAt(at) }}) + done <- server.HandshakeContext(ctx) + }() + + peer := tls.Client(clientSide, &tls.Config{MinVersion: tls.VersionTLS13, RootCAs: roots, ServerName: "127.0.0.1", Time: func() time.Time { return at }}) + err := peer.HandshakeContext(ctx) + require.Equal(t, want, err == nil, "handshake at %s: %v", at, err) + + if want { + require.Len(t, peer.ConnectionState().PeerCertificates, length) + } + + clientSide.Close() + + err = <-done + if want { + require.NoError(t, err) + } +} + +type temporaryAcceptError struct{} + +func (temporaryAcceptError) Error() string { return "temporary accept failure" } + +func (temporaryAcceptError) Timeout() bool { return false } + +func (temporaryAcceptError) Temporary() bool { return true } + +func TestTransportTemporaryAcceptRecoversTLS(t *testing.T) { + f := newServingFixture(t) + + raw, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + + first := true + listener := teardownListener{Listener: raw, close: raw.Close, accept: func() (net.Conn, error) { + if first { + first = false + return nil, fmt.Errorf("accept: %w", temporaryAcceptError{}) + } + + return raw.Accept() + }} + l := newTransportListener(f.ctx, listener, f.a.Server.tlsConfig(f.ctx, f.serverCertificate), Limits{MaxConnections: 1, MaxConcurrentHandshakes: 1, HandshakeTimeout: time.Second, WriteTimeout: 30 * time.Second}) + + t.Cleanup(func() { + if err := l.Close(); err != nil { + t.Error(err) + } + + select { + case <-l.acceptDone: + case <-time.After(time.Second): + t.Error("accept pump leaked") + } + + awaitTransport(t, l, 0, 0) + }) + + ctx, cancel := context.WithTimeout(f.ctx, 3*time.Second) + defer cancel() + + client := tls.Client(dialTransport(t, l), &tls.Config{RootCAs: f.roots, ServerName: "127.0.0.1", MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{f.certificate}}) + if err := client.HandshakeContext(ctx); err != nil { + t.Fatal(err) + } + + conn, err := l.Accept() + if err != nil { + t.Fatal(err) + } + + secured, ok := conn.(*tls.Conn) + if !ok || !secured.ConnectionState().HandshakeComplete || len(secured.ConnectionState().VerifiedChains) == 0 { + t.Fatal("recovered accept lost TLS or client certificate") + } + + closeTransport(conn) + awaitTransport(t, l, 0, 0) +} + +func TestTransportTemporaryAcceptBackoffCapResetAndTerminal(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + accepted, peer := net.Pipe() + defer accepted.Close() + defer peer.Close() + + terminal := errors.New("terminal accept failure") + + var calls []time.Time + + listener := teardownListener{close: func() error { return nil }, accept: func() (net.Conn, error) { + calls = append(calls, time.Now()) + switch len(calls) { + case 11: + return accepted, nil + case 13: + return nil, terminal + default: + return nil, temporaryAcceptError{} + } + }} + // Zero connection capacity rejects the successful accept without a TLS + // worker; even rejected sockets must reset the accept-error backoff. + l := newTransportListener(t.Context(), listener, &tls.Config{}, Limits{}) + defer l.Close() + + <-l.acceptDone + + want := []time.Duration{5 * time.Millisecond, 10 * time.Millisecond, 20 * time.Millisecond, 40 * time.Millisecond, 80 * time.Millisecond, 160 * time.Millisecond, 320 * time.Millisecond, 640 * time.Millisecond, time.Second, time.Second, 0, 5 * time.Millisecond} + if len(calls) != len(want)+1 { + t.Fatalf("accept calls=%d", len(calls)) + } + + for i, delay := range want { + if got := calls[i+1].Sub(calls[i]); got != delay { + t.Fatalf("retry %d delay=%s want=%s", i, got, delay) + } + } + + for range 2 { + if conn, err := l.Accept(); conn != nil || !errors.Is(err, terminal) { + t.Fatalf("terminal accept=%v, %v", conn, err) + } + } + + if len(l.connections) != 0 || len(l.handshakes) != 0 { + t.Fatal("accept errors leaked admission") + } + }) +} + +func TestTransportTemporaryAcceptCancellation(t *testing.T) { + for _, closeListener := range []bool{false, true} { + t.Run(fmt.Sprintf("close=%v", closeListener), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + calls := 0 + listener := teardownListener{close: func() error { return nil }, accept: func() (net.Conn, error) { + calls++ + return nil, temporaryAcceptError{} + }} + + l := newTransportListener(ctx, listener, &tls.Config{}, Limits{}) + defer l.Close() + // Reach the capped retry window. Virtual time and exact call + // counts prove the pump sleeps rather than spinning/spawning work. + time.Sleep(1500 * time.Millisecond) + synctest.Wait() + + require.Equal(t, 9, calls, "accept calls") + + start := time.Now() + + if closeListener { + require.NoError(t, l.Close()) + } else { + cancel() + } + + synctest.Wait() + + select { + case <-l.acceptDone: + default: + t.Fatal("cancellation left accept pump sleeping") + } + + conn, err := l.Accept() + require.Nil(t, conn) + require.ErrorIs(t, err, net.ErrClosed) + require.Zero(t, time.Since(start)) + require.Equal(t, 9, calls) + require.Empty(t, l.connections) + require.Empty(t, l.handshakes) + }) + }) + } +} + +func testTransport(t *testing.T, f *servingFixture, connections, handshakes int, deadline time.Duration) *transportListener { + t.Helper() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + + l := newTransportListenerWithMetrics(f.ctx, listener, f.a.Server.tlsConfig(f.ctx, f.serverCertificate), Limits{MaxConnections: connections, MaxConcurrentHandshakes: handshakes, HandshakeTimeout: deadline, WriteTimeout: 30 * time.Second}, newTransportMetrics(prometheus.NewPedanticRegistry())) + + t.Cleanup(func() { + if err := l.Close(); err != nil { + t.Error(err) + } + + awaitTransport(t, l, 0, 0) + }) + + return l +} + +func awaitTransport(t *testing.T, l *transportListener, connections, handshakes int) { + t.Helper() + + deadline := time.Now().Add(3 * time.Second) + for len(l.connections) != connections || len(l.handshakes) != handshakes { + if time.Now().After(deadline) { + t.Fatalf("transport slots: connections=%d handshakes=%d, want %d/%d", len(l.connections), len(l.handshakes), connections, handshakes) + } + + time.Sleep(time.Millisecond) + } +} + +func dialTransport(t *testing.T, l *transportListener) net.Conn { + t.Helper() + + conn, err := net.DialTimeout("tcp", l.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + + t.Cleanup(func() { closeTransport(conn) }) + + return conn +} + +func expectTransportClosed(t *testing.T, conn net.Conn) { + t.Helper() + + if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + + _, err := conn.Read(make([]byte, 1)) + + var timeout net.Error + if err == nil || errors.As(err, &timeout) && timeout.Timeout() { + t.Fatalf("excess or closed transport retained: %v", err) + } +} + +func TestTransportSilentPeersBoundAndRelease(t *testing.T) { + for _, limits := range []struct { + name string + connections, handshakes int + }{ + {"connections", 2, 3}, {"handshakes", 3, 2}, + } { + t.Run(limits.name, func(t *testing.T) { + f := newServingFixture(t) + l := testTransport(t, f, limits.connections, limits.handshakes, time.Minute) + first, second := dialTransport(t, l), dialTransport(t, l) + awaitTransport(t, l, 2, 2) + assertTransportMetrics(t, l, 2, 2, 0, 0, 0) + // Rejection never enters a per-peer waiter or handshake goroutine. + for range 16 { + expectTransportClosed(t, dialTransport(t, l)) + } + + awaitTransport(t, l, 2, 2) + + var connectionRejected, handshakeRejected float64 + if limits.name == "connections" { + connectionRejected = 16 + } else { + handshakeRejected = 16 + } + + assertTransportMetrics(t, l, 2, 2, connectionRejected, handshakeRejected, 0) + closeTransport(first) + awaitTransport(t, l, 1, 1) + assertTransportMetrics(t, l, 1, 1, connectionRejected, handshakeRejected, 0) + dialTransport(t, l) + awaitTransport(t, l, 2, 2) + + if err := l.Close(); err != nil { + t.Fatal(err) + } + + expectTransportClosed(t, second) + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, connectionRejected, handshakeRejected, 0) + }) + } +} + +func TestTransportHandshakeDeadlineAndFailure(t *testing.T) { + f := newServingFixture(t) + l := testTransport(t, f, 1, 1, 100*time.Millisecond) + conn := dialTransport(t, l) + awaitTransport(t, l, 1, 1) + assertTransportMetrics(t, l, 1, 1, 0, 0, 0) + expectTransportClosed(t, conn) + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, 0, 0, 1) + + conn = dialTransport(t, l) + if _, err := io.WriteString(conn, "not a TLS record"); err != nil { + t.Fatal(err) + } + + expectTransportClosed(t, conn) + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, 0, 0, 1) +} + +func TestTransportPartialFlightTimeoutAndRecovery(t *testing.T) { + f := newServingFixture(t) + l := testTransport(t, f, 1, 1, 200*time.Millisecond) + entered, resume := make(chan struct{}), make(chan struct{}) + + unblock := sync.OnceFunc(func() { close(resume) }) + defer unblock() + + client := tls.Client(dialTransport(t, l), &tls.Config{ + RootCAs: f.roots, ServerName: "127.0.0.1", MinVersion: tls.VersionTLS13, + GetClientCertificate: func(*tls.CertificateRequestInfo) (*tls.Certificate, error) { + close(entered) + <-resume + + return &f.certificate, nil + }, + }) + done := make(chan error, 1) + + go func() { done <- client.HandshakeContext(f.ctx) }() + + select { + case <-entered: + case <-time.After(3 * time.Second): + t.Fatal("client did not receive the server flight") + } + + awaitTransport(t, l, 1, 1) + assertTransportMetrics(t, l, 1, 1, 0, 0, 0) + expectTransportClosed(t, dialTransport(t, l)) + // No final client flight arrives. Both budgets must recover without waiting + // for the independent 30-second response write budget. + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, 1, 0, 1) + unblock() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("partial client handshake leaked") + } + + fresh := tls.Client(dialTransport(t, l), &tls.Config{RootCAs: f.roots, ServerName: "127.0.0.1", MinVersion: tls.VersionTLS13}) + require.NoError(t, fresh.HandshakeContext(f.ctx), "admission did not recover") + + conn, err := l.Accept() + require.NoError(t, err) + awaitTransport(t, l, 1, 0) + + var closes sync.WaitGroup + for range 8 { + closes.Go(func() { closeTransport(conn) }) + } + + closes.Go(func() { + if err := l.Close(); err != nil { + t.Error(err) + } + }) + closes.Wait() + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, 1, 0, 1) +} + +func TestTransportHandshakeBudgetThroughClientCertificate(t *testing.T) { + f := newServingFixture(t) + l := testTransport(t, f, 2, 1, time.Minute) + + entered, resume := make(chan struct{}), make(chan struct{}) + defer close(resume) + + raw := dialTransport(t, l) + conn := tls.Client(raw, &tls.Config{ + RootCAs: f.roots, ServerName: "127.0.0.1", MinVersion: tls.VersionTLS13, + GetClientCertificate: func(*tls.CertificateRequestInfo) (*tls.Certificate, error) { + close(entered) + <-resume + + return &f.certificate, nil + }, + }) + done := make(chan error, 1) + + go func() { done <- conn.HandshakeContext(f.ctx) }() + + select { + case <-entered: + case <-time.After(3 * time.Second): + t.Fatal("client certificate callback not reached") + } + // ServerHello was already delivered: GetConfigForClient has returned, but + // the server must retain admission while waiting for the client's flight. + awaitTransport(t, l, 1, 1) + assertTransportMetrics(t, l, 1, 1, 0, 0, 0) + expectTransportClosed(t, dialTransport(t, l)) + assertTransportMetrics(t, l, 1, 1, 0, 1, 0) + + if err := l.Close(); err != nil { + t.Fatal(err) + } + + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, 0, 1, 0) + // The client callback is deliberately parked; resume it during cleanup. + t.Cleanup(func() { + select { + case <-done: + case <-time.After(3 * time.Second): + t.Error("client handshake leaked") + } + }) +} + +func TestTransportCompletedTLSRetainsOnlyConnectionSlot(t *testing.T) { + f := newServingFixture(t) + l := testTransport(t, f, 1, 1, time.Second) + + client := tls.Client(dialTransport(t, l), &tls.Config{RootCAs: f.roots, ServerName: "127.0.0.1", MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{f.certificate}}) + if err := client.HandshakeContext(f.ctx); err != nil { + t.Fatal(err) + } + + awaitTransport(t, l, 1, 0) + assertTransportMetrics(t, l, 1, 0, 0, 0, 0) + + conn, err := l.Accept() + if err != nil { + t.Fatal(err) + } + + secured, ok := conn.(*tls.Conn) + if !ok || !secured.ConnectionState().HandshakeComplete || len(secured.ConnectionState().VerifiedChains) == 0 { + t.Fatal("listener lost concrete TLS type or client certificate authentication") + } + + expectTransportClosed(t, dialTransport(t, l)) + assertTransportMetrics(t, l, 1, 0, 1, 0, 0) + // Both TLS Close and force-close may happen concurrently; release once. + closeTransport(conn) + closeTransport(conn) + awaitTransport(t, l, 0, 0) + assertTransportMetrics(t, l, 0, 0, 1, 0, 0) +} + +func TestTransportProductionLongPollAndIdleAdmission(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { + c.Limits.MaxConnections = 2 + c.Limits.MaxConcurrentHandshakes = 1 + c.Limits.HandshakeTimeout = 100 * time.Millisecond + c.Limits.WriteTimeout = time.Second + }) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + + t.Cleanup(func() { _ = listener.Close() }) + + done := make(chan error, 1) + + go func() { done <- f.a.Server.serve(f.ctx, listener, f.a.Server.tlsConfig(f.ctx, f.serverCertificate)) }() + + t.Cleanup(func() { + f.cancel() + + select { + case err := <-done: + if err != nil { + t.Error(err) + } + case <-time.After(3 * time.Second): + t.Error("serve shutdown blocked") + } + }) + client := f.client(t, &f.certificate) + + publication, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + requestDone := make(chan error, 1) + + go func() { + response, err := client.Get(fmt.Sprintf("https://%s/v1/snapshot?after=%d", listener.Addr(), publication.Sequence())) + if response != nil { + response.Body.Close() + } + + requestDone <- err + }() + + awaitServerPolls(t, f.a.Server, 1) + // A long poll outlives the handshake deadline and does not monopolize it. + time.Sleep(150 * time.Millisecond) + + select { + case err := <-requestDone: + t.Fatalf("handshake timeout interrupted an established poll: %v", err) + default: + } + + other := f.client(t, &f.certificate) + response, err := other.Get("https://" + listener.Addr().String() + "/invalid") + responseBody(t, response, err, http.StatusBadRequest) + // Both the poll and the HTTP keep-alive consume their connection slots. + raw, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + defer raw.Close() + + expectTransportClosed(t, raw) + f.cancel() + + select { + case <-requestDone: + case <-time.After(3 * time.Second): + t.Fatal("poll retained on shutdown") + } + + awaitServerPolls(t, f.a.Server, 0) +} + +func TestTransportLimitsValidation(t *testing.T) { + for _, field := range []string{"connections", "handshakes"} { + for _, value := range []int{0, -1} { + cfg := testConfig(t).ServerConfig + + if field == "connections" { + cfg.Limits.MaxConnections = value + } else { + cfg.Limits.MaxConcurrentHandshakes = value + } + + if cfg.Validate() == nil { + t.Fatalf("accepted %s=%d", field, value) + } + } + } + + for _, value := range []time.Duration{0, -time.Nanosecond, time.Nanosecond, 5 * time.Second} { + cfg := testConfig(t).ServerConfig + + cfg.Limits.HandshakeTimeout = value + if value <= 0 { + require.ErrorIs(t, cfg.Validate(), wire.InvalidRequest) + } else { + require.NoError(t, cfg.Validate()) + cfg.Limits.WriteTimeout = 0 + require.ErrorIs(t, cfg.Validate(), wire.InvalidRequest, "handshake timeout must not replace write validation") + } + } +} + +var _ net.Listener = (*transportListener)(nil) + +func TestServingReadinessRequiresConfiguredHostname(t *testing.T) { + for _, names := range [][]string{nil, {"other-service.racer.svc"}, {"racer-controller.racer.svc"}, {"*.racer.svc"}} { + t.Run("names="+strings.Join(names, ","), func(t *testing.T) { + f := newServingFixture(t) + certificate := servingTestCertificate(t, 1, time.Now().Add(-time.Minute), time.Now().Add(time.Hour), nil, false) + leaf := *certificate.Leaf + leaf.DNSNames = names + // A matching CommonName alone must not bypass SAN validation. + leaf.Subject.CommonName = f.a.Server.config.ReplicationServerName + + der, err := x509.CreateCertificate(rand.Reader, &leaf, &leaf, leaf.PublicKey, certificate.PrivateKey) + if err != nil { + t.Fatal(err) + } + + certificate = tls.Certificate{Certificate: [][]byte{der}, PrivateKey: certificate.PrivateKey} + dir := t.TempDir() + writeServingTestPair(t, dir, certificate) + + r, err := newServingCertificateReloader(filepath.Join(dir, "tls.crt"), filepath.Join(dir, "tls.key")) + if err != nil { + t.Fatal(err) + } + + f.a.Server.servingCertificate.Store(r) + + wantReady := len(names) != 0 && names[0] != "other-service.racer.svc" + if ready := f.a.Server.Ready(nil) == nil; ready != wantReady { + t.Fatalf("ready=%v want=%v", ready, wantReady) + } + // Reloading a correctly named certificate repairs readiness without + // restart; mismatched names do not corrupt issuer trust/readiness. + writeServingTestPair(t, dir, f.serverCertificate) + + if err := r.reload(); err != nil { + t.Fatal(err) + } + + if err := f.a.Server.Ready(nil); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestLeaderCancellationClosesActiveTLSPoll(t *testing.T) { + f := newServingFixture(t) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + + done := make(chan error, 1) + + go func() { done <- f.a.Server.serve(f.ctx, listener, f.a.Server.tlsConfig(f.ctx, f.serverCertificate)) }() + + c := f.client(t, &f.certificate) + + publication, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + requestDone := make(chan error, 1) + + go func() { + response, err := c.Get(fmt.Sprintf("https://%s/v1/snapshot?after=%d", listener.Addr(), publication.Sequence())) + if response != nil { + response.Body.Close() + } + + requestDone <- err + }() + + awaitServerPolls(t, f.a.Server, 1) + f.cancel() + + select { + case <-requestDone: + case <-time.After(time.Second): + t.Fatal("leader cancellation left active poll") + } + + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("server shutdown blocked") + } + + awaitServerPolls(t, f.a.Server, 0) +} + +func TestTLSPollExpirationAndRequestCancellation(t *testing.T) { + for _, expiration := range []bool{true, false} { + t.Run(fmt.Sprint(expiration), func(t *testing.T) { + f := newServingFixture(t) + + cert := f.certificate + if expiration { + cert = f.signLeaf(t, func(c *x509.Certificate) { c.NotAfter = time.Now().Add(2 * time.Second).Truncate(time.Second) }) + } + + endpoint := f.start(t) + c := f.client(t, &cert) + + publication, err := f.a.authority.Current() + require.NoError(t, err) + + ctx, cancel := context.WithCancel(f.ctx) + defer cancel() + + r, err := http.NewRequestWithContext(ctx, "GET", fmt.Sprintf("%s/v1/snapshot?after=%d", endpoint, publication.Sequence()), nil) + require.NoError(t, err) + + done := make(chan struct{}) + + go func() { + defer close(done) + + response, err := c.Do(r) + if expiration { + responseBody(t, response, err, 401) + } else { + if response != nil { + response.Body.Close() + } + + if err == nil { + t.Error("canceled poll succeeded") + } + } + }() + + awaitServerPolls(t, f.a.Server, 1) + + if !expiration { + cancel() + } + + select { + case <-done: + case <-time.After(4 * time.Second): + t.Fatal("poll outlived expiration/cancellation") + } + + awaitServerPolls(t, f.a.Server, 0) + }) + } +} + +func TestBootstrapReadDeadlineAndChunkedBound(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.WriteTimeout = 100 * time.Millisecond }) + endpoint := f.start(t) + c := f.client(t, nil) + // A body of unknown length must still be bounded by the wire decoder. + r, err := http.NewRequestWithContext(f.ctx, "POST", endpoint+wire.BootstrapPath, io.NopCloser(strings.NewReader(strings.Repeat("x", wire.MaxBootstrapBytes+1)))) + if err != nil { + t.Fatal(err) + } + + r.Header.Set("Content-Type", "application/json") + response, err := c.Do(r) + responseBody(t, response, err, 413) + // A client that never completes its body cannot hold bootstrap admission. + address := strings.TrimPrefix(endpoint, "https://") + + conn, err := tls.Dial("tcp", address, &tls.Config{RootCAs: f.roots, MinVersion: tls.VersionTLS13}) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + if _, err := fmt.Fprintf(conn, "POST /v1/bootstrap HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: 100\r\n\r\n{"); err != nil { + t.Fatal(err) + } + + if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + + _, _ = io.Copy(io.Discard, conn) + deadline := time.After(time.Second) + + for len(f.a.Server.bootstrapSlots) != 0 { + select { + case <-deadline: + t.Fatal("slow body retained admission") + default: + time.Sleep(time.Millisecond) + } + } +} + +func TestTLSSlowSnapshotWriteDeadline(t *testing.T) { + f := newServingFixture(t) + f.configureServer(func(c *Config) { c.Limits.WriteTimeout = 200 * time.Millisecond }) + // Exercise socket backpressure without constructing a large topology. The + // immutable publication remains valid JSON with bounded trailing whitespace. + largeFixturePublication(t, f) + endpoint := f.start(t) + + conn, err := tls.Dial("tcp", strings.TrimPrefix(endpoint, "https://"), &tls.Config{RootCAs: f.roots, MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{f.certificate}}) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + + if _, err := io.WriteString(conn, "GET /v1/snapshot HTTP/1.1\r\nHost: localhost\r\n\r\n"); err != nil { + t.Fatal(err) + } + + deadline := time.After(3 * time.Second) + + for len(f.a.Server.writes) == 0 { + select { + case <-deadline: + t.Fatal("write not admitted") + default: + time.Sleep(time.Millisecond) + } + } + // Do not read response bytes. The write deadline must release both slots. + for { + n := f.a.Server.polls.count() + + if n == 0 && len(f.a.Server.writes) == 0 { + break + } + + select { + case <-deadline: + t.Fatal("slow socket bypassed write deadline") + default: + time.Sleep(time.Millisecond) + } + } +} + +func TestTLSNodeExclusionRemovesRoutingMembershipWhilePolling(t *testing.T) { + f := newServingFixture(t) + endpoint := f.start(t) + c := f.client(t, &f.certificate) + + publication, err := f.a.authority.Current() + if err != nil { + t.Fatal(err) + } + + done := make(chan struct{}) + + go func() { + defer close(done) + + response, err := c.Get(fmt.Sprintf("%s/v1/snapshot?after=%d", endpoint, publication.Sequence())) + body := responseBody(t, response, err, 200) + + updated, decodeErr := wire.DecodePublication(bytes.NewReader(body)) + if decodeErr != nil || len(updated.Members) != 0 { + t.Errorf("exclusion must remove routing membership: %v", decodeErr) + } + }() + + awaitServerPolls(t, f.a.Server, 1) + + node := &corev1.Node{} + if err := f.a.Topology.Get(f.ctx, client.ObjectKey{Name: "worker"}, node); err != nil { + t.Fatal(err) + } + + node.Labels = map[string]string{wire.ExclusionLabel: ""} + if err := f.a.Topology.Update(f.ctx, node); err != nil { + t.Fatal(err) + } + + reconcileTopology(t, f.a.Topology, f.ctx) + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("routing membership poll did not wake after exclusion") + } +} diff --git a/internal/racer/testutil/workload.go b/internal/racer/testutil/workload.go new file mode 100644 index 000000000..717eff867 --- /dev/null +++ b/internal/racer/testutil/workload.go @@ -0,0 +1,469 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// Package testutil provides external workload fixtures for controller tests. +// It is not imported by the controller or used to deploy workloads. +package testutil + +import ( + "encoding/json" + "fmt" + "net/url" + "slices" + "strconv" + "strings" + + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/util/intstr" + "k8s.io/apimachinery/pkg/util/validation" + "k8s.io/utils/ptr" + + "github.com/Azure/unbounded/internal/racer/wire" +) + +// Workload configuration and construction. + +const ( + DataplaneDaemonSetName = "racer-dataplane" + PodNetworkDaemonSetName = "racer-dataplane-podnet" +) + +// ManagedNames contains only the explicitly configured workload. +func ManagedNames(daemonSetName string) []string { + return []string{daemonSetName} +} + +// Config contains only the inputs needed to build the dataplane DaemonSet. +// Controller limits, rotation policy, serving TLS, and durable state are independent. +type Config struct { + Cluster wire.ClusterID + Namespace string + ControlURL string + BootstrapTrustConfigMap string + DataplaneImage string + PeerPort uint16 + HostNetwork bool + PodNetworkNodes []string + // Zero preserves the legacy automatic diagnostics port (9090, or 9091). + DiagnosticsPort uint16 + DataplaneServiceAccount string + DaemonSetName string +} + +// ConfigFromLookup reads operator deployment wiring without loading or +// validating controller runtime configuration. Shared settings retain the same defaults. +func ConfigFromLookup(lookup func(string) (string, bool)) (Config, error) { + env := func(key, fallback string) string { + if value, ok := lookup(key); ok { + return value + } + + return fallback + } + + port, err := strconv.ParseUint(env("RACER_PEER_PORT", "8082"), 10, 16) + if err != nil { + return Config{}, fmt.Errorf("RACER_PEER_PORT: %w", wire.InvalidRequest) + } + + hostNetwork := env("RACER_HOST_NETWORK", "false") + if hostNetwork != "true" && hostNetwork != "false" { + return Config{}, fmt.Errorf("RACER_HOST_NETWORK must be true or false: %w", wire.InvalidRequest) + } + + var diagnosticsPort uint64 + if value, ok := lookup("RACER_DIAGNOSTICS_PORT"); ok { + diagnosticsPort, err = strconv.ParseUint(value, 10, 16) + if err != nil || diagnosticsPort < 1024 { + return Config{}, fmt.Errorf("RACER_DIAGNOSTICS_PORT must be 1024..65535: %w", wire.InvalidRequest) + } + } + + var podNetworkNodes []string + if value, ok := lookup("RACER_POD_NETWORK_NODES"); ok { + if err := json.Unmarshal([]byte(value), &podNetworkNodes); err != nil || podNetworkNodes == nil { + return Config{}, fmt.Errorf("RACER_POD_NETWORK_NODES must be a JSON array: %w", wire.InvalidRequest) + } + } + + cfg := Config{ + Cluster: wire.ClusterID(env("RACER_CLUSTER_ID", "")), + Namespace: env("POD_NAMESPACE", "unbounded-system"), + ControlURL: env("RACER_CONTROL_URL", ""), + BootstrapTrustConfigMap: env("RACER_BOOTSTRAP_TRUST_CONFIGMAP", "racer-bootstrap-trust"), + DataplaneImage: env("RACER_DATAPLANE_IMAGE", ""), + PeerPort: uint16(port), + HostNetwork: hostNetwork == "true", + PodNetworkNodes: podNetworkNodes, + DiagnosticsPort: uint16(diagnosticsPort), + DataplaneServiceAccount: env("RACER_DATAPLANE_SERVICE_ACCOUNT", "racer-dataplane"), + DaemonSetName: env("RACER_DAEMONSET_NAME", "racer-dataplane"), + } + + return cfg, cfg.Validate() +} + +func (c Config) Validate() error { + if err := c.validateNetworkNodes(); err != nil { + return err + } + + if err := c.validateIdentityAndPorts(); err != nil { + return err + } + + return c.validateEndpoint() +} + +func (c Config) validateNetworkNodes() error { + if len(c.PodNetworkNodes) != 0 && (!c.HostNetwork || c.DaemonSetName != DataplaneDaemonSetName) { + return fmt.Errorf("pod network exceptions require host networking and the fixed dataplane name: %w", wire.InvalidRequest) + } + + seen := make(map[string]bool, len(c.PodNetworkNodes)) + for _, node := range c.PodNetworkNodes { + if len(validation.IsDNS1123Subdomain(node)) != 0 || seen[node] { + return fmt.Errorf("pod network nodes must be unique valid node names: %w", wire.InvalidRequest) + } + + seen[node] = true + } + + return nil +} + +func (c Config) validateIdentityAndPorts() error { + if !wire.ValidUUID(string(c.Cluster)) || len(validation.IsDNS1123Label(c.Namespace)) != 0 || c.PeerPort < 1024 { + return fmt.Errorf("cluster, namespace, or peer port: %w", wire.InvalidRequest) + } + + if c.DiagnosticsPort != 0 && (c.DiagnosticsPort < 1024 || c.DiagnosticsPort == c.PeerPort) { + return fmt.Errorf("diagnostics port must be 1024..65535 and distinct from peer port: %w", wire.InvalidRequest) + } + + for _, name := range []string{c.DaemonSetName, c.BootstrapTrustConfigMap, c.DataplaneServiceAccount} { + if len(validation.IsDNS1123Subdomain(name)) != 0 { + return fmt.Errorf("resource name: %w", wire.InvalidRequest) + } + } + + // The workload name is also the immutable instance selector label value. + if len(validation.IsValidLabelValue(c.DaemonSetName)) != 0 { + return fmt.Errorf("DaemonSet name must fit a label value: %w", wire.InvalidRequest) + } + + return nil +} + +func (c Config) validateEndpoint() error { + u, err := url.Parse(c.ControlURL) + if err != nil { + return fmt.Errorf("workload endpoint or image: %w", wire.InvalidRequest) + } + + validOrigin := u.Scheme == "https" && u.Hostname() != "" && u.User == nil + + plainRoot := u.RawQuery == "" && !u.ForceQuery && u.Fragment == "" && u.RawPath == "" && (u.Path == "" || u.Path == "/") + if !validOrigin || !plainRoot || strings.TrimSpace(c.DataplaneImage) == "" { + return fmt.Errorf("workload endpoint or image: %w", wire.InvalidRequest) + } + + if u.Port() != "" { + port, err := strconv.ParseUint(u.Port(), 10, 16) + if err != nil || port == 0 { + return fmt.Errorf("workload endpoint port: %w", wire.InvalidRequest) + } + } + + return nil +} + +// DesiredDaemonSet declares the token audience, controller trust projection, +// node-private identity and slab storage, socket mounts, and exclusion affinity. +// It must never introduce a per-node Secret or trust a node-name as a Node UID. +// This fixture models an externally provisioned workload, never controller startup. +func DesiredDaemonSet(c Config) (*appsv1.DaemonSet, error) { + if err := c.Validate(); err != nil { + return nil, err + } + + // Callers must implement drain-before-admit before opting into two workloads. + if len(c.PodNetworkNodes) != 0 { + return nil, fmt.Errorf("mixed networking requires the drain-aware two-workload planner: %w", wire.InvalidRequest) + } + + return buildDaemonSet(c), nil +} + +func buildDaemonSet(c Config) *appsv1.DaemonSet { + // Use workload identity, not the manager's identity, for the Pod selector. + labels := map[string]string{"app.kubernetes.io/name": "racer-dataplane", "app.kubernetes.io/instance": c.DaemonSetName} + + return &appsv1.DaemonSet{ + TypeMeta: metav1.TypeMeta{APIVersion: "apps/v1", Kind: "DaemonSet"}, + ObjectMeta: metav1.ObjectMeta{Name: c.DaemonSetName, Namespace: c.Namespace, Labels: labels}, + Spec: appsv1.DaemonSetSpec{ + Selector: &metav1.LabelSelector{MatchLabels: labels}, + UpdateStrategy: appsv1.DaemonSetUpdateStrategy{ + Type: appsv1.RollingUpdateDaemonSetStrategyType, + RollingUpdate: &appsv1.RollingUpdateDaemonSet{ + MaxUnavailable: ptr.To(intstr.FromInt32(1)), + MaxSurge: ptr.To(intstr.FromInt32(0)), + }, + }, + // Require sustained readiness across probe periods before advancing a rollout. + MinReadySeconds: 10, + Template: corev1.PodTemplateSpec{ObjectMeta: metav1.ObjectMeta{Labels: labels}, Spec: dataplanePod(c)}, + }, + } +} + +func dataplanePod(c Config) corev1.PodSpec { + pod := corev1.PodSpec{ + ServiceAccountName: c.DataplaneServiceAccount, + AutomountServiceAccountToken: ptr.To(false), + RestartPolicy: corev1.RestartPolicyAlways, + DNSPolicy: corev1.DNSClusterFirst, + SchedulerName: corev1.DefaultSchedulerName, + EnableServiceLinks: ptr.To(false), + PreemptionPolicy: ptr.To(corev1.PreemptLowerPriority), + TerminationGracePeriodSeconds: ptr.To(int64(30)), + SecurityContext: &corev1.PodSecurityContext{RunAsUser: ptr.To(int64(0))}, + Affinity: dataplaneAffinity(), + Containers: []corev1.Container{dataplaneContainer(c)}, + Volumes: dataplaneVolumes(c), + } + if c.HostNetwork { + pod.HostNetwork = true + pod.DNSPolicy = corev1.DNSClusterFirstWithHostNet + } + + return pod +} + +func diagnosticsPort(c Config) int32 { + // Preserve the automatic port for existing installations; explicit ports + // are validated rather than silently moved on collision. + diagnosticsPort := int32(c.DiagnosticsPort) + if diagnosticsPort == 0 { + diagnosticsPort = 9090 + if c.PeerPort == uint16(diagnosticsPort) { + diagnosticsPort++ + } + } + + return diagnosticsPort +} + +func dataplaneAffinity() *corev1.Affinity { + return &corev1.Affinity{ + NodeAffinity: &corev1.NodeAffinity{ + RequiredDuringSchedulingIgnoredDuringExecution: &corev1.NodeSelector{ + NodeSelectorTerms: []corev1.NodeSelectorTerm{{ + MatchExpressions: []corev1.NodeSelectorRequirement{ + { + Key: wire.ExclusionLabel, + Operator: corev1.NodeSelectorOpDoesNotExist, + }, + { + Key: "kubernetes.io/os", + Operator: corev1.NodeSelectorOpIn, + Values: []string{"linux"}, + }, + }, + }}, + }, + }, + } +} + +func dataplaneVolumes(c Config) []corev1.Volume { + volumes := []corev1.Volume{ + { + Name: "devices", + VolumeSource: corev1.VolumeSource{ + HostPath: &corev1.HostPathVolumeSource{Path: "/dev", Type: ptr.To(corev1.HostPathDirectory)}, + }, + }, + { + Name: "token", + VolumeSource: corev1.VolumeSource{ + Projected: &corev1.ProjectedVolumeSource{ + DefaultMode: ptr.To(int32(0o400)), + Sources: []corev1.VolumeProjection{{ + ServiceAccountToken: &corev1.ServiceAccountTokenProjection{ + Audience: wire.TokenAudience, + ExpirationSeconds: ptr.To(int64(3600)), + Path: "token", + }, + }}, + }, + }, + }, + { + Name: "bootstrap", + VolumeSource: corev1.VolumeSource{ + ConfigMap: &corev1.ConfigMapVolumeSource{ + LocalObjectReference: corev1.LocalObjectReference{Name: c.BootstrapTrustConfigMap}, + DefaultMode: ptr.To(int32(0o444)), + Items: []corev1.KeyToPath{{ + Key: "ca.crt", + Path: "ca.crt", + }}, + }, + }, + }, + } + + // DirectoryOrCreate also permits HTTP-only nodes without RDMA hardware. + // Kubelet creates an empty /dev/infiniband; it does not create device nodes. + for _, mount := range []struct{ name, path string }{{"identity", "/var/lib/racer/identity"}, {"slabs", "/var/lib/racer/slabs"}, {"sockets", "/run/racer"}, {"infiniband", "/dev/infiniband"}} { + volumes = append(volumes, corev1.Volume{Name: mount.name, VolumeSource: corev1.VolumeSource{HostPath: &corev1.HostPathVolumeSource{Path: mount.path, Type: ptr.To(corev1.HostPathDirectoryOrCreate)}}}) + } + + return volumes +} + +func dataplaneContainer(c Config) corev1.Container { + diagnosticsPort := diagnosticsPort(c) + + return corev1.Container{ + Name: "dataplane", + Image: c.DataplaneImage, + ImagePullPolicy: corev1.PullIfNotPresent, + TerminationMessagePath: corev1.TerminationMessagePathDefault, + TerminationMessagePolicy: corev1.TerminationMessageReadFile, + Env: dataplaneEnvironment(c, diagnosticsPort), + Ports: []corev1.ContainerPort{ + {Name: "peer", ContainerPort: int32(c.PeerPort), Protocol: corev1.ProtocolTCP}, + {Name: "diagnostics", ContainerPort: diagnosticsPort, Protocol: corev1.ProtocolTCP}, + }, + // Readiness may wait on enrollment/recovery indefinitely without probe + // restarts. Membership must continue to include unready Pods. + ReadinessProbe: &corev1.Probe{ + ProbeHandler: corev1.ProbeHandler{HTTPGet: &corev1.HTTPGetAction{ + Path: "/readyz", Port: intstr.FromString("diagnostics"), Scheme: corev1.URISchemeHTTP, + }}, + PeriodSeconds: 5, TimeoutSeconds: 2, SuccessThreshold: 1, FailureThreshold: 1, + }, + SecurityContext: &corev1.SecurityContext{ + // Native verbs require host device access. Privileged mode + // implies escalation and all capabilities; do not claim otherwise. + Privileged: ptr.To(true), AllowPrivilegeEscalation: ptr.To(true), ReadOnlyRootFilesystem: ptr.To(true), + }, + VolumeMounts: []corev1.VolumeMount{ + {Name: "token", MountPath: "/var/run/racer-token", ReadOnly: true}, + {Name: "bootstrap", MountPath: "/etc/racer/bootstrap", ReadOnly: true}, + {Name: "identity", MountPath: "/var/lib/racer/identity"}, + {Name: "slabs", MountPath: "/var/lib/racer/slabs"}, + {Name: "sockets", MountPath: "/run/racer"}, + // Protect directory entries, not device I/O or the host + // from this privileged container. + {Name: "infiniband", MountPath: "/dev/infiniband", ReadOnly: true}, + {Name: "devices", MountPath: "/host/dev", ReadOnly: true}, + }, + } +} + +func dataplaneEnvironment(c Config, diagnosticsPort int32) []corev1.EnvVar { + return []corev1.EnvVar{ + { + Name: "RACER_CLUSTER_ID", + Value: string(c.Cluster), + }, + { + Name: "RACER_CONTROL_ENDPOINT", + Value: c.ControlURL, + }, + // This is only a bind address, never an authority for node identity. + // Define it first so kubelet expands either Pod IP family below. + { + Name: "RACER_POD_IP", + ValueFrom: &corev1.EnvVarSource{ + FieldRef: &corev1.ObjectFieldSelector{ + APIVersion: "v1", + FieldPath: "status.podIP", + }, + }, + }, + { + Name: "RACER_PEER_LISTEN", + Value: "[$(RACER_POD_IP)]:" + strconv.Itoa(int(c.PeerPort)), + }, + { + Name: "RACER_DIAGNOSTICS_LISTEN", + Value: "[$(RACER_POD_IP)]:" + strconv.Itoa(int(diagnosticsPort)), + }, + { + Name: "RACER_TRUST_BUNDLE", + Value: "/etc/racer/bootstrap/ca.crt", + }, + { + Name: "RACER_SERVICE_ACCOUNT_TOKEN", + Value: "/var/run/racer-token/token", + }, + // Kubelet creates the hostPath mount with mode 0755. Let the + // dataplane create its private 0700 directory beneath it. + { + Name: "RACER_IDENTITY_DIRECTORY", + Value: "/var/lib/racer/identity/private", + }, + { + Name: "RACER_SLAB_DIRECTORY", + Value: "/var/lib/racer/slabs", + }, + { + Name: "RACER_DEVICE_DIRECTORY", + Value: "/host/dev", + }, + } +} + +// DesiredDaemonSets builds steady-state placement, not a safe migration plan. +// The operator must additionally exclude occupied destination nodes until every +// source Pod, including terminating Pods, has disappeared. +func DesiredDaemonSets(c Config) ([]*appsv1.DaemonSet, error) { + if err := c.Validate(); err != nil { + return nil, err + } + + nodes := slices.Clone(c.PodNetworkNodes) + slices.Sort(nodes) + + c.PodNetworkNodes = nil + + host := buildDaemonSet(c) + + if len(nodes) == 0 { + return []*appsv1.DaemonSet{host}, nil + } + + c.HostNetwork = false + c.DaemonSetName = PodNetworkDaemonSetName + + pod := buildDaemonSet(c) + + // The existing selector is immutable. Use a distinct app value, not an + // additional label that would still match the original workload selector. + pod.Labels["app.kubernetes.io/name"] = PodNetworkDaemonSetName + pod.Spec.Selector.MatchLabels["app.kubernetes.io/name"] = PodNetworkDaemonSetName + pod.Spec.Template.Labels["app.kubernetes.io/name"] = PodNetworkDaemonSetName + hostSelector := host.Spec.Template.Spec.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution + podSelector := pod.Spec.Template.Spec.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution + base := podSelector.NodeSelectorTerms[0].DeepCopy() + podSelector.NodeSelectorTerms = nil + // Field selectors accept one value per requirement. Host exclusions are + // ANDed; each pod-network node gets an OR term retaining the base constraints. + for _, node := range nodes { + hostSelector.NodeSelectorTerms[0].MatchFields = append(hostSelector.NodeSelectorTerms[0].MatchFields, corev1.NodeSelectorRequirement{ + Key: "metadata.name", Operator: corev1.NodeSelectorOpNotIn, Values: []string{node}, + }) + term := base.DeepCopy() + term.MatchFields = []corev1.NodeSelectorRequirement{{Key: "metadata.name", Operator: corev1.NodeSelectorOpIn, Values: []string{node}}} + podSelector.NodeSelectorTerms = append(podSelector.NodeSelectorTerms, *term) + } + + return []*appsv1.DaemonSet{host, pod}, nil +} diff --git a/internal/racer/wire/testdata/bootstrap-block-devices.json b/internal/racer/wire/testdata/bootstrap-block-devices.json new file mode 100644 index 000000000..268319f96 --- /dev/null +++ b/internal/racer/wire/testdata/bootstrap-block-devices.json @@ -0,0 +1,9 @@ +[ + {"name":"absent","fields":"","pattern":"","code":""}, + {"name":"empty","fields":",\"block_devices\":\"\"","pattern":"","code":""}, + {"name":"valid","fields":",\"block_devices\":\"^nvme-eui\\\\.[0-9a-f]+$\"","pattern":"^nvme-eui\\.[0-9a-f]+$","code":""}, + {"name":"null","fields":",\"block_devices\":null","pattern":"","code":"invalid_request"}, + {"name":"array","fields":",\"block_devices\":[]","pattern":"","code":"invalid_request"}, + {"name":"duplicate","fields":",\"block_devices\":\"a\",\"block_devices\":\"b\"","pattern":"","code":"invalid_request"}, + {"name":"wrong case","fields":",\"BlockDevices\":\"a\"","pattern":"","code":"invalid_request"} +] diff --git a/internal/racer/wire/testdata/bootstrap-request.json b/internal/racer/wire/testdata/bootstrap-request.json new file mode 100644 index 000000000..24f010fc2 --- /dev/null +++ b/internal/racer/wire/testdata/bootstrap-request.json @@ -0,0 +1 @@ +{"shares":4,"rdma_nics":[],"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","enrollment":"55555555-5555-4555-8555-555555555555","csr_der":"MIGWMEoCAQAwFzEVMBMGA1UEAxMMd2lyZSBmaXh0dXJlMCowBQYDK2VwAyEArvwqpxb5Obl8EI7xTvFeRVFowzucRaHW7prXgx4k1s+gADAFBgMrZXADQQDSTQKGW5jhUBgOR94nvtS/uvCoTtmLzs+rP/gbcET6MtUP/DqTZ4IUzajFPMA89sHWRFECjQNiRhuz21Moa7cO"} diff --git a/internal/racer/wire/testdata/bootstrap-response.json b/internal/racer/wire/testdata/bootstrap-response.json new file mode 100644 index 000000000..c3bd7a7a5 --- /dev/null +++ b/internal/racer/wire/testdata/bootstrap-response.json @@ -0,0 +1 @@ +{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","node":"22222222-2222-4222-8222-222222222222","enrollment":"55555555-5555-4555-8555-555555555555","certificate_chain":["MIIBHjCB0aADAgECAgEBMAUGAytlcDAXMRUwEwYDVQQDEwx3aXJlIGZpeHR1cmUwHhcNNzAwMTAxMDAwMDAwWhcNMzMwNTE4MDMzMzIwWjAXMRUwEwYDVQQDEwx3aXJlIGZpeHR1cmUwKjAFBgMrZXADIQCu/CqnFvk5uXwQjvFO8V5FUWjDO5xFodbumteDHiTWz6NCMEAwDgYDVR0PAQH/BAQDAgIEMA8GA1UdEwEB/wQFMAMBAf8wHQYDVR0OBBYEFIm/J6egWBKSQx6RvXajS4GX5SFOMAUGAytlcANBACtFVQHIakNS+zw6KW5xa3cujIWPET8Kfal1bYBfGhc9x4eEkrwkC13508hPQ24lyg04O19fjp6wPmOK/rZ/UQI="]} diff --git a/internal/racer/wire/testdata/bundle.json b/internal/racer/wire/testdata/bundle.json new file mode 100644 index 000000000..b84f3583b --- /dev/null +++ b/internal/racer/wire/testdata/bundle.json @@ -0,0 +1 @@ +{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","generation":"18446744073709551615","peer_trust_roots":["MIIBHjCB0aADAgECAgEBMAUGAytlcDAXMRUwEwYDVQQDEwx3aXJlIGZpeHR1cmUwHhcNNzAwMTAxMDAwMDAwWhcNMzMwNTE4MDMzMzIwWjAXMRUwEwYDVQQDEwx3aXJlIGZpeHR1cmUwKjAFBgMrZXADIQCu/CqnFvk5uXwQjvFO8V5FUWjDO5xFodbumteDHiTWz6NCMEAwDgYDVR0PAQH/BAQDAgIEMA8GA1UdEwEB/wQFMAMBAf8wHQYDVR0OBBYEFIm/J6egWBKSQx6RvXajS4GX5SFOMAUGAytlcANBACtFVQHIakNS+zw6KW5xa3cujIWPET8Kfal1bYBfGhc9x4eEkrwkC13508hPQ24lyg04O19fjp6wPmOK/rZ/UQI="],"cache_keys":[{"cache":"44444444-4444-4444-8444-444444444444","id":"UktHMQAAAAAAAAABAAAAAA==","purpose":"page","state":"active","material":"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA="},{"cache":"44444444-4444-4444-8444-444444444444","id":"UktHMQAAAAAAAAACAAAAAA==","purpose":"page","state":"prepared","material":"AQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQE="},{"cache":"44444444-4444-4444-8444-444444444444","id":"UktHMQAAAAAAAAABAAAAAA==","purpose":"origin_credentials","state":"active","material":"AgICAgICAgICAgICAgICAgICAgICAgICAgICAgICAgI="}]} diff --git a/internal/racer/wire/testdata/content.json b/internal/racer/wire/testdata/content.json new file mode 100644 index 000000000..a7fd50651 --- /dev/null +++ b/internal/racer/wire/testdata/content.json @@ -0,0 +1 @@ +{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","members":[{"node":"22222222-2222-4222-8222-222222222222","shares":4,"peer_endpoint":"192.0.2.1:7443","rdma_nics":[],"site":""},{"node":"33333333-3333-4333-8333-333333333333","shares":4294967295,"peer_endpoint":"[2001:db8::1]:7443","rdma_nics":[{"device":"fabric-a","port":1,"rail":0},{"device":"β<&>\u2028","port":1,"rail":65535,"numa_node":4294967295}],"site":""}],"caches":[{"id":"44444444-4444-4444-8444-444444444444","name":"cache-a","client_socket":"/run/racer/cache-a/client/socket","origin_socket":"/run/racer/cache-a/origin/socket"}]} diff --git a/internal/racer/wire/testdata/delta.json b/internal/racer/wire/testdata/delta.json new file mode 100644 index 000000000..0eb310ae1 --- /dev/null +++ b/internal/racer/wire/testdata/delta.json @@ -0,0 +1 @@ +{"delta_version":1,"cluster":"11111111-1111-4111-8111-111111111111","base_sequence":"1","base_hash":"3d59ee5856c8f1998f01f947e058caf4382b08b7efbcdeb0b8fb8e234d6ed144","sequence":"2","membership_version":"2","content_hash":"0929db4d1b89c4fdb5c6c20d938a77b286318a830df0b1ac823ee62916af4e20","upsert_members":[{"node":"22222222-2222-4222-8222-222222222222","shares":9,"peer_endpoint":"127.0.0.1:7443","rdma_nics":[],"site":""},{"node":"44444444-4444-4444-8444-444444444444","shares":7,"peer_endpoint":"127.0.0.4:7443","rdma_nics":[],"site":""}],"remove_members":["33333333-3333-4333-8333-333333333333"],"caches":[]} diff --git a/internal/racer/wire/testdata/hashes.json b/internal/racer/wire/testdata/hashes.json new file mode 100644 index 000000000..ebe0c21b0 --- /dev/null +++ b/internal/racer/wire/testdata/hashes.json @@ -0,0 +1 @@ +{"content":"1f68b286bce90f488529367057850804d24be4dc5e81c811cc06ebfa376c0a0e","membership":"c523f9a1b8f503628a8dbcef9fa0ef35553640b7c226515f5f5d70a9f3abe0ac"} diff --git a/internal/racer/wire/testdata/membership.json b/internal/racer/wire/testdata/membership.json new file mode 100644 index 000000000..e13775999 --- /dev/null +++ b/internal/racer/wire/testdata/membership.json @@ -0,0 +1 @@ +{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","members":[{"node":"22222222-2222-4222-8222-222222222222","shares":4,"peer_endpoint":"192.0.2.1:7443","rdma_nics":[],"site":""},{"node":"33333333-3333-4333-8333-333333333333","shares":4294967295,"peer_endpoint":"[2001:db8::1]:7443","rdma_nics":[{"device":"fabric-a","port":1,"rail":0},{"device":"β<&>\u2028","port":1,"rail":65535,"numa_node":4294967295}],"site":""}]} diff --git a/internal/racer/wire/testdata/publication.json b/internal/racer/wire/testdata/publication.json new file mode 100644 index 000000000..957431f2a --- /dev/null +++ b/internal/racer/wire/testdata/publication.json @@ -0,0 +1 @@ +{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","sequence":"18446744073709551615","membership_version":"9007199254740993","members":[{"node":"33333333-3333-4333-8333-333333333333","shares":4294967295,"peer_endpoint":"[2001:db8::1]:7443","rdma_nics":[{"device":"β\u003c\u0026\u003e\u2028","port":1,"rail":65535,"numa_node":4294967295},{"device":"fabric-a","port":1,"rail":0}],"site":""},{"node":"22222222-2222-4222-8222-222222222222","shares":4,"peer_endpoint":"192.0.2.1:7443","rdma_nics":[],"site":""}],"caches":[{"id":"44444444-4444-4444-8444-444444444444","name":"cache-a","client_socket":"/run/racer/cache-a/client/socket","origin_socket":"/run/racer/cache-a/origin/socket"}]} diff --git a/internal/racer/wire/testdata/rejections.json b/internal/racer/wire/testdata/rejections.json new file mode 100644 index 000000000..b7fe354d6 --- /dev/null +++ b/internal/racer/wire/testdata/rejections.json @@ -0,0 +1,401 @@ +[ + { + "name": "legacy key id", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAABAAAAAA==", + "new": "AAAAAAAAAAAAAAAAAAAAAA==", + "code": "invalid_request" + }, + { + "name": "zero key generation", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAABAAAAAA==", + "new": "UktHMQAAAAAAAAAAAAAAAA==", + "code": "invalid_request" + }, + { + "name": "future key generation", + "file": "bundle.json", + "old": "\"generation\":\"18446744073709551615\"", + "new": "\"generation\":\"1\"", + "code": "invalid_request" + }, + { + "name": "unpaired surrogate", + "file": "publication.json", + "old": "fabric-a", + "new": "\\ud800", + "code": "invalid_request" + }, + { + "name": "unknown unpaired surrogate", + "file": "publication.json", + "old": "\"schema_version\":1", + "new": "\"schema_version\":1,\"\\udc00\":0", + "code": "invalid_request" + }, + { + "name": "zoned IPv6", + "file": "publication.json", + "old": "[2001:db8::1]:7443", + "new": "[fe80::1%eth0]:7443", + "code": "invalid_request" + }, + { + "name": "negative zero number", + "file": "publication.json", + "old": "\"rail\":0", + "new": "\"rail\":-0", + "code": "invalid_request" + }, + { + "name": "duplicate top field", + "file": "publication.json", + "old": "\"schema_version\":1", + "new": "\"schema_version\":1,\"schema_version\":1", + "code": "invalid_request" + }, + { + "name": "escaped duplicate", + "file": "publication.json", + "old": "\"schema_version\":1", + "new": "\"schema_version\":1,\"schema_versi\\u006fn\":1", + "code": "invalid_request" + }, + { + "name": "unknown duplicate", + "file": "publication.json", + "old": "\"schema_version\":1", + "new": "\"schema_version\":1,\"future\":{\"x\":1,\"x\":2}", + "code": "invalid_request" + }, + { + "name": "unknown top field", + "file": "publication.json", + "old": "\"schema_version\":1", + "new": "\"schema_version\":1,\"future\":0", + "code": "invalid_request" + }, + { + "name": "unknown member field", + "file": "publication.json", + "old": "\"shares\":4,", + "new": "\"shares\":4,\"future\":0,", + "code": "invalid_request" + }, + { + "name": "unknown rail field", + "file": "publication.json", + "old": "\"rail\":0", + "new": "\"rail\":0,\"future\":0", + "code": "invalid_request" + }, + { + "name": "missing site", + "file": "publication.json", + "old": ",\"site\":\"\"", + "new": "", + "code": "invalid_request" + }, + { + "name": "missing shares", + "file": "bootstrap-request.json", + "old": "\"shares\":4,", + "new": "", + "code": "invalid_request" + }, + { + "name": "zero bootstrap shares", + "file": "bootstrap-request.json", + "old": "\"shares\":4,", + "new": "\"shares\":0,", + "code": "invalid_request" + }, + { + "name": "unknown version", + "file": "publication.json", + "old": "\"schema_version\":1", + "new": "\"schema_version\":2", + "code": "unsupported_version" + }, + { + "name": "numeric counter", + "file": "publication.json", + "old": "\"sequence\":\"18446744073709551615\"", + "new": "\"sequence\":1", + "code": "invalid_request" + }, + { + "name": "overflow counter", + "file": "publication.json", + "old": "18446744073709551615", + "new": "18446744073709551616", + "code": "invalid_request" + }, + { + "name": "leading zero counter", + "file": "publication.json", + "old": "18446744073709551615", + "new": "01", + "code": "invalid_request" + }, + { + "name": "zero counter", + "file": "publication.json", + "old": "18446744073709551615", + "new": "0", + "code": "invalid_request" + }, + { + "name": "signed counter", + "file": "publication.json", + "old": "18446744073709551615", + "new": "+1", + "code": "invalid_request" + }, + { + "name": "null members", + "file": "publication.json", + "old": "\"members\":[", + "new": "\"members\":null,\"ignored\":[", + "code": "invalid_request" + }, + { + "name": "missing required field", + "file": "publication.json", + "old": "\"rdma_nics\":[]", + "new": "\"ignored_nics\":[]", + "code": "invalid_request" + }, + { + "name": "null optional field", + "file": "publication.json", + "old": "\"numa_node\":4294967295", + "new": "\"numa_node\":null", + "code": "invalid_request" + }, + { + "name": "uppercase UUID", + "file": "publication.json", + "old": "11111111-1111-4111-8111-111111111111", + "new": "AAAAAAAA-1111-4111-8111-111111111111", + "code": "invalid_request" + }, + { + "name": "duplicate node", + "file": "publication.json", + "old": "22222222-2222-4222-8222-222222222222", + "new": "33333333-3333-4333-8333-333333333333", + "code": "invalid_request" + }, + { + "name": "zero shares", + "file": "publication.json", + "old": "\"shares\":4,", + "new": "\"shares\":0,", + "code": "invalid_request" + }, + { + "name": "float shares", + "file": "publication.json", + "old": "\"shares\":4,", + "new": "\"shares\":4.0,", + "code": "invalid_request" + }, + { + "name": "overflow shares", + "file": "publication.json", + "old": "4294967295", + "new": "4294967296", + "code": "invalid_request" + }, + { + "name": "zero NIC port", + "file": "publication.json", + "old": "\"port\":1", + "new": "\"port\":0", + "code": "invalid_request" + }, + { + "name": "overflow rail", + "file": "publication.json", + "old": "\"rail\":65535", + "new": "\"rail\":65536", + "code": "invalid_request" + }, + { + "name": "empty device", + "file": "publication.json", + "old": "fabric-a", + "new": "", + "code": "invalid_request" + }, + { + "name": "hostname endpoint", + "file": "publication.json", + "old": "192.0.2.1:7443", + "new": "example.com:7443", + "code": "invalid_request" + }, + { + "name": "zero port", + "file": "publication.json", + "old": "192.0.2.1:7443", + "new": "192.0.2.1:0", + "code": "invalid_request" + }, + { + "name": "unbracketed IPv6", + "file": "publication.json", + "old": "[2001:db8::1]:7443", + "new": "2001:db8::1:7443", + "code": "invalid_request" + }, + { + "name": "path traversal", + "file": "publication.json", + "old": "cache-a", + "new": "..", + "code": "invalid_request" + }, + { + "name": "wrong path", + "file": "publication.json", + "old": "/run/racer/cache-a/client/socket", + "new": "/tmp/client.sock", + "code": "invalid_request" + }, + { + "name": "truncated DER CSR", + "file": "bootstrap-request.json", + "old": "MIGWMEoCAQAwFzEVMBMGA1UEAxMMd2lyZSBmaXh0dXJlMCowBQYDK2VwAyEArvwqpxb5Obl8EI7xTvFeRVFowzucRaHW7prXgx4k1s+gADAFBgMrZXADQQDSTQKGW5jhUBgOR94nvtS/uvCoTtmLzs+rP/gbcET6MtUP/DqTZ4IUzajFPMA89sHWRFECjQNiRhuz21Moa7cO", + "new": "MAA=", + "code": "invalid_request" + }, + { + "name": "unknown purpose", + "file": "bundle.json", + "old": "\"purpose\":\"page\"", + "new": "\"purpose\":\"future\"", + "code": "invalid_request" + }, + { + "name": "unknown state", + "file": "bundle.json", + "old": "\"state\":\"prepared\"", + "new": "\"state\":\"future\"", + "code": "invalid_request" + }, + { + "name": "two active keys", + "file": "bundle.json", + "old": "\"state\":\"prepared\"", + "new": "\"state\":\"active\"", + "code": "invalid_request" + }, + { + "name": "no active key", + "file": "bundle.json", + "old": "\"state\":\"active\"", + "new": "\"state\":\"retiring\"", + "code": "invalid_request" + }, + { + "name": "duplicate key identity", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAACAAAAAA==", + "new": "UktHMQAAAAAAAAABAAAAAA==", + "code": "invalid_request" + }, + { + "name": "short key id", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAABAAAAAA==", + "new": "AA==", + "code": "invalid_request" + }, + { + "name": "short material", + "file": "bundle.json", + "old": "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=", + "new": "AA==", + "code": "invalid_request" + }, + { + "name": "unpadded base64", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAABAAAAAA==", + "new": "AAAAAAAAAAAAAAAAAAAAAA", + "code": "invalid_request" + }, + { + "name": "nonzero padding bits", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAABAAAAAA==", + "new": "AAAAAAAAAAAAAAAAAAAAAB==", + "code": "invalid_request" + }, + { + "name": "base64 newline", + "file": "bundle.json", + "old": "UktHMQAAAAAAAAABAAAAAA==", + "new": "AAAAAAAAAA\\nAAAAAAAAAAAA==", + "code": "invalid_request" + }, + { + "name": "missing bootstrap NIC report", + "file": "bootstrap-request.json", + "old": "\"rdma_nics\":[],", + "new": "", + "code": "invalid_request" + }, + { + "name": "null bootstrap NIC report", + "file": "bootstrap-request.json", + "old": "\"rdma_nics\":[]", + "new": "\"rdma_nics\":null", + "code": "invalid_request" + }, + { + "name": "legacy bootstrap rails", + "file": "bootstrap-request.json", + "old": "\"rdma_nics\":[]", + "new": "\"rails\":[]", + "code": "invalid_request" + }, + { + "name": "legacy member rails", + "file": "publication.json", + "old": "\"rdma_nics\":[]", + "new": "\"rails\":[],\"alignment_enabled\":true", + "code": "invalid_request" + }, + { + "name": "overflow NIC port", + "file": "publication.json", + "old": "\"port\":1", + "new": "\"port\":256", + "code": "invalid_request" + }, + { + "name": "null NIC GID", + "file": "publication.json", + "old": "\"port\":1", + "new": "\"port\":1,\"gid\":null", + "code": "invalid_request" + }, + { + "name": "empty NIC GID", + "file": "publication.json", + "old": "\"port\":1", + "new": "\"port\":1,\"gid\":\"\"", + "code": "invalid_request" + }, + { + "name": "uppercase NIC GID", + "file": "publication.json", + "old": "\"port\":1", + "new": "\"port\":1,\"gid\":\"ABCDEF0123456789abcdef0123456789\"", + "code": "invalid_request" + } +] diff --git a/internal/racer/wire/testdata/site-vectors.json b/internal/racer/wire/testdata/site-vectors.json new file mode 100644 index 000000000..04e6cebc3 --- /dev/null +++ b/internal/racer/wire/testdata/site-vectors.json @@ -0,0 +1,41 @@ +[ + { + "name": "absent", + "site": "", + "publication": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"sequence\":\"1\",\"membership_version\":\"1\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "content": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "membership": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}]}", + "content_hash": "1f68b286bce90f488529367057850804d24be4dc5e81c811cc06ebfa376c0a0e", + "membership_hash": "c523f9a1b8f503628a8dbcef9fa0ef35553640b7c226515f5f5d70a9f3abe0ac" + }, + { + "name": "added", + "site": "Site_1.west-2", + "publication": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"sequence\":\"2\",\"membership_version\":\"2\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_1.west-2\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "content": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_1.west-2\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "membership": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_1.west-2\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}]}", + "content_hash": "7862289d5981ec3f27d55f751673ad230fc91c83d8fefbd2118985f20658f3c2", + "membership_hash": "33c3fd9bd9ee0e48785a4eb9a6bfdce21c755ad2b047a2c0e1e94a66f507460c", + "delta": "{\"delta_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"base_sequence\":\"1\",\"base_hash\":\"1f68b286bce90f488529367057850804d24be4dc5e81c811cc06ebfa376c0a0e\",\"sequence\":\"2\",\"membership_version\":\"2\",\"content_hash\":\"7862289d5981ec3f27d55f751673ad230fc91c83d8fefbd2118985f20658f3c2\",\"upsert_members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_1.west-2\"}],\"remove_members\":[],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}" + }, + { + "name": "changed", + "site": "Site_2.east-1", + "publication": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"sequence\":\"3\",\"membership_version\":\"3\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_2.east-1\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "content": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_2.east-1\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "membership": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_2.east-1\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}]}", + "content_hash": "30c83458f60a81a4b1cd357e2663d732e5e68c1be64f658a9f052d3025872c93", + "membership_hash": "39d71ee75ae9dab3d40a43cbcc0f86812c5c76089290ced0ddc84e854a56f2cf", + "delta": "{\"delta_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"base_sequence\":\"2\",\"base_hash\":\"7862289d5981ec3f27d55f751673ad230fc91c83d8fefbd2118985f20658f3c2\",\"sequence\":\"3\",\"membership_version\":\"3\",\"content_hash\":\"30c83458f60a81a4b1cd357e2663d732e5e68c1be64f658a9f052d3025872c93\",\"upsert_members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"Site_2.east-1\"}],\"remove_members\":[],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}" + }, + { + "name": "removed", + "site": "", + "publication": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"sequence\":\"4\",\"membership_version\":\"4\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "content": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}", + "membership": "{\"schema_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"},{\"node\":\"33333333-3333-4333-8333-333333333333\",\"shares\":4294967295,\"peer_endpoint\":\"[2001:db8::1]:7443\",\"rdma_nics\":[{\"device\":\"fabric-a\",\"port\":1,\"rail\":0},{\"device\":\"β\u003c\u0026\u003e\\u2028\",\"port\":1,\"rail\":65535,\"numa_node\":4294967295}],\"site\":\"\"}]}", + "content_hash": "1f68b286bce90f488529367057850804d24be4dc5e81c811cc06ebfa376c0a0e", + "membership_hash": "c523f9a1b8f503628a8dbcef9fa0ef35553640b7c226515f5f5d70a9f3abe0ac", + "delta": "{\"delta_version\":1,\"cluster\":\"11111111-1111-4111-8111-111111111111\",\"base_sequence\":\"3\",\"base_hash\":\"30c83458f60a81a4b1cd357e2663d732e5e68c1be64f658a9f052d3025872c93\",\"sequence\":\"4\",\"membership_version\":\"4\",\"content_hash\":\"1f68b286bce90f488529367057850804d24be4dc5e81c811cc06ebfa376c0a0e\",\"upsert_members\":[{\"node\":\"22222222-2222-4222-8222-222222222222\",\"shares\":4,\"peer_endpoint\":\"192.0.2.1:7443\",\"rdma_nics\":[],\"site\":\"\"}],\"remove_members\":[],\"caches\":[{\"id\":\"44444444-4444-4444-8444-444444444444\",\"name\":\"cache-a\",\"client_socket\":\"/run/racer/cache-a/client/socket\",\"origin_socket\":\"/run/racer/cache-a/origin/socket\"}]}" + } +] diff --git a/internal/racer/wire/wire.go b/internal/racer/wire/wire.go new file mode 100644 index 000000000..3391b9889 --- /dev/null +++ b/internal/racer/wire/wire.go @@ -0,0 +1,1336 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +// Package wire declares Racer's HTTPS/JSON contracts. +// Codecs validate bounded inputs before returning usable protocol state. +package wire + +import ( + "bytes" + "cmp" + "crypto/sha256" + "crypto/x509" + "encoding/base64" + "encoding/binary" + "encoding/hex" + "encoding/json" + "io" + "net/netip" + "reflect" + "slices" + "strconv" + "strings" + "time" + "unicode/utf8" + + "k8s.io/apimachinery/pkg/util/validation" +) + +const ( + SchemaVersion = 1 + BootstrapPath = "/v1/bootstrap" + SnapshotPath = "/v1/snapshot" + KeyringPath = "/v1/keyring" + TokenAudience = "racer-control" + MaxBootstrapBytes = 64 * 1024 + MaxBundleBytes = 512 * 1024 + MaxPublicationBytes = 64 * 1024 * 1024 + MaxMembers = 100_000 + PollWait = 30 * time.Second + CertificateLifetime = 24 * time.Hour + DefaultShares = 4 + SharesAnnotation = "racer.unbounded-cloud.io/shares" + BlockDevicesAnnotation = "racer.unbounded-cloud.io/block-devices" + RDMANICsAnnotation = "racer.unbounded-cloud.io/rdma-nics" + MaxRDMANICs = 64 + ExclusionLabel = "racer.unbounded-cloud.io/exclude" +) + +type ( + ClusterID string + NodeID string + CacheID string + EnrollmentID string + Sequence uint64 + MembershipVersion uint64 + Generation uint64 +) + +type RDMANIC struct { + Device string `json:"device"` + Port uint8 `json:"port"` + Rail uint16 `json:"rail"` + GID string `json:"gid,omitempty"` + NUMANode *uint32 `json:"numa_node,omitempty"` +} + +type Member struct { + Node NodeID `json:"node"` + Shares uint32 `json:"shares"` + PeerEndpoint string `json:"peer_endpoint"` + RDMANICs []RDMANIC `json:"rdma_nics"` + // Site is the RDMA boundary. Empty means HTTP-only, not a shared default site. + Site string `json:"site"` +} + +type CacheDefinition struct { + ID CacheID `json:"id"` + Name string `json:"name"` + ClientSocket string `json:"client_socket"` + OriginSocket string `json:"origin_socket"` +} + +type Publication struct { + SchemaVersion uint32 `json:"schema_version"` + Cluster ClusterID `json:"cluster"` + Sequence Sequence `json:"sequence,string"` + MembershipVersion MembershipVersion `json:"membership_version,string"` + Members []Member `json:"members"` + Caches []CacheDefinition `json:"caches"` +} + +// BootstrapRequest contains no bearer token. The transport reads the projected +// token for each issuance attempt and supplies it in Authorization. +type BootstrapRequest struct { + Shares uint32 `json:"shares"` + RDMANICs []RDMANIC `json:"rdma_nics"` + SchemaVersion uint32 `json:"schema_version"` + Cluster ClusterID `json:"cluster"` + Enrollment EnrollmentID `json:"enrollment"` + CSRDER []byte `json:"csr_der"` +} + +type BootstrapResponse struct { + SchemaVersion uint32 `json:"schema_version"` + Cluster ClusterID `json:"cluster"` + Node NodeID `json:"node"` + Enrollment EnrollmentID `json:"enrollment"` + CertificateChain [][]byte `json:"certificate_chain"` + BlockDevices string `json:"block_devices,omitempty"` +} + +type ( + KeyPurpose string + KeyState string +) + +const ( + PageKey KeyPurpose = "page" + OriginCredentialsKey KeyPurpose = "origin_credentials" + PreparedKey KeyState = "prepared" + ActiveKey KeyState = "active" +) + +type CacheKeyRef struct { + Cache CacheID `json:"cache"` + ID []byte `json:"id"` // Exactly 16 bytes, padded standard base64 on the wire. + Purpose KeyPurpose `json:"purpose"` +} + +// CacheKey intentionally hides material from ordinary formatting and JSON. +// Only the bounded bundle codec may access and encode the 32-byte material. +type CacheKey struct { + Key CacheKeyRef + State KeyState + material [32]byte +} + +func (CacheKey) String() string { return "" } +func (CacheKey) GoString() string { return "" } + +// NewCacheKey is the only material ingress besides bounded bundle decoding. +func NewCacheKey(ref CacheKeyRef, state KeyState, material [32]byte) (CacheKey, error) { + if !ValidUUID(string(ref.Cache)) || len(ref.ID) != 16 || string(ref.ID[:4]) != "RKG1" || binary.BigEndian.Uint64(ref.ID[4:12]) == 0 || (ref.Purpose != PageKey && ref.Purpose != OriginCredentialsKey) || (state != PreparedKey && state != ActiveKey) { + return CacheKey{}, InvalidRequest + } + + ref.ID = bytes.Clone(ref.ID) + + return CacheKey{Key: ref, State: state, material: material}, nil +} + +// EqualMaterial compares private key bytes without exposing them to consumers. +// It is intended for bundle replay validation, not authentication. +func (k CacheKey) EqualMaterial(other CacheKey) bool { + return bytes.Equal(k.material[:], other.material[:]) +} + +type KeyringBundle struct { + SchemaVersion uint32 + Cluster ClusterID + Generation Generation + PeerTrustRoots [][]byte + CacheKeys []CacheKey +} + +// MarshalJSON prevents bypassing validation through the standard JSON encoder. +func (b KeyringBundle) MarshalJSON() ([]byte, error) { return EncodeBundle(b) } + +type ErrorCode string + +const ( + InvalidRequest ErrorCode = "invalid_request" + Unauthenticated ErrorCode = "unauthenticated" + Forbidden ErrorCode = "forbidden" + Conflict ErrorCode = "conflict" + TooLarge ErrorCode = "too_large" + UnsupportedVersion ErrorCode = "unsupported_version" + Overloaded ErrorCode = "overloaded" + Unavailable ErrorCode = "unavailable" +) + +type ErrorResponse struct { + Code ErrorCode `json:"code"` +} + +// Error returns only a protocol code, never input or secret material. +func (c ErrorCode) Error() string { return string(c) } + +func validError(c ErrorCode) bool { + switch c { + case InvalidRequest, Unauthenticated, Forbidden, Conflict, TooLarge, UnsupportedVersion, Overloaded, Unavailable: + return true + default: + return false + } +} + +// ValidUUID checks canonical Kubernetes identities before constructing wire state. +func ValidUUID(s string) bool { + if len(s) != 36 { + return false + } + + for i, c := range []byte(s) { + if i == 8 || i == 13 || i == 18 || i == 23 { + if c != '-' { + return false + } + + continue + } + + if (c < '0' || c > '9') && (c < 'a' || c > 'f') { + return false + } + } + + return true +} + +func validRDMANIC(r RDMANIC) bool { + if r.Device == "" || !utf8.ValidString(r.Device) || strings.ContainsAny(r.Device, "\x00\r\n") || r.Port == 0 { + return false + } + + if r.GID != "" { + if len(r.GID) != 32 || strings.ToLower(r.GID) != r.GID { + return false + } + + if _, err := hex.DecodeString(r.GID); err != nil { + return false + } + } + + return true +} + +func validateRDMANICs(nics []RDMANIC) error { + if len(nics) > MaxRDMANICs { + return TooLarge + } + + type physical struct { + device string + port uint8 + } + + seen := map[physical]bool{} + + for _, nic := range nics { + key := physical{nic.Device, nic.Port} + if !validRDMANIC(nic) || seen[key] { + return InvalidRequest + } + + seen[key] = true + } + + return nil +} + +func validateHeader(version uint32, cluster ClusterID) error { + if version != SchemaVersion { + return UnsupportedVersion + } + + if !ValidUUID(string(cluster)) { + return InvalidRequest + } + + return nil +} + +// ValidateBootstrapRequest checks the wire fields, CSR syntax, and full encoded +// size without serializing. Proof of possession and identity binding belong to +// the issuer. Direct callers have the same size bound as EncodeBootstrapRequest. +func ValidateBootstrapRequest(v BootstrapRequest) error { + if err := validateHeader(v.SchemaVersion, v.Cluster); err != nil { + return err + } + + if !ValidUUID(string(v.Enrollment)) || v.Shares == 0 { + return InvalidRequest + } + + if len(v.CSRDER) > MaxBootstrapBytes { + return TooLarge + } + + if err := validateRDMANICs(v.RDMANICs); err != nil { + return err + } + + remaining := MaxBootstrapBytes + for _, nic := range v.RDMANICs { + if len(nic.Device) > remaining { + return TooLarge + } + + remaining -= len(nic.Device) + } + + if _, err := x509.ParseCertificateRequest(v.CSRDER); err != nil { + return InvalidRequest + } + + // The validated version is 1 and UUIDs are unescaped ASCII. Only padded + // base64 contributes variable framing size; no encoder newline is on the wire. + nics, err := encode(CanonicalRDMANICs(v.RDMANICs), MaxBootstrapBytes) + if err != nil { + return err + } + + const framing = len(`{"shares":,"rdma_nics":,"schema_version":1,"cluster":"","enrollment":"","csr_der":""}`) + if framing+len(nics)+len(strconv.FormatUint(uint64(v.Shares), 10))+len(v.Cluster)+len(v.Enrollment)+base64.StdEncoding.EncodedLen(len(v.CSRDER)) > MaxBootstrapBytes { + return TooLarge + } + + return nil +} + +func validateBootstrapResponse(v BootstrapResponse) error { + if err := validateHeader(v.SchemaVersion, v.Cluster); err != nil { + return err + } + + if !ValidUUID(string(v.Node)) || !ValidUUID(string(v.Enrollment)) { + return InvalidRequest + } + + return validateCertificates(v.CertificateChain, MaxBootstrapBytes) +} + +func validateCertificates(certs [][]byte, limit int) error { + if len(certs) == 0 { + return InvalidRequest + } + + total := 0 + for _, cert := range certs { + if len(cert) > limit-total { + return TooLarge + } + + total += len(cert) + if _, err := x509.ParseCertificate(cert); err != nil { + return InvalidRequest + } + } + + return nil +} + +// CanonicalSocketPaths returns Linux pathname sockets including room for NUL in +// sockaddr_un.sun_path (108 bytes). Names are ASCII Kubernetes DNS subdomains. +func CanonicalSocketPaths(name string) (client, origin string, err error) { + if name == "" || len(name) > 253 { + return "", "", InvalidRequest + } + + for _, label := range strings.Split(name, ".") { + if !validDNSLabel(label) { + return "", "", InvalidRequest + } + } + + client = "/run/racer/" + name + "/client/socket" + + origin = "/run/racer/" + name + "/origin/socket" + if len(client) > 107 || len(origin) > 107 { + return "", "", InvalidRequest + } + + return client, origin, nil +} + +func validDNSLabel(label string) bool { + if len(label) == 0 || len(label) > 63 { + return false + } + + for i, c := range []byte(label) { + if c >= 'a' && c <= 'z' || c >= '0' && c <= '9' { + continue + } + + if c != '-' || i == 0 || i == len(label)-1 { + return false + } + } + + return true +} + +func validatePublication(v Publication, counters bool) error { + if err := validateHeader(v.SchemaVersion, v.Cluster); err != nil { + return err + } + + if counters && (v.Sequence == 0 || v.MembershipVersion == 0) { + return InvalidRequest + } + + if len(v.Members) > MaxMembers { + return TooLarge + } + + if err := checkPublicationSize(v); err != nil { + return err + } + + if err := validateMembers(v.Members); err != nil { + return err + } + + return validateCaches(v.Caches) +} + +func checkPublicationSize(v Publication) error { + // A cheap lower bound prevents copying/encoding caller-owned oversized state. + remaining := MaxPublicationBytes + + consume := func(n int) bool { + if n > remaining { + return false + } + + remaining -= n + + return true + } + for _, m := range v.Members { + if !consume(len(m.Node)) || !consume(len(m.PeerEndpoint)) || !consume(len(m.Site)) { + return TooLarge + } + + for _, r := range m.RDMANICs { + if !consume(len(r.Device) + len(r.GID) + 1) { + return TooLarge + } + } + } + + for _, c := range v.Caches { + if !consume(len(c.ID)) || !consume(len(c.Name)) || !consume(len(c.ClientSocket)) || !consume(len(c.OriginSocket)) { + return TooLarge + } + } + + return nil +} + +func validateMembers(members []Member) error { + nodes := map[NodeID]bool{} + for _, m := range members { + if !ValidUUID(string(m.Node)) || nodes[m.Node] || m.Shares == 0 || len(validation.IsValidLabelValue(m.Site)) != 0 { + return InvalidRequest + } + + nodes[m.Node] = true + + ap, err := netip.ParseAddrPort(m.PeerEndpoint) + if err != nil || ap.Port() == 0 || ap.Addr().Zone() != "" { + return InvalidRequest + } + + if err := validateRDMANICs(m.RDMANICs); err != nil { + return err + } + } + + return nil +} + +func validateCaches(caches []CacheDefinition) error { + ids := map[CacheID]bool{} + + names := map[string]bool{} + for _, c := range caches { + if !ValidUUID(string(c.ID)) || ids[c.ID] || names[c.Name] { + return InvalidRequest + } + + ids[c.ID], names[c.Name] = true, true + + client, origin, err := CanonicalSocketPaths(c.Name) + if err != nil || c.ClientSocket != client || c.OriginSocket != origin { + return InvalidRequest + } + } + + return nil +} + +func validateBundle(v KeyringBundle) error { + if err := validateHeader(v.SchemaVersion, v.Cluster); err != nil { + return err + } + + if v.Generation == 0 { + return InvalidRequest + } + + if len(v.CacheKeys) > MaxBundleBytes/32 { + return TooLarge + } + + if err := validateCertificates(v.PeerTrustRoots, MaxBundleBytes); err != nil { + return err + } + + roots := map[string]bool{} + for _, r := range v.PeerTrustRoots { + if roots[string(r)] { + return InvalidRequest + } + + roots[string(r)] = true + } + + return validateCacheKeys(v.CacheKeys, v.Generation) +} + +func validateCacheKeys(keys []CacheKey, generation Generation) error { + type scope struct { + cache CacheID + purpose KeyPurpose + } + + type identity struct { + scope + id string + } + + seen := map[identity]bool{} + materials := map[[32]byte]bool{} + active := map[scope]int{} + + for _, k := range keys { + if _, err := NewCacheKey(k.Key, k.State, k.material); err != nil { + return err + } + + if binary.BigEndian.Uint64(k.Key.ID[4:12]) > uint64(generation) { + return InvalidRequest + } + + s := scope{k.Key.Cache, k.Key.Purpose} + + id := identity{s, string(k.Key.ID)} + if seen[id] || materials[k.material] { + return InvalidRequest + } + + seen[id] = true + materials[k.material] = true + + if _, ok := active[s]; !ok { + active[s] = 0 + } + + if k.State == ActiveKey { + active[s]++ + } + } + + for _, count := range active { + if count != 1 { + return InvalidRequest + } + } + + return nil +} + +// Bounded protocol codecs. + +// DecodeBootstrap bounds the entire document before allocating decoded state. +func DecodeBootstrap(r io.Reader) (BootstrapRequest, error) { + var v BootstrapRequest + if err := decode(r, MaxBootstrapBytes, &v); err != nil { + return BootstrapRequest{}, err + } + + if err := ValidateBootstrapRequest(v); err != nil { + return BootstrapRequest{}, err + } + + v.RDMANICs = CanonicalRDMANICs(v.RDMANICs) + + return v, nil +} + +func EncodeBootstrap(v BootstrapResponse) ([]byte, error) { + if err := validateBootstrapResponse(v); err != nil { + return nil, err + } + + return encode(v, MaxBootstrapBytes) +} + +func EncodeBootstrapRequest(v BootstrapRequest) ([]byte, error) { + if err := ValidateBootstrapRequest(v); err != nil { + return nil, err + } + + v.RDMANICs = CanonicalRDMANICs(v.RDMANICs) + + return encode(v, MaxBootstrapBytes) +} + +func DecodeBootstrapResponse(r io.Reader) (BootstrapResponse, error) { + var v BootstrapResponse + if err := decode(r, MaxBootstrapBytes, &v); err != nil { + return BootstrapResponse{}, err + } + + if err := validateBootstrapResponse(v); err != nil { + return BootstrapResponse{}, err + } + + return v, nil +} + +// DecodeRDMANICs strictly decodes a bounded NIC annotation and canonicalizes it. +func DecodeRDMANICs(r io.Reader) ([]RDMANIC, error) { + var nics []RDMANIC + if err := decode(r, 256*1024, &nics); err != nil { + return nil, err + } + + if err := validateRDMANICs(nics); err != nil { + return nil, err + } + + return CanonicalRDMANICs(nics), nil +} + +// DecodeAdmittedMember validates current restart hints with the strict wire +// shape, including required arrays and rejection of duplicate fields. +func DecodeAdmittedMember(r io.Reader) (Member, error) { + var member Member + if err := decode(r, MaxBootstrapBytes, &member); err != nil { + return Member{}, err + } + + probe := Publication{SchemaVersion: SchemaVersion, Cluster: ClusterID(member.Node), Sequence: 1, MembershipVersion: 1, Members: []Member{member}} + if err := validatePublication(probe, true); err != nil { + return Member{}, err + } + + member.RDMANICs = CanonicalRDMANICs(member.RDMANICs) + + return member, nil +} + +type bundleJSON struct { + SchemaVersion uint32 `json:"schema_version"` + Cluster ClusterID `json:"cluster"` + Generation Generation `json:"generation,string"` + PeerTrustRoots [][]byte `json:"peer_trust_roots"` + CacheKeys []keyJSON `json:"cache_keys"` +} + +type keyJSON struct { + Cache CacheID `json:"cache"` + ID []byte `json:"id"` + Purpose KeyPurpose `json:"purpose"` + State KeyState `json:"state"` + Material []byte `json:"material"` +} + +func DecodeBundle(r io.Reader) (KeyringBundle, error) { + var raw bundleJSON + if err := decode(r, MaxBundleBytes, &raw); err != nil { + return KeyringBundle{}, err + } + + v := KeyringBundle{SchemaVersion: raw.SchemaVersion, Cluster: raw.Cluster, Generation: raw.Generation, PeerTrustRoots: raw.PeerTrustRoots, CacheKeys: make([]CacheKey, 0, len(raw.CacheKeys))} + for _, k := range raw.CacheKeys { + if len(k.Material) != 32 { + return KeyringBundle{}, InvalidRequest + } + + key, err := NewCacheKey(CacheKeyRef{Cache: k.Cache, ID: k.ID, Purpose: k.Purpose}, k.State, [32]byte(k.Material)) + if err != nil { + return KeyringBundle{}, err + } + + v.CacheKeys = append(v.CacheKeys, key) + } + + if err := validateBundle(v); err != nil { + return KeyringBundle{}, err + } + + return v, nil +} + +func EncodeBundle(v KeyringBundle) ([]byte, error) { + if err := validateBundle(v); err != nil { + return nil, err + } + + raw := bundleJSON{SchemaVersion: v.SchemaVersion, Cluster: v.Cluster, Generation: v.Generation, PeerTrustRoots: v.PeerTrustRoots, CacheKeys: make([]keyJSON, 0, len(v.CacheKeys))} + for _, k := range v.CacheKeys { + raw.CacheKeys = append(raw.CacheKeys, keyJSON{Cache: k.Key.Cache, ID: k.Key.ID, Purpose: k.Key.Purpose, State: k.State, Material: k.material[:]}) + } + + return encode(raw, MaxBundleBytes) +} + +func DecodeError(r io.Reader) (ErrorResponse, error) { + var v ErrorResponse + if err := decode(r, MaxBootstrapBytes, &v); err != nil { + return ErrorResponse{}, err + } + + if !validError(v.Code) { + return ErrorResponse{}, InvalidRequest + } + + return v, nil +} + +func EncodeError(v ErrorResponse) ([]byte, error) { + if !validError(v.Code) { + return nil, InvalidRequest + } + + return encode(v, MaxBootstrapBytes) +} + +// limitedWriter bounds encoder output too, including base64 expansion and escaping. +type limitedWriter struct { + bytes.Buffer + limit int +} + +func (w *limitedWriter) Write(p []byte) (int, error) { + if len(p) > w.limit-w.Len() { + return 0, TooLarge + } + + return w.Buffer.Write(p) +} + +func encode(v any, limit int) ([]byte, error) { + w := &limitedWriter{limit: limit + 1} // Encoder adds a newline that is not on the wire. + e := json.NewEncoder(w) + e.SetEscapeHTML(false) + + if err := e.Encode(v); err != nil { + return nil, TooLarge + } + + return bytes.TrimSuffix(w.Bytes(), []byte{'\n'}), nil +} + +func decode(r io.Reader, limit int, v any) error { + b, err := io.ReadAll(io.LimitReader(r, int64(limit)+1)) + if err != nil { + return InvalidRequest + } + + if len(b) > limit { + return TooLarge + } + + if !utf8.Valid(b) || !validSurrogates(b) { + return InvalidRequest + } + + d := json.NewDecoder(bytes.NewReader(b)) + d.UseNumber() + + if err := checkValue(d, reflect.TypeOf(v).Elem(), false, 0); err != nil { + return err + } + + if _, err = d.Token(); err != io.EOF { + return InvalidRequest + } + // The original bytes are safe only after duplicate, exact-name, shape, and + // primitive checks above. No generic JSON tree is retained by validation. + if err = json.Unmarshal(b, v); err != nil { + return InvalidRequest + } + + return nil +} + +// checkValue validates against the wire type while consuming tokens. Reject +// wrong shapes before descending and oversized arrays before their next element; +// never accumulate input-sized maps or slices just to validate the document. +func checkValue(d *json.Decoder, t reflect.Type, quoted bool, depth int) error { + for t.Kind() == reflect.Pointer { + t = t.Elem() + } + + v, err := d.Token() + if err != nil { + return InvalidRequest + } + + if _, container := v.(json.Delim); container && depth >= 64 { + return InvalidRequest + } + + switch t.Kind() { + case reflect.Struct: + if v != json.Delim('{') { + return InvalidRequest + } + + return checkObject(d, t, depth) + case reflect.Slice: + if t.Elem().Kind() == reflect.Uint8 { + return checkPrimitive(v, t, quoted) + } + + if v != json.Delim('[') { + return InvalidRequest + } + + return checkArray(d, t.Elem(), depth) + default: + return checkPrimitive(v, t, quoted) + } +} + +func checkArray(d *json.Decoder, element reflect.Type, depth int) error { + for count := 0; d.More(); count++ { + if (element == reflect.TypeFor[Member]() && count >= MaxMembers) || + (element == reflect.TypeFor[RDMANIC]() && count >= MaxRDMANICs) { + return TooLarge + } + + if err := checkValue(d, element, false, depth+1); err != nil { + return err + } + } + + return checkEnd(d, ']') +} + +func checkEnd(d *json.Decoder, want json.Delim) error { + if end, err := d.Token(); err != nil || end != want { + return InvalidRequest + } + + return nil +} + +func fieldIndex(t reflect.Type, key any) int { + for i := range t.NumField() { + name, _, _ := strings.Cut(t.Field(i).Tag.Get("json"), ",") + if key == name { + return i + } + } + + return -1 +} + +func checkField(d *json.Decoder, t reflect.Type, index, depth int) error { + field := t.Field(index) + name, option, _ := strings.Cut(field.Tag.Get("json"), ",") + // An omitted GID is valid, but an explicitly empty GID is not. + if t == reflect.TypeFor[RDMANIC]() && name == "gid" { + value, err := d.Token() + + text, ok := value.(string) + if err != nil || !ok || text == "" { + return InvalidRequest + } + + return nil + } + + return checkValue(d, field.Type, option == "string", depth+1) +} + +func checkObject(d *json.Decoder, t reflect.Type, depth int) error { + // Tracking is sized by the schema, not by attacker-supplied field names. + seen := make([]bool, t.NumField()) + + for d.More() { + key, err := d.Token() + if err != nil { + return InvalidRequest + } + + index := fieldIndex(t, key) + if index < 0 || seen[index] { + return InvalidRequest + } + + seen[index] = true + if err := checkField(d, t, index, depth); err != nil { + return err + } + } + + if err := checkEnd(d, '}'); err != nil { + return err + } + + for i, present := range seen { + _, option, _ := strings.Cut(t.Field(i).Tag.Get("json"), ",") + if !present && option != "omitempty" { + return InvalidRequest + } + } + + return nil +} + +// encoding/json replaces unpaired UTF-16 surrogates with U+FFFD. Reject that +// lossy conversion so Go and Rust agree on both accepted text and hashes. +func validSurrogates(b []byte) bool { + quoted := false + + for i := 0; i < len(b); i++ { + if b[i] == '"' { + quoted = !quoted + continue + } + + if !quoted || b[i] != '\\' { + continue + } + + i++ + if i >= len(b) { + return false + } + + if b[i] != 'u' { + continue + } + + end, ok := unicodeEscapeEnd(b, i) + if !ok { + return false + } + + i = end + } + + return true +} + +// i points at the u in an escape; the result includes a required surrogate pair. +func unicodeEscapeEnd(b []byte, i int) (int, bool) { + if i+4 >= len(b) { + return i, false + } + + n, err := strconv.ParseUint(string(b[i+1:i+5]), 16, 16) + if err != nil || n >= 0xdc00 && n <= 0xdfff { + return i, false + } + + i += 4 + if n < 0xd800 || n > 0xdbff { + return i, true + } + + if i+6 >= len(b) || b[i+1] != '\\' || b[i+2] != 'u' { + return i, false + } + + n, err = strconv.ParseUint(string(b[i+3:i+7]), 16, 16) + + return i + 6, err == nil && n >= 0xdc00 && n <= 0xdfff +} + +// checkPrimitive rejects null and noncanonical encodings before typed decoding. +func checkPrimitive(v any, t reflect.Type, quoted bool) error { + switch t.Kind() { + case reflect.Slice: + s, ok := v.(string) + if !ok || t.Elem().Kind() != reflect.Uint8 { + return InvalidRequest + } + + b, err := base64.StdEncoding.Strict().DecodeString(s) + if err != nil || base64.StdEncoding.EncodeToString(b) != s { + return InvalidRequest + } + case reflect.String: + if _, ok := v.(string); !ok { + return InvalidRequest + } + case reflect.Bool: + if _, ok := v.(bool); !ok { + return InvalidRequest + } + case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + return checkUnsigned(v, t.Bits(), quoted) + default: + return InvalidRequest + } + + return nil +} + +func checkUnsigned(v any, bits int, quoted bool) error { + var text string + + if quoted { + var ok bool + + text, ok = v.(string) + if !ok { + return InvalidRequest + } + } else { + n, ok := v.(json.Number) + if !ok { + return InvalidRequest + } + + text = string(n) + } + + n, err := strconv.ParseUint(text, 10, bits) + if err != nil || strconv.FormatUint(n, 10) != text { + return InvalidRequest + } + + return nil +} + +// Canonical publications and counter-free hashes. + +// EncodePublication sorts copies of the input collections, never caller state. +func EncodePublication(v Publication) ([]byte, error) { + c, err := newCanonicalCandidate(v, true) + if err != nil { + return nil, err + } + + return c.EncodePublication(v.Sequence, v.MembershipVersion) +} + +func DecodePublication(r io.Reader) (Publication, error) { + var v Publication + if err := decode(r, MaxPublicationBytes, &v); err != nil { + return Publication{}, err + } + + if err := validatePublication(v, true); err != nil { + return Publication{}, err + } + + return canonicalPublication(v), nil +} + +func canonicalPublication(v Publication) Publication { + v.Members = append([]Member{}, v.Members...) + + v.Caches = append([]CacheDefinition{}, v.Caches...) + for i := range v.Members { + v.Members[i].RDMANICs = CanonicalRDMANICs(v.Members[i].RDMANICs) + } + + slices.SortFunc(v.Members, func(a, b Member) int { return cmp.Compare(a.Node, b.Node) }) + slices.SortFunc(v.Caches, func(a, b CacheDefinition) int { return cmp.Compare(a.ID, b.ID) }) + + return v +} + +// CanonicalRDMANICs copies NICs, including nested pointers, in rail/device/port order. +func CanonicalRDMANICs(nics []RDMANIC) []RDMANIC { + nics = append([]RDMANIC{}, nics...) + for i := range nics { + if n := nics[i].NUMANode; n != nil { + value := *n + nics[i].NUMANode = &value + } + } + + slices.SortFunc(nics, func(a, b RDMANIC) int { + if n := cmp.Compare(a.Rail, b.Rail); n != 0 { + return n + } + + if n := cmp.Compare(a.Device, b.Device); n != 0 { + return n + } + + return cmp.Compare(a.Port, b.Port) + }) + + return nics +} + +// CanonicalCandidate owns validated, sorted publication content, including nested +// pointers. Its private state can be reused for hashing and encoding after counter +// assignment without retaining caller-owned mutable state. The zero value is invalid. +type CanonicalCandidate struct { + publication Publication +} + +// NewCanonicalCandidate validates and copies content once. Input counters are +// ignored; final encoding checks the assigned counters and complete byte bound. +func NewCanonicalCandidate(v Publication) (CanonicalCandidate, error) { + return newCanonicalCandidate(v, false) +} + +func newCanonicalCandidate(v Publication, counters bool) (CanonicalCandidate, error) { + if err := validatePublication(v, counters); err != nil { + return CanonicalCandidate{}, err + } + + return CanonicalCandidate{publication: canonicalPublication(v)}, nil +} + +// EncodePublication encodes the candidate with nonzero counters. It does not +// mutate the candidate, and each call returns independently owned bytes. +func (c CanonicalCandidate) EncodePublication(sequence Sequence, membership MembershipVersion) ([]byte, error) { + v := c.publication + if err := validateHeader(v.SchemaVersion, v.Cluster); err != nil { + return nil, err + } + + if sequence == 0 || membership == 0 { + return nil, InvalidRequest + } + + v.Sequence, v.MembershipVersion = sequence, membership + + return encode(v, MaxPublicationBytes) +} + +// canonicalContent returns counter-free canonical JSON for durable version CAS. +// The membership document includes schema and cluster, and every member input. +// Input counters may be zero because callers hash candidates before assigning them. +func (c CanonicalCandidate) canonicalContent() (content, membership []byte, err error) { + v := c.publication + if err := validateHeader(v.SchemaVersion, v.Cluster); err != nil { + return nil, nil, err + } + + m := struct { + SchemaVersion uint32 `json:"schema_version"` + Cluster ClusterID `json:"cluster"` + Members []Member `json:"members"` + }{v.SchemaVersion, v.Cluster, v.Members} + p := struct { + SchemaVersion uint32 `json:"schema_version"` + Cluster ClusterID `json:"cluster"` + Members []Member `json:"members"` + Caches []CacheDefinition `json:"caches"` + }{v.SchemaVersion, v.Cluster, v.Members, v.Caches} + + content, err = encode(p, MaxPublicationBytes) + if err != nil { + return nil, nil, err + } + + membership, err = encode(m, MaxPublicationBytes) + + return content, membership, err +} + +// ContentHashes returns lowercase SHA-256 hex, suitable for VersionRecord fields. +func ContentHashes(v Publication) (content, membership string, err error) { + c, err := NewCanonicalCandidate(v) + if err != nil { + return "", "", err + } + + return c.ContentHashes() +} + +// ContentHashes returns the same counter-free hashes as ContentHashes without +// repeating validation, copying, or sorting of the candidate's collections. +func (c CanonicalCandidate) ContentHashes() (content, membership string, err error) { + p, m, err := c.canonicalContent() + if err != nil { + return "", "", err + } + + ph, mh := sha256.Sum256(p), sha256.Sum256(m) + + return hex.EncodeToString(ph[:]), hex.EncodeToString(mh[:]), nil +} + +// Snapshot deltas use the same validation and hashes as full publications. + +const ( + DeltaHeader = "X-Racer-Delta-Base" + MaxDeltaBytes = 4 * 1024 * 1024 +) + +// Delta is transported only on the authenticated snapshot endpoint. Hashes use +// ContentHashes, not a second canonical representation or placement authority. +type Delta struct { + DeltaVersion uint32 `json:"delta_version"` + Cluster ClusterID `json:"cluster"` + BaseSequence Sequence `json:"base_sequence,string"` + BaseHash string `json:"base_hash"` + Sequence Sequence `json:"sequence,string"` + MembershipVersion MembershipVersion `json:"membership_version,string"` + ContentHash string `json:"content_hash"` + UpsertMembers []Member `json:"upsert_members"` + RemoveMembers []NodeID `json:"remove_members"` + Caches []CacheDefinition `json:"caches"` +} + +func EncodeDelta(base, next Publication) ([]byte, error) { + b, err := newCanonicalCandidate(base, true) + if err != nil { + return nil, err + } + + n, err := newCanonicalCandidate(next, true) + if err != nil { + return nil, err + } + + if base.Cluster != next.Cluster || next.Sequence <= base.Sequence { + return nil, Conflict + } + + base, next = b.publication, n.publication + + bh, _, err := b.ContentHashes() + if err != nil { + return nil, err + } + + nh, _, err := n.ContentHashes() + if err != nil { + return nil, err + } + + d := Delta{DeltaVersion: 1, Cluster: next.Cluster, BaseSequence: base.Sequence, BaseHash: bh, Sequence: next.Sequence, MembershipVersion: next.MembershipVersion, ContentHash: nh, UpsertMembers: []Member{}, RemoveMembers: []NodeID{}, Caches: next.Caches} + + old := make(map[NodeID]Member, len(base.Members)) + for _, m := range base.Members { + old[m.Node] = m + } + + for _, m := range next.Members { + if previous, ok := old[m.Node]; !ok || !reflect.DeepEqual(previous, m) { + d.UpsertMembers = append(d.UpsertMembers, m) + } + + delete(old, m.Node) + } + + for _, m := range base.Members { + if _, ok := old[m.Node]; ok { + d.RemoveMembers = append(d.RemoveMembers, m.Node) + } + } + + return encode(d, MaxDeltaBytes) +} + +func ApplyDelta(base Publication, reader io.Reader) (Publication, error) { + var d Delta + if err := decode(reader, MaxDeltaBytes, &d); err != nil { + return Publication{}, err + } + + hash, _, err := ContentHashes(base) + if err != nil { + return Publication{}, err + } + + if d.DeltaVersion != 1 || d.Cluster != base.Cluster || d.BaseSequence != base.Sequence || d.BaseHash != hash || d.Sequence <= base.Sequence || d.MembershipVersion < base.MembershipVersion { + return Publication{}, Conflict + } + + members, err := applyMemberChanges(base.Members, d) + if err != nil { + return Publication{}, err + } + + next := Publication{SchemaVersion: SchemaVersion, Cluster: d.Cluster, Sequence: d.Sequence, MembershipVersion: d.MembershipVersion, Caches: d.Caches, Members: members} + + candidate, err := newCanonicalCandidate(next, true) + if err != nil { + return Publication{}, err + } + + hash, _, err = candidate.ContentHashes() + if err != nil { + return Publication{}, err + } + + if hash != d.ContentHash { + return Publication{}, Conflict + } + + return candidate.publication, nil +} + +func applyMemberChanges(base []Member, d Delta) ([]Member, error) { + members := make(map[NodeID]Member, len(base)) + for _, m := range base { + members[m.Node] = m + } + // A node may occur only once across both change lists. + seen := make(map[NodeID]bool) + for _, id := range d.RemoveMembers { + if _, ok := members[id]; !ok || seen[id] { + return nil, InvalidRequest + } + + seen[id] = true + delete(members, id) + } + + for _, m := range d.UpsertMembers { + if seen[m.Node] { + return nil, InvalidRequest + } + + seen[m.Node] = true + members[m.Node] = m + } + + next := make([]Member, 0, len(members)) + for _, m := range members { + next = append(next, m) + } + + return next, nil +} diff --git a/internal/racer/wire/wire_test.go b/internal/racer/wire/wire_test.go new file mode 100644 index 000000000..61667c702 --- /dev/null +++ b/internal/racer/wire/wire_test.go @@ -0,0 +1,1682 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 + +package wire + +import ( + "bytes" + "crypto/ed25519" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "math/big" + "os" + "reflect" + "slices" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func vectorRoundTrip(name string, b []byte) ([]byte, error) { + r := bytes.NewReader(b) + + switch name { + case "publication.json": + v, err := DecodePublication(r) + if err != nil { + return nil, err + } + + return EncodePublication(v) + case "bootstrap-request.json": + v, err := DecodeBootstrap(r) + if err != nil { + return nil, err + } + + return EncodeBootstrapRequest(v) + case "bootstrap-response.json": + v, err := DecodeBootstrapResponse(r) + if err != nil { + return nil, err + } + + return EncodeBootstrap(v) + case "bundle.json": + v, err := DecodeBundle(r) + if err != nil { + return nil, err + } + + return EncodeBundle(v) + default: + return nil, InvalidRequest + } +} + +func TestSharedVectors(t *testing.T) { + for _, name := range []string{"bootstrap-request.json", "bootstrap-response.json", "bundle.json"} { + t.Run(name, func(t *testing.T) { + b := fixture(t, name) + + encoded, err := vectorRoundTrip(name, b) + if err != nil { + t.Fatal(err) + } + + if !bytes.Equal(b, encoded) { + t.Fatal("wire bytes changed") + } + }) + } + + var cases []struct { + Name, File, Old, New string + Code ErrorCode + } + if err := json.Unmarshal(fixture(t, "rejections.json"), &cases); err != nil { + t.Fatal(err) + } + + for _, tc := range cases { + t.Run(tc.Name, func(t *testing.T) { + b := fixture(t, tc.File) + + changed := bytes.Replace(b, []byte(tc.Old), []byte(tc.New), 1) + if bytes.Equal(b, changed) { + t.Fatal("mutation did not match") + } + + if _, err := vectorRoundTrip(tc.File, changed); !errors.Is(err, tc.Code) { + t.Fatalf("got %v, want %v", err, tc.Code) + } + }) + } +} + +func TestBootstrapBlockDevicesVectors(t *testing.T) { + var vectors []struct { + Name, Fields, Pattern string + Code ErrorCode + } + require.NoError(t, json.Unmarshal(fixture(t, "bootstrap-block-devices.json"), &vectors)) + + base := fixture(t, "bootstrap-response.json") + for _, vector := range vectors { + t.Run(vector.Name, func(t *testing.T) { + raw := string(base[:len(base)-1]) + vector.Fields + "}" + + response, err := DecodeBootstrapResponse(strings.NewReader(raw)) + if vector.Code != "" { + require.ErrorIs(t, err, vector.Code) + return + } + + require.NoError(t, err) + require.Equal(t, vector.Pattern, response.BlockDevices) + encoded, err := EncodeBootstrap(response) + require.NoError(t, err) + + if vector.Pattern == "" { + require.Equal(t, base, encoded) + } else { + require.Equal(t, raw, string(encoded)) + } + }) + } +} + +func TestCanonicalHashSemantics(t *testing.T) { + v, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + ph, mh, err := ContentHashes(v) + require.NoError(t, err) + + v.Sequence, v.MembershipVersion = 0, 0 + // Reorder all three collections. Hashing must not mutate caller-owned slices. + v.Members[0], v.Members[1] = v.Members[1], v.Members[0] + v.Members[0].RDMANICs[0], v.Members[0].RDMANICs[1] = v.Members[0].RDMANICs[1], v.Members[0].RDMANICs[0] + + before, err := json.Marshal(v) + require.NoError(t, err) + + p2, m2, err := ContentHashes(v) + require.NoError(t, err) + require.Equal(t, ph, p2, "ordering or counters changed content hash") + require.Equal(t, mh, m2, "ordering or counters changed membership hash") + + after, err := json.Marshal(v) + require.NoError(t, err) + require.Equal(t, before, after, "hash mutated input") + + v.Caches[0].ID = "66666666-6666-4666-8666-666666666666" + + p2, m2, err = ContentHashes(v) + require.NoError(t, err) + require.NotEqual(t, ph, p2, "cache changes content hash") + require.Equal(t, mh, m2, "cache does not change membership hash") + + v.Members[0].PeerEndpoint = "[2001:db8::2]:7443" + + p3, m3, err := ContentHashes(v) + require.NoError(t, err) + require.NotEqual(t, p2, p3, "endpoint changes content hash") + require.NotEqual(t, m2, m3, "endpoint changes membership hash") + + for _, mutate := range []func(*Publication){ + func(p *Publication) { p.Members[0].Shares-- }, + func(p *Publication) { p.Members[0].RDMANICs[0].Port++ }, + func(p *Publication) { p.Members[0].RDMANICs[0].Device += "-new" }, + func(p *Publication) { p.Cluster = "aaaaaaaa-1111-4111-8111-111111111111" }, + } { + mutate(&v) + + nextP, nextM, err := ContentHashes(v) + if err != nil || p3 == nextP || m3 == nextM { + t.Fatal("member/cluster content not hashed", err) + } + + p3, m3 = nextP, nextM + } +} + +type repeatedReader struct{ remaining, read int } + +func (r *repeatedReader) Read(p []byte) (int, error) { + if r.remaining == 0 { + return 0, io.EOF + } + + n := min(len(p), r.remaining) + clear(p[:n]) + r.remaining -= n + r.read += n + + return n, nil +} + +func TestByteBoundsAndMalformedDocuments(t *testing.T) { + for _, tc := range []struct { + name string + limit int + }{{"bootstrap-request.json", MaxBootstrapBytes}, {"bootstrap-response.json", MaxBootstrapBytes}, {"bundle.json", MaxBundleBytes}, {"publication.json", MaxPublicationBytes}} { + t.Run(tc.name, func(t *testing.T) { + b := fixture(t, tc.name) + + padded := append(bytes.Clone(b), bytes.Repeat([]byte{' '}, tc.limit-len(b))...) + if _, err := vectorRoundTrip(tc.name, padded); err != nil { + t.Fatalf("exact bound: %v", err) + } + + if _, err := vectorRoundTrip(tc.name, append(padded, ' ')); !errors.Is(err, TooLarge) { + t.Fatalf("over bound: %v", err) + } + + for _, bad := range [][]byte{nil, []byte("null"), []byte("[]"), b[:len(b)-1], append(bytes.Clone(b), []byte("{}")...), append(bytes.Clone(b), 0xff), []byte(strings.Repeat("[", 1000) + strings.Repeat("]", 1000))} { + if _, err := vectorRoundTrip(tc.name, bad); !errors.Is(err, InvalidRequest) { + t.Fatalf("malformed: %v", err) + } + } + }) + } + + r := &repeatedReader{remaining: MaxBootstrapBytes * 10} + if _, err := DecodeBootstrap(r); !errors.Is(err, TooLarge) || r.read != MaxBootstrapBytes+1 { + t.Fatalf("unbounded read: %d, %v", r.read, err) + } +} + +func TestPublicationValidationAndLimits(t *testing.T) { + v := Publication{SchemaVersion: 1, Cluster: "11111111-1111-4111-8111-111111111111", Sequence: 1, MembershipVersion: 1} + + b, err := EncodePublication(v) + if err != nil || !bytes.Contains(b, []byte(`"members":[],"caches":[]`)) { + t.Fatal("nil collections not normalized", err) + } + + v.Members = make([]Member, MaxMembers+1) + if _, err := EncodePublication(v); !errors.Is(err, TooLarge) { + t.Fatal("member limit", err) + } + + for _, name := range []string{".", "..", "a/b", "a\\b", "a\x00b", "A", "-a", "a-", "a..b", strings.Repeat("a", 64)} { + if _, _, err := CanonicalSocketPaths(name); err == nil { + t.Fatalf("accepted unsafe name %q", name) + } + } + + name := strings.Repeat("a", 63) + "." + strings.Repeat("b", 18) + + client, _, err := CanonicalSocketPaths(name) + if err != nil || len(client) != 107 { + t.Fatalf("UDS boundary: %d %v", len(client), err) + } + + if _, _, err := CanonicalSocketPaths(name + "b"); err == nil { + t.Fatal("UDS overflow") + } + + v, err = DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + if err != nil { + t.Fatal(err) + } + + v.Caches = append(v.Caches, v.Caches[0]) + if _, err := EncodePublication(v); err == nil { + t.Fatal("duplicate cache") + } + + v.Caches[1].ID = "66666666-6666-4666-8666-666666666666" + if _, err := EncodePublication(v); err == nil { + t.Fatal("duplicate cache name") + } + + v.Caches = nil + + v.Members[0].RDMANICs = []RDMANIC{{Device: strings.Repeat("x", MaxPublicationBytes), Port: 1}} + if _, err := EncodePublication(v); !errors.Is(err, TooLarge) { + t.Fatal("encoded byte limit", err) + } +} + +func TestUnknownFieldsAndExactNames(t *testing.T) { + b := fixture(t, "publication.json") + + b = bytes.Replace(b, []byte(`"schema_version":1`), []byte(`"schema_version":1,"SCHEMA_VERSION":42`), 1) + if _, err := DecodePublication(bytes.NewReader(b)); err == nil { + t.Fatal("unknown case variant accepted") + } + + b = bytes.Replace(b, []byte(`"schema_version":1,`), nil, 1) + if _, err := DecodePublication(bytes.NewReader(b)); err == nil { + t.Fatal("case variant supplied required field") + } +} + +func TestValidatedOriginalBytesPreserveEscapes(t *testing.T) { + for _, name := range []string{"publication.json", "bootstrap-request.json", "bootstrap-response.json", "bundle.json"} { + t.Run(name, func(t *testing.T) { + original := fixture(t, name) + + want, err := vectorRoundTrip(name, original) + if err != nil { + t.Fatal(err) + } + + // Escaped exact field names and scalar text are valid, but escaped + // aliases must still be rejected as duplicates before typed decoding. + escaped := bytes.ReplaceAll(original, []byte(`"schema_version"`), []byte(`"schema_versi\u006fn"`)) + escaped = bytes.ReplaceAll(escaped, []byte(`"18446744073709551615"`), []byte(`"\u00318446744073709551615"`)) + escaped = bytes.ReplaceAll(escaped, []byte(`11111111-`), []byte(`\u00311111111-`)) + + got, err := vectorRoundTrip(name, escaped) + if err != nil || !bytes.Equal(got, want) { + t.Fatalf("escaped document changed semantics: %v", err) + } + }) + } +} + +func TestBundleValidationAndKeyIsolation(t *testing.T) { + b, err := DecodeBundle(bytes.NewReader(fixture(t, "bundle.json"))) + if err != nil { + t.Fatal(err) + } + + ref := b.CacheKeys[0].Key + + key, err := NewCacheKey(ref, ActiveKey, [32]byte{}) + if err != nil { + t.Fatal(err) + } + + ref.ID[0] = 123 + if key.Key.ID[0] != 'R' { + t.Fatal("material ingress retained mutable id") + } + + if !key.EqualMaterial(b.CacheKeys[0]) || key.EqualMaterial(b.CacheKeys[1]) { + t.Fatal("material equality") + } + + b.PeerTrustRoots = append(b.PeerTrustRoots, b.PeerTrustRoots[0]) + if _, err := EncodeBundle(b); err == nil { + t.Fatal("duplicate trust root") + } + + for _, purpose := range []KeyPurpose{"", "future"} { + ref.Purpose = purpose + if _, err := NewCacheKey(ref, ActiveKey, [32]byte{}); err == nil { + t.Fatal("invalid purpose") + } + } + + for _, code := range []ErrorCode{InvalidRequest, Unauthenticated, Forbidden, Conflict, TooLarge, UnsupportedVersion, Overloaded, Unavailable} { + encoded, err := EncodeError(ErrorResponse{Code: code}) + if err != nil { + t.Fatal(err) + } + + decoded, err := DecodeError(bytes.NewReader(encoded)) + if err != nil || decoded.Code != code { + t.Fatal("error codec", err) + } + } + + if _, err := DecodeError(strings.NewReader(`{"code":"future"}`)); err == nil { + t.Fatal("unknown error enum") + } +} + +func FuzzDecodePublication(f *testing.F) { + f.Add([]byte(`{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","sequence":"1","membership_version":"1","members":[],"caches":[]}`)) + f.Add([]byte(`{"x":1,"x":2}`)) + f.Fuzz(func(t *testing.T, b []byte) { + v, err := DecodePublication(bytes.NewReader(b)) + if err != nil { + if !reflect.DeepEqual(v, Publication{}) { + t.Fatal("partial result on error") + } + + return + } + + encoded, err := EncodePublication(v) + if err != nil { + t.Fatal("decoded invalid publication", err) + } + + if _, err := DecodePublication(bytes.NewReader(encoded)); err != nil { + t.Fatal("invalid re-encoding", err) + } + }) +} + +func TestMaximumMembership(t *testing.T) { + v := Publication{SchemaVersion: 1, Cluster: "11111111-1111-4111-8111-111111111111", Sequence: 1, MembershipVersion: 1, Members: make([]Member, MaxMembers)} + for i := range v.Members { + v.Members[i] = Member{Node: NodeID(fmt.Sprintf("%08x-1111-4111-8111-111111111111", i)), Shares: 1, PeerEndpoint: "192.0.2.1:1"} + } + + b, err := EncodePublication(v) + if err != nil { + t.Fatal(err) + } + + decoded, err := DecodePublication(bytes.NewReader(b)) + if err != nil || len(decoded.Members) != MaxMembers { + t.Fatal("exact member limit", err) + } + // Input remains below the byte bound but exceeds the independent member cap. + start := bytes.Index(b, []byte(`"members":[`)) + len(`"members":[`) + end := bytes.IndexByte(b[start:], '}') + start + 1 + tooMany := append(bytes.Clone(b[:start]), append(bytes.Clone(b[start:end]), ',')...) + + tooMany = append(tooMany, b[start:]...) + if _, err := DecodePublication(bytes.NewReader(tooMany)); !errors.Is(err, TooLarge) { + t.Fatal("decoded member cap", err) + } +} + +func FuzzDecodeSecretAndEnrollment(f *testing.F) { + f.Add([]byte(`{"schema_version":1}`)) + f.Add([]byte(`{"cache_keys":[{"material":"AA=="}]}`)) + f.Fuzz(func(t *testing.T, b []byte) { + if v, err := DecodeBundle(bytes.NewReader(b)); err == nil { + encoded, err := EncodeBundle(v) + if err != nil { + t.Fatal("decoded invalid bundle", err) + } + + if _, err := DecodeBundle(bytes.NewReader(encoded)); err != nil { + t.Fatal("bundle re-encoding", err) + } + } + + if v, err := DecodeBootstrap(bytes.NewReader(b)); err == nil { + if _, err := EncodeBootstrapRequest(v); err != nil { + t.Fatal("decoded invalid request", err) + } + } + + if v, err := DecodeBootstrapResponse(bytes.NewReader(b)); err == nil { + if _, err := EncodeBootstrap(v); err != nil { + t.Fatal("decoded invalid response", err) + } + } + }) +} + +// Fixtures and public certificate inputs. + +func fixture(t *testing.T, name string) []byte { + t.Helper() + + b, err := os.ReadFile("testdata/" + name) + require.NoError(t, err) + + return bytes.TrimSuffix(b, []byte{'\n'}) +} + +func TestSharedPublicationVector(t *testing.T) { + v, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + candidate, err := NewCanonicalCandidate(v) + require.NoError(t, err) + p, m, err := candidate.canonicalContent() + require.NoError(t, err) + require.Equal(t, fixture(t, "content.json"), p) + require.Equal(t, fixture(t, "membership.json"), m) + + ph, mh, err := ContentHashes(v) + require.NoError(t, err) + + var hashes struct{ Content, Membership string } + require.NoError(t, json.Unmarshal(fixture(t, "hashes.json"), &hashes)) + require.Equal(t, hashes.Content, ph) + require.Equal(t, hashes.Membership, mh) + + encoded, err := EncodePublication(v) + require.NoError(t, err) + replay, err := DecodePublication(bytes.NewReader(encoded)) + require.NoError(t, err) + again, err := EncodePublication(replay) + require.NoError(t, err) + require.Equal(t, encoded, again, "unstable round trip") +} + +// The key is ephemeral and never printed or persisted. Wire validation checks +// DER syntax, not trust or CSR authority. +func publicDER(t *testing.T) (cert, csr []byte) { + t.Helper() + + pub, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "wire fixture"}, NotBefore: time.Unix(0, 0), NotAfter: time.Unix(2000000000, 0), IsCA: true, BasicConstraintsValid: true, KeyUsage: x509.KeyUsageCertSign} + cert, err = x509.CreateCertificate(rand.Reader, template, template, pub, key) + require.NoError(t, err) + csr, err = x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{Subject: template.Subject}, key) + require.NoError(t, err) + + return cert, csr +} + +func TestPublicDEREncoding(t *testing.T) { + cert, csr := publicDER(t) + cluster := ClusterID("11111111-1111-4111-8111-111111111111") + enrollment := EnrollmentID("55555555-5555-4555-8555-555555555555") + request, err := EncodeBootstrapRequest(BootstrapRequest{Shares: DefaultShares, SchemaVersion: 1, Cluster: cluster, Enrollment: enrollment, CSRDER: csr}) + require.NoError(t, err) + _, err = DecodeBootstrap(bytes.NewReader(request)) + require.NoError(t, err) + response, err := EncodeBootstrap(BootstrapResponse{SchemaVersion: 1, Cluster: cluster, Enrollment: enrollment, Node: "22222222-2222-4222-8222-222222222222", CertificateChain: [][]byte{cert}}) + require.NoError(t, err) + _, err = DecodeBootstrapResponse(bytes.NewReader(response)) + require.NoError(t, err) + _, err = json.Marshal(KeyringBundle{}) + require.Error(t, err, "invalid bundle accepted") +} + +// Bootstrap bounds include JSON framing and base64 expansion. + +func TestBootstrapRequestEncodedBoundary(t *testing.T) { + request, err := DecodeBootstrap(bytes.NewReader(fixture(t, "bootstrap-request.json"))) + require.NoError(t, err) + + request.CSRDER = []byte{} + request.RDMANICs = []RDMANIC{{Device: "mlx5_0", Port: 1, Rail: 1, GID: "abcdef0123456789abcdef0123456789"}} + framing, err := json.Marshal(request) + require.NoError(t, err) + + maxDER := (MaxBootstrapBytes - len(framing)) / 4 * 3 + _, key, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + + makeCSR := func(padding int) []byte { + t.Helper() + + der, err := x509.CreateCertificateRequest(rand.Reader, &x509.CertificateRequest{Subject: pkix.Name{CommonName: strings.Repeat("x", padding)}}, key) + require.NoError(t, err) + + return der + } + padding := maxDER - 256 + padding += maxDER - len(makeCSR(padding)) + // Cover every base64 padding case and the first DER byte that overflows. + for _, delta := range []int{-2, -1, 0, 1} { + t.Run(fmt.Sprint(delta), func(t *testing.T) { + request.CSRDER = makeCSR(padding + delta) + require.Len(t, request.CSRDER, maxDER+delta) + raw, err := json.Marshal(request) + require.NoError(t, err) + + validationErr := ValidateBootstrapRequest(request) + encoded, encodeErr := EncodeBootstrapRequest(request) + + _, decodeErr := DecodeBootstrap(bytes.NewReader(raw)) + if delta == 1 { + require.Greater(t, len(raw), MaxBootstrapBytes) + require.Less(t, len(request.CSRDER), MaxBootstrapBytes) + require.ErrorIs(t, validationErr, TooLarge) + require.ErrorIs(t, encodeErr, TooLarge) + require.ErrorIs(t, decodeErr, TooLarge) + require.Nil(t, encoded) + + return + } + + require.Less(t, MaxBootstrapBytes-len(raw), 4) + require.NoError(t, validationErr) + require.NoError(t, encodeErr) + require.NoError(t, decodeErr) + require.Equal(t, raw, encoded) + }) + } +} + +func TestBootstrapResponseEncodedBoundary(t *testing.T) { + response, err := DecodeBootstrapResponse(bytes.NewReader(fixture(t, "bootstrap-response.json"))) + require.NoError(t, err) + + cert := response.CertificateChain[0] + response.CertificateChain = nil + + var previous []byte + + for { + response.CertificateChain = append(response.CertificateChain, cert) + raw, err := json.Marshal(response) + require.NoError(t, err) + + encoded, encodeErr := EncodeBootstrap(response) + if len(raw) <= MaxBootstrapBytes { + require.NoError(t, encodeErr) + + previous = encoded + + continue + } + + require.Less(t, len(response.CertificateChain)*len(cert), MaxBootstrapBytes, "fixture must overflow only after encoding") + require.ErrorIs(t, encodeErr, TooLarge) + require.Nil(t, encoded) + + _, err = DecodeBootstrapResponse(bytes.NewReader(raw)) + require.ErrorIs(t, err, TooLarge) + _, err = DecodeBootstrapResponse(bytes.NewReader(previous)) + require.NoError(t, err, "last fitting response") + + break + } +} + +// Bundle identity, generation, and material isolation. + +func TestKeyringDeliveryEncodingBoundsAndGeneration(t *testing.T) { + bundle, err := DecodeBundle(bytes.NewReader(fixture(t, "bundle.json"))) + require.NoError(t, err) + + bundle.Generation = ^Generation(0) + encoded, err := EncodeBundle(bundle) + require.NoError(t, err) + require.Contains(t, string(encoded), `"generation":"18446744073709551615"`) + decoded, err := DecodeBundle(bytes.NewReader(encoded)) + require.NoError(t, err) + require.Equal(t, bundle.Generation, decoded.Generation) + require.True(t, decoded.CacheKeys[0].EqualMaterial(bundle.CacheKeys[0])) + + for n := uint64(1); n <= 4000; n++ { + ref := bundle.CacheKeys[0].Key + ref.ID = make([]byte, 16) + copy(ref.ID, "RKG1") + binary.BigEndian.PutUint64(ref.ID[4:12], n+10) + + var material [32]byte + binary.BigEndian.PutUint64(material[:8], n+10) + key, err := NewCacheKey(ref, PreparedKey, material) + require.NoError(t, err) + + bundle.CacheKeys = append(bundle.CacheKeys, key) + } + + encoded, err = EncodeBundle(bundle) + require.ErrorIs(t, err, TooLarge) + require.Nil(t, encoded) +} + +func TestBundleEncodingCannotBypassCodec(t *testing.T) { + _, err := json.Marshal(KeyringBundle{}) + require.ErrorIs(t, err, UnsupportedVersion) +} + +// Rust rejects material shared by distinct refs regardless of scope or state. +func TestBundleDuplicateMaterialParity(t *testing.T) { + for _, scope := range []string{"id", "purpose", "cache"} { + t.Run(scope, func(t *testing.T) { + bundle, err := DecodeBundle(bytes.NewReader(fixture(t, "bundle.json"))) + require.NoError(t, err) + + first := bundle.CacheKeys[0] + second := first + second.Key.ID = bytes.Clone(first.Key.ID) + + switch scope { + case "id": + second.Key.ID[15]++ + second.State = PreparedKey + case "purpose": + second.Key.Purpose = OriginCredentialsKey + if first.Key.Purpose == OriginCredentialsKey { + second.Key.Purpose = PageKey + } + case "cache": + second.Key.Cache = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" + } + + second.material[0] ^= 0xff + bundle.CacheKeys = []CacheKey{first, second} + valid, err := EncodeBundle(bundle) + require.NoError(t, err, "distinct material rejected") + + var document bundleJSON + require.NoError(t, json.Unmarshal(valid, &document)) + document.CacheKeys[1].Material = document.CacheKeys[0].Material + corrupt, err := json.Marshal(document) + require.NoError(t, err) + _, err = DecodeBundle(bytes.NewReader(corrupt)) + require.ErrorIs(t, err, InvalidRequest, "decoder accepted duplicate material") + + bundle.CacheKeys[1].material = first.material + _, err = EncodeBundle(bundle) + require.ErrorIs(t, err, InvalidRequest, "encoder accepted duplicate material") + }) + } +} + +func TestRetiringKeyStateIsRejected(t *testing.T) { + raw := fixture(t, "bundle.json") + // Leave an active key so this checks the enum, not the active-key count. + changed := bytes.Replace(raw, []byte(`"state":"prepared"`), []byte(`"state":"retiring"`), 1) + require.NotEqual(t, raw, changed, "mutation did not match") + _, err := DecodeBundle(bytes.NewReader(changed)) + require.ErrorIs(t, err, InvalidRequest) + bundle, err := DecodeBundle(bytes.NewReader(raw)) + require.NoError(t, err) + _, err = NewCacheKey(bundle.CacheKeys[0].Key, KeyState("retiring"), [32]byte{}) + require.ErrorIs(t, err, InvalidRequest) + + bundle.CacheKeys[1].State = KeyState("retiring") + _, err = EncodeBundle(bundle) + require.ErrorIs(t, err, InvalidRequest) +} + +func TestKeyIDsRequireCurrentNamespaceAndValidGeneration(t *testing.T) { + bundle, err := DecodeBundle(bytes.NewReader(fixture(t, "bundle.json"))) + require.NoError(t, err) + + bundle.Generation = 2 + for _, generation := range []uint64{0, 1, 2, 3} { + ref := bundle.CacheKeys[0].Key + ref.ID = bytes.Clone(ref.ID) + binary.BigEndian.PutUint64(ref.ID[4:12], generation) + + key, err := NewCacheKey(ref, ActiveKey, [32]byte{}) + if generation == 0 { + require.Error(t, err, "zero generation accepted") + continue + } + + require.NoError(t, err) + + candidate := bundle + candidate.CacheKeys = []CacheKey{key} + _, err = EncodeBundle(candidate) + require.Equal(t, generation <= 2, err == nil, "generation %d: %v", generation, err) + } + + for _, id := range [][]byte{make([]byte, 16), []byte("RKG0abcdefgh1234"), []byte("RKG1")} { + ref := bundle.CacheKeys[0].Key + ref.ID = id + _, err := NewCacheKey(ref, ActiveKey, [32]byte{}) + require.Error(t, err, "legacy ID accepted") + } +} + +func TestKeyDiagnosticsRedactMaterial(t *testing.T) { + key := CacheKey{material: [32]byte{1, 2, 3}} + for _, format := range []string{"%v", "%+v", "%#v"} { + require.Equal(t, "", fmt.Sprintf(format, key)) + } + + encoded, err := json.Marshal(key) + require.NoError(t, err) + require.NotContains(t, string(encoded), "material") +} + +// Canonical candidates own their input and output independently of callers. + +func TestCanonicalCandidateEquivalence(t *testing.T) { + vector, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + vector.Caches = append(vector.Caches, CacheDefinition{ID: "00000000-0000-4000-8000-000000000000", Name: "cache-b", ClientSocket: "/run/racer/cache-b/client/socket", OriginSocket: "/run/racer/cache-b/origin/socket"}) + slices.Reverse(vector.Members) + + for _, v := range []Publication{ + vector, + {SchemaVersion: SchemaVersion, Cluster: vector.Cluster}, + {SchemaVersion: SchemaVersion, Cluster: vector.Cluster, Members: []Member{}, Caches: []CacheDefinition{}}, + } { + before, err := json.Marshal(v) + require.NoError(t, err) + candidate, err := NewCanonicalCandidate(v) + require.NoError(t, err) + wantContent, wantMembership, err := ContentHashes(v) + require.NoError(t, err) + + sequence, membershipVersion := v.Sequence, v.MembershipVersion + for _, counters := range [][2]uint64{{1, 1}, {10, 9}, {1, 2}, {^uint64(0), ^uint64(0)}} { + v.Sequence, v.MembershipVersion = Sequence(counters[0]), MembershipVersion(counters[1]) + want, err := EncodePublication(v) + require.NoError(t, err) + got, err := candidate.EncodePublication(v.Sequence, v.MembershipVersion) + require.NoError(t, err) + require.Equal(t, want, got, "counters %v", counters) + + content, membership, err := candidate.ContentHashes() + require.NoError(t, err) + require.Equal(t, wantContent, content) + require.Equal(t, wantMembership, membership) + } + + v.Sequence, v.MembershipVersion = sequence, membershipVersion + after, err := json.Marshal(v) + require.NoError(t, err) + require.Equal(t, before, after, "candidate mutated input") + } +} + +func TestCanonicalCandidateOwnsNestedStateAndOutput(t *testing.T) { + v, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + want, err := EncodePublication(v) + require.NoError(t, err) + candidate, err := NewCanonicalCandidate(v) + require.NoError(t, err) + // Mutate all caller collections, including the nested pointer, before use. + *v.Members[1].RDMANICs[1].NUMANode = 7 + v.Members[1].RDMANICs[0].Device = "changed" + v.Members[0].Shares = 0 + v.Caches[0].Name = "changed" + clear(v.Members) + clear(v.Caches) + got, err := candidate.EncodePublication(v.Sequence, v.MembershipVersion) + require.NoError(t, err) + require.Equal(t, want, got) + + var hashes struct{ Content, Membership string } + require.NoError(t, json.Unmarshal(fixture(t, "hashes.json"), &hashes)) + + content, membership, err := candidate.ContentHashes() + require.NoError(t, err) + require.Equal(t, hashes.Content, content) + require.Equal(t, hashes.Membership, membership) + clear(got) + got, err = candidate.EncodePublication(v.Sequence, v.MembershipVersion) + require.NoError(t, err) + require.Equal(t, want, got, "encoding retained returned bytes") +} + +func TestCanonicalCandidateValidation(t *testing.T) { + for _, tc := range []struct { + name string + edit func(*Publication) + want error + }{ + {"schema", func(v *Publication) { v.SchemaVersion++ }, UnsupportedVersion}, + {"cluster", func(v *Publication) { v.Cluster = "invalid" }, InvalidRequest}, + {"node", func(v *Publication) { v.Members[0].Node = "invalid" }, InvalidRequest}, + {"duplicate node", func(v *Publication) { v.Members = append(v.Members, v.Members[0]) }, InvalidRequest}, + {"shares", func(v *Publication) { v.Members[0].Shares = 0 }, InvalidRequest}, + {"endpoint", func(v *Publication) { v.Members[0].PeerEndpoint = "192.0.2.1:0" }, InvalidRequest}, + {"duplicate physical NIC", func(v *Publication) { v.Members[1].RDMANICs = append(v.Members[1].RDMANICs, v.Members[1].RDMANICs[0]) }, InvalidRequest}, + {"device", func(v *Publication) { v.Members[1].RDMANICs[0].Device = "\xff" }, InvalidRequest}, + {"duplicate cache", func(v *Publication) { v.Caches = append(v.Caches, v.Caches[0]) }, InvalidRequest}, + {"socket", func(v *Publication) { v.Caches[0].ClientSocket += "x" }, InvalidRequest}, + {"member limit", func(v *Publication) { v.Members = make([]Member, MaxMembers+1) }, TooLarge}, + {"byte lower bound", func(v *Publication) { v.Members[1].RDMANICs[0].Device = strings.Repeat("x", MaxPublicationBytes) }, TooLarge}, + } { + t.Run(tc.name, func(t *testing.T) { + v, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + tc.edit(&v) + candidate, err := NewCanonicalCandidate(v) + require.ErrorIs(t, err, tc.want) + require.Equal(t, CanonicalCandidate{}, candidate) + + _, err = EncodePublication(v) + require.ErrorIs(t, err, tc.want) + }) + } + + var zero CanonicalCandidate + + p, m, err := zero.ContentHashes() + require.ErrorIs(t, err, UnsupportedVersion) + require.Empty(t, p) + require.Empty(t, m) + + b, err := zero.EncodePublication(1, 1) + require.ErrorIs(t, err, UnsupportedVersion) + require.Nil(t, b) + + candidate, err := NewCanonicalCandidate(Publication{SchemaVersion: SchemaVersion, Cluster: "11111111-1111-4111-8111-111111111111"}) + require.NoError(t, err) + + for _, counters := range [][2]uint64{{0, 0}, {0, 1}, {1, 0}} { + b, err := candidate.EncodePublication(Sequence(counters[0]), MembershipVersion(counters[1])) + require.ErrorIs(t, err, InvalidRequest) + require.Nil(t, b) + } +} + +func TestCanonicalCandidateFinalByteBound(t *testing.T) { + v := Publication{ + SchemaVersion: SchemaVersion, Cluster: "11111111-1111-4111-8111-111111111111", Sequence: 1, MembershipVersion: 1, + Members: []Member{{Node: "22222222-2222-4222-8222-222222222222", Shares: 1, PeerEndpoint: "192.0.2.1:1", RDMANICs: []RDMANIC{{Device: "x", Port: 1}}}}, + } + b, err := EncodePublication(v) + require.NoError(t, err) + // Tabs expand to two bytes but remain below the cheap input bound. + padding := MaxPublicationBytes - len(b) + v.Members[0].RDMANICs[0].Device += strings.Repeat("\t", padding/2) + strings.Repeat("x", padding%2) + candidate, err := NewCanonicalCandidate(v) + require.NoError(t, err) + _, _, err = candidate.ContentHashes() + require.NoError(t, err, "counter-free content fits") + b, err = candidate.EncodePublication(1, 1) + require.NoError(t, err) + require.Len(t, b, MaxPublicationBytes) + + for _, counters := range [][2]uint64{{10, 1}, {1, 10}, {^uint64(0), ^uint64(0)}} { + b, err := candidate.EncodePublication(Sequence(counters[0]), MembershipVersion(counters[1])) + require.ErrorIs(t, err, TooLarge) + require.Nil(t, b) + } + + v.Members[0].RDMANICs[0].Device += strings.Repeat("\t", 100) + candidate, err = NewCanonicalCandidate(v) + require.NoError(t, err) + p, m, err := candidate.ContentHashes() + require.ErrorIs(t, err, TooLarge) + require.Empty(t, p) + require.Empty(t, m) +} + +// Token validation must reject bad shapes without traversing their contents. + +func TestTokenValidationRejectsBeforeDescending(t *testing.T) { + for _, tc := range []struct { + name, prefix, suffix string + typ reflect.Type + }{ + {"unknown field", `{"unknown"`, `:[{"ignored":[]}]}`, reflect.TypeFor[Publication]()}, + {"case variant", `{"Schema_version"`, `:[{}]}`, reflect.TypeFor[Publication]()}, + {"escaped duplicate", `{"schema_version":1,"schema_versi\u006fn"`, `:[{}]}`, reflect.TypeFor[Publication]()}, + {"wrong root", `[`, `{"ignored":[]}]`, reflect.TypeFor[Publication]()}, + {"wrong collection", `{"members":{`, `"ignored":[]}}`, reflect.TypeFor[Publication]()}, + {"wrong element", `{"members":[[`, `{"ignored":[]}]]}`, reflect.TypeFor[Publication]()}, + {"wrong primitive", `{"schema_version":[`, `{"ignored":[]}]}`, reflect.TypeFor[Publication]()}, + {"wrong bytes", `{"csr_der":[`, `{"ignored":[]}]}`, reflect.TypeFor[BootstrapRequest]()}, + {"nested unknown", `[{"unknown"`, `:[{}]}]`, reflect.TypeFor[[]RDMANIC]()}, + } { + t.Run(tc.name, func(t *testing.T) { + d := json.NewDecoder(strings.NewReader(tc.prefix + tc.suffix)) + d.UseNumber() + require.ErrorIs(t, checkValue(d, tc.typ, false, 0), InvalidRequest) + // InputOffset counts consumed tokens, not decoder read-ahead. + require.Equal(t, int64(len(tc.prefix)), d.InputOffset(), "traversed rejected subtree") + }) + } +} + +func TestTokenValidationCollectionLimitsBeforeNextElement(t *testing.T) { + for _, tc := range []struct { + name, element string + typ reflect.Type + limit int + }{ + {"members", `{"node":"","shares":0,"peer_endpoint":"","rdma_nics":[],"site":""}`, reflect.TypeFor[[]Member](), MaxMembers}, + {"NICs", `{"device":"a","port":1,"rail":0}`, reflect.TypeFor[[]RDMANIC](), MaxRDMANICs}, + } { + t.Run(tc.name, func(t *testing.T) { + prefix := "[" + strings.Repeat(tc.element+",", tc.limit-1) + tc.element + for _, suffix := range []string{"]", `,{"unvisited":[{}]}]`} { + d := json.NewDecoder(strings.NewReader(prefix + suffix)) + d.UseNumber() + + err := checkValue(d, tc.typ, false, 0) + if suffix == "]" { + require.NoError(t, err, "exact limit") + continue + } + + require.ErrorIs(t, err, TooLarge) + require.Equal(t, int64(len(prefix)), d.InputOffset(), "consumed excess element") + } + }) + } +} + +func TestTokenValidationPrimitiveContract(t *testing.T) { + type primitives struct { + Count uint64 `json:"count,string"` + Byte uint8 `json:"byte"` + Flag bool `json:"flag"` + Text string `json:"text"` + Bytes []byte `json:"bytes"` + Opt *uint32 `json:"opt,omitempty"` + } + + const valid = `{"count":"18446744073709551615","byte":255,"flag":true,"text":"\ud83d\ude00","bytes":"AA=="}` + for _, tc := range []struct{ old, replacement string }{ + {`"count":"18446744073709551615"`, `"count":"18446744073709551616"`}, + {`"count":"18446744073709551615"`, `"count":1`}, + {`"count":"18446744073709551615"`, `"count":"01"`}, + {`"count":"18446744073709551615"`, `"count":"+1"`}, + {`255`, `256`}, + {`255`, `-0`}, + {`255`, `1.0`}, + {`255`, `1e0`}, + {`255`, `"1"`}, + {`true`, `1`}, + {`true`, `null`}, + {`"\ud83d\ude00"`, `"\ud83d"`}, + {`"\ud83d\ude00"`, `"\ude00"`}, + {`"\ud83d\ude00"`, `null`}, + {`"\ud83d\ude00"`, `"` + "\xff" + `"`}, + {`"AA=="`, `"AB=="`}, + {`"AA=="`, `"AA"`}, + {`"AA=="`, `"AA==\n"`}, + {`"AA=="`, `null`}, + {`"AA=="`, `[]`}, + {`"flag":true,`, ``}, + {`"bytes":"AA=="`, `"bytes":"AA==","opt":null`}, + {`"bytes":"AA=="`, `"bytes":"AA==","opt":4294967296`}, + } { + raw := strings.Replace(valid, tc.old, tc.replacement, 1) + + var v primitives + require.ErrorIs(t, decode(strings.NewReader(raw), 4096, &v), InvalidRequest, "%s", raw) + } + + for _, raw := range []string{valid, `{"bytes":"AA==","text":"\ud83d\ude00","flag":true,"byte":255,"count":"18446744073709551615","opt":0}`} { + var v primitives + require.NoError(t, decode(strings.NewReader(raw), 4096, &v)) + require.Equal(t, ^uint64(0), v.Count) + require.Equal(t, "😀", v.Text) + + for _, suffix := range []string{`{}`, `null`, `!`} { + require.ErrorIs(t, decode(strings.NewReader(raw+suffix), 4096, &v), InvalidRequest, "%s", suffix) + } + } +} + +func TestTokenValidationDepthBoundary(t *testing.T) { + for _, depth := range []int{64, 65} { + typ := reflect.TypeFor[string]() + for range depth { + typ = reflect.SliceOf(typ) + } + + raw := strings.Repeat("[", depth) + `"leaf"` + strings.Repeat("]", depth) + + err := decode(strings.NewReader(raw), 4096, reflect.New(typ).Interface()) + if depth == 64 { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, InvalidRequest) + } + } +} + +// NIC reports, persisted member migration, and Site boundaries. + +func TestRDMANICStrictFieldsAndBounds(t *testing.T) { + valid := `{"device":"mlx5_0","port":1,"rail":0}` + for _, bad := range []string{ + `{"port":1,"rail":0}`, `{"device":"a","rail":0}`, `{"device":"a","port":1}`, + strings.Replace(valid, `"port":1`, `"port":0`, 1), + strings.Replace(valid, `"port":1`, `"port":256`, 1), + strings.Replace(valid, `"port":1`, `"port":1.0`, 1), + strings.Replace(valid, `"port":1`, `"port":"1"`, 1), + strings.Replace(valid, `"port":1`, `"port":-1`, 1), + strings.Replace(valid, `"port":1`, `"port":1,"port":2`, 1), + strings.Replace(valid, `"rail":0`, `"rail":65536`, 1), + strings.Replace(valid, `"rail":0`, `"rail":0,"gid":""`, 1), + strings.Replace(valid, `"rail":0`, `"rail":0,"gid":null`, 1), + strings.Replace(valid, `"rail":0`, `"rail":0,"gid":"ABCDEF0123456789abcdef0123456789"`, 1), + strings.Replace(valid, `"rail":0`, `"rail":0,"gid":"0123"`, 1), + strings.Replace(valid, `"rail":0`, `"rail":0,"gid":"gggggggggggggggggggggggggggggggg"`, 1), + strings.Replace(valid, `"rail":0`, `"rail":0,"fabric":"old"`, 1), + } { + t.Run(bad, func(t *testing.T) { + _, err := DecodeRDMANICs(strings.NewReader("[" + bad + "]")) + require.ErrorIs(t, err, InvalidRequest) + }) + } + + nics, err := DecodeRDMANICs(strings.NewReader(`[{"device":"mlx5_0","port":255,"rail":65535,"gid":"abcdef0123456789abcdef0123456789","numa_node":4294967295}]`)) + require.NoError(t, err) + require.Equal(t, uint8(255), nics[0].Port) + + for _, count := range []int{64, 65} { + nics := make([]RDMANIC, count) + for i := range nics { + nics[i] = RDMANIC{Device: fmt.Sprintf("mlx5_%d", i), Port: 1} + } + + raw, err := json.Marshal(nics) + require.NoError(t, err) + + _, err = DecodeRDMANICs(bytes.NewReader(raw)) + if count == 64 { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, TooLarge) + } + } +} + +func TestBootstrapRDMANICRequiredCanonicalAndBounded(t *testing.T) { + request, err := DecodeBootstrap(bytes.NewReader(fixture(t, "bootstrap-request.json"))) + require.NoError(t, err) + + request.RDMANICs = []RDMANIC{{Device: "b", Port: 2, Rail: 1}, {Device: "a", Port: 1, Rail: 1}} + raw, err := EncodeBootstrapRequest(request) + require.NoError(t, err) + decoded, err := DecodeBootstrap(bytes.NewReader(raw)) + require.NoError(t, err) + require.Equal(t, "a", decoded.RDMANICs[0].Device) + require.Equal(t, "b", request.RDMANICs[0].Device) + + for _, replacement := range []string{`"rails":`, `"Rdma_nics":`} { + _, err := DecodeBootstrap(bytes.NewReader(bytes.Replace(raw, []byte(`"rdma_nics":`), []byte(replacement), 1))) + require.ErrorIs(t, err, InvalidRequest) + } + + request.RDMANICs = append(request.RDMANICs, request.RDMANICs[0]) + _, err = EncodeBootstrapRequest(request) + require.ErrorIs(t, err, InvalidRequest) + + request.RDMANICs = []RDMANIC{{Device: strings.Repeat("x", MaxBootstrapBytes), Port: 1}} + require.ErrorIs(t, ValidateBootstrapRequest(request), TooLarge) + request.RDMANICs = make([]RDMANIC, 65) + require.ErrorIs(t, ValidateBootstrapRequest(request), TooLarge) + request.RDMANICs = nil + raw, err = EncodeBootstrapRequest(request) + require.NoError(t, err) + require.Contains(t, string(raw), `"rdma_nics":[]`) + _, err = DecodeBootstrap(bytes.NewReader(bytes.Replace(raw, []byte(`"rdma_nics":[],`), nil, 1))) + require.ErrorIs(t, err, InvalidRequest) +} + +func TestAdmittedMemberMigrationAndHardBreak(t *testing.T) { + legacy := `{"node":"22222222-2222-4222-8222-222222222222","shares":4,"peer_endpoint":"192.0.2.1:7443","rails":[{"rail":0,"fabric":"old"}],"alignment_enabled":false,"site":""}` + member, err := DecodeAdmittedMember(strings.NewReader(legacy)) + require.ErrorIs(t, err, InvalidRequest) + require.Zero(t, member) + member = Member{Node: "22222222-2222-4222-8222-222222222222", Shares: 4, PeerEndpoint: "192.0.2.1:7443", RDMANICs: []RDMANIC{{Device: "a", Port: 1}}} + raw, err := json.Marshal(member) + require.NoError(t, err) + _, err = DecodeAdmittedMember(bytes.NewReader(raw)) + require.NoError(t, err) + + for _, bad := range []string{ + strings.Replace(string(raw), `"port":1`, `"port":0`, 1), + strings.Replace(string(raw), `"port":1`, `"port":1,"port":2`, 1), + strings.Replace(string(raw), `"rdma_nics":[{"device":"a","port":1,"rail":0}]`, `"rdma_nics":null`, 1), + strings.Replace(string(raw), `"device":"a"`, `"device":"a","unknown":1`, 1), + } { + _, err = DecodeAdmittedMember(strings.NewReader(bad)) + require.ErrorIs(t, err, InvalidRequest) + } + + _, err = DecodePublication(strings.NewReader(`{"schema_version":1,"cluster":"11111111-1111-4111-8111-111111111111","sequence":"1","membership_version":"1","members":[` + legacy + `],"caches":[]}`)) + require.ErrorIs(t, err, InvalidRequest) +} + +func TestMemberSiteWireValidation(t *testing.T) { + for _, site := range []string{"", "a", "0", "Site_1.a-b", strings.Repeat("a", 63)} { + t.Run("valid/"+site, func(t *testing.T) { + v, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + v.Members[0].Site = site + encoded, err := EncodePublication(v) + require.NoError(t, err) + + if site == "" { + require.Contains(t, string(encoded), `"site":""`) + } else { + require.Contains(t, string(encoded), `"rdma_nics":[],"site":"`+site+`"`) + } + + decoded, err := DecodePublication(bytes.NewReader(encoded)) + require.NoError(t, err) + require.Equal(t, v, decoded) + }) + } + + for _, site := range []string{strings.Repeat("a", 64), "-a", "a-", ".a", "a.", "_a", "a_", "a/b", "a b", "a\n", "a\x00", "β", "\xff"} { + t.Run("invalid/"+site, func(t *testing.T) { + v, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + v.Members[0].Site = site + _, err = EncodePublication(v) + require.ErrorIs(t, err, InvalidRequest) + _, err = NewCanonicalCandidate(v) + require.ErrorIs(t, err, InvalidRequest) + encoded, err := json.Marshal(v) + require.NoError(t, err) + _, err = DecodePublication(bytes.NewReader(encoded)) + require.ErrorIs(t, err, InvalidRequest) + }) + } +} + +func TestMemberSiteJSONShape(t *testing.T) { + for _, value := range []string{`null`, `1`, `true`, `[]`, `{}`, `"a","site":"b"`} { + raw := strings.Replace(string(fixture(t, "publication.json")), `"site":""`, `"site":`+value, 1) + _, err := DecodePublication(strings.NewReader(raw)) + require.ErrorIs(t, err, InvalidRequest, "%s", value) + } + + raw := string(fixture(t, "publication.json")) + v, err := DecodePublication(strings.NewReader(raw)) + require.NoError(t, err) + encoded, err := EncodePublication(v) + require.NoError(t, err) + require.Contains(t, string(encoded), `"site":""`) +} + +func TestMemberSiteDeltaValidation(t *testing.T) { + base, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + base.Sequence, base.MembershipVersion = 1, 1 + next := canonicalPublication(base) + next.Sequence, next.MembershipVersion = 2, 2 + next.Members[0].Site = "site-a" + encoded, err := EncodeDelta(base, next) + require.NoError(t, err) + applied, err := ApplyDelta(base, bytes.NewReader(encoded)) + require.NoError(t, err) + require.Equal(t, next, applied) + + for _, site := range []string{"-invalid", "site-b"} { + bad := strings.Replace(string(encoded), `"site":"site-a"`, `"site":"`+site+`"`, 1) + + _, err := ApplyDelta(base, strings.NewReader(bad)) + if site == "-invalid" { + require.ErrorIs(t, err, InvalidRequest) + } else { + require.ErrorIs(t, err, Conflict) + } + } +} + +func TestDeltaAddRemoveUpdateAndRejectedBase(t *testing.T) { + base := Publication{ + SchemaVersion: 1, Cluster: "11111111-1111-4111-8111-111111111111", Sequence: 1, MembershipVersion: 1, + Members: []Member{ + {Node: "22222222-2222-4222-8222-222222222222", Shares: 4, PeerEndpoint: "127.0.0.1:7443", RDMANICs: []RDMANIC{}}, + {Node: "33333333-3333-4333-8333-333333333333", Shares: 4, PeerEndpoint: "127.0.0.2:7443", RDMANICs: []RDMANIC{}}, + }, Caches: []CacheDefinition{}, + } + next := base + next.Sequence, next.MembershipVersion = 2, 2 + next.Members = []Member{base.Members[0], {Node: "44444444-4444-4444-8444-444444444444", Shares: 7, PeerEndpoint: "127.0.0.4:7443", RDMANICs: []RDMANIC{}}} + next.Members[0].Shares = 9 + encoded, err := EncodeDelta(base, next) + require.NoError(t, err) + + if os.Getenv("RACER_UPDATE_SITE_VECTORS") == "1" { + require.NoError(t, os.WriteFile("testdata/delta.json", append(bytes.Clone(encoded), '\n'), 0o644)) + } + + golden, err := os.ReadFile("testdata/delta.json") + require.NoError(t, err) + require.Equal(t, bytes.TrimSpace(golden), encoded) + got, err := ApplyDelta(base, bytes.NewReader(encoded)) + require.NoError(t, err) + a, err := EncodePublication(got) + require.NoError(t, err) + b, err := EncodePublication(next) + require.NoError(t, err) + require.Equal(t, b, a) + + _, err = ApplyDelta(next, bytes.NewReader(encoded)) + require.Error(t, err, "replay accepted") + + bad := strings.Replace(string(encoded), `"shares":9`, `"shares":8`, 1) + _, err = ApplyDelta(base, strings.NewReader(bad)) + require.Error(t, err, "bad target hash accepted") + + wrong := base + wrong.Sequence = 3 + _, err = ApplyDelta(wrong, bytes.NewReader(encoded)) + require.Error(t, err, "wrong base accepted") +} + +func TestDeltaValidationFailures(t *testing.T) { + base, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + base.Sequence, base.MembershipVersion = 1, 2 + next := canonicalPublication(base) + next.Sequence = 2 + encoded, err := EncodeDelta(base, next) + require.NoError(t, err) + + for _, tc := range []struct { + name string + edit func(*Delta) + want error + }{ + {"version", func(d *Delta) { d.DeltaVersion++ }, Conflict}, + {"cluster", func(d *Delta) { d.Cluster = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" }, Conflict}, + {"sequence", func(d *Delta) { d.Sequence = base.Sequence }, Conflict}, + {"membership rollback", func(d *Delta) { d.MembershipVersion = 1 }, Conflict}, + {"unknown removal", func(d *Delta) { d.RemoveMembers = []NodeID{"aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"} }, InvalidRequest}, + {"duplicate removal", func(d *Delta) { d.RemoveMembers = []NodeID{base.Members[0].Node, base.Members[0].Node} }, InvalidRequest}, + {"duplicate upsert", func(d *Delta) { d.UpsertMembers = []Member{base.Members[0], base.Members[0]} }, InvalidRequest}, + {"remove and upsert", func(d *Delta) { + d.RemoveMembers = []NodeID{base.Members[0].Node} + d.UpsertMembers = []Member{base.Members[0]} + }, InvalidRequest}, + } { + t.Run(tc.name, func(t *testing.T) { + var delta Delta + require.NoError(t, json.Unmarshal(encoded, &delta)) + tc.edit(&delta) + raw, err := json.Marshal(delta) + require.NoError(t, err) + got, err := ApplyDelta(base, bytes.NewReader(raw)) + require.ErrorIs(t, err, tc.want) + require.Equal(t, Publication{}, got, "no partial state on error") + }) + } + + for _, tc := range []struct { + name string + base, next Publication + want error + }{ + {"invalid base", Publication{}, next, UnsupportedVersion}, + {"invalid next", base, Publication{}, UnsupportedVersion}, + {"no advance", base, base, Conflict}, + } { + t.Run(tc.name, func(t *testing.T) { + raw, err := EncodeDelta(tc.base, tc.next) + require.ErrorIs(t, err, tc.want) + require.Nil(t, raw) + }) + } + + got, err := ApplyDelta(Publication{}, bytes.NewReader(encoded)) + require.ErrorIs(t, err, UnsupportedVersion) + require.Equal(t, Publication{}, got) + got, err = ApplyDelta(base, strings.NewReader("{")) + require.ErrorIs(t, err, InvalidRequest) + require.Equal(t, Publication{}, got) +} + +func TestUnicodeEscapeValidation(t *testing.T) { + for _, tc := range []struct { + raw string + valid bool + }{ + {`"\u0041"`, true}, + {`"\ud800\udc00"`, true}, + {`"\udbff\udfff"`, true}, + {`"\\ud800"`, true}, + {`"\"\u0041"`, true}, + {`"\`, false}, + {`"\u12`, false}, + {`"\uxxxx"`, false}, + {`"\ud800\u0041"`, false}, + {`"\ud800\uxxxx"`, false}, + {`"\ud800abcdef"`, false}, + {`"\udfff"`, false}, + } { + t.Run(tc.raw, func(t *testing.T) { + require.Equal(t, tc.valid, validSurrogates([]byte(tc.raw))) + }) + } +} + +type failingReader struct{} + +func (failingReader) Read([]byte) (int, error) { return 0, io.ErrUnexpectedEOF } + +func TestAdmittedMemberReadFailures(t *testing.T) { + for _, tc := range []struct { + name string + reader io.Reader + want error + }{ + {"read error", failingReader{}, InvalidRequest}, + {"over bound", &repeatedReader{remaining: MaxBootstrapBytes + 1}, TooLarge}, + {"malformed JSON", strings.NewReader("{"), InvalidRequest}, + {"invalid legacy", strings.NewReader(`{"rails":[]}`), InvalidRequest}, + } { + t.Run(tc.name, func(t *testing.T) { + got, err := DecodeAdmittedMember(tc.reader) + require.ErrorIs(t, err, tc.want) + require.Equal(t, Member{}, got) + }) + } + + _, err := DecodeBootstrap(failingReader{}) + require.ErrorIs(t, err, InvalidRequest) +} + +// Shared Site vectors and opt-in regeneration. Ordinary tests never write goldens. + +type siteVector struct { + Name string `json:"name"` + Site string `json:"site"` + Publication string `json:"publication"` + Content string `json:"content"` + Membership string `json:"membership"` + ContentHash string `json:"content_hash"` + MembershipHash string `json:"membership_hash"` + Delta string `json:"delta,omitempty"` +} + +// Regenerate with RACER_UPDATE_SITE_VECTORS=1 and a bounded Go test command +// selecting TestGenerateSharedSiteVectors. Rust consumes the same output. +func TestGenerateSharedSiteVectors(t *testing.T) { + if os.Getenv("RACER_UPDATE_SITE_VECTORS") != "1" { + t.Skip("set RACER_UPDATE_SITE_VECTORS=1 to regenerate shared Site vectors") + } + + regenerateSharedVectors(t) + p, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + var ( + previous Publication + vectors []siteVector + ) + + for i, step := range siteSteps { + p.Sequence, p.MembershipVersion = Sequence(i+1), MembershipVersion(i+1) + p.Members[0].Site = step.site + v := siteVector{Name: step.name, Site: step.site} + publication, err := EncodePublication(p) + require.NoError(t, err) + + v.Publication = string(publication) + candidate, err := NewCanonicalCandidate(p) + require.NoError(t, err) + content, membership, err := candidate.canonicalContent() + require.NoError(t, err) + + v.Content, v.Membership = string(content), string(membership) + v.ContentHash, v.MembershipHash, err = candidate.ContentHashes() + require.NoError(t, err) + + if i > 0 { + delta, err := EncodeDelta(previous, p) + require.NoError(t, err) + + v.Delta = string(delta) + } + + vectors = append(vectors, v) + previous = canonicalPublication(p) + } + + encoded, err := json.MarshalIndent(vectors, "", " ") + require.NoError(t, err) + require.NoError(t, os.WriteFile("testdata/site-vectors.json", append(encoded, '\n'), 0o644)) +} + +var siteSteps = []struct{ name, site string }{ + {"absent", ""}, + {"added", "Site_1.west-2"}, + {"changed", "Site_2.east-1"}, + {"removed", ""}, +} + +func TestSharedSiteVectors(t *testing.T) { + var vectors []siteVector + require.NoError(t, json.Unmarshal(fixture(t, "site-vectors.json"), &vectors)) + require.Len(t, vectors, 4) + legacy, err := DecodePublication(bytes.NewReader(fixture(t, "publication.json"))) + require.NoError(t, err) + + var previous Publication + + for i, step := range siteSteps { + v := vectors[i] + + t.Run(step.name, func(t *testing.T) { + require.Equal(t, step.name, v.Name) + require.Equal(t, step.site, v.Site) + p, err := DecodePublication(bytes.NewBufferString(v.Publication)) + require.NoError(t, err) + + want := canonicalPublication(legacy) + want.Sequence, want.MembershipVersion = Sequence(i+1), MembershipVersion(i+1) + want.Members[0].Site = step.site + require.Equal(t, want, p, "only Site and counters change") + encoded, err := EncodePublication(p) + require.NoError(t, err) + require.Equal(t, []byte(v.Publication), encoded) + + if step.site == "" { + require.Contains(t, string(encoded), `"site":""`) + } else { + require.Contains(t, string(encoded), `"rdma_nics":[],"site":"`+step.site+`"`) + } + + candidate, err := NewCanonicalCandidate(p) + require.NoError(t, err) + content, membership, err := candidate.canonicalContent() + require.NoError(t, err) + require.Equal(t, []byte(v.Content), content) + require.Equal(t, []byte(v.Membership), membership) + + ph, mh, err := ContentHashes(p) + require.NoError(t, err) + require.Equal(t, v.ContentHash, ph) + require.Equal(t, v.MembershipHash, mh) + + if i == 0 { + require.Empty(t, v.Delta) + return + } + + require.NotEqual(t, vectors[i-1].ContentHash, ph) + require.NotEqual(t, vectors[i-1].MembershipHash, mh) + + delta, err := EncodeDelta(previous, p) + require.NoError(t, err) + require.Equal(t, []byte(v.Delta), delta) + applied, err := ApplyDelta(previous, bytes.NewBufferString(v.Delta)) + require.NoError(t, err) + require.Equal(t, p, applied) + + var d Delta + require.NoError(t, json.Unmarshal([]byte(v.Delta), &d)) + require.Equal(t, []Member{p.Members[0]}, d.UpsertMembers) + require.Empty(t, d.RemoveMembers, "removing Site replaces the member, not the node") + d.UpsertMembers[0].Site = "tampered-site" + bad, err := json.Marshal(d) + require.NoError(t, err) + _, err = ApplyDelta(previous, bytes.NewReader(bad)) + require.ErrorIs(t, err, Conflict, "Site is bound to the target hash") + + wrongBase := canonicalPublication(previous) + wrongBase.Members[0].Site = "stale-site" + _, err = ApplyDelta(wrongBase, bytes.NewBufferString(v.Delta)) + require.ErrorIs(t, err, Conflict, "Site is bound to the base hash") + _, err = ApplyDelta(p, bytes.NewBufferString(v.Delta)) + require.ErrorIs(t, err, Conflict, "replay rejected") + }) + + previous, err = DecodePublication(bytes.NewBufferString(v.Publication)) + require.NoError(t, err) + } + + require.Equal(t, vectors[0].Content, vectors[3].Content) + require.Equal(t, vectors[0].Membership, vectors[3].Membership) + require.Equal(t, vectors[0].ContentHash, vectors[3].ContentHash) + require.Equal(t, vectors[0].MembershipHash, vectors[3].MembershipHash) +} + +// Only public fixture inputs are persisted, never private certificate keys. +func regenerateSharedVectors(t *testing.T) { + t.Helper() + + var p Publication + require.NoError(t, json.Unmarshal(fixture(t, "publication.json"), &p)) + // Preserve unsorted members/NICs and escaped text when upgrading old input. + var old struct { + Members []struct { + Rails []struct { + Rail uint16 + Fabric string + NUMANode *uint32 `json:"numa_node"` + } + } + } + require.NoError(t, json.Unmarshal(fixture(t, "publication.json"), &old)) + + for i := range p.Members { + if p.Members[i].RDMANICs == nil { + p.Members[i].RDMANICs = []RDMANIC{} + for _, rail := range old.Members[i].Rails { + p.Members[i].RDMANICs = append(p.Members[i].RDMANICs, RDMANIC{Device: rail.Fabric, Port: 1, Rail: rail.Rail, NUMANode: rail.NUMANode}) + } + } + } + + b, err := json.Marshal(p) + require.NoError(t, err) + writeSharedVector(t, "publication.json", b) + + candidate, err := NewCanonicalCandidate(p) + require.NoError(t, err) + c, m, err := candidate.canonicalContent() + require.NoError(t, err) + writeSharedVector(t, "content.json", c) + writeSharedVector(t, "membership.json", m) + + ph, mh, err := candidate.ContentHashes() + require.NoError(t, err) + b, err = json.Marshal(map[string]string{"content": ph, "membership": mh}) + require.NoError(t, err) + writeSharedVector(t, "hashes.json", b) + + var request BootstrapRequest + require.NoError(t, json.Unmarshal(fixture(t, "bootstrap-request.json"), &request)) + request.Shares = DefaultShares + b, err = EncodeBootstrapRequest(request) + require.NoError(t, err) + writeSharedVector(t, "bootstrap-request.json", b) + regenerateRejections(t) + writeSharedVector(t, "bootstrap-response.json", fixture(t, "bootstrap-response.json")) + + var bundle bundleJSON + require.NoError(t, json.Unmarshal(fixture(t, "bundle.json"), &bundle)) + + for i := range bundle.CacheKeys { + id := make([]byte, 16) + copy(id, "RKG1") + binary.BigEndian.PutUint64(id[4:12], uint64(i%2+1)) + bundle.CacheKeys[i].ID = id + // Deterministic public test material must be unique across refs too. + bundle.CacheKeys[i].Material = bytes.Repeat([]byte{byte(i)}, 32) + } + + b, err = json.Marshal(bundle) + require.NoError(t, err) + _, err = DecodeBundle(bytes.NewReader(b)) + require.NoError(t, err) + writeSharedVector(t, "bundle.json", b) +} + +func writeSharedVector(t *testing.T, name string, b []byte) { + t.Helper() + require.NoError(t, os.WriteFile("testdata/"+name, append(b, '\n'), 0o644)) + + if name == "publication.json" || name == "bootstrap-request.json" || name == "bootstrap-response.json" || name == "bundle.json" { + // Rust reads this directory directly. Remove old generated copies too. + err := os.Remove("../../../cmd/racer-dataplane/src/control/testdata/" + name) + if !os.IsNotExist(err) { + require.NoError(t, err) + } + } +} + +type rejectionVector struct { + Name string `json:"name"` + File string `json:"file"` + Old string `json:"old"` + New string `json:"new"` + Code ErrorCode `json:"code"` +} + +func regenerateRejections(t *testing.T) { + t.Helper() + + var rejections []rejectionVector + require.NoError(t, json.Unmarshal(fixture(t, "rejections.json"), &rejections)) + + for i := range rejections { + switch rejections[i].Name { + case "missing required field": + rejections[i].Old, rejections[i].New = `"rdma_nics":[]`, `"ignored_nics":[]` + case "duplicate rail": + rejections[i].Name = "zero NIC port" + rejections[i].Old, rejections[i].New = `"port":1`, `"port":0` + case "empty fabric": + rejections[i].Name = "empty device" + } + } + + for _, addition := range []struct{ name, file, old, new string }{ + {"missing bootstrap NIC report", "bootstrap-request.json", `"rdma_nics":[],`, ""}, + {"null bootstrap NIC report", "bootstrap-request.json", `"rdma_nics":[]`, `"rdma_nics":null`}, + {"legacy bootstrap rails", "bootstrap-request.json", `"rdma_nics":[]`, `"rails":[]`}, + {"legacy member rails", "publication.json", `"rdma_nics":[]`, `"rails":[],"alignment_enabled":true`}, + {"overflow NIC port", "publication.json", `"port":1`, `"port":256`}, + {"null NIC GID", "publication.json", `"port":1`, `"port":1,"gid":null`}, + {"empty NIC GID", "publication.json", `"port":1`, `"port":1,"gid":""`}, + {"uppercase NIC GID", "publication.json", `"port":1`, `"port":1,"gid":"ABCDEF0123456789abcdef0123456789"`}, + } { + found := slices.ContainsFunc(rejections, func(rejection rejectionVector) bool { return rejection.Name == addition.name }) + if !found { + rejections = append(rejections, rejectionVector{addition.name, addition.file, addition.old, addition.new, InvalidRequest}) + } + } + + b, err := json.MarshalIndent(rejections, "", " ") + require.NoError(t, err) + writeSharedVector(t, "rejections.json", b) +}