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
71 changes: 57 additions & 14 deletions cmd/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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.")
Expand Down Expand Up @@ -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, ","),
Comment thread
kerthcet marked this conversation as resolved.
"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,
}
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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] {
Comment thread
kerthcet marked this conversation as resolved.
setupLog.Info("provider disabled by --providers", "provider", name)
Comment on lines +629 to +630
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
Expand All @@ -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
Expand All @@ -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
Expand Down
60 changes: 60 additions & 0 deletions cmd/main_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}
4 changes: 4 additions & 0 deletions config/manager/manager.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
18 changes: 9 additions & 9 deletions docs/add-a-provider.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading