diff --git a/.env.example b/.env.example index 7a8bcb0..3f749a7 100644 --- a/.env.example +++ b/.env.example @@ -45,6 +45,14 @@ MODAL_ENVIRONMENT= AWS_ACCESS_KEY_ID= AWS_SECRET_ACCESS_KEY= +# --- RunPod provider ------------------------------------------------------- +# One API key, from https://console.runpod.io/user/settings. It needs READ+WRITE on +# Pods: Nebula creates and deletes them, and also creates container-registry-auth +# objects when a workload uses an imagePullSecret. A read-only key registers fine +# and then fails every provision with an auth error, which blocklists the whole +# provider until it is replaced. +RUNPOD_API_KEY= + # --- Additional providers (add as adapters land) --------------------------- -# Each provider gets its OWN secret (see hack/deploy.sh PROVIDER_SECRETS), e.g.: -# RUNPOD_API_KEY= +# Each provider gets its OWN secret; see hack/deploy.sh PROVIDER_SECRETS and the +# RunPod block above for the shape. diff --git a/README.md b/README.md index c8d2181..f81d38e 100644 --- a/README.md +++ b/README.md @@ -62,6 +62,9 @@ metadata: spec: providers: - name: modal # NeoCloud; regions omitted = place anywhere (cheapest) + - name: runpod # NeoCloud, OnDemand only; a region is a geography + regions: # ("us") or one data center ("EU-RO-1") + - us - name: aws # hyperscaler; "us" expands to every US region regions: - us diff --git a/api/v1alpha1/nodepool_types.go b/api/v1alpha1/nodepool_types.go index 7cf369f..8f4921b 100644 --- a/api/v1alpha1/nodepool_types.go +++ b/api/v1alpha1/nodepool_types.go @@ -181,7 +181,7 @@ const ( ) // CapacityType is the purchase model (the outer axis). Each provider maps it to -// its own concept — e.g. RunPod Spot -> interruptible/podRentInterruptable. +// its own concept — e.g. AWS Spot -> a spot-market CreateFleet request. // +kubebuilder:validation:Enum=Spot;OnDemand type CapacityType string diff --git a/cmd/main.go b/cmd/main.go index dbb5da2..c27cf23 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -58,6 +58,7 @@ import ( awsprovider "github.com/InftyAI/Nebula/pkg/provider/aws" "github.com/InftyAI/Nebula/pkg/provider/fake" "github.com/InftyAI/Nebula/pkg/provider/modal" + "github.com/InftyAI/Nebula/pkg/provider/runpod" "github.com/InftyAI/Nebula/pkg/vnode" // +kubebuilder:scaffold:imports ) @@ -544,7 +545,7 @@ func setupKubeletServer(mgr ctrl.Manager, addr, clientCA string, servingTLSBoots // +kubebuilder:rbac:groups=certificates.k8s.io,resources=certificatesigningrequests,resourceNames=nebula-kubelet-serving,verbs=delete;get // +kubebuilder:rbac:groups=certificates.k8s.io,resources=certificatesigningrequests/approval,resourceNames=nebula-kubelet-serving,verbs=update // +kubebuilder:rbac:groups=certificates.k8s.io,resources=signers,resourceNames=kubernetes.io/kubelet-serving,verbs=approve -// +kubebuilder:rbac:groups="",resources=users,resourceNames={"system:node:nebula-aws","system:node:nebula-modal","system:node:nebula-fake"},verbs=impersonate +// +kubebuilder:rbac:groups="",resources=users,resourceNames={"system:node:nebula-aws","system:node:nebula-modal","system:node:nebula-runpod","system:node:nebula-fake"},verbs=impersonate // +kubebuilder:rbac:groups="",resources=groups,resourceNames="system:nodes",verbs=impersonate // addServingCertificateBootstrap requests a trusted serving certificate for the kubelet @@ -647,21 +648,14 @@ func registerProviders(ctx context.Context, c client.Client, enabled map[string] 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 - // region source below, so a pool added at runtime widens the fan-out without a - // restart. One AWS provider spans every such region (per-region clients are built - // lazily). The adapter is otherwise self-configuring: it resolves each region's - // GPU AMI and default-VPC subnets itself, so no launch template or pre-created - // infra is needed. Credentials are secrets and are NEVER read here: the SDK client - // uses the default credential chain (IRSA / instance-role / AWS_ACCESS_KEY_ID - // 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. register(provider.ProviderAWS, func() (provider.Provider, error) { return awsprovider.NewSDKClient(ctx, awsRegionSource(c)) }) + register(provider.ProviderRunPod, func() (provider.Provider, error) { + return runpod.NewSDKClient(ctx) + }) + // 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 // the binary but registers ONLY when explicitly enabled, so it can never place @@ -675,7 +669,7 @@ func registerProviders(ctx context.Context, c client.Client, enabled map[string] // 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} +var knownProviders = []string{provider.ProviderModal, provider.ProviderAWS, provider.ProviderRunPod} // 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. diff --git a/cmd/main_test.go b/cmd/main_test.go index c86aa4c..cf2091c 100644 --- a/cmd/main_test.go +++ b/cmd/main_test.go @@ -18,9 +18,15 @@ package main import ( "maps" + "os" "slices" "strings" "testing" + + rbacv1 "k8s.io/api/rbac/v1" + "sigs.k8s.io/yaml" + + "github.com/InftyAI/Nebula/pkg/vnode" ) func TestParseProviders(t *testing.T) { @@ -58,3 +64,30 @@ func TestParseProviders(t *testing.T) { }) } } + +// TestImpersonateGrantCoversKnownProviders guards the serving-certificate bootstrap: it +// impersonates whichever provider registers first, and a missing grant only surfaces at +// runtime as a Forbidden retry loop. Reads the generated role, so a stale `make manifests` +// fails too. +func TestImpersonateGrantCoversKnownProviders(t *testing.T) { + raw, err := os.ReadFile("../config/rbac/role.yaml") + if err != nil { + t.Fatal(err) + } + var role rbacv1.ClusterRole + if err := yaml.Unmarshal(raw, &role); err != nil { + t.Fatal(err) + } + var users []string + for _, r := range role.Rules { + if slices.Contains(r.Resources, "users") && slices.Contains(r.Verbs, "impersonate") { + users = append(users, r.ResourceNames...) + } + } + for _, name := range knownProviders { + if id := vnode.NodeIdentity(vnode.NodeName(name)); !slices.Contains(users, id) { + t.Errorf("role.yaml grants no impersonate on %s; add it to the marker in main.go "+ + "and run `make manifests`", id) + } + } +} diff --git a/config/catalog/kustomization.yaml b/config/catalog/kustomization.yaml index 39d2f64..2862055 100644 --- a/config/catalog/kustomization.yaml +++ b/config/catalog/kustomization.yaml @@ -17,6 +17,7 @@ configMapGenerator: files: - modal.csv=../../pkg/provider/catalog/data/modal.csv - aws.csv=../../pkg/provider/catalog/data/aws.csv + - runpod.csv=../../pkg/provider/catalog/data/runpod.csv generatorOptions: # Stable name (no content-hash suffix) so `kubectl edit` and the volume diff --git a/config/crd/bases/nebula.inftyai.com_nodepools.yaml b/config/crd/bases/nebula.inftyai.com_nodepools.yaml index 144fb6d..5b097f0 100644 --- a/config/crd/bases/nebula.inftyai.com_nodepools.yaml +++ b/config/crd/bases/nebula.inftyai.com_nodepools.yaml @@ -90,7 +90,7 @@ spec: items: description: |- CapacityType is the purchase model (the outer axis). Each provider maps it to - its own concept — e.g. RunPod Spot -> interruptible/podRentInterruptable. + its own concept — e.g. AWS Spot -> a spot-market CreateFleet request. enum: - Spot - OnDemand diff --git a/config/manager/manager.yaml b/config/manager/manager.yaml index 743b58a..0543b5a 100644 --- a/config/manager/manager.yaml +++ b/config/manager/manager.yaml @@ -67,7 +67,7 @@ spec: # 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 + # - --providers=runpod image: controller:latest name: manager imagePullPolicy: IfNotPresent @@ -131,10 +131,14 @@ spec: - secretRef: name: nebula-aws-credentials optional: true - # Add one secretRef per provider as adapters land, e.g.: - # - secretRef: - # name: nebula-runpod-credentials - # optional: true + # RunPod: a single API key (RUNPOD_API_KEY). Unlike AWS there is no ambient + # identity to fall back on, so an absent Secret means the provider is simply + # skipped at registration. Regions come from the NodePool, so this is the only + # RunPod config here. + - secretRef: + name: nebula-runpod-credentials + optional: true + # Add one secretRef per provider as adapters land, following the pattern above. ports: # The kubelet API the API server dials for `kubectl logs` (10250, like a real # kubelet). Declaring it is documentation and NetworkPolicy surface; the diff --git a/config/rbac/role.yaml b/config/rbac/role.yaml index 8706076..cce39fc 100644 --- a/config/rbac/role.yaml +++ b/config/rbac/role.yaml @@ -66,6 +66,7 @@ rules: - system:node:nebula-aws - system:node:nebula-fake - system:node:nebula-modal + - system:node:nebula-runpod resources: - users verbs: diff --git a/config/samples/deployment.yaml b/config/samples/deployment.yaml index c399a98..d061546 100644 --- a/config/samples/deployment.yaml +++ b/config/samples/deployment.yaml @@ -39,7 +39,7 @@ spec: app: gpu-workload-sample nebula.inftyai.com/enabled: "true" nebula.inftyai.com/nodepool: sample - nebula.inftyai.com/accelerator-type: t4 + nebula.inftyai.com/accelerator-type: l4 spec: # Do NOT set nodeName or a provider nodeSelector yourself — the placement # controller fills the nodeSelector in when it ungates the Pod. Setting diff --git a/config/samples/nodepool.yaml b/config/samples/nodepool.yaml index f68529f..a24f100 100644 --- a/config/samples/nodepool.yaml +++ b/config/samples/nodepool.yaml @@ -19,7 +19,10 @@ spec: - us - eu - ap-melbourne - # - name: runpod + - name: runpod + regions: + - us + - eu # Outer axis: try OnDemand on every provider first, fall back to Spot. capacityTypes: - OnDemand @@ -32,6 +35,11 @@ spec: # sandbox stays reachable through its connect URL and token under every mode — it just # cannot call out. # + # Setting ANY mode other than Open narrows this pool to the providers that can enforce + # it: placement skips a provider whose Capabilities report SupportsEgressPolicy=false + # (RunPod, which exposes no outbound knob at all), rather than provisioning something + # with open internet access under a policy that says otherwise. + # # Blocked permits nothing: # egress: # mode: Blocked diff --git a/docs/deploy.md b/docs/deploy.md index 807adb4..ba53f71 100644 --- a/docs/deploy.md +++ b/docs/deploy.md @@ -118,6 +118,7 @@ itself at startup (see [Webhook TLS](#webhook-tls-no-cert-manager)). | `MODAL_ENVIRONMENT` | Modal | no | Modal Environment to create sandboxes in. Blank omits the key and the SDK uses the token profile's default. See [Modal Environments](#modal-environments). | | `AWS_ACCESS_KEY_ID` | AWS | dev only | Prefer IRSA / instance role in production and leave blank — the SDK's default credential chain finds the role. Set only for local/dev. | | `AWS_SECRET_ACCESS_KEY` | AWS | dev only | Pairs with `AWS_ACCESS_KEY_ID`; both required together or both blank. | +| `RUNPOD_API_KEY` | RunPod | yes | From [console.runpod.io/user/settings](https://console.runpod.io/user/settings). Needs **read+write on Pods**: a read-only key registers fine and then fails every provision with an auth error, which blocklists the whole provider. Unlike AWS there is no ambient identity, so blank skips RunPod entirely. | Non-secret config, passed as `make` variables: diff --git a/docs/status.md b/docs/status.md index ac469e9..ec6b650 100644 --- a/docs/status.md +++ b/docs/status.md @@ -25,6 +25,7 @@ enters the system. - [Provider mappings](#provider-mappings) - [AWS](#aws) - [Modal](#modal) + - [RunPod](#runpod) - [fake](#fake) - [Logs and exec](#logs-and-exec) - [What is not observable](#what-is-not-observable) @@ -81,8 +82,9 @@ therefore emits without storing, and stores only once `Provision` returns. `Provision` returns `(id, reserved, error)`. `reserved` means the provider committed capacity, not merely accepted the request: AWS always does (`CreateFleet` with -`FleetTypeInstant` is synchronous), a fresh Modal sandbox never does (the GPU may -still be queued). Only a reserved instance advances to `Initializing`; an unreserved +`FleetTypeInstant` is synchronous), RunPod always does for the same reason (`POST +/pods` allocates a host before it answers, and a shortage comes back as an error), a +fresh Modal sandbox never does (the GPU may still be queued). Only a reserved instance advances to `Initializing`; an unreserved id holds at `Provisioning`, which is still exactly true — the id is real and must be reclaimed, but nothing is allocated. Either way the Pod is now tracked: `reserved` constrains what the status may claim, not what is owed, and says nothing about @@ -234,6 +236,44 @@ only two signals and has to record a third fact itself. first poll tick. An *adopted* sandbox has been observed, so a `running` one is known to be reserved. See below for why queued is not reported distinctly. +### RunPod + +One RunPod Pod per NodeClaim, read through REST v2's `status`, which (unlike v1's +`desiredStatus`) is the state the Pod has *reached*. + +| `status` | `InstanceState` | Pod | +|---|---|---| +| `RUNNING` | `Running` | `Running` / `Ready=True` | +| `PROVISIONING`, `STARTING` | `Pending` | `Pending` / `Initializing` | +| `EXITED`, `TERMINATED` | `Terminated` | `Failed` / `Terminated` | +| `ERROR` | `Failed` | `Failed` | +| anything else | `Pending` | `Pending` / `Initializing` | +| absent from `List` | `Terminated` | `Failed` / `Terminated` | + +- **There is no readiness concept beyond `RUNNING`.** RunPod has no probe, so "started" + is the strongest signal available; a container that is up but not yet serving reads + `Running`. Contrast Modal, which has a real probe and latches it. +- **There is no queueing**, as with AWS: `POST /v2/pods` allocates a host before it + answers, and a shortfall is a synchronous error (`ErrNoCapacity`) that drives region + failover. So `Provision` always returns `reserved`. +- **`EXITED` hides crashes.** It covers a clean exit and a crash alike, with no exit + code, so a workload that died reads as `Terminated`, indistinguishable from teardown. +- **OnDemand only.** v2 has no interruptible tier, so nothing is reclaimed and the + default poll cadence applies. +- **Identity rides the Pod name**, not tags: RunPod Pods have none, so a Pod is named + after its claim (`-`) and `List` reads the name back as the claim. + There is no ownership marker, so **the account must be dedicated to Nebula**: a Pod + someone else names like a claim is adopted and later terminated. A Pod whose name + would exceed RunPod's 191-character cap is refused at `Provision` rather than + truncated — two truncated claims would collide onto one Pod. +- The endpoint is **derived, not read back**: `https://-.proxy.runpod.net` + is known at create time, so it is published from `CreatePod` like Modal's, but with + no token — that proxy is unauthenticated. Every TCP port is exposed as `/http`, so a + raw-TCP service is not reachable through it; UDP and SCTP ports are not exposed. +- **Neither `kubectl logs` nor `kubectl exec` works yet.** v2 streams logs over SSE + (`/v2/pods/{id}/logs`), which a `LogStreamer` could wrap; the only way into a + container is SSH, so exec answers NotFound. + ### fake The in-memory e2e provider reports `InstanceRunning` as soon as an instance is diff --git a/hack/deploy.sh b/hack/deploy.sh index e79d45b..6c4014b 100755 --- a/hack/deploy.sh +++ b/hack/deploy.sh @@ -102,7 +102,12 @@ PROVIDER_SECRETS=( # instance role (the preferred path). Region is NON-SECRET (on the manager # Deployment); the adapter self-configures the rest (GPU AMI + subnets). "nebula-aws-credentials|AWS_ACCESS_KEY_ID AWS_SECRET_ACCESS_KEY|" - # "nebula-runpod-credentials|RUNPOD_API_KEY|" + # RunPod: one API key, and the only credential it has — there is no ambient identity to + # fall back on as AWS has, so a blank key skips the Secret AND the provider. Mint it at + # https://console.runpod.io/user/settings with read+write on Pods; a read-only key + # registers fine and then fails every create with an auth error, which blocklists the + # whole provider. + "nebula-runpod-credentials|RUNPOD_API_KEY|" ) # create_provider_secret diff --git a/internal/controller/nodeclaim_controller_test.go b/internal/controller/nodeclaim_controller_test.go index 292d391..abaa6f1 100644 --- a/internal/controller/nodeclaim_controller_test.go +++ b/internal/controller/nodeclaim_controller_test.go @@ -52,13 +52,14 @@ type fakeProvider struct { gpus []string // accelerators MapAccelerator offers; nil = offer any spot bool // Capabilities().SupportsSpot (placement skips Spot without it) egress bool // Capabilities().SupportsEgressPolicy (placement skips restricted pools without it) + gpuOnly bool // !Capabilities().SupportsCPUOnly (placement skips CPU-only Pods) // expandRegions overrides ResolveRegions; nil = pass the declaration through. expandRegions func([]string) []string } func (f *fakeProvider) Name() string { return f.name } func (f *fakeProvider) Capabilities() provider.Capabilities { - return provider.Capabilities{SupportsSpot: f.spot, SupportsEgressPolicy: f.egress} + return provider.Capabilities{SupportsSpot: f.spot, SupportsEgressPolicy: f.egress, SupportsCPUOnly: !f.gpuOnly} } func (f *fakeProvider) Provision(context.Context, *corev1.Pod, provider.ProvisionRequest) (provider.ProvisionResult, error) { return provider.ProvisionResult{}, nil diff --git a/internal/controller/placement_metrics_test.go b/internal/controller/placement_metrics_test.go index 4e4cb71..de555fd 100644 --- a/internal/controller/placement_metrics_test.go +++ b/internal/controller/placement_metrics_test.go @@ -32,7 +32,6 @@ import ( "github.com/InftyAI/Nebula/pkg/failover" "github.com/InftyAI/Nebula/pkg/metrics" "github.com/InftyAI/Nebula/pkg/provider" - awsprovider "github.com/InftyAI/Nebula/pkg/provider/aws" "github.com/InftyAI/Nebula/pkg/util" ) @@ -247,10 +246,10 @@ func TestPlacement_NarrowingToNoRegionFilesNoAvailableRegions(t *testing.T) { noRegions := skipLabels(provider.ProviderAWS, nebulav1alpha1.CapacityOnDemand, "", metrics.SkipNoAvailableRegions) before := counterVal(t, metrics.CandidateSkips, noRegions) - pod := gatedPod("r1", "default", "uid-r1", "pool", "") + pod := gatedPod("r1", "default", "uid-r1", "pool", "T4") pod.Annotations = map[string]string{nebulav1alpha1.RegionsAnnotation: "af"} pool := poolWith("pool", []nebulav1alpha1.CapacityType{nebulav1alpha1.CapacityOnDemand}, provider.ProviderAWS) - r, _ := newPlacementReconciler(t, []client.Object{pod, pool}, awsprovider.New(nil, nil, nil)) + r, _ := newPlacementReconciler(t, []client.Object{pod, pool}, catalogAWS(t)) reconcilePod(t, r, "default", "r1") if got := counterVal(t, metrics.CandidateSkips, noRegions) - before; got != 1 { diff --git a/internal/controller/pod_placement_controller.go b/internal/controller/pod_placement_controller.go index d0b1bd1..533c882 100644 --- a/internal/controller/pod_placement_controller.go +++ b/internal/controller/pod_placement_controller.go @@ -152,8 +152,6 @@ func (r *PodPlacementReconciler) Reconcile(ctx context.Context, req ctrl.Request "pod", pod.Name, "pool", pool.Name, "retryAfter", retryAfter.String()) return ctrl.Result{RequeueAfter: retryAfter}, nil } - log.Info("no provider in pool can serve the Pod; leaving it gated", - "pod", pod.Name, "pool", pool.Name) return ctrl.Result{}, nil } diff --git a/internal/controller/pod_placement_controller_test.go b/internal/controller/pod_placement_controller_test.go index 98b5696..0da7557 100644 --- a/internal/controller/pod_placement_controller_test.go +++ b/internal/controller/pod_placement_controller_test.go @@ -36,6 +36,7 @@ import ( "github.com/InftyAI/Nebula/pkg/failover" "github.com/InftyAI/Nebula/pkg/provider" awsprovider "github.com/InftyAI/Nebula/pkg/provider/aws" + "github.com/InftyAI/Nebula/pkg/provider/catalog" "github.com/InftyAI/Nebula/pkg/util" ) @@ -244,8 +245,9 @@ func TestPlacement_NoMatchingProviderLeavesPodGated(t *testing.T) { } } -func TestPlacement_CPUOnlyPodMatchesAnyProvider(t *testing.T) { - // No GPU annotation => any provider matches; even one offering nothing. +func TestPlacement_CPUOnlyPodMatchesAnyCPUOnlyProvider(t *testing.T) { + // No GPU annotation => any provider that runs CPU-only Pods matches; even one offering + // no GPUs. pod := gatedPod("p1", "default", "uid-1", "pool-a", "") pool := poolWith("pool-a", []nebulav1alpha1.CapacityType{nebulav1alpha1.CapacityOnDemand}, provider.ProviderModal) modal := &fakeProvider{name: provider.ProviderModal, gpus: []string{}} // offers no GPUs @@ -267,6 +269,23 @@ func TestPlacement_CPUOnlyPodMatchesAnyProvider(t *testing.T) { } } +func TestPlacement_CPUOnlyPodSkipsGPUOnlyProvider(t *testing.T) { + // RunPod is listed first but runs GPU Pods only, so the CPU-only Pod lands on Modal. + pod := gatedPod("p1", "default", "uid-1", "pool-a", "") + pool := poolWith("pool-a", []nebulav1alpha1.CapacityType{nebulav1alpha1.CapacityOnDemand}, + provider.ProviderRunPod, provider.ProviderModal) + runpod := &fakeProvider{name: provider.ProviderRunPod, gpuOnly: true} + modal := &fakeProvider{name: provider.ProviderModal} + r, c := newPlacementReconciler(t, []client.Object{pod, pool}, runpod, modal) + + reconcilePod(t, r, "default", "p1") + + got := getPod(t, c, "default", "p1") + if got.Spec.NodeSelector[nebulav1alpha1.ProviderLabel] != provider.ProviderModal { + t.Fatalf("expected the CPU-only Pod to skip runpod for modal, got %v", got.Spec.NodeSelector) + } +} + func TestPlacement_GPUCountWithoutAcceleratorTypeLeavesPodGated(t *testing.T) { // A Pod that requests nvidia.com/gpu but omits accelerator-type is malformed, // not CPU-only. Placement must not silently route it as a CPU-only workload. @@ -572,19 +591,28 @@ func TestRequestedGeographies(t *testing.T) { } } +// catalogAWS is the real AWS adapter with no client: enough for placement, which reads only +// its catalog and region table. AWS runs GPU Pods only, so a test Pod must request one. +func catalogAWS(t *testing.T) *awsprovider.Provider { + t.Helper() + cat, err := catalog.Load() + if err != nil { + t.Fatalf("loading the embedded catalog: %v", err) + } + return awsprovider.New(nil, cat, nil) +} + func TestPlacement_PodAnnotationNarrowsToTheRequestedJurisdiction(t *testing.T) { // The pool is unconstrained, so AWS offers all 17 default-enabled regions. The Pod asks // for "uk", which AWS serves from London alone — so the claim must carry eu-west-2 and // not the first region of the walk. This is the whole point of the annotation: data // residency for ONE workload, without an operator carving out a per-jurisdiction pool. // - // CPU-only on purpose: it keeps MapAccelerator (and so the catalog) out of the path, so - // the real adapter's region table can be exercised with no client and no CSV. - pod := gatedPod("p1", "default", "uid-1", "pool-a", "") + pod := gatedPod("p1", "default", "uid-1", "pool-a", "T4") pod.Annotations = map[string]string{nebulav1alpha1.RegionsAnnotation: "uk"} pool := poolWith("pool-a", []nebulav1alpha1.CapacityType{nebulav1alpha1.CapacityOnDemand}, provider.ProviderAWS) - r, c := newPlacementReconciler(t, []client.Object{pod, pool}, awsprovider.New(nil, nil, nil)) + r, c := newPlacementReconciler(t, []client.Object{pod, pool}, catalogAWS(t)) reconcilePod(t, r, "default", "p1") diff --git a/internal/controller/pod_placement_helpers.go b/internal/controller/pod_placement_helpers.go index 8f9a1c1..1dc9662 100644 --- a/internal/controller/pod_placement_helpers.go +++ b/internal/controller/pod_placement_helpers.go @@ -117,6 +117,7 @@ func (r *PodPlacementReconciler) selectPlacement(ctx context.Context, pod *corev } var soonest time.Duration // 0 = no blocked-but-servable candidate seen + skipped := map[string]string{} // provider -> why, for the exhausted-walk log for _, tier := range capacityTiers(pool) { // outer: capacity for _, ref := range pool.Spec.Providers { // provider (Ordered = listed order) prov, ok := r.provider(ref.Name) @@ -124,24 +125,27 @@ func (r *PodPlacementReconciler) selectPlacement(ctx context.Context, pod *corev // No region on this skip and the two below: they are decided before the // walk reaches the region axis, so they rule out every region at once. metrics.RecordCandidateSkip(ref.Name, tier, "", metrics.SkipProviderUnregistered) + skipped[ref.Name] = metrics.SkipProviderUnregistered log.V(1).Info("skipping candidate: provider not registered", "provider", ref.Name, "capacityType", tier) continue // unregistered; NodePool status surfaces this separately } if !prov.Capabilities().ServesCapacityTier(tier) { metrics.RecordCandidateSkip(ref.Name, tier, "", metrics.SkipCapacityUnsupported) + skipped[ref.Name] = metrics.SkipCapacityUnsupported log.V(1).Info("skipping candidate: provider does not offer the capacity tier", "provider", ref.Name, "capacityType", tier) continue } if !prov.Capabilities().ServesEgress(pool.Spec.Egress) { metrics.RecordCandidateSkip(ref.Name, tier, "", metrics.SkipEgressUnsupported) + skipped[ref.Name] = metrics.SkipEgressUnsupported log.V(1).Info("skipping candidate: provider cannot enforce the pool's egress policy", "provider", ref.Name, "egressMode", pool.Spec.Egress.ModeOrOpen()) continue } - // A CPU-only Pod (no accelerator) matches any provider; an accelerator - // Pod only matches a provider whose catalog serves that (type, count). + // A CPU-only Pod (no accelerator) matches a provider that SupportsCPUOnly; an + // accelerator Pod only matches a provider whose catalog serves that (type, count). // MapAccelerator is consulted only for that servability check — the block // key and the reported identity are the POOL (type:count), not the // provider's SKU, so a launch spanning alternates and a post-launch SKU @@ -150,16 +154,23 @@ func (r *PodPlacementReconciler) selectPlacement(ctx context.Context, pod *corev if accel != "" { if _, offered := prov.MapAccelerator(accel, count); !offered { metrics.RecordCandidateSkip(ref.Name, tier, "", metrics.SkipAcceleratorUnsupported) + skipped[ref.Name] = metrics.SkipAcceleratorUnsupported log.V(1).Info("skipping candidate: provider does not offer the accelerator", "provider", ref.Name, "accelerator", accel, "count", count) continue } + } else if !prov.Capabilities().SupportsCPUOnly { + metrics.RecordCandidateSkip(ref.Name, tier, "", metrics.SkipAcceleratorUnsupported) + skipped[ref.Name] = metrics.SkipAcceleratorUnsupported + log.V(1).Info("skipping candidate: provider does not run CPU-only Pods", "provider", ref.Name) + continue } // Empty means the pool's declaration, or the Pod's narrowing of it, reaches // no region this provider can place in. regions := prov.ResolveRegions(ref.Regions, narrowTo) if len(regions) == 0 { metrics.RecordCandidateSkip(ref.Name, tier, "", metrics.SkipNoAvailableRegions) + skipped[ref.Name] = metrics.SkipNoAvailableRegions log.V(1).Info("skipping candidate: no available region serves the requested geographies", "provider", ref.Name, "capacityType", tier, "regions", narrowTo) continue @@ -196,6 +207,8 @@ func (r *PodPlacementReconciler) selectPlacement(ctx context.Context, pod *corev metrics.RecordDeferral(pool.Name, metrics.DeferAllBlocked) } else { metrics.RecordDeferral(pool.Name, metrics.DeferNoCandidate) + log.Info("no provider in pool can serve the Pod; leaving it gated", + "accelerator", util.AcceleratorPool(accel, count), "skipped", skipped) } return placement{}, false, soonest } diff --git a/pkg/provider/aws/aws.go b/pkg/provider/aws/aws.go index c50e8a3..7eac846 100644 --- a/pkg/provider/aws/aws.go +++ b/pkg/provider/aws/aws.go @@ -407,6 +407,7 @@ func (p *Provider) Capabilities() provider.Capabilities { // which routes to an internet gateway. Enforcing a pool's policy means managing SG // egress rules (and no NAT for the Blocked case), so it is unsupported until then. SupportsEgressPolicy: false, + SupportsCPUOnly: false, // the catalog is GPU instance types only NativeTags: true, // EC2 tags carry identity PreemptionNotice: preemptionNotice, // Spot 2-minute warning PollInterval: spotPollInterval, // Spot reclaims are abrupt; poll faster than default diff --git a/pkg/provider/catalog/data/pricing.go b/pkg/provider/catalog/data/pricing.go index 3265563..867dc68 100644 --- a/pkg/provider/catalog/data/pricing.go +++ b/pkg/provider/catalog/data/pricing.go @@ -69,6 +69,11 @@ func ModalMemoryCostPerHour(memoryMiB int) float64 { // regions differ by a few cents per GB-month. const AWSGP3PricePerGBHour = 0.08 / hoursPerMonth +// RunPodContainerDiskPricePerGBHour is the pricing page's $0.10/GB/month for a running Pod's +// container disk, spread over an average month. It is RunPod's only extra on a GPU Pod: vCPU +// and RAM are bundled into the GPU price. +const RunPodContainerDiskPricePerGBHour = 0.10 / hoursPerMonth + // hoursPerMonth is 365 days / 12, the conversion for a rate quoted per month. const hoursPerMonth = 730 @@ -77,3 +82,9 @@ const hoursPerMonth = 730 func AWSRootVolumeCostPerHour(diskGiB int) float64 { return float64(diskGiB) * AWSGP3PricePerGBHour } + +// RunPodContainerDiskCostPerHour takes the disk size the Pod is created with, not its raw +// request. +func RunPodContainerDiskCostPerHour(diskGB int) float64 { + return float64(diskGB) * RunPodContainerDiskPricePerGBHour +} diff --git a/pkg/provider/catalog/data/runpod.csv b/pkg/provider/catalog/data/runpod.csv new file mode 100644 index 0000000..a11d66b --- /dev/null +++ b/pkg/provider/catalog/data/runpod.csv @@ -0,0 +1,54 @@ +# RunPod price/availability catalog — community-maintained. +# +# Prices are `price.secure` from GET https://api.runpod.io/v2/catalog/gpus, copied by +# hand: the API has no canonical-name mapping, so there is no refresh tool. Treat them +# as a starting point, not a billing source. Only list an id whose `secure` is true. +# +# Prices are per GPU-HOUR, as Modal's are — not per instance-hour like aws.csv. +# RunPod bills per GPU, and the adapter passes the count as a runtime parameter. +# +# Only SECURE cloud is priced here. RunPod's COMMUNITY cloud rents the same GPUs +# from peer hosts for roughly 30-50% less, but this file has no cloud-type column +# to say which tier a row prices, so listing both would make the number the +# optimizer reads a coin flip. The adapter pins cloudType=SECURE to match; adding +# COMMUNITY means adding that column first (see pkg/provider/runpod's package doc). +# +# Columns (shared header across all provider CSVs; unused cells left blank): +# accelerator_type canonical Nebula accelerator type, matched case-insensitively +# against the nebula.inftyai.com/accelerator-type label +# accelerator_id ALWAYS SET here, unlike modal.csv: RunPod's ids are marketing +# strings ("NVIDIA H100 80GB HBM3") that share nothing with the +# canonical names, so every row carries its own translation. +# +# SEVERAL rows may share one accelerator_type, and the order +# matters: MapAccelerator returns them in file order, so the +# FIRST is the primary and the rest are interchangeable +# alternates. A REST v2 create takes ONE gpu id, so only the +# primary is launched; an alternate matters only once the rows +# above it are flipped to available=false. Put the variant with +# the best interconnect first (SXM before NVL before PCIe). +# gpu_count BLANK. RunPod takes the GPU count as a request parameter, so it +# is not a lookup dimension (contrast aws.csv, where the count is +# baked into the instance type). A blank row matches any count. +# capacity_type OnDemand only. REST v2 has no interruptible tier, so a Spot row +# would place Pods the adapter cannot launch. +# price_per_hour approximate USD per GPU-hour on SECURE cloud +# available whether Nebula may schedule onto it. Flipping a row to false +# removes it from placement everywhere without touching Go — the +# escape hatch for a GPU type RunPod has stopped offering. +# region BLANK. RunPod's prices are not partitioned by data center, so +# one row prices every region. Region is still a real placement +# axis for this provider (a pool's regions become dataCenterIds); +# it is just not a pricing one. +# updated documentation only, ignored by the parser +accelerator_type,accelerator_id,gpu_count,capacity_type,price_per_hour,available,region,updated +L4,NVIDIA L4,,OnDemand,0.49,true,,2026-10-04 +A40,NVIDIA A40,,OnDemand,0.49,true,,2026-10-04 +RTX4090,NVIDIA GeForce RTX 4090,,OnDemand,0.74,true,,2026-10-04 +L40S,NVIDIA L40S,,OnDemand,1.09,true,,2026-10-04 +A100-80GB,NVIDIA A100-SXM4-80GB,,OnDemand,1.59,true,,2026-10-04 +A100-80GB,NVIDIA A100 80GB PCIe,,OnDemand,1.59,true,,2026-10-04 +H100,NVIDIA H100 80GB HBM3,,OnDemand,3.49,true,,2026-10-04 +H100,NVIDIA H100 NVL,,OnDemand,3.19,true,,2026-10-04 +H100,NVIDIA H100 PCIe,,OnDemand,2.89,true,,2026-10-04 +H200,NVIDIA H200,,OnDemand,4.59,true,,2026-10-04 diff --git a/pkg/provider/errors.go b/pkg/provider/errors.go index aba8213..e3337b8 100644 --- a/pkg/provider/errors.go +++ b/pkg/provider/errors.go @@ -22,6 +22,7 @@ import ( "strings" nebulav1alpha1 "github.com/InftyAI/Nebula/api/v1alpha1" + "github.com/InftyAI/Nebula/pkg/util" ) // Provision failure categories, shared by every adapter. The CATEGORIES are @@ -169,7 +170,7 @@ func categorize(err error) failureCategory { // — a manager shutdown, a leader handoff — and the provider may well have accepted the // request. Nothing about the candidate was learned, so blocklisting would punish it for // our own exit. Only the block scope; the caller still fails the Pod. - if errors.Is(err, context.Canceled) || containsAny(msg, "context canceled") { + if errors.Is(err, context.Canceled) || util.ContainsAny(msg, "context canceled") { return catUnattributable } @@ -189,7 +190,7 @@ func categorize(err error) failureCategory { // - "image build for": the Modal SDK's remote build verdict, the only thing it says when // the build itself reached a decision. An API error from the same call does NOT carry // it, which is what keeps an expired workspace token classifiable as auth below. - if containsAny(msg, "image pull credential", "image build for") { + if util.ContainsAny(msg, "image pull credential", "image build for") { return catRequest } @@ -198,14 +199,14 @@ func categorize(err error) failureCategory { // while every replacement Pod retried the same broken provider. if strings.Contains(msg, "rpc error") { switch { - case containsAny(msg, "code = unauthenticated", "code = permissiondenied"): + case util.ContainsAny(msg, "code = unauthenticated", "code = permissiondenied"): return catAuth case strings.Contains(msg, "code = resourceexhausted"): return catCapacity } } - if containsAny(msg, + if util.ContainsAny(msg, "rpc error", "connection refused", "connection reset", "broken pipe", "no such host", "i/o timeout", "eof", "tls handshake", "service unavailable", "bad gateway", "gateway timeout", "internal server error") { @@ -213,24 +214,14 @@ func categorize(err error) failureCategory { } switch { - case containsAny(msg, "unauthorized", "forbidden", "authentication", + case util.ContainsAny(msg, "unauthorized", "forbidden", "authentication", "unauthenticated", "invalid token", "api key"): return catAuth - case containsAny(msg, "quota", "limit exceeded", "rate limit"): + case util.ContainsAny(msg, "quota", "limit exceeded", "rate limit"): return catCapacity - case containsAny(msg, "no capacity", "capacity", "unavailable", "out of", "no gpu"): + case util.ContainsAny(msg, "no capacity", "capacity", "unavailable", "out of", "no gpu"): return catCapacity default: return catUnattributable } } - -// containsAny reports whether s contains any of subs. -func containsAny(s string, subs ...string) bool { - for _, sub := range subs { - if strings.Contains(s, sub) { - return true - } - } - return false -} diff --git a/pkg/provider/fake/fake.go b/pkg/provider/fake/fake.go index 062e9cb..1a091e8 100644 --- a/pkg/provider/fake/fake.go +++ b/pkg/provider/fake/fake.go @@ -86,6 +86,7 @@ func (p *Provider) Capabilities() provider.Capabilities { SupportsStop: false, SupportsSpot: false, SupportsEgressPolicy: false, // nothing to enforce against; egress pools skip it + SupportsCPUOnly: true, NativeTags: true, PreemptionNotice: 0, PollInterval: 0, // use the vnode default cadence diff --git a/pkg/provider/modal/client.go b/pkg/provider/modal/client.go index b84e30f..f5ab5a4 100644 --- a/pkg/provider/modal/client.go +++ b/pkg/provider/modal/client.go @@ -714,7 +714,7 @@ func (c *sdkClient) ListSandboxes(ctx context.Context) ([]Sandbox, error) { } // FindSandbox implements Client. Modal's Tags filter matches exact key=value pairs, which is -// useless for ListSandboxes (every claim value differs) but exactly one claim's lookup. +// useless for ListSandboxes (every claim value differs) but is ideal for one claim's lookup. func (c *sdkClient) FindSandbox(ctx context.Context, claimName string) (*Sandbox, error) { app, err := c.app(ctx) if err != nil { diff --git a/pkg/provider/modal/modal.go b/pkg/provider/modal/modal.go index 4470fbe..cde6171 100644 --- a/pkg/provider/modal/modal.go +++ b/pkg/provider/modal/modal.go @@ -435,9 +435,10 @@ func (p *Provider) PricePerHour(req provider.PriceRequest) (float64, error) { // trait is set the way it is. func (p *Provider) Capabilities() provider.Capabilities { return provider.Capabilities{ - SupportsStop: false, // create/terminate only - SupportsSpot: false, // no user-facing preemptible tier - SupportsEgressPolicy: true, // outbound allowlists on the sandbox itself + SupportsStop: false, // create/terminate only + SupportsSpot: false, // no user-facing preemptible tier + SupportsEgressPolicy: true, // outbound allowlists on the sandbox itself + SupportsCPUOnly: true, NativeTags: true, // sandbox tags carry identity PreemptionNotice: 0, // no push; poll-based detection PollInterval: 0, // OnDemand-only (never preempts) → the default cadence is fine diff --git a/pkg/provider/provider.go b/pkg/provider/provider.go index 37e4e7f..529dccb 100644 --- a/pkg/provider/provider.go +++ b/pkg/provider/provider.go @@ -220,10 +220,10 @@ type ProvisionRequest struct { // nowhere to live on the Pod. CapacityType nebulav1alpha1.CapacityType // Region is the ONE candidate placement chose, exactly as this provider's own - // ResolveRegions minted it — already resolved (never a group token) and opaque to the - // control plane. Usually one concrete region (AWS "us-east-1"), which is what lets a - // capacity failure blocklist just that region; a provider that cannot fail over - // (Modal) may encode several for its own scheduler, and only that adapter parses it. + // ResolveRegions minted it, opaque to the control plane: only that adapter parses it. + // Usually one concrete region (AWS "us-east-1"), which is what lets a capacity failure + // blocklist just that region; Modal encodes several for its own scheduler, and RunPod + // keeps a geography token that it expands to data centers at create time. // // Empty means "no region constraint". Region string @@ -321,6 +321,8 @@ type Capabilities struct { // a policy that says otherwise (AWS: false — its instances land in the default VPC, so // enforcement needs security-group egress rules and no NAT, not one API field). SupportsEgressPolicy bool + // SupportsCPUOnly is true if the provider runs a Pod with no accelerator. + SupportsCPUOnly bool // NativeTags is true if the provider has real instance tags/labels; when // false, identity is encoded in the instance name (RunPod: false). NativeTags bool @@ -429,9 +431,16 @@ type Offering struct { PricePerHour float64 Available bool // Region is the provider region this row prices, in the provider's own - // vocabulary (e.g. AWS "us-east-1"). Empty for region-simple providers whose - // catalog is not region-partitioned (Modal, RunPod); a region-aware provider - // emits one row per {accelerator, capacityType, region}. + // vocabulary (e.g. AWS "us-east-1"). Empty when a provider's catalog is not + // region-partitioned; a region-aware provider emits one row per {accelerator, + // capacityType, region}. + // + // Empty here is about PRICING, and says nothing about whether the provider has a + // region axis at all — the two are independent. Modal has neither. RunPod prices + // every data center alike, so its rows carry no region, yet region IS a real + // placement axis for it (a pool's regions become RunPod dataCenterIds). AWS's rows + // are blank for a third reason: its per-region truth is probed live rather than + // hand-maintained. Region string // AcceleratorID is this provider's own name for what serves the canonical // AcceleratorType (AWS "p5.48xlarge" for H100) — the lookup data MapAccelerator @@ -473,9 +482,9 @@ type BlockScope struct { Accelerator *string // CapacityType empty => blocks all capacity types. CapacityType nebulav1alpha1.CapacityType - // Region: nil => the provider has no region axis (Modal/RunPod, whose candidates - // carry an empty region too); &"us-east-1" => that region only, so a shortage there - // does not disqualify us-west-2. + // Region: nil => the provider has no region axis (Modal, whose candidates carry an + // empty region too); &"us-east-1" => that region only, so a shortage there does not + // disqualify us-west-2. Region *string // DenyAll true => block everything on this provider (auth/quota errors), ignoring the // fields above. Still scoped to this one provider; it never spans providers. diff --git a/pkg/provider/runpod/client.go b/pkg/provider/runpod/client.go new file mode 100644 index 0000000..723a09e --- /dev/null +++ b/pkg/provider/runpod/client.go @@ -0,0 +1,557 @@ +/* +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 runpod + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "os" + "strconv" + "strings" + "sync" + "time" + + "github.com/InftyAI/Nebula/pkg/provider" + "github.com/InftyAI/Nebula/pkg/provider/catalog" + "github.com/InftyAI/Nebula/pkg/util" +) + +// The RunPod REST API v2 (https://docs.runpod.io/api-reference-v2/overview; v1 retires +// 2026-11-15). +const ( + // defaultBaseURL is RunPod's API host; every path carries its own /v2 prefix. + // Overridable only in tests (see newClient). + defaultBaseURL = "https://api.runpod.io" + // apiKeyEnv is where the credential comes from. It is delivered by the per-provider + // Secret the manager mounts via envFrom; absent means the provider is skipped at + // registration rather than failing the process. + apiKeyEnv = "RUNPOD_API_KEY" + // requestTimeout bounds one HTTP call. Generous because a create allocates a machine + // server-side, but finite: a hung call would otherwise pin a Provision until the + // caller's own context expired. + requestTimeout = 60 * time.Second + // maxResponseBytes caps how much of a response is read, so a malformed or hostile + // response cannot exhaust memory. A full Pod list is a few KB per Pod. + maxResponseBytes = 8 << 20 + // maxErrorBodyChars caps how much of an error response reaches the error string, which + // is logged and may land on a Pod condition. + maxErrorBodyChars = 512 + // cloudTypeSecure is the only cloud type Nebula requests; see the package doc for why + // COMMUNITY is out until the catalog can price it. + cloudTypeSecure = "SECURE" + // listPageSize is the largest page GET /v2/pods serves. + listPageSize = 1000 +) + +// restClient is the real Client, backed by RunPod's REST API. Every RunPod-specific HTTP +// call lives here so the adapter and its tests stay transport-free. +type restClient struct { + http *http.Client + baseURL string + apiKey string + // registryIDs caches registryAuthName -> RunPod credential id. The key is derived from + // the credential, so an entry can never serve the wrong one; see EnsureRegistryAuth. + registryIDs sync.Map + // registryMu serializes the uncached list-then-create; see EnsureRegistryAuth. + registryMu sync.Mutex +} + +// compile-time assertion that restClient satisfies the adapter's Client seam. +var _ Client = (*restClient)(nil) + +// NewSDKClient builds a RunPod-backed Provider, reading the API key from RUNPOD_API_KEY. +// An absent key is an ERROR rather than a client that fails on first use, so +// registerProviders can log and skip RunPod the same way it skips Modal and AWS. +// +// The context is accepted for symmetry with the other adapters' constructors (and so a +// future availability probe can use it); nothing here makes a call. +func NewSDKClient(_ context.Context) (*Provider, error) { + apiKey := strings.TrimSpace(os.Getenv(apiKeyEnv)) + if apiKey == "" { + return nil, fmt.Errorf("runpod: %s is not set", apiKeyEnv) + } + cat, err := catalog.Load() + if err != nil { + return nil, fmt.Errorf("runpod: load price catalog: %w", err) + } + return New(newClient(defaultBaseURL, apiKey), cat), nil +} + +// newClient builds a restClient against baseURL. Separate from NewSDKClient so a test can +// point it at an httptest.Server. +func newClient(baseURL, apiKey string) *restClient { + return &restClient{ + http: &http.Client{Timeout: requestTimeout}, + baseURL: strings.TrimSuffix(baseURL, "/"), + apiKey: apiKey, + } +} + +// apiError is one non-2xx RunPod response. It carries the status and RunPod's own message +// so the classify helpers can key on both, and so an operator reading a log sees what +// RunPod actually said. +// +// It deliberately holds NOTHING from the REQUEST body: that body carries the workload's +// resolved environment and, on a registry-auth create, a registry password. Only the method +// and path are echoed back. +type apiError struct { + status int + method string + path string + message string +} + +func (e *apiError) Error() string { + return fmt.Sprintf("runpod: %s %s: HTTP %d: %s", e.method, e.path, e.status, e.message) +} + +// notFound reports whether err is a 404. Both Get and Terminate treat that as "already +// gone" rather than a failure, which is what makes Terminate idempotent for the finalizer. +func notFound(err error) bool { + var ae *apiError + return errors.As(err, &ae) && ae.status == http.StatusNotFound +} + +// do performs one API call: body is JSON-encoded when non-nil, out is JSON-decoded when +// non-nil, and any non-2xx becomes an *apiError. +func (c *restClient) do(ctx context.Context, method, path string, body, out any) error { + var payload io.Reader + if body != nil { + encoded, err := json.Marshal(body) + if err != nil { + return fmt.Errorf("runpod: encode %s %s request: %w", method, path, err) + } + payload = bytes.NewReader(encoded) + } + req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, payload) + if err != nil { + return fmt.Errorf("runpod: build %s %s request: %w", method, path, err) + } + req.Header.Set("Authorization", "Bearer "+c.apiKey) + req.Header.Set("Accept", "application/json") + if body != nil { + req.Header.Set("Content-Type", "application/json") + } + + resp, err := c.http.Do(req) + if err != nil { + // A transport failure, wrapped WITHOUT a sentinel on purpose: nobody knows whether + // RunPod acted on the request, so it must stay unattributable and blocklist nothing. + return fmt.Errorf("runpod: %s %s: %w", method, path, err) + } + defer func() { _ = resp.Body.Close() }() + + raw, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseBytes)) + if err != nil { + return fmt.Errorf("runpod: %s %s: read response: %w", method, path, err) + } + if resp.StatusCode >= http.StatusMultipleChoices { + return &apiError{ + status: resp.StatusCode, + method: method, + path: path, + message: errorMessage(raw), + } + } + if out == nil { + return nil + } + if err := json.Unmarshal(raw, out); err != nil { + return fmt.Errorf("runpod: %s %s: decode response: %w", method, path, err) + } + return nil +} + +// errorMessage pulls the human-readable part out of an error response: v2's RFC 9457 body +// ({title, status, detail, errors}), falling back to the raw text — some gateway errors are +// HTML, and an empty message would leave the classifier nothing to read. +// +// The per-field errors are appended to the detail because a 422 puts the reason ONLY there. +func errorMessage(raw []byte) string { + var problem struct { + Title string `json:"title"` + Detail string `json:"detail"` + Errors []string `json:"errors"` + } + if err := json.Unmarshal(raw, &problem); err == nil { + msg := problem.Detail + if msg == "" { + msg = problem.Title + } + if len(problem.Errors) > 0 { + msg = strings.TrimSpace(msg + ": " + strings.Join(problem.Errors, "; ")) + } + if msg != "" { + return truncate(msg) + } + } + return truncate(strings.TrimSpace(string(raw))) +} + +// truncate bounds a message so an error string stays loggable. +func truncate(s string) string { + if len(s) <= maxErrorBodyChars { + return s + } + return s[:maxErrorBodyChars] + "…" +} + +// classifyCreate wraps a create failure with the shared sentinel that matches it, which is +// what lets the control plane act on the failure without knowing anything about RunPod (see +// docs/add-a-provider.md, "Wrap the errors your Provision returns"). +// +// The status table follows RunPod's own guidance for POST /v2/pods. Its gotcha is 400: it +// means both "this GPU and data center could not be placed" and "the body breaks a +// cross-field rule", with no machine-readable code to tell them apart, so only a detail +// that reads as capacity is treated as one. +func classifyCreate(err error) error { + var ae *apiError + if !errors.As(err, &ae) { + return err // a transport/encode failure: unattributable, and already wrapped + } + msg := strings.ToLower(ae.message) + + switch { + case ae.status >= http.StatusInternalServerError: + // Left UNWRAPPED, deliberately. A 5xx says RunPod failed to answer, not that it + // said no, and it may well have created the Pod before falling over. + return err + + case ae.status == http.StatusUnauthorized: + // Whole-provider: nothing succeeds until the key is fixed. + return fmt.Errorf("%w: %w", err, provider.ErrAuth) + + case ae.status == http.StatusForbidden: + // NOT auth, on create: RunPod documents it as "your account cannot access the + // requested pool", to be skipped for the next candidate. DenyAll would fence off + // every other pool the account can use. + return fmt.Errorf("%w: %w", err, provider.ErrUnsupportedAccelerator) + + // Money, not capacity, but scoped the same way: it is transient, it is not an + // authentication problem, and ErrQuota is the sentinel for "a limit stopped this". + case ae.status == http.StatusPaymentRequired, ae.status == http.StatusTooManyRequests, + util.ContainsAny(msg, "insufficient funds", "insufficient balance", "not enough credit"): + return fmt.Errorf("%w: %w", err, provider.ErrQuota) + + case util.ContainsAny(msg, "registry", "image", "pull", "manifest"): + // Belongs to the REQUEST, not the candidate, so it must blocklist NOTHING. The phrase + // is what provider.ClassifyError keys on; left bare, a registry's "unauthorized" would + // read as OUR auth failing and fence the whole provider. Before the capacity and GPU + // cases, whose generic "unavailable"/"unsupported" also match an image message. + return fmt.Errorf("runpod: image pull credential or image rejected: %w", err) + + case util.ContainsAny(msg, "no longer any instances available", "no instances available", + "no instance available", "out of capacity", "no capacity", "not available", + "unavailable", "sold out", "could not be placed"): + return fmt.Errorf("%w: %w", err, provider.ErrNoCapacity) + + case util.ContainsAny(msg, "invalid gpu", "unknown gpu", "gpu type", "unsupported"): + // A GPU id RunPod does not recognize: durable until runpod.csv is corrected, and + // accelerator-scoped so the rest of the provider stays usable. + return fmt.Errorf("%w: %w", err, provider.ErrUnsupportedAccelerator) + + default: + // A 422 or an unrecognized 400. Left unwrapped rather than guessed at: every + // available sentinel is worse — ErrAuth would fence off the whole provider, a + // capacity wrap would evict a healthy candidate. + return err + } +} + +// createPodRequest is RunPod's POST /v2/pods body. Only the fields Nebula sets are present; +// everything omitted takes RunPod's own default, which is the point of the omitempty tags — +// a zero we did not mean would override a sane default with 0. +// +// No mounts: a Nebula instance is cattle with nothing to persist, and v2 attaches no +// persistent volume unless asked (v1 defaulted to a billable 20 GiB one). +type createPodRequest struct { + Name string `json:"name"` + Image string `json:"image"` + Cloud string `json:"cloud"` + GPU *gpuRequest `json:"gpu,omitempty"` + + Disk int `json:"disk,omitempty"` + Env map[string]string `json:"env,omitempty"` + Entrypoint []string `json:"entrypoint,omitempty"` + Cmd []string `json:"cmd,omitempty"` + Ports []string `json:"ports,omitempty"` + + DataCenterIDs []string `json:"dataCenterIds,omitempty"` + Registry string `json:"registry,omitempty"` +} + +// gpuRequest is a GPU Pod's compute. The per-GPU minimums are placement filters, not +// reservations: RunPod may hand out more, never less. +type gpuRequest struct { + ID string `json:"id"` + Count int32 `json:"count"` + MinVCPUCountPerGPU int `json:"minVcpuCountPerGpu,omitempty"` + MinRAMPerGPU int `json:"minRamPerGpu,omitempty"` +} + +// podResponse is the subset of RunPod's Pod object this adapter reads. Fields it ignores +// (runtime, ssh, cost, template, mounts) are omitted rather than carried, so the struct states exactly +// what the adapter's behaviour depends on. +type podResponse struct { + ID string `json:"id"` + Name string `json:"name"` + Status string `json:"status"` + Ports []string `json:"ports"` + // DataCenterID is null until the scheduler assigns one, which decodes to "" — Region + // then stays empty rather than reporting a placement we did not observe. + DataCenterID string `json:"dataCenterId"` +} + +// toPod converts the wire shape into the adapter's view. +func (r podResponse) toPod() Pod { + return Pod{ + ID: r.ID, + Name: r.Name, + Status: r.Status, + Ports: r.Ports, + DataCenterID: r.DataCenterID, + } +} + +// CreatePod implements Client. +func (c *restClient) CreatePod(ctx context.Context, spec PodSpec) (string, error) { + body := createPodRequest{ + Name: spec.Name, + Image: spec.Image, + Cloud: cloudTypeSecure, + Disk: spec.ContainerDiskGiB, + Env: spec.Env, + Entrypoint: spec.Entrypoint, + Cmd: spec.StartCmd, + Ports: spec.Ports, + DataCenterIDs: spec.DataCenterIDs, + Registry: spec.RegistryAuthID, + GPU: &gpuRequest{ + ID: spec.GPUTypeID, + Count: spec.GPUCount, + MinVCPUCountPerGPU: spec.VCPUPerGPU, + MinRAMPerGPU: spec.RAMPerGPUGiB, + }, + } + + var out podResponse + if err := c.do(ctx, http.MethodPost, "/v2/pods", body, &out); err != nil { + if spec.RegistryAuthID != "" { + // The cached id may name a credential deleted out from under us. Evicted on ANY + // failure, since RunPod's wording for that is undocumented and a spurious evict + // costs one list. + c.forgetRegistryID(spec.RegistryAuthID) + } + return "", classifyCreate(err) + } + if out.ID == "" { + // A 2xx with no id is unusable and, worse, ambiguous: a Pod may exist that we can + // never name to terminate. Reported as an error with no sentinel, so the Pod retries + // (Provision is idempotent on the claim name, and the name lookup will find any Pod + // this call did create). + return "", fmt.Errorf("runpod: create pod %q: response carried no id", spec.Name) + } + return out.ID, nil +} + +// TerminatePod implements Client. Idempotent: a 404 means the Pod is already gone, which is +// success for the caller (the NodeClaim finalizer retries against this). +func (c *restClient) TerminatePod(ctx context.Context, id string) error { + err := c.do(ctx, http.MethodDelete, "/v2/pods/"+url.PathEscape(id), nil, nil) + if err != nil && !notFound(err) { + return err + } + return nil +} + +// GetPod implements Client, returning (nil, nil) for a Pod that no longer exists. +func (c *restClient) GetPod(ctx context.Context, id string) (*Pod, error) { + var out podResponse + err := c.do(ctx, http.MethodGet, "/v2/pods/"+url.PathEscape(id), nil, &out) + if notFound(err) { + return nil, nil + } + if err != nil { + return nil, err + } + pd := out.toPod() + return &pd, nil +} + +// ListPods implements Client. Filtering to Nebula's own Pods is the adapter's job, since it +// owns the naming scheme (see Provider.List). +// +// Every page is walked: a Pod missing from List reads as terminated, so stopping at the +// first page would report every Pod past it dead. One call below listPageSize Pods. +func (c *restClient) ListPods(ctx context.Context) ([]Pod, error) { + var pods []Pod + cursor := "" + for { + q := url.Values{"limit": {strconv.Itoa(listPageSize)}} + if cursor != "" { + q.Set("cursor", cursor) + } + var out struct { + Pods []podResponse `json:"pods"` + Pagination struct { + NextCursor *string `json:"nextCursor"` + HasNextPage bool `json:"hasNextPage"` + } `json:"pagination"` + } + if err := c.do(ctx, http.MethodGet, "/v2/pods?"+q.Encode(), nil, &out); err != nil { + return nil, err + } + for _, r := range out.Pods { + pods = append(pods, r.toPod()) + } + next := out.Pagination.NextCursor + if !out.Pagination.HasNextPage || next == nil || *next == "" || *next == cursor { + return pods, nil + } + cursor = *next + } +} + +// registryAuthPath is the collection RunPod stores image-pull credentials in. +const registryAuthPath = "/v2/registries" + +// registryAuthResponse is the subset of a registry credential this adapter reads. +// Notably NOT the password: RunPod does not return it, and nothing here needs it back. +type registryAuthResponse struct { + ID string `json:"id"` + Name string `json:"name"` +} + +// EnsureRegistryAuth implements Client. RunPod's create takes a registry credential ID, +// never an inline username/password, so a credential has to become an OBJECT in RunPod's +// account before a Pod can use it. +// +// The object is CONTENT-ADDRESSED — its name is a hash of the credential (see +// registryAuthName) — which is what makes this safe to call on every Provision: +// +// - Idempotent. The same credential resolves to the same name, so the list-then-create +// finds the existing object instead of accumulating one object per Pod. +// - Correct across rotation. A changed password hashes differently, so it becomes a new +// object rather than silently reusing a stale one that would 401 at pull time. +// +// Resolved ids are cached, so only a credential's first Provision per process calls the API. +// CreatePod evicts on failure, which bounds a stale entry to one failed Pod. +// +// Objects are never DELETED, and that is deliberate: one object is shared by every Pod using +// that credential, so deleting it on any single teardown would break the others' next pull. +// The population is bounded by the number of distinct credentials, not by the number of Pods. +func (c *restClient) EnsureRegistryAuth(ctx context.Context, auth *provider.RegistryAuth) (string, error) { + if auth == nil { + return "", errors.New("runpod: nil image pull credential") + } + if auth.Basic == nil { + return "", auth.Unsupported("runpod") + } + name := registryAuthName(auth.Basic.Username, auth.Basic.Password) + if id, ok := c.registryIDs.Load(name); ok { + return id.(string), nil + } + + // RunPod names are unique, so two concurrent first Provisions that both miss the list + // would race to create and one would fail its Pod. Re-check once the lock is held: a + // waiter reuses the winner's id. + c.registryMu.Lock() + defer c.registryMu.Unlock() + if id, ok := c.registryIDs.Load(name); ok { + return id.(string), nil + } + if id, err := c.findRegistryAuth(ctx, name); err != nil || id != "" { + return id, err + } + + body := struct { + Name string `json:"name"` + Username string `json:"username"` + Password string `json:"password"` + }{Name: name, Username: auth.Basic.Username, Password: auth.Basic.Password} + + var created registryAuthResponse + if err := c.do(ctx, http.MethodPost, registryAuthPath, body, &created); err != nil { + // The lock is per process; another replica (or a leader handoff) may have created + // the same object in between, so a name clash resolves by listing again. + if id, lerr := c.findRegistryAuth(ctx, name); lerr == nil && id != "" { + return id, nil + } + // Worded as an image-pull failure, which blocklists nothing: a credential RunPod + // would not store is a fact about this Pod's imagePullSecret, not about the + // accelerator or region the Pod was headed for. + return "", fmt.Errorf("runpod: store image pull credential %q: %w", name, err) + } + if created.ID == "" { + return "", fmt.Errorf("runpod: store image pull credential %q: response carried no id", name) + } + c.registryIDs.Store(name, created.ID) + return created.ID, nil +} + +// findRegistryAuth returns the id of the object named name, caching it, or "" if absent. +func (c *restClient) findRegistryAuth(ctx context.Context, name string) (string, error) { + var existing struct { + Registries []registryAuthResponse `json:"registries"` + } + if err := c.do(ctx, http.MethodGet, registryAuthPath, nil, &existing); err != nil { + return "", err + } + for _, e := range existing.Registries { + if e.Name == name && e.ID != "" { + c.registryIDs.Store(name, e.ID) + return e.ID, nil + } + } + return "", nil +} + +// forgetRegistryID evicts every cache entry resolving to id. +func (c *restClient) forgetRegistryID(id string) { + c.registryIDs.Range(func(name, cached any) bool { + if cached == id { + c.registryIDs.Delete(name) + } + return true + }) +} + +// registryAuthName is the content-addressed name of a stored credential: a fixed prefix +// (so Nebula's objects are recognizable in the RunPod console) plus a hash of the +// credential itself. +// +// The hash is what makes EnsureRegistryAuth idempotent, and it is over BOTH fields so a +// rotated password yields a new object. Hashed rather than named after the registry or the +// claim for two reasons: a RunPod object name is not a secret and appears in its UI, so the +// username must not be in it; and naming it after the claim would create one object per +// NodeClaim for a credential every claim shares. +// +// Truncated to 16 hex characters — 64 bits, which for a per-account population of at most a +// handful of credentials is far past any collision concern, and keeps the name readable. +func registryAuthName(username, password string) string { + // The NUL separator keeps ("ab", "c") from hashing the same as ("a", "bc"). + sum := sha256.Sum256([]byte(username + "\x00" + password)) + return registryAuthPrefix + hex.EncodeToString(sum[:])[:16] +} diff --git a/pkg/provider/runpod/client_test.go b/pkg/provider/runpod/client_test.go new file mode 100644 index 0000000..c5e212e --- /dev/null +++ b/pkg/provider/runpod/client_test.go @@ -0,0 +1,586 @@ +/* +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 runpod + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "sync" + "testing" + + "github.com/InftyAI/Nebula/pkg/provider" +) + +// recordedRequest is what the fake RunPod saw, so a test can assert on the wire form the +// adapter produced rather than only on what it got back. +type recordedRequest struct { + method string + path string + query string + auth string + body map[string]any +} + +// testServer stands in for RunPod's REST API. handler answers each call; every request is +// recorded first. Returns the client under test and a pointer to the log. +func testServer(t *testing.T, handler http.HandlerFunc) (*restClient, *[]recordedRequest) { + t.Helper() + var ( + mu sync.Mutex + seen []recordedRequest + ) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + rec := recordedRequest{ + method: r.Method, + path: r.URL.Path, + query: r.URL.RawQuery, + auth: r.Header.Get("Authorization"), + } + if raw, err := io.ReadAll(r.Body); err == nil && len(raw) > 0 { + _ = json.Unmarshal(raw, &rec.body) + } + mu.Lock() + seen = append(seen, rec) + mu.Unlock() + handler(w, r) + })) + t.Cleanup(srv.Close) + return newClient(srv.URL, "test-key"), &seen +} + +// jsonReply writes one canned response. +func jsonReply(status int, body string) http.HandlerFunc { + return func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _, _ = io.WriteString(w, body) + } +} + +// problem is a v2 error body (RFC 9457) carrying msg as its detail. +func problem(status int, msg string) string { + return fmt.Sprintf(`{"title":%q,"status":%d,"detail":%q}`, http.StatusText(status), status, msg) +} + +func TestClassifyCreate(t *testing.T) { + // Every one of these must carry a sentinel or deliberately carry none: unwrapped, a + // rejection lands on nebula_provision_failures_total{reason="other"}. + cases := []struct { + name string + status int + message string + + want error // the sentinel the error must wrap, nil for "none" + // blocksNothing: the zero BlockScope, the contract for a REQUEST-scoped failure. + blocksNothing bool + }{{ + name: "401 is auth", status: 401, message: "invalid token", want: provider.ErrAuth, + }, { + // RunPod: "your account cannot access the requested pool — skip this candidate". + // DenyAll would fence off every pool the account CAN use. + name: "403 is scoped to the candidate", status: 403, message: "Forbidden", + want: provider.ErrUnsupportedAccelerator, + }, { + // A 5xx says RunPod failed to ANSWER, not that it said no — and it may have created + // the Pod before falling over. + name: "500 stays unwrapped", status: 500, message: "internal error", want: nil, + }, { + name: "502 stays unwrapped", status: 502, message: "bad gateway", want: nil, + }, { + name: "429 is quota", status: 429, message: "Too Many Requests", want: provider.ErrQuota, + }, { + name: "402 is quota", status: 402, message: "Insufficient balance", want: provider.ErrQuota, + }, { + name: "insufficient funds is quota", + status: 400, + message: "Insufficient funds to start this pod", + want: provider.ErrQuota, + }, { + name: "no instances available is capacity", + status: 400, + message: "There are no longer any instances available with the requested specifications", + want: provider.ErrNoCapacity, + }, { + name: "an unknown gpu type is an accelerator problem", + status: 400, + message: "invalid gpu type id", + want: provider.ErrUnsupportedAccelerator, + }, { + // Belongs to the REQUEST, not the candidate: one Pod's bad credential must not exclude + // an accelerator serving every other Pod — and a registry's "unauthorized" must not + // read as OUR auth failing. + name: "a registry failure blocks nothing", + status: 400, + message: "registry unauthorized: could not pull image", + blocksNothing: true, + }, { + // The capacity and GPU phrases are generic enough to match these too. + name: "an unavailable image blocks nothing", status: 400, message: "image not available", + blocksNothing: true, + }, { + name: "an unavailable manifest blocks nothing", status: 400, message: "manifest unavailable", + blocksNothing: true, + }, { + name: "an unsupported image blocks nothing", status: 400, message: "unsupported image format", + blocksNothing: true, + }, { + // v2's 400 also covers cross-field rule violations, with no code to tell them from + // capacity; guessing either way is worse than leaving it unattributable. + name: "an unrecognized 400 stays unwrapped", status: 400, message: "something new", want: nil, + }, { + name: "a 422 stays unwrapped", status: 422, message: "gpu.count: must be >= 1", want: nil, + }} + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + c, _ := testServer(t, jsonReply(tc.status, problem(tc.status, tc.message))) + + _, err := c.CreatePod(context.Background(), PodSpec{ + Name: "nebula-claim-a", Image: "img", GPUCount: 1, GPUTypeID: "NVIDIA H100 80GB HBM3", + }) + if err == nil { + t.Fatal("CreatePod succeeded against an error response") + } + if tc.want != nil && !errors.Is(err, tc.want) { + t.Errorf("error %v does not wrap %v", err, tc.want) + } + if tc.want == nil { + for _, s := range []error{ + provider.ErrAuth, provider.ErrQuota, provider.ErrNoCapacity, + provider.ErrUnsupportedAccelerator, + } { + if errors.Is(err, s) { + t.Errorf("error %v wraps %v; it must stay unattributable", err, s) + } + } + } + scope := provider.ClassifyError(err, "", "H100:1") + if tc.blocksNothing && scope != (provider.BlockScope{}) { + t.Errorf("scope = %+v, want the zero scope", scope) + } + // RunPod's own words survive into the message an operator reads off a Pod condition. + if !strings.Contains(err.Error(), tc.message) { + t.Errorf("error %q dropped RunPod's message %q", err, tc.message) + } + }) + } +} + +func TestErrorMessage_KeepsValidationErrors(t *testing.T) { + // A 422 carries its reason ONLY in errors[]; dropping it leaves "Unprocessable Entity". + got := errorMessage([]byte(`{"title":"Unprocessable Entity","status":422, + "detail":"validation failed","errors":["cpu.vcpuCount: must be a power of two"]}`)) + if want := "validation failed: cpu.vcpuCount: must be a power of two"; got != want { + t.Errorf("errorMessage = %q, want %q", got, want) + } + if got := errorMessage([]byte("bad gateway")); got != "bad gateway" { + t.Errorf("non-JSON body = %q, want it verbatim", got) + } +} + +func TestCreatePod_WireForm(t *testing.T) { + c, seen := testServer(t, jsonReply(201, `{"id":"pod-1"}`)) + + id, err := c.CreatePod(context.Background(), PodSpec{ + Name: "nebula-claim-a", + Image: "myimg:latest", + Entrypoint: []string{"python"}, + StartCmd: []string{"serve.py"}, + Env: map[string]string{"K": "v"}, + GPUTypeID: "NVIDIA H100 80GB HBM3", + GPUCount: 2, + VCPUPerGPU: 5, + RAMPerGPUGiB: 50, + DataCenterIDs: []string{"US-KS-2", "US-TX-3"}, + }) + if err != nil { + t.Fatalf("CreatePod: %v", err) + } + if id != "pod-1" { + t.Fatalf("id = %q, want pod-1", id) + } + req := (*seen)[0] + if req.method != http.MethodPost || req.path != "/v2/pods" { + t.Errorf("%s %s, want POST /v2/pods", req.method, req.path) + } + if req.auth != "Bearer test-key" { + t.Errorf("Authorization = %q", req.auth) + } + want := map[string]any{ + "name": "nebula-claim-a", + "image": "myimg:latest", + // SECURE only: COMMUNITY prices the same GPU differently and the catalog has no + // cloud-type axis to express that. + "cloud": cloudTypeSecure, + "entrypoint": []any{"python"}, + "cmd": []any{"serve.py"}, + "env": map[string]any{"K": "v"}, + "gpu": map[string]any{ + "id": "NVIDIA H100 80GB HBM3", "count": float64(2), + "minVcpuCountPerGpu": float64(5), "minRamPerGpu": float64(50), + }, + "dataCenterIds": []any{"US-KS-2", "US-TX-3"}, + } + if !reflect.DeepEqual(req.body, want) { + t.Errorf("body =\n%v\nwant\n%v", req.body, want) + } +} + +func TestCreatePod_SuccessWithNoID(t *testing.T) { + // A 2xx with no id is worse than an error: a Pod may exist that we can never name to + // terminate. It must fail WITHOUT a sentinel so nothing is blocklisted — Provision is + // idempotent on the claim name, and the name lookup will find whatever this call created. + c, _ := testServer(t, jsonReply(201, `{}`)) + + _, err := c.CreatePod(context.Background(), PodSpec{Name: "nebula-a", Image: "img", GPUCount: 1}) + if err == nil { + t.Fatal("CreatePod accepted a response with no id") + } + if scope := provider.ClassifyError(err, "", "H100:1"); scope != (provider.BlockScope{}) { + t.Errorf("scope = %+v; a missing id says nothing about the candidate", scope) + } +} + +func TestTerminatePod_404IsSuccess(t *testing.T) { + // The NodeClaim finalizer retries Terminate, so an already-gone Pod has to be success — + // otherwise the finalizer never clears and the Pod is stuck deleting forever. + c, seen := testServer(t, jsonReply(404, problem(404, "pod not found"))) + if err := c.TerminatePod(context.Background(), "pod-gone"); err != nil { + t.Fatalf("TerminatePod on a missing Pod = %v, want nil", err) + } + if req := (*seen)[0]; req.method != http.MethodDelete || req.path != "/v2/pods/pod-gone" { + t.Errorf("%s %s, want DELETE /v2/pods/pod-gone", req.method, req.path) + } + + // Any OTHER failure must still surface: swallowing a 500 would drop the teardown + // obligation and leak a billing instance. + c2, _ := testServer(t, jsonReply(500, problem(500, "boom"))) + if err := c2.TerminatePod(context.Background(), "pod-1"); err == nil { + t.Error("TerminatePod swallowed a 500; the instance would leak") + } +} + +func TestGetPod_404IsGone(t *testing.T) { + // Absent means terminated, per the interface contract. + c, seen := testServer(t, jsonReply(404, problem(404, "not found"))) + pd, err := c.GetPod(context.Background(), "pod-gone") + if err != nil || pd != nil { + t.Fatalf("GetPod(missing) = %v, %v; want nil, nil", pd, err) + } + if req := (*seen)[0]; req.path != "/v2/pods/pod-gone" { + t.Errorf("GET %s, want /v2/pods/pod-gone", req.path) + } +} + +func TestListPods_DecodesPlacementAndPorts(t *testing.T) { + c, _ := testServer(t, jsonReply(200, `{"pods":[ + {"id":"pod-1","name":"nebula-claim-a","status":"RUNNING","dataCenterId":"EU-RO-1", + "ports":["8000/http","22/tcp"], + "runtime":{"ports":[{"private":22,"public":34446,"type":"tcp","ip":"195.26.233.3"}, + {"private":8000,"public":null,"type":"http","ip":null}]}}, + {"id":"pod-2","name":"nebula-claim-b","status":"PROVISIONING","dataCenterId":null,"runtime":null} + ],"pagination":{"nextCursor":null,"hasNextPage":false}}`)) + + pods, err := c.ListPods(context.Background()) + if err != nil { + t.Fatalf("ListPods: %v", err) + } + if len(pods) != 2 { + t.Fatalf("got %d pods, want 2", len(pods)) + } + if pods[0].DataCenterID != "EU-RO-1" || pods[0].Status != "RUNNING" { + t.Errorf("pod-1 = %+v", pods[0]) + } + if !reflect.DeepEqual(pods[0].Ports, []string{"8000/http", "22/tcp"}) { + t.Errorf("pod-1 ports = %v", pods[0].Ports) + } + // A null data center must leave the region EMPTY rather than reporting a placement that + // was never observed. + if pods[1].DataCenterID != "" || pods[1].Ports != nil { + t.Errorf("pod-2 = %+v, want no region and no ports", pods[1]) + } +} + +func TestListPods_WalksEveryPage(t *testing.T) { + // A Pod missing from List reads as terminated, so stopping at page one would report every + // Pod past it dead. + c, seen := testServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Query().Get("cursor") == "" { + jsonReply(200, `{"pods":[{"id":"pod-1","name":"nebula-a"}], + "pagination":{"nextCursor":"c2","hasNextPage":true}}`)(w, r) + return + } + jsonReply(200, `{"pods":[{"id":"pod-2","name":"nebula-b"}], + "pagination":{"nextCursor":null,"hasNextPage":false}}`)(w, r) + }) + + pods, err := c.ListPods(context.Background()) + if err != nil { + t.Fatalf("ListPods: %v", err) + } + if len(pods) != 2 || pods[1].ID != "pod-2" { + t.Fatalf("pods = %+v, want both pages", pods) + } + if q := (*seen)[1].query; !strings.Contains(q, "cursor=c2") { + t.Errorf("second page query = %q, want the cursor passed through", q) + } +} + +func TestEnsureRegistryAuth(t *testing.T) { + auth := &provider.RegistryAuth{ + Registry: "ghcr.io", + Basic: &provider.BasicAuth{Username: "u", Password: "p4ssw0rd"}, + } + // Content-addressed: the SAME credential always resolves to the same object name, which + // is what makes calling this on every Provision safe. + name := registryAuthName("u", "p4ssw0rd") + + t.Run("reuses an existing object", func(t *testing.T) { + // One object is shared by every Pod using that credential. Creating a second per Pod + // would accumulate objects without bound, and RunPod never garbage-collects them. + c, seen := testServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost { + t.Error("POST issued although a matching object already exists") + } + jsonReply(200, fmt.Sprintf(`{"registries":[{"id":"cra-existing","name":%q}, + {"id":"cra-other","name":"nebula-deadbeefdeadbeef"}]}`, name))(w, r) + }) + + id, err := c.EnsureRegistryAuth(context.Background(), auth) + if err != nil { + t.Fatalf("EnsureRegistryAuth: %v", err) + } + if id != "cra-existing" { + t.Errorf("id = %q, want cra-existing", id) + } + if len(*seen) != 1 { + t.Errorf("made %d calls, want 1 (the list)", len(*seen)) + } + }) + + t.Run("creates when absent, and never sends the password back on the list", func(t *testing.T) { + c, seen := testServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + jsonReply(200, `{"registries":[]}`)(w, r) + return + } + jsonReply(201, `{"id":"cra-new","name":"whatever"}`)(w, r) + }) + + id, err := c.EnsureRegistryAuth(context.Background(), auth) + if err != nil { + t.Fatalf("EnsureRegistryAuth: %v", err) + } + if id != "cra-new" { + t.Errorf("id = %q, want cra-new", id) + } + if len(*seen) != 2 { + t.Fatalf("made %d calls, want 2 (list then create)", len(*seen)) + } + post := (*seen)[1] + if post.method != http.MethodPost || post.path != registryAuthPath { + t.Errorf("%s %s, want POST %s", post.method, post.path, registryAuthPath) + } + // The object NAME must be the hash, not the username or the claim: a RunPod object + // name is not secret and shows in its UI, and naming it after the claim would create + // one object per NodeClaim for a credential every claim shares. + if post.body["name"] != name { + t.Errorf("name = %v, want the content-addressed %q", post.body["name"], name) + } + if post.body["username"] != "u" || post.body["password"] != "p4ssw0rd" { + t.Errorf("credential did not reach the create body: %v", post.body["username"]) + } + }) + + t.Run("a rotated password becomes a new object", func(t *testing.T) { + // The hash covers BOTH fields, so a rotation cannot silently reuse a stale object + // that would 401 at pull time. + if registryAuthName("u", "old") == registryAuthName("u", "new") { + t.Error("a rotated password hashes to the same object name") + } + // And the NUL separator keeps ("ab","c") from colliding with ("a","bc"). + if registryAuthName("ab", "c") == registryAuthName("a", "bc") { + t.Error("username/password boundary is not separated in the hash") + } + if !strings.HasPrefix(name, registryAuthPrefix) { + t.Errorf("name %q lacks the %q prefix that makes Nebula's objects recognizable", + name, registryAuthPrefix) + } + }) + + t.Run("cached after the first call, evicted when a create using it fails", func(t *testing.T) { + c, seen := testServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/v2/pods" { + jsonReply(400, problem(400, "registry not found"))(w, r) + return + } + jsonReply(200, fmt.Sprintf(`{"registries":[{"id":"cra-existing","name":%q}]}`, name))(w, r) + }) + ctx := context.Background() + + for range 2 { + if _, err := c.EnsureRegistryAuth(ctx, auth); err != nil { + t.Fatalf("EnsureRegistryAuth: %v", err) + } + } + if len(*seen) != 1 { + t.Fatalf("made %d calls for two resolves, want 1 (the second is cached)", len(*seen)) + } + + // A deleted credential leaves a stale cached id; the failed create must drop it so + // the next Provision re-lists rather than failing every Pod until restart. + if _, err := c.CreatePod(ctx, PodSpec{Name: "claim-a", RegistryAuthID: "cra-existing"}); err == nil { + t.Fatal("CreatePod succeeded against an error response") + } + if _, err := c.EnsureRegistryAuth(ctx, auth); err != nil { + t.Fatalf("EnsureRegistryAuth: %v", err) + } + if last := (*seen)[len(*seen)-1]; last.path != registryAuthPath { + t.Errorf("last call %s %s, want a re-list after eviction", last.method, last.path) + } + }) + + t.Run("a create failure blocklists nothing", func(t *testing.T) { + // A credential RunPod will not store is a fact about this Pod's imagePullSecret, not + // about the accelerator or region it was headed for — so the zero BlockScope. + c, _ := testServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + jsonReply(200, `{"registries":[]}`)(w, r) + return + } + jsonReply(400, problem(400, "invalid credential"))(w, r) + }) + + _, err := c.EnsureRegistryAuth(context.Background(), auth) + if err == nil { + t.Fatal("EnsureRegistryAuth succeeded against an error response") + } + if scope := provider.ClassifyError(err, "", "H100:1"); scope != (provider.BlockScope{}) { + t.Errorf("scope = %+v, want the zero scope", scope) + } + }) + + t.Run("a kind RunPod cannot express is refused, not silently dropped", func(t *testing.T) { + // The adapter vets the kind first, so this is a programming error — but it must never + // become a silent ANONYMOUS pull, which either 401s opaquely or succeeds against a + // PUBLIC image of the same name. + c, seen := testServer(t, jsonReply(200, `{"registries":[]}`)) + _, err := c.EnsureRegistryAuth(context.Background(), &provider.RegistryAuth{ + Registry: "1234.dkr.ecr.us-east-1.amazonaws.com", + AWSRole: &provider.AWSRoleAuth{RoleARN: "arn:aws:iam::1234:role/pull", Region: "us-east-1"}, + }) + if scope := provider.ClassifyError(err, "", "H100:1"); err == nil || scope != (provider.BlockScope{}) { + t.Errorf("error = %v, scope = %+v; want a refusal that blocks nothing", err, scope) + } + if len(*seen) != 0 { + t.Errorf("made %d API calls for a credential it cannot express", len(*seen)) + } + }) +} + +// TestEnsureRegistryAuthRace pins that a name clash on RunPod's unique names never fails a Pod. +func TestEnsureRegistryAuthRace(t *testing.T) { + auth := &provider.RegistryAuth{ + Registry: "ghcr.io", + Basic: &provider.BasicAuth{Username: "u", Password: "p4ssw0rd"}, + } + name := registryAuthName("u", "p4ssw0rd") + + t.Run("concurrent first calls create once", func(t *testing.T) { + var mu sync.Mutex + var created bool + posts := 0 + c, _ := testServer(t, func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + defer mu.Unlock() + if r.Method == http.MethodPost { + posts++ + created = true + jsonReply(201, `{"id":"cra-new","name":"whatever"}`)(w, r) + return + } + if created { + jsonReply(200, fmt.Sprintf(`{"registries":[{"id":"cra-new","name":%q}]}`, name))(w, r) + return + } + jsonReply(200, `{"registries":[]}`)(w, r) + }) + + const n = 8 + ids := make([]string, n) + var wg sync.WaitGroup + for i := range n { + wg.Go(func() { + id, err := c.EnsureRegistryAuth(context.Background(), auth) + if err != nil { + t.Errorf("EnsureRegistryAuth: %v", err) + } + ids[i] = id + }) + } + wg.Wait() + if posts != 1 { + t.Errorf("issued %d creates, want 1", posts) + } + for i, id := range ids { + if id != "cra-new" { + t.Errorf("call %d got id %q, want cra-new", i, id) + } + } + }) + + t.Run("a create lost to another process resolves by re-listing", func(t *testing.T) { + lists := 0 + c, _ := testServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost { + jsonReply(400, problem(400, "name already exists"))(w, r) + return + } + lists++ + if lists == 1 { + jsonReply(200, `{"registries":[]}`)(w, r) + return + } + jsonReply(200, fmt.Sprintf(`{"registries":[{"id":"cra-theirs","name":%q}]}`, name))(w, r) + }) + + id, err := c.EnsureRegistryAuth(context.Background(), auth) + if err != nil { + t.Fatalf("EnsureRegistryAuth: %v", err) + } + if id != "cra-theirs" { + t.Errorf("id = %q, want cra-theirs", id) + } + }) +} + +func TestNewSDKClient_MissingKeyIsSkippable(t *testing.T) { + // An absent key must be an ERROR rather than a client that fails on first use, so + // registerProviders can log and skip RunPod the way it skips Modal and AWS — an operator + // who configured only Modal must not get a fatal boot. + t.Setenv(apiKeyEnv, "") + if _, err := NewSDKClient(context.Background()); err == nil { + t.Fatal("NewSDKClient succeeded with no API key; registration would silently register a dead provider") + } +} diff --git a/pkg/provider/runpod/runpod.go b/pkg/provider/runpod/runpod.go new file mode 100644 index 0000000..5213ca4 --- /dev/null +++ b/pkg/provider/runpod/runpod.go @@ -0,0 +1,740 @@ +/* +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 runpod implements the provider.Provider interface for RunPod +// (https://runpod.io), a NeoCloud renting GPU containers by the second. One NodeClaim +// maps to one RunPod Pod. +// +// RunPod sits between the two adapters that came before it, and the differences are +// what drive the decisions here: +// +// - Lifecycle is create/terminate only for our purposes, so SupportsStop=false. RunPod +// does expose stop/start, but a stopped Pod still bills for its disk and releases the +// GPU, so it is neither free nor resumable in the sense the capability promises. +// - OnDemand only, so SupportsSpot=false: REST v2 has no interruptible tier (v1's +// `interruptible` flag has no successor). +// - Create FAILS SYNCHRONOUSLY when capacity is short ("no longer any instances +// available..."), the AWS behaviour rather than Modal's queue-and-accept. That is +// what makes region failover meaningful, so ResolveRegions mints one candidate per +// declared region, each blocklistable on its own. +// - A create takes ONE GPU type, so interchangeable alternates cannot widen a launch the +// way AWS's fleet does; only the primary is sent. +// - Pods have NO tags. NativeTags=false, and Nebula's identity rides the Pod NAME, +// which is the claim name itself (see podName). +// - There is no outbound-allowlist knob at all, so SupportsEgressPolicy=false and +// placement skips RunPod for any pool that restricts egress rather than provisioning +// something with open internet access under a policy that says otherwise. +// - Only SECURE cloud is used. RunPod's cheaper COMMUNITY cloud prices the same GPU +// differently, and the catalog CSV has no cloud-type axis to express that, so +// offering both would make the price the optimizer reads a guess. +// +// kubectl logs and exec are NOT served yet. v2 does stream logs (GET /v2/pods/{id}/logs, as +// SSE), which a LogStreamer could wrap; its only way into a container is SSH, which needs key +// material this adapter has nowhere to put. Both are optional halves of provider.Provider resolved by type +// assertion, so leaving them out costs nothing but a NotFound. +// +// The concrete HTTP API lives behind the Client seam, so this file holds only +// provider-agnostic translation and is unit-testable without network access. +package runpod + +import ( + "context" + "errors" + "fmt" + "slices" + "strings" + + corev1 "k8s.io/api/core/v1" + "k8s.io/apimachinery/pkg/api/resource" + + nebulav1alpha1 "github.com/InftyAI/Nebula/api/v1alpha1" + "github.com/InftyAI/Nebula/pkg/provider" + "github.com/InftyAI/Nebula/pkg/provider/catalog" + "github.com/InftyAI/Nebula/pkg/provider/catalog/data" + "github.com/InftyAI/Nebula/pkg/util" +) + +// registryAuthPrefix makes Nebula's registry-auth objects recognizable in RunPod's UI. +const registryAuthPrefix = "nebula-" + +// maxNameLen is RunPod's own cap on the Pod name (191 chars). Nebula refuses a claim +// whose name would exceed it rather than truncating — see podName. +const maxNameLen = 191 + +// compile-time assertion that Provider satisfies the interface. LogStreamer and Executor +// are deliberately absent; see the package doc. +var _ provider.Provider = (*Provider)(nil) + +// Client is the narrow seam over RunPod's REST API: only the operations the adapter +// needs, in provider-agnostic terms, so the real HTTP implementation and a test fake are +// interchangeable. +type Client interface { + // CreatePod launches one Pod from spec and returns its RunPod id. A capacity + // shortage is an ERROR here, not a queued Pod, and must be wrapped with + // provider.ErrNoCapacity. + CreatePod(ctx context.Context, spec PodSpec) (id string, err error) + // TerminatePod deletes a Pod by id. Must be idempotent: deleting an already-gone + // Pod returns nil, since the NodeClaim finalizer retries against it. + TerminatePod(ctx context.Context, id string) error + // GetPod returns one Pod, or (nil, nil) if it no longer exists. + GetPod(ctx context.Context, id string) (*Pod, error) + // ListPods returns every Pod in the account, in as few calls as possible. Filtering + // down to Nebula's own is the ADAPTER's job (it owns the naming scheme), so this + // must not filter. + ListPods(ctx context.Context) ([]Pod, error) + // EnsureRegistryAuth resolves auth to a RunPod containerRegistryAuth id, creating + // the object if this credential has no id yet. RunPod's create takes an id, never an + // inline username/password, so this indirection is unavoidable. + EnsureRegistryAuth(ctx context.Context, auth *provider.RegistryAuth) (id string, err error) +} + +// PodSpec is the resolved, RunPod-shaped request the Client turns into a Pod. The +// adapter builds it from the Pod (source of truth) plus the resolved accelerator ids. +type PodSpec struct { + // Name is the claim name (`-`). RunPod has no tags, so the name is the + // only claim identity; see podName. + Name string + // Image is the container image, from the Pod's first container. + Image string + // Entrypoint and StartCmd are the container's command and args, mapped onto RunPod's + // dockerEntrypoint/dockerStartCmd. They stay SEPARATE, unlike Modal where both + // concatenate into one command: RunPod's two fields are ENTRYPOINT and CMD, so + // Kubernetes' command/args land on their exact Docker counterparts. + Entrypoint []string + StartCmd []string + // Env is the environment, taken whole from provider.ProvisionRequest.Env: literals + // plus everything envFrom/valueFrom referenced, already resolved by the caller. + // + // SECRET-BEARING, hence the redacting String below. + Env map[string]string + // GPUTypeID is RunPod's own id for the accelerator — MapAccelerator's PRIMARY, since a + // v2 create takes exactly one. + GPUTypeID string + GPUCount int32 + // VCPUPerGPU and RAMPerGPUGiB are the Pod's cpu/memory requests expressed RunPod's + // way — PER GPU, not in total, so the adapter divides by GPUCount and rounds UP + // (rounding down would hand the workload less than it asked for). Zero sends no + // filter, so any host with the GPU qualifies. + VCPUPerGPU int + RAMPerGPUGiB int + // ContainerDiskGiB is the writable container disk: util.PodEphemeralStorageGiB, at least + // defaultContainerDiskGiB. + // + // No persistent volume is ever requested: RunPod defaults to a billable 20 GiB one, + // and a Nebula instance is cattle with nothing to persist, so the Client pins it to 0. + ContainerDiskGiB int + // Ports are the container ports to expose, in RunPod's "/" form (see + // containerPorts). Empty leaves RunPod's default exposure. + Ports []string + // DataCenterIDs is the placement constraint, split out of the ONE region candidate this + // request carries (see ResolveRegions). Empty means unconstrained — the widest capacity + // pool, and the normal case for a pool that declares no regions. + DataCenterIDs []string + // RegistryAuthID is the RunPod registry credential authenticating the image + // pull, already resolved from the canonical credential by Client.EnsureRegistryAuth. + // Empty is an anonymous pull. + RegistryAuthID string +} + +// String redacts Env so a spec can be logged or wrapped in an error safely: key names +// print (they are in the Pod spec already), values never do. RegistryAuthID is an opaque +// object id, not a credential, so it prints as-is. +func (s PodSpec) String() string { + return fmt.Sprintf("PodSpec{Name:%s Image:%s Entrypoint:%v StartCmd:%v Env:%s "+ + "GPUTypeID:%s GPUCount:%d VCPUPerGPU:%d RAMPerGPUGiB:%d "+ + "ContainerDiskGiB:%d Ports:%v DataCenterIDs:%v "+ + "RegistryAuthID:%s}", + s.Name, s.Image, s.Entrypoint, s.StartCmd, provider.RedactedEnv(s.Env), + s.GPUTypeID, s.GPUCount, s.VCPUPerGPU, s.RAMPerGPUGiB, + s.ContainerDiskGiB, s.Ports, s.DataCenterIDs, + s.RegistryAuthID) +} + +// GoString implements fmt.GoStringer so %#v is redacted too. +func (s PodSpec) GoString() string { return s.String() } + +// Pod is the adapter-level view of a RunPod Pod as observed. +type Pod struct { + ID string + Name string + // Status is RunPod's observed lifecycle status; see toState. + Status string + // DataCenterID is where RunPod placed it, in RunPod's own vocabulary. + DataCenterID string + // Ports are the exposed ports RunPod echoes back, in the same "/" form + // PodSpec.Ports sends; the proxy URL routes to the first. + Ports []string +} + +// Provider is the RunPod implementation of provider.Provider. It embeds catalog.Base for +// the generic catalog methods — Name, Offerings and the catalog-driven +// MapAccelerator, which does real work here: RunPod's accelerator ids are marketing +// strings ("NVIDIA H100 80GB HBM3") that share nothing with Nebula's canonical names, so +// every row in runpod.csv carries an accelerator_id. +type Provider struct { + catalog.Base + client Client +} + +// New returns a RunPod Provider backed by client and price catalog. Both must be +// non-nil; use catalog.Load() to build the catalog from the CSV/ConfigMap data. cat is +// the catalog.Lookup seam, so tests can inject a fake. +func New(client Client, cat catalog.Lookup) *Provider { + return &Provider{ + Base: catalog.Base{ProviderName: provider.ProviderRunPod, Catalog: cat}, + client: client, + } +} + +// Capabilities implements provider.Provider. See the package doc for why each trait is +// set the way it is. +func (p *Provider) Capabilities() provider.Capabilities { + return provider.Capabilities{ + SupportsStop: false, // a stopped Pod still bills and loses its GPU + SupportsSpot: false, // v2 has no interruptible tier + SupportsEgressPolicy: false, // no outbound allowlist in the API at all + SupportsCPUOnly: false, // CPU flavors are rarely available + NativeTags: false, // identity rides the Pod name + PreemptionNotice: 0, // OnDemand-only: nothing is reclaimed + PollInterval: 0, // OnDemand-only → the default cadence is fine + // No ProvisionTimeout: one create call, no internal sweep across capacity pools + // to bound (contrast AWS, which walks a region's AZs itself). + } +} + +// Provision implements provider.Provider. The Pod is the source of truth for the +// workload; req carries only the claim identity, the capacity tier and the region. +// +// A successful create is RESERVED, unlike Modal: RunPod allocates a host machine before +// it answers, and a shortage comes back as an error rather than a queued Pod. So an id +// here means real capacity, and the Pod may honestly move on from "provisioning". +// +// The connect URL is RunPod's HTTP proxy for the first declared port — deterministic from +// the Pod id, so no read-back is needed. There is no token: the proxy is unauthenticated, +// which is why ConnectToken stays empty rather than carrying a placeholder. +func (p *Provider) Provision( + ctx context.Context, pod *corev1.Pod, req provider.ProvisionRequest, +) (provider.ProvisionResult, error) { + if pod == nil { + return provider.ProvisionResult{}, errors.New("runpod: nil pod") + } + if req.ClaimName == "" { + return provider.ProvisionResult{}, errors.New("runpod: empty ClaimName in ProvisionRequest") + } + // Refuse a restrictive egress policy rather than dropping it. Placement checks + // SupportsEgressPolicy and should never route such a pool here, but the request can be + // built by anyone, and silently ignoring it would put the workload on the open + // internet under a policy that says otherwise. + if mode := req.Egress.ModeOrOpen(); mode != nebulav1alpha1.EgressOpen { + return provider.ProvisionResult{}, fmt.Errorf( + "runpod: cannot enforce egress mode %q; RunPod exposes no outbound policy", mode) + } + + // Idempotency: RunPod has no tags, so the claim is looked up by the Pod NAME this + // adapter mints. A repeat after a partial create returns the existing Pod rather than + // paying for a second. + // + // Reserved is true for the same reason as a fresh create: the Pod exists, so a machine + // was allocated. That is the capacity question; readiness is separate and observed + // through List. No credential comes back — the interface forbids it on a re-Provision, + // and here there is nothing to re-mint anyway, since the proxy URL is derivable. + // + // Only a Pod that is still coming up or running is adopted: claim names are reused across + // Pod restarts, and an EXITED one reads as Terminated, which would fail the new Pod + // instead of creating its replacement. + existing, err := p.findByClaim(ctx, req.ClaimName, func(pd Pod) bool { + st := toState(pd) + return st == provider.InstancePending || st == provider.InstanceRunning + }) + if err != nil { + return provider.ProvisionResult{}, err + } + if existing != nil { + return provider.ProvisionResult{InstanceID: existing.ID, Reserved: true}, nil + } + + spec, err := p.podSpecFromPod(pod, req) + if err != nil { + return provider.ProvisionResult{}, err + } + // Resolve the pull credential LAST among the request-shaping steps, because unlike + // everything else in the spec it costs API calls; a spec that was going to be rejected + // for its accelerator or its name has already failed by now. + if req.RegistryAuth != nil { + if err := checkRegistryAuth(req.RegistryAuth); err != nil { + return provider.ProvisionResult{}, err + } + authID, err := p.client.EnsureRegistryAuth(ctx, req.RegistryAuth) + if err != nil { + return provider.ProvisionResult{}, err + } + spec.RegistryAuthID = authID + } + + id, err := p.client.CreatePod(ctx, spec) + if err != nil { + return provider.ProvisionResult{}, err + } + return provider.ProvisionResult{ + InstanceID: id, + Reserved: true, + ConnectURL: proxyURL(id, spec.Ports), + }, nil +} + +// Terminate implements provider.Provider. Idempotent by the Client contract. The region +// is ignored: RunPod's API is global and a Pod id addresses it from anywhere, so region is +// only ever a placement input. +func (p *Provider) Terminate(ctx context.Context, instanceID, _ string) error { + if instanceID == "" { + return nil // nothing provisioned yet; treat as already gone + } + return p.client.TerminatePod(ctx, instanceID) +} + +// Get implements provider.Provider. The region is ignored, as in Terminate. +func (p *Provider) Get(ctx context.Context, instanceID, _ string) (*provider.Instance, error) { + pd, err := p.client.GetPod(ctx, instanceID) + if err != nil { + return nil, err + } + if pd == nil { + return nil, nil // absent => terminated, per interface contract + } + inst := toInstance(*pd) + return &inst, nil +} + +// List implements provider.Provider. It reports EVERY Pod in the account: nothing marks a +// Pod as Nebula's, so the account must be dedicated to Nebula (see podName). +func (p *Provider) List(ctx context.Context) ([]provider.Instance, error) { + pods, err := p.client.ListPods(ctx) + if err != nil { + return nil, err + } + out := make([]provider.Instance, 0, len(pods)) + for _, pd := range pods { + out = append(out, toInstance(pd)) + } + return out, nil +} + +// regionsByGeography maps a geography token to the RunPod data centers it encompasses. The +// ids are GET /v2/catalog/datacenters as of 2026-09-29; the grouping is ours, by each id's +// country, since the catalog's continent is coarser (it puts Canada with the US). Empty +// entries are geographies RunPod has no data center in. +var regionsByGeography = map[string][]string{ + "us": { + "US-CA-2", "US-CO-1", "US-GA-2", "US-IL-1", "US-KS-2", "US-MD-1", "US-MO-1", "US-MO-2", + "US-NC-1", "US-NC-2", "US-NE-1", "US-PA-1", "US-TX-3", "US-TX-4", "US-WA-1", "US-WA-2", + }, + "ca": {"CA-MTL-1", "CA-MTL-3", "CA-MTL-4"}, + // Iceland and Norway are EEA members, which is what "eu" means (see provider.Geographies). + "eu": { + "EU-CZ-1", "EU-FR-1", "EU-NL-1", "EU-RO-1", "EU-SE-1", + "EUR-IS-1", "EUR-IS-2", "EUR-IS-3", "EUR-IS-4", "EUR-IS-5", "EUR-NO-1", "EUR-NO-2", + }, + "ap": {"AP-IN-1", "AP-JP-1", "OC-AU-1"}, + "uk": {}, + "sa": {}, + "af": {}, + "me": {}, + "mx": {}, +} + +// ResolveRegions implements provider.Provider. Each declared token is ONE candidate, and a +// geography stays a geography: dataCentersOf expands it at create time, where RunPod places +// by availability within the dataCenterIds it is given. So a geography costs one create, +// while failover still walks the declared tokens one by one. +// +// nil/[] => [""], unpinned: RunPod's widest pool +// ["US"] => ["us"] +// ["EU-RO-1"] => itself, verbatim and unvalidated +// +// narrowTo keeps the candidates inside those geographies; an unconstrained pool becomes one +// candidate per requested geography. +func (p *Provider) ResolveRegions(declared, narrowTo []string) []string { + var within map[string]bool // nil => no narrowing + if len(narrowTo) > 0 { + within = make(map[string]bool) + for _, t := range narrowTo { + if t = strings.ToLower(strings.TrimSpace(t)); provider.IsGeography(t) { + within[t] = true + } + } + } + tokens := declared + if len(declared) == 0 { + if within == nil { + return []string{""} + } + tokens = narrowTo + } + + seen := make(map[string]bool) + var out []string + for _, d := range tokens { + d = strings.TrimSpace(d) + c := "" + if g := strings.ToLower(d); provider.IsGeography(g) { + // A geography with no data center here is no candidate, not a literal id. + if (within == nil || within[g]) && len(regionsByGeography[g]) > 0 { + c = g + } + } else if within == nil || geographyOf(d, within) { + c = d + } + if c != "" && !seen[c] { + seen[c] = true + out = append(out, c) + } + } + return out +} + +// PricePerHour overrides catalog.Base to add the container disk, which RunPod bills on top +// of the GPU. vCPU and RAM add nothing: they come with the GPU. +func (p *Provider) PricePerHour(req provider.PriceRequest) (float64, error) { + gpu, err := p.Base.PricePerHour(req) + if err != nil { + return 0, err + } + return gpu + data.RunPodContainerDiskCostPerHour(max(defaultContainerDiskGiB, req.DiskGiB)), nil +} + +// geographyOf reports whether data center dc is listed under any of geographies. +func geographyOf(dc string, geographies map[string]bool) bool { + for g := range geographies { + if slices.Contains(regionsByGeography[g], dc) { + return true + } + } + return false +} + +// ClassifyProvisionError implements provider.Provider. The categories and the +// scope-derivation rule are shared (provider.ClassifyError and the sentinels the Client +// wraps), so this supplies only the RunPod-specific fact: the region axis. +func (p *Provider) ClassifyProvisionError(err error, accelerator, region string) provider.BlockScope { + // No failure, no block. ClassifyError already returns the zero scope, but the region + // decoration below would repopulate it into a scope recordBlock would install. + if err == nil { + return provider.BlockScope{} + } + scope := provider.ClassifyError(err, nebulav1alpha1.CapacityOnDemand, accelerator) + // The zero scope means BLOCK NOTHING — a rejection of this request that says nothing + // about the candidate, such as an image credential RunPod cannot use. Stamping a region + // onto it would make it non-empty, and recordBlock would install a region-wide block + // across every accelerator: the same trap as the err == nil guard above. + if scope == (provider.BlockScope{}) { + return scope + } + // DenyAll already covers every region (auth fails everywhere), so narrowing it would + // contradict the category. An empty region leaves Region nil, which per BlockScope + // matches only candidates that carry no region either — the unconstrained pool — so + // the block never leaks onto region-pinned candidates. + if region != "" && !scope.DenyAll { + scope.Region = ®ion + } + return scope +} + +// FindByClaim implements provider.Provider. The region is ignored, as in Terminate. +// EXITED and ERROR Pods match: they still bill their disk, so teardown must reach them. +func (p *Provider) FindByClaim(ctx context.Context, claimName, _ string) (*provider.Instance, error) { + if _, err := podName(claimName); err != nil { + return nil, nil // Provision refuses such a name, so no Pod was ever created for it + } + return p.findByClaim(ctx, claimName, func(pd Pod) bool { + return strings.ToUpper(pd.Status) != statusTerminated + }) +} + +// findByClaim returns the first Nebula-owned Pod for claimName that match accepts, or nil +// if none. RunPod has no server-side tag filter, so this is List plus a name comparison. +func (p *Provider) findByClaim( + ctx context.Context, claimName string, match func(Pod) bool, +) (*provider.Instance, error) { + name, err := podName(claimName) + if err != nil { + return nil, err + } + pods, err := p.client.ListPods(ctx) + if err != nil { + return nil, err + } + for _, pd := range pods { + if pd.Name == name && match(pd) { + inst := toInstance(pd) + return &inst, nil + } + } + return nil, nil +} + +// podName is the RunPod Pod name for a NodeClaim: the claim name itself, which is what +// makes FindByClaim work on a backend with no tags. There is no ownership marker, so a Pod +// someone else names like a claim ("default-web-0") would be adopted and later terminated; +// the account must be dedicated to Nebula. +// +// A name that would exceed RunPod's cap is an ERROR, never a truncation. Truncating would +// map two long claim names onto one Pod name, and every consequence of that collision is +// severe: findByClaim adopts the other claim's instance, so one Pod is billed twice and +// the other claim's teardown reaps the survivor. Refusing is loud and fixable (claim names +// derive from the Pod's, so the workload can be renamed); a collision is silent. +func podName(claimName string) (string, error) { + if len(claimName) > maxNameLen { + return "", fmt.Errorf( + "runpod: claim name %q is too long: RunPod caps a pod name at %d characters and identity "+ + "rides that name, so it cannot be shortened", claimName, maxNameLen) + } + return claimName, nil +} + +// podSpecFromPod reads the workload off the Pod (source of truth) and the placement +// decisions off req, then maps the accelerator to RunPod's own ids. +func (p *Provider) podSpecFromPod(pod *corev1.Pod, req provider.ProvisionRequest) (PodSpec, error) { + if len(pod.Spec.Containers) == 0 { + return PodSpec{}, errors.New("runpod: pod has no containers") + } + c := pod.Spec.Containers[0] + if c.Image == "" { + return PodSpec{}, errors.New("runpod: pod's first container has no image") + } + name, err := podName(req.ClaimName) + if err != nil { + return PodSpec{}, err + } + + spec := PodSpec{ + Name: name, + Image: c.Image, + // command → ENTRYPOINT, args → CMD: Kubernetes' two fields land on the Docker + // fields they are defined in terms of, so an image whose own CMD supplies the args + // keeps working when only command is overridden. + Entrypoint: c.Command, + StartCmd: c.Args, + // The caller's resolved map is the whole environment — the Pod's literals plus + // everything envFrom/valueFrom referenced. pod.Spec.Containers[0].Env is NOT read + // here: it holds references this adapter has no cluster access to follow. + Env: req.Env, + ContainerDiskGiB: max(defaultContainerDiskGiB, util.PodEphemeralStorageGiB(pod)), + Ports: containerPorts(&c), + DataCenterIDs: dataCentersOf(req.Region), + } + + // Accelerator type comes from the AcceleratorTypeLabel; count from the container's + // nvidia.com/gpu resource (see util.AcceleratorRequest). + canonical, count, err := util.AcceleratorRequest(pod) + if err != nil { + return PodSpec{}, fmt.Errorf("runpod: %w", err) + } + if canonical == "" || count <= 0 { + // Placement never routes one here (see Capabilities.SupportsCPUOnly). + return PodSpec{}, errors.New("runpod: pod requests no accelerator; RunPod runs GPU Pods only") + } + ids, ok := p.MapAccelerator(canonical, count) + if !ok { + return PodSpec{}, fmt.Errorf("runpod: unsupported accelerator %q: %w", + canonical, provider.ErrUnsupportedAccelerator) + } + spec.GPUTypeID = ids[0] + spec.GPUCount = count + // RunPod sizes a GPU Pod's cpu/memory PER GPU, so the Pod's totals are divided by + // the count (see PodSpec.VCPUPerGPU). + spec.VCPUPerGPU = perGPU(cores(resourceQty(&c, corev1.ResourceCPU)), count) + spec.RAMPerGPUGiB = perGPU(gib(resourceQty(&c, corev1.ResourceMemory)), count) + return spec, nil +} + +// checkRegistryAuth reports whether RunPod can honour a pull credential, so Provision only +// pays for an EnsureRegistryAuth call on a kind that can work. +// +// A refusal, never a fallback: an anonymous pull of a private image either 401s opaquely or +// succeeds against a PUBLIC image of the same name. +func checkRegistryAuth(a *provider.RegistryAuth) error { + switch { + case a.Basic != nil: + if err := a.Validate(); err != nil { + return fmt.Errorf("runpod: %w", err) // every error out of this adapter is prefixed + } + return nil + default: + // AWSRole is the kind the canonical form carries that RunPod has no equivalent for: + // its registry credentials are a static username/password object, with nothing that + // assumes an IAM role on the workload's behalf. An ECR "password" is a 12-hour + // token, so smuggling one in as Basic would provision a Pod that stops being able + // to pull halfway through the day. + return a.Unsupported("runpod") + } +} + +// dataCentersOf expands a candidate ResolveRegions minted into RunPod data centers. Empty is +// unconstrained and yields none. +func dataCentersOf(region string) []string { + if region == "" { + return nil + } + if dcs, ok := regionsByGeography[region]; ok { + return slices.Clone(dcs) + } + return []string{region} +} + +// containerPorts renders the container's declared ports in RunPod's "/" form. +// +// Every port goes out as /http, not /tcp, and that is a real choice: RunPod's http scheme +// publishes a proxy URL derivable from the Pod id (see proxyURL), so the workload is +// reachable the moment it comes up, whereas /tcp reaches it only through a randomly +// assigned public port that must be read back after boot. A Pod serving raw TCP is +// therefore not addressable today; a containerPort carries no hint of its protocol above +// TCP/UDP, so the common case is what gets served. +// +// Only TCP ports are exposed (an unset Protocol is TCP). UDP and SCTP, which RunPod cannot +// expose at all, are skipped rather than refused: a port the workload never serves +// externally should not block creating the Pod. Skipping also keeps one out of proxyURL. +func containerPorts(c *corev1.Container) []string { + var ports []string + for _, p := range c.Ports { + if p.Protocol == "" || p.Protocol == corev1.ProtocolTCP { + ports = append(ports, fmt.Sprintf("%d/http", p.ContainerPort)) + } + } + return ports +} + +// proxyURL is RunPod's HTTP proxy address for a Pod's first declared port. It is derived, +// not read back: the form is fixed, so the URL is known at create time and survives a +// manager restart without an API call. +// +// Empty when the container declares no port — there is then no port to route to, and a +// guess would publish an endpoint that answers nothing. +func proxyURL(id string, ports []string) string { + if id == "" || len(ports) == 0 { + return "" + } + port, _, _ := strings.Cut(ports[0], "/") + return fmt.Sprintf("https://%s-%s.proxy.runpod.net", id, port) +} + +// resourceQty returns the container's request for name, falling back to its limit, or nil +// when neither is present. Requests first because that is the floor the workload declared; +// RunPod has no separate ceiling to set, so limits are only a fallback source of a number. +func resourceQty(c *corev1.Container, name corev1.ResourceName) *resource.Quantity { + if q, ok := c.Resources.Requests[name]; ok { + return &q + } + if q, ok := c.Resources.Limits[name]; ok { + return &q + } + return nil +} + +// cores converts a CPU quantity to whole vCPUs, RunPod's unit, rounding UP: a request of +// 500m is one vCPU, and 1500m is two. Rounding down would hand the workload less CPU than +// it asked for, and a fractional vCPU is not something RunPod can express. Nil (unset) is 0, +// which leaves RunPod's own default. +func cores(q *resource.Quantity) int { + if q == nil { + return 0 + } + return ceilDiv(int(q.MilliValue()), 1000) +} + +// gib converts a memory quantity to whole GiB, RunPod's unit, rounding up as cores does and +// for the same reason. Nil is 0. +func gib(q *resource.Quantity) int { + if q == nil { + return 0 + } + const giB = 1024 * 1024 * 1024 + return ceilDiv(int(q.Value()), giB) +} + +// defaultContainerDiskGiB sizes the container disk of a Pod that requests no ephemeral +// storage, and PricePerHour charges the same floor. It must be sent: despite the schema +// marking disk optional, a create without it fails with "You must either provide a template +// id or pod configuration parameters". +const defaultContainerDiskGiB = 5 + +// perGPU divides a Pod-wide total by the accelerator count, rounding up, because RunPod +// sizes cpu and memory PER GPU. Rounding up keeps the total at or above what the Pod asked +// for; rounding down would under-provision every request that does not divide evenly. +// +// A zero total stays zero (unset → RunPod's default), and a zero count is treated as one so +// a malformed request cannot divide by zero. +func perGPU(total int, count int32) int { + if total <= 0 { + return 0 + } + if count <= 0 { + return total + } + return ceilDiv(total, int(count)) +} + +// ceilDiv divides rounding away from zero for positive inputs. Its own function because +// every conversion above rounds the same way, and an inlined `(a+b-1)/b` is easy to get +// subtly wrong once. +func ceilDiv(a, b int) int { + if a <= 0 || b <= 0 { + return 0 + } + return (a + b - 1) / b +} + +// RunPod's observed Pod statuses, as v2 documents them. PROVISIONING and STARTING are the +// other two, both Pending. +const ( + statusRunning = "RUNNING" // container healthy + // statusExited is RunPod's stopped state, clean exit and crash alike, so it maps to + // Terminated — "gone", with no claim about why. + statusExited = "EXITED" + statusError = "ERROR" // unrecoverable + statusTerminated = "TERMINATED" +) + +// toState maps an observed Pod to the provider-agnostic lifecycle state. Everything +// unrecognized falls to Pending, so a status this adapter has not seen keeps the poll loop +// watching rather than going terminal on a live, billing Pod. +func toState(pd Pod) provider.InstanceState { + switch strings.ToUpper(pd.Status) { + case statusRunning: + return provider.InstanceRunning + case statusExited, statusTerminated: + return provider.InstanceTerminated + case statusError: + return provider.InstanceFailed + default: + return provider.InstancePending + } +} + +// toInstance normalizes an observed RunPod Pod into the provider-agnostic Instance. +// +// Endpoint is always the proxy URL: every port goes out /http, and runtime.ports carries +// only RunPod-internal 100.64/10 addresses plus a port RunPod injects, none of them +// public. An empty value never clears what is already on the Pod (the write paths skip ""). +func toInstance(pd Pod) provider.Instance { + return provider.Instance{ + ID: pd.ID, + ClaimName: pd.Name, + State: toState(pd), + CapacityType: nebulav1alpha1.CapacityOnDemand, + Region: pd.DataCenterID, + Endpoint: proxyURL(pd.ID, pd.Ports), + } +} diff --git a/pkg/provider/runpod/runpod_test.go b/pkg/provider/runpod/runpod_test.go new file mode 100644 index 0000000..f216582 --- /dev/null +++ b/pkg/provider/runpod/runpod_test.go @@ -0,0 +1,772 @@ +/* +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 runpod + +import ( + "context" + "errors" + "fmt" + "math" + "reflect" + "strings" + "testing" + + corev1 "k8s.io/api/core/v1" + "k8s.io/apimachinery/pkg/api/resource" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + + nebulav1alpha1 "github.com/InftyAI/Nebula/api/v1alpha1" + "github.com/InftyAI/Nebula/pkg/provider" + "github.com/InftyAI/Nebula/pkg/provider/catalog/data" + "github.com/InftyAI/Nebula/pkg/util" +) + +// fakeClient is an in-memory Client. It records the last CreatePod spec — the thing the +// adapter actually decides — and lets a test seed existing Pods or inject an error. +type fakeClient struct { + pods []Pod + lastSpec PodSpec + createCnt int + createErr error + createID string + + terminated []string + + // authID is what EnsureRegistryAuth resolves to, authFor the credential it was handed, + // and authCnt how many times it was called — Provision must not pay for it when the + // spec was going to be refused anyway. + authID string + authErr error + authFor *provider.RegistryAuth + authCnt int +} + +func (f *fakeClient) CreatePod(_ context.Context, spec PodSpec) (string, error) { + f.createCnt++ + f.lastSpec = spec + if f.createErr != nil { + return "", f.createErr + } + id := f.createID + if id == "" { + id = "pod-new" + } + f.pods = append(f.pods, Pod{ID: id, Name: spec.Name, Status: statusRunning}) + return id, nil +} + +func (f *fakeClient) TerminatePod(_ context.Context, id string) error { + f.terminated = append(f.terminated, id) + return nil +} + +func (f *fakeClient) GetPod(_ context.Context, id string) (*Pod, error) { + for i := range f.pods { + if f.pods[i].ID == id { + pd := f.pods[i] + return &pd, nil + } + } + return nil, nil +} + +func (f *fakeClient) ListPods(_ context.Context) ([]Pod, error) { return f.pods, nil } + +func (f *fakeClient) EnsureRegistryAuth(_ context.Context, a *provider.RegistryAuth) (string, error) { + f.authCnt++ + f.authFor = a + if f.authErr != nil { + return "", f.authErr + } + return f.authID, nil +} + +// fakeCatalog is a trivial catalog.Lookup. +type fakeCatalog struct{ rows []provider.Offering } + +func (c fakeCatalog) Offerings(_ string) []provider.Offering { return c.rows } + +// newTestProvider builds a Provider over a fake client and a catalog shaped like +// runpod.csv: H100 carries THREE interchangeable RunPod ids (so MapAccelerator's +// primary is observable), A100-80GB one, and L4 one. +func newTestProvider(f *fakeClient) *Provider { + od := nebulav1alpha1.CapacityOnDemand + row := func(typ, id string, tier nebulav1alpha1.CapacityType, price float64) provider.Offering { + return provider.Offering{ + AcceleratorType: typ, AcceleratorID: id, CapacityType: tier, + PricePerHour: price, Available: true, + } + } + return New(f, fakeCatalog{rows: []provider.Offering{ + row("H100", "NVIDIA H100 80GB HBM3", od, 2.99), + row("H100", "NVIDIA H100 NVL", od, 2.79), + row("H100", "NVIDIA H100 PCIe", od, 2.39), + row("A100-80GB", "NVIDIA A100-SXM4-80GB", od, 1.74), + row("L4", "NVIDIA L4", od, 0.39), + }}) +} + +// gpuPod builds a Pod whose accelerator type rides on the label and whose count rides on +// the container's nvidia.com/gpu limit. count<=0 means CPU-only (neither is set). accel is +// passed through verbatim so a test can also exercise non-canonical casing. +func gpuPod(accel string, count int64) *corev1.Pod { + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{Name: "p", Namespace: "default"}, + Spec: corev1.PodSpec{Containers: []corev1.Container{{ + Name: "main", + Image: "myimg:latest", + Command: []string{"/entry.sh"}, + Args: []string{"--flag", "v"}, + }}}, + } + if accel != "" && count > 0 { + pod.Labels = map[string]string{nebulav1alpha1.AcceleratorTypeLabel: accel} + pod.Spec.Containers[0].Resources.Limits = corev1.ResourceList{ + util.NvidiaGPUResource: *resource.NewQuantity(count, resource.DecimalSI), + } + } + return pod +} + +// requests sets the container's resource requests from a k8s-notation map. +func requests(pod *corev1.Pod, m map[corev1.ResourceName]string) *corev1.Pod { + c := &pod.Spec.Containers[0] + if c.Resources.Requests == nil { + c.Resources.Requests = corev1.ResourceList{} + } + for k, v := range m { + c.Resources.Requests[k] = resource.MustParse(v) + } + return pod +} + +// compile-time check that the fake really satisfies the seam the adapter is written to. +var _ Client = (*fakeClient)(nil) + +func TestProvision_GPUPod(t *testing.T) { + f := &fakeClient{createID: "pod-1"} + p := newTestProvider(f) + + pod := requests(gpuPod("H100", 2), map[corev1.ResourceName]string{ + corev1.ResourceCPU: "9", + corev1.ResourceMemory: "100Gi", + corev1.ResourceEphemeralStorage: "80Gi", + }) + pod.Spec.Containers[0].Ports = []corev1.ContainerPort{ + {ContainerPort: 5353, Protocol: corev1.ProtocolUDP}, // skipped, so not the ConnectURL port + {ContainerPort: 8000}, + {ContainerPort: 9090, Protocol: corev1.ProtocolTCP}, + } + + res, err := p.Provision(context.Background(), pod, provider.ProvisionRequest{ + ClaimName: "claim-a", + CapacityType: nebulav1alpha1.CapacityOnDemand, + Region: "US-KS-2", + Env: map[string]string{"HF_TOKEN": "hf_secret"}, + }) + if err != nil { + t.Fatalf("Provision: %v", err) + } + // A RunPod create allocates a machine before it answers — a shortage comes back as an + // error, not a queued Pod — so an id here means real capacity was reserved. This is the + // one place RunPod differs from Modal, and getting it wrong would report capacity that + // was never granted. + if !res.Reserved { + t.Error("Reserved = false; a successful RunPod create means a host was allocated") + } + if res.InstanceID != "pod-1" { + t.Errorf("InstanceID = %q, want pod-1", res.InstanceID) + } + // Derived from the id and the FIRST declared port, so it needs no read-back. + if want := "https://pod-1-8000.proxy.runpod.net"; res.ConnectURL != want { + t.Errorf("ConnectURL = %q, want %q", res.ConnectURL, want) + } + // RunPod's HTTP proxy is unauthenticated: there is no credential to hand back, and a + // placeholder would look like one the Pod could authenticate with. + if res.ConnectToken != "" { + t.Errorf("ConnectToken = %q, want empty (the proxy is unauthenticated)", res.ConnectToken) + } + + s := f.lastSpec + if s.Name != "claim-a" { + t.Errorf("Name = %q, want claim-a", s.Name) + } + if s.Image != "myimg:latest" { + t.Errorf("Image = %q", s.Image) + } + // command → ENTRYPOINT and args → CMD stay SEPARATE, unlike Modal where both + // concatenate into one command. + if strings.Join(s.Entrypoint, " ") != "/entry.sh" || strings.Join(s.StartCmd, " ") != "--flag v" { + t.Errorf("Entrypoint = %v, StartCmd = %v", s.Entrypoint, s.StartCmd) + } + if s.Env["HF_TOKEN"] != "hf_secret" { + t.Errorf("Env = %v; the caller's resolved env must go out whole", provider.RedactedEnv(s.Env)) + } + // v2 takes ONE gpu id per create, so only the primary goes out. + if s.GPUTypeID != "NVIDIA H100 80GB HBM3" { + t.Errorf("GPUTypeID = %q, want the primary H100 id", s.GPUTypeID) + } + if s.GPUCount != 2 { + t.Errorf("GPUCount = %d, want 2", s.GPUCount) + } + // RunPod sizes cpu/memory PER GPU, so the Pod's totals divide by 2 and round UP: + // 9 vCPU → 5, 100 GiB → 50. Rounding down would under-provision the request. + if s.VCPUPerGPU != 5 || s.RAMPerGPUGiB != 50 { + t.Errorf("VCPUPerGPU = %d, RAMPerGPUGiB = %d, want 5/50 (per-GPU, rounded up)", + s.VCPUPerGPU, s.RAMPerGPUGiB) + } + if s.ContainerDiskGiB != 80 { + t.Errorf("ContainerDiskGiB = %d, want 80", s.ContainerDiskGiB) + } + if strings.Join(s.Ports, ",") != "8000/http,9090/http" { + t.Errorf("Ports = %v, want the TCP ports as /http, UDP skipped", s.Ports) + } + if strings.Join(s.DataCenterIDs, ",") != "US-KS-2" { + t.Errorf("DataCenterIDs = %v, want [US-KS-2]", s.DataCenterIDs) + } +} + +func TestProvision_GeographyRegion(t *testing.T) { + f := &fakeClient{createID: "pod-ca"} + p := newTestProvider(f) + + // A geography reaches Provision as the bare token; it must expand into the exact data + // centers, or RunPod is sent an id it has never heard of. + region := p.ResolveRegions([]string{"ca"}, nil)[0] + if _, err := p.Provision(context.Background(), gpuPod("l4", 1), provider.ProvisionRequest{ + ClaimName: "claim-g", + CapacityType: nebulav1alpha1.CapacityOnDemand, + Region: region, + }); err != nil { + t.Fatalf("Provision: %v", err) + } + s := f.lastSpec + if !reflect.DeepEqual(s.DataCenterIDs, regionsByGeography["ca"]) { + t.Errorf("DataCenterIDs = %v, want %v", s.DataCenterIDs, regionsByGeography["ca"]) + } + // The label's casing is the user's; the catalog id that goes out is not. + if s.GPUTypeID != "NVIDIA L4" { + t.Errorf("GPUTypeID = %q, want NVIDIA L4 from a lowercase label", s.GPUTypeID) + } + // No ephemeral-storage request: the disk must still go out, or RunPod rejects the create. + if s.ContainerDiskGiB != defaultContainerDiskGiB { + t.Errorf("ContainerDiskGiB = %d, want the %d GiB default", s.ContainerDiskGiB, defaultContainerDiskGiB) + } +} + +func TestProvision_CPUOnlyPodRefused(t *testing.T) { + f := &fakeClient{createID: "pod-cpu"} + p := newTestProvider(f) + + pod := requests(gpuPod("", 0), map[corev1.ResourceName]string{corev1.ResourceCPU: "2500m"}) + if _, err := p.Provision(context.Background(), pod, provider.ProvisionRequest{ + ClaimName: "claim-c", + CapacityType: nebulav1alpha1.CapacityOnDemand, + }); err == nil { + t.Fatal("Provision of a CPU-only Pod succeeded, want refusal") + } + if f.createCnt != 0 { + t.Errorf("CreatePod called %d times for a refused Pod", f.createCnt) + } +} + +func TestProvision_Idempotent(t *testing.T) { + // A Pod already carrying this claim's name is the ONLY record of ownership RunPod + // offers, so a repeat after a partial create must find it rather than pay twice. + f := &fakeClient{pods: []Pod{{ + ID: "pod-existing", Name: "claim-a", Status: statusRunning, + }}} + p := newTestProvider(f) + + res, err := p.Provision(context.Background(), gpuPod("H100", 1), provider.ProvisionRequest{ + ClaimName: "claim-a", + CapacityType: nebulav1alpha1.CapacityOnDemand, + }) + if err != nil { + t.Fatalf("Provision: %v", err) + } + if res.InstanceID != "pod-existing" { + t.Errorf("InstanceID = %q, want the existing pod-existing", res.InstanceID) + } + if !res.Reserved { + t.Error("Reserved = false; the Pod exists, so a machine was allocated") + } + if f.createCnt != 0 { + t.Errorf("CreatePod called %d times; a second Pod would be billed twice", f.createCnt) + } + // The interface forbids re-minting a credential on a repeat, and there is nothing to + // mint here anyway — the proxy URL is derivable from the id. + if res.ConnectToken != "" { + t.Errorf("ConnectToken = %q, want empty on a re-Provision", res.ConnectToken) + } +} + +func TestProvision_DoesNotAdoptExitedPod(t *testing.T) { + // Claim names are reused across Pod restarts. Adopting a leftover EXITED Pod would hand + // the new Pod an instance toState reads as Terminated, failing it instead of replacing. + f := &fakeClient{pods: []Pod{{ID: "pod-exited", Name: "claim-a", Status: statusExited}}} + p := newTestProvider(f) + + res, err := p.Provision(context.Background(), gpuPod("H100", 1), provider.ProvisionRequest{ + ClaimName: "claim-a", + CapacityType: nebulav1alpha1.CapacityOnDemand, + }) + if err != nil { + t.Fatalf("Provision: %v", err) + } + if res.InstanceID == "pod-exited" || f.createCnt != 1 { + t.Fatalf("adopted the EXITED Pod (id %q, creates %d), want a fresh create", res.InstanceID, f.createCnt) + } +} + +func TestFindByClaim(t *testing.T) { + f := &fakeClient{pods: []Pod{ + {ID: "pod-exited", Name: "claim-a", Status: statusExited}, + {ID: "pod-gone", Name: "claim-b", Status: statusTerminated}, + }} + p := newTestProvider(f) + ctx := context.Background() + + // An EXITED Pod still bills its disk, so teardown must find it. + if got, err := p.FindByClaim(ctx, "claim-a", ""); err != nil || got == nil || got.ID != "pod-exited" { + t.Fatalf("FindByClaim(EXITED) = %+v, %v; want pod-exited", got, err) + } + if got, err := p.FindByClaim(ctx, "claim-b", ""); err != nil || got != nil { + t.Fatalf("FindByClaim(TERMINATED) = %+v, %v; want nil, nil", got, err) + } + // Provision refuses an overlong name, so nothing exists for it. An error here would + // wedge the NodeClaim finalizer on every retry. + long := strings.Repeat("x", maxNameLen+1) + if got, err := p.FindByClaim(ctx, long, ""); err != nil || got != nil { + t.Fatalf("FindByClaim(overlong) = %+v, %v; want nil, nil", got, err) + } +} + +func TestProvision_RefusesOverlongClaimName(t *testing.T) { + // RunPod caps a pod name at 191 chars and the name is Nebula's ONLY carrier of + // identity, so a name that does not fit is refused rather than truncated: two + // truncated claims would collide onto one Pod, which bills one twice and lets the + // other's teardown reap the survivor. + f := &fakeClient{} + p := newTestProvider(f) + + long := strings.Repeat("a", maxNameLen+1) + _, err := p.Provision(context.Background(), gpuPod("H100", 1), provider.ProvisionRequest{ + ClaimName: long, + CapacityType: nebulav1alpha1.CapacityOnDemand, + }) + if err == nil { + t.Fatal("Provision succeeded with an over-long claim name; the name would have collided") + } + if f.createCnt != 0 { + t.Errorf("CreatePod called %d times despite the refusal", f.createCnt) + } + + // One char shorter fits exactly, so the boundary is not off by one. + if _, err := podName(strings.Repeat("a", maxNameLen)); err != nil { + t.Errorf("podName rejected a name that fits exactly: %v", err) + } +} + +func TestProvision_RefusesRestrictedEgress(t *testing.T) { + // RunPod has no outbound-allowlist knob at all. Placement should never route such a + // pool here, but a request can be built by anyone, and provisioning it anyway would + // put the workload on the open internet under a policy that says otherwise. + f := &fakeClient{} + p := newTestProvider(f) + + _, err := p.Provision(context.Background(), gpuPod("H100", 1), provider.ProvisionRequest{ + ClaimName: "claim-e", + CapacityType: nebulav1alpha1.CapacityOnDemand, + Egress: &nebulav1alpha1.EgressPolicy{Mode: nebulav1alpha1.EgressBlocked}, + }) + if err == nil { + t.Fatal("Provision accepted a restricted egress policy RunPod cannot enforce") + } + if f.createCnt != 0 { + t.Errorf("CreatePod called %d times despite the refusal", f.createCnt) + } +} + +func TestProvision_RegistryAuth(t *testing.T) { + basic := &provider.RegistryAuth{ + Registry: "ghcr.io", + Basic: &provider.BasicAuth{Username: "u", Password: "p4ssw0rd"}, + } + + t.Run("basic resolves to an auth id", func(t *testing.T) { + f := &fakeClient{createID: "pod-auth", authID: "cra-123"} + p := newTestProvider(f) + + if _, err := p.Provision(context.Background(), gpuPod("H100", 1), provider.ProvisionRequest{ + ClaimName: "claim-r", + CapacityType: nebulav1alpha1.CapacityOnDemand, + RegistryAuth: basic, + }); err != nil { + t.Fatalf("Provision: %v", err) + } + // RunPod's create takes an OBJECT ID, never an inline username/password, so the + // indirection has to happen before the create. + if f.lastSpec.RegistryAuthID != "cra-123" { + t.Errorf("RegistryAuthID = %q, want cra-123", f.lastSpec.RegistryAuthID) + } + if f.authCnt != 1 { + t.Errorf("EnsureRegistryAuth called %d times, want 1", f.authCnt) + } + }) + + t.Run("aws role is refused without an API call", func(t *testing.T) { + // An ECR "password" is a 12-hour token, so there is no honest way to flatten a role + // into RunPod's static credential object: a Pod pulling from it would stop being able + // to pull halfway through the day. Refuse, and never fall back to an anonymous pull. + f := &fakeClient{} + p := newTestProvider(f) + + _, err := p.Provision(context.Background(), gpuPod("H100", 1), provider.ProvisionRequest{ + ClaimName: "claim-ecr", + CapacityType: nebulav1alpha1.CapacityOnDemand, + RegistryAuth: &provider.RegistryAuth{ + Registry: "1234.dkr.ecr.us-east-1.amazonaws.com", + AWSRole: &provider.AWSRoleAuth{RoleARN: "arn:aws:iam::1234:role/pull", Region: "us-east-1"}, + }, + }) + if err == nil { + t.Fatal("Provision accepted an AWSRole credential RunPod has no equivalent for") + } + // A rejection of the REQUEST, not of the candidate: the Pod fails with the reason + // instead of retrying, and nothing gets blocklisted. + if scope := p.ClassifyProvisionError(err, "H100:1", ""); scope != (provider.BlockScope{}) { + t.Errorf("scope = %+v, want the zero scope", scope) + } + if f.authCnt != 0 || f.createCnt != 0 { + t.Errorf("authCnt = %d, createCnt = %d; a refused credential must cost no API calls", + f.authCnt, f.createCnt) + } + }) +} + +func TestProvision_UnsupportedAccelerator(t *testing.T) { + // A type the catalog has no RunPod id for cannot be launched, and the sentinel is what + // turns this into a capacity-class block rather than nebula_provision_failures_total + // {reason="other"}. + f := &fakeClient{} + p := newTestProvider(f) + + _, err := p.Provision(context.Background(), gpuPod("TPUv5", 1), provider.ProvisionRequest{ + ClaimName: "claim-x", + CapacityType: nebulav1alpha1.CapacityOnDemand, + }) + if !errors.Is(err, provider.ErrUnsupportedAccelerator) { + t.Fatalf("error = %v, want it to wrap ErrUnsupportedAccelerator", err) + } + if f.createCnt != 0 { + t.Errorf("CreatePod called %d times for an accelerator with no RunPod id", f.createCnt) + } +} + +func TestPodSpecString_Redacts(t *testing.T) { + // A spec reaches logs and error strings, and Env holds everything envFrom/valueFrom + // resolved — Secret values included. Key NAMES are already in the Pod spec, so they + // may print; values never may. + s := PodSpec{ + Name: "claim-a", + Image: "myimg:latest", + Env: map[string]string{"HF_TOKEN": "hf_supersecret", "PLAIN": "visible"}, + } + for _, form := range []string{s.String(), fmt.Sprintf("%v", s), fmt.Sprintf("%#v", s)} { + if strings.Contains(form, "hf_supersecret") || strings.Contains(form, "visible") { + t.Errorf("rendered spec leaks an env VALUE: %s", form) + } + if !strings.Contains(form, "HF_TOKEN") { + t.Errorf("rendered spec dropped the env key names, which are safe: %s", form) + } + } +} + +func TestList_NameIsTheClaim(t *testing.T) { + // No ownership marker: every Pod in the account is reported, its name read as the claim. + f := &fakeClient{pods: []Pod{ + {ID: "pod-1", Name: "claim-a", Status: statusRunning}, + {ID: "pod-2", Name: "claim-b", Status: statusExited}, + }} + p := newTestProvider(f) + + got, err := p.List(context.Background()) + if err != nil { + t.Fatalf("List: %v", err) + } + if len(got) != 2 { + t.Fatalf("List returned %d instances, want 2: %+v", len(got), got) + } + if got[0].ClaimName != "claim-a" || got[1].ClaimName != "claim-b" { + t.Errorf("claim names = %q/%q, want claim-a/claim-b", got[0].ClaimName, got[1].ClaimName) + } + if got[0].State != provider.InstanceRunning || got[1].State != provider.InstanceTerminated { + t.Errorf("states = %v/%v", got[0].State, got[1].State) + } +} + +func TestGetAndTerminate(t *testing.T) { + f := &fakeClient{pods: []Pod{{ + ID: "pod-1", Name: "claim-a", Status: statusRunning, DataCenterID: "EU-RO-1", + Ports: []string{"8000/http"}, + }}} + p := newTestProvider(f) + + inst, err := p.Get(context.Background(), "pod-1", "") + if err != nil { + t.Fatalf("Get: %v", err) + } + if inst == nil { + t.Fatal("Get returned nil for a live Pod") + } + // Region is where RunPod actually put it, which a multi-DC candidate does not predict. + if inst.CapacityType != nebulav1alpha1.CapacityOnDemand || inst.Region != "EU-RO-1" { + t.Errorf("CapacityType = %q, Region = %q", inst.CapacityType, inst.Region) + } + // An /http-only Pod has no public IP or port mapping ever, so the derived proxy URL is + // its only address. + if want := "https://pod-1-8000.proxy.runpod.net"; inst.Endpoint != want { + t.Errorf("Endpoint = %q, want %q", inst.Endpoint, want) + } + + // A Pod that is gone reports (nil, nil) — absent means terminated, per the interface. + gone, err := p.Get(context.Background(), "pod-missing", "") + if err != nil || gone != nil { + t.Errorf("Get(missing) = %v, %v; want nil, nil", gone, err) + } + + if err := p.Terminate(context.Background(), "pod-1", "EU-RO-1"); err != nil { + t.Fatalf("Terminate: %v", err) + } + if len(f.terminated) != 1 || f.terminated[0] != "pod-1" { + t.Errorf("terminated = %v, want [pod-1]", f.terminated) + } + // Nothing was ever provisioned: there is no id to delete and no call to make. + if err := p.Terminate(context.Background(), "", ""); err != nil { + t.Errorf("Terminate(\"\") = %v, want nil", err) + } + if len(f.terminated) != 1 { + t.Errorf("Terminate(\"\") called the API: %v", f.terminated) + } +} + +func TestToState(t *testing.T) { + cases := []struct { + status string + want provider.InstanceState + }{ + // v2's status is OBSERVED, unlike v1's desiredStatus: RUNNING means the container + // started, so no started-at check is needed to keep an image pull from reading ready. + {"RUNNING", provider.InstanceRunning}, + {"running", provider.InstanceRunning}, // casing is not load-bearing + {"PROVISIONING", provider.InstancePending}, + {"STARTING", provider.InstancePending}, + {"EXITED", provider.InstanceTerminated}, + {"TERMINATED", provider.InstanceTerminated}, + {"ERROR", provider.InstanceFailed}, + // A status this adapter has never seen keeps the poll loop WATCHING rather than going + // terminal on a live, billing Pod. + {"SOMETHING_NEW", provider.InstancePending}, + {"", provider.InstancePending}, + } + for _, tc := range cases { + if got := toState(Pod{Status: tc.status}); got != tc.want { + t.Errorf("toState(%q) = %v, want %v", tc.status, got, tc.want) + } + } +} + +func TestToInstance_EndpointIsProxyURL(t *testing.T) { + pd := Pod{ID: "pod-1", Ports: []string{"8080/http"}} + if got := toInstance(pd).Endpoint; got != "https://pod-1-8080.proxy.runpod.net" { + t.Errorf("Endpoint = %q, want the proxy URL", got) + } + // No port at all: no address to publish, and a guess would advertise something that + // answers nothing. + if got := toInstance(Pod{ID: "pod-2"}).Endpoint; got != "" { + t.Errorf("Endpoint(no ports) = %q, want empty", got) + } +} + +func TestClassifyProvisionError(t *testing.T) { + p := newTestProvider(&fakeClient{}) + const accel, region = "H100:8", "EU-RO-1" + + t.Run("nil error blocks nothing", func(t *testing.T) { + // ClassifyError already returns the zero scope here, but the region decoration + // below would repopulate it into a scope recordBlock installs. + if got := p.ClassifyProvisionError(nil, accel, region); got != (provider.BlockScope{}) { + t.Errorf("scope = %+v, want the zero scope", got) + } + }) + + t.Run("capacity is scoped to this accelerator, tier and region", func(t *testing.T) { + err := fmt.Errorf("no instances available: %w", provider.ErrNoCapacity) + got := p.ClassifyProvisionError(err, accel, region) + if got.DenyAll { + t.Error("DenyAll = true for a capacity shortage; only this candidate ran out") + } + if got.Accelerator == nil || *got.Accelerator != accel { + t.Errorf("Accelerator = %v, want %q", got.Accelerator, accel) + } + if got.CapacityType != nebulav1alpha1.CapacityOnDemand { + t.Errorf("CapacityType = %q, want OnDemand", got.CapacityType) + } + if got.Region == nil || *got.Region != region { + t.Errorf("Region = %v, want %q — a shortage in one DC must not disqualify another", + got.Region, region) + } + }) + + t.Run("auth denies the whole provider and is not narrowed", func(t *testing.T) { + // A bad API key fails in every region, so narrowing DenyAll to one would contradict + // the category and keep trying the other candidates against the same dead key. + err := fmt.Errorf("401: %w", provider.ErrAuth) + got := p.ClassifyProvisionError(err, accel, region) + if !got.DenyAll { + t.Fatal("DenyAll = false for an auth failure") + } + if got.Region != nil { + t.Errorf("Region = %v on a DenyAll scope, want nil", got.Region) + } + }) + + t.Run("a request rejection is not decorated into a block", func(t *testing.T) { + // An unusable pull credential says nothing about the CANDIDATE. Stamping a region + // onto the zero scope would make it non-empty, and recordBlock would then install a + // region-wide block across every accelerator — excluding that DC for every other Pod + // until the TTL lapsed. + err := fmt.Errorf("runpod: image pull credential or image rejected: %w", errors.New("401")) + if got := p.ClassifyProvisionError(err, accel, region); got != (provider.BlockScope{}) { + t.Errorf("scope = %+v, want the zero scope so nothing is blocklisted", got) + } + }) + + t.Run("an empty region leaves Region nil", func(t *testing.T) { + // nil matches only candidates that carry no region either — the unconstrained pool — + // so the block never leaks onto region-pinned candidates. + err := fmt.Errorf("no capacity: %w", provider.ErrNoCapacity) + if got := p.ClassifyProvisionError(err, accel, ""); got.Region != nil { + t.Errorf("Region = %v, want nil", got.Region) + } + }) +} + +func TestCapabilities(t *testing.T) { + c := newTestProvider(&fakeClient{}).Capabilities() + // REST v2 has no interruptible tier, so advertising Spot would place Pods it cannot launch. + if c.SupportsSpot { + t.Error("SupportsSpot = true") + } + // A stopped RunPod Pod still bills for its disk and releases its GPU, so it is neither + // free nor resumable in the sense the capability promises. + if c.SupportsStop { + t.Error("SupportsStop = true; a stopped Pod still bills and has lost its GPU") + } + // No outbound-allowlist knob exists in the API, so placement must skip RunPod for an + // egress-restricted pool rather than have Provision refuse it after the fact. + if c.SupportsEgressPolicy { + t.Error("SupportsEgressPolicy = true; RunPod exposes no outbound policy") + } + // No tags: identity rides the Pod name, which is what List filters on. + if c.NativeTags { + t.Error("NativeTags = true; RunPod Pods have no tags") + } + if c.PollInterval != 0 { + t.Errorf("PollInterval = %v, want 0 (the default cadence; nothing is reclaimed)", c.PollInterval) + } +} + +func TestRegionsByGeography_IsResolvable(t *testing.T) { + for geography := range regionsByGeography { + if !provider.IsGeography(geography) { + t.Errorf("%q is not a provider.Geographies token, so nothing can resolve to it", geography) + } + } + // Every geography needs an entry, even an empty one, or a NodePool declaring it would be + // passed through as a literal data center RunPod has never heard of. + for _, g := range provider.Geographies { + if _, ok := regionsByGeography[g]; !ok { + t.Errorf("geography %q has no RunPod entry", g) + } + } +} + +func TestResolveRegions(t *testing.T) { + p := newTestProvider(&fakeClient{}) + cases := []struct { + name string + declared, narrowTo []string + want []string + }{ + {name: "unconstrained is unpinned", want: []string{""}}, + {name: "a geography is one candidate", declared: []string{"us"}, want: []string{"us"}}, + {name: "a data center passes through", declared: []string{"EU-RO-1"}, want: []string{"EU-RO-1"}}, + // No RunPod data center is in the UK, so the token resolves to no candidate at all + // rather than to the literal "uk". + {name: "an empty geography yields nothing", declared: []string{"uk"}, want: nil}, + {name: "duplicates collapse", declared: []string{"us", "US", "EU-RO-1", "EU-RO-1"}, + want: []string{"us", "EU-RO-1"}}, + {name: "narrowing an unpinned pool", narrowTo: []string{"eu"}, want: []string{"eu"}}, + {name: "narrowing drops other geographies", declared: []string{"us", "eu"}, + narrowTo: []string{"eu"}, want: []string{"eu"}}, + {name: "narrowing keeps a data center inside it", declared: []string{"EU-RO-1", "US-KS-2"}, + narrowTo: []string{"eu"}, want: []string{"EU-RO-1"}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := p.ResolveRegions(tc.declared, tc.narrowTo); !reflect.DeepEqual(got, tc.want) { + t.Errorf("ResolveRegions(%v, %v) = %q, want %q", tc.declared, tc.narrowTo, got, tc.want) + } + }) + } +} + +func TestPricePerHour(t *testing.T) { + p := newTestProvider(&fakeClient{}) + defaultDisk := data.RunPodContainerDiskCostPerHour(defaultContainerDiskGiB) + cases := []struct { + name string + req provider.PriceRequest + want float64 + }{ + // vCPU and RAM ride the GPU price, so only the disk is added. + {name: "gpu pod adds the default disk", + req: provider.PriceRequest{AcceleratorType: "L4", Count: 2, CPUCores: 8, MemoryMiB: 65536}, + want: 2*0.39 + defaultDisk}, + // The disk is priced at what is created, the request or the default floor. + {name: "requested disk is priced", + req: provider.PriceRequest{AcceleratorType: "L4", Count: 1, DiskGiB: 50}, + want: 0.39 + data.RunPodContainerDiskCostPerHour(50)}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := p.PricePerHour(tc.req) + if err != nil { + t.Fatalf("PricePerHour: %v", err) + } + if math.Abs(got-tc.want) > 1e-9 { + t.Errorf("PricePerHour = %v, want %v", got, tc.want) + } + }) + } + _, err := p.PricePerHour(provider.PriceRequest{AcceleratorType: "T4", Count: 1}) + if !errors.Is(err, provider.ErrNoPrice) { + t.Errorf("unlisted accelerator: err = %v, want ErrNoPrice", err) + } +} diff --git a/pkg/util/strings.go b/pkg/util/strings.go new file mode 100644 index 0000000..18aedbf --- /dev/null +++ b/pkg/util/strings.go @@ -0,0 +1,29 @@ +/* +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 util + +import "strings" + +// ContainsAny reports whether s contains any of subs. +func ContainsAny(s string, subs ...string) bool { + for _, sub := range subs { + if strings.Contains(s, sub) { + return true + } + } + return false +}