diff --git a/cmd/main.go b/cmd/main.go index 55aee60..dbb5da2 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -24,6 +24,8 @@ import ( "fmt" "os" "path/filepath" + "slices" + "strings" // Import all Kubernetes client auth plugins (e.g. Azure, GCP, OIDC, etc.) // to ensure that exec-entrypoint and run can make use of them. @@ -105,6 +107,7 @@ func main() { var kubeletAddr, kubeletClientCA string var costLabels string var kubeletServingTLSBootstrap bool + var providers string var tlsOpts []func(*tls.Config) 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.") @@ -147,6 +150,10 @@ func main() { "x509: certificate signed by unknown authority. Off by default because it needs the "+ "RBAC to impersonate one virtual node identity — the signer signs for nobody else "+ "(see addServingCertificateBootstrap).") + flag.StringVar(&providers, "providers", strings.Join(knownProviders, ","), + "Comma-separated production providers to enable; all by default. A listed provider is "+ + "skipped if it fails to initialize. Drain a provider before dropping it: nothing "+ + "terminates its running instances afterwards.") opts := zap.Options{ Development: true, } @@ -162,6 +169,11 @@ func main() { setupLog.Error(err, "invalid --cost-labels") os.Exit(1) } + enabledProviders, err := parseProviders(providers) + if err != nil { + setupLog.Error(err, "invalid --providers") + os.Exit(1) + } if err := nebulametrics.InitCost(attribution); err != nil { setupLog.Error(err, "configuring the cost metric") os.Exit(1) @@ -314,7 +326,7 @@ func main() { // startup already sees its provider. The manager's client backs the AWS region // source (regions are read from NodePools at call time, not env), so it is // threaded in; the client is only queried at runtime, after the cache has synced. - registerProviders(context.Background(), mgr.GetClient()) + registerProviders(context.Background(), mgr.GetClient(), enabledProviders) // One shared failover blocklist, written by the virtual kubelet handlers on a // Provision failure and read by the placement controller to skip a candidate @@ -610,17 +622,30 @@ func setupVirtualNodes(mgr ctrl.Manager, blocklist vnode.Blocklister, kubeletSrv // still run for the providers that ARE configured, and a pool referencing an // unregistered provider surfaces as a clear NodePool condition rather than a // crash loop. -func registerProviders(ctx context.Context, c client.Client) { +// +// Only providers in enabled are attempted (see parseProviders). +func registerProviders(ctx context.Context, c client.Client, enabled map[string]bool) { + register := func(name string, build func() (provider.Provider, error)) { + if !enabled[name] { + setupLog.Info("provider disabled by --providers", "provider", name) + return + } + p, err := build() + if err != nil { + setupLog.Info("skipping provider registration", "provider", name, "reason", err.Error()) + return + } + provider.Register(p) + setupLog.Info("registered provider", "provider", p.Name()) + } + appName := os.Getenv("MODAL_APP_NAME") if appName == "" { appName = "nebula" } - if p, err := modal.NewSDKClient(ctx, appName, os.Getenv("MODAL_ENVIRONMENT")); err != nil { - setupLog.Info("skipping Modal provider registration", "reason", err.Error()) - } else { - provider.Register(p) - setupLog.Info("registered provider", "provider", p.Name()) - } + register(provider.ProviderModal, func() (provider.Provider, error) { + return modal.NewSDKClient(ctx, appName, os.Getenv("MODAL_ENVIRONMENT")) + }) // AWS. There is NO region env/flag: the regions this provider may use are declared // per-pool in the NodePool (ProviderSpec.Regions) and read at call time via the @@ -633,12 +658,9 @@ func registerProviders(ctx context.Context, c client.Client) { // delivered via a Secret), and one account-global credential authorizes every // region. Registration only fails (and is a non-fatal skip) if the price catalog // cannot load — region config can no longer make it fail. - if p, err := awsprovider.NewSDKClient(ctx, awsRegionSource(c)); err != nil { - setupLog.Info("skipping AWS provider registration", "reason", err.Error()) - } else { - provider.Register(p) - setupLog.Info("registered provider", "provider", p.Name()) - } + register(provider.ProviderAWS, func() (provider.Provider, error) { + return awsprovider.NewSDKClient(ctx, awsRegionSource(c)) + }) // The fake provider is an in-memory backend used only by the e2e suite to // exercise the full control-plane loop without cloud credentials. It ships in @@ -651,6 +673,27 @@ func registerProviders(ctx context.Context, c client.Client) { } } +// knownProviders are the names --providers accepts, one per register call above. The fake +// provider is not among them: it stays gated on its env var alone. +var knownProviders = []string{provider.ProviderModal, provider.ProviderAWS} + +// parseProviders turns --providers into the enabled set. An unknown name is an error rather +// than ignored, so a typo cannot silently leave a provider off. +func parseProviders(s string) (map[string]bool, error) { + enabled := map[string]bool{} + for _, name := range strings.Split(s, ",") { + name = strings.TrimSpace(name) + if name == "" { + continue + } + if !slices.Contains(knownProviders, name) { + return nil, fmt.Errorf("unknown provider %q, want one of %v", name, knownProviders) + } + enabled[name] = true + } + return enabled, nil +} + // awsRegionSource returns the AWS adapter's RegionSource: ProviderSpec.Regions of every // NodePool referencing the "aws" provider, one entry per pool and unexpanded — the // adapter resolves them, as it does for placement. No env/flag needed — regions are the diff --git a/cmd/main_test.go b/cmd/main_test.go new file mode 100644 index 0000000..c86aa4c --- /dev/null +++ b/cmd/main_test.go @@ -0,0 +1,60 @@ +/* +Copyright 2026 The InftyAI Team. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package main + +import ( + "maps" + "slices" + "strings" + "testing" +) + +func TestParseProviders(t *testing.T) { + cases := []struct { + name string + in string + want []string + wantErr bool + }{ + {name: "default enables all", in: strings.Join(knownProviders, ","), want: knownProviders}, + {name: "subset", in: "aws", want: []string{"aws"}}, + {name: "spaces and empties ignored", in: " modal , ,aws ", want: []string{"aws", "modal"}}, + {name: "empty disables all", in: "", want: nil}, + // A typo must not silently leave a provider off. + {name: "unknown name", in: "modal,moddal", wantErr: true}, + {name: "fake is not selectable", in: "fake", wantErr: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := parseProviders(tc.in) + if tc.wantErr { + if err == nil { + t.Fatalf("parseProviders(%q) = %v, want an error", tc.in, got) + } + return + } + if err != nil { + t.Fatalf("parseProviders(%q): %v", tc.in, err) + } + names := slices.Sorted(maps.Keys(got)) + want := slices.Sorted(slices.Values(tc.want)) + if !slices.Equal(names, want) { + t.Fatalf("parseProviders(%q) = %v, want %v", tc.in, names, want) + } + }) + } +} diff --git a/config/manager/manager.yaml b/config/manager/manager.yaml index 404f25f..743b58a 100644 --- a/config/manager/manager.yaml +++ b/config/manager/manager.yaml @@ -64,6 +64,10 @@ spec: - --leader-elect - --health-probe-bind-address=:8081 # - --kubelet-serving-tls-bootstrap=true + # Every provider is registered by default. Restrict with --providers; it is the + # only way to turn AWS off, which registers even without credentials. Drain a + # provider first: once dropped, nothing terminates its running instances. + # - --providers=modal image: controller:latest name: manager imagePullPolicy: IfNotPresent diff --git a/docs/add-a-provider.md b/docs/add-a-provider.md index 0ebd3ef..67b0202 100644 --- a/docs/add-a-provider.md +++ b/docs/add-a-provider.md @@ -112,19 +112,19 @@ ProviderRunPod = "runpod" ## 4. Wire it into the manager -In `registerProviders` (`cmd/main.go`), build the adapter and register it. A -provider whose credentials are absent must be **logged and skipped, not fatal** — -follow the existing Modal/AWS pattern: +In `registerProviders` (`cmd/main.go`), build the adapter through `register`. A +provider whose credentials are absent must return an error from its constructor, which +is **logged and skipped, not fatal**: ```go -if p, err := runpod.NewSDKClient(ctx); err != nil { - setupLog.Info("skipping RunPod provider registration", "reason", err.Error()) -} else { - provider.Register(p) - setupLog.Info("registered provider", "provider", p.Name()) -} +register(provider.ProviderRunPod, func() (provider.Provider, error) { + return runpod.NewSDKClient(ctx) +}) ``` +Then add its name to `knownProviders`, so `--providers` accepts it and enables it by +default. + ## 5. Wire its credentials Credentials live in **one Kubernetes Secret per provider**, mounted via an optional