From 69eef9f560086417920d7b37d092d180df987728 Mon Sep 17 00:00:00 2001 From: Aditya Choudhari Date: Tue, 4 Aug 2026 11:31:58 -0400 Subject: [PATCH] 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()) + }) + }) })