Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 9 additions & 7 deletions cmd/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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, 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.")
Expand All @@ -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(&managedSpecCutoverEnabled, "managed-spec-enabled", false, "Enable managed spec cutover")

opts := zap.Options{
Development: true,
Expand Down Expand Up @@ -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,
ManagedSpecCutoverEnabled: managedSpecCutoverEnabled,
}).SetupWithManager(mgr); err != nil {
setupLog.Error(err, "unable to create controller", "controller", "WeightsAndBiases")
os.Exit(1)
Expand Down
180 changes: 180 additions & 0 deletions internal/controller/managed_spec.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,180 @@
package controller

import (
"context"
"encoding/json"
"fmt"

"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{}
}

type baseSpecSelection struct {
selectedSpec *spec.Spec
shouldCompleteCutover bool
}

func (r *WeightsAndBiasesReconciler) selectBaseSpec(
ctx context.Context,
namespace string,
getDeployerSpec func() (*spec.Spec, error),
) (baseSpecSelection, error) {
if !r.ManagedSpecCutoverEnabled {
deployerSpec, err := getDeployerSpec()
return baseSpecSelection{selectedSpec: deployerSpec}, err
}

log := ctrllog.FromContext(ctx)
cutoverComplete, err := r.isManagedSpecCutoverComplete(ctx, namespace)
if err != nil {
return baseSpecSelection{}, err
}
if cutoverComplete {
log.Info("Managed spec cutover is active; skipping Deployer")
managedSpec, err := r.getManagedSpec(ctx, namespace)
if err != nil {
return baseSpecSelection{}, err
}
return baseSpecSelection{selectedSpec: managedSpec.spec}, nil
}

deployerSpec, err := getDeployerSpec()
if err != nil {
return baseSpecSelection{}, err
}

managedSpec, err := r.getManagedSpec(ctx, namespace)
if apierrors.IsNotFound(err) {
return baseSpecSelection{selectedSpec: deployerSpec}, nil
}
if err != nil {
log.Info("Managed spec is invalid; continuing with Deployer", "error", err)
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 baseSpecSelection{selectedSpec: deployerSpec}, nil
}
if matches {
log.Info("Managed spec matches Deployer; cutover is pending successful apply")
return baseSpecSelection{
selectedSpec: managedSpec.spec,
shouldCompleteCutover: true,
}, nil
}

log.Info("Managed spec does not match Deployer; continuing with Deployer")
return baseSpecSelection{selectedSpec: deployerSpec}, nil
}

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) {
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) markManagedSpecCutoverComplete(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) 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}
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
}
171 changes: 171 additions & 0 deletions internal/controller/managed_spec_comparison.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
Loading
Loading