From 69eef9f560086417920d7b37d092d180df987728 Mon Sep 17 00:00:00 2001 From: Aditya Choudhari Date: Tue, 4 Aug 2026 11:31:58 -0400 Subject: [PATCH 1/4] feat: Add one-way managed spec cutover --- internal/controller/managed_spec.go | 213 +++++++++++++ internal/controller/managed_spec_test.go | 296 ++++++++++++++++++ .../controller/weightsandbiases_controller.go | 83 ++++- .../weightsandbiases_controller_test.go | 107 +++++++ 4 files changed, 687 insertions(+), 12 deletions(-) create mode 100644 internal/controller/managed_spec.go create mode 100644 internal/controller/managed_spec_test.go diff --git a/internal/controller/managed_spec.go b/internal/controller/managed_spec.go new file mode 100644 index 00000000..47eb5b67 --- /dev/null +++ b/internal/controller/managed_spec.go @@ -0,0 +1,213 @@ +package controller + +import ( + "context" + "encoding/json" + "fmt" + "reflect" + + "github.com/wandb/operator/pkg/wandb/spec" + "github.com/wandb/operator/pkg/wandb/spec/charts" + corev1 "k8s.io/api/core/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "sigs.k8s.io/controller-runtime/pkg/client" + ctrllog "sigs.k8s.io/controller-runtime/pkg/log" +) + +const ( + managedSpecConfigMapName = "wandb-spec-managed" + managedSpecStateConfigMapName = "wandb-managed-spec-state" + managedSpecStateKey = "managed" +) + +type managedSpecSource struct { + spec *spec.Spec + rawChart interface{} + rawValues map[string]interface{} +} + +func (r *WeightsAndBiasesReconciler) selectBaseSpec( + ctx context.Context, + namespace string, + getDeployerSpec func() (*spec.Spec, error), +) (*spec.Spec, bool, error) { + log := ctrllog.FromContext(ctx) + managed, err := r.managedSpecEnabled(ctx, namespace) + if err != nil { + return nil, false, err + } + if managed { + log.Info("Managed spec cutover is active; skipping Deployer") + managedSpec, err := r.getManagedSpec(ctx, namespace) + if err != nil { + return nil, false, err + } + return managedSpec.spec, false, nil + } + + deployerSpec, err := getDeployerSpec() + if err != nil { + return nil, false, err + } + + managedSpec, err := r.getManagedSpec(ctx, namespace) + if apierrors.IsNotFound(err) { + return deployerSpec, false, nil + } + if err != nil { + log.Info("Managed spec is invalid; continuing with Deployer", "error", err) + return deployerSpec, false, nil + } + matches, err := managedSpecConfigurationMatches(managedSpec, deployerSpec) + if err != nil { + log.Info("Managed spec could not be compared; continuing with Deployer", "error", err) + return deployerSpec, false, nil + } + if matches { + log.Info("Managed spec matches Deployer; cutover is pending successful apply") + return managedSpec.spec, true, nil + } + + log.Info("Managed spec does not match Deployer; continuing with Deployer") + return deployerSpec, false, nil +} + +func (r *WeightsAndBiasesReconciler) managedSpecEnabled(ctx context.Context, namespace string) (bool, error) { + state := &corev1.ConfigMap{} + err := r.Get(ctx, client.ObjectKey{Name: managedSpecStateConfigMapName, Namespace: namespace}, state) + if apierrors.IsNotFound(err) { + return false, nil + } + if err != nil { + return false, err + } + + value, ok := state.Data[managedSpecStateKey] + if !ok { + return false, nil + } + if value == "true" { + return true, nil + } + return false, fmt.Errorf("invalid %s value %q in ConfigMap %s", managedSpecStateKey, value, managedSpecStateConfigMapName) +} + +func (r *WeightsAndBiasesReconciler) setManagedSpecEnabled(ctx context.Context, namespace string) error { + key := client.ObjectKey{Name: managedSpecStateConfigMapName, Namespace: namespace} + state := &corev1.ConfigMap{} + err := r.Get(ctx, key, state) + if apierrors.IsNotFound(err) { + return r.Create(ctx, &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: key.Name, Namespace: key.Namespace}, + Data: map[string]string{managedSpecStateKey: "true"}, + }) + } + if err != nil { + return err + } + if state.Data == nil { + state.Data = make(map[string]string) + } + if state.Data[managedSpecStateKey] == "true" { + return nil + } + state.Data[managedSpecStateKey] = "true" + return r.Update(ctx, state) +} + +func (r *WeightsAndBiasesReconciler) getManagedSpec(ctx context.Context, namespace string) (*managedSpecSource, error) { + configMap := &corev1.ConfigMap{} + key := client.ObjectKey{Name: managedSpecConfigMapName, Namespace: namespace} + if err := r.Get(ctx, key, configMap); err != nil { + return nil, err + } + + valuesJSON, ok := configMap.Data["values"] + if !ok { + return nil, fmt.Errorf("ConfigMap %s/%s does not have a values key", namespace, managedSpecConfigMapName) + } + rawValues := map[string]interface{}{} + if err := json.Unmarshal([]byte(valuesJSON), &rawValues); err != nil { + return nil, fmt.Errorf("decode values from ConfigMap %s/%s: %w", namespace, managedSpecConfigMapName, err) + } + + chartJSON, ok := configMap.Data["chart"] + if !ok { + return nil, fmt.Errorf("ConfigMap %s/%s does not have a chart key", namespace, managedSpecConfigMapName) + } + var rawChart interface{} + if err := json.Unmarshal([]byte(chartJSON), &rawChart); err != nil { + return nil, fmt.Errorf("decode chart from ConfigMap %s/%s: %w", namespace, managedSpecConfigMapName, err) + } + chart := charts.Get(rawChart) + if chart == nil { + return nil, fmt.Errorf("ConfigMap %s/%s contains an unsupported chart", namespace, managedSpecConfigMapName) + } + + return &managedSpecSource{ + spec: &spec.Spec{Chart: chart, Values: spec.Values(rawValues)}, + rawChart: rawChart, + rawValues: rawValues, + }, nil +} + +func managedSpecConfigurationMatches(managed *managedSpecSource, deployer *spec.Spec) (bool, error) { + if managed == nil || deployer == nil { + return false, nil + } + + deployerChart, err := normalizeJSONValue(deployer.Chart) + if err != nil { + return false, fmt.Errorf("normalize Deployer chart: %w", err) + } + deployerValues, err := normalizeJSONValue(deployer.Values) + if err != nil { + return false, fmt.Errorf("normalize Deployer values: %w", err) + } + + return managedJSONSubsetEqual(managed.rawChart, deployerChart) && + managedJSONSubsetEqual(managed.rawValues, deployerValues), nil +} + +func normalizeJSONValue(value interface{}) (interface{}, error) { + data, err := json.Marshal(value) + if err != nil { + return nil, err + } + var normalized interface{} + if err := json.Unmarshal(data, &normalized); err != nil { + return nil, err + } + return normalized, nil +} + +func managedJSONSubsetEqual(managed, deployer interface{}) bool { + switch managedValue := managed.(type) { + case map[string]interface{}: + deployerValue, ok := deployer.(map[string]interface{}) + if !ok { + return false + } + for key, managedChild := range managedValue { + deployerChild, ok := deployerValue[key] + if !ok || !managedJSONSubsetEqual(managedChild, deployerChild) { + return false + } + } + return true + case []interface{}: + deployerValue, ok := deployer.([]interface{}) + if !ok || len(managedValue) != len(deployerValue) { + return false + } + for index, managedChild := range managedValue { + if !managedJSONSubsetEqual(managedChild, deployerValue[index]) { + return false + } + } + return true + default: + return reflect.DeepEqual(managed, deployer) + } +} diff --git a/internal/controller/managed_spec_test.go b/internal/controller/managed_spec_test.go new file mode 100644 index 00000000..aabcb874 --- /dev/null +++ b/internal/controller/managed_spec_test.go @@ -0,0 +1,296 @@ +package controller + +import ( + "context" + "errors" + "reflect" + "testing" + + appsv1 "github.com/wandb/operator/api/v1" + "github.com/wandb/operator/pkg/wandb/spec" + "github.com/wandb/operator/pkg/wandb/spec/charts" + corev1 "k8s.io/api/core/v1" + 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" +) + +func TestManagedSpecSelection(t *testing.T) { + ctx := context.Background() + namespace := "default" + deployerSpec := testManagedSpec(map[string]interface{}{ + "global": map[string]interface{}{"enabled": true}, + }) + + t.Run("uses Deployer when the managed spec does not exist", func(t *testing.T) { + reconciler := testManagedSpecReconciler(t) + calls := 0 + + selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + calls++ + return deployerSpec, nil + }) + + if err != nil { + t.Fatalf("selectBaseSpec returned an error: %v", err) + } + if selected != deployerSpec { + t.Fatal("selectBaseSpec did not return the Deployer spec") + } + if pendingCutover { + t.Fatal("cutover must not be pending without a managed spec") + } + if calls != 1 { + t.Fatalf("Deployer was called %d times, want 1", calls) + } + }) + + t.Run("uses Deployer when the managed spec differs", func(t *testing.T) { + reconciler := testManagedSpecReconciler(t, testManagedSpecConfigMap(namespace, map[string]interface{}{ + "global": map[string]interface{}{"enabled": false}, + })) + + selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + return deployerSpec, nil + }) + + if err != nil { + t.Fatalf("selectBaseSpec returned an error: %v", err) + } + if selected != deployerSpec { + t.Fatal("selectBaseSpec did not return the Deployer spec") + } + if pendingCutover { + t.Fatal("cutover must not be pending for a mismatched managed spec") + } + }) + + t.Run("selects matching managed-owned configuration and requests cutover", func(t *testing.T) { + reconciler := testManagedSpecReconciler(t, testManagedSpecConfigMap(namespace, deployerSpec.Values)) + + selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + withMetadata := *testManagedSpec(map[string]interface{}{ + "global": map[string]interface{}{ + "enabled": true, + "image": map[string]interface{}{"tag": "deployer-only"}, + }, + "legacy": map[string]interface{}{"enabled": true}, + }) + metadata := spec.Metadata{"releaseId": "release-1"} + withMetadata.Metadata = &metadata + withMetadata.Chart.(*charts.RepoRelease).Debug = true + return &withMetadata, nil + }) + + if err != nil { + t.Fatalf("selectBaseSpec returned an error: %v", err) + } + if !pendingCutover { + t.Fatal("matching managed configuration must request cutover") + } + if selected == nil || selected.Metadata != nil { + t.Fatal("selectBaseSpec did not return the managed ConfigMap spec") + } + if !reflect.DeepEqual(selected.Values, deployerSpec.Values) { + t.Fatal("selectBaseSpec did not preserve the managed-owned values") + } + }) + + t.Run("uses managed configuration without calling Deployer after cutover", func(t *testing.T) { + reconciler := testManagedSpecReconciler( + t, + testManagedSpecConfigMap(namespace, deployerSpec.Values), + &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecStateConfigMapName, Namespace: namespace}, + Data: map[string]string{managedSpecStateKey: "true"}, + }, + ) + calls := 0 + + selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + calls++ + return nil, errors.New("Deployer must not be called") + }) + + if err != nil { + t.Fatalf("selectBaseSpec returned an error: %v", err) + } + if pendingCutover { + t.Fatal("cutover cannot be pending after it is active") + } + if selected == nil || !selected.IsEqual(deployerSpec) { + t.Fatal("selectBaseSpec did not return the managed spec") + } + if calls != 0 { + t.Fatalf("Deployer was called %d times after cutover, want 0", calls) + } + }) + + t.Run("fails closed when managed configuration is missing after cutover", func(t *testing.T) { + reconciler := testManagedSpecReconciler(t, &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecStateConfigMapName, Namespace: namespace}, + Data: map[string]string{managedSpecStateKey: "true"}, + }) + calls := 0 + + selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + calls++ + return deployerSpec, nil + }) + + if err == nil { + t.Fatal("selectBaseSpec succeeded without the managed spec after cutover") + } + if selected != nil || pendingCutover { + t.Fatal("selectBaseSpec returned a spec while failing closed") + } + if calls != 0 { + t.Fatalf("Deployer was called %d times after cutover, want 0", calls) + } + }) + + for _, value := range []string{"false", "1"} { + t.Run("rejects managed state "+value+" without calling Deployer", func(t *testing.T) { + reconciler := testManagedSpecReconciler(t, &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecStateConfigMapName, Namespace: namespace}, + Data: map[string]string{managedSpecStateKey: value}, + }) + calls := 0 + + _, _, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + calls++ + return deployerSpec, nil + }) + + if err == nil { + t.Fatalf("selectBaseSpec accepted managed state %q", value) + } + if calls != 0 { + t.Fatalf("Deployer was called %d times with invalid managed state, want 0", calls) + } + }) + } +} + +func TestSetManagedSpecEnabled(t *testing.T) { + ctx := context.Background() + namespace := "default" + reconciler := testManagedSpecReconciler(t) + + if err := reconciler.setManagedSpecEnabled(ctx, namespace); err != nil { + t.Fatalf("setManagedSpecEnabled returned an error: %v", err) + } + + state := &corev1.ConfigMap{} + key := client.ObjectKey{Name: managedSpecStateConfigMapName, Namespace: namespace} + if err := reconciler.Get(ctx, key, state); err != nil { + t.Fatalf("could not read managed spec state: %v", err) + } + if state.Data[managedSpecStateKey] != "true" { + t.Fatalf("managed state is %q, want true", state.Data[managedSpecStateKey]) + } +} + +func TestManagedJSONSubsetEqual(t *testing.T) { + tests := []struct { + name string + managed interface{} + deployer interface{} + matches bool + }{ + { + name: "ignores Deployer-only object keys", + managed: map[string]interface{}{"api": map[string]interface{}{"enabled": true}}, + deployer: map[string]interface{}{"api": map[string]interface{}{"enabled": true, "tag": "latest"}}, + matches: true, + }, + { + name: "rejects a different managed-owned value", + managed: map[string]interface{}{"api": map[string]interface{}{"enabled": true}}, + deployer: map[string]interface{}{"api": map[string]interface{}{"enabled": false}}, + matches: false, + }, + { + name: "requires arrays to have the same length", + managed: []interface{}{map[string]interface{}{"name": "first"}}, + deployer: []interface{}{map[string]interface{}{"name": "first"}, map[string]interface{}{"name": "second"}}, + matches: false, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + if got := managedJSONSubsetEqual(test.managed, test.deployer); got != test.matches { + t.Fatalf("managedJSONSubsetEqual() = %t, want %t", got, test.matches) + } + }) + } +} + +func TestManagedSpecConfigMapRequests(t *testing.T) { + scheme := runtime.NewScheme() + if err := corev1.AddToScheme(scheme); err != nil { + t.Fatalf("could not register core Kubernetes types: %v", err) + } + if err := appsv1.AddToScheme(scheme); err != nil { + t.Fatalf("could not register WeightsAndBiases types: %v", err) + } + + reconciler := &WeightsAndBiasesReconciler{ + Client: fake.NewClientBuilder().WithScheme(scheme).WithObjects( + &appsv1.WeightsAndBiases{ObjectMeta: metav1.ObjectMeta{Name: "first", Namespace: "default"}}, + &appsv1.WeightsAndBiases{ObjectMeta: metav1.ObjectMeta{Name: "other", Namespace: "other"}}, + ).Build(), + Scheme: scheme, + } + + requests := reconciler.managedSpecConfigMapRequests(context.Background(), &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: "default"}, + }) + if len(requests) != 1 { + t.Fatalf("managedSpecConfigMapRequests returned %d requests, want 1", len(requests)) + } + want := types.NamespacedName{Name: "first", Namespace: "default"} + if requests[0].NamespacedName != want { + t.Fatalf("managedSpecConfigMapRequests returned %v, want %v", requests[0].NamespacedName, want) + } +} + +func testManagedSpecReconciler(t *testing.T, objects ...client.Object) *WeightsAndBiasesReconciler { + t.Helper() + scheme := runtime.NewScheme() + if err := corev1.AddToScheme(scheme); err != nil { + t.Fatalf("could not register core Kubernetes types: %v", err) + } + return &WeightsAndBiasesReconciler{ + Client: fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).Build(), + Scheme: scheme, + } +} + +func testManagedSpecConfigMap(namespace string, values map[string]interface{}) *corev1.ConfigMap { + valuesJSON := `{"global":{"enabled":true}}` + if enabled, ok := values["global"].(map[string]interface{})["enabled"].(bool); ok && !enabled { + valuesJSON = `{"global":{"enabled":false}}` + } + return &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: namespace}, + Data: map[string]string{ + "chart": `{"name":"operator-wandb","url":"https://charts.wandb.ai","version":"0.43.5"}`, + "values": valuesJSON, + }, + } +} + +func testManagedSpec(values map[string]interface{}) *spec.Spec { + return &spec.Spec{ + Chart: &charts.RepoRelease{ + Name: "operator-wandb", + URL: "https://charts.wandb.ai", + Version: "0.43.5", + }, + Values: values, + } +} diff --git a/internal/controller/weightsandbiases_controller.go b/internal/controller/weightsandbiases_controller.go index 76396eb2..ebefee03 100644 --- a/internal/controller/weightsandbiases_controller.go +++ b/internal/controller/weightsandbiases_controller.go @@ -33,8 +33,10 @@ import ( "sigs.k8s.io/controller-runtime/pkg/client" "sigs.k8s.io/controller-runtime/pkg/controller/controllerutil" "sigs.k8s.io/controller-runtime/pkg/event" + "sigs.k8s.io/controller-runtime/pkg/handler" ctrllog "sigs.k8s.io/controller-runtime/pkg/log" "sigs.k8s.io/controller-runtime/pkg/predicate" + "sigs.k8s.io/controller-runtime/pkg/reconcile" corev1 "k8s.io/api/core/v1" @@ -147,9 +149,12 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req license := utils.GetLicense(ctx, r.Client, wandb, crdSpec, userInputSpec) - var deployerSpec *spec.Spec - if !r.IsAirgapped { - deployerSpec, err = r.DeployerClient.GetSpec(deployer.GetSpecOptions{ + getDeployerSpec := func() (*spec.Spec, error) { + if r.IsAirgapped { + return nil, nil + } + + deployerSpec, err := r.DeployerClient.GetSpec(deployer.GetSpecOptions{ License: license, ActiveState: currentActiveSpec, ReleaseId: releaseID, @@ -160,12 +165,13 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req // This scenario may occur if the user disables networking, or if the deployer // is not operational, and a version has been deployed successfully. Rather than // reverting to the container defaults, we've stored the most recent successful - // deployer release in the cache - // Attempt to retrieve the cached release - if deployerSpec, err = specManager.Get("latest-cached-release"); err != nil { + // deployer release in the cache. + deployerSpec, err = specManager.Get("latest-cached-release") + if err != nil { log.Info("No cached release found", "error", err.Error()) + deployerSpec = nil } - if r.Debug { + if r.Debug && deployerSpec != nil { log.Info("Using cached deployer spec", "spec", deployerSpec.SensitiveValuesMasked()) } } @@ -177,9 +183,20 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req if err := specManager.Set("latest-cached-release", deployerSpec); err != nil { r.Recorder.Event(wandb, corev1.EventTypeNormal, "SecretWriteFailed", "Unable to write secret to kubernetes") log.Error(err, "Unable to save latest release.") - return ctrlqueue.DoNotRequeue() + return nil, err } } + return deployerSpec, nil + } + + baseSpec := currentActiveSpec + pendingManagedCutover := false + if wandb.ObjectMeta.DeletionTimestamp.IsZero() { + baseSpec, pendingManagedCutover, err = r.selectBaseSpec(ctx, wandb.Namespace, getDeployerSpec) + if err != nil { + log.Error(err, "Failed to select Deployer or managed spec") + return ctrlqueue.RequeueWithError(err) + } } desiredSpec := new(spec.Spec) @@ -205,13 +222,13 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req log.Info("Desired spec after merging userInputSpec", "spec", desiredSpec.SensitiveValuesMasked()) } - if err := desiredSpec.Merge(deployerSpec); err != nil { - log.Error(err, "Failed to merge deployer spec into desired spec") + if err := desiredSpec.Merge(baseSpec); err != nil { + log.Error(err, "Failed to merge selected base spec into desired spec") return ctrlqueue.RequeueWithError(err) } if r.Debug { - log.Info("Desired spec after merging deployerSpec", "spec", desiredSpec.SensitiveValuesMasked()) + log.Info("Desired spec after merging selected base spec", "spec", desiredSpec.SensitiveValuesMasked()) } if err := desiredSpec.Merge(operator.Defaults(wandb, r.Scheme)); err != nil { @@ -230,6 +247,13 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req log.Info("Active spec found", "spec", currentActiveSpec.SensitiveValuesMasked()) if currentActiveSpec.IsEqual(desiredSpec) { log.Info("No changes found") + if pendingManagedCutover { + if err := r.setManagedSpecEnabled(ctx, wandb.Namespace); err != nil { + log.Error(err, "Failed to persist managed spec cutover") + return ctrlqueue.RequeueWithError(err) + } + log.Info("Managed spec cutover completed") + } statusManager.Set(status.Completed) return ctrlqueue.Requeue(desiredSpec) } else { @@ -279,6 +303,13 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req if r.Debug { log.Info("Successfully saved active spec", "spec", desiredSpec.SensitiveValuesMasked()) } + if pendingManagedCutover { + if err := r.setManagedSpecEnabled(ctx, wandb.Namespace); err != nil { + log.Error(err, "Failed to persist managed spec cutover") + return ctrlqueue.RequeueWithError(err) + } + log.Info("Managed spec cutover completed") + } r.Recorder.Event(wandb, corev1.EventTypeNormal, "Completed", "Completed reconcile successfully") if err := r.discoverAndPatchResources(ctx, wandb); err != nil { @@ -373,10 +404,38 @@ func (r *WeightsAndBiasesReconciler) SetupWithManager(mgr ctrl.Manager) error { builder := ctrl.NewControllerManagedBy(mgr). For(&apiv1.WeightsAndBiases{}, builder.WithPredicates(filterWBEvents{})). Owns(&corev1.Secret{}, builder.WithPredicates(filterSecretEvents{})). - Owns(&corev1.ConfigMap{}) + Owns(&corev1.ConfigMap{}). + Watches( + &corev1.ConfigMap{}, + handler.EnqueueRequestsFromMapFunc(r.managedSpecConfigMapRequests), + builder.WithPredicates(predicate.NewPredicateFuncs(isManagedSpecConfigMap)), + ) return builder.Complete(r) } +func isManagedSpecConfigMap(object client.Object) bool { + return object.GetName() == managedSpecConfigMapName || object.GetName() == managedSpecStateConfigMapName +} + +func (r *WeightsAndBiasesReconciler) managedSpecConfigMapRequests( + ctx context.Context, + object client.Object, +) []reconcile.Request { + instances := &apiv1.WeightsAndBiasesList{} + if err := r.List(ctx, instances, client.InNamespace(object.GetNamespace())); err != nil { + ctrllog.FromContext(ctx).Error(err, "Failed to list WeightsAndBiases instances for managed spec ConfigMap") + return nil + } + + requests := make([]reconcile.Request, 0, len(instances.Items)) + for _, instance := range instances.Items { + requests = append(requests, reconcile.Request{ + NamespacedName: client.ObjectKeyFromObject(&instance), + }) + } + return requests +} + type filterWBEvents struct { predicate.Funcs } diff --git a/internal/controller/weightsandbiases_controller_test.go b/internal/controller/weightsandbiases_controller_test.go index ce38b6a0..0fee638f 100644 --- a/internal/controller/weightsandbiases_controller_test.go +++ b/internal/controller/weightsandbiases_controller_test.go @@ -2,6 +2,7 @@ package controller import ( "context" + "encoding/json" "time" . "github.com/onsi/ginkgo/v2" @@ -14,11 +15,13 @@ import ( "github.com/wandb/operator/pkg/wandb/spec/state/secrets" v1 "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" "k8s.io/client-go/kubernetes/scheme" "k8s.io/client-go/tools/record" ctrl "sigs.k8s.io/controller-runtime" + "sigs.k8s.io/controller-runtime/pkg/client" ) var deployerSpec = spec.Spec{ @@ -320,4 +323,108 @@ var _ = Describe("WeightsandbiasesController", func() { // }) // }) //}) + Describe("Managed spec cutover", Label("managed spec cutover"), func() { + const name = "test-managed-cutover" + var deployerClient *deployerfakes.FakeDeployerInterface + + BeforeEach(func() { + ctx := context.Background() + recorder = record.NewFakeRecorder(10) + deployerClient = &deployerfakes.FakeDeployerInterface{} + deployerClient.GetSpecReturns(&deployerSpec, nil) + reconciler = &WeightsAndBiasesReconciler{ + Client: k8sClient, + IsAirgapped: false, + DeployerClient: deployerClient, + Scheme: scheme.Scheme, + Recorder: recorder, + DryRun: true, + } + + wandb := &wandbcomv1.WeightsAndBiases{ + ObjectMeta: metav1.ObjectMeta{Name: name, Namespace: "default"}, + Spec: wandbcomv1.WeightsAndBiasesSpec{ + Chart: wandbcomv1.Object{Object: map[string]interface{}{}}, + Values: wandbcomv1.Object{Object: map[string]interface{}{}}, + }, + } + Expect(k8sClient.Create(ctx, wandb)).To(Succeed()) + + chartJSON, err := json.Marshal(deployerSpec.Chart) + Expect(err).NotTo(HaveOccurred()) + valuesJSON, err := json.Marshal(deployerSpec.Values) + Expect(err).NotTo(HaveOccurred()) + managedSpec := &v1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: "default"}, + Data: map[string]string{ + "chart": string(chartJSON), + "values": string(valuesJSON), + }, + } + Expect(k8sClient.Create(ctx, managedSpec)).To(Succeed()) + }) + + AfterEach(func() { + ctx := context.Background() + wandb := &wandbcomv1.WeightsAndBiases{} + key := types.NamespacedName{Name: name, Namespace: "default"} + if err := k8sClient.Get(ctx, key, wandb); err == nil { + Expect(k8sClient.Delete(ctx, wandb)).To(Succeed()) + _, err = reconciler.Reconcile(ctx, ctrl.Request{NamespacedName: key}) + Expect(err).NotTo(HaveOccurred()) + } + + objects := []client.Object{ + &v1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: "default"}}, + &v1.ConfigMap{ObjectMeta: metav1.ObjectMeta{Name: managedSpecStateConfigMapName, Namespace: "default"}}, + &v1.Secret{ObjectMeta: metav1.ObjectMeta{Name: name + "-spec-user", Namespace: "default"}}, + &v1.Secret{ObjectMeta: metav1.ObjectMeta{Name: name + "-spec-active", Namespace: "default"}}, + &v1.Secret{ObjectMeta: metav1.ObjectMeta{Name: name + "-latest-cached-release", Namespace: "default"}}, + } + for _, object := range objects { + Expect(client.IgnoreNotFound(k8sClient.Delete(ctx, object))).To(Succeed()) + } + }) + + It("persists cutover and does not call Deployer again", func() { + ctx := context.Background() + request := ctrl.Request{NamespacedName: types.NamespacedName{Name: name, Namespace: "default"}} + + _, err := reconciler.Reconcile(ctx, request) + Expect(err).NotTo(HaveOccurred()) + + state := &v1.ConfigMap{} + stateKey := types.NamespacedName{Name: managedSpecStateConfigMapName, Namespace: "default"} + Expect(k8sClient.Get(ctx, stateKey, state)).To(Succeed()) + Expect(state.Data).To(HaveKeyWithValue(managedSpecStateKey, "true")) + Expect(deployerClient.GetSpecCallCount()).To(Equal(1)) + + _, err = reconciler.Reconcile(ctx, request) + Expect(err).NotTo(HaveOccurred()) + Expect(deployerClient.GetSpecCallCount()).To(Equal(1)) + }) + + It("does not let a missing managed spec block deletion", func() { + ctx := context.Background() + key := types.NamespacedName{Name: name, Namespace: "default"} + request := ctrl.Request{NamespacedName: key} + + _, err := reconciler.Reconcile(ctx, request) + Expect(err).NotTo(HaveOccurred()) + + managedSpec := &v1.ConfigMap{} + Expect(k8sClient.Get(ctx, types.NamespacedName{ + Name: managedSpecConfigMapName, Namespace: "default", + }, managedSpec)).To(Succeed()) + Expect(k8sClient.Delete(ctx, managedSpec)).To(Succeed()) + + wandb := &wandbcomv1.WeightsAndBiases{} + Expect(k8sClient.Get(ctx, key, wandb)).To(Succeed()) + Expect(k8sClient.Delete(ctx, wandb)).To(Succeed()) + + _, err = reconciler.Reconcile(ctx, request) + Expect(err).NotTo(HaveOccurred()) + Expect(apierrors.IsNotFound(k8sClient.Get(ctx, key, &wandbcomv1.WeightsAndBiases{}))).To(BeTrue()) + }) + }) }) From 7f7990aa47c2e1886f52234002725874131f65ec Mon Sep 17 00:00:00 2001 From: Aditya Choudhari Date: Tue, 18 Aug 2026 16:28:11 -0400 Subject: [PATCH 2/4] feat: gate managed spec cutover behind environment flag --- cmd/main.go | 16 ++++--- internal/controller/managed_spec.go | 5 ++ internal/controller/managed_spec_test.go | 47 +++++++++++++++++-- .../controller/weightsandbiases_controller.go | 27 +++++++---- .../weightsandbiases_controller_test.go | 13 ++--- 5 files changed, 82 insertions(+), 26 deletions(-) diff --git a/cmd/main.go b/cmd/main.go index d69ac631..0ff5a279 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -69,7 +69,7 @@ func main() { var enableHTTP2 bool var tlsOpts []func(*tls.Config) var deployerAPI, isolationNamespaces string - var debug, airgapped bool + var debug, airgapped, managedSpecEnabled bool flag.StringVar(&metricsAddr, "metrics-bind-address", "0", "The address the metrics endpoint binds to. "+ "Use :8443 for HTTPS or :8080 for HTTP, or leave as 0 to disable the metrics service.") @@ -94,6 +94,7 @@ func main() { flag.StringVar(&isolationNamespaces, "isolation-namespaces", "", "Specify namespaces (as a comma separated string) that the controller should monitor when operating in namespace isolation mode.") flag.BoolVar(&debug, "debug", false, "Enable debug mode") + flag.BoolVar(&managedSpecEnabled, "managed-spec-enabled", false, "Enable managed spec cutover") opts := zap.Options{ Development: true, @@ -229,12 +230,13 @@ func main() { } if err = (&controller.WeightsAndBiasesReconciler{ - IsAirgapped: airgapped, - Recorder: mgr.GetEventRecorderFor("weightsandbiases"), - Client: mgr.GetClient(), - Scheme: mgr.GetScheme(), - DeployerClient: &deployer.DeployerClient{DeployerAPI: deployerAPI}, - Debug: debug, + IsAirgapped: airgapped, + Recorder: mgr.GetEventRecorderFor("weightsandbiases"), + Client: mgr.GetClient(), + Scheme: mgr.GetScheme(), + DeployerClient: &deployer.DeployerClient{DeployerAPI: deployerAPI}, + Debug: debug, + ManagedSpecEnabled: managedSpecEnabled, }).SetupWithManager(mgr); err != nil { setupLog.Error(err, "unable to create controller", "controller", "WeightsAndBiases") os.Exit(1) diff --git a/internal/controller/managed_spec.go b/internal/controller/managed_spec.go index 47eb5b67..b1ae8c8b 100644 --- a/internal/controller/managed_spec.go +++ b/internal/controller/managed_spec.go @@ -32,6 +32,11 @@ func (r *WeightsAndBiasesReconciler) selectBaseSpec( namespace string, getDeployerSpec func() (*spec.Spec, error), ) (*spec.Spec, bool, error) { + if !r.ManagedSpecEnabled { + deployerSpec, err := getDeployerSpec() + return deployerSpec, false, err + } + log := ctrllog.FromContext(ctx) managed, err := r.managedSpecEnabled(ctx, namespace) if err != nil { diff --git a/internal/controller/managed_spec_test.go b/internal/controller/managed_spec_test.go index aabcb874..44af2e2d 100644 --- a/internal/controller/managed_spec_test.go +++ b/internal/controller/managed_spec_test.go @@ -24,6 +24,37 @@ func TestManagedSpecSelection(t *testing.T) { "global": map[string]interface{}{"enabled": true}, }) + t.Run("uses Deployer when managed spec is disabled even after cutover", func(t *testing.T) { + reconciler := testManagedSpecReconciler( + t, + testManagedSpecConfigMap(namespace, deployerSpec.Values), + &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecStateConfigMapName, Namespace: namespace}, + Data: map[string]string{managedSpecStateKey: "true"}, + }, + ) + reconciler.ManagedSpecEnabled = false + calls := 0 + + selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + calls++ + return deployerSpec, nil + }) + + if err != nil { + t.Fatalf("selectBaseSpec returned an error: %v", err) + } + if selected != deployerSpec { + t.Fatal("selectBaseSpec did not return the Deployer spec while managed spec was disabled") + } + if pendingCutover { + t.Fatal("cutover must not be pending while managed spec is disabled") + } + if calls != 1 { + t.Fatalf("Deployer was called %d times, want 1", calls) + } + }) + t.Run("uses Deployer when the managed spec does not exist", func(t *testing.T) { reconciler := testManagedSpecReconciler(t) calls := 0 @@ -243,7 +274,8 @@ func TestManagedSpecConfigMapRequests(t *testing.T) { &appsv1.WeightsAndBiases{ObjectMeta: metav1.ObjectMeta{Name: "first", Namespace: "default"}}, &appsv1.WeightsAndBiases{ObjectMeta: metav1.ObjectMeta{Name: "other", Namespace: "other"}}, ).Build(), - Scheme: scheme, + Scheme: scheme, + ManagedSpecEnabled: true, } requests := reconciler.managedSpecConfigMapRequests(context.Background(), &corev1.ConfigMap{ @@ -256,6 +288,14 @@ func TestManagedSpecConfigMapRequests(t *testing.T) { if requests[0].NamespacedName != want { t.Fatalf("managedSpecConfigMapRequests returned %v, want %v", requests[0].NamespacedName, want) } + + reconciler.ManagedSpecEnabled = false + requests = reconciler.managedSpecConfigMapRequests(context.Background(), &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: "default"}, + }) + if len(requests) != 0 { + t.Fatalf("managedSpecConfigMapRequests returned %d requests while disabled, want 0", len(requests)) + } } func testManagedSpecReconciler(t *testing.T, objects ...client.Object) *WeightsAndBiasesReconciler { @@ -265,8 +305,9 @@ func testManagedSpecReconciler(t *testing.T, objects ...client.Object) *WeightsA t.Fatalf("could not register core Kubernetes types: %v", err) } return &WeightsAndBiasesReconciler{ - Client: fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).Build(), - Scheme: scheme, + Client: fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).Build(), + Scheme: scheme, + ManagedSpecEnabled: true, } } diff --git a/internal/controller/weightsandbiases_controller.go b/internal/controller/weightsandbiases_controller.go index ebefee03..c21b1c17 100644 --- a/internal/controller/weightsandbiases_controller.go +++ b/internal/controller/weightsandbiases_controller.go @@ -57,12 +57,13 @@ const resFinalizer = "finalizer.app.wandb.com" // WeightsAndBiasesReconciler reconciles a WeightsAndBiases object type WeightsAndBiasesReconciler struct { client.Client - IsAirgapped bool - DeployerClient deployer.DeployerInterface - Scheme *runtime.Scheme - Recorder record.EventRecorder - DryRun bool - Debug bool + IsAirgapped bool + DeployerClient deployer.DeployerInterface + Scheme *runtime.Scheme + Recorder record.EventRecorder + DryRun bool + Debug bool + ManagedSpecEnabled bool } //+kubebuilder:rbac:groups=apps.wandb.com,resources=weightsandbiases,verbs=get;list;watch;create;update;patch;delete @@ -401,16 +402,18 @@ func (r *WeightsAndBiasesReconciler) Delete(e event.DeleteEvent) bool { // SetupWithManager sets up the controller with the Manager. func (r *WeightsAndBiasesReconciler) SetupWithManager(mgr ctrl.Manager) error { - builder := ctrl.NewControllerManagedBy(mgr). + controllerBuilder := ctrl.NewControllerManagedBy(mgr). For(&apiv1.WeightsAndBiases{}, builder.WithPredicates(filterWBEvents{})). Owns(&corev1.Secret{}, builder.WithPredicates(filterSecretEvents{})). - Owns(&corev1.ConfigMap{}). - Watches( + Owns(&corev1.ConfigMap{}) + if r.ManagedSpecEnabled { + controllerBuilder = controllerBuilder.Watches( &corev1.ConfigMap{}, handler.EnqueueRequestsFromMapFunc(r.managedSpecConfigMapRequests), builder.WithPredicates(predicate.NewPredicateFuncs(isManagedSpecConfigMap)), ) - return builder.Complete(r) + } + return controllerBuilder.Complete(r) } func isManagedSpecConfigMap(object client.Object) bool { @@ -421,6 +424,10 @@ func (r *WeightsAndBiasesReconciler) managedSpecConfigMapRequests( ctx context.Context, object client.Object, ) []reconcile.Request { + if !r.ManagedSpecEnabled { + return nil + } + instances := &apiv1.WeightsAndBiasesList{} if err := r.List(ctx, instances, client.InNamespace(object.GetNamespace())); err != nil { ctrllog.FromContext(ctx).Error(err, "Failed to list WeightsAndBiases instances for managed spec ConfigMap") diff --git a/internal/controller/weightsandbiases_controller_test.go b/internal/controller/weightsandbiases_controller_test.go index 0fee638f..31dddced 100644 --- a/internal/controller/weightsandbiases_controller_test.go +++ b/internal/controller/weightsandbiases_controller_test.go @@ -333,12 +333,13 @@ var _ = Describe("WeightsandbiasesController", func() { deployerClient = &deployerfakes.FakeDeployerInterface{} deployerClient.GetSpecReturns(&deployerSpec, nil) reconciler = &WeightsAndBiasesReconciler{ - Client: k8sClient, - IsAirgapped: false, - DeployerClient: deployerClient, - Scheme: scheme.Scheme, - Recorder: recorder, - DryRun: true, + Client: k8sClient, + IsAirgapped: false, + DeployerClient: deployerClient, + Scheme: scheme.Scheme, + Recorder: recorder, + DryRun: true, + ManagedSpecEnabled: true, } wandb := &wandbcomv1.WeightsAndBiases{ From ceb9ae923a8a4954f357aa207bbfabeb59875e6f Mon Sep 17 00:00:00 2001 From: Aditya Choudhari Date: Tue, 18 Aug 2026 16:48:01 -0400 Subject: [PATCH 3/4] refactor: clarify managed spec cutover flow --- cmd/main.go | 18 +++--- internal/controller/managed_spec.go | 55 +++++++++++----- internal/controller/managed_spec_test.go | 63 ++++++++----------- .../controller/weightsandbiases_controller.go | 44 ++++++------- .../weightsandbiases_controller_test.go | 14 ++--- 5 files changed, 100 insertions(+), 94 deletions(-) diff --git a/cmd/main.go b/cmd/main.go index 0ff5a279..d0c2a6c9 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -69,7 +69,7 @@ func main() { var enableHTTP2 bool var tlsOpts []func(*tls.Config) var deployerAPI, isolationNamespaces string - var debug, airgapped, managedSpecEnabled bool + var debug, airgapped, managedSpecCutoverEnabled bool flag.StringVar(&metricsAddr, "metrics-bind-address", "0", "The address the metrics endpoint binds to. "+ "Use :8443 for HTTPS or :8080 for HTTP, or leave as 0 to disable the metrics service.") @@ -94,7 +94,7 @@ func main() { flag.StringVar(&isolationNamespaces, "isolation-namespaces", "", "Specify namespaces (as a comma separated string) that the controller should monitor when operating in namespace isolation mode.") flag.BoolVar(&debug, "debug", false, "Enable debug mode") - flag.BoolVar(&managedSpecEnabled, "managed-spec-enabled", false, "Enable managed spec cutover") + flag.BoolVar(&managedSpecCutoverEnabled, "managed-spec-enabled", false, "Enable managed spec cutover") opts := zap.Options{ Development: true, @@ -230,13 +230,13 @@ func main() { } if err = (&controller.WeightsAndBiasesReconciler{ - IsAirgapped: airgapped, - Recorder: mgr.GetEventRecorderFor("weightsandbiases"), - Client: mgr.GetClient(), - Scheme: mgr.GetScheme(), - DeployerClient: &deployer.DeployerClient{DeployerAPI: deployerAPI}, - Debug: debug, - ManagedSpecEnabled: managedSpecEnabled, + IsAirgapped: airgapped, + Recorder: mgr.GetEventRecorderFor("weightsandbiases"), + Client: mgr.GetClient(), + Scheme: mgr.GetScheme(), + DeployerClient: &deployer.DeployerClient{DeployerAPI: deployerAPI}, + Debug: debug, + ManagedSpecCutoverEnabled: managedSpecCutoverEnabled, }).SetupWithManager(mgr); err != nil { setupLog.Error(err, "unable to create controller", "controller", "WeightsAndBiases") os.Exit(1) diff --git a/internal/controller/managed_spec.go b/internal/controller/managed_spec.go index b1ae8c8b..faddbc8e 100644 --- a/internal/controller/managed_spec.go +++ b/internal/controller/managed_spec.go @@ -27,58 +27,66 @@ type managedSpecSource struct { rawValues map[string]interface{} } +type baseSpecSelection struct { + selectedSpec *spec.Spec + shouldCompleteCutover bool +} + func (r *WeightsAndBiasesReconciler) selectBaseSpec( ctx context.Context, namespace string, getDeployerSpec func() (*spec.Spec, error), -) (*spec.Spec, bool, error) { - if !r.ManagedSpecEnabled { +) (baseSpecSelection, error) { + if !r.ManagedSpecCutoverEnabled { deployerSpec, err := getDeployerSpec() - return deployerSpec, false, err + return baseSpecSelection{selectedSpec: deployerSpec}, err } log := ctrllog.FromContext(ctx) - managed, err := r.managedSpecEnabled(ctx, namespace) + cutoverComplete, err := r.isManagedSpecCutoverComplete(ctx, namespace) if err != nil { - return nil, false, err + return baseSpecSelection{}, err } - if managed { + if cutoverComplete { log.Info("Managed spec cutover is active; skipping Deployer") managedSpec, err := r.getManagedSpec(ctx, namespace) if err != nil { - return nil, false, err + return baseSpecSelection{}, err } - return managedSpec.spec, false, nil + return baseSpecSelection{selectedSpec: managedSpec.spec}, nil } deployerSpec, err := getDeployerSpec() if err != nil { - return nil, false, err + return baseSpecSelection{}, err } managedSpec, err := r.getManagedSpec(ctx, namespace) if apierrors.IsNotFound(err) { - return deployerSpec, false, nil + return baseSpecSelection{selectedSpec: deployerSpec}, nil } if err != nil { log.Info("Managed spec is invalid; continuing with Deployer", "error", err) - return deployerSpec, false, nil + return baseSpecSelection{selectedSpec: deployerSpec}, nil } matches, err := managedSpecConfigurationMatches(managedSpec, deployerSpec) if err != nil { log.Info("Managed spec could not be compared; continuing with Deployer", "error", err) - return deployerSpec, false, nil + return baseSpecSelection{selectedSpec: deployerSpec}, nil } if matches { log.Info("Managed spec matches Deployer; cutover is pending successful apply") - return managedSpec.spec, true, nil + return baseSpecSelection{ + selectedSpec: managedSpec.spec, + shouldCompleteCutover: true, + }, nil } log.Info("Managed spec does not match Deployer; continuing with Deployer") - return deployerSpec, false, nil + return baseSpecSelection{selectedSpec: deployerSpec}, nil } -func (r *WeightsAndBiasesReconciler) managedSpecEnabled(ctx context.Context, namespace string) (bool, error) { +func (r *WeightsAndBiasesReconciler) isManagedSpecCutoverComplete(ctx context.Context, namespace string) (bool, error) { state := &corev1.ConfigMap{} err := r.Get(ctx, client.ObjectKey{Name: managedSpecStateConfigMapName, Namespace: namespace}, state) if apierrors.IsNotFound(err) { @@ -98,7 +106,7 @@ func (r *WeightsAndBiasesReconciler) managedSpecEnabled(ctx context.Context, nam return false, fmt.Errorf("invalid %s value %q in ConfigMap %s", managedSpecStateKey, value, managedSpecStateConfigMapName) } -func (r *WeightsAndBiasesReconciler) setManagedSpecEnabled(ctx context.Context, namespace string) error { +func (r *WeightsAndBiasesReconciler) markManagedSpecCutoverComplete(ctx context.Context, namespace string) error { key := client.ObjectKey{Name: managedSpecStateConfigMapName, Namespace: namespace} state := &corev1.ConfigMap{} err := r.Get(ctx, key, state) @@ -121,6 +129,21 @@ func (r *WeightsAndBiasesReconciler) setManagedSpecEnabled(ctx context.Context, return r.Update(ctx, state) } +func (r *WeightsAndBiasesReconciler) completeManagedSpecCutoverIfNeeded( + ctx context.Context, + namespace string, + shouldComplete bool, +) error { + if !shouldComplete { + return nil + } + if err := r.markManagedSpecCutoverComplete(ctx, namespace); err != nil { + return err + } + ctrllog.FromContext(ctx).Info("Managed spec cutover completed") + return nil +} + func (r *WeightsAndBiasesReconciler) getManagedSpec(ctx context.Context, namespace string) (*managedSpecSource, error) { configMap := &corev1.ConfigMap{} key := client.ObjectKey{Name: managedSpecConfigMapName, Namespace: namespace} diff --git a/internal/controller/managed_spec_test.go b/internal/controller/managed_spec_test.go index 44af2e2d..db274aa9 100644 --- a/internal/controller/managed_spec_test.go +++ b/internal/controller/managed_spec_test.go @@ -33,10 +33,10 @@ func TestManagedSpecSelection(t *testing.T) { Data: map[string]string{managedSpecStateKey: "true"}, }, ) - reconciler.ManagedSpecEnabled = false + reconciler.ManagedSpecCutoverEnabled = false calls := 0 - selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + selection, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { calls++ return deployerSpec, nil }) @@ -44,10 +44,10 @@ func TestManagedSpecSelection(t *testing.T) { if err != nil { t.Fatalf("selectBaseSpec returned an error: %v", err) } - if selected != deployerSpec { + if selection.selectedSpec != deployerSpec { t.Fatal("selectBaseSpec did not return the Deployer spec while managed spec was disabled") } - if pendingCutover { + if selection.shouldCompleteCutover { t.Fatal("cutover must not be pending while managed spec is disabled") } if calls != 1 { @@ -59,7 +59,7 @@ func TestManagedSpecSelection(t *testing.T) { reconciler := testManagedSpecReconciler(t) calls := 0 - selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + selection, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { calls++ return deployerSpec, nil }) @@ -67,10 +67,10 @@ func TestManagedSpecSelection(t *testing.T) { if err != nil { t.Fatalf("selectBaseSpec returned an error: %v", err) } - if selected != deployerSpec { + if selection.selectedSpec != deployerSpec { t.Fatal("selectBaseSpec did not return the Deployer spec") } - if pendingCutover { + if selection.shouldCompleteCutover { t.Fatal("cutover must not be pending without a managed spec") } if calls != 1 { @@ -83,17 +83,17 @@ func TestManagedSpecSelection(t *testing.T) { "global": map[string]interface{}{"enabled": false}, })) - selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + selection, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { return deployerSpec, nil }) if err != nil { t.Fatalf("selectBaseSpec returned an error: %v", err) } - if selected != deployerSpec { + if selection.selectedSpec != deployerSpec { t.Fatal("selectBaseSpec did not return the Deployer spec") } - if pendingCutover { + if selection.shouldCompleteCutover { t.Fatal("cutover must not be pending for a mismatched managed spec") } }) @@ -101,7 +101,7 @@ func TestManagedSpecSelection(t *testing.T) { t.Run("selects matching managed-owned configuration and requests cutover", func(t *testing.T) { reconciler := testManagedSpecReconciler(t, testManagedSpecConfigMap(namespace, deployerSpec.Values)) - selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + selection, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { withMetadata := *testManagedSpec(map[string]interface{}{ "global": map[string]interface{}{ "enabled": true, @@ -118,13 +118,13 @@ func TestManagedSpecSelection(t *testing.T) { if err != nil { t.Fatalf("selectBaseSpec returned an error: %v", err) } - if !pendingCutover { + if !selection.shouldCompleteCutover { t.Fatal("matching managed configuration must request cutover") } - if selected == nil || selected.Metadata != nil { + if selection.selectedSpec == nil || selection.selectedSpec.Metadata != nil { t.Fatal("selectBaseSpec did not return the managed ConfigMap spec") } - if !reflect.DeepEqual(selected.Values, deployerSpec.Values) { + if !reflect.DeepEqual(selection.selectedSpec.Values, deployerSpec.Values) { t.Fatal("selectBaseSpec did not preserve the managed-owned values") } }) @@ -140,7 +140,7 @@ func TestManagedSpecSelection(t *testing.T) { ) calls := 0 - selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + selection, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { calls++ return nil, errors.New("Deployer must not be called") }) @@ -148,10 +148,10 @@ func TestManagedSpecSelection(t *testing.T) { if err != nil { t.Fatalf("selectBaseSpec returned an error: %v", err) } - if pendingCutover { + if selection.shouldCompleteCutover { t.Fatal("cutover cannot be pending after it is active") } - if selected == nil || !selected.IsEqual(deployerSpec) { + if selection.selectedSpec == nil || !selection.selectedSpec.IsEqual(deployerSpec) { t.Fatal("selectBaseSpec did not return the managed spec") } if calls != 0 { @@ -166,7 +166,7 @@ func TestManagedSpecSelection(t *testing.T) { }) calls := 0 - selected, pendingCutover, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + selection, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { calls++ return deployerSpec, nil }) @@ -174,7 +174,7 @@ func TestManagedSpecSelection(t *testing.T) { if err == nil { t.Fatal("selectBaseSpec succeeded without the managed spec after cutover") } - if selected != nil || pendingCutover { + if selection.selectedSpec != nil || selection.shouldCompleteCutover { t.Fatal("selectBaseSpec returned a spec while failing closed") } if calls != 0 { @@ -190,7 +190,7 @@ func TestManagedSpecSelection(t *testing.T) { }) calls := 0 - _, _, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { + _, err := reconciler.selectBaseSpec(ctx, namespace, func() (*spec.Spec, error) { calls++ return deployerSpec, nil }) @@ -205,13 +205,13 @@ func TestManagedSpecSelection(t *testing.T) { } } -func TestSetManagedSpecEnabled(t *testing.T) { +func TestMarkManagedSpecCutoverComplete(t *testing.T) { ctx := context.Background() namespace := "default" reconciler := testManagedSpecReconciler(t) - if err := reconciler.setManagedSpecEnabled(ctx, namespace); err != nil { - t.Fatalf("setManagedSpecEnabled returned an error: %v", err) + if err := reconciler.markManagedSpecCutoverComplete(ctx, namespace); err != nil { + t.Fatalf("markManagedSpecCutoverComplete returned an error: %v", err) } state := &corev1.ConfigMap{} @@ -274,8 +274,7 @@ func TestManagedSpecConfigMapRequests(t *testing.T) { &appsv1.WeightsAndBiases{ObjectMeta: metav1.ObjectMeta{Name: "first", Namespace: "default"}}, &appsv1.WeightsAndBiases{ObjectMeta: metav1.ObjectMeta{Name: "other", Namespace: "other"}}, ).Build(), - Scheme: scheme, - ManagedSpecEnabled: true, + Scheme: scheme, } requests := reconciler.managedSpecConfigMapRequests(context.Background(), &corev1.ConfigMap{ @@ -288,14 +287,6 @@ func TestManagedSpecConfigMapRequests(t *testing.T) { if requests[0].NamespacedName != want { t.Fatalf("managedSpecConfigMapRequests returned %v, want %v", requests[0].NamespacedName, want) } - - reconciler.ManagedSpecEnabled = false - requests = reconciler.managedSpecConfigMapRequests(context.Background(), &corev1.ConfigMap{ - ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: "default"}, - }) - if len(requests) != 0 { - t.Fatalf("managedSpecConfigMapRequests returned %d requests while disabled, want 0", len(requests)) - } } func testManagedSpecReconciler(t *testing.T, objects ...client.Object) *WeightsAndBiasesReconciler { @@ -305,9 +296,9 @@ func testManagedSpecReconciler(t *testing.T, objects ...client.Object) *WeightsA t.Fatalf("could not register core Kubernetes types: %v", err) } return &WeightsAndBiasesReconciler{ - Client: fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).Build(), - Scheme: scheme, - ManagedSpecEnabled: true, + Client: fake.NewClientBuilder().WithScheme(scheme).WithObjects(objects...).Build(), + Scheme: scheme, + ManagedSpecCutoverEnabled: true, } } diff --git a/internal/controller/weightsandbiases_controller.go b/internal/controller/weightsandbiases_controller.go index c21b1c17..9950ba79 100644 --- a/internal/controller/weightsandbiases_controller.go +++ b/internal/controller/weightsandbiases_controller.go @@ -57,13 +57,13 @@ const resFinalizer = "finalizer.app.wandb.com" // WeightsAndBiasesReconciler reconciles a WeightsAndBiases object type WeightsAndBiasesReconciler struct { client.Client - IsAirgapped bool - DeployerClient deployer.DeployerInterface - Scheme *runtime.Scheme - Recorder record.EventRecorder - DryRun bool - Debug bool - ManagedSpecEnabled bool + IsAirgapped bool + DeployerClient deployer.DeployerInterface + Scheme *runtime.Scheme + Recorder record.EventRecorder + DryRun bool + Debug bool + ManagedSpecCutoverEnabled bool } //+kubebuilder:rbac:groups=apps.wandb.com,resources=weightsandbiases,verbs=get;list;watch;create;update;patch;delete @@ -191,13 +191,15 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req } baseSpec := currentActiveSpec - pendingManagedCutover := false + shouldCompleteManagedSpecCutover := false if wandb.ObjectMeta.DeletionTimestamp.IsZero() { - baseSpec, pendingManagedCutover, err = r.selectBaseSpec(ctx, wandb.Namespace, getDeployerSpec) + selection, err := r.selectBaseSpec(ctx, wandb.Namespace, getDeployerSpec) if err != nil { log.Error(err, "Failed to select Deployer or managed spec") return ctrlqueue.RequeueWithError(err) } + baseSpec = selection.selectedSpec + shouldCompleteManagedSpecCutover = selection.shouldCompleteCutover } desiredSpec := new(spec.Spec) @@ -248,12 +250,9 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req log.Info("Active spec found", "spec", currentActiveSpec.SensitiveValuesMasked()) if currentActiveSpec.IsEqual(desiredSpec) { log.Info("No changes found") - if pendingManagedCutover { - if err := r.setManagedSpecEnabled(ctx, wandb.Namespace); err != nil { - log.Error(err, "Failed to persist managed spec cutover") - return ctrlqueue.RequeueWithError(err) - } - log.Info("Managed spec cutover completed") + if err := r.completeManagedSpecCutoverIfNeeded(ctx, wandb.Namespace, shouldCompleteManagedSpecCutover); err != nil { + log.Error(err, "Failed to persist managed spec cutover") + return ctrlqueue.RequeueWithError(err) } statusManager.Set(status.Completed) return ctrlqueue.Requeue(desiredSpec) @@ -304,12 +303,9 @@ func (r *WeightsAndBiasesReconciler) Reconcile(ctx context.Context, req ctrl.Req if r.Debug { log.Info("Successfully saved active spec", "spec", desiredSpec.SensitiveValuesMasked()) } - if pendingManagedCutover { - if err := r.setManagedSpecEnabled(ctx, wandb.Namespace); err != nil { - log.Error(err, "Failed to persist managed spec cutover") - return ctrlqueue.RequeueWithError(err) - } - log.Info("Managed spec cutover completed") + if err := r.completeManagedSpecCutoverIfNeeded(ctx, wandb.Namespace, shouldCompleteManagedSpecCutover); err != nil { + log.Error(err, "Failed to persist managed spec cutover") + return ctrlqueue.RequeueWithError(err) } r.Recorder.Event(wandb, corev1.EventTypeNormal, "Completed", "Completed reconcile successfully") @@ -406,7 +402,7 @@ func (r *WeightsAndBiasesReconciler) SetupWithManager(mgr ctrl.Manager) error { For(&apiv1.WeightsAndBiases{}, builder.WithPredicates(filterWBEvents{})). Owns(&corev1.Secret{}, builder.WithPredicates(filterSecretEvents{})). Owns(&corev1.ConfigMap{}) - if r.ManagedSpecEnabled { + if r.ManagedSpecCutoverEnabled { controllerBuilder = controllerBuilder.Watches( &corev1.ConfigMap{}, handler.EnqueueRequestsFromMapFunc(r.managedSpecConfigMapRequests), @@ -424,10 +420,6 @@ func (r *WeightsAndBiasesReconciler) managedSpecConfigMapRequests( ctx context.Context, object client.Object, ) []reconcile.Request { - if !r.ManagedSpecEnabled { - return nil - } - instances := &apiv1.WeightsAndBiasesList{} if err := r.List(ctx, instances, client.InNamespace(object.GetNamespace())); err != nil { ctrllog.FromContext(ctx).Error(err, "Failed to list WeightsAndBiases instances for managed spec ConfigMap") diff --git a/internal/controller/weightsandbiases_controller_test.go b/internal/controller/weightsandbiases_controller_test.go index 31dddced..0682aa22 100644 --- a/internal/controller/weightsandbiases_controller_test.go +++ b/internal/controller/weightsandbiases_controller_test.go @@ -333,13 +333,13 @@ var _ = Describe("WeightsandbiasesController", func() { deployerClient = &deployerfakes.FakeDeployerInterface{} deployerClient.GetSpecReturns(&deployerSpec, nil) reconciler = &WeightsAndBiasesReconciler{ - Client: k8sClient, - IsAirgapped: false, - DeployerClient: deployerClient, - Scheme: scheme.Scheme, - Recorder: recorder, - DryRun: true, - ManagedSpecEnabled: true, + Client: k8sClient, + IsAirgapped: false, + DeployerClient: deployerClient, + Scheme: scheme.Scheme, + Recorder: recorder, + DryRun: true, + ManagedSpecCutoverEnabled: true, } wandb := &wandbcomv1.WeightsAndBiases{ From e1e1c78ea44a5d77ac0c027adb7b2ea0bd9809b4 Mon Sep 17 00:00:00 2001 From: Aditya Choudhari Date: Fri, 21 Aug 2026 13:45:40 -0400 Subject: [PATCH 4/4] cleanup normamization --- internal/controller/managed_spec.go | 61 ------ .../controller/managed_spec_comparison.go | 171 +++++++++++++++ .../managed_spec_comparison_test.go | 201 ++++++++++++++++++ internal/controller/managed_spec_test.go | 38 +++- .../weightsandbiases_controller_test.go | 22 +- 5 files changed, 422 insertions(+), 71 deletions(-) create mode 100644 internal/controller/managed_spec_comparison.go create mode 100644 internal/controller/managed_spec_comparison_test.go diff --git a/internal/controller/managed_spec.go b/internal/controller/managed_spec.go index faddbc8e..00f6b57c 100644 --- a/internal/controller/managed_spec.go +++ b/internal/controller/managed_spec.go @@ -4,7 +4,6 @@ import ( "context" "encoding/json" "fmt" - "reflect" "github.com/wandb/operator/pkg/wandb/spec" "github.com/wandb/operator/pkg/wandb/spec/charts" @@ -179,63 +178,3 @@ func (r *WeightsAndBiasesReconciler) getManagedSpec(ctx context.Context, namespa rawValues: rawValues, }, nil } - -func managedSpecConfigurationMatches(managed *managedSpecSource, deployer *spec.Spec) (bool, error) { - if managed == nil || deployer == nil { - return false, nil - } - - deployerChart, err := normalizeJSONValue(deployer.Chart) - if err != nil { - return false, fmt.Errorf("normalize Deployer chart: %w", err) - } - deployerValues, err := normalizeJSONValue(deployer.Values) - if err != nil { - return false, fmt.Errorf("normalize Deployer values: %w", err) - } - - return managedJSONSubsetEqual(managed.rawChart, deployerChart) && - managedJSONSubsetEqual(managed.rawValues, deployerValues), nil -} - -func normalizeJSONValue(value interface{}) (interface{}, error) { - data, err := json.Marshal(value) - if err != nil { - return nil, err - } - var normalized interface{} - if err := json.Unmarshal(data, &normalized); err != nil { - return nil, err - } - return normalized, nil -} - -func managedJSONSubsetEqual(managed, deployer interface{}) bool { - switch managedValue := managed.(type) { - case map[string]interface{}: - deployerValue, ok := deployer.(map[string]interface{}) - if !ok { - return false - } - for key, managedChild := range managedValue { - deployerChild, ok := deployerValue[key] - if !ok || !managedJSONSubsetEqual(managedChild, deployerChild) { - return false - } - } - return true - case []interface{}: - deployerValue, ok := deployer.([]interface{}) - if !ok || len(managedValue) != len(deployerValue) { - return false - } - for index, managedChild := range managedValue { - if !managedJSONSubsetEqual(managedChild, deployerValue[index]) { - return false - } - } - return true - default: - return reflect.DeepEqual(managed, deployer) - } -} diff --git a/internal/controller/managed_spec_comparison.go b/internal/controller/managed_spec_comparison.go new file mode 100644 index 00000000..26cbf9d9 --- /dev/null +++ b/internal/controller/managed_spec_comparison.go @@ -0,0 +1,171 @@ +package controller + +import ( + "encoding/json" + "fmt" + "reflect" + "strings" + + "github.com/wandb/operator/pkg/wandb/spec" +) + +const ( + orbUsageEventReporterEnvironment = "GORILLA_ORB_USAGE_EVENT_REPORTER_SECRET" +) + +func managedSpecConfigurationMatches(managed *managedSpecSource, deployer *spec.Spec) (bool, error) { + if managed == nil || deployer == nil { + return false, nil + } + + managedChart, err := normalizeJSONValue(managed.rawChart) + if err != nil { + return false, fmt.Errorf("normalize managed chart: %w", err) + } + deployerChart, err := normalizeJSONValue(deployer.Chart) + if err != nil { + return false, fmt.Errorf("normalize Deployer chart: %w", err) + } + managedValues, err := normalizeJSONValue(managed.rawValues) + if err != nil { + return false, fmt.Errorf("normalize managed values: %w", err) + } + deployerValues, err := normalizeJSONValue(deployer.Values) + if err != nil { + return false, fmt.Errorf("normalize Deployer values: %w", err) + } + + managedValuesMap, ok := managedValues.(map[string]interface{}) + if !ok { + return false, fmt.Errorf("normalized managed values are not an object") + } + deployerValuesMap, ok := deployerValues.(map[string]interface{}) + if !ok { + return false, fmt.Errorf("normalized Deployer values are not an object") + } + if err := normalizeExpectedManagedSpecDifferences(managedValuesMap, deployerValuesMap); err != nil { + return false, err + } + + return managedJSONSubsetEqual(managedChart, deployerChart) && + managedJSONSubsetEqual(managedValuesMap, deployerValuesMap), nil +} + +func normalizeExpectedManagedSpecDifferences(managed, deployer map[string]interface{}) error { + value, ok := nestedValue(managed, "global", "cloudProvider") + if !ok { + return fmt.Errorf("managed global.cloudProvider is required") + } + cloudProvider, ok := value.(string) + if !ok || !isSupportedManagedCloudProvider(cloudProvider) { + return fmt.Errorf("managed global.cloudProvider must be aws, gcp, or azure") + } + deployerCloud, hasDeployerCloud := nestedValue(deployer, "global", "extraEnv", "TAG_CLOUD") + if hasDeployerCloud { + deployerCloudString, ok := deployerCloud.(string) + if !ok || strings.ToLower(deployerCloudString) != cloudProvider { + return fmt.Errorf("managed global.cloudProvider %q does not match Deployer TAG_CLOUD", cloudProvider) + } + } + deleteNestedValue(managed, "global", "cloudProvider") + + value, ok = nestedValue(managed, "global", "extraEnv", "TAG_CUSTOMER_NS") + if !ok { + return fmt.Errorf("managed TAG_CUSTOMER_NS is required") + } + customerNamespace, ok := value.(string) + if !ok || strings.TrimSpace(customerNamespace) == "" { + return fmt.Errorf("managed TAG_CUSTOMER_NS must be a non-empty string") + } + deleteNestedValue(managed, "global", "extraEnv", "TAG_CUSTOMER_NS") + + deleteNestedValue(managed, "otel", "daemonset", "config", "exporters", "datadog", "api", "key") + deleteNestedValue(managed, "otel", "daemonset", "extraEnvFrom", "DD_API_KEY") + deleteNestedValue(managed, "app", "env", orbUsageEventReporterEnvironment) + deleteNestedValue(managed, "glue", "env", orbUsageEventReporterEnvironment) + return nil +} + +func isSupportedManagedCloudProvider(cloudProvider string) bool { + return cloudProvider == "aws" || cloudProvider == "gcp" || cloudProvider == "azure" +} + +func nestedValue(root map[string]interface{}, path ...string) (interface{}, bool) { + var current interface{} = root + for _, key := range path { + object, ok := current.(map[string]interface{}) + if !ok { + return nil, false + } + current, ok = object[key] + if !ok { + return nil, false + } + } + return current, true +} + +func deleteNestedValue(root map[string]interface{}, path ...string) { + if len(path) == 0 { + return + } + objects := []map[string]interface{}{root} + current := root + for _, key := range path[:len(path)-1] { + next, ok := current[key].(map[string]interface{}) + if !ok { + return + } + objects = append(objects, next) + current = next + } + delete(current, path[len(path)-1]) + for index := len(objects) - 1; index > 0; index-- { + if len(objects[index]) != 0 { + break + } + delete(objects[index-1], path[index-1]) + } +} + +func normalizeJSONValue(value interface{}) (interface{}, error) { + data, err := json.Marshal(value) + if err != nil { + return nil, err + } + var normalized interface{} + if err := json.Unmarshal(data, &normalized); err != nil { + return nil, err + } + return normalized, nil +} + +func managedJSONSubsetEqual(managed, deployer interface{}) bool { + switch managedValue := managed.(type) { + case map[string]interface{}: + deployerValue, ok := deployer.(map[string]interface{}) + if !ok { + return false + } + for key, managedChild := range managedValue { + deployerChild, ok := deployerValue[key] + if !ok || !managedJSONSubsetEqual(managedChild, deployerChild) { + return false + } + } + return true + case []interface{}: + deployerValue, ok := deployer.([]interface{}) + if !ok || len(managedValue) != len(deployerValue) { + return false + } + for index, managedChild := range managedValue { + if !managedJSONSubsetEqual(managedChild, deployerValue[index]) { + return false + } + } + return true + default: + return reflect.DeepEqual(managed, deployer) + } +} diff --git a/internal/controller/managed_spec_comparison_test.go b/internal/controller/managed_spec_comparison_test.go new file mode 100644 index 00000000..275cc28e --- /dev/null +++ b/internal/controller/managed_spec_comparison_test.go @@ -0,0 +1,201 @@ +package controller + +import ( + "testing" + + "github.com/wandb/operator/pkg/wandb/spec" +) + +func TestManagedSpecConfigurationMatchesExpectedDifferences(t *testing.T) { + managedSpec, deployerSpec := testExpectedDifferenceSpecs() + + matches, err := managedSpecConfigurationMatches(managedSpec, deployerSpec) + if err != nil { + t.Fatalf("managedSpecConfigurationMatches returned an error: %v", err) + } + if !matches { + t.Fatal("expected managed metadata and ESO credential references to match Deployer literals") + } +} + +func TestManagedSpecConfigurationRejectsInvalidMetadata(t *testing.T) { + tests := []struct { + name string + mutate func(*testing.T, map[string]interface{}) + }{ + { + name: "unsupported cloud provider", + mutate: func(t *testing.T, managed map[string]interface{}) { + testSetNestedValue(t, managed, "digitalocean", "global", "cloudProvider") + }, + }, + { + name: "missing cloud provider", + mutate: func(t *testing.T, managed map[string]interface{}) { + testDeleteNestedValue(t, managed, "global", "cloudProvider") + }, + }, + { + name: "cloud provider disagrees with Deployer cloud tag", + mutate: func(t *testing.T, managed map[string]interface{}) { + testSetNestedValue(t, managed, "aws", "global", "cloudProvider") + }, + }, + { + name: "empty customer namespace", + mutate: func(t *testing.T, managed map[string]interface{}) { + testSetNestedValue(t, managed, "", "global", "extraEnv", "TAG_CUSTOMER_NS") + }, + }, + { + name: "missing customer namespace", + mutate: func(t *testing.T, managed map[string]interface{}) { + testDeleteNestedValue(t, managed, "global", "extraEnv", "TAG_CUSTOMER_NS") + }, + }, + { + name: "non-string customer namespace", + mutate: func(t *testing.T, managed map[string]interface{}) { + testSetNestedValue(t, managed, true, "global", "extraEnv", "TAG_CUSTOMER_NS") + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + managedSpec, deployerSpec := testExpectedDifferenceSpecs() + test.mutate(t, managedSpec.rawValues) + + matches, err := managedSpecConfigurationMatches(managedSpec, deployerSpec) + if err == nil { + t.Fatal("managedSpecConfigurationMatches accepted invalid managed metadata") + } + if matches { + t.Fatal("managedSpecConfigurationMatches matched invalid managed metadata") + } + }) + } +} + +func testExpectedDifferenceSpecs() (*managedSpecSource, *spec.Spec) { + managedValues := map[string]interface{}{ + "app": map[string]interface{}{ + "enabled": true, + "env": map[string]interface{}{ + "GORILLA_ORB_USAGE_EVENT_REPORTER_SECRET": testSecretKeyRef("orb-api-key", "api-key"), + }, + }, + "global": map[string]interface{}{ + "cloudProvider": "gcp", + "extraEnv": map[string]interface{}{ + "SERVER_FLAG_ENABLE_CORE_WEAVE_OBSERVABILITY": "true", + "TAG_CLOUD": "GCP", + "TAG_CUSTOMER_NS": "wandb-abridge", + }, + }, + "glue": map[string]interface{}{ + "env": map[string]interface{}{ + "GORILLA_ORB_USAGE_EVENT_REPORTER_SECRET": testSecretKeyRef("orb-api-key", "api-key"), + }, + }, + "otel": map[string]interface{}{ + "daemonset": map[string]interface{}{ + "config": map[string]interface{}{ + "exporters": map[string]interface{}{ + "datadog": map[string]interface{}{ + "api": map[string]interface{}{ + "key": "${env:DD_API_KEY}", + "site": "us5.datadoghq.com", + }, + }, + }, + }, + "extraEnvFrom": map[string]interface{}{ + "DD_API_KEY": map[string]interface{}{ + "secretKeyRef": map[string]interface{}{ + "key": "api-key", + "name": "datadog-api-key", + }, + }, + }, + }, + }, + } + deployerValues := map[string]interface{}{ + "app": map[string]interface{}{ + "enabled": true, + "extraEnv": map[string]interface{}{ + "GORILLA_ORB_USAGE_EVENT_REPORTER_SECRET": "deployer-orb-credential", + }, + "image": map[string]interface{}{"tag": "deployer-owned"}, + }, + "global": map[string]interface{}{ + "extraEnv": map[string]interface{}{ + "SERVER_FLAG_ENABLE_CORE_WEAVE_OBSERVABILITY": "true", + "TAG_CLOUD": "GCP", + }, + }, + "glue": map[string]interface{}{ + "env": map[string]interface{}{ + "GORILLA_ORB_USAGE_EVENT_REPORTER_SECRET": "deployer-orb-credential", + }, + }, + "otel": map[string]interface{}{ + "daemonset": map[string]interface{}{ + "config": map[string]interface{}{ + "exporters": map[string]interface{}{ + "datadog": map[string]interface{}{ + "api": map[string]interface{}{ + "key": "deployer-datadog-credential", + "site": "us5.datadoghq.com", + }, + }, + }, + }, + }, + }, + } + managedSpec := &managedSpecSource{ + spec: testManagedSpec(managedValues), + rawChart: map[string]interface{}{"name": "operator-wandb", "url": "https://charts.wandb.ai", "version": "0.43.5"}, + rawValues: managedValues, + } + return managedSpec, testManagedSpec(deployerValues) +} + +func testSecretKeyRef(name, key string) map[string]interface{} { + return map[string]interface{}{ + "valueFrom": map[string]interface{}{ + "secretKeyRef": map[string]interface{}{ + "name": name, + "key": key, + }, + }, + } +} + +func testSetNestedValue(t *testing.T, root map[string]interface{}, value interface{}, path ...string) { + t.Helper() + current := root + for _, key := range path[:len(path)-1] { + next, ok := current[key].(map[string]interface{}) + if !ok { + t.Fatalf("test fixture path %v does not contain an object at %q", path, key) + } + current = next + } + current[path[len(path)-1]] = value +} + +func testDeleteNestedValue(t *testing.T, root map[string]interface{}, path ...string) { + t.Helper() + current := root + for _, key := range path[:len(path)-1] { + next, ok := current[key].(map[string]interface{}) + if !ok { + t.Fatalf("test fixture path %v does not contain an object at %q", path, key) + } + current = next + } + delete(current, path[len(path)-1]) +} diff --git a/internal/controller/managed_spec_test.go b/internal/controller/managed_spec_test.go index db274aa9..eb8bba96 100644 --- a/internal/controller/managed_spec_test.go +++ b/internal/controller/managed_spec_test.go @@ -2,6 +2,7 @@ package controller import ( "context" + "encoding/json" "errors" "reflect" "testing" @@ -21,7 +22,10 @@ func TestManagedSpecSelection(t *testing.T) { ctx := context.Background() namespace := "default" deployerSpec := testManagedSpec(map[string]interface{}{ - "global": map[string]interface{}{"enabled": true}, + "global": map[string]interface{}{ + "enabled": true, + "extraEnv": map[string]interface{}{"TAG_CLOUD": "GCP"}, + }, }) t.Run("uses Deployer when managed spec is disabled even after cutover", func(t *testing.T) { @@ -106,6 +110,9 @@ func TestManagedSpecSelection(t *testing.T) { "global": map[string]interface{}{ "enabled": true, "image": map[string]interface{}{"tag": "deployer-only"}, + "extraEnv": map[string]interface{}{ + "TAG_CLOUD": "GCP", + }, }, "legacy": map[string]interface{}{"enabled": true}, }) @@ -124,7 +131,7 @@ func TestManagedSpecSelection(t *testing.T) { if selection.selectedSpec == nil || selection.selectedSpec.Metadata != nil { t.Fatal("selectBaseSpec did not return the managed ConfigMap spec") } - if !reflect.DeepEqual(selection.selectedSpec.Values, deployerSpec.Values) { + if !reflect.DeepEqual(selection.selectedSpec.Values, spec.Values(testManagedValues(true))) { t.Fatal("selectBaseSpec did not preserve the managed-owned values") } }) @@ -151,7 +158,7 @@ func TestManagedSpecSelection(t *testing.T) { if selection.shouldCompleteCutover { t.Fatal("cutover cannot be pending after it is active") } - if selection.selectedSpec == nil || !selection.selectedSpec.IsEqual(deployerSpec) { + if selection.selectedSpec == nil || !selection.selectedSpec.IsEqual(testManagedSpec(testManagedValues(true))) { t.Fatal("selectBaseSpec did not return the managed spec") } if calls != 0 { @@ -303,15 +310,32 @@ func testManagedSpecReconciler(t *testing.T, objects ...client.Object) *WeightsA } func testManagedSpecConfigMap(namespace string, values map[string]interface{}) *corev1.ConfigMap { - valuesJSON := `{"global":{"enabled":true}}` - if enabled, ok := values["global"].(map[string]interface{})["enabled"].(bool); ok && !enabled { - valuesJSON = `{"global":{"enabled":false}}` + enabled := true + if value, ok := values["global"].(map[string]interface{})["enabled"].(bool); ok { + enabled = value + } + valuesJSON, err := json.Marshal(testManagedValues(enabled)) + if err != nil { + panic(err) } return &corev1.ConfigMap{ ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: namespace}, Data: map[string]string{ "chart": `{"name":"operator-wandb","url":"https://charts.wandb.ai","version":"0.43.5"}`, - "values": valuesJSON, + "values": string(valuesJSON), + }, + } +} + +func testManagedValues(enabled bool) map[string]interface{} { + return map[string]interface{}{ + "global": map[string]interface{}{ + "cloudProvider": "gcp", + "enabled": enabled, + "extraEnv": map[string]interface{}{ + "TAG_CLOUD": "GCP", + "TAG_CUSTOMER_NS": "wandb-test", + }, }, } } diff --git a/internal/controller/weightsandbiases_controller_test.go b/internal/controller/weightsandbiases_controller_test.go index 0682aa22..1a39a08d 100644 --- a/internal/controller/weightsandbiases_controller_test.go +++ b/internal/controller/weightsandbiases_controller_test.go @@ -331,7 +331,16 @@ var _ = Describe("WeightsandbiasesController", func() { ctx := context.Background() recorder = record.NewFakeRecorder(10) deployerClient = &deployerfakes.FakeDeployerInterface{} - deployerClient.GetSpecReturns(&deployerSpec, nil) + deployerValuesJSON, err := json.Marshal(deployerSpec.Values) + Expect(err).NotTo(HaveOccurred()) + var deployerValues spec.Values + Expect(json.Unmarshal(deployerValuesJSON, &deployerValues)).To(Succeed()) + deployerValues["global"] = map[string]interface{}{ + "extraEnv": map[string]interface{}{"TAG_CLOUD": "GCP"}, + } + deployerSpecForCutover := deployerSpec + deployerSpecForCutover.Values = deployerValues + deployerClient.GetSpecReturns(&deployerSpecForCutover, nil) reconciler = &WeightsAndBiasesReconciler{ Client: k8sClient, IsAirgapped: false, @@ -351,9 +360,16 @@ var _ = Describe("WeightsandbiasesController", func() { } Expect(k8sClient.Create(ctx, wandb)).To(Succeed()) - chartJSON, err := json.Marshal(deployerSpec.Chart) + chartJSON, err := json.Marshal(deployerSpecForCutover.Chart) + Expect(err).NotTo(HaveOccurred()) + managedValuesJSON, err := json.Marshal(deployerValues) Expect(err).NotTo(HaveOccurred()) - valuesJSON, err := json.Marshal(deployerSpec.Values) + var managedValues spec.Values + Expect(json.Unmarshal(managedValuesJSON, &managedValues)).To(Succeed()) + managedGlobal := managedValues["global"].(map[string]interface{}) + managedGlobal["cloudProvider"] = "gcp" + managedGlobal["extraEnv"].(map[string]interface{})["TAG_CUSTOMER_NS"] = "wandb-test" + valuesJSON, err := json.Marshal(managedValues) Expect(err).NotTo(HaveOccurred()) managedSpec := &v1.ConfigMap{ ObjectMeta: metav1.ObjectMeta{Name: managedSpecConfigMapName, Namespace: "default"},